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..f3f62fbd 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -123,7 +123,10 @@ 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. 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 @@ -280,12 +283,13 @@ 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()`); `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 @@ -331,6 +335,19 @@ 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. 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()`, `copy()` and `mv()` 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..2a0d0912 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,393 +1729,99 @@ 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, and so does + 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. + + 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. 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 - 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) 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, - 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, + # 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 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..9c84b516 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,89 +548,46 @@ 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. 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: - 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. 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 - 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 - ) - upload_id = cast(str, multipart_upload.upload_id) + upload_id: str 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,22 +599,26 @@ 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: 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) @@ -680,10 +629,27 @@ 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), + cast(str, creation.result().upload_id), + 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. + 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(). @@ -696,10 +662,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 +694,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..53369ba5 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 ``ranges``, the parameters and + ``annotations`` 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. + """Typed S3 operations 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,333 @@ 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 = 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( + 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_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_params}, + } + 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..17bbb858 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 @@ -1507,10 +1508,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 +1830,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 +1842,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 +1865,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 +1876,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 +1888,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): @@ -1917,13 +1906,10 @@ 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( - bucket1=bucket1, - key1="src", - size1=MULTIPART_COPY_SIZE, - bucket2="bucket", - key2="dst", + S3Path(source_bucket, "src"), + S3Path("bucket", "dst"), block_size=MULTIPART_COPY_BLOCK_SIZE, **kwargs, ) @@ -2024,17 +2010,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 +2032,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 +2068,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 +2076,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 +2103,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", @@ -2148,7 +2115,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): @@ -3408,14 +3375,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). @@ -3433,6 +3398,84 @@ 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() + waiting = threading.Event() + interrupted = threading.Event() + + def create_multipart_upload(*args, **kw): + 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) + 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() + 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(): + # 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) + + thread = threading.Thread(target=interrupt, daemon=True) + previous_handler = signal.signal(signal.SIGINT, handle_interrupt) + try: + thread.start() + with pytest.raises(KeyboardInterrupt): + fs._copy_object_with_multipart_upload( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + 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() + if thread.ident is not None: + thread.join() + 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() + @pytest.mark.parametrize( "block_size", [ @@ -3449,12 +3492,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..9663602c 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", @@ -506,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", [ @@ -524,12 +563,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 +1043,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 +1059,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 +1080,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..bf6425a0 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,11 +25,17 @@ S3ListBucketsPage, S3ListObjectsPage, S3ListObjectVersionsPage, + S3MultipartCopyPlan, S3ObjectSummary, ) 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) @@ -611,6 +619,286 @@ 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 = MULTIPART_COPY_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"}, + ) + with stubber: + plan = core.plan_multipart_copy( + S3Path("bucket", "src"), + S3Path("bucket", "dst"), + MULTIPART_COPY_BLOCK_SIZE, + **MULTIPART_COPY_KWARGS, + ) + 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, MULTIPART_COPY_BLOCK_SIZE), (MULTIPART_COPY_BLOCK_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):