From 9187b6f203aca5bd2f9bcbe1abdf7c5172e63779 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:13:32 +0900 Subject: [PATCH 1/9] Move the copy requests and the multipart copy plan into S3Core Add S3Core.copy_object(), plan_multipart_copy() with the frozen S3MultipartCopyPlan, list_object_annotations() and copy_object_annotation(). The planner owns the reads of the source (HeadObject, GetObjectTagging, ListObjectAnnotations), the block size validation, the version pinning, the CopyObject directives, the copy-source parameter mapping and the directory-bucket rule, and filters the parameters of each multipart request. The sync and aio multipart copies now only schedule a plan: sync with the executor and _finish_multipart_upload(), aio with to_thread and its unchanged cancellation handling. A HeadObject response without a size now raises ValueError instead of falling back to the cached size. Part of #1063 (sub-step 3). Co-Authored-By: Claude Opus 5.5 --- docs/api/filesystem.rst | 3 + docs/filesystem.md | 18 +- pyathena/filesystem/s3.py | 421 ++------------------- pyathena/filesystem/s3_async.py | 111 ++---- pyathena/filesystem/s3_core.py | 392 ++++++++++++++++++- tests/pyathena/filesystem/test_s3.py | 109 ++---- tests/pyathena/filesystem/test_s3_async.py | 101 ++--- tests/pyathena/filesystem/test_s3_core.py | 292 ++++++++++++++ 8 files changed, 843 insertions(+), 604 deletions(-) diff --git a/docs/api/filesystem.rst b/docs/api/filesystem.rst index de5b733c..170e1c2f 100644 --- a/docs/api/filesystem.rst +++ b/docs/api/filesystem.rst @@ -82,6 +82,9 @@ S3 Core .. autoclass:: pyathena.filesystem.s3_core.S3DeleteError :members: +.. autoclass:: pyathena.filesystem.s3_core.S3MultipartCopyPlan + :members: + S3 Objects ---------- diff --git a/docs/filesystem.md b/docs/filesystem.md index 0202a8a7..39ad8f1a 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -280,9 +280,10 @@ directories below the bucket level) and is always a no-op. ## Typed S3 operations `S3FileSystem.core` is an `S3Core`, the typed operations that the filesystem sends -its listing, lookup, delete and multipart upload requests with. It can also be built -on a boto3 S3 client. Each operation sends one request (one per page for the -iterators) with the retry policy, raises `FileNotFoundError` for a missing bucket or +its listing, lookup, delete, multipart upload and copy requests with. It can also be +built on a boto3 S3 client. Each operation sends one request (one per page for the +iterators and `list_object_annotations()`), except the copy operations described below, +with the retry policy, raises `FileNotFoundError` for a missing bucket or multipart upload, or for a missing object or version that it reads, and caches nothing. Requests sent through `fs.core` do not invalidate the filesystem's cache: call @@ -331,6 +332,17 @@ ranges of the parts that copy it, by the part limits `MULTIPART_UPLOAD_MIN_PART_ `MULTIPART_UPLOAD_MAX_PART_SIZE` (5 GiB) and `MULTIPART_UPLOAD_MAX_PARTS` (10,000) of `S3Core`. +`copy_object()` copies an object with one CopyObject request, which accepts objects +up to `MULTIPART_UPLOAD_MAX_PART_SIZE`. For a larger object, `plan_multipart_copy()` +reads the source and returns an `S3MultipartCopyPlan`: the version to copy, the byte +ranges of the parts, the parameters of each multipart upload request, and the +annotations to copy, so that the multipart upload writes the metadata, tags and +annotations that CopyObject would. It sends HeadObject, then GetObjectTagging and +ListObjectAnnotations unless the directives or the source exclude them, and writes +nothing. `copy_object_annotation()` copies one annotation onto the destination after +the upload completes, with GetObjectAnnotation and PutObjectAnnotation. The +filesystems' `cp_file()` and `copy()` run these plans. + ## Async filesystem `AioS3FileSystem` provides the same functionality on top of fsspec's diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 846f56cb..cb601d11 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -19,7 +19,7 @@ from multiprocessing import cpu_count from re import Pattern from typing import Any, BinaryIO, cast -from urllib.parse import unquote_plus, urlencode +from urllib.parse import unquote_plus import botocore.exceptions from boto3 import Session @@ -190,18 +190,6 @@ class S3FileSystem(AbstractFileSystem): "SSECustomerKeyMD5", } ) - # https://docs.aws.amazon.com/AmazonS3/latest/API/API_CopyObject.html - # The metadata that CopyObject copies from the source with the COPY - # metadata directive, which ignores the values given in the request. - _COPY_METADATA_PARAMS: tuple[str, ...] = ( - "CacheControl", - "ContentDisposition", - "ContentEncoding", - "ContentLanguage", - "ContentType", - "Expires", - "Metadata", - ) PATTERN_PATH: Pattern[str] = S3Path.PATTERN protocol = ("s3", "s3a") @@ -1726,22 +1714,11 @@ def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: size1 = info1.get("size", 0) try: if size1 <= self.core.MULTIPART_UPLOAD_MAX_PART_SIZE: - self._copy_object( - bucket1=source.bucket, - key1=source.key, - version_id1=source.version_id, - bucket2=destination.bucket, - key2=destination.key, - **kwargs, - ) + self.core.copy_object(source, destination, **kwargs) else: self._copy_object_with_multipart_upload( - bucket1=source.bucket, - key1=source.key, - version_id1=source.version_id, - size1=size1, - bucket2=destination.bucket, - key2=destination.key, + source, + destination, max_workers=max_workers, block_size=block_size, **kwargs, @@ -1752,66 +1729,30 @@ def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: self.invalidate_cache(path2) return True - def _copy_object( - self, - bucket1: str, - key1: str, - version_id1: str | None, - bucket2: str, - key2: str, - **kwargs, - ) -> None: - copy_source = { - "Bucket": bucket1, - "Key": key1, - } - if version_id1: - copy_source.update({"VersionId": version_id1}) - request = { - "CopySource": copy_source, - "Bucket": bucket2, - "Key": key2, - } - - _logger.debug( - f"Copy object from {S3Path(bucket1, key1, version_id1).uri} " - f"to {S3Path(bucket2, key2).uri}." - ) - self._call(self._client.copy_object, **request, **kwargs) - def _copy_object_with_multipart_upload( self, - bucket1: str, - key1: str, - size1: int, - bucket2: str, - key2: str, + source: S3Path, + destination: S3Path, max_workers: int | None = None, block_size: int | None = None, - version_id1: str | None = None, **kwargs, ) -> None: """Copy an object with a multipart upload of its byte ranges. - The parts are copied in parallel with UploadPartCopy. The upload - gets the metadata and tags that CopyObject would copy (see - :meth:`_get_multipart_copy_kwargs`). The annotations of the source - are listed before the upload is created and copied onto the - destination after it completes. - A failed part or completion aborts the upload; a failed annotation - copy is raised and leaves the destination in place. If HeadObject - reports a size that fits in a single CopyObject request, the - reported version is copied with CopyObject instead. - - Args: - bucket1: Source S3 bucket name. - key1: Source object key. - size1: Size of the source object in bytes. - bucket2: Destination S3 bucket name. - key2: Destination object key. + Runs the plan of :meth:`S3Core.plan_multipart_copy`, which reads the + source and lists its annotations before anything is written. The + parts are copied in parallel with UploadPartCopy, and the + annotations are copied onto the destination after the upload + completes. A failed part or completion aborts the upload; a failed + annotation copy is raised and leaves the destination in place. If + HeadObject reports a size that fits in a single CopyObject request, + the reported version is copied with CopyObject instead. + + Args: + source: Source S3 path, with the version ID to copy, if any. + destination: Destination S3 path. max_workers: Maximum number of parallel requests. block_size: Size in bytes of the copied ranges. - version_id1: Source version ID, if any. **kwargs: The CopyObject parameters of the copy; each request receives those that it accepts. @@ -1820,325 +1761,47 @@ def _copy_object_with_multipart_upload( directive has an invalid value. """ max_workers = max_workers if max_workers else self.max_workers - block_size = block_size if block_size else self.core.MULTIPART_UPLOAD_MAX_PART_SIZE - if ( - block_size < self.core.MULTIPART_UPLOAD_MIN_PART_SIZE - or block_size > self.core.MULTIPART_UPLOAD_MAX_PART_SIZE - ): - raise ValueError( - "Block size must be between " - f"5 MiB ({self.core.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and " - f"5 GiB ({self.core.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), " - f"inclusive: {block_size}." - ) - - create_kwargs, version_id1, head_size = self._get_multipart_copy_kwargs( - bucket1, key1, version_id1, kwargs - ) - if head_size is not None and head_size <= self.core.MULTIPART_UPLOAD_MAX_PART_SIZE: + plan = self.core.plan_multipart_copy(source, destination, block_size, **kwargs) + if plan.fits_single_request: # The size that the caller found, which may come from a cached # listing, was larger than the copied version, which fits in a # single CopyObject request. - self._copy_object( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, - **kwargs, - ) + self.core.copy_object(plan.source, plan.destination, **kwargs) return - # The size of the copied version, not the one that the caller found. - ranges = self.core.part_ranges(size1 if head_size is None else head_size, block_size) - source = S3Path(bucket1, key1, version_id1) - destination = S3Path(bucket2, key2) - # The annotations are listed before anything is written, so that a - # missing permission fails first. - annotations = ( - self._list_object_annotations(bucket1, key1, version_id1, kwargs) - if self._copies_annotations(bucket1, kwargs) - else [] - ) - multipart_upload = self.core.create_multipart_upload(destination, **create_kwargs) + multipart_upload = self.core.create_multipart_upload(plan.destination, **plan.create_params) + upload_id = cast(str, multipart_upload.upload_id) with self._create_executor(max_workers=max_workers) as executor: futures = [ executor.submit( self.core.upload_part_copy, - path=destination, - upload_id=cast(str, multipart_upload.upload_id), + path=plan.destination, + upload_id=upload_id, part_number=i + 1, - source=source, + source=plan.source, range_=range_, - **self.core.operation_params("upload_part_copy", kwargs), + **plan.part_params, ) - for i, range_ in enumerate(ranges) + for i, range_ in enumerate(plan.ranges) ] completed = self._finish_multipart_upload( - bucket=bucket2, - key=key2, - upload_id=cast(str, multipart_upload.upload_id), + bucket=plan.destination.bucket, + key=cast(str, plan.destination.key), + upload_id=upload_id, futures=futures, - request_kwargs=kwargs, + # The completion and the abort each receive those that they + # accept, as filtered for the plan. + request_kwargs={**plan.complete_params, **plan.abort_params}, ) - for name in annotations: - self._copy_object_annotation( - name, bucket1, key1, version_id1, bucket2, key2, completed, kwargs + for name in plan.annotations: + self.core.copy_object_annotation( + name, + plan.source, + plan.destination, + completed.version_id, + completed.etag, + **kwargs, ) - @staticmethod - def _is_directory_bucket(bucket: str) -> bool: - """Return whether the bucket is a directory bucket (S3 Express One Zone). - - Directory bucket names end with ``--x-s3``. - - Args: - bucket: S3 bucket name. - - Returns: - True if the bucket is a directory bucket. - """ - return bucket.endswith("--x-s3") - - @staticmethod - def _get_copy_source_kwargs(kwargs: Mapping[str, Any]) -> dict[str, Any]: - """Map the parameters of a copy to those of the requests that read its source. - - Args: - kwargs: The CopyObject parameters of the copy. - - Returns: - ``RequestPayer``, and the source's expected bucket owner and SSE-C - parameters under the names of the requests that read the source - (``ExpectedBucketOwner`` and ``SSECustomer*``), where given. - """ - source_kwargs = { - "RequestPayer": kwargs.get("RequestPayer"), - "ExpectedBucketOwner": kwargs.get("ExpectedSourceBucketOwner"), - "SSECustomerAlgorithm": kwargs.get("CopySourceSSECustomerAlgorithm"), - "SSECustomerKey": kwargs.get("CopySourceSSECustomerKey"), - "SSECustomerKeyMD5": kwargs.get("CopySourceSSECustomerKeyMD5"), - } - return {k: v for k, v in source_kwargs.items() if v is not None} - - def _get_multipart_copy_kwargs( - self, bucket: str, key: str, version_id: str | None, kwargs: Mapping[str, Any] - ) -> tuple[dict[str, Any], str | None, int | None]: - """Build the CreateMultipartUpload parameters of a multipart copy. - - The source is read with HeadObject. Without a given version, a - version ID other than ``null`` that it reports is the version to - copy, so that the parts, the tags and the annotations come from the - same object even if the source is replaced during the copy. A - ``null`` version, which a write can replace, is not pinned. - - No multipart request accepts the directives of CopyObject, so they - are implemented here as CopyObject applies them. With the COPY - metadata directive (the default), the content headers and the - user-defined metadata of the source are used, and the values of the - copy are ignored. With the COPY tagging directive - (the default), the tags are read with GetObjectTagging, and the - ``Tagging`` of the copy is ignored. REPLACE uses the values of the - copy instead. CopyObject parameters that CreateMultipartUpload does - not accept, such as the source conditions, which go to the part - copies, are left out. - - Args: - bucket: Source S3 bucket name. - key: Source object key. - version_id: Source version ID, if any. - kwargs: The CopyObject parameters of the copy. - - Returns: - The parameters for CreateMultipartUpload, the version of the - source to copy (the given one, the one that HeadObject reported, - or None), and the size of that version from HeadObject. The - parameters are empty, without reading the tags, if the size fits - in a single CopyObject request. - - Raises: - ValueError: If a directive has a value that CopyObject does not - accept. - """ - metadata_directive = kwargs.get("MetadataDirective", "COPY") - tagging_directive = kwargs.get("TaggingDirective", "COPY") - annotation_directive = kwargs.get("AnnotationDirective", "COPY") - if metadata_directive not in ("COPY", "REPLACE"): - raise ValueError(f"Invalid MetadataDirective: {metadata_directive}.") - if tagging_directive not in ("COPY", "REPLACE"): - raise ValueError(f"Invalid TaggingDirective: {tagging_directive}.") - if annotation_directive not in ("COPY", "EXCLUDE"): - raise ValueError(f"Invalid AnnotationDirective: {annotation_directive}.") - - request = dict(kwargs) - source_kwargs = self._get_copy_source_kwargs(kwargs) - source = {"Bucket": bucket, "Key": key} - if version_id: - source.update({"VersionId": version_id}) - _logger.debug(f"Head object to copy: {S3Path(bucket, key, version_id).uri}") - head = self.core.head_object( - S3Path(bucket, key, version_id), - **self.core.operation_params("head_object", source_kwargs), - ) - if not version_id and head.version_id and head.version_id != "null": - version_id = head.version_id - source.update({"VersionId": version_id}) - if ( - head.content_length is not None - and head.content_length <= self.core.MULTIPART_UPLOAD_MAX_PART_SIZE - ): - # Copied with CopyObject instead, which applies the directives. - return {}, version_id, head.content_length - if metadata_directive == "COPY": - for name in self._COPY_METADATA_PARAMS: - request.pop(name, None) - copied = { - "CacheControl": head.cache_control, - "ContentDisposition": head.content_disposition, - "ContentEncoding": head.content_encoding, - "ContentLanguage": head.content_language, - "ContentType": head.content_type, - "Expires": head.expires, - "Metadata": head.user_metadata, - } - request.update({k: v for k, v in copied.items() if v is not None}) - if tagging_directive == "COPY": - request.pop("Tagging", None) - # Directory buckets do not support GetObjectTagging, and their - # objects have no tags. - if not self._is_directory_bucket(bucket): - _logger.debug(f"Get tags to copy: {S3Path(bucket, key, version_id).uri}") - response = self._call( - self._client.get_object_tagging, - **self.core.operation_params("get_object_tagging", source_kwargs), - **source, - ) - tags = [(t["Key"], t["Value"]) for t in response["TagSet"]] - if tags: - request.update({"Tagging": urlencode(tags)}) - copy_members = self._client.meta.service_model.operation_model( - "CopyObject" - ).input_shape.members - return ( - { - **self.core.operation_params("create_multipart_upload", request), - # A parameter that CopyObject does not accept either is sent as - # is, so that botocore rejects it as it does for CopyObject. - **{k: v for k, v in request.items() if k not in copy_members}, - }, - version_id, - head.content_length, - ) - - def _copies_annotations(self, bucket: str, kwargs: Mapping[str, Any]) -> bool: - """Return whether a multipart copy copies the annotations of its source. - - Args: - bucket: Source S3 bucket name. - kwargs: The CopyObject parameters of the copy. - - Returns: - True unless the ``AnnotationDirective`` is EXCLUDE or the source - cannot have annotations: an object encrypted with SSE-C, or an - object in a directory bucket. - """ - return ( - kwargs.get("AnnotationDirective", "COPY") == "COPY" - and "CopySourceSSECustomerAlgorithm" not in kwargs - and not self._is_directory_bucket(bucket) - ) - - def _list_object_annotations( - self, bucket: str, key: str, version_id: str | None, kwargs: Mapping[str, Any] - ) -> list[str]: - """List the names of the annotations of the source of a copy. - - Args: - bucket: Source S3 bucket name. - key: Source object key. - version_id: Source version ID, if any. - kwargs: The CopyObject parameters of the copy. - - Returns: - The annotation names, across all pages of ListObjectAnnotations. - """ - request: dict[str, Any] = { - **self.core.operation_params( - "list_object_annotations", self._get_copy_source_kwargs(kwargs) - ), - "Bucket": bucket, - "Key": key, - } - if version_id: - request.update({"VersionId": version_id}) - names: list[str] = [] - while True: - _logger.debug(f"List object annotations: {S3Path(bucket, key, version_id).uri}") - response = self._call(self._client.list_object_annotations, **request) - names.extend(a["AnnotationName"] for a in response.get("Annotations", [])) - token = response.get("NextContinuationToken") - if not token: - return names - request.update({"ContinuationToken": token}) - - def _copy_object_annotation( - self, - name: str, - bucket1: str, - key1: str, - version_id1: str | None, - bucket2: str, - key2: str, - completed: S3CompleteMultipartUpload, - kwargs: Mapping[str, Any], - ) -> None: - """Copy an annotation of the source of a copy onto its destination. - - The annotation is written to the version that the copy created, if - the bucket is versioned, and only if the destination still has the - ETag of the copy, so that it is not attached to an object written - over the copy. - - Args: - name: The annotation name. - bucket1: Source S3 bucket name. - key1: Source object key. - version_id1: Source version ID, if any. - bucket2: Destination S3 bucket name. - key2: Destination object key. - completed: The completion of the multipart upload of the copy. - kwargs: The CopyObject parameters of the copy. - """ - source: dict[str, Any] = {"Bucket": bucket1, "Key": key1, "AnnotationName": name} - if version_id1: - source.update({"VersionId": version_id1}) - _logger.debug( - f"Copy object annotation {name} from {S3Path(bucket1, key1, version_id1).uri} " - f"to s3://{bucket2}/{key2}." - ) - response = self._call( - self._client.get_object_annotation, - **self.core.operation_params( - "get_object_annotation", self._get_copy_source_kwargs(kwargs) - ), - **source, - ) - destination: dict[str, Any] = { - "Bucket": bucket2, - "Key": key2, - "AnnotationName": name, - "AnnotationPayload": response["AnnotationPayload"].read(), - } - if completed.version_id: - destination.update({"VersionId": completed.version_id}) - if completed.etag: - destination.update({"ObjectIfMatch": completed.etag}) - self._call( - self._client.put_object_annotation, - # The fields of the request take precedence over inherited - # parameters of the same name. - **{**self.core.operation_params("put_object_annotation", kwargs), **destination}, - ) - def _check_multipart_upload_size(self, path: str, size: int, block_size: int) -> None: """Check that data fits in a multipart upload before uploading it. diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 88f2d897..fc4a8fb8 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -532,23 +532,11 @@ async def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: size1 = info1.get("size", 0) try: if size1 <= self.core.MULTIPART_UPLOAD_MAX_PART_SIZE: - await asyncio.to_thread( - self._sync_fs._copy_object, - bucket1=source.bucket, - key1=source.key, - version_id1=source.version_id, - bucket2=destination.bucket, - key2=destination.key, - **kwargs, - ) + await asyncio.to_thread(self.core.copy_object, source, destination, **kwargs) else: await self._copy_object_with_multipart_upload( - bucket1=source.bucket, - key1=source.key, - version_id1=source.version_id, - size1=size1, - bucket2=destination.bucket, - key2=destination.key, + source, + destination, max_workers=max_workers, block_size=block_size, **kwargs, @@ -560,14 +548,10 @@ async def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: async def _copy_object_with_multipart_upload( self, - bucket1: str, - key1: str, - size1: int, - bucket2: str, - key2: str, + source: S3Path, + destination: S3Path, max_workers: int | None = None, block_size: int | None = None, - version_id1: str | None = None, **kwargs, ) -> None: """Copy an object with a multipart upload of its byte ranges. @@ -581,14 +565,10 @@ async def _copy_object_with_multipart_upload( cleanup. Args: - bucket1: Source S3 bucket name. - key1: Source object key. - size1: Size of the source object in bytes. - bucket2: Destination S3 bucket name. - key2: Destination object key. + source: Source S3 path, with the version ID to copy, if any. + destination: Destination S3 path. max_workers: Maximum number of parallel requests. block_size: Size in bytes of the copied ranges. - version_id1: Source version ID, if any. **kwargs: The CopyObject parameters of the copy; each request receives those that it accepts. @@ -597,52 +577,19 @@ async def _copy_object_with_multipart_upload( directive has an invalid value. """ max_workers = max_workers if max_workers else self._sync_fs.max_workers - block_size = block_size if block_size else self.core.MULTIPART_UPLOAD_MAX_PART_SIZE - if ( - block_size < self.core.MULTIPART_UPLOAD_MIN_PART_SIZE - or block_size > self.core.MULTIPART_UPLOAD_MAX_PART_SIZE - ): - raise ValueError( - "Block size must be between " - f"5 MiB ({self.core.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and " - f"5 GiB ({self.core.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), " - f"inclusive: {block_size}." - ) - - create_kwargs, version_id1, head_size = await asyncio.to_thread( - self._sync_fs._get_multipart_copy_kwargs, bucket1, key1, version_id1, kwargs + plan = await asyncio.to_thread( + self.core.plan_multipart_copy, source, destination, block_size, **kwargs ) - if head_size is not None and head_size <= self.core.MULTIPART_UPLOAD_MAX_PART_SIZE: + if plan.fits_single_request: # See S3FileSystem._copy_object_with_multipart_upload. - await asyncio.to_thread( - self._sync_fs._copy_object, - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, - **kwargs, - ) + await asyncio.to_thread(self.core.copy_object, plan.source, plan.destination, **kwargs) return - # The size of the copied version; see S3FileSystem. - ranges = self.core.part_ranges(size1 if head_size is None else head_size, block_size) - source = S3Path(bucket1, key1, version_id1) - destination = S3Path(bucket2, key2) - # Listed before anything is written; see S3FileSystem. - annotations = ( - await asyncio.to_thread( - self._sync_fs._list_object_annotations, bucket1, key1, version_id1, kwargs - ) - if self._sync_fs._copies_annotations(bucket1, kwargs) - else [] - ) multipart_upload = await asyncio.to_thread( - self.core.create_multipart_upload, destination, **create_kwargs + self.core.create_multipart_upload, plan.destination, **plan.create_params ) upload_id = cast(str, multipart_upload.upload_id) semaphore = asyncio.Semaphore(max_workers) - part_kwargs = self.core.operation_params("upload_part_copy", kwargs) failed = False async def _upload_part(i: int, range_: tuple[int, int]) -> S3MultipartUploadPart | None: @@ -654,19 +601,19 @@ async def _upload_part(i: int, range_: tuple[int, int]) -> S3MultipartUploadPart try: return await asyncio.to_thread( self.core.upload_part_copy, - path=destination, + path=plan.destination, upload_id=upload_id, part_number=i + 1, - source=source, + source=plan.source, range_=range_, - **part_kwargs, + **plan.part_params, ) except Exception: # Set before the semaphore lets a waiting part start. failed = True raise - tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(ranges)] + tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(plan.ranges)] completion: asyncio.Task[S3CompleteMultipartUpload] | None = None async def _abort() -> None: @@ -680,7 +627,11 @@ async def _abort() -> None: # there is nothing to abort. return await asyncio.to_thread( - self._sync_fs._abort_multipart_upload, bucket2, key2, upload_id, kwargs + self._sync_fs._abort_multipart_upload, + plan.destination.bucket, + cast(str, plan.destination.key), + upload_id, + plan.abort_params, ) try: @@ -696,10 +647,10 @@ async def _abort() -> None: completion = asyncio.ensure_future( asyncio.to_thread( self.core.complete_multipart_upload, - destination, + plan.destination, upload_id, cast(list[S3MultipartUploadPart], parts), - **self.core.operation_params("complete_multipart_upload", kwargs), + **plan.complete_params, ) ) # shield keeps a cancellation from cancelling the completion, whose @@ -728,22 +679,20 @@ async def _copy_annotation(name: str) -> None: return try: await asyncio.to_thread( - self._sync_fs._copy_object_annotation, + self.core.copy_object_annotation, name, - bucket1, - key1, - version_id1, - bucket2, - key2, - completed, - kwargs, + plan.source, + plan.destination, + completed.version_id, + completed.etag, + **kwargs, ) except Exception: failed = True raise results = await asyncio.gather( - *[_copy_annotation(name) for name in annotations], return_exceptions=True + *[_copy_annotation(name) for name in plan.annotations], return_exceptions=True ) for result in results: if isinstance(result, BaseException): diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index 0a29c195..f167e8e9 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -12,9 +12,10 @@ import logging import math from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import datetime from typing import Any, ClassVar, cast +from urllib.parse import urlencode import botocore.exceptions from botocore.client import BaseClient @@ -400,11 +401,57 @@ def from_response(cls, bucket: str, response: Mapping[str, Any]) -> S3DeleteResu ) +@dataclass(frozen=True) +class S3MultipartCopyPlan: + """The requests of a copy with a multipart upload, as CopyObject copies. + + :meth:`S3Core.plan_multipart_copy` reads the source and builds the plan; + the caller schedules the requests: CreateMultipartUpload with + ``create_params``, one UploadPartCopy per range with ``part_params``, + then CompleteMultipartUpload with ``complete_params``, or + AbortMultipartUpload with ``abort_params`` after a failure, and finally + the copy of each annotation with :meth:`S3Core.copy_object_annotation`. + If ``fits_single_request`` is true, the source is copied with + :meth:`S3Core.copy_object` instead, and the fields after ``size`` are + empty. + + Attributes: + source: The object to copy, with the version that HeadObject + reported unless the path has one or the version is ``null``. + destination: The object that the copy writes. + size: The size in bytes of the source, from HeadObject. + ranges: The ``(start, end)`` byte ranges of the source that the parts + copy, with an exclusive end, in part-number order. + create_params: The parameters of CreateMultipartUpload, with the + metadata and the tags that CopyObject would write. + part_params: The parameters of each UploadPartCopy. + complete_params: The parameters of CompleteMultipartUpload. + abort_params: The parameters of AbortMultipartUpload. + annotations: The names of the annotations to copy; empty if the + directive excludes them or the source cannot have any. + fits_single_request: Whether the source fits in a single CopyObject + request. + """ + + source: S3Path + destination: S3Path + size: int + ranges: tuple[tuple[int, int], ...] = () + create_params: Mapping[str, Any] = field(default_factory=dict) + part_params: Mapping[str, Any] = field(default_factory=dict) + complete_params: Mapping[str, Any] = field(default_factory=dict) + abort_params: Mapping[str, Any] = field(default_factory=dict) + annotations: tuple[str, ...] = () + fits_single_request: bool = False + + class S3Core: """Typed S3 operations, one request each, on a boto3 S3 client. Each operation sends one request, or one per page for the iterators, - with the retry policy, and translates S3 errors into ``OSError`` + except those that say otherwise, such as :meth:`plan_multipart_copy` + and :meth:`copy_object_annotation`. The requests are sent + with the retry policy, and S3 errors are translated into ``OSError`` subclasses (see :class:`~pyathena.filesystem.s3_errors.S3ClientError`): a missing bucket or multipart upload, or a missing object or version that an operation reads, raises ``FileNotFoundError``, and a denied request @@ -425,6 +472,18 @@ class S3Core: MULTIPART_UPLOAD_MAX_PART_SIZE: int = 5 * 2**30 # 5GiB # The maximum number of parts per multipart upload is 10,000. MULTIPART_UPLOAD_MAX_PARTS: int = 10_000 + # https://docs.aws.amazon.com/AmazonS3/latest/API/API_CopyObject.html + # The metadata that CopyObject copies from the source with the COPY + # metadata directive, which ignores the values given in the request. + _COPY_METADATA_PARAMS: ClassVar[tuple[str, ...]] = ( + "CacheControl", + "ContentDisposition", + "ContentEncoding", + "ContentLanguage", + "ContentType", + "Expires", + "Metadata", + ) def __init__( self, @@ -789,6 +848,335 @@ def part_ranges(self, size: int, block_size: int) -> list[tuple[int, int]]: starts.append(starts[-1] + (size - starts[-1]) // 2) return list(zip(starts, [*starts[1:], size], strict=True)) + def copy_object(self, source: S3Path, destination: S3Path, **params) -> None: + """Copy an object, or a version of it, with CopyObject. + + Args: + source: The path of the object to copy, with the version ID to + copy, if any. + destination: The path of the object to write, without a version + ID. + **params: Additional request parameters, sent as given. + + Raises: + ValueError: If the source or the destination has no key, or the + destination has a version ID, which a write cannot replace. + """ + if not source.key: + raise ValueError(f"The source has no key: {source.uri}.") + if not destination.key: + raise ValueError(f"The path has no key: {destination.uri}.") + if destination.version_id: + raise ValueError(f"Cannot write to a version: {destination.uri}.") + copy_source: dict[str, Any] = {"Bucket": source.bucket, "Key": source.key} + if source.version_id: + copy_source.update({"VersionId": source.version_id}) + request: dict[str, Any] = { + "CopySource": copy_source, + "Bucket": destination.bucket, + "Key": destination.key, + } + _logger.debug(f"Copy object from {source.uri} to {destination.uri}.") + self.call(self._client.copy_object, **request, **params) + + def plan_multipart_copy( + self, + source: S3Path, + destination: S3Path, + block_size: int | None = None, + **params, + ) -> S3MultipartCopyPlan: + """Plan a copy with a multipart upload that copies as CopyObject does. + + The source is read with HeadObject. Without a version in the path, a + version ID other than ``null`` that it reports is the version to + copy, so that the parts, the tags and the annotations come from the + same object even if the source is replaced during the copy. A + ``null`` version, which a write can replace, is not pinned. If the + reported size fits in a single CopyObject request, nothing else is + read and the plan says so. + + No multipart request accepts the directives of CopyObject, so the + plan applies them as CopyObject does. With the COPY metadata + directive (the default), the content headers and the user-defined + metadata of the source are used, and the values of ``params`` are + ignored. With the COPY tagging directive (the default), the tags are + read with GetObjectTagging, and the ``Tagging`` of ``params`` is + ignored. REPLACE uses the values of ``params`` instead. With the COPY + annotation directive (the default), the annotations are listed with + ListObjectAnnotations, unless the source is encrypted with SSE-C or + is in a directory bucket, which cannot have annotations; an object in + a directory bucket has no tags either. The requests that read the + source receive ``RequestPayer``, and the source's expected bucket + owner and SSE-C parameters (``ExpectedSourceBucketOwner`` and + ``CopySourceSSECustomer*``) under their names in those requests. + + Args: + source: The path of the object to copy, with the version ID to + copy, if any. + destination: The path of the object to write, without a version + ID. + block_size: The size in bytes of the copied ranges, between + ``MULTIPART_UPLOAD_MIN_PART_SIZE`` and + ``MULTIPART_UPLOAD_MAX_PART_SIZE`` (the default); see + :meth:`part_ranges`. + **params: The CopyObject parameters of the copy. Each request + receives those that it accepts; CreateMultipartUpload also + receives those that CopyObject does not accept either, so + that botocore rejects them as it would for CopyObject. + + Returns: + The plan. + + Raises: + ValueError: If the source or the destination has no key, the + destination has a version ID, ``block_size`` is out of the + part size limits, a directive has a value that CopyObject + does not accept, or HeadObject reports no size. + """ + if not source.key: + raise ValueError(f"The source has no key: {source.uri}.") + if not destination.key: + raise ValueError(f"The path has no key: {destination.uri}.") + if destination.version_id: + raise ValueError(f"Cannot write to a version: {destination.uri}.") + block_size = block_size if block_size else self.MULTIPART_UPLOAD_MAX_PART_SIZE + if ( + block_size < self.MULTIPART_UPLOAD_MIN_PART_SIZE + or block_size > self.MULTIPART_UPLOAD_MAX_PART_SIZE + ): + 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), " + f"inclusive: {block_size}." + ) + metadata_directive = params.get("MetadataDirective", "COPY") + tagging_directive = params.get("TaggingDirective", "COPY") + annotation_directive = params.get("AnnotationDirective", "COPY") + if metadata_directive not in ("COPY", "REPLACE"): + raise ValueError(f"Invalid MetadataDirective: {metadata_directive}.") + if tagging_directive not in ("COPY", "REPLACE"): + raise ValueError(f"Invalid TaggingDirective: {tagging_directive}.") + if annotation_directive not in ("COPY", "EXCLUDE"): + raise ValueError(f"Invalid AnnotationDirective: {annotation_directive}.") + + source_params = self._copy_source_params(params) + _logger.debug(f"Head object to copy: {source.uri}") + head = self.head_object(source, **self.operation_params("head_object", source_params)) + if head.content_length is None: + raise ValueError(f"HeadObject reported no size for {source.uri}.") + if not source.version_id and head.version_id and head.version_id != "null": + source = S3Path(source.bucket, source.key, head.version_id) + if head.content_length <= self.MULTIPART_UPLOAD_MAX_PART_SIZE: + # Copied with CopyObject instead, which applies the directives. + return S3MultipartCopyPlan( + source=source, + destination=destination, + size=head.content_length, + fits_single_request=True, + ) + + request = dict(params) + if metadata_directive == "COPY": + for name in self._COPY_METADATA_PARAMS: + request.pop(name, None) + copied = { + "CacheControl": head.cache_control, + "ContentDisposition": head.content_disposition, + "ContentEncoding": head.content_encoding, + "ContentLanguage": head.content_language, + "ContentType": head.content_type, + "Expires": head.expires, + "Metadata": head.user_metadata, + } + request.update({k: v for k, v in copied.items() if v is not None}) + if tagging_directive == "COPY": + request.pop("Tagging", None) + # Directory buckets do not support GetObjectTagging, and their + # objects have no tags. + if not self._is_directory_bucket(source.bucket): + _logger.debug(f"Get tags to copy: {source.uri}") + tagging_request: dict[str, Any] = {"Bucket": source.bucket, "Key": source.key} + if source.version_id: + tagging_request.update({"VersionId": source.version_id}) + response = self.call( + self._client.get_object_tagging, + **self.operation_params("get_object_tagging", source_params), + **tagging_request, + ) + tags = [(t["Key"], t["Value"]) for t in response["TagSet"]] + if tags: + request.update({"Tagging": urlencode(tags)}) + copy_members = self._client.meta.service_model.operation_model( + "CopyObject" + ).input_shape.members + create_params = { + **self.operation_params("create_multipart_upload", request), + # A parameter that CopyObject does not accept either is sent as + # is, so that botocore rejects it as it does for CopyObject. + **{k: v for k, v in request.items() if k not in copy_members}, + } + ranges = tuple(self.part_ranges(head.content_length, block_size)) + # The annotations are listed before the caller writes anything, so + # that a missing permission fails first. + annotations = ( + tuple( + self.list_object_annotations( + source, **self.operation_params("list_object_annotations", source_params) + ) + ) + if annotation_directive == "COPY" + and "CopySourceSSECustomerAlgorithm" not in params + and not self._is_directory_bucket(source.bucket) + else () + ) + return S3MultipartCopyPlan( + source=source, + destination=destination, + size=head.content_length, + ranges=ranges, + create_params=create_params, + part_params=self.operation_params("upload_part_copy", params), + complete_params=self.operation_params("complete_multipart_upload", params), + abort_params=self.operation_params("abort_multipart_upload", params), + annotations=annotations, + ) + + def list_object_annotations(self, path: S3Path, **params) -> list[str]: + """List the names of the annotations of an object with ListObjectAnnotations. + + Sends one request per page. + + Args: + path: The path of the object, with the version ID to list, if + any. + **params: Additional request parameters. The fields that the + other arguments set take precedence over parameters of the + same name. + + Returns: + The annotation names, across all pages. + + Raises: + ValueError: If the path has no key. + """ + if not path.key: + raise ValueError(f"The path has no key: {path.uri}.") + request: dict[str, Any] = {**params, "Bucket": path.bucket, "Key": path.key} + if path.version_id: + request.update({"VersionId": path.version_id}) + names: list[str] = [] + while True: + _logger.debug(f"List object annotations: {path.uri}") + response = self.call(self._client.list_object_annotations, **request) + names.extend(a["AnnotationName"] for a in response.get("Annotations", [])) + token = response.get("NextContinuationToken") + if not token: + return names + request.update({"ContinuationToken": token}) + + def copy_object_annotation( + self, + name: str, + source: S3Path, + destination: S3Path, + version_id: str | None, + etag: str | None, + **params, + ) -> None: + """Copy an annotation of an object onto the object that a copy wrote. + + Reads the annotation with GetObjectAnnotation and writes it with + PutObjectAnnotation. The annotation is written to the version that + the copy created, if the bucket is versioned, and only if the + destination still has the ETag of the copy, so that it is not + attached to an object written over the copy. + + Args: + name: The annotation name. + source: The path of the copied object, with the version ID that + was copied, if any. + destination: The path of the object that the copy wrote, without + a version ID. + version_id: The version ID that the copy created, if any. + etag: The ETag of the object that the copy wrote, if any. + **params: The CopyObject parameters of the copy. + GetObjectAnnotation receives those of the source, mapped as + for :meth:`plan_multipart_copy`, and PutObjectAnnotation + those that it accepts; the fields that the other arguments + set take precedence. + + Raises: + ValueError: If the source or the destination has no key. + """ + if not source.key: + raise ValueError(f"The source has no key: {source.uri}.") + if not destination.key: + raise ValueError(f"The path has no key: {destination.uri}.") + get_request: dict[str, Any] = { + "Bucket": source.bucket, + "Key": source.key, + "AnnotationName": name, + } + if source.version_id: + get_request.update({"VersionId": source.version_id}) + _logger.debug(f"Copy object annotation {name} from {source.uri} to {destination.uri}.") + response = self.call( + self._client.get_object_annotation, + **self.operation_params("get_object_annotation", self._copy_source_params(params)), + **get_request, + ) + put_request: dict[str, Any] = { + "Bucket": destination.bucket, + "Key": destination.key, + "AnnotationName": name, + "AnnotationPayload": response["AnnotationPayload"].read(), + } + if version_id: + put_request.update({"VersionId": version_id}) + if etag: + put_request.update({"ObjectIfMatch": etag}) + self.call( + self._client.put_object_annotation, + **{**self.operation_params("put_object_annotation", params), **put_request}, + ) + + @staticmethod + def _is_directory_bucket(bucket: str) -> bool: + """Return whether the bucket is a directory bucket (S3 Express One Zone). + + Directory bucket names end with ``--x-s3``. + + Args: + bucket: S3 bucket name. + + Returns: + True if the bucket is a directory bucket. + """ + return bucket.endswith("--x-s3") + + @staticmethod + def _copy_source_params(params: Mapping[str, Any]) -> dict[str, Any]: + """Map the parameters of a copy to those of the requests that read its source. + + Args: + params: The CopyObject parameters of the copy. + + Returns: + ``RequestPayer``, and the source's expected bucket owner and SSE-C + parameters under the names of the requests that read the source + (``ExpectedBucketOwner`` and ``SSECustomer*``), where given. + """ + source_params = { + "RequestPayer": params.get("RequestPayer"), + "ExpectedBucketOwner": params.get("ExpectedSourceBucketOwner"), + "SSECustomerAlgorithm": params.get("CopySourceSSECustomerAlgorithm"), + "SSECustomerKey": params.get("CopySourceSSECustomerKey"), + "SSECustomerKeyMD5": params.get("CopySourceSSECustomerKeyMD5"), + } + return {k: v for k, v in source_params.items() if v is not None} + def list_objects_page( self, bucket: str, diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index d6c0084d..9bef7b8f 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1507,10 +1507,10 @@ def test_cp_file_directory(self): # sent to CopyObject and fail with NoSuchKey. fs = self._make_fs() fs.info = mock.MagicMock(return_value=S3FileSystem._directory_object("bucket", "src")) - fs._copy_object = mock.MagicMock() + fs.core.copy_object = mock.MagicMock() fs.cp_file("s3://bucket/src", "s3://bucket/dst") - fs._copy_object.assert_not_called() + fs.core.copy_object.assert_not_called() fs._call.assert_not_called() @pytest.mark.parametrize( @@ -1829,7 +1829,7 @@ def test_cp_file_multipart_parameters(self, size): key="src", ) ) - fs._copy_object = mock.MagicMock() + fs.core.copy_object = mock.MagicMock() fs._copy_object_with_multipart_upload = mock.MagicMock() fs.cp_file( @@ -1841,30 +1841,21 @@ def test_cp_file_multipart_parameters(self, size): ) if size <= fs.core.MULTIPART_UPLOAD_MAX_PART_SIZE: - fs._copy_object.assert_called_once_with( - bucket1="bucket", - key1="src", - version_id1=None, - bucket2="bucket", - key2="dst", - RequestPayer="requester", + fs.core.copy_object.assert_called_once_with( + S3Path("bucket", "src"), S3Path("bucket", "dst"), RequestPayer="requester" ) else: fs._copy_object_with_multipart_upload.assert_called_once_with( - bucket1="bucket", - key1="src", - version_id1=None, - size1=size, - bucket2="bucket", - key2="dst", + S3Path("bucket", "src"), + S3Path("bucket", "dst"), max_workers=2, block_size=fs.core.MULTIPART_UPLOAD_MIN_PART_SIZE, RequestPayer="requester", ) def test_copy_object_with_multipart_upload_request_parameters(self): - # GH-946: the part copies receive the parameters of the copy that - # they accept, and the completion and the abort get them all. + # GH-946: the part copies, the completion and the abort receive the + # parameters of the copy that they accept. fs = self._make_fs() fs.core.create_multipart_upload = mock.MagicMock( return_value=SimpleNamespace(upload_id="uploadid") @@ -1873,7 +1864,7 @@ def test_copy_object_with_multipart_upload_request_parameters(self): side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) ) fs._finish_multipart_upload = mock.MagicMock() - fs._call.return_value = {} + fs._call.return_value = {"ContentLength": 5 * 2**30 + 2**20} # The directives make the copy use the given values without reading # the source's metadata, tags and annotations (GH-973). directives = { @@ -1884,12 +1875,7 @@ def test_copy_object_with_multipart_upload_request_parameters(self): kwargs = {"ContentType": "text/csv", "RequestPayer": "requester", **directives} fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=5 * 2**30 + 2**20, - bucket2="bucket", - key2="dst", - **kwargs, + S3Path("bucket", "src"), S3Path("bucket", "dst"), **kwargs ) fs.core.create_multipart_upload.assert_called_once_with( @@ -1901,7 +1887,9 @@ def test_copy_object_with_multipart_upload_request_parameters(self): c.kwargs["RequestPayer"] == "requester" and "ContentType" not in c.kwargs for c in fs.core.upload_part_copy.call_args_list ) - assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == kwargs + assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == { + "RequestPayer": "requester" + } @staticmethod def _stubbed_copy_fs(**kwargs): @@ -1919,11 +1907,8 @@ def _stubbed_copy_fs(**kwargs): @staticmethod def _multipart_copy(fs, bucket1="bucket", **kwargs): fs._copy_object_with_multipart_upload( - bucket1=bucket1, - key1="src", - size1=MULTIPART_COPY_SIZE, - bucket2="bucket", - key2="dst", + S3Path(bucket1, "src"), + S3Path("bucket", "dst"), block_size=MULTIPART_COPY_BLOCK_SIZE, **kwargs, ) @@ -2024,17 +2009,14 @@ def test_copy_object_with_multipart_upload_small_head_object_size(self, size): # an empty object, the reported version is copied with CopyObject. fs = self._make_fs() fs._call.return_value = {"ContentLength": size, "VersionId": "v1"} - fs._copy_object = mock.MagicMock() + fs.core.copy_object = mock.MagicMock() fs.core.create_multipart_upload = mock.MagicMock() self._multipart_copy(fs, ContentType="text/csv", RequestPayer="requester") - fs._copy_object.assert_called_once_with( - bucket1="bucket", - key1="src", - version_id1="v1", - bucket2="bucket", - key2="dst", + fs.core.copy_object.assert_called_once_with( + S3Path("bucket", "src", "v1"), + S3Path("bucket", "dst"), ContentType="text/csv", RequestPayer="requester", ) @@ -2049,7 +2031,11 @@ def test_copy_object_with_multipart_upload_replace_directives(self): with Stubber(fs._client) as stubber: # Read only for the version, which a bucket without versioning # does not report. - stubber.add_response("head_object", {"ContentType": "text/csv"}, None) + stubber.add_response( + "head_object", + {"ContentLength": MULTIPART_COPY_SIZE, "ContentType": "text/csv"}, + None, + ) stubber.add_response( "create_multipart_upload", {"UploadId": "u"}, @@ -2081,28 +2067,6 @@ def test_copy_object_with_multipart_upload_invalid_directive(self, directive): with Stubber(fs._client), pytest.raises(ValueError, match="Invalid"): self._multipart_copy(fs, **directive) - def test_copy_object_with_multipart_upload_unknown_parameter(self): - # A parameter that CopyObject does not accept is passed on to - # CreateMultipartUpload, so that botocore still rejects it. - fs = self._stubbed_copy_fs() - with Stubber(fs._client) as stubber: - stubber.add_response("head_object", {}, None) - create_kwargs, version_id, size = fs._get_multipart_copy_kwargs( - "bucket", - "src", - None, - { - "ContentTyp": "text/csv", - "MetadataDirective": "REPLACE", - "TaggingDirective": "REPLACE", - }, - ) - assert create_kwargs == {"ContentTyp": "text/csv"} - assert version_id is None - assert size is None - with pytest.raises(botocore.exceptions.ParamValidationError, match="ContentTyp"): - fs._client.create_multipart_upload(Bucket="bucket", Key="dst", **create_kwargs) - def test_copy_object_with_multipart_upload_sse_c_source(self): # GH-973: the source's SSE-C key reaches its HeadObject, and an SSE-C # object, which cannot have annotations, is not listed for them. @@ -2111,7 +2075,7 @@ def test_copy_object_with_multipart_upload_sse_c_source(self): with Stubber(fs._client) as stubber: stubber.add_response( "head_object", - {"ContentType": "text/csv"}, + {"ContentLength": MULTIPART_COPY_SIZE, "ContentType": "text/csv"}, { "Bucket": "bucket", "Key": "src", @@ -2138,7 +2102,9 @@ def test_copy_object_with_multipart_upload_directory_bucket_source(self): bucket = "bucket--usw2-az1--x-s3" with Stubber(fs._client) as stubber: stubber.add_response( - "head_object", {"ContentType": "text/csv"}, {"Bucket": bucket, "Key": "src"} + "head_object", + {"ContentLength": MULTIPART_COPY_SIZE, "ContentType": "text/csv"}, + {"Bucket": bucket, "Key": "src"}, ) stubber.add_response( "create_multipart_upload", @@ -3408,14 +3374,12 @@ def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): ) fs.core.upload_part_copy = mock.MagicMock() fs._finish_multipart_upload = mock.MagicMock() - fs._call.return_value = {} + # The HeadObject of the source. + fs._call.return_value = {"ContentLength": 5 * 2**30 + 2**20} fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=5 * 2**30 + 2**20, - bucket2="bucket", - key2="dst", + S3Path("bucket", "src"), + S3Path("bucket", "dst"), max_workers=max_workers, # Copy without reading the metadata, tags and annotations of the # source (GH-973). @@ -3449,12 +3413,7 @@ def test_copy_object_with_multipart_upload_invalid_block_size(self, block_size): 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, + S3Path("bucket", "src"), S3Path("bucket", "dst"), block_size=block_size ) fs._call.assert_not_called() diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 107d53ac..720ee136 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -35,7 +35,6 @@ from tests.pyathena.util import ( MULTIPART_COPY_BLOCK_SIZE, MULTIPART_COPY_KWARGS, - MULTIPART_COPY_SIZE, stub_multipart_copy, ) @@ -166,14 +165,14 @@ async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) ) sync_fs.core.complete_multipart_upload = mock.MagicMock() - sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={}) + # The HeadObject of the source. + sync_fs._call = sync_fs._core.call = mock.MagicMock( + return_value={"ContentLength": 5 * 2**30 + 2**20} + ) await fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=5 * 2**30 + 2**20, - bucket2="bucket", - key2="dst", + S3Path("bucket", "src"), + S3Path("bucket", "dst"), # Copy without reading the metadata, tags and annotations of the # source (GH-973). MetadataDirective="REPLACE", @@ -204,11 +203,8 @@ async def _multipart_copy(fs=None, **kwargs): stub_multipart_copy(stubber, **kwargs) try: await fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=MULTIPART_COPY_SIZE, - bucket2="bucket", - key2="dst", + S3Path("bucket", "src"), + S3Path("bucket", "dst"), block_size=MULTIPART_COPY_BLOCK_SIZE, **MULTIPART_COPY_KWARGS, ) @@ -225,25 +221,19 @@ async def test_copy_object_with_multipart_upload_small_head_object_size(self, si sync_fs._call = sync_fs._core.call = mock.MagicMock( return_value={"ContentLength": size, "VersionId": "v1"} ) - sync_fs._copy_object = mock.MagicMock() + sync_fs.core.copy_object = mock.MagicMock() sync_fs.core.create_multipart_upload = mock.MagicMock() await fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=MULTIPART_COPY_SIZE, - bucket2="bucket", - key2="dst", + S3Path("bucket", "src"), + S3Path("bucket", "dst"), ContentType="text/csv", RequestPayer="requester", ) - sync_fs._copy_object.assert_called_once_with( - bucket1="bucket", - key1="src", - version_id1="v1", - bucket2="bucket", - key2="dst", + sync_fs.core.copy_object.assert_called_once_with( + S3Path("bucket", "src", "v1"), + S3Path("bucket", "dst"), ContentType="text/csv", RequestPayer="requester", ) @@ -284,7 +274,7 @@ async def test_copy_object_with_multipart_upload_failed_annotation(self): sync_fs = fs._sync_fs with ( mock.patch.object( - sync_fs, "_copy_object_annotation", wraps=sync_fs._copy_object_annotation + sync_fs.core, "copy_object_annotation", wraps=sync_fs.core.copy_object_annotation ) as copy_annotation, pytest.raises(PermissionError), ): @@ -335,20 +325,17 @@ def upload_part_copy(**kw): sync_fs.core.upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) sync_fs.core.complete_multipart_upload = mock.MagicMock() - # The HeadObject of the source, for its version. - sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={}) + # The HeadObject of the source, with 3 parts of the default block size. + size = 3 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE + sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={"ContentLength": size}) sync_fs._abort_multipart_upload = mock.MagicMock( side_effect=lambda *args: events.append("abort") ) with pytest.raises(OSError, match="part failed"): await fs._copy_object_with_multipart_upload( - bucket1="bucket", - key1="src", - size1=3 * S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, - bucket2="bucket", - key2="dst", - block_size=S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, + S3Path("bucket", "src"), + S3Path("bucket", "dst"), MetadataDirective="REPLACE", TaggingDirective="REPLACE", AnnotationDirective="EXCLUDE", @@ -394,18 +381,15 @@ def abort_multipart_upload(*args): sync_fs.core.upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) sync_fs.core.complete_multipart_upload = mock.MagicMock() - # The HeadObject of the source, for its version. - sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={}) + # The HeadObject of the source, with 3 parts of the default block size. + size = 3 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE + sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={"ContentLength": size}) 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 * S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, - bucket2="bucket", - key2="dst", - block_size=S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, + S3Path("bucket", "src"), + S3Path("bucket", "dst"), MetadataDirective="REPLACE", TaggingDirective="REPLACE", AnnotationDirective="EXCLUDE", @@ -469,20 +453,17 @@ def complete_multipart_upload(*args, **kw): sync_fs.core.complete_multipart_upload = mock.MagicMock( side_effect=complete_multipart_upload ) - # The HeadObject of the source, for its version. - sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={}) + # The HeadObject of the source, with 2 parts of the default block size. + size = 2 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE + sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={"ContentLength": size}) 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 * S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, - bucket2="bucket", - key2="dst", - block_size=S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE, + S3Path("bucket", "src"), + S3Path("bucket", "dst"), MetadataDirective="REPLACE", TaggingDirective="REPLACE", AnnotationDirective="EXCLUDE", @@ -524,12 +505,7 @@ async def test_copy_object_with_multipart_upload_invalid_block_size(self, block_ 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, + S3Path("bucket", "src"), S3Path("bucket", "dst"), block_size=block_size ) fs._sync_fs._call.assert_not_called() @@ -1009,7 +985,7 @@ async def test_cp_file_multipart_parameters(self, size): ) ) sync_fs = fs._sync_fs - sync_fs._copy_object = mock.MagicMock() + sync_fs.core.copy_object = mock.MagicMock() sync_fs.core.create_multipart_upload = mock.MagicMock( return_value=SimpleNamespace(upload_id="uploadid") ) @@ -1025,8 +1001,8 @@ def upload_part_copy(**kw): sync_fs.core.upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) sync_fs.core.complete_multipart_upload = mock.MagicMock() - # The HeadObject of the source, for its version. - sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={}) + # The HeadObject of the source. + sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={"ContentLength": size}) directives = { "MetadataDirective": "REPLACE", "TaggingDirective": "REPLACE", @@ -1046,12 +1022,9 @@ def upload_part_copy(**kw): ) if size <= S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE: - sync_fs._copy_object.assert_called_once_with( - bucket1="bucket", - key1="src", - version_id1=None, - bucket2="bucket", - key2="dst", + sync_fs.core.copy_object.assert_called_once_with( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), RequestPayer="requester", ContentType="text/csv", **directives, diff --git a/tests/pyathena/filesystem/test_s3_core.py b/tests/pyathena/filesystem/test_s3_core.py index b635783a..a4c43a74 100644 --- a/tests/pyathena/filesystem/test_s3_core.py +++ b/tests/pyathena/filesystem/test_s3_core.py @@ -5,12 +5,14 @@ # # SPDX-License-Identifier: MIT +import io from datetime import UTC, datetime from itertools import pairwise import boto3 import botocore.exceptions import pytest +from botocore.response import StreamingBody from botocore.stub import Stubber from pyathena.filesystem.s3_core import ( @@ -23,6 +25,7 @@ S3ListBucketsPage, S3ListObjectsPage, S3ListObjectVersionsPage, + S3MultipartCopyPlan, S3ObjectSummary, ) from pyathena.filesystem.s3_object import S3MultipartUploadPart @@ -611,6 +614,295 @@ def test_part_ranges_max_parts(self, size, num_ranges): for start, end in ranges ) + def test_copy_object(self): + core, stubber = _make_core() + stubber.add_response( + "copy_object", + {}, + { + "CopySource": {"Bucket": "src-bucket", "Key": "src", "VersionId": "v1"}, + "Bucket": "bucket", + "Key": "dst", + "MetadataDirective": "REPLACE", + }, + ) + stubber.add_response( + "copy_object", + {}, + {"CopySource": {"Bucket": "bucket", "Key": "src"}, "Bucket": "bucket", "Key": "dst"}, + ) + with stubber: + core.copy_object( + S3Path("src-bucket", "src", "v1"), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + ) + core.copy_object(S3Path("bucket", "src"), S3Path("bucket", "dst")) + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + ("method", "args", "match"), + [ + ("copy_object", (S3Path("bucket"), S3Path("bucket", "dst")), "has no key"), + ("copy_object", (S3Path("bucket", "src"), S3Path("bucket")), "has no key"), + ( + "copy_object", + (S3Path("bucket", "src"), S3Path("bucket", "dst", "v1")), + "Cannot write to a version", + ), + ("plan_multipart_copy", (S3Path("bucket"), S3Path("bucket", "dst")), "has no key"), + ("plan_multipart_copy", (S3Path("bucket", "src"), S3Path("bucket")), "has no key"), + ( + "plan_multipart_copy", + (S3Path("bucket", "src"), S3Path("bucket", "dst", "v1")), + "Cannot write to a version", + ), + ("list_object_annotations", (S3Path("bucket"),), "has no key"), + ( + "copy_object_annotation", + ("a", S3Path("bucket"), S3Path("bucket", "dst"), None, None), + "has no key", + ), + ( + "copy_object_annotation", + ("a", S3Path("bucket", "src"), S3Path("bucket"), None, None), + "has no key", + ), + ], + ) + def test_copy_rejects_paths(self, method, args, match): + core, stubber = _make_core() + with stubber, pytest.raises(ValueError, match=match): + getattr(core, method)(*args) + + @staticmethod + def _stub_head(stubber, response, version_id=None, **params): + expected = {"Bucket": "bucket", "Key": "src", **params} + if version_id: + expected.update({"VersionId": version_id}) + stubber.add_response("head_object", response, expected) + + def test_plan_multipart_copy(self): + # The source is read as CopyObject would read it: the version that + # HeadObject reports is pinned, its metadata and tags replace those of + # the parameters, and its annotations are listed on every page. The + # source's expected owner reaches the reads under their own names. + core, stubber = _make_core() + size = core.MULTIPART_UPLOAD_MAX_PART_SIZE + core.MULTIPART_UPLOAD_MIN_PART_SIZE + source_params = {"RequestPayer": "requester", "ExpectedBucketOwner": "222222222222"} + self._stub_head( + stubber, + { + "ContentLength": size, + "ContentType": "text/csv", + "Metadata": {"owner": "etl"}, + "VersionId": "v-src", + }, + **source_params, + ) + source = {"Bucket": "bucket", "Key": "src", "VersionId": "v-src", **source_params} + stubber.add_response("get_object_tagging", {"TagSet": [{"Key": "t", "Value": "1"}]}, source) + stubber.add_response( + "list_object_annotations", + { + "Annotations": [{"AnnotationName": "a1", "LastModified": MODIFIED, "Size": 1}], + "NextContinuationToken": "next", + }, + source, + ) + stubber.add_response( + "list_object_annotations", + {"Annotations": [{"AnnotationName": "a2", "LastModified": MODIFIED, "Size": 1}]}, + {**source, "ContinuationToken": "next"}, + ) + params = { + "ContentType": "text/plain", + "Tagging": "ignored=1", + "StorageClass": "STANDARD_IA", + "RequestPayer": "requester", + "ExpectedBucketOwner": "111111111111", + "ExpectedSourceBucketOwner": "222222222222", + "CopySourceIfMatch": '"src"', + } + with stubber: + plan = core.plan_multipart_copy( + S3Path("bucket", "src"), S3Path("bucket", "dst"), **params + ) + stubber.assert_no_pending_responses() + + destination_params = {"RequestPayer": "requester", "ExpectedBucketOwner": "111111111111"} + assert plan == S3MultipartCopyPlan( + source=S3Path("bucket", "src", "v-src"), + destination=S3Path("bucket", "dst"), + size=size, + ranges=( + (0, core.MULTIPART_UPLOAD_MAX_PART_SIZE), + (core.MULTIPART_UPLOAD_MAX_PART_SIZE, size), + ), + create_params={ + **destination_params, + "ContentType": "text/csv", + "Metadata": {"owner": "etl"}, + "Tagging": "t=1", + "StorageClass": "STANDARD_IA", + }, + part_params={ + **destination_params, + "ExpectedSourceBucketOwner": "222222222222", + "CopySourceIfMatch": '"src"', + }, + complete_params=destination_params, + abort_params=destination_params, + annotations=("a1", "a2"), + ) + + @pytest.mark.parametrize( + ("version_id", "head_version_id", "expected"), + [ + # The version that HeadObject reports is pinned, + (None, "v1", "v1"), + # except a "null" version, which a write can replace, + (None, "null", None), + # and a version given with the path is kept. + ("null", "null", "null"), + ("v0", "v0", "v0"), + ], + ) + @pytest.mark.parametrize("size", [0, S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE + 1]) + def test_plan_multipart_copy_source_version(self, version_id, head_version_id, expected, size): + core, stubber = _make_core() + self._stub_head(stubber, {"ContentLength": size, "VersionId": head_version_id}, version_id) + with stubber: + plan = core.plan_multipart_copy( + S3Path("bucket", "src", version_id), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + stubber.assert_no_pending_responses() + assert plan.source == S3Path("bucket", "src", expected) + assert plan.fits_single_request is (size == 0) + + @pytest.mark.parametrize("size", [0, S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE]) + def test_plan_multipart_copy_fits_single_request(self, size): + # GH-973: a source that fits in a single CopyObject request, such as + # one whose cached size was stale, is not read any further. + core, stubber = _make_core() + self._stub_head(stubber, {"ContentLength": size}, RequestPayer="requester") + with stubber: + plan = core.plan_multipart_copy( + S3Path("bucket", "src"), S3Path("bucket", "dst"), RequestPayer="requester" + ) + stubber.assert_no_pending_responses() + assert plan == S3MultipartCopyPlan( + source=S3Path("bucket", "src"), + destination=S3Path("bucket", "dst"), + size=size, + fits_single_request=True, + ) + + def test_plan_multipart_copy_without_size(self): + core, stubber = _make_core() + self._stub_head(stubber, {}) + with stubber, pytest.raises(ValueError, match="no size"): + core.plan_multipart_copy(S3Path("bucket", "src"), S3Path("bucket", "dst")) + + def test_plan_multipart_copy_unknown_parameter(self): + # A parameter that CopyObject does not accept is passed on to + # CreateMultipartUpload, so that botocore still rejects it. + core, stubber = _make_core() + self._stub_head(stubber, {"ContentLength": core.MULTIPART_UPLOAD_MAX_PART_SIZE + 1}) + with stubber: + plan = core.plan_multipart_copy( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + ContentTyp="text/csv", + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + assert plan.create_params == {"ContentTyp": "text/csv"} + assert plan.part_params == {} + with pytest.raises(botocore.exceptions.ParamValidationError, match="ContentTyp"): + core.client.create_multipart_upload(Bucket="bucket", Key="dst", **plan.create_params) + + def test_list_object_annotations(self): + core, stubber = _make_core() + stubber.add_response( + "list_object_annotations", + { + "Annotations": [{"AnnotationName": "a1", "LastModified": MODIFIED, "Size": 1}], + "NextContinuationToken": "next", + }, + {"Bucket": "bucket", "Key": "key", "VersionId": "v1", "RequestPayer": "requester"}, + ) + stubber.add_response( + "list_object_annotations", + {}, + { + "Bucket": "bucket", + "Key": "key", + "VersionId": "v1", + "RequestPayer": "requester", + "ContinuationToken": "next", + }, + ) + with stubber: + # The key of the path takes precedence over a parameter. + names = core.list_object_annotations( + S3Path("bucket", "key", "v1"), RequestPayer="requester", Key="other" + ) + stubber.assert_no_pending_responses() + assert names == ["a1"] + + def test_copy_object_annotation(self): + # The source is read with the source's parameters of the copy, and + # the annotation is written to the version and the ETag that the copy + # wrote, with the parameters that PutObjectAnnotation accepts. + core, stubber = _make_core() + stubber.add_response( + "get_object_annotation", + {"AnnotationPayload": StreamingBody(io.BytesIO(b"payload"), 7)}, + { + "Bucket": "bucket", + "Key": "src", + "VersionId": "v-src", + "AnnotationName": "a1", + "RequestPayer": "requester", + "ExpectedBucketOwner": "222222222222", + }, + ) + stubber.add_response( + "put_object_annotation", + {}, + { + "Bucket": "bucket", + "Key": "dst", + "AnnotationName": "a1", + "AnnotationPayload": b"payload", + "VersionId": "v-dst", + "ObjectIfMatch": '"dst"', + "RequestPayer": "requester", + "ExpectedBucketOwner": "111111111111", + }, + ) + with stubber: + core.copy_object_annotation( + "a1", + S3Path("bucket", "src", "v-src"), + S3Path("bucket", "dst"), + "v-dst", + '"dst"', + RequestPayer="requester", + ExpectedBucketOwner="111111111111", + ExpectedSourceBucketOwner="222222222222", + ContentType="text/csv", + # A field of the request takes precedence. + ObjectIfMatch='"other"', + ) + stubber.assert_no_pending_responses() + class TestS3DeleteBatch: def test_from_paths(self): From 3a184ae2f5d11a8f81ece7d3f6a4ae285d00fb5b Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:20:00 +0900 Subject: [PATCH 2/9] State which core operations send more than one request, and the missing-size error Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 4 ++-- pyathena/filesystem/s3.py | 5 +++-- pyathena/filesystem/s3_async.py | 5 +++-- pyathena/filesystem/s3_core.py | 2 +- 4 files changed, 9 insertions(+), 7 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index 39ad8f1a..6353f282 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -282,8 +282,8 @@ directories below the bucket level) and is always a no-op. `S3FileSystem.core` is an `S3Core`, the typed operations that the filesystem sends its listing, lookup, delete, multipart upload and copy requests with. It can also be built on a boto3 S3 client. Each operation sends one request (one per page for the -iterators and `list_object_annotations()`), except the copy operations described below, -with the retry policy, raises `FileNotFoundError` for a missing bucket or +iterators and `list_object_annotations()`), except `plan_multipart_copy()` and +`copy_object_annotation()`, described below, with the retry policy, raises `FileNotFoundError` for a missing bucket or multipart upload, or for a missing object or version that it reads, and caches nothing. Requests sent through `fs.core` do not invalidate the filesystem's cache: call diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index cb601d11..4e5d5c98 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1757,8 +1757,9 @@ def _copy_object_with_multipart_upload( receives those that it accepts. Raises: - ValueError: If ``block_size`` is out of the part size limits or a - directive has an invalid value. + ValueError: If ``block_size`` is out of the part size limits, a + directive has an invalid value, or HeadObject reports no + size. """ max_workers = max_workers if max_workers else self.max_workers plan = self.core.plan_multipart_copy(source, destination, block_size, **kwargs) diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index fc4a8fb8..7fc5da59 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -573,8 +573,9 @@ async def _copy_object_with_multipart_upload( receives those that it accepts. Raises: - ValueError: If ``block_size`` is out of the part size limits or a - directive has an invalid value. + ValueError: If ``block_size`` is out of the part size limits, a + directive has an invalid value, or HeadObject reports no + size. """ max_workers = max_workers if max_workers else self._sync_fs.max_workers plan = await asyncio.to_thread( diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index f167e8e9..486d41a1 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -446,7 +446,7 @@ class S3MultipartCopyPlan: class S3Core: - """Typed S3 operations, one request each, on a boto3 S3 client. + """Typed S3 operations on a boto3 S3 client. Each operation sends one request, or one per page for the iterators, except those that say otherwise, such as :meth:`plan_multipart_copy` From 567b21616b93f323177188ec49f2d12f9433247a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:22:46 +0900 Subject: [PATCH 3/9] Correct which plan fields are empty and split the docs exception into its own sentence Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 8 ++++---- pyathena/filesystem/s3_core.py | 4 ++-- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index 6353f282..ce776347 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -282,11 +282,11 @@ directories below the bucket level) and is always a no-op. `S3FileSystem.core` is an `S3Core`, the typed operations that the filesystem sends its listing, lookup, delete, multipart upload and copy requests with. It can also be built on a boto3 S3 client. Each operation sends one request (one per page for the -iterators and `list_object_annotations()`), except `plan_multipart_copy()` and -`copy_object_annotation()`, described below, with the retry policy, raises `FileNotFoundError` for a missing bucket or +iterators and `list_object_annotations()`); `plan_multipart_copy()` and +`copy_object_annotation()`, described below, send several. The requests are sent with +the retry policy. An operation raises `FileNotFoundError` for a missing bucket or multipart upload, or for a missing object or version that it reads, and caches -nothing. Requests sent -through `fs.core` do not invalidate the filesystem's cache: call +nothing. Requests sent through `fs.core` do not invalidate the filesystem's cache: call `fs.invalidate_cache()` after a change, or make it through the filesystem. ```python diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index 486d41a1..5d7e5f20 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -412,8 +412,8 @@ class S3MultipartCopyPlan: AbortMultipartUpload with ``abort_params`` after a failure, and finally the copy of each annotation with :meth:`S3Core.copy_object_annotation`. If ``fits_single_request`` is true, the source is copied with - :meth:`S3Core.copy_object` instead, and the fields after ``size`` are - empty. + :meth:`S3Core.copy_object` instead, and ``ranges``, the parameters and + ``annotations`` are empty. Attributes: source: The object to copy, with the version that HeadObject From 3e6d7c534fcddeed02878ca78ab6eefb65e4885f Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:27:12 +0900 Subject: [PATCH 4/9] Reuse S3Path.with_version_id, operation_params and the shared copy fixtures Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 4 ++-- pyathena/filesystem/s3_core.py | 8 +++---- tests/pyathena/filesystem/test_s3.py | 6 +++--- tests/pyathena/filesystem/test_s3_core.py | 26 ++++++++++------------- 4 files changed, 19 insertions(+), 25 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 4e5d5c98..63403351 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1789,8 +1789,8 @@ def _copy_object_with_multipart_upload( key=cast(str, plan.destination.key), upload_id=upload_id, futures=futures, - # The completion and the abort each receive those that they - # accept, as filtered for the plan. + # Filtered again for the completion and the abort, which + # leaves the plan's parameters of each unchanged. request_kwargs={**plan.complete_params, **plan.abort_params}, ) for name in plan.annotations: diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index 5d7e5f20..53369ba5 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -967,7 +967,7 @@ def plan_multipart_copy( if head.content_length is None: raise ValueError(f"HeadObject reported no size for {source.uri}.") if not source.version_id and head.version_id and head.version_id != "null": - source = S3Path(source.bucket, source.key, head.version_id) + source = source.with_version_id(head.version_id) if head.content_length <= self.MULTIPART_UPLOAD_MAX_PART_SIZE: # Copied with CopyObject instead, which applies the directives. return S3MultipartCopyPlan( @@ -1008,14 +1008,12 @@ def plan_multipart_copy( tags = [(t["Key"], t["Value"]) for t in response["TagSet"]] if tags: request.update({"Tagging": urlencode(tags)}) - copy_members = self._client.meta.service_model.operation_model( - "CopyObject" - ).input_shape.members + copy_params = self.operation_params("copy_object", request) create_params = { **self.operation_params("create_multipart_upload", request), # A parameter that CopyObject does not accept either is sent as # is, so that botocore rejects it as it does for CopyObject. - **{k: v for k, v in request.items() if k not in copy_members}, + **{k: v for k, v in request.items() if k not in copy_params}, } ranges = tuple(self.part_ranges(head.content_length, block_size)) # The annotations are listed before the caller writes anything, so diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 9bef7b8f..fb8449a9 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1905,9 +1905,9 @@ def _stubbed_copy_fs(**kwargs): ) @staticmethod - def _multipart_copy(fs, bucket1="bucket", **kwargs): + def _multipart_copy(fs, source_bucket="bucket", **kwargs): fs._copy_object_with_multipart_upload( - S3Path(bucket1, "src"), + S3Path(source_bucket, "src"), S3Path("bucket", "dst"), block_size=MULTIPART_COPY_BLOCK_SIZE, **kwargs, @@ -2114,7 +2114,7 @@ def test_copy_object_with_multipart_upload_directory_bucket_source(self): for _ in (1, 2): stubber.add_response("upload_part_copy", {"CopyPartResult": {"ETag": '"p"'}}, None) stubber.add_response("complete_multipart_upload", {"ETag": '"dst"'}, None) - self._multipart_copy(fs, bucket1=bucket) + self._multipart_copy(fs, source_bucket=bucket) stubber.assert_no_pending_responses() def test_pipe_file_invalid_path_raises(self): diff --git a/tests/pyathena/filesystem/test_s3_core.py b/tests/pyathena/filesystem/test_s3_core.py index a4c43a74..bf6425a0 100644 --- a/tests/pyathena/filesystem/test_s3_core.py +++ b/tests/pyathena/filesystem/test_s3_core.py @@ -31,6 +31,11 @@ from pyathena.filesystem.s3_object import S3MultipartUploadPart from pyathena.filesystem.s3_path import S3Path from pyathena.util import RetryConfig +from tests.pyathena.util import ( + MULTIPART_COPY_BLOCK_SIZE, + MULTIPART_COPY_KWARGS, + MULTIPART_COPY_SIZE, +) MODIFIED = datetime(2026, 10, 4, tzinfo=UTC) @@ -688,7 +693,7 @@ def test_plan_multipart_copy(self): # the parameters, and its annotations are listed on every page. The # source's expected owner reaches the reads under their own names. core, stubber = _make_core() - size = core.MULTIPART_UPLOAD_MAX_PART_SIZE + core.MULTIPART_UPLOAD_MIN_PART_SIZE + size = MULTIPART_COPY_SIZE source_params = {"RequestPayer": "requester", "ExpectedBucketOwner": "222222222222"} self._stub_head( stubber, @@ -715,18 +720,12 @@ def test_plan_multipart_copy(self): {"Annotations": [{"AnnotationName": "a2", "LastModified": MODIFIED, "Size": 1}]}, {**source, "ContinuationToken": "next"}, ) - params = { - "ContentType": "text/plain", - "Tagging": "ignored=1", - "StorageClass": "STANDARD_IA", - "RequestPayer": "requester", - "ExpectedBucketOwner": "111111111111", - "ExpectedSourceBucketOwner": "222222222222", - "CopySourceIfMatch": '"src"', - } with stubber: plan = core.plan_multipart_copy( - S3Path("bucket", "src"), S3Path("bucket", "dst"), **params + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MULTIPART_COPY_BLOCK_SIZE, + **MULTIPART_COPY_KWARGS, ) stubber.assert_no_pending_responses() @@ -735,10 +734,7 @@ def test_plan_multipart_copy(self): source=S3Path("bucket", "src", "v-src"), destination=S3Path("bucket", "dst"), size=size, - ranges=( - (0, core.MULTIPART_UPLOAD_MAX_PART_SIZE), - (core.MULTIPART_UPLOAD_MAX_PART_SIZE, size), - ), + ranges=((0, MULTIPART_COPY_BLOCK_SIZE), (MULTIPART_COPY_BLOCK_SIZE, size)), create_params={ **destination_params, "ContentType": "text/csv", From 2f7a891c77c088c0afb529c5cf5aca4aeff67e6a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:28:55 +0900 Subject: [PATCH 5/9] Say in the docs that a source whose HeadObject size fits one CopyObject is not read further Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index ce776347..1b9b607a 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -339,7 +339,8 @@ ranges of the parts, the parameters of each multipart upload request, and the annotations to copy, so that the multipart upload writes the metadata, tags and annotations that CopyObject would. It sends HeadObject, then GetObjectTagging and ListObjectAnnotations unless the directives or the source exclude them, and writes -nothing. `copy_object_annotation()` copies one annotation onto the destination after +nothing. If HeadObject reports a size that fits in one CopyObject request, nothing else +is read, and the plan's `fits_single_request` says to copy with `copy_object()` instead. `copy_object_annotation()` copies one annotation onto the destination after the upload completes, with GetObjectAnnotation and PutObjectAnnotation. The filesystems' `cp_file()` and `copy()` run these plans. From 1d84d89532c3ebe11efb1ea83ece665b547e8050 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:50:42 +0900 Subject: [PATCH 6/9] Wrap the copy paragraph of the docs and name mv() among its callers Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index 1b9b607a..a67bc306 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -340,9 +340,10 @@ annotations to copy, so that the multipart upload writes the metadata, tags and annotations that CopyObject would. It sends HeadObject, then GetObjectTagging and ListObjectAnnotations unless the directives or the source exclude them, and writes nothing. If HeadObject reports a size that fits in one CopyObject request, nothing else -is read, and the plan's `fits_single_request` says to copy with `copy_object()` instead. `copy_object_annotation()` copies one annotation onto the destination after +is read, and the plan's `fits_single_request` says to copy with `copy_object()` +instead. `copy_object_annotation()` copies one annotation onto the destination after the upload completes, with GetObjectAnnotation and PutObjectAnnotation. The -filesystems' `cp_file()` and `copy()` run these plans. +filesystems' `cp_file()`, `copy()` and `mv()` run these plans. ## Async filesystem From 748e48be6eac69b8a05bfd5a92d49714ff85472d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 16:23:33 +0900 Subject: [PATCH 7/9] Abort a multipart copy's upload when an interrupt arrives during its creation The sync copy now creates the upload on its executor and, on an interrupt while waiting, waits for the creation and aborts the upload that it created. The aio copy runs the creation as its own task behind asyncio.shield(), and the existing cleanup waits for it before deciding whether there is an upload to abort. Closes #1076. Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 4 +- pyathena/filesystem/s3.py | 31 +++++++++--- pyathena/filesystem/s3_async.py | 33 ++++++++---- tests/pyathena/filesystem/test_s3.py | 46 +++++++++++++++++ tests/pyathena/filesystem/test_s3_async.py | 58 ++++++++++++++++++++++ 5 files changed, 155 insertions(+), 17 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index a67bc306..35c48ad5 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -123,7 +123,9 @@ from that version, which needs `s3:GetObjectVersion` and, to copy the tags, `s3:GetObjectVersionTagging` on the source. A `null` version is not pinned. The annotations are listed before anything is written and copied after the upload completes, so the destination exists without them until the last one is written. If an annotation fails to copy, the error is raised and the destination is -kept. A failed part copy aborts the multipart upload. +kept. A failed part copy aborts the multipart upload. So does an interrupt, or the +cancellation of an `AioS3FileSystem` copy, unless the upload has already completed; the +requests in flight, including a CreateMultipartUpload request, finish first. Paths are normalized as in fsspec, which drops a trailing slash, so `info`, `isfile`, and `open` treat `s3://YOUR_S3_BUCKET/dir/` as `s3://YOUR_S3_BUCKET/dir`: the object diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 63403351..9f273840 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1743,10 +1743,12 @@ def _copy_object_with_multipart_upload( source and lists its annotations before anything is written. The parts are copied in parallel with UploadPartCopy, and the annotations are copied onto the destination after the upload - completes. A failed part or completion aborts the upload; a failed - annotation copy is raised and leaves the destination in place. If - HeadObject reports a size that fits in a single CopyObject request, - the reported version is copied with CopyObject instead. + completes. A failed part or completion aborts the upload, and so does + an interrupt, including one while the upload is being created, once + the running requests have finished; a failed annotation copy is + raised and leaves the destination in place. If HeadObject reports a + size that fits in a single CopyObject request, the reported version + is copied with CopyObject instead. Args: source: Source S3 path, with the version ID to copy, if any. @@ -1769,9 +1771,26 @@ def _copy_object_with_multipart_upload( # single CopyObject request. self.core.copy_object(plan.source, plan.destination, **kwargs) return - multipart_upload = self.core.create_multipart_upload(plan.destination, **plan.create_params) - upload_id = cast(str, multipart_upload.upload_id) with self._create_executor(max_workers=max_workers) as executor: + # Created on the executor, so that an interrupt while it is being + # created lets the request finish and the upload be aborted. + creation = executor.submit( + self.core.create_multipart_upload, plan.destination, **plan.create_params + ) + try: + multipart_upload = creation.result() + except BaseException: + if not creation.cancel(): + wait([creation]) + if creation.exception() is None: + self._abort_multipart_upload( + plan.destination.bucket, + cast(str, plan.destination.key), + cast(str, creation.result().upload_id), + plan.abort_params, + ) + raise + upload_id = cast(str, multipart_upload.upload_id) futures = [ executor.submit( self.core.upload_part_copy, diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 7fc5da59..6bec9aa2 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -558,11 +558,11 @@ async def _copy_object_with_multipart_upload( See :meth:`S3FileSystem._copy_object_with_multipart_upload`. The part 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. + ``asyncio.to_thread``. On a cancellation, the creation of the upload, + the running part copies and the completion are waited for, the + upload is aborted if it was created and has not completed, and the + cancellation is re-raised. A repeated cancellation returns without + stopping this cleanup. Args: source: Source S3 path, with the version ID to copy, if any. @@ -585,10 +585,14 @@ async def _copy_object_with_multipart_upload( # See S3FileSystem._copy_object_with_multipart_upload. await asyncio.to_thread(self.core.copy_object, plan.source, plan.destination, **kwargs) return - multipart_upload = await asyncio.to_thread( - self.core.create_multipart_upload, plan.destination, **plan.create_params + # A task, so that _abort() can wait for an upload that is created + # after a cancellation. + creation = asyncio.ensure_future( + asyncio.to_thread( + self.core.create_multipart_upload, plan.destination, **plan.create_params + ) ) - upload_id = cast(str, multipart_upload.upload_id) + upload_id: str semaphore = asyncio.Semaphore(max_workers) failed = False @@ -614,10 +618,14 @@ async def _upload_part(i: int, range_: tuple[int, int]) -> S3MultipartUploadPart failed = True raise - tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(plan.ranges)] + tasks: list[asyncio.Future[S3MultipartUploadPart | None]] = [] completion: asyncio.Task[S3CompleteMultipartUpload] | None = None async def _abort() -> None: + await asyncio.wait([creation]) + if creation.cancelled() or creation.exception() is not None: + # No upload was created. + return # 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) @@ -631,11 +639,16 @@ async def _abort() -> None: self._sync_fs._abort_multipart_upload, plan.destination.bucket, cast(str, plan.destination.key), - upload_id, + cast(str, creation.result().upload_id), plan.abort_params, ) try: + # shield keeps a cancellation from cancelling the creation, whose + # thread would keep running, so that _abort() can wait for it. + multipart_upload = await asyncio.shield(creation) + upload_id = cast(str, multipart_upload.upload_id) + tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(plan.ranges)] # 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(). diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index fb8449a9..001edc31 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -8,6 +8,7 @@ import lzma import os import re +import signal import sys import tempfile import threading @@ -3397,6 +3398,51 @@ 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.skipif( + threading.current_thread() is not threading.main_thread(), + reason="SIGINT interrupts the main thread.", + ) + def test_copy_object_with_multipart_upload_interrupted_creation(self): + # An interrupt during CreateMultipartUpload waits for it, aborts the + # upload that it created before any part is copied, and is re-raised. + # The created upload used to be left incomplete. + fs = self._make_fs() + started = threading.Event() + + def create_multipart_upload(*args, **kw): + started.set() + # Still running when the interrupt arrives. + time.sleep(0.5) + return SimpleNamespace(upload_id="uploadid") + + fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) + fs.core.upload_part_copy = mock.MagicMock() + # The HeadObject of the source. + fs._call.return_value = {"ContentLength": 2 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE} + fs._abort_multipart_upload = mock.MagicMock() + + def interrupt(): + if started.wait(5): + # Lets the main thread return from submit() and wait for the + # creation, where an interrupt during the request arrives. + time.sleep(0.1) + signal.pthread_kill(threading.main_thread().ident, signal.SIGINT) + + thread = threading.Thread(target=interrupt, daemon=True) + thread.start() + with pytest.raises(KeyboardInterrupt): + fs._copy_object_with_multipart_upload( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + thread.join(5) + + fs._abort_multipart_upload.assert_called_once_with("bucket", "dst", "uploadid", {}) + fs.core.upload_part_copy.assert_not_called() + @pytest.mark.parametrize( "block_size", [ diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 720ee136..9663602c 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -487,6 +487,64 @@ def complete_multipart_upload(*args, **kw): assert events == (["complete", "abort"] if completion_fails else ["complete"]) + @pytest.mark.parametrize("creation_fails", [False, True]) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_cancelled_creation(self, creation_fails): + # A cancellation during CreateMultipartUpload waits for it, aborts + # the upload that it created before any part is copied, and is + # re-raised. The created upload used to be left incomplete. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + sync_fs = fs._sync_fs + events = [] + started = threading.Event() + release = threading.Event() + + def create_multipart_upload(*args, **kw): + started.set() + # The finally blocks of the test always release it. + release.wait() + events.append("create") + if creation_fails: + raise OSError("creation failed") + return SimpleNamespace(upload_id="uploadid") + + sync_fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) + sync_fs.core.upload_part_copy = mock.MagicMock() + # The HeadObject of the source, with 2 parts of the default block size. + size = 2 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE + sync_fs._call = sync_fs._core.call = mock.MagicMock(return_value={"ContentLength": size}) + sync_fs._abort_multipart_upload = mock.MagicMock( + side_effect=lambda *args: events.append(("abort", args[2])) + ) + + task = asyncio.ensure_future( + fs._copy_object_with_multipart_upload( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + ) + try: + assert await asyncio.to_thread(started.wait, 5) + task.cancel() + # Gives the cleanup time to return, which it must not do while + # the creation 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 == (["create"] if creation_fails else ["create", ("abort", "uploadid")]) + sync_fs.core.upload_part_copy.assert_not_called() + @pytest.mark.parametrize( "block_size", [ From aa90e742b76799cff31f912ebaf2442bd28d41b9 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 16:34:42 +0900 Subject: [PATCH 8/9] Schedule the aio creation right before its cleanup guard, synchronize the interrupt test, and scope the wait claims Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 5 ++- pyathena/filesystem/s3.py | 6 +-- pyathena/filesystem/s3_async.py | 15 ++++---- tests/pyathena/filesystem/test_s3.py | 57 ++++++++++++++++++++-------- 4 files changed, 55 insertions(+), 28 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index 35c48ad5..f3f62fbd 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -124,8 +124,9 @@ from that version, which needs `s3:GetObjectVersion` and, to copy the tags, upload completes, so the destination exists without them until the last one is written. If an annotation fails to copy, the error is raised and the destination is kept. A failed part copy aborts the multipart upload. So does an interrupt, or the -cancellation of an `AioS3FileSystem` copy, unless the upload has already completed; the -requests in flight, including a CreateMultipartUpload request, finish first. +cancellation of an `AioS3FileSystem` copy, unless the upload has already completed. A +CreateMultipartUpload request and the part copies in flight finish first, and so does +the CompleteMultipartUpload request of an `AioS3FileSystem` copy. Paths are normalized as in fsspec, which drops a trailing slash, so `info`, `isfile`, and `open` treat `s3://YOUR_S3_BUCKET/dir/` as `s3://YOUR_S3_BUCKET/dir`: the object diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 9f273840..2a0d0912 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1744,9 +1744,9 @@ def _copy_object_with_multipart_upload( parts are copied in parallel with UploadPartCopy, and the annotations are copied onto the destination after the upload completes. A failed part or completion aborts the upload, and so does - an interrupt, including one while the upload is being created, once - the running requests have finished; a failed annotation copy is - raised and leaves the destination in place. If HeadObject reports a + an interrupt, after the creation of the upload and the running part + copies have finished; a failed annotation copy is raised and leaves + the destination in place. If HeadObject reports a size that fits in a single CopyObject request, the reported version is copied with CopyObject instead. diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 6bec9aa2..9c84b516 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -585,13 +585,6 @@ async def _copy_object_with_multipart_upload( # See S3FileSystem._copy_object_with_multipart_upload. await asyncio.to_thread(self.core.copy_object, plan.source, plan.destination, **kwargs) return - # A task, so that _abort() can wait for an upload that is created - # after a cancellation. - creation = asyncio.ensure_future( - asyncio.to_thread( - self.core.create_multipart_upload, plan.destination, **plan.create_params - ) - ) upload_id: str semaphore = asyncio.Semaphore(max_workers) @@ -643,6 +636,14 @@ async def _abort() -> None: plan.abort_params, ) + # A task, so that _abort() can wait for an upload that is created + # after a cancellation; scheduled right before the try, so that + # nothing can fail between the two. + creation = asyncio.ensure_future( + asyncio.to_thread( + self.core.create_multipart_upload, plan.destination, **plan.create_params + ) + ) try: # shield keeps a cancellation from cancelling the creation, whose # thread would keep running, so that _abort() can wait for it. diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 001edc31..63d3584d 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -3407,12 +3407,12 @@ def test_copy_object_with_multipart_upload_interrupted_creation(self): # upload that it created before any part is copied, and is re-raised. # The created upload used to be left incomplete. fs = self._make_fs() - started = threading.Event() + waiting = threading.Event() + interrupted = threading.Event() def create_multipart_upload(*args, **kw): - started.set() # Still running when the interrupt arrives. - time.sleep(0.5) + interrupted.wait(5) return SimpleNamespace(upload_id="uploadid") fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) @@ -3420,25 +3420,50 @@ def create_multipart_upload(*args, **kw): # The HeadObject of the source. fs._call.return_value = {"ContentLength": 2 * S3Core.MULTIPART_UPLOAD_MAX_PART_SIZE} fs._abort_multipart_upload = mock.MagicMock() + executor = S3ThreadPoolExecutor(max_workers=2) + submit = executor.submit + + def submit_creation(fn, *args, **kwargs): + future = submit(fn, *args, **kwargs) + if fn is fs.core.create_multipart_upload: + result = future.result + + def wait_for_result(timeout=None): + # The interrupt is sent once the copy waits for the creation. + waiting.set() + return result(timeout) + + future.result = wait_for_result # type: ignore[method-assign] + return future + + executor.submit = submit_creation # type: ignore[method-assign] + fs._create_executor = mock.MagicMock(return_value=executor) + + def handle_interrupt(signum, frame): + interrupted.set() + raise KeyboardInterrupt def interrupt(): - if started.wait(5): - # Lets the main thread return from submit() and wait for the - # creation, where an interrupt during the request arrives. - time.sleep(0.1) + if waiting.wait(5): signal.pthread_kill(threading.main_thread().ident, signal.SIGINT) + previous_handler = signal.signal(signal.SIGINT, handle_interrupt) thread = threading.Thread(target=interrupt, daemon=True) thread.start() - with pytest.raises(KeyboardInterrupt): - fs._copy_object_with_multipart_upload( - S3Path("bucket", "src"), - S3Path("bucket", "dst"), - MetadataDirective="REPLACE", - TaggingDirective="REPLACE", - AnnotationDirective="EXCLUDE", - ) - thread.join(5) + try: + with pytest.raises(KeyboardInterrupt): + fs._copy_object_with_multipart_upload( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + finally: + waiting.set() + thread.join(5) + signal.signal(signal.SIGINT, previous_handler) + interrupted.set() fs._abort_multipart_upload.assert_called_once_with("bucket", "dst", "uploadid", {}) fs.core.upload_part_copy.assert_not_called() From 733cb483bdb0f7527229c2a1806e87f4efafa9c2 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 16:43:56 +0900 Subject: [PATCH 9/9] Interrupt the creation in its test only once it runs, and keep a late interrupt from reaching later tests Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/filesystem/test_s3.py | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 63d3584d..17bbb858 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -3407,12 +3407,14 @@ def test_copy_object_with_multipart_upload_interrupted_creation(self): # upload that it created before any part is copied, and is re-raised. # The created upload used to be left incomplete. fs = self._make_fs() + started = threading.Event() waiting = threading.Event() interrupted = threading.Event() def create_multipart_upload(*args, **kw): - # Still running when the interrupt arrives. - interrupted.wait(5) + started.set() + # Still running when the interrupt arrives, which releases it. + interrupted.wait(30) return SimpleNamespace(upload_id="uploadid") fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) @@ -3444,13 +3446,15 @@ def handle_interrupt(signum, frame): raise KeyboardInterrupt def interrupt(): - if waiting.wait(5): + # Sent only while the creation is running and the copy waits for + # it, which the creation cannot stop doing before the interrupt. + if started.wait(5) and waiting.wait(5): signal.pthread_kill(threading.main_thread().ident, signal.SIGINT) - previous_handler = signal.signal(signal.SIGINT, handle_interrupt) thread = threading.Thread(target=interrupt, daemon=True) - thread.start() + previous_handler = signal.signal(signal.SIGINT, handle_interrupt) try: + thread.start() with pytest.raises(KeyboardInterrupt): fs._copy_object_with_multipart_upload( S3Path("bucket", "src"), @@ -3460,8 +3464,12 @@ def interrupt(): AnnotationDirective="EXCLUDE", ) finally: + # A late interrupt is ignored instead of reaching a later test. + signal.signal(signal.SIGINT, lambda signum, frame: None) + started.set() waiting.set() - thread.join(5) + if thread.ident is not None: + thread.join() signal.signal(signal.SIGINT, previous_handler) interrupted.set()