diff --git a/docs/filesystem.md b/docs/filesystem.md index 35f4401d..e6c7521f 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -104,6 +104,26 @@ the existing object also count toward the limit. A write with `open` that reache limit raises `ValueError` and aborts its multipart upload. Multipart copies with `cp` use parts large enough to stay within the limit. +`cp` copies an object larger than 5 GiB with a multipart upload instead of a single +CopyObject request, with the same result as CopyObject. If HeadObject reports a size +of at most 5 GiB when the copy starts, for example because the size that `cp` found +came from a cached listing, the object is copied with CopyObject instead. CopyObject parameters given as +keyword arguments are sent to the multipart requests that accept them, such as +`CopySourceIfMatch` to each part copy. With the default `COPY` value of +`MetadataDirective`, `TaggingDirective`, and `AnnotationDirective`, the content headers +(such as `ContentType`) and user-defined metadata, the tags, and the annotations of the +source are copied, and the values given for them are ignored, as CopyObject does. A +`REPLACE` directive uses the given values instead, and `AnnotationDirective="EXCLUDE"` +skips the annotations. Copying the tags needs `s3:GetObjectTagging` on the source, and +copying the annotations needs `s3:ListObjectAnnotations` and `s3:GetObjectAnnotation` on +the source and `s3:PutObjectAnnotation` on the destination. A source without a +`?versionId=` suffix whose HeadObject reports a version ID other than `null` is copied +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. + 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 `dir` if it exists, and otherwise the directory `dir`. An object whose key ends in a diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index fc573173..1adaa894 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -20,7 +20,7 @@ from multiprocessing import cpu_count from re import Pattern from typing import Any, BinaryIO, cast -from urllib.parse import unquote_plus +from urllib.parse import unquote_plus, urlencode import botocore.exceptions from boto3 import Session @@ -200,6 +200,18 @@ 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] = re.compile( r"(^s3://|^s3a://|^)(?P[a-zA-Z0-9.\-_]+)(/(?P[^?]+)|/)?" r"($|\?version(Id|ID|id|_id)=(?P.+)$)" @@ -1661,12 +1673,15 @@ def cp_file( control a multipart copy and are not sent to S3. Raises: - ValueError: If trying to copy to a versioned file or copy buckets. + ValueError: If trying to copy to a versioned file or copy buckets, + or if a directive of a multipart copy has an invalid value. Note: Uses multipart copy for objects larger than the maximum part size to optimize performance for large files. The copy operation is performed entirely on the S3 service without data transfer. + The multipart copy applies the metadata, tagging and annotation + directives as CopyObject does (default ``COPY``). A directory ``path1``, which recursive ``copy()`` passes along with the files under it, is skipped. """ @@ -1710,28 +1725,32 @@ def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: # to copy for them. return False size1 = info1.get("size", 0) - if size1 <= self.MULTIPART_UPLOAD_MAX_PART_SIZE: - self._copy_object( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, - **kwargs, - ) - else: - self._copy_object_with_multipart_upload( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - size1=size1, - bucket2=bucket2, - key2=key2, - max_workers=max_workers, - block_size=block_size, - **kwargs, - ) - self.invalidate_cache(path2) + try: + if size1 <= self.MULTIPART_UPLOAD_MAX_PART_SIZE: + self._copy_object( + bucket1=bucket1, + key1=key1, + version_id1=version_id1, + bucket2=bucket2, + key2=key2, + **kwargs, + ) + else: + self._copy_object_with_multipart_upload( + bucket1=bucket1, + key1=key1, + version_id1=version_id1, + size1=size1, + bucket2=bucket2, + key2=key2, + max_workers=max_workers, + block_size=block_size, + **kwargs, + ) + finally: + # A multipart copy that fails to copy an annotation has already + # written the destination. + self.invalidate_cache(path2) return True def _copy_object( @@ -1773,6 +1792,34 @@ def _copy_object_with_multipart_upload( 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. + 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. + """ max_workers = max_workers if max_workers else self.max_workers block_size = block_size if block_size else self.MULTIPART_UPLOAD_MAX_PART_SIZE if ( @@ -1785,18 +1832,41 @@ def _copy_object_with_multipart_upload( f"5 GiB ({self.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), 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.MULTIPART_UPLOAD_MAX_PART_SIZE: + # 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, + ) + return + # The size of the copied version, not the one that the caller found. + ranges = self._get_copy_ranges(size1 if head_size is None else head_size, block_size) copy_source = { "Bucket": bucket1, "Key": key1, } if version_id1: copy_source.update({"VersionId": version_id1}) - - ranges = self._get_copy_ranges(size1, block_size) + # 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._create_multipart_upload( bucket=bucket2, key=key2, - **kwargs, + **create_kwargs, ) with self._create_executor(max_workers=max_workers) as executor: futures = [ @@ -1812,13 +1882,274 @@ def _copy_object_with_multipart_upload( ) for i, range_ in enumerate(ranges) ] - self._finish_multipart_upload( + completed = self._finish_multipart_upload( bucket=bucket2, key=key2, upload_id=cast(str, multipart_upload.upload_id), futures=futures, request_kwargs=kwargs, ) + for name in annotations: + self._copy_object_annotation( + name, bucket1, key1, version_id1, bucket2, key2, completed, 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: s3://{bucket}/{key}?versionId={version_id}") + head = S3Metadata( + self._call( + self._client.head_object, + **self._get_operation_kwargs("head_object", source_kwargs), + **source, + ) + ) + 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.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: s3://{bucket}/{key}?versionId={version_id}") + response = self._call( + self._client.get_object_tagging, + **self._get_operation_kwargs("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._get_operation_kwargs("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._get_operation_kwargs( + "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: s3://{bucket}/{key}?versionId={version_id}") + 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 s3://{bucket1}/{key1}?versionId={version_id1} " + f"to s3://{bucket2}/{key2}." + ) + response = self._call( + self._client.get_object_annotation, + **self._get_operation_kwargs( + "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._get_operation_kwargs("put_object_annotation", kwargs), **destination}, + ) def _get_copy_ranges(self, size: int, block_size: int) -> list[tuple[int, int]]: """Split an object into the source ranges of a multipart copy. @@ -2050,22 +2381,39 @@ def _finish_multipart_upload( # be stored after the abort, so wait for the parts that could not # be cancelled first. wait([future for future in futures if not future.cancel()]) - try: - self._call( - self._client.abort_multipart_upload, - **{ - **self._get_operation_kwargs("abort_multipart_upload", request_kwargs), - "Bucket": bucket, - "Key": key, - "UploadId": upload_id, - }, - ) - except Exception: - _logger.exception( - f"Failed to abort multipart upload {upload_id} to s3://{bucket}/{key}." - ) + self._abort_multipart_upload(bucket, key, upload_id, request_kwargs) raise + def _abort_multipart_upload( + self, bucket: str, key: str, upload_id: str, request_kwargs: Mapping[str, Any] + ) -> None: + """Abort a failed multipart upload, logging an error of the abort. + + The error of the abort is logged instead of raised, so that the + caller can re-raise the error that made the upload fail. + + Args: + bucket: S3 bucket name. + key: Object key being uploaded. + upload_id: Unique identifier for the multipart upload. + request_kwargs: Parameters of the upload; the abort receives + those that it accepts. + """ + try: + self._call( + self._client.abort_multipart_upload, + **{ + **self._get_operation_kwargs("abort_multipart_upload", request_kwargs), + "Bucket": bucket, + "Key": key, + "UploadId": upload_id, + }, + ) + except Exception: + _logger.exception( + f"Failed to abort multipart upload {upload_id} to s3://{bucket}/{key}." + ) + def cat_file( self, path: str, start: int | None = None, end: int | None = None, **kwargs ) -> bytes: diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index fdbd0121..17a81772 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -448,29 +448,32 @@ async def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: # S3FileSystem.cp_file. return False size1 = info1.get("size", 0) - if size1 <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: - await asyncio.to_thread( - self._sync_fs._copy_object, - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, - **kwargs, - ) - else: - await self._copy_object_with_multipart_upload( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - size1=size1, - bucket2=bucket2, - key2=key2, - max_workers=max_workers, - block_size=block_size, - **kwargs, - ) - self._sync_fs.invalidate_cache(path2) + try: + if size1 <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: + await asyncio.to_thread( + self._sync_fs._copy_object, + bucket1=bucket1, + key1=key1, + version_id1=version_id1, + bucket2=bucket2, + key2=key2, + **kwargs, + ) + else: + await self._copy_object_with_multipart_upload( + bucket1=bucket1, + key1=key1, + version_id1=version_id1, + size1=size1, + bucket2=bucket2, + key2=key2, + max_workers=max_workers, + block_size=block_size, + **kwargs, + ) + finally: + # See S3FileSystem._copy_file. + self._sync_fs.invalidate_cache(path2) return True async def _copy_object_with_multipart_upload( @@ -485,6 +488,28 @@ async def _copy_object_with_multipart_upload( 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 with ``asyncio.gather`` and + ``asyncio.to_thread``. + + 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. + 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. + """ max_workers = max_workers if max_workers else self._sync_fs.max_workers block_size = block_size if block_size else S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE if ( @@ -498,52 +523,130 @@ async def _copy_object_with_multipart_upload( 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 + ) + if head_size is not None and head_size <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: + # 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, + ) + return + # The size of the copied version; see S3FileSystem. + ranges = self._sync_fs._get_copy_ranges( + size1 if head_size is None else head_size, block_size + ) copy_source: dict[str, Any] = { "Bucket": bucket1, "Key": key1, } if version_id1: copy_source["VersionId"] = version_id1 - - ranges = self._sync_fs._get_copy_ranges(size1, block_size) + # 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._sync_fs._create_multipart_upload, bucket=bucket2, key=key2, - **kwargs, + **create_kwargs, ) + upload_id = cast(str, multipart_upload.upload_id) semaphore = asyncio.Semaphore(max_workers) part_kwargs = self._sync_fs._get_operation_kwargs("upload_part_copy", kwargs) + failed = False - async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any]: + async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any] | None: + nonlocal failed async with semaphore: - result = await asyncio.to_thread( - self._sync_fs._upload_part_copy, - bucket=bucket2, - key=key2, - copy_source=copy_source, - upload_id=cast(str, multipart_upload.upload_id), - part_number=i + 1, - copy_source_ranges=range_, - **part_kwargs, - ) + if failed: + # The upload is being aborted; do not start more parts. + return None + try: + result = await asyncio.to_thread( + self._sync_fs._upload_part_copy, + bucket=bucket2, + key=key2, + copy_source=copy_source, + upload_id=upload_id, + part_number=i + 1, + copy_source_ranges=range_, + **part_kwargs, + ) + except Exception: + # Set before the semaphore lets a waiting part start. + failed = True + raise return { "ETag": result.etag, "PartNumber": result.part_number, } - parts = await asyncio.gather(*[_upload_part(i, r) for i, r in enumerate(ranges)]) - parts_list = sorted(parts, key=lambda x: x["PartNumber"]) + tasks = [asyncio.ensure_future(_upload_part(i, r)) for i, r in enumerate(ranges)] + try: + # gather keeps the part-number order of the tasks. + parts = await asyncio.gather(*tasks) + completed = await asyncio.to_thread( + self._sync_fs._complete_multipart_upload, + bucket=bucket2, + key=key2, + upload_id=upload_id, + parts=cast(list[dict[str, Any]], parts), + **self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs), + ) + except Exception: + failed = True + # 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) + await asyncio.to_thread( + self._sync_fs._abort_multipart_upload, bucket2, key2, upload_id, kwargs + ) + raise - await asyncio.to_thread( - self._sync_fs._complete_multipart_upload, - bucket=bucket2, - key=key2, - upload_id=cast(str, multipart_upload.upload_id), - parts=parts_list, - **self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs), + failed = False + + async def _copy_annotation(name: str) -> None: + nonlocal failed + async with semaphore: + if failed: + # Do not start more copies after one failed, as + # S3FileSystem does. + return + try: + await asyncio.to_thread( + self._sync_fs._copy_object_annotation, + name, + bucket1, + key1, + version_id1, + bucket2, + key2, + completed, + kwargs, + ) + except Exception: + failed = True + raise + + results = await asyncio.gather( + *[_copy_annotation(name) for name in annotations], return_exceptions=True ) + for result in results: + if isinstance(result, BaseException): + raise result async def _find( self, diff --git a/pyproject.toml b/pyproject.toml index f56cbae1..37465714 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,8 +8,8 @@ maintainers = [ {name = "laughingman7743", email = "laughingman7743@gmail.com"}, ] dependencies = [ - "boto3>=1.41.2", - "botocore>=1.41.2", + "boto3>=1.43.31", + "botocore>=1.43.31", "tenacity>=4.1.0", "fsspec", "python-dateutil", diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 600372aa..0ef4bc68 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -41,6 +41,12 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.util import ( + MULTIPART_COPY_BLOCK_SIZE, + MULTIPART_COPY_KWARGS, + MULTIPART_COPY_SIZE, + stub_multipart_copy, +) # A client that sends no requests; its service model selects the parameters # that each S3 operation accepts. @@ -1714,7 +1720,15 @@ 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() - kwargs = {"ContentType": "text/csv", "RequestPayer": "requester"} + fs._call.return_value = {} + # The directives make the copy use the given values without reading + # the source's metadata, tags and annotations (GH-973). + directives = { + "MetadataDirective": "REPLACE", + "TaggingDirective": "REPLACE", + "AnnotationDirective": "EXCLUDE", + } + kwargs = {"ContentType": "text/csv", "RequestPayer": "requester", **directives} fs._copy_object_with_multipart_upload( bucket1="bucket", @@ -1725,13 +1739,265 @@ def test_copy_object_with_multipart_upload_request_parameters(self): **kwargs, ) - fs._create_multipart_upload.assert_called_once_with(bucket="bucket", key="dst", **kwargs) + fs._create_multipart_upload.assert_called_once_with( + bucket="bucket", key="dst", ContentType="text/csv", RequestPayer="requester" + ) + # Only the HeadObject of the source, for its version. + assert fs._call.call_count == 1 assert all( c.kwargs["RequestPayer"] == "requester" and "ContentType" not in c.kwargs for c in fs._upload_part_copy.call_args_list ) assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == kwargs + @staticmethod + def _stubbed_copy_fs(**kwargs): + # max_workers=1 copies the parts in the order of the stubbed + # responses. + return S3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + skip_instance_cache=True, + max_workers=1, + **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", + block_size=MULTIPART_COPY_BLOCK_SIZE, + **kwargs, + ) + + def test_copy_object_with_multipart_upload_copies_source(self): + # GH-973: as CopyObject does by default, the multipart copy copies + # the content headers, the user-defined metadata, the tags and the + # annotations of the source, ignoring the values of the copy; the + # source condition goes to the part copies, and the source's lookups + # get the source's expected bucket owner. + fs = self._stubbed_copy_fs() + with Stubber(fs._client) as stubber: + stub_multipart_copy(stubber) + self._multipart_copy(fs, **MULTIPART_COPY_KWARGS) + stubber.assert_no_pending_responses() + + def test_copy_object_with_multipart_upload_failed_listing(self): + # GH-973: the annotations are listed before the upload is created, so + # a caller without s3:ListObjectAnnotations fails before anything is + # written. + fs = self._stubbed_copy_fs() + with Stubber(fs._client) as stubber: + stub_multipart_copy(stubber, fail_list=True) + with pytest.raises(PermissionError): + self._multipart_copy(fs, **MULTIPART_COPY_KWARGS) + stubber.assert_no_pending_responses() + + def test_copy_object_with_multipart_upload_failed_part(self): + fs = self._stubbed_copy_fs() + with Stubber(fs._client) as stubber: + stub_multipart_copy(stubber, fail_part=True) + with pytest.raises(OSError, match="part failed"): + self._multipart_copy(fs, **MULTIPART_COPY_KWARGS) + stubber.assert_no_pending_responses() + + def test_copy_object_with_multipart_upload_failed_annotation(self): + # GH-973: a failed annotation copy is raised; the completed + # destination is neither aborted nor deleted. + fs = self._stubbed_copy_fs() + with Stubber(fs._client) as stubber: + stub_multipart_copy(stubber, fail_annotation=True) + with pytest.raises(PermissionError): + self._multipart_copy(fs, **MULTIPART_COPY_KWARGS) + stubber.assert_no_pending_responses() + + def test_cp_file_failed_multipart_copy_invalidates_cache(self): + # GH-973: a multipart copy that fails to copy an annotation has + # written the destination, so its cached entries are removed. + fs = self._make_fs() + fs.info = mock.MagicMock(return_value=self._file_object("src")) + fs.info.return_value.size = S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE + 1 + fs._copy_object_with_multipart_upload = mock.MagicMock(side_effect=PermissionError) + fs.dircache["bucket/dst"] = [self._file_object("dst")] + + with pytest.raises(PermissionError): + fs.cp_file("s3://bucket/src", "s3://bucket/dst") + assert "bucket/dst" not in fs.dircache + + def test_copy_object_with_multipart_upload_head_object_size(self): + # GH-973: the ranges cover the size that HeadObject reports for the + # copied object, not a cached size, and the "null" version of a + # bucket with versioning suspended is not pinned. + fs = self._make_fs() + block_size = MULTIPART_COPY_BLOCK_SIZE + fs._call.return_value = { + "ContentLength": 2 * block_size + S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + "VersionId": "null", + } + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._upload_part_copy = mock.MagicMock() + fs._finish_multipart_upload = mock.MagicMock() + + self._multipart_copy( + fs, + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + + # The parts are copied in parallel, in any order. + parts = sorted( + (c.kwargs["copy_source_ranges"], c.kwargs["copy_source"]) + for c in fs._upload_part_copy.call_args_list + ) + source = {"Bucket": "bucket", "Key": "src"} + assert parts == [ + ((0, block_size), source), + ((block_size, 2 * block_size), source), + ((2 * block_size, fs._call.return_value["ContentLength"]), source), + ] + + @pytest.mark.parametrize("size", [0, 10]) + def test_copy_object_with_multipart_upload_small_head_object_size(self, size): + # GH-973: when a cached size over 5 GiB is stale and HeadObject + # reports a size that fits in a single CopyObject request, including + # 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._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", + ContentType="text/csv", + RequestPayer="requester", + ) + fs._create_multipart_upload.assert_not_called() + # Only HeadObject; the tags are not read for the multipart upload. + assert fs._call.call_count == 1 + + def test_copy_object_with_multipart_upload_replace_directives(self): + # GH-973: REPLACE uses the values of the copy without reading the + # source, and EXCLUDE skips the annotations. + fs = self._stubbed_copy_fs() + 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( + "create_multipart_upload", + {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "ContentType": "text/plain", "Tagging": "a=1"}, + ) + 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, + ContentType="text/plain", + Tagging="a=1", + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + "directive", + [ + {"MetadataDirective": "EXCLUDE"}, + {"TaggingDirective": "copy"}, + {"AnnotationDirective": "REPLACE"}, + ], + ) + def test_copy_object_with_multipart_upload_invalid_directive(self, directive): + fs = self._stubbed_copy_fs() + 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. + fs = self._stubbed_copy_fs() + sse_c = {"CopySourceSSECustomerAlgorithm": "AES256", "CopySourceSSECustomerKey": "k" * 32} + with Stubber(fs._client) as stubber: + stubber.add_response( + "head_object", + {"ContentType": "text/csv"}, + { + "Bucket": "bucket", + "Key": "src", + "SSECustomerAlgorithm": "AES256", + "SSECustomerKey": "k" * 32, + }, + ) + stubber.add_response( + "create_multipart_upload", + {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "ContentType": "text/csv", "Metadata": {}}, + ) + 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, TaggingDirective="REPLACE", **sse_c) + stubber.assert_no_pending_responses() + + def test_copy_object_with_multipart_upload_directory_bucket_source(self): + # GH-973: objects in a directory bucket have neither tags nor + # annotations, and the bucket supports neither GetObjectTagging nor + # ListObjectAnnotations. + fs = self._stubbed_copy_fs() + bucket = "bucket--usw2-az1--x-s3" + with Stubber(fs._client) as stubber: + stubber.add_response( + "head_object", {"ContentType": "text/csv"}, {"Bucket": bucket, "Key": "src"} + ) + stubber.add_response( + "create_multipart_upload", + {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "ContentType": "text/csv", "Metadata": {}}, + ) + 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) + stubber.assert_no_pending_responses() + def test_pipe_file_invalid_path_raises(self): fs = self._make_fs() with pytest.raises(ValueError, match="Cannot write to a bucket"): @@ -2993,6 +3259,7 @@ def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): ) fs._upload_part_copy = mock.MagicMock() fs._finish_multipart_upload = mock.MagicMock() + fs._call.return_value = {} fs._copy_object_with_multipart_upload( bucket1="bucket", @@ -3001,6 +3268,11 @@ def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): bucket2="bucket", key2="dst", max_workers=max_workers, + # Copy without reading the metadata, tags and annotations of the + # source (GH-973). + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", ) parts = sorted( diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index c8d3db98..bdb22913 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -17,6 +17,7 @@ import boto3 import fsspec import pytest +from botocore.stub import Stubber from fsspec import Callback from pyathena.filesystem.s3 import S3File, S3FileSystem @@ -29,6 +30,12 @@ ) from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.util import ( + MULTIPART_COPY_BLOCK_SIZE, + MULTIPART_COPY_KWARGS, + MULTIPART_COPY_SIZE, + stub_multipart_copy, +) @pytest.fixture(scope="class") @@ -158,6 +165,7 @@ 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._complete_multipart_upload = mock.MagicMock() + sync_fs._call = mock.MagicMock(return_value={}) await fs._copy_object_with_multipart_upload( bucket1="bucket", @@ -165,6 +173,11 @@ async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): size1=5 * 2**30 + 2**20, bucket2="bucket", key2="dst", + # Copy without reading the metadata, tags and annotations of the + # source (GH-973). + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", ) parts = sorted( @@ -176,6 +189,173 @@ async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): (2, (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)), ] + @staticmethod + async def _multipart_copy(fs=None, **kwargs): + # max_workers=1 runs the stubbed requests in a deterministic order. + fs = fs or AioS3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + with Stubber(fs._sync_fs._client) as stubber: + 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", + block_size=MULTIPART_COPY_BLOCK_SIZE, + **MULTIPART_COPY_KWARGS, + ) + finally: + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize("size", [0, 10]) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_small_head_object_size(self, size): + # GH-973: see + # TestS3FileSystem.test_copy_object_with_multipart_upload_small_head_object_size. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + sync_fs = fs._sync_fs + sync_fs._call = mock.MagicMock(return_value={"ContentLength": size, "VersionId": "v1"}) + sync_fs._copy_object = mock.MagicMock() + sync_fs._create_multipart_upload = mock.MagicMock() + + await fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=MULTIPART_COPY_SIZE, + bucket2="bucket", + key2="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", + ContentType="text/csv", + RequestPayer="requester", + ) + sync_fs._create_multipart_upload.assert_not_called() + # Only HeadObject; the tags are not read for the multipart upload. + assert sync_fs._call.call_count == 1 + + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_copies_source(self): + # GH-973: the same requests as S3FileSystem; see + # TestS3FileSystem.test_copy_object_with_multipart_upload_copies_source. + await self._multipart_copy() + + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_failed_listing(self): + # GH-973: nothing is written when the annotations cannot be listed. + with pytest.raises(PermissionError): + await self._multipart_copy(fail_list=True) + + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_failed_part(self): + # GH-973: a failed part copy aborts the upload, as in S3FileSystem. + with pytest.raises(OSError, match="part failed"): + await self._multipart_copy(fail_part=True) + + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_failed_annotation(self): + # GH-973: a failed annotation copy is raised; the completed + # destination is neither aborted nor deleted, and no other + # annotation is copied. + fs = AioS3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + sync_fs = fs._sync_fs + with ( + mock.patch.object( + sync_fs, "_copy_object_annotation", wraps=sync_fs._copy_object_annotation + ) as copy_annotation, + pytest.raises(PermissionError), + ): + await self._multipart_copy(fs, fail_annotation=True) + assert [c.args[0] for c in copy_annotation.call_args_list] == ["a1"] + + @pytest.mark.asyncio + async def test_cp_file_failed_multipart_copy_invalidates_cache(self): + # GH-973: see TestS3FileSystem.test_cp_file_failed_multipart_copy_invalidates_cache. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._info = mock.AsyncMock( + return_value=S3Object( + init={"ContentLength": S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE + 1}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="src", + ) + ) + fs._copy_object_with_multipart_upload = mock.AsyncMock(side_effect=PermissionError) + fs._sync_fs.dircache["bucket/dst"] = [] + + with pytest.raises(PermissionError): + await fs._cp_file("s3://bucket/src", "s3://bucket/dst") + assert "bucket/dst" not in fs._sync_fs.dircache + + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_waits_for_running_parts(self): + # GH-973: the abort waits for the part copies that are running when + # one fails, and no part starts after the failure. + fs = AioS3FileSystem(connection=mock.MagicMock(), max_workers=2, skip_instance_cache=True) + sync_fs = fs._sync_fs + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + events = [] + failed = threading.Event() + + def upload_part_copy(**kw): + part_number = kw["part_number"] + events.append(f"start {part_number}") + if part_number == 1: + failed.wait(5) + time.sleep(0.05) + events.append("end 1") + return SimpleNamespace(etag='"e"', part_number=part_number) + failed.set() + raise OSError("part failed") + + sync_fs._upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) + sync_fs._complete_multipart_upload = mock.MagicMock() + # The HeadObject of the source, for its version. + sync_fs._call = mock.MagicMock(return_value={}) + sync_fs._abort_multipart_upload = mock.MagicMock( + side_effect=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 * S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + bucket2="bucket", + key2="dst", + block_size=S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + MetadataDirective="REPLACE", + TaggingDirective="REPLACE", + AnnotationDirective="EXCLUDE", + ) + + # Part 3 waits for a worker and is not started after the failure. + assert sorted(events[:2]) == ["start 1", "start 2"] + assert events[2:] == ["end 1", "abort"] + sync_fs._complete_multipart_upload.assert_not_called() + @pytest.mark.parametrize( "block_size", [ @@ -641,6 +821,13 @@ def upload_part_copy(**kw): sync_fs._upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) sync_fs._complete_multipart_upload = mock.MagicMock() + # The HeadObject of the source, for its version. + sync_fs._call = mock.MagicMock(return_value={}) + directives = { + "MetadataDirective": "REPLACE", + "TaggingDirective": "REPLACE", + "AnnotationDirective": "EXCLUDE", + } await fs._cp_file( "s3://bucket/src", @@ -649,6 +836,9 @@ def upload_part_copy(**kw): max_workers=1, RequestPayer="requester", ContentType="text/csv", + # Copy without reading the metadata, tags and annotations of the + # source (GH-973). + **directives, ) if size <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: @@ -660,6 +850,7 @@ def upload_part_copy(**kw): key2="dst", RequestPayer="requester", ContentType="text/csv", + **directives, ) else: sync_fs._create_multipart_upload.assert_called_once_with( diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index b8479d36..85974109 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -5,14 +5,16 @@ # # SPDX-License-Identifier: MIT +import io import time from concurrent.futures import wait -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from pathlib import Path from unittest.mock import patch from botocore.config import Config from botocore.exceptions import ClientError +from botocore.response import StreamingBody from dateutil.tz import gettz from jinja2 import Environment, FileSystemLoader from sqlalchemy import types @@ -240,3 +242,155 @@ def interrupting_wait(futures, timeout=None): return wait(futures, timeout) return patch("pyathena.common.wait", side_effect=interrupting_wait), raised + + +# A source object larger than a single CopyObject request allows, copied +# from bucket/src to bucket/dst by a multipart copy of two parts with the +# maximum block size (GH-973). +MULTIPART_COPY_BLOCK_SIZE = S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE +MULTIPART_COPY_SIZE = MULTIPART_COPY_BLOCK_SIZE + S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE +MULTIPART_COPY_EXPIRES = datetime(2030, 1, 1, tzinfo=UTC) +MULTIPART_COPY_KWARGS = { + # Ignored with the default COPY directives, as CopyObject ignores them. + "ContentType": "text/plain", + "Tagging": "ignored=1", + # Not metadata: sent to CreateMultipartUpload. + "StorageClass": "STANDARD_IA", + "RequestPayer": "requester", + "ExpectedBucketOwner": "111111111111", + "ExpectedSourceBucketOwner": "222222222222", + # A source condition: sent to the part copies only. + "CopySourceIfMatch": '"src"', +} + + +def _annotation(name): + return {"AnnotationName": name, "LastModified": MULTIPART_COPY_EXPIRES, "Size": 1} + + +def stub_multipart_copy(stubber, fail_list=False, fail_part=False, fail_annotation=False): + """Queue the requests of a multipart copy with MULTIPART_COPY_KWARGS. + + The source, in a versioned bucket, has user-defined metadata, a tag and + two annotations listed on two pages, which are copied with the default + COPY directives from the version that HeadObject reports. + + Args: + stubber: The Stubber of the S3 client. + fail_list: Deny the listing of the annotations; nothing is written + then. + fail_part: Fail the second part copy; the upload is then aborted. + fail_annotation: Fail the write of the first annotation; no + other annotation is copied then. + """ + source = { + "Bucket": "bucket", + "Key": "src", + "RequestPayer": "requester", + "ExpectedBucketOwner": "222222222222", + } + destination = { + "Bucket": "bucket", + "Key": "dst", + "RequestPayer": "requester", + "ExpectedBucketOwner": "111111111111", + } + stubber.add_response( + "head_object", + { + "ContentLength": MULTIPART_COPY_SIZE, + "ContentType": "text/csv", + "CacheControl": "max-age=60", + "Expires": MULTIPART_COPY_EXPIRES, + "StorageClass": "GLACIER_IR", + "Metadata": {"owner": "etl"}, + "VersionId": "v-src", + }, + source, + ) + # The version that HeadObject reports is read from then on. + source = {**source, "VersionId": "v-src"} + stubber.add_response("get_object_tagging", {"TagSet": [{"Key": "t 1", "Value": "v1"}]}, source) + # The annotations are listed before anything is written. + if fail_list: + stubber.add_client_error( + "list_object_annotations", "AccessDenied", 403, expected_params=source + ) + return + stubber.add_response( + "list_object_annotations", + {"Annotations": [_annotation("a1")], "NextContinuationToken": "next"}, + source, + ) + stubber.add_response( + "list_object_annotations", + {"Annotations": [_annotation("a2")]}, + {**source, "ContinuationToken": "next"}, + ) + stubber.add_response( + "create_multipart_upload", + {"Bucket": "bucket", "Key": "dst", "UploadId": "u"}, + { + **destination, + "CacheControl": "max-age=60", + "ContentType": "text/csv", + "Expires": MULTIPART_COPY_EXPIRES, + "Metadata": {"owner": "etl"}, + "Tagging": "t+1=v1", + "StorageClass": "STANDARD_IA", + }, + ) + ranges = { + 1: (0, MULTIPART_COPY_BLOCK_SIZE - 1), + 2: (MULTIPART_COPY_BLOCK_SIZE, MULTIPART_COPY_SIZE - 1), + } + for part_number in (1, 2): + part = { + **destination, + "CopySource": {"Bucket": "bucket", "Key": "src", "VersionId": "v-src"}, + "UploadId": "u", + "PartNumber": part_number, + "CopySourceRange": "bytes={}-{}".format(*ranges[part_number]), + "CopySourceIfMatch": '"src"', + "ExpectedSourceBucketOwner": "222222222222", + } + if fail_part and part_number == 2: + stubber.add_client_error( + "upload_part_copy", "InternalError", "part failed", 500, expected_params=part + ) + stubber.add_response("abort_multipart_upload", {}, {**destination, "UploadId": "u"}) + return + stubber.add_response( + "upload_part_copy", {"CopyPartResult": {"ETag": f'"p{part_number}"'}}, part + ) + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"dst"', "VersionId": "v-dst"}, + { + **destination, + "UploadId": "u", + "MultipartUpload": { + "Parts": [{"ETag": '"p1"', "PartNumber": 1}, {"ETag": '"p2"', "PartNumber": 2}] + }, + }, + ) + for name in ("a1", "a2"): + payload = f"payload of {name}".encode() + stubber.add_response( + "get_object_annotation", + {"AnnotationPayload": StreamingBody(io.BytesIO(payload), len(payload))}, + {**source, "AnnotationName": name}, + ) + put = { + **destination, + "AnnotationName": name, + "AnnotationPayload": payload, + "VersionId": "v-dst", + "ObjectIfMatch": '"dst"', + } + if fail_annotation: + stubber.add_client_error( + "put_object_annotation", "AccessDenied", 403, expected_params=put + ) + return + stubber.add_response("put_object_annotation", {}, put) diff --git a/uv.lock b/uv.lock index 4a980c35..f736824b 100644 --- a/uv.lock +++ b/uv.lock @@ -976,8 +976,8 @@ dev = [ [package.metadata] requires-dist = [ - { name = "boto3", specifier = ">=1.41.2" }, - { name = "botocore", specifier = ">=1.41.2" }, + { name = "boto3", specifier = ">=1.43.31" }, + { name = "botocore", specifier = ">=1.43.31" }, { name = "fsspec" }, { name = "pandas", marker = "extra == 'pandas'", specifier = ">=3.0.0" }, { name = "polars", marker = "extra == 'polars'", specifier = ">=1.39.0" },