diff --git a/docs/aio.md b/docs/aio.md index 492d50186..76d917d02 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -280,6 +280,8 @@ parallel operations. Two implementations are provided: `AioS3FileSystem` automatically uses `S3AioExecutor` for file handles, so multipart uploads and parallel range reads are dispatched through the event loop with `asyncio.to_thread()` instead of a separate `ThreadPoolExecutor` per file. +At most `max_workers` of them run at once. +An instance created with `asynchronous=True` has no event loop of its own, so its file handles use `S3ThreadPoolExecutor`. ### Usage with AioS3FSCursor diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 8d1177a6c..7006c6e47 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -11,6 +11,8 @@ import asyncio import logging +import mimetypes +import os from multiprocessing import cpu_count from typing import TYPE_CHECKING, Any, cast @@ -18,7 +20,7 @@ from fsspec.callbacks import _DEFAULT_CALLBACK from pyathena.filesystem.s3 import S3File, S3FileSystem -from pyathena.filesystem.s3_executor import S3AioExecutor +from pyathena.filesystem.s3_executor import S3AioExecutor, S3Executor, S3ThreadPoolExecutor from pyathena.filesystem.s3_object import ( S3Metadata, S3MultipartUpload, @@ -48,6 +50,8 @@ class AioS3FileSystem(AsyncFileSystem): File handles created by ``_open`` use ``S3AioExecutor`` so that parallel operations (range reads, multipart uploads) are dispatched through the event loop with ``asyncio.to_thread`` instead of a ``ThreadPoolExecutor`` per file. + An instance created with ``asynchronous=True`` has no event loop of its own, + so its file handles use a ``ThreadPoolExecutor``. Attributes: _sync_fs: The internal synchronous S3FileSystem instance. @@ -161,11 +165,74 @@ async def _rm_file(self, path: str, **kwargs) -> None: async def _pipe_file( self, path: str, value: bytes | bytearray | memoryview, mode: str = "overwrite", **kwargs ) -> None: + if self._intrans: + # The transaction belongs to this filesystem, not to the internal + # S3FileSystem, so write through open() to defer the commit to it. + await asyncio.to_thread(self._pipe_file_in_transaction, path, value, mode, **kwargs) + return await asyncio.to_thread(self._sync_fs.pipe_file, path, value, mode=mode, **kwargs) + def _pipe_file_in_transaction( + self, path: str, value: bytes | bytearray | memoryview, mode: str, **kwargs + ) -> None: + """Write bytes into the path as a file of this filesystem's transaction. + + Args: + path: S3 path (s3://bucket/key) to write to. + value: The bytes to write. + mode: "overwrite" or "create". With "create", raise + FileExistsError when the object already exists. + **kwargs: Additional parameters passed to ``open()``. + + Raises: + FileExistsError: If the mode is "create" and the path already + exists. + """ + if mode == "create" and self._sync_fs.exists(path): + raise FileExistsError(path) + with self.open(path, "wb", **kwargs) as f: + f.write(value) + async def _put_file(self, lpath: str, rpath: str, callback=_DEFAULT_CALLBACK, **kwargs) -> None: + if self._intrans: + # See _pipe_file. + await asyncio.to_thread(self._put_file_in_transaction, lpath, rpath, callback, **kwargs) + return await asyncio.to_thread(self._sync_fs.put_file, lpath, rpath, callback=callback, **kwargs) + def _put_file_in_transaction(self, lpath: str, rpath: str, callback, **kwargs) -> None: + """Upload a local file as a file of this filesystem's transaction. + + Mirrors :meth:`S3FileSystem.put_file`, but writes through ``open()`` + of this filesystem. + + Args: + lpath: Local file path to upload. + rpath: S3 destination path (s3://bucket/key). + callback: Progress callback for tracking upload progress. + **kwargs: Additional S3 parameters (e.g., ContentType, StorageClass). + """ + if os.path.isdir(lpath): + return + _, key, _ = self.parse_path(rpath) + if not key: + return + + callback.set_size(os.path.getsize(lpath)) + if "ContentType" not in kwargs: + content_type, _ = mimetypes.guess_type(lpath) + if content_type is not None: + kwargs["ContentType"] = content_type + + with ( + self.open(rpath, "wb", s3_additional_kwargs=kwargs) as remote, + open(lpath, "rb") as local, + ): + while data := local.read(remote.blocksize): + remote.write(data) + callback.relative_update(len(data)) + self.invalidate_cache(rpath) + async def _get_file(self, rpath: str, lpath: str, callback=_DEFAULT_CALLBACK, **kwargs) -> None: await asyncio.to_thread(self._sync_fs.get_file, rpath, lpath, callback=callback, **kwargs) @@ -344,6 +411,23 @@ async def _find( return {f.name: f for f in files} return [f.name for f in files] + def _create_executor(self, max_workers: int) -> S3Executor: + """Create the executor for the parallel operations of a file. + + An instance created with ``asynchronous=True`` has no event loop of + its own, so its files run the operations in a thread pool. + + Args: + max_workers: The maximum number of operations that run at once. + + Returns: + An ``S3AioExecutor`` on the event loop of this filesystem, or an + ``S3ThreadPoolExecutor`` if it has none. + """ + if self._loop is None: + return S3ThreadPoolExecutor(max_workers=max_workers) + return S3AioExecutor(loop=self._loop, max_workers=max_workers) + def _open( self, path: str, @@ -367,7 +451,7 @@ def _open( path, mode, max_workers=max_workers, - executor=S3AioExecutor(loop=self._loop), + executor=self._create_executor(max_workers=max_workers), block_size=block_size, cache_type=cache_type, autocommit=autocommit, @@ -576,8 +660,24 @@ def invalidate_cache(self, path: str | None = None) -> None: """ self._sync_fs.invalidate_cache(path) - async def _touch(self, path: str, truncate: bool = True, **kwargs) -> None: - await asyncio.to_thread(self._sync_fs.touch, path, truncate=truncate, **kwargs) + async def _touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]: + return await asyncio.to_thread(self._sync_fs.touch, path, truncate=truncate, **kwargs) + + def touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]: + """Create an empty object with PutObject. + + See :meth:`S3FileSystem.touch`. + + Args: + path: S3 path (s3://bucket/key) of the object. + truncate: If True, replace an existing object with an empty one; + if False, raise if the object exists. + **kwargs: Additional parameters passed to the PutObject API. + + Returns: + The PutObject response as a dictionary. + """ + return self._sync_fs.touch(path, truncate=truncate, **kwargs) class AioS3File(S3File): @@ -589,4 +689,6 @@ class AioS3File(S3File): through the ``S3Executor`` interface — the ``S3AioExecutor`` provided by ``AioS3FileSystem`` dispatches them through the event loop with ``asyncio.to_thread`` instead of a ``ThreadPoolExecutor`` per file. + For an ``AioS3FileSystem`` created with ``asynchronous=True``, it is an + ``S3ThreadPoolExecutor``. """ diff --git a/pyathena/filesystem/s3_executor.py b/pyathena/filesystem/s3_executor.py index 475e21387..c0f6adf2d 100644 --- a/pyathena/filesystem/s3_executor.py +++ b/pyathena/filesystem/s3_executor.py @@ -15,6 +15,7 @@ from collections.abc import Callable from concurrent.futures import Future from concurrent.futures.thread import ThreadPoolExecutor +from multiprocessing import cpu_count from typing import Any, TypeVar from pyathena.util import override @@ -80,26 +81,54 @@ class S3AioExecutor(S3Executor): ``concurrent.futures.Future`` objects that are compatible with ``as_completed()``, ``wait()`` and ``Future.cancel()``. As with ``ThreadPoolExecutor``, a future cannot be cancelled once its function has - started. + started. At most ``max_workers`` of the submitted functions run at once. This avoids thread-in-thread nesting when ``S3File`` is used from within ``asyncio.to_thread()`` calls (the pattern used by ``AioS3FileSystem``). Args: loop: A running asyncio event loop. + max_workers: The maximum number of submitted functions that run at once. Raises: RuntimeError: If the event loop is not running when ``submit`` is called. """ - def __init__(self, loop: asyncio.AbstractEventLoop | None = None) -> None: + def __init__( + self, + loop: asyncio.AbstractEventLoop | None = None, + max_workers: int = (cpu_count() or 1) * 5, + ) -> None: """Initialize the executor with the event loop to schedule work on. Args: loop: The asyncio event loop. ``submit`` raises ``RuntimeError`` if it is None or not running. + max_workers: The maximum number of submitted functions that run + at once. + + Raises: + ValueError: If ``max_workers`` is not positive. """ + if max_workers <= 0: + # As ThreadPoolExecutor does; a semaphore of 0 would never run anything. + raise ValueError("max_workers must be greater than 0") self._loop = loop + self._semaphore = asyncio.Semaphore(max_workers) + + async def _run(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> T: + """Run the function in a thread once fewer than ``max_workers`` run. + + Args: + fn: The blocking function to run. + *args: Positional arguments passed to the function. + **kwargs: Keyword arguments passed to the function. + + Returns: + The return value of the function. + """ + async with self._semaphore: + return await asyncio.to_thread(fn, *args, **kwargs) @override def submit(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> Future[T]: @@ -142,7 +171,7 @@ def settle(task: Future[None]) -> None: elif future.set_running_or_notify_cancel(): future.set_exception(task.exception()) - task = asyncio.run_coroutine_threadsafe(asyncio.to_thread(run), self._loop) + task = asyncio.run_coroutine_threadsafe(self._run(run), self._loop) task.add_done_callback(settle) return future raise RuntimeError( diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index e85438d4e..9ba1e3cb7 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -1,5 +1,7 @@ +import asyncio import os import tempfile +import threading import time import urllib.parse import urllib.request @@ -16,7 +18,12 @@ from pyathena.filesystem.s3 import S3File, S3FileSystem from pyathena.filesystem.s3_async import AioS3File, AioS3FileSystem -from pyathena.filesystem.s3_object import S3Object, S3ObjectType, S3StorageClass +from pyathena.filesystem.s3_object import ( + S3MultipartUploadPart, + S3Object, + S3ObjectType, + S3StorageClass, +) from tests import ENV from tests.pyathena.conftest import connect @@ -193,6 +200,60 @@ async def test_copy_object_with_multipart_upload_invalid_block_size(self, block_ ) fs._sync_fs._call.assert_not_called() + @pytest.mark.parametrize("commit", [True, False]) + def test_transaction_pipe_put_file(self, tmp_path, commit): + # GH-977: pipe_file() and put_file() join the transaction of this + # filesystem; they used to write through the internal S3FileSystem, + # which is not in the transaction, and were not rolled back. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + put_object = fs._sync_fs._put_object = mock.MagicMock() + local = tmp_path / "local.txt" + local.write_bytes(b"local") + + def write(): + with fs.transaction: + fs.pipe_file("s3://bucket/k1", b"data") + fs.put_file(str(local), "s3://bucket/k2") + put_object.assert_not_called() + if not commit: + raise RuntimeError("rollback") + + if commit: + write() + assert [ + (c.kwargs["key"], c.kwargs["body"], c.kwargs.get("ContentType")) + for c in put_object.call_args_list + ] == [("k1", b"data", None), ("k2", b"local", "text/plain")] + else: + with pytest.raises(RuntimeError, match="rollback"): + write() + put_object.assert_not_called() + + def test_transaction_pipe_file_create_existing(self): + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._sync_fs.exists = mock.MagicMock(return_value=True) + with fs.transaction, pytest.raises(FileExistsError): + fs.pipe_file("s3://bucket/key", b"data", mode="create") + fs._sync_fs.exists.assert_called_once_with("s3://bucket/key") + + def test_touch_sync_wrapper(self): + # GH-977: touch() used to be fsspec's open()-based default, which + # dropped the PutObject parameters and returned None. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._sync_fs._call = mock.MagicMock(return_value={"ETag": '"e"'}) + + actual = fs.touch("s3://bucket/key", ContentType="text/plain") + assert isinstance(actual, dict) + assert fs._sync_fs._call.call_args.kwargs == { + "Bucket": "bucket", + "Key": "key", + "ContentType": "text/plain", + } + + fs._sync_fs.exists = mock.MagicMock(return_value=True) + with pytest.raises(ValueError, match="Cannot touch the existing file"): + fs.touch("s3://bucket/key", truncate=False) + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): @@ -968,6 +1029,82 @@ def test_open_max_workers(self): assert isinstance(f, AioS3File) assert f.max_workers == 2 + @pytest.mark.parametrize("asynchronous", [False, True]) + @pytest.mark.asyncio + async def test_open_parallel_requests(self, asynchronous): + # GH-954: max_workers bounds the parallel part uploads and range + # reads. GH-977: a filesystem created with asynchronous=True has no + # event loop of its own and used to fail to run them. + fs = AioS3FileSystem( + connection=mock.MagicMock(), asynchronous=asynchronous, skip_instance_cache=True + ) + block_size = S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE + size = block_size * 4 + condition = threading.Condition() + state = {"active": 0, "peak": 0} + + def track(result): + def call(**kwargs): + with condition: + state["active"] += 1 + state["peak"] = max(state["peak"], state["active"]) + condition.notify_all() + # Hold the call until another one overlaps it, then + # briefly for a third one, which only an executor + # without the limit runs. + condition.wait_for(lambda: state["active"] >= 2, timeout=5) + condition.wait_for(lambda: state["active"] >= 3, timeout=0.2) + with condition: + state["active"] -= 1 + return result(**kwargs) + + return mock.MagicMock(side_effect=call) + + sync_fs = fs._sync_fs + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + sync_fs._upload_part = track( + lambda **kw: S3MultipartUploadPart(kw["part_number"], {"ETag": '"e"'}) + ) + sync_fs._complete_multipart_upload = mock.MagicMock() + sync_fs._get_object = track( + lambda **kw: (kw["ranges"][0], b"a" * (kw["ranges"][1] - kw["ranges"][0])) + ) + sync_fs.info = mock.MagicMock( + return_value=S3Object( + init={"Key": "key"}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="key", + ) + ) + sync_fs.info.return_value.size = size + + def write(): + with fs.open("s3://bucket/key", "wb", block_size=block_size, max_workers=2) as f: + f.write(b"a" * size) + + await asyncio.to_thread(write) + assert sync_fs._upload_part.call_count == 4 + assert state["peak"] == 2 + + def read(): + with fs.open( + "s3://bucket/key", "rb", block_size=block_size, cache_type="none", max_workers=2 + ) as f: + return f.read() + + state["peak"] = 0 + assert await asyncio.to_thread(read) == b"a" * size + assert sync_fs._get_object.call_count == 4 + assert state["peak"] == 2 + + def test_open_invalid_max_workers(self): + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + with pytest.raises(ValueError, match="max_workers must be greater than 0"): + fs.open("s3://bucket/key", "wb", max_workers=0) + def test_open_version_id(self): fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) fs._sync_fs.info = mock.MagicMock( diff --git a/tests/pyathena/filesystem/test_s3_executor.py b/tests/pyathena/filesystem/test_s3_executor.py index 4168c04e9..9b1d005f6 100644 --- a/tests/pyathena/filesystem/test_s3_executor.py +++ b/tests/pyathena/filesystem/test_s3_executor.py @@ -16,6 +16,12 @@ class TestS3AioExecutor: + def test_init(self): + # max_workers is optional, as before it was added. + S3AioExecutor(loop=None) + with pytest.raises(ValueError, match="max_workers must be greater than 0"): + S3AioExecutor(loop=None, max_workers=0) + def test_submit(self): async def main(): executor = S3AioExecutor(loop=asyncio.get_running_loop()) @@ -73,7 +79,14 @@ async def main(): assert not running.cancelled() assert pending.cancelled() - def test_loop_shutdown(self): + @pytest.mark.parametrize( + ("threads", "max_workers"), + [ + (1, 5), # The pending function waits for a thread. + (2, 1), # The pending function waits for a permit. + ], + ) + def test_loop_shutdown(self, threads, max_workers): # A function that has not started when the event loop shuts down is # never run, and its future is cancelled and settled, so that wait() # returns, instead of left pending. @@ -90,8 +103,8 @@ def work(): async def main(): loop = asyncio.get_running_loop() - loop.set_default_executor(ThreadPoolExecutor(max_workers=1)) - executor = S3AioExecutor(loop=loop) + loop.set_default_executor(ThreadPoolExecutor(max_workers=threads)) + executor = S3AioExecutor(loop=loop, max_workers=max_workers) running = executor.submit(work) pending = executor.submit(events.append, "pending finished") pending.add_done_callback(lambda _: settled.set())