diff --git a/docs/filesystem.md b/docs/filesystem.md index b19f93d43..4f8135cf0 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -81,6 +81,10 @@ through the buffered file path. Inside an [fsspec transaction](https://filesystem-spec.readthedocs.io/en/latest/features.html#transactions), writes are deferred until the transaction commits and are discarded on rollback. +The block size for writing, given by the `block_size` argument of `open` or by the +filesystem's `default_block_size`, must be between 5 MiB and 5 GiB, inclusive, the part +size limits of a multipart upload. Otherwise, `open` raises `ValueError`. + ## Error translation S3 error responses are translated into standard Python exceptions, so filesystem diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 3d039832c..2bbda039c 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -7,7 +7,7 @@ import os.path import re from collections.abc import Callable, Iterator -from concurrent.futures import Future, as_completed +from concurrent.futures import Future, as_completed, wait from copy import deepcopy from datetime import datetime from multiprocessing import cpu_count @@ -1275,7 +1275,11 @@ def _copy_object_with_multipart_upload( block_size < self.MULTIPART_UPLOAD_MIN_PART_SIZE or block_size > self.MULTIPART_UPLOAD_MAX_PART_SIZE ): - raise ValueError("Block size must be greater than 5MiB and less than 5GiB.") + raise ValueError( + "Block size must be between " + f"5 MiB ({self.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and " + f"5 GiB ({self.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), inclusive: {block_size}." + ) copy_source = { "Bucket": bucket1, @@ -1407,9 +1411,10 @@ def _finish_multipart_upload( ) -> S3CompleteMultipartUpload: """Collect the uploaded parts and complete the multipart upload. - When any part fails, the remaining parts are cancelled and the - multipart upload is aborted so that no incomplete upload is left - behind, then the original error is re-raised. + When any part or the completion fails, the parts that have not + started are cancelled, the running ones are waited for, and the + multipart upload is aborted so that no incomplete upload or part is + left behind. The original error is then re-raised. Args: bucket: S3 bucket name. @@ -1431,8 +1436,10 @@ def _finish_multipart_upload( parts=parts, ) except Exception: - for future in futures: - future.cancel() + # A part that is still uploading when the upload is aborted may + # be stored after the abort, so wait for the parts that could not + # be cancelled first. + wait([future for future in futures if not future.cancel()]) try: self._call( self._client.abort_multipart_upload, @@ -2249,8 +2256,9 @@ def __init__( part copies. executor: The executor for parallel operations. If None, a new ``S3ThreadPoolExecutor`` is created. - block_size: The block size for reads and writes. Must be at least - ``MULTIPART_UPLOAD_MIN_PART_SIZE`` unless reading. + block_size: The block size for reads and writes. Must be between + ``MULTIPART_UPLOAD_MIN_PART_SIZE`` and + ``MULTIPART_UPLOAD_MAX_PART_SIZE``, inclusive, unless reading. cache_type: The fsspec cache type for reads. autocommit: Whether to commit the written data when the file is closed. If False, :meth:`commit` must be called. @@ -2263,13 +2271,16 @@ def __init__( Raises: ValueError: If the path has no key, the version IDs do not match, - a version is given for writing, or the block size is too small - for writing. + a version is given for writing, or the block size is not + between ``MULTIPART_UPLOAD_MIN_PART_SIZE`` and + ``MULTIPART_UPLOAD_MAX_PART_SIZE`` for writing. """ self.max_workers = max_workers - self._executor: S3Executor = executor or S3ThreadPoolExecutor(max_workers=max_workers) self.s3_additional_kwargs = s3_additional_kwargs if s3_additional_kwargs else {} + # The arguments are validated, and the objects looked up, before the + # base class initializer: a file that fails here is never opened, so + # its garbage collection does not close (flush and commit) it. bucket, key, path_version_id = S3FileSystem.parse_path(path) self.bucket = bucket if not key: @@ -2292,8 +2303,20 @@ def __init__( # Carry the version in the path, as with the ?versionId= suffix, # so that a reopened (e.g., unpickled) file reads the same version. path = f"{path}?versionId={self.version_id}" + if "r" not in mode and not ( + fs.MULTIPART_UPLOAD_MIN_PART_SIZE <= block_size <= fs.MULTIPART_UPLOAD_MAX_PART_SIZE + ): + # When writing, every full block is uploaded as a part of a + # multipart upload. + raise ValueError( + "Block size for writing must be between " + f"5 MiB ({fs.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and " + f"5 GiB ({fs.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), inclusive: {block_size}." + ) self._details: S3Object | dict[str, Any] = {} + append_info: S3Object | None = None + append_data: bytes | None = None if "r" in mode: # Looked up before the base class initializer, which would # otherwise take the size from the latest version of the object. @@ -2308,7 +2331,14 @@ def __init__( self._details = info if size is None: size = info.get("size") + elif "a" in mode and fs.exists(path): + append_info = fs.info(path) + if append_info.get("size", 0) < fs.MULTIPART_UPLOAD_MIN_PART_SIZE: + # Too small to be a part of a multipart upload: rewritten + # from the buffer. + append_data = fs.cat(path) + self._executor: S3Executor = executor or S3ThreadPoolExecutor(max_workers=max_workers) super().__init__( fs=fs, path=path, @@ -2319,28 +2349,19 @@ def __init__( cache_options=cache_options, size=size, ) - if "r" not in mode and block_size < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: - # When writing occurs, the block size should not be smaller - # than the minimum size of a part in a multipart upload. - raise ValueError(f"Block size must be >= {self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE}MB.") self.append_block = False - if "a" in mode and self.fs.exists(path): - info = self.fs.info(self.path, version_id=self.version_id) - loc = info.get("size", 0) - if loc < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: - # Too small to be a part of a multipart upload: rewrite it - # from the buffer. - self.write(self.fs.cat(self.path)) + self.multipart_upload: S3MultipartUpload | None = None + self.multipart_upload_parts: list[Future[S3MultipartUploadPart]] = [] + if append_info is not None: + if append_data is not None: + self.write(append_data) else: # Copied with UploadPartCopy as the leading part(s). self.append_block = True - self.loc = loc - self.s3_additional_kwargs.update(info.to_api_repr()) - self._details = info - - self.multipart_upload: S3MultipartUpload | None = None - self.multipart_upload_parts: list[Future[S3MultipartUploadPart]] = [] + self.loc = append_info.get("size", 0) + self.s3_additional_kwargs.update(append_info.to_api_repr()) + self._details = append_info def close(self) -> None: """Close the file, flushing any written data, and shut down its executor.""" @@ -2497,10 +2518,16 @@ def commit(self) -> None: self.fs.invalidate_cache(self.path) def discard(self) -> None: - """Cancel pending part uploads and abort the multipart upload, if any.""" + """Abort the multipart upload, if any. + + The part uploads that have not started are cancelled, and the + running ones are waited for before the abort. + """ if self.multipart_upload: - for f in self.multipart_upload_parts: - f.cancel() + # A part that is still uploading when the upload is aborted may + # be stored after the abort, so wait for the parts that could not + # be cancelled first. + wait([f for f in self.multipart_upload_parts if not f.cancel()]) # s3_additional_kwargs also holds object parameters (e.g., the # existing object's metadata in append mode) that # AbortMultipartUpload rejects. diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 77deccc2a..8d1177a6c 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -281,7 +281,12 @@ async def _copy_object_with_multipart_upload( block_size < S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE or block_size > S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE ): - raise ValueError("Block size must be greater than 5MiB and less than 5GiB.") + raise ValueError( + "Block size must be between " + f"5 MiB ({S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and " + f"5 GiB ({S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), " + f"inclusive: {block_size}." + ) copy_source: dict[str, Any] = { "Bucket": bucket1, diff --git a/pyathena/filesystem/s3_executor.py b/pyathena/filesystem/s3_executor.py index 4cf80093b..475e21387 100644 --- a/pyathena/filesystem/s3_executor.py +++ b/pyathena/filesystem/s3_executor.py @@ -10,6 +10,7 @@ from __future__ import annotations import asyncio +import threading from abc import ABCMeta, abstractmethod from collections.abc import Callable from concurrent.futures import Future @@ -77,7 +78,9 @@ class S3AioExecutor(S3Executor): Uses ``asyncio.run_coroutine_threadsafe(asyncio.to_thread(fn), loop)`` to dispatch blocking functions onto the event loop's thread pool, returning ``concurrent.futures.Future`` objects that are compatible with - ``as_completed()`` and ``Future.cancel()``. + ``as_completed()``, ``wait()`` and ``Future.cancel()``. As with + ``ThreadPoolExecutor``, a future cannot be cancelled once its function has + started. This avoids thread-in-thread nesting when ``S3File`` is used from within ``asyncio.to_thread()`` calls (the pattern used by ``AioS3FileSystem``). @@ -101,9 +104,47 @@ def __init__(self, loop: asyncio.AbstractEventLoop | None = None) -> None: @override def submit(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> Future[T]: if self._loop is not None and self._loop.is_running(): - return asyncio.run_coroutine_threadsafe( - asyncio.to_thread(fn, *args, **kwargs), self._loop - ) + # The future of run_coroutine_threadsafe can be cancelled while + # the function keeps running in its thread, so the returned future + # is started and resolved by the function's thread instead. + future: Future[T] = Future() + # Acquired once, by run() or by settle(), whichever comes first, + # so that the future is started or settled exactly once. + claim = threading.Lock() + + def run() -> None: + """Run the function and resolve the future unless it was cancelled.""" + if not claim.acquire(blocking=False) or not future.set_running_or_notify_cancel(): + return + try: + result = fn(*args, **kwargs) + except BaseException as e: + future.set_exception(e) + else: + future.set_result(result) + + def settle(task: Future[None]) -> None: + """Resolve the future if the task ended before the function started. + + This happens, for example, when the event loop shuts down. + + Args: + task: The finished future of the task that runs the function. + """ + if not claim.acquire(blocking=False): + # run() has claimed the future and resolves it. + return + if task.cancelled(): + future.cancel() + # Notify the waiters of the cancellation, as an executor + # does when it drops a cancelled function. + future.set_running_or_notify_cancel() + elif future.set_running_or_notify_cancel(): + future.set_exception(task.exception()) + + task = asyncio.run_coroutine_threadsafe(asyncio.to_thread(run), self._loop) + task.add_done_callback(settle) + return future raise RuntimeError( "S3AioExecutor requires a running event loop. " "Use S3ThreadPoolExecutor for synchronous usage." diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2131eb1e3..95b47d016 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1,12 +1,16 @@ +import asyncio import functools +import gc import io import os +import sys import tempfile +import threading import time import urllib.parse import urllib.request import uuid -from concurrent.futures import Future, ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor, wait from datetime import UTC, datetime from itertools import chain from pathlib import Path @@ -19,6 +23,7 @@ import pyathena from pyathena.filesystem import register_s3_filesystem from pyathena.filesystem.s3 import S3File, S3FileSystem +from pyathena.filesystem.s3_executor import S3AioExecutor from pyathena.filesystem.s3_object import S3Object, S3ObjectType, S3StorageClass from pyathena.util import RetryConfig from tests import ENV @@ -551,6 +556,69 @@ def test_open_version_id_for_writing(self, mode, path, kwargs): with pytest.raises(ValueError, match="version specified"): fs.open(path, mode, **kwargs) + @pytest.mark.parametrize("mode", ["wb", "ab", "xb"]) + @pytest.mark.parametrize( + ("path", "block_size", "match"), + [ + # GH-926: the message states the accepted range. + ( + "s3://bucket/key", + S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE - 1, + r"between 5 MiB \(5242880 bytes\) and 5 GiB \(5368709120 bytes\), inclusive", + ), + # GH-952: a part cannot be larger than the maximum part size. + ("s3://bucket/key", S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE + 1, "between"), + ("s3://bucket", S3FileSystem.DEFAULT_BLOCK_SIZE, "does not contain a key"), + ], + ) + def test_open_invalid_for_writing(self, monkeypatch, mode, path, block_size, match): + # GH-976: an open() that fails validation sends no request and leaves + # no half-initialized file, whose garbage collection would close it. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._call.side_effect = AssertionError("No request is expected.") + unraisable = [] + monkeypatch.setattr(sys, "unraisablehook", unraisable.append) + + with pytest.raises(ValueError, match=match): + fs.open(path, mode, block_size=block_size) + gc.collect() + + assert unraisable == [] + + def test_open_append_lookup_failure(self, monkeypatch): + # GH-976: an append whose lookup of the existing object fails leaves + # no half-initialized file, whose garbage collection would close it. + fs = self._make_fs() + fs.default_cache_type = "bytes" + + def exists(path): + # A new exception each time: one kept by a mock would keep its + # traceback, and the file, alive. + raise PermissionError("denied") + + fs.exists = exists + unraisable = [] + monkeypatch.setattr(sys, "unraisablehook", unraisable.append) + + with pytest.raises(PermissionError, match="denied"): + fs.open("s3://bucket/key", "ab") + gc.collect() + + assert unraisable == [] + fs._call.assert_not_called() + + @pytest.mark.parametrize( + "block_size", + [S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE], + ) + def test_open_block_size_limits_for_writing(self, block_size): + fs = self._make_fs() + fs.default_cache_type = "bytes" + + with fs.open("s3://bucket/key", "wb", block_size=block_size) as f: + assert f.blocksize == block_size + @pytest.mark.parametrize( ("path", "expected"), [ @@ -637,6 +705,80 @@ def test_finish_multipart_upload_abort_failure_does_not_mask_the_original_error( bucket="bucket", key="key", upload_id="uploadid", futures=[future] ) + def test_finish_multipart_upload_waits_for_running_parts(self): + # GH-976: a part that is still uploading when the upload is aborted + # may be stored after the abort, so the abort waits for it. The + # parts that have not started are cancelled. + fs = self._make_fs() + fs._complete_multipart_upload = mock.MagicMock() + events = [] + fs._call.side_effect = lambda *args, **kwargs: events.append("abort") + failed: Future[SimpleNamespace] = Future() + failed.set_exception(RuntimeError("upload failed")) + started = threading.Event() + release = threading.Event() + + def upload_part(): + started.set() + # Uploading until the abort waits for it, so that an abort that + # does not wait comes first. + release.wait(5) + events.append("part 2 stored") + + def wait_parts(futures): + release.set() + return wait(futures) + + with ( + ThreadPoolExecutor(max_workers=1) as executor, + mock.patch("pyathena.filesystem.s3.wait", side_effect=wait_parts) as waited, + ): + running = executor.submit(upload_part) + pending = executor.submit(events.append, "part 3 stored") + started.wait(5) + with pytest.raises(RuntimeError, match="upload failed"): + fs._finish_multipart_upload( + bucket="bucket", + key="key", + upload_id="uploadid", + futures=[failed, running, pending], + ) + + assert events == ["part 2 stored", "abort"] + waited.assert_called_once_with([failed, running]) + assert pending.cancelled() + + def test_finish_multipart_upload_does_not_wait_for_cancelled_parts(self): + # GH-976: a cancelled part is not waited for, as nothing may + # acknowledge its cancellation, e.g., an event loop blocked by the + # caller. + fs = self._make_fs() + fs._complete_multipart_upload = mock.MagicMock() + failed: Future[SimpleNamespace] = Future() + failed.set_exception(RuntimeError("upload failed")) + never_started: Future[SimpleNamespace] = Future() + errors = [] + + def finish(): + try: + fs._finish_multipart_upload( + bucket="bucket", + key="key", + upload_id="uploadid", + futures=[failed, never_started], + ) + except RuntimeError as e: + errors.append(e) + + thread = threading.Thread(target=finish, daemon=True) + thread.start() + thread.join(5) + + assert not thread.is_alive() + assert [str(e) for e in errors] == ["upload failed"] + assert never_started.cancelled() + fs._call.assert_called_once() + @pytest.mark.parametrize( ("size", "block_size", "ranges"), [ @@ -694,6 +836,31 @@ def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): (2, (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)), ] + @pytest.mark.parametrize( + "block_size", + [ + S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE - 1, + S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE + 1, + ], + ) + def test_copy_object_with_multipart_upload_invalid_block_size(self, block_size): + # GH-926: the message states the accepted range. + fs = self._make_fs() + + with pytest.raises( + ValueError, + match=r"between 5 MiB \(5242880 bytes\) and 5 GiB \(5368709120 bytes\), inclusive", + ): + fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=5 * 2**30 + 2**20, + bucket2="bucket", + key2="dst", + block_size=block_size, + ) + fs._call.assert_not_called() + def test_head_object_version_aware(self): fs = self._make_fs() fs._call.return_value = {"ContentLength": 4, "ETag": '"etag"', "VersionId": "v1"} @@ -2305,6 +2472,66 @@ def test_discard(self, multipart): assert file.multipart_upload is None assert file.multipart_upload_parts == [] + def test_discard_waits_for_running_parts(self): + # GH-976: a part that is still uploading when the upload is aborted + # may be stored after the abort, so the abort waits for it. The + # parts that have not started are cancelled. + file = self._make_write_file(b"", autocommit=False) + file.multipart_upload = SimpleNamespace(upload_id="uploadid") + events = [] + file.fs._call.side_effect = lambda *args, **kwargs: events.append("abort") + started = threading.Event() + release = threading.Event() + + def upload_part(): + started.set() + # Uploading until the abort waits for it, so that an abort that + # does not wait comes first. + release.wait(5) + events.append("part 1 stored") + + def wait_parts(futures): + release.set() + return wait(futures) + + with ( + ThreadPoolExecutor(max_workers=1) as executor, + mock.patch("pyathena.filesystem.s3.wait", side_effect=wait_parts) as waited, + ): + running = executor.submit(upload_part) + pending = executor.submit(events.append, "part 2 stored") + file.multipart_upload_parts = [running, pending] + started.wait(5) + file.discard() + + assert events == ["part 1 stored", "abort"] + waited.assert_called_once_with([running]) + assert pending.cancelled() + + def test_discard_on_event_loop_thread(self): + # GH-976: the parts that have not started are cancelled and not + # waited for, so a rollback on the thread of the event loop that + # would run them does not block. + file = self._make_write_file(b"", autocommit=False) + file.multipart_upload = SimpleNamespace(upload_id="uploadid") + parts = [] + + async def rollback(): + executor = S3AioExecutor(loop=asyncio.get_running_loop()) + parts.extend(executor.submit(file.fs._upload_part) for _ in range(2)) + file.multipart_upload_parts = list(parts) + file.discard() + + thread = threading.Thread(target=asyncio.run, args=(rollback(),), daemon=True) + thread.start() + thread.join(5) + + assert not thread.is_alive() + assert all(part.cancelled() for part in parts) + file.fs._upload_part.assert_not_called() + file.fs._call.assert_called_once() + assert file.fs._call.call_args.args[0] == "abort_multipart_upload" + @pytest.mark.parametrize("autocommit", [True, False]) def test_upload_chunk_multipart(self, autocommit): # Multipart upload (CompleteMultipartUpload), completed from the uploaded diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 7f5851288..e85438d4e 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -166,6 +166,33 @@ async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): (2, (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)), ] + @pytest.mark.parametrize( + "block_size", + [ + S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE - 1, + S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE + 1, + ], + ) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_invalid_block_size(self, block_size): + # GH-926: the message states the accepted range. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._sync_fs._call = mock.MagicMock() + + with pytest.raises( + ValueError, + match=r"between 5 MiB \(5242880 bytes\) and 5 GiB \(5368709120 bytes\), inclusive", + ): + await fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=5 * 2**30 + 2**20, + bucket2="bucket", + key2="dst", + block_size=block_size, + ) + fs._sync_fs._call.assert_not_called() + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): diff --git a/tests/pyathena/filesystem/test_s3_executor.py b/tests/pyathena/filesystem/test_s3_executor.py new file mode 100644 index 000000000..4168c04e9 --- /dev/null +++ b/tests/pyathena/filesystem/test_s3_executor.py @@ -0,0 +1,124 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +import asyncio +import threading +import time +from concurrent.futures import Future, ThreadPoolExecutor, wait + +import pytest + +from pyathena.filesystem.s3_executor import S3AioExecutor + + +class TestS3AioExecutor: + def test_submit(self): + async def main(): + executor = S3AioExecutor(loop=asyncio.get_running_loop()) + succeeded = executor.submit(sum, [1, 2]) + failed = executor.submit(int, "x") + await asyncio.to_thread(wait, [succeeded, failed]) + return succeeded, failed + + succeeded, failed = asyncio.run(main()) + + assert succeeded.result() == 3 + with pytest.raises(ValueError, match="invalid literal"): + failed.result() + + def test_submit_without_running_loop(self): + with pytest.raises(RuntimeError, match="requires a running event loop"): + S3AioExecutor().submit(sum, [1, 2]) + + def test_cancel(self): + # GH-976: as with ThreadPoolExecutor, a future can be cancelled only + # before its function starts, so that wait() waits for a running one. + events = [] + started = threading.Event() + cancelled = threading.Event() + + def work(): + started.set() + cancelled.wait(5) + time.sleep(0.05) + events.append("finished") + + def cancel(executor: S3AioExecutor) -> tuple[Future[None], Future[None]]: + running = executor.submit(work) + pending = executor.submit(events.append, "pending finished") + started.wait(5) + assert not running.cancel() + assert pending.cancel() + cancelled.set() + _, not_done = wait([running, pending], timeout=5) + assert not not_done + events.append("waited") + return running, pending + + async def main(): + loop = asyncio.get_running_loop() + loop.set_default_executor(ThreadPoolExecutor(max_workers=2)) + executor = S3AioExecutor(loop=loop) + # The function runs in one of the two threads, and cancel() in the other. + return await asyncio.to_thread(cancel, executor) + + running, pending = asyncio.run(main()) + + assert events == ["finished", "waited"] + assert running.done() + assert not running.cancelled() + assert pending.cancelled() + + def test_loop_shutdown(self): + # 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. + events = [] + started = threading.Event() + settled = threading.Event() + + def work(): + started.set() + # Running until the pending future is settled, so that the + # pending function cannot start before the shutdown. + settled.wait(5) + events.append("finished") + + async def main(): + loop = asyncio.get_running_loop() + loop.set_default_executor(ThreadPoolExecutor(max_workers=1)) + executor = S3AioExecutor(loop=loop) + running = executor.submit(work) + pending = executor.submit(events.append, "pending finished") + pending.add_done_callback(lambda _: settled.set()) + while not started.is_set(): + await asyncio.sleep(0.01) + return running, pending + + running, pending = asyncio.run(main()) + + _, not_done = wait([running, pending], timeout=5) + assert not not_done + assert events == ["finished"] + assert not running.cancelled() + assert pending.cancelled() + + def test_executor_shut_down(self): + # An error raised before the function starts is set on the future. + async def main(): + loop = asyncio.get_running_loop() + default_executor = ThreadPoolExecutor(max_workers=1) + loop.set_default_executor(default_executor) + default_executor.shutdown() + future = S3AioExecutor(loop=loop).submit(sum, [1, 2]) + with pytest.raises(RuntimeError, match="after shutdown"): + await asyncio.wrap_future(future) + return future + + future = asyncio.run(main()) + + assert not future.cancelled()