diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 0d98a28a..588932d8 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -23,6 +23,7 @@ from pyathena.filesystem.s3 import CompressedBuffer, S3File, S3FileSystem from pyathena.filesystem.s3_executor import S3AioExecutor, S3Executor, S3ThreadPoolExecutor from pyathena.filesystem.s3_object import ( + S3CompleteMultipartUpload, S3Metadata, S3MultipartUpload, S3Object, @@ -37,6 +38,8 @@ from pyathena.connection import Connection _logger = logging.getLogger(__name__) +# The running cleanups of cancelled multipart copies. +_cleanup_tasks: set[asyncio.Future[None]] = set() class AioS3FileSystem(AsyncFileSystem): @@ -485,8 +488,12 @@ async def _copy_object_with_multipart_upload( """Copy an object with a multipart upload of its byte ranges. See :meth:`S3FileSystem._copy_object_with_multipart_upload`. The part - and annotation copies run in parallel with ``asyncio.gather`` and - ``asyncio.to_thread``. + and annotation copies run in parallel as asyncio tasks with + ``asyncio.to_thread``. On a cancellation after the upload is + created, the running part copies and the completion are waited for, + the upload is aborted unless it has completed, and the cancellation + is re-raised. A repeated cancellation returns without stopping this + cleanup. Args: bucket1: Source S3 bucket name. @@ -589,25 +596,55 @@ async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any] | None } tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(ranges)] - try: - # gather keeps the part-number order of the tasks. - parts = await asyncio.gather(*tasks) - completed = await asyncio.to_thread( - self._sync_fs._complete_multipart_upload, - bucket=bucket2, - key=key2, - upload_id=upload_id, - parts=cast(list[dict[str, Any]], parts), - **self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs), - ) - except Exception: - failed = True + completion: asyncio.Task[S3CompleteMultipartUpload] | None = None + + async def _abort() -> None: # A part that is still copying when the upload is aborted may be # stored after the abort, so wait for the running parts first. await asyncio.gather(*tasks, return_exceptions=True) + if completion is not None: + await asyncio.wait([completion]) + if not completion.cancelled() and completion.exception() is None: + # The upload completed despite the cancellation, so + # there is nothing to abort. + return await asyncio.to_thread( self._sync_fs._abort_multipart_upload, bucket2, key2, upload_id, kwargs ) + + try: + # Unlike gather, wait does not cancel the parts when this task is + # cancelled; their threads would keep copying, so they are waited + # for in _abort(). + done, _ = await asyncio.wait(tasks, return_when=asyncio.FIRST_EXCEPTION) + for task in done: + if (error := task.exception()) is not None: + raise error + # The tasks are in part-number order. + parts = [task.result() for task in tasks] + completion = asyncio.ensure_future( + asyncio.to_thread( + self._sync_fs._complete_multipart_upload, + bucket=bucket2, + key=key2, + upload_id=upload_id, + parts=cast(list[dict[str, Any]], parts), + **self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs), + ) + ) + # shield keeps a cancellation from cancelling the completion, whose + # thread would keep running, so that _abort() can wait for it. + completed = await asyncio.shield(completion) + except BaseException: + # Also on cancellation, as S3FileSystem._finish_multipart_upload + # does on an interrupt. + failed = True + cleanup = asyncio.ensure_future(_abort()) + # A repeated cancellation of this task returns without stopping + # the cleanup; the event loop keeps only weak references to tasks. + _cleanup_tasks.add(cleanup) + cleanup.add_done_callback(_cleanup_tasks.discard) + await asyncio.shield(cleanup) raise failed = False diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index b35eb1c2..339d0e7e 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -356,6 +356,151 @@ def upload_part_copy(**kw): assert events[2:] == ["end 1", "abort"] sync_fs._complete_multipart_upload.assert_not_called() + @pytest.mark.parametrize("cancellations", [1, 2]) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_cancelled(self, cancellations): + # GH-1046: a cancellation waits for the part copies that are running, + # aborts the upload, and is re-raised, as S3FileSystem does on an + # interrupt. A repeated cancellation returns without stopping the + # cleanup. + fs = AioS3FileSystem(connection=mock.MagicMock(), max_workers=2, skip_instance_cache=True) + sync_fs = fs._sync_fs + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + events = [] + lock = threading.Lock() + started = threading.Semaphore(0) + release = threading.Event() + aborted = threading.Event() + + def upload_part_copy(**kw): + part_number = kw["part_number"] + with lock: + events.append(f"start {part_number}") + started.release() + # The finally blocks of the test always release it. + release.wait() + with lock: + events.append(f"end {part_number}") + return SimpleNamespace(etag='"e"', part_number=part_number) + + def abort_multipart_upload(*args): + events.append("abort") + aborted.set() + + sync_fs._upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) + sync_fs._complete_multipart_upload = mock.MagicMock() + # The HeadObject of the source, for its version. + sync_fs._call = mock.MagicMock(return_value={}) + sync_fs._abort_multipart_upload = mock.MagicMock(side_effect=abort_multipart_upload) + + task = asyncio.ensure_future( + fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=3 * S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + bucket2="bucket", + key2="dst", + block_size=S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + ) + try: + for _ in range(2): + assert await asyncio.to_thread(started.acquire, timeout=5) + # The running parts are held until the copy has been cancelled. + for _ in range(cancellations): + task.cancel() + # Lets the copy enter its cleanup. + await asyncio.sleep(0) + if cancellations > 1: + # The repeated cancellation returns while the parts still run. + assert task.done() + assert events[2:] == [] + else: + assert not task.done() + except BaseException: + task.cancel() + raise + finally: + release.set() + with pytest.raises(asyncio.CancelledError): + await task + assert await asyncio.to_thread(aborted.wait, 5) + + # Part 3 waits for a worker and is not started after the cancellation. + assert sorted(events[:2]) == ["start 1", "start 2"] + assert sorted(events[2:4]) == ["end 1", "end 2"] + assert events[4:] == ["abort"] + sync_fs._complete_multipart_upload.assert_not_called() + + @pytest.mark.parametrize("completion_fails", [False, True]) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_cancelled_completion(self, completion_fails): + # GH-1046: a cancellation during CompleteMultipartUpload waits for it, + # aborts the upload only if it failed, and is re-raised. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + sync_fs = fs._sync_fs + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + events = [] + started = threading.Event() + release = threading.Event() + + def complete_multipart_upload(**kw): + started.set() + # The finally blocks of the test always release it. + release.wait() + events.append("complete") + if completion_fails: + raise OSError("completion failed") + return SimpleNamespace() + + sync_fs._upload_part_copy = mock.MagicMock( + side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + ) + sync_fs._complete_multipart_upload = mock.MagicMock(side_effect=complete_multipart_upload) + # The HeadObject of the source, for its version. + sync_fs._call = mock.MagicMock(return_value={}) + sync_fs._abort_multipart_upload = mock.MagicMock( + side_effect=lambda *args: events.append("abort") + ) + + task = asyncio.ensure_future( + fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=2 * S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + bucket2="bucket", + key2="dst", + block_size=S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + ) + try: + assert await asyncio.to_thread(started.wait, 5) + task.cancel() + # Gives the cleanup time to abort early or to return, which it + # must not do while the completion is held. + await asyncio.sleep(0.1) + assert not task.done() + assert events == [] + except BaseException: + task.cancel() + raise + finally: + release.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert events == (["complete", "abort"] if completion_fails else ["complete"]) + @pytest.mark.parametrize( "block_size", [