diff --git a/docs/filesystem.md b/docs/filesystem.md index f3f62fbd..765cf0c8 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -348,6 +348,29 @@ instead. `copy_object_annotation()` copies one annotation onto the destination a the upload completes, with GetObjectAnnotation and PutObjectAnnotation. The filesystems' `cp_file()`, `copy()` and `mv()` run these plans. +The multipart primitives take the `S3MultipartUpload` returned by creation, +keeping its bucket, key, upload ID and checksum configuration together. +The core retains no upload state. +`upload_part()` passes the upload's checksum algorithm to the SDK; +completion selects the matching part checksum and sends the upload's +`ChecksumType` when present, including `FULL_OBJECT`. +Without a creation algorithm, completion sends the ETag and part number without +checksums that the SDK may add to part uploads. +`abort_multipart_upload()` also accepts uploads returned by +`list_multipart_uploads()`. + +```python +upload = core.create_multipart_upload( + S3Path("YOUR_S3_BUCKET", "path/to/object"), ChecksumAlgorithm="SHA256" +) +try: + part = core.upload_part(upload, 1, b"data") + result = core.complete_multipart_upload(upload, [part]) +except Exception: + core.abort_multipart_upload(upload) + raise +``` + ## 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 2a0d0912..387c814d 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1784,18 +1784,14 @@ def _copy_object_with_multipart_upload( 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), + creation.result(), plan.abort_params, ) raise - upload_id = cast(str, multipart_upload.upload_id) futures = [ executor.submit( self.core.upload_part_copy, - path=plan.destination, - upload_id=upload_id, + upload=multipart_upload, part_number=i + 1, source=plan.source, range_=range_, @@ -1804,9 +1800,7 @@ def _copy_object_with_multipart_upload( for i, range_ in enumerate(plan.ranges) ] completed = self._finish_multipart_upload( - bucket=plan.destination.bucket, - key=cast(str, plan.destination.key), - upload_id=upload_id, + upload=multipart_upload, futures=futures, # Filtered again for the completion and the abort, which # leaves the plan's parameters of each unchanged. @@ -1930,9 +1924,7 @@ def pipe_file( def _finish_multipart_upload( self, - bucket: str, - key: str, - upload_id: str, + upload: S3MultipartUpload, futures: list[Future[S3MultipartUploadPart]], request_kwargs: Mapping[str, Any] | None = None, abort: bool = True, @@ -1946,9 +1938,7 @@ def _finish_multipart_upload( false. The original error is then re-raised. Args: - bucket: S3 bucket name. - key: Object key being uploaded. - upload_id: Unique identifier for the multipart upload. + upload: The multipart upload returned by creation. futures: Futures of the part uploads, in part-number order. request_kwargs: Parameters of the upload, such as ``RequestPayer`` or the SSE-C parameters; the completion and @@ -1964,8 +1954,7 @@ def _finish_multipart_upload( # The futures are in part-number order. parts = [future.result() for future in futures] return self.core.complete_multipart_upload( - S3Path(bucket, key), - upload_id, + upload, parts, **self.core.operation_params("complete_multipart_upload", request_kwargs), ) @@ -1976,11 +1965,11 @@ 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()]) - self._abort_multipart_upload(bucket, key, upload_id, request_kwargs) + self._abort_multipart_upload(upload, request_kwargs) raise def _abort_multipart_upload( - self, bucket: str, key: str, upload_id: str, request_kwargs: Mapping[str, Any] + self, upload: S3MultipartUpload, request_kwargs: Mapping[str, Any] ) -> None: """Abort a failed multipart upload, logging an error of the abort. @@ -1988,21 +1977,19 @@ def _abort_multipart_upload( 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. + upload: The multipart upload returned by creation. request_kwargs: Parameters of the upload; the abort receives those that it accepts. """ try: self.core.abort_multipart_upload( - S3Path(bucket, key), - upload_id, + upload, **self.core.operation_params("abort_multipart_upload", request_kwargs), ) except Exception: _logger.exception( - f"Failed to abort multipart upload {upload_id} to s3://{bucket}/{key}." + f"Failed to abort multipart upload {upload.upload_id} " + f"to s3://{upload.bucket}/{upload.key}." ) def cat_file( @@ -2623,6 +2610,9 @@ def object_version_info( def clear_multipart_uploads(self, path: str) -> None: """Abort any incomplete multipart uploads in the bucket. + Uploads completed or aborted after listing are already cleared. + All abort results are checked before another abort error is raised. + Args: path: S3 bucket or key path (e.g., "bucket", "s3://bucket" or "s3://bucket/prefix"). If the path contains a key, only the @@ -2632,17 +2622,30 @@ def clear_multipart_uploads(self, path: str) -> None: uploads = self.list_multipart_uploads(path) if not uploads: return + error: Exception | None = None with self._create_executor(max_workers=self.max_workers) as executor: futures = [ executor.submit( self.core.abort_multipart_upload, - S3Path(cast(str, upload.bucket), cast(str, upload.key)), - cast(str, upload.upload_id), + upload, ) for upload in uploads ] for future in as_completed(futures): - future.result() + try: + future.result() + except Exception as e: + cause = e.__cause__ + if ( + isinstance(e, FileNotFoundError) + and isinstance(cause, botocore.exceptions.ClientError) + and S3ClientError(cause).code == "NoSuchUpload" + ): + continue + if error is None: + error = e + if error is not None: + raise error def created(self, path: str) -> datetime: """Return the creation time of the path. @@ -3287,8 +3290,7 @@ def _initiate_upload(self) -> None: self.multipart_upload_parts.append( self._executor.submit( self.fs.core.upload_part_copy, - path=path, - upload_id=cast(str, self.multipart_upload.upload_id), + upload=self.multipart_upload, part_number=i + 1, # The existing object is copied into the upload. source=path, @@ -3300,8 +3302,7 @@ def _initiate_upload(self) -> None: self.multipart_upload_parts.append( self._executor.submit( self.fs.core.upload_part_copy, - path=path, - upload_id=cast(str, self.multipart_upload.upload_id), + upload=self.multipart_upload, part_number=1, source=path, **self._get_request_kwargs("upload_part_copy"), @@ -3365,8 +3366,7 @@ def _upload_chunk(self, final: bool = False) -> bool: self.multipart_upload_parts.append( self._executor.submit( self.fs.core.upload_part, - path=S3Path(self.bucket, self.key), - upload_id=cast(str, self.multipart_upload.upload_id), + upload=self.multipart_upload, part_number=part_number, body=upload, **self._get_request_kwargs("upload_part"), @@ -3421,9 +3421,7 @@ def commit(self) -> None: upload_id = cast(str, self.multipart_upload.upload_id) try: self.fs._finish_multipart_upload( - bucket=self.bucket, - key=self.key, - upload_id=upload_id, + upload=self.multipart_upload, futures=self.multipart_upload_parts, request_kwargs=self.s3_additional_kwargs, abort=False, @@ -3456,8 +3454,7 @@ def discard(self) -> None: # be cancelled first. wait([f for f in self.multipart_upload_parts if not f.cancel()]) self.fs.core.abort_multipart_upload( - S3Path(self.bucket, self.key), - cast(str, self.multipart_upload.upload_id), + self.multipart_upload, **self._get_request_kwargs("abort_multipart_upload"), ) diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 9c84b516..75f95cf3 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -585,7 +585,7 @@ async def _copy_object_with_multipart_upload( # See S3FileSystem._copy_object_with_multipart_upload. await asyncio.to_thread(self.core.copy_object, plan.source, plan.destination, **kwargs) return - upload_id: str + multipart_upload: S3MultipartUpload semaphore = asyncio.Semaphore(max_workers) failed = False @@ -599,8 +599,7 @@ async def _upload_part(i: int, range_: tuple[int, int]) -> S3MultipartUploadPart try: return await asyncio.to_thread( self.core.upload_part_copy, - path=plan.destination, - upload_id=upload_id, + upload=multipart_upload, part_number=i + 1, source=plan.source, range_=range_, @@ -630,9 +629,7 @@ async def _abort() -> None: return await asyncio.to_thread( self._sync_fs._abort_multipart_upload, - plan.destination.bucket, - cast(str, plan.destination.key), - cast(str, creation.result().upload_id), + creation.result(), plan.abort_params, ) @@ -648,7 +645,6 @@ async def _abort() -> None: # 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 @@ -662,8 +658,7 @@ async def _abort() -> None: completion = asyncio.ensure_future( asyncio.to_thread( self.core.complete_multipart_upload, - plan.destination, - upload_id, + multipart_upload, cast(list[S3MultipartUploadPart], parts), **plan.complete_params, ) diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index 53369ba5..43246884 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -662,7 +662,9 @@ def create_multipart_upload(self, path: S3Path, **params) -> S3MultipartUpload: the path take precedence over parameters of the same name. Returns: - The multipart upload. + The upload retaining the path's bucket and key, including an + access point alias or ARN, and the response's upload ID and + checksum configuration. Raises: ValueError: If the path has no key, or has a version ID, which a @@ -675,16 +677,29 @@ def create_multipart_upload(self, path: S3Path, **params) -> S3MultipartUpload: request: dict[str, Any] = {"Bucket": path.bucket, "Key": path.key} _logger.debug(f"Create multipart upload to {path.uri}.") response = self.call(self._client.create_multipart_upload, **{**params, **request}) - return S3MultipartUpload(response) + # S3 returns the bucket name even when creation uses an access point. + # Keep the request identity for every subsequent upload operation. + return S3MultipartUpload({**response, "Bucket": path.bucket, "Key": path.key}) + + @staticmethod + def _multipart_upload_request(upload: S3MultipartUpload) -> dict[str, Any]: + """Return the upload identity, rejecting incomplete upload objects.""" + if not upload.bucket: + raise ValueError("The multipart upload has no bucket.") + if not upload.key: + raise ValueError("The multipart upload has no key.") + if not upload.upload_id: + raise ValueError("The multipart upload has no upload ID.") + return {"Bucket": upload.bucket, "Key": upload.key, "UploadId": upload.upload_id} def upload_part( - self, path: S3Path, upload_id: str, part_number: int, body: bytes, **params + self, upload: S3MultipartUpload, part_number: int, body: bytes, **params ) -> S3MultipartUploadPart: """Upload a part of a multipart upload with UploadPart. Args: - path: The path of the object that the upload writes. - upload_id: The ID of the multipart upload. + upload: The multipart upload returned by creation. Its checksum + algorithm is passed to the SDK to calculate the part checksum. part_number: The number of the part, from 1. body: The data of the part. **params: Additional request parameters. The fields that the @@ -695,25 +710,22 @@ def upload_part( The uploaded part. Raises: - ValueError: If the path has no key. + ValueError: If the upload has no bucket, key, or upload ID. """ - if not path.key: - raise ValueError(f"The path has no key: {path.uri}.") request: dict[str, Any] = { - "Bucket": path.bucket, - "Key": path.key, - "UploadId": upload_id, + **self._multipart_upload_request(upload), "PartNumber": part_number, "Body": body, } - _logger.debug(f"Upload part of {upload_id} to {path.uri} as part {part_number}.") + if upload.checksum_algorithm is not None: + request["ChecksumAlgorithm"] = upload.checksum_algorithm + _logger.debug(f"Upload part of {upload.upload_id} as part {part_number}.") response = self.call(self._client.upload_part, **{**params, **request}) return S3MultipartUploadPart(part_number, response) def upload_part_copy( self, - path: S3Path, - upload_id: str, + upload: S3MultipartUpload, part_number: int, source: S3Path, range_: tuple[int, int] | None = None, @@ -722,8 +734,7 @@ def upload_part_copy( """Copy a part of a multipart upload from an object with UploadPartCopy. Args: - path: The path of the object that the upload writes. - upload_id: The ID of the multipart upload. + upload: The multipart upload returned by creation. part_number: The number of the part, from 1. source: The path of the object to copy, with the version ID to copy, if any. @@ -739,36 +750,42 @@ def upload_part_copy( The copied part. Raises: - ValueError: If the path or the source has no key. + ValueError: If the upload has no bucket, key, or upload ID, or the + source has no key. """ - if not path.key: - raise ValueError(f"The path has no key: {path.uri}.") + request = self._multipart_upload_request(upload) if not source.key: raise ValueError(f"The source has no key: {source.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] = { - "Bucket": path.bucket, - "Key": path.key, - "CopySource": copy_source, - "UploadId": upload_id, - "PartNumber": part_number, - } + request.update( + { + "CopySource": copy_source, + "PartNumber": part_number, + } + ) if range_: request.update({"CopySourceRange": f"bytes={range_[0]}-{range_[1] - 1}"}) - _logger.debug(f"Upload part copy from {source.uri} to {path.uri} as part {part_number}.") + _logger.debug( + f"Copy part from {source.uri} to upload {upload.upload_id} as part {part_number}." + ) response = self.call(self._client.upload_part_copy, **{**params, **request}) return S3MultipartUploadPart(part_number, response) def complete_multipart_upload( - self, path: S3Path, upload_id: str, parts: Sequence[S3MultipartUploadPart], **params + self, + upload: S3MultipartUpload, + parts: Sequence[S3MultipartUploadPart], + **params, ) -> S3CompleteMultipartUpload: """Complete a multipart upload with CompleteMultipartUpload. Args: - path: The path of the object that the upload writes. - upload_id: The ID of the multipart upload. + upload: The multipart upload returned by creation. Its checksum + algorithm selects the matching part checksum, and its checksum + type is sent when present. Without an algorithm, only the ETag + and part number are sent, even if the SDK added a part checksum. parts: The uploaded parts, in part-number order. **params: Additional request parameters. The fields that the other arguments set take precedence over parameters of the @@ -778,40 +795,39 @@ def complete_multipart_upload( The completed upload. Raises: - ValueError: If the path has no key. + ValueError: If the upload has no bucket, key, or upload ID. """ - if not path.key: - raise ValueError(f"The path has no key: {path.uri}.") - request: dict[str, Any] = { - "Bucket": path.bucket, - "Key": path.key, - "UploadId": upload_id, - "MultipartUpload": { - "Parts": [{"ETag": p.etag, "PartNumber": p.part_number} for p in parts] - }, + request = self._multipart_upload_request(upload) + part_fields = {"ETag", "PartNumber"} + if upload.checksum_algorithm is not None: + part_fields.add(f"Checksum{upload.checksum_algorithm}") + if upload.checksum_type is not None: + request["ChecksumType"] = upload.checksum_type + request["MultipartUpload"] = { + "Parts": [ + {key: value for key, value in part.to_api_repr().items() if key in part_fields} + for part in parts + ] } - _logger.debug(f"Complete multipart upload {upload_id} to {path.uri}.") + _logger.debug(f"Complete multipart upload {upload.upload_id}.") response = self.call(self._client.complete_multipart_upload, **{**params, **request}) return S3CompleteMultipartUpload(response) - def abort_multipart_upload(self, path: S3Path, upload_id: str, **params) -> None: + def abort_multipart_upload(self, upload: S3MultipartUpload, **params) -> None: """Abort a multipart upload with AbortMultipartUpload. Args: - path: The path of the object that the upload writes. - upload_id: The ID of the multipart upload. + upload: The multipart upload returned by creation or listing. **params: Additional request parameters. The fields that the other arguments set take precedence over parameters of the same name. Raises: - ValueError: If the path has no key. + ValueError: If the upload has no bucket, key, or upload ID. FileNotFoundError: If the upload does not exist, for example because it was completed or aborted. """ - if not path.key: - raise ValueError(f"The path has no key: {path.uri}.") - request: dict[str, Any] = {"Bucket": path.bucket, "Key": path.key, "UploadId": upload_id} + request = self._multipart_upload_request(upload) self.call(self._client.abort_multipart_upload, **{**params, **request}) def part_ranges(self, size: int, block_size: int) -> list[tuple[int, int]]: diff --git a/pyathena/filesystem/s3_object.py b/pyathena/filesystem/s3_object.py index 8745542a..1c9874ca 100644 --- a/pyathena/filesystem/s3_object.py +++ b/pyathena/filesystem/s3_object.py @@ -868,6 +868,7 @@ def __init__(self, response: dict[str, Any]) -> None: self._bucket_key_enabled = response.get("BucketKeyEnabled") self._request_charged = response.get("RequestCharged") self._checksum_algorithm = response.get("ChecksumAlgorithm") + self._checksum_type = response.get("ChecksumType") # The following fields are returned by the ListMultipartUploads API. self._initiated: datetime | None = response.get("Initiated") self._storage_class: str | None = response.get("StorageClass") @@ -941,6 +942,11 @@ def checksum_algorithm(self) -> str | None: """The ``ChecksumAlgorithm`` of the upload.""" return self._checksum_algorithm + @property + def checksum_type(self) -> str | None: + """The ``ChecksumType`` of the upload: COMPOSITE or FULL_OBJECT.""" + return self._checksum_type + @property def initiated(self) -> datetime | None: """The ``Initiated`` time of the upload, returned by ListMultipartUploads.""" @@ -993,20 +999,21 @@ def __init__(self, part_number: int, response: dict[str, Any]) -> None: self._part_number = part_number self._copy_source_version_id: str | None = response.get("CopySourceVersionId") copy_part_result = response.get("CopyPartResult") - if copy_part_result: - self._last_modified: datetime | None = copy_part_result.get("LastModified") - self._etag: str | None = copy_part_result.get("ETag") - self._checksum_crc32: str | None = copy_part_result.get("ChecksumCRC32") - self._checksum_crc32c: str | None = copy_part_result.get("ChecksumCRC32C") - self._checksum_sha1: str | None = copy_part_result.get("ChecksumSHA1") - self._checksum_sha256: str | None = copy_part_result.get("ChecksumSHA256") - else: - self._last_modified = None - self._etag = response.get("ETag") - self._checksum_crc32 = response.get("ChecksumCRC32") - self._checksum_crc32c = response.get("ChecksumCRC32C") - self._checksum_sha1 = response.get("ChecksumSHA1") - self._checksum_sha256 = response.get("ChecksumSHA256") + self._last_modified: datetime | None = ( + copy_part_result.get("LastModified") if copy_part_result else None + ) + part_result = copy_part_result or response + self._etag: str | None = part_result.get("ETag") + self._checksum_crc32: str | None = part_result.get("ChecksumCRC32") + self._checksum_crc32c: str | None = part_result.get("ChecksumCRC32C") + self._checksum_crc64nvme: str | None = part_result.get("ChecksumCRC64NVME") + self._checksum_sha1: str | None = part_result.get("ChecksumSHA1") + self._checksum_sha256: str | None = part_result.get("ChecksumSHA256") + self._checksum_sha512: str | None = part_result.get("ChecksumSHA512") + self._checksum_md5: str | None = part_result.get("ChecksumMD5") + self._checksum_xxhash64: str | None = part_result.get("ChecksumXXHASH64") + self._checksum_xxhash3: str | None = part_result.get("ChecksumXXHASH3") + self._checksum_xxhash128: str | None = part_result.get("ChecksumXXHASH128") self._server_side_encryption: str | None = response.get("ServerSideEncryption") self._sse_customer_algorithm: str | None = response.get("SSECustomerAlgorithm") self._sse_customer_key_md5: str | None = response.get("SSECustomerKeyMD5") @@ -1054,6 +1061,36 @@ def checksum_sha256(self) -> str | None: """The ``ChecksumSHA256`` of the part.""" return self._checksum_sha256 + @property + def checksum_crc64nvme(self) -> str | None: + """The ``ChecksumCRC64NVME`` of the part.""" + return self._checksum_crc64nvme + + @property + def checksum_sha512(self) -> str | None: + """The ``ChecksumSHA512`` of the part.""" + return self._checksum_sha512 + + @property + def checksum_md5(self) -> str | None: + """The ``ChecksumMD5`` of the part.""" + return self._checksum_md5 + + @property + def checksum_xxhash64(self) -> str | None: + """The ``ChecksumXXHASH64`` of the part.""" + return self._checksum_xxhash64 + + @property + def checksum_xxhash3(self) -> str | None: + """The ``ChecksumXXHASH3`` of the part.""" + return self._checksum_xxhash3 + + @property + def checksum_xxhash128(self) -> str | None: + """The ``ChecksumXXHASH128`` of the part.""" + return self._checksum_xxhash128 + @property def server_side_encryption(self) -> str | None: """The ``ServerSideEncryption`` algorithm of the part.""" @@ -1089,16 +1126,23 @@ def to_api_repr(self) -> dict[str, Any]: Returns: Dictionary with the ``ETag``, checksum and ``PartNumber`` fields of - the part. + the part, omitting fields whose value is None. """ - return { + fields = { "ETag": self.etag, "ChecksumCRC32": self.checksum_crc32, "ChecksumCRC32C": self.checksum_crc32c, + "ChecksumCRC64NVME": self.checksum_crc64nvme, "ChecksumSHA1": self.checksum_sha1, "ChecksumSHA256": self.checksum_sha256, + "ChecksumSHA512": self.checksum_sha512, + "ChecksumMD5": self.checksum_md5, + "ChecksumXXHASH64": self.checksum_xxhash64, + "ChecksumXXHASH3": self.checksum_xxhash3, + "ChecksumXXHASH128": self.checksum_xxhash128, "PartNumber": self.part_number, } + return {key: value for key, value in fields.items() if value is not None} class S3CompleteMultipartUpload: diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 17bbb858..9bdd00fd 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -16,12 +16,14 @@ import urllib.parse import urllib.request import uuid +from base64 import b64encode from concurrent.futures import Future, ThreadPoolExecutor, wait from datetime import UTC, datetime from itertools import chain from pathlib import Path from types import SimpleNamespace from unittest import mock +from zlib import crc32 import boto3 import botocore.exceptions @@ -39,7 +41,7 @@ from pyathena.filesystem.s3_core import S3Core, S3DeleteBatch from pyathena.filesystem.s3_errors import S3ClientError from pyathena.filesystem.s3_executor import S3AioExecutor, S3ThreadPoolExecutor -from pyathena.filesystem.s3_object import S3Object, S3ObjectType, S3StorageClass +from pyathena.filesystem.s3_object import S3MultipartUpload, S3Object, S3ObjectType, S3StorageClass from pyathena.filesystem.s3_path import S3Path from pyathena.util import RetryConfig from tests import ENV @@ -1186,7 +1188,9 @@ def test_pipe_file_buffered_s3_parameters(self, transaction): key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True ) fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -1281,7 +1285,9 @@ def call(method, **request): name, ) raise S3ClientError(error).os_error from error - return {"UploadId": "uploadid", "ETag": '"e"'} + if name == "create_multipart_upload": + return {"Bucket": request["Bucket"], "Key": request["Key"], "UploadId": "uploadid"} + return {"ETag": '"e"'} fs._call.side_effect = call return requests @@ -1434,7 +1440,12 @@ def call(method, **request): requests.append((name, request)) if name == "upload_part" and fail: raise OSError("upload failed") - return {"UploadId": "uploadid", "ETag": '"e"'} + return { + "Bucket": request["Bucket"], + "Key": request["Key"], + "UploadId": "uploadid", + "ETag": '"e"', + } fs._call.side_effect = call block_size = fs.core.MULTIPART_UPLOAD_MIN_PART_SIZE @@ -1465,6 +1476,7 @@ def test_finish_multipart_upload_request_parameters(self): # GH-946: the completion and the abort receive the parameters of the # upload that they accept. fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() kwargs = { "ContentType": "text/csv", @@ -1474,23 +1486,18 @@ def test_finish_multipart_upload_request_parameters(self): part: Future[SimpleNamespace] = Future() part.set_result(SimpleNamespace(etag='"e1"', part_number=1)) - fs._finish_multipart_upload( - bucket="bucket", key="key", upload_id="uploadid", futures=[part], request_kwargs=kwargs - ) + fs._finish_multipart_upload(upload=upload, futures=[part], request_kwargs=kwargs) failed: Future[SimpleNamespace] = Future() failed.set_exception(RuntimeError("upload failed")) with pytest.raises(RuntimeError, match="upload failed"): fs._finish_multipart_upload( - bucket="bucket", - key="key", - upload_id="uploadid", + upload=upload, futures=[failed], request_kwargs=kwargs, ) fs.core.complete_multipart_upload.assert_called_once_with( - S3Path("bucket", "key"), - "uploadid", + upload, [part.result()], RequestPayer="requester", SSECustomerAlgorithm="AES256", @@ -1859,7 +1866,9 @@ def test_copy_object_with_multipart_upload_request_parameters(self): # parameters of the copy that they accept. fs = self._make_fs() fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) fs.core.upload_part_copy = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -1979,7 +1988,9 @@ def test_copy_object_with_multipart_upload_head_object_size(self): "VersionId": "null", } fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) fs.core.upload_part_copy = mock.MagicMock() fs._finish_multipart_upload = mock.MagicMock() @@ -2039,7 +2050,7 @@ def test_copy_object_with_multipart_upload_replace_directives(self): ) stubber.add_response( "create_multipart_upload", - {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "UploadId": "u"}, {"Bucket": "bucket", "Key": "dst", "ContentType": "text/plain", "Tagging": "a=1"}, ) for _ in (1, 2): @@ -2086,7 +2097,7 @@ def test_copy_object_with_multipart_upload_sse_c_source(self): ) stubber.add_response( "create_multipart_upload", - {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "UploadId": "u"}, {"Bucket": "bucket", "Key": "dst", "ContentType": "text/csv", "Metadata": {}}, ) for _ in (1, 2): @@ -2109,7 +2120,7 @@ def test_copy_object_with_multipart_upload_directory_bucket_source(self): ) stubber.add_response( "create_multipart_upload", - {"UploadId": "u"}, + {"Bucket": "bucket", "Key": "dst", "UploadId": "u"}, {"Bucket": "bucket", "Key": "dst", "ContentType": "text/csv", "Metadata": {}}, ) for _ in (1, 2): @@ -2147,7 +2158,9 @@ def test_pipe_file_trailing_slash(self, intrans, size): fs._transaction = None fs._put_object = mock.MagicMock() fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -2170,7 +2183,9 @@ def test_pipe_file_memoryview_routed_by_bytes(self): fs.default_cache_type = "bytes" fs._put_object = mock.MagicMock() fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -2199,7 +2214,9 @@ def test_pipe_file_buffered_non_contiguous_memoryview(self): fs = self._make_fs() fs.default_cache_type = "bytes" fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -2273,7 +2290,9 @@ def test_pipe_file_compression_multipart(self, intrans): fs.default_cache_type = "bytes" fs._transaction = None fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -2347,7 +2366,9 @@ def test_pipe_file_failed_write_aborts_multipart_upload(self): fs = self._make_fs() fs.default_cache_type = "bytes" fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs._finish_multipart_upload = mock.MagicMock() executor = mock.MagicMock() @@ -2402,7 +2423,9 @@ def test_put_file_failed_write_aborts_multipart_upload(self, tmp_path, intrans): fs.default_cache_type = "bytes" fs._transaction = None fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) fs.core.upload_part = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -3216,6 +3239,7 @@ def test_open_directory(self, cache_type): def test_finish_multipart_upload(self): fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() futures = [] for part_number in (1, 2): @@ -3223,11 +3247,10 @@ def test_finish_multipart_upload(self): future.set_result(SimpleNamespace(etag=f'"e{part_number}"', part_number=part_number)) futures.append(future) - fs._finish_multipart_upload( - bucket="bucket", key="key", upload_id="uploadid", futures=futures - ) + fs._finish_multipart_upload(upload=upload, futures=futures) fs.core.complete_multipart_upload.assert_called_once_with( - S3Path("bucket", "key"), "uploadid", [f.result() for f in futures] + upload, + [f.result() for f in futures], ) fs._call.assert_not_called() @@ -3236,14 +3259,13 @@ def test_finish_multipart_upload(self): @pytest.mark.parametrize("error", [RuntimeError, KeyboardInterrupt]) def test_finish_multipart_upload_aborts_on_failure(self, error): fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() future: Future[SimpleNamespace] = Future() future.set_exception(error("upload failed")) with pytest.raises(error, match="upload failed"): - fs._finish_multipart_upload( - bucket="bucket", key="key", upload_id="uploadid", futures=[future] - ) + fs._finish_multipart_upload(upload=upload, futures=[future]) fs.core.complete_multipart_upload.assert_not_called() fs._call.assert_called_once_with( fs._client.abort_multipart_upload, @@ -3256,6 +3278,7 @@ def test_finish_multipart_upload_without_abort(self): # A caller that aborts the upload itself, as S3File.commit() does, # gets the original error with the parts and the upload left alone. fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() failed: Future[SimpleNamespace] = Future() failed.set_exception(RuntimeError("upload failed")) @@ -3263,9 +3286,7 @@ def test_finish_multipart_upload_without_abort(self): with pytest.raises(RuntimeError, match="upload failed"): fs._finish_multipart_upload( - bucket="bucket", - key="key", - upload_id="uploadid", + upload=upload, futures=[failed, pending], abort=False, ) @@ -3274,6 +3295,7 @@ def test_finish_multipart_upload_without_abort(self): def test_finish_multipart_upload_abort_failure_does_not_mask_the_original_error(self, caplog): fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() # The abort is sent through the core, whose call is the same mock. fs._call.side_effect = RuntimeError("abort failed") @@ -3282,9 +3304,7 @@ def test_finish_multipart_upload_abort_failure_does_not_mask_the_original_error( # The abort failure is logged, and the original error propagates. with pytest.raises(RuntimeError, match="upload failed"): - fs._finish_multipart_upload( - bucket="bucket", key="key", upload_id="uploadid", futures=[future] - ) + fs._finish_multipart_upload(upload=upload, futures=[future]) fs._call.assert_called_once_with( fs._client.abort_multipart_upload, Bucket="bucket", Key="key", UploadId="uploadid" ) @@ -3295,6 +3315,7 @@ def test_finish_multipart_upload_waits_for_running_parts(self): # may be stored after the abort, so the abort waits for it. The # parts that have not started are cancelled. fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() events = [] fs._call.side_effect = lambda *args, **kwargs: events.append("abort") @@ -3323,9 +3344,7 @@ def wait_parts(futures): started.wait(5) with pytest.raises(RuntimeError, match="upload failed"): fs._finish_multipart_upload( - bucket="bucket", - key="key", - upload_id="uploadid", + upload=upload, futures=[failed, running, pending], ) @@ -3338,6 +3357,7 @@ def test_finish_multipart_upload_does_not_wait_for_cancelled_parts(self): # acknowledge its cancellation, e.g., an event loop blocked by the # caller. fs = self._make_fs() + upload = S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"}) fs.core.complete_multipart_upload = mock.MagicMock() failed: Future[SimpleNamespace] = Future() failed.set_exception(RuntimeError("upload failed")) @@ -3347,9 +3367,7 @@ def test_finish_multipart_upload_does_not_wait_for_cancelled_parts(self): def finish(): try: fs._finish_multipart_upload( - bucket="bucket", - key="key", - upload_id="uploadid", + upload=upload, futures=[failed, never_started], ) except RuntimeError as e: @@ -3371,7 +3389,9 @@ def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): # as one part larger than 5 GiB. fs = self._make_fs() fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) fs.core.upload_part_copy = mock.MagicMock() fs._finish_multipart_upload = mock.MagicMock() @@ -3415,7 +3435,7 @@ 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") + return S3MultipartUpload({"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"}) fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) fs.core.upload_part_copy = mock.MagicMock() @@ -3473,7 +3493,14 @@ def interrupt(): signal.signal(signal.SIGINT, previous_handler) interrupted.set() - fs._abort_multipart_upload.assert_called_once_with("bucket", "dst", "uploadid", {}) + fs._abort_multipart_upload.assert_called_once() + upload, params = fs._abort_multipart_upload.call_args.args + assert (upload.bucket, upload.key, upload.upload_id, params) == ( + "bucket", + "dst", + "uploadid", + {}, + ) fs.core.upload_part_copy.assert_not_called() @pytest.mark.parametrize( @@ -3953,12 +3980,375 @@ def call(method, **request): fs.clear_multipart_uploads(path) assert sorted(aborted) == expected + @pytest.mark.parametrize( + ("code", "exception"), + [ + (None, None), + ("NoSuchUpload", None), + ("NoSuchBucket", FileNotFoundError), + ("AccessDenied", PermissionError), + ("InternalError", OSError), + ], + ) + def test_clear_multipart_uploads_race(self, code, exception): + fs = S3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + retry_config=RetryConfig(attempt=1), + skip_instance_cache=True, + ) + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "list_multipart_uploads", + { + "Uploads": [ + {"Key": "prefix/gone", "UploadId": "gone"}, + { + "Key": "prefix/pending", + "UploadId": "pending", + "ChecksumAlgorithm": "CRC32", + "ChecksumType": "FULL_OBJECT", + }, + ], + "IsTruncated": False, + }, + {"Bucket": "bucket", "Prefix": "prefix/"}, + ) + request = {"Bucket": "bucket", "Key": "prefix/gone", "UploadId": "gone"} + if code: + stubber.add_client_error( + "abort_multipart_upload", + service_error_code=code, + http_status_code=404 if code.startswith("NoSuch") else 500, + expected_params=request, + ) + else: + stubber.add_response("abort_multipart_upload", {}, request) + stubber.add_response( + "abort_multipart_upload", + {}, + {"Bucket": "bucket", "Key": "prefix/pending", "UploadId": "pending"}, + ) + expected = pytest.raises(exception) if exception else contextlib.nullcontext() + with expected: + fs.clear_multipart_uploads("s3://bucket/prefix/") + stubber.assert_no_pending_responses() + + def test_clear_multipart_uploads_empty(self): + fs = S3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "list_multipart_uploads", + {"Uploads": [], "IsTruncated": False}, + {"Bucket": "bucket", "Prefix": "prefix/"}, + ) + fs.clear_multipart_uploads("s3://bucket/prefix/") + stubber.assert_no_pending_responses() + + def test_clear_multipart_uploads_checks_all_results(self, monkeypatch): + fs = S3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) + fs.list_multipart_uploads = mock.MagicMock( + return_value=[ + S3MultipartUpload({"Bucket": "bucket", "Key": f"prefix/{n}", "UploadId": str(n)}) + for n in range(3) + ] + ) + error = PermissionError("denied") + futures = [mock.Mock(), mock.Mock(), mock.Mock()] + futures[0].result.side_effect = error + futures[2].result.side_effect = FileNotFoundError("unclassified missing resource") + executor = mock.MagicMock() + executor.__enter__.return_value.submit.side_effect = futures + monkeypatch.setattr(fs, "_create_executor", mock.Mock(return_value=executor)) + monkeypatch.setattr("pyathena.filesystem.s3.as_completed", lambda pending: iter(pending)) + with pytest.raises(PermissionError) as raised: + fs.clear_multipart_uploads("s3://bucket/prefix/") + assert raised.value is error + for future in futures: + future.result.assert_called_once_with() + + def test_clear_multipart_uploads_preserves_unclassified_file_not_found(self): + fs = S3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) + fs.list_multipart_uploads = mock.MagicMock( + return_value=[ + S3MultipartUpload({"Bucket": "bucket", "Key": "prefix/key", "UploadId": "u"}) + ] + ) + error = FileNotFoundError("unclassified missing resource") + fs.core.abort_multipart_upload = mock.MagicMock(side_effect=error) + with pytest.raises(FileNotFoundError) as raised: + fs.clear_multipart_uploads("s3://bucket/prefix/") + assert raised.value is error + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_multipart_copy_uses_creation_algorithm(self, algorithm, checksum_type): + fs = S3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + block_size = 5 * 2**30 + size = 2 * block_size + checksum_kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + checksum_kwargs["ChecksumType"] = checksum_type + expected_parts = [] + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "head_object", {"ContentLength": size}, {"Bucket": "bucket", "Key": "src"} + ) + stubber.add_response( + "create_multipart_upload", + {"Bucket": "bucket", "Key": "dst", "UploadId": "u", **checksum_kwargs}, + {"Bucket": "bucket", "Key": "dst", **checksum_kwargs}, + ) + for number, range_ in ( + (1, f"bytes=0-{block_size - 1}"), + (2, f"bytes={block_size}-{size - 1}"), + ): + result = {"ETag": f'"p{number}"', f"Checksum{algorithm or 'CRC32'}": "checksum"} + stubber.add_response( + "upload_part_copy", + {"CopyPartResult": result}, + { + "Bucket": "bucket", + "Key": "dst", + "UploadId": "u", + "PartNumber": number, + "CopySource": {"Bucket": "bucket", "Key": "src"}, + "CopySourceRange": range_, + }, + ) + expected_parts.append( + { + "ETag": result["ETag"], + "PartNumber": number, + **({f"Checksum{algorithm}": "checksum"} if algorithm else {}), + } + ) + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"done"'}, + { + "Bucket": "bucket", + "Key": "dst", + "UploadId": "u", + "MultipartUpload": {"Parts": expected_parts}, + **({"ChecksumType": checksum_type} if checksum_type else {}), + }, + ) + kwargs = { + "source": S3Path("bucket", "src"), + "destination": S3Path("bucket", "dst"), + "block_size": block_size, + "MetadataDirective": "REPLACE", + "TaggingDirective": "REPLACE", + "AnnotationDirective": "EXCLUDE", + **checksum_kwargs, + } + fs._copy_object_with_multipart_upload(**kwargs) + stubber.assert_no_pending_responses() + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): request.param = {} return S3FileSystem(connect(), **request.param) + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_open_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_open_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + with fs.open(path, "wb", block_size=block_size, s3_additional_kwargs=kwargs) as file: + file.write(data) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_put_file_multipart_with_checksum(self, fs, tmp_path, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_put_file_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + lpath = tmp_path / "data" + lpath.write_bytes(data) + fs.put_file(str(lpath), path, block_size=block_size, s3_additional_kwargs=kwargs) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_pipe_file_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_pipe_file_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + fs.pipe_file(path, data, block_size=block_size, s3_additional_kwargs=kwargs) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_append_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_append_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + original = b"y" * block_size + fs.pipe_file(path, original) + with fs.open(path, "ab", block_size=block_size, s3_additional_kwargs=kwargs) as file: + file.write(data) + assert fs.cat_file(path) == original + data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", "COMPOSITE"), ("CRC32", "FULL_OBJECT")], + ) + @pytest.mark.parametrize("copy", [False, True]) + def test_core_multipart_upload_with_checksum(self, fs, algorithm, checksum_type, copy): + prefix = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_core_multipart_with_checksum/{uuid.uuid4()}/" + ) + destination = S3Path.parse(f"{prefix}destination") + source = S3Path.parse(f"{prefix}source") + data = b"x" * S3Core.MULTIPART_UPLOAD_MIN_PART_SIZE + kwargs = ( + {"ChecksumAlgorithm": algorithm, "ChecksumType": checksum_type} if algorithm else {} + ) + try: + if copy: + fs.pipe_file(source.uri, data) + upload = fs.core.create_multipart_upload(destination, **kwargs) + if algorithm: + assert upload.checksum_algorithm == algorithm + assert upload.checksum_type == checksum_type + listed = fs.list_multipart_uploads(destination.uri) + assert len(listed) == 1 + assert listed[0].upload_id == upload.upload_id + assert listed[0].checksum_algorithm == upload.checksum_algorithm + assert listed[0].checksum_type == upload.checksum_type + if copy: + first = fs.core.upload_part_copy(upload, 1, source) + else: + first = fs.core.upload_part(upload, 1, data) + last = fs.core.upload_part(upload, 2, b"end") + fs.core.complete_multipart_upload(upload, [first, last]) + if checksum_type == "FULL_OBJECT": + metadata = fs.core.call( + "head_object", Bucket=upload.bucket, Key=upload.key, ChecksumMode="ENABLED" + ) + assert metadata["ChecksumType"] == "FULL_OBJECT" + assert ( + metadata["ChecksumCRC32"] + == b64encode(crc32(data + b"end").to_bytes(4, "big")).decode() + ) + fs.invalidate_cache(destination.uri) + assert fs.cat_file(destination.uri) == data + b"end" + assert fs.list_multipart_uploads(destination.uri) == [] + finally: + fs.clear_multipart_uploads(prefix) + for path in (source, destination): + if fs.exists(path.uri): + fs.rm(path.uri) + + def test_clear_multipart_uploads_after_listed_upload_is_aborted(self, fs): + prefix = ( + f"{ENV.s3_staging_key}{ENV.schema}/filesystem/test_clear_multipart_race/{uuid.uuid4()}/" + ) + path = f"s3://{ENV.s3_staging_bucket}/{prefix}" + gone_path = S3Path(ENV.s3_staging_bucket, f"{prefix}gone") + gone = fs.core.create_multipart_upload(gone_path) + try: + fs.core.create_multipart_upload(S3Path(ENV.s3_staging_bucket, f"{prefix}pending")) + list_uploads = fs.list_multipart_uploads + + def list_then_abort(path): + uploads = list_uploads(path) + assert len(uploads) == 2 + assert any(upload.upload_id == gone.upload_id for upload in uploads) + fs.core.abort_multipart_upload(gone) + return uploads + + with mock.patch.object( + fs, "list_multipart_uploads", side_effect=list_then_abort + ) as list_mock: + fs.clear_multipart_uploads(path) + list_mock.assert_called_once_with(path) + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + @pytest.mark.parametrize( ("fs", "start", "end", "target_data"), list( @@ -5120,7 +5510,11 @@ def test_error_translation_permission_error(self): with pytest.raises(PermissionError): anon_fs.info(f"s3://{ENV.s3_staging_bucket}/{ENV.s3_filesystem_test_file_key}") - def test_list_and_clear_multipart_uploads(self, fs): + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", "COMPOSITE"), ("CRC32", "FULL_OBJECT")], + ) + def test_list_and_clear_multipart_uploads(self, fs, algorithm, checksum_type): # Scope the list/clear to a unique prefix so that parallel test # workers' in-flight multipart uploads in the shared bucket are # not aborted. @@ -5130,7 +5524,10 @@ def test_list_and_clear_multipart_uploads(self, fs): ) prefix_path = f"s3://{bucket}/{prefix}" key = f"{prefix}/file" - upload = fs.core.create_multipart_upload(S3Path(bucket, key)) + kwargs = ( + {"ChecksumAlgorithm": algorithm, "ChecksumType": checksum_type} if algorithm else {} + ) + upload = fs.core.create_multipart_upload(S3Path(bucket, key), **kwargs) # A sibling key that starts with the same characters as the prefix. sibling = fs.core.create_multipart_upload(S3Path(bucket, f"{prefix}2/file")) try: @@ -5140,6 +5537,8 @@ def test_list_and_clear_multipart_uploads(self, fs): assert listed.bucket == bucket assert listed.key == key assert listed.initiated + assert listed.checksum_algorithm == upload.checksum_algorithm + assert listed.checksum_type == upload.checksum_type assert not any(u.upload_id == sibling.upload_id for u in uploads) fs.clear_multipart_uploads(prefix_path) @@ -5148,6 +5547,7 @@ def test_list_and_clear_multipart_uploads(self, fs): uploads = fs.list_multipart_uploads(f"{prefix_path}2") assert any(u.upload_id == sibling.upload_id for u in uploads) finally: + fs.clear_multipart_uploads(prefix_path) fs.clear_multipart_uploads(f"{prefix_path}2") def test_object_version_info(self, fs): @@ -5351,7 +5751,9 @@ def _make_multipart_write_file(data: bytes, autocommit: bool): file.blocksize = 4 file.fs.core.MULTIPART_UPLOAD_MIN_PART_SIZE = 4 file.fs.core.MULTIPART_UPLOAD_MAX_PART_SIZE = 8 - file.multipart_upload = SimpleNamespace(upload_id="uploadid") + file.multipart_upload = S3MultipartUpload( + {"Bucket": "bucket", "Key": "key.txt", "UploadId": "uploadid"} + ) file._executor = ThreadPoolExecutor(max_workers=1) file.fs.core.upload_part.side_effect = lambda **kw: SimpleNamespace( etag=f'"e{kw["part_number"]}"', part_number=kw["part_number"] @@ -5374,7 +5776,9 @@ def _make_append_fs(existing: bytes): key="key.txt", ) fs.cat_file.return_value = existing - fs.core.create_multipart_upload.return_value = SimpleNamespace(upload_id="uploadid") + fs.core.create_multipart_upload.return_value = S3MultipartUpload( + {"Bucket": "bucket", "Key": "key.txt", "UploadId": "uploadid"} + ) def part(**kw): return SimpleNamespace(etag=f'"e{kw["part_number"]}"', part_number=kw["part_number"]) @@ -5672,16 +6076,25 @@ def test_multipart_write_request_parameters(self): } assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == kwargs - def test_multipart_write_keyword_named_as_argument(self): + @pytest.mark.parametrize("parameter", ["key", "upload"]) + def test_multipart_write_keyword_named_as_argument(self, parameter): # A keyword parameter of the file named like a helper argument does # not break the completion, which takes the parameters as a mapping. fs = self._make_append_fs(b"") - with S3File(fs, "s3://bucket/key.txt", mode="wb", block_size=4, key="other") as f: + with S3File( + fs, "s3://bucket/key.txt", mode="wb", block_size=4, **{parameter: "other"} + ) as f: f.write(b"x" * 8) - assert fs._finish_multipart_upload.call_args.kwargs["key"] == "key.txt" - assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == {"key": "other"} + fs.core.create_multipart_upload.assert_called_once_with(S3Path("bucket", "key.txt")) + assert ( + fs._finish_multipart_upload.call_args.kwargs["upload"] + is fs.core.create_multipart_upload.return_value + ) + assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == { + parameter: "other" + } def test_append_discard(self): # Rolling back an append aborts its multipart upload without the @@ -5820,7 +6233,9 @@ def test_discard_waits_for_running_parts(self): # may be stored after the abort, so the abort waits for it. The # parts that have not started are cancelled. file = self._make_write_file(b"", autocommit=False) - file.multipart_upload = SimpleNamespace(upload_id="uploadid") + file.multipart_upload = S3MultipartUpload( + {"Bucket": "bucket", "Key": "key.txt", "UploadId": "uploadid"} + ) events = [] file.fs._call.side_effect = lambda *args, **kwargs: events.append("abort") started = threading.Event() @@ -5915,7 +6330,9 @@ def test_discard_on_event_loop_thread(self): # waited for, so a rollback on the thread of the event loop that # would run them does not block. file = self._make_write_file(b"", autocommit=False) - file.multipart_upload = SimpleNamespace(upload_id="uploadid") + file.multipart_upload = S3MultipartUpload( + {"Bucket": "bucket", "Key": "key.txt", "UploadId": "uploadid"} + ) parts = [] async def rollback(): diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 9663602c..19f442c5 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -24,12 +24,14 @@ from pyathena.filesystem.s3_async import AioS3File, AioS3FileSystem from pyathena.filesystem.s3_core import S3Core from pyathena.filesystem.s3_object import ( + S3MultipartUpload, S3MultipartUploadPart, S3Object, S3ObjectType, S3StorageClass, ) from pyathena.filesystem.s3_path import S3Path +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect from tests.pyathena.util import ( @@ -159,7 +161,9 @@ async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): ) sync_fs = fs._sync_fs sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) sync_fs.core.upload_part_copy = mock.MagicMock( side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) @@ -307,7 +311,9 @@ async def test_copy_object_with_multipart_upload_waits_for_running_parts(self): fs = AioS3FileSystem(connection=mock.MagicMock(), max_workers=2, skip_instance_cache=True) sync_fs = fs._sync_fs sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) events = [] failed = threading.Event() @@ -356,7 +362,9 @@ async def test_copy_object_with_multipart_upload_cancelled(self, cancellations): fs = AioS3FileSystem(connection=mock.MagicMock(), max_workers=2, skip_instance_cache=True) sync_fs = fs._sync_fs sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) events = [] lock = threading.Lock() @@ -432,7 +440,9 @@ async def test_copy_object_with_multipart_upload_cancelled_completion(self, comp fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) sync_fs = fs._sync_fs sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"} + ) ) events = [] started = threading.Event() @@ -506,7 +516,7 @@ def create_multipart_upload(*args, **kw): events.append("create") if creation_fails: raise OSError("creation failed") - return SimpleNamespace(upload_id="uploadid") + return S3MultipartUpload({"Bucket": "bucket", "Key": "dst", "UploadId": "uploadid"}) sync_fs.core.create_multipart_upload = mock.MagicMock(side_effect=create_multipart_upload) sync_fs.core.upload_part_copy = mock.MagicMock() @@ -514,7 +524,7 @@ def create_multipart_upload(*args, **kw): 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])) + side_effect=lambda upload, params: events.append(("abort", upload.upload_id)) ) task = asyncio.ensure_future( @@ -1045,7 +1055,9 @@ async def test_cp_file_multipart_parameters(self, size): sync_fs = fs._sync_fs sync_fs.core.copy_object = mock.MagicMock() sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) running = [] concurrency = [] @@ -1119,12 +1131,287 @@ def test_internal_file_system_not_cached(self): for fs in (fs1, fs2, fs3): assert fs._sync_fs not in S3FileSystem._cache.values() + @pytest.mark.parametrize( + ("code", "exception"), + [ + (None, None), + ("NoSuchUpload", None), + ("NoSuchBucket", FileNotFoundError), + ("AccessDenied", PermissionError), + ("InternalError", OSError), + ], + ) + def test_clear_multipart_uploads_race(self, code, exception): + fs = AioS3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + retry_config=RetryConfig(attempt=1), + skip_instance_cache=True, + ) + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "list_multipart_uploads", + { + "Uploads": [ + {"Key": "prefix/gone", "UploadId": "gone"}, + { + "Key": "prefix/pending", + "UploadId": "pending", + "ChecksumAlgorithm": "CRC32", + "ChecksumType": "FULL_OBJECT", + }, + ], + "IsTruncated": False, + }, + {"Bucket": "bucket", "Prefix": "prefix/"}, + ) + request = {"Bucket": "bucket", "Key": "prefix/gone", "UploadId": "gone"} + if code: + stubber.add_client_error( + "abort_multipart_upload", + service_error_code=code, + http_status_code=404 if code.startswith("NoSuch") else 500, + expected_params=request, + ) + else: + stubber.add_response("abort_multipart_upload", {}, request) + stubber.add_response( + "abort_multipart_upload", + {}, + {"Bucket": "bucket", "Key": "prefix/pending", "UploadId": "pending"}, + ) + expected = pytest.raises(exception) if exception else contextlib.nullcontext() + with expected: + fs.clear_multipart_uploads("s3://bucket/prefix/") + stubber.assert_no_pending_responses() + + def test_clear_multipart_uploads_empty(self): + fs = AioS3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "list_multipart_uploads", + {"Uploads": [], "IsTruncated": False}, + {"Bucket": "bucket", "Prefix": "prefix/"}, + ) + fs.clear_multipart_uploads("s3://bucket/prefix/") + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + @pytest.mark.asyncio + async def test_multipart_copy_uses_creation_algorithm(self, algorithm, checksum_type): + fs = AioS3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + max_workers=1, + skip_instance_cache=True, + ) + block_size = 5 * 2**30 + size = 2 * block_size + checksum_kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + checksum_kwargs["ChecksumType"] = checksum_type + expected_parts = [] + with Stubber(fs.core.client) as stubber: + stubber.add_response( + "head_object", {"ContentLength": size}, {"Bucket": "bucket", "Key": "src"} + ) + stubber.add_response( + "create_multipart_upload", + {"Bucket": "bucket", "Key": "dst", "UploadId": "u", **checksum_kwargs}, + {"Bucket": "bucket", "Key": "dst", **checksum_kwargs}, + ) + for number, range_ in ( + (1, f"bytes=0-{block_size - 1}"), + (2, f"bytes={block_size}-{size - 1}"), + ): + result = {"ETag": f'"p{number}"', f"Checksum{algorithm or 'CRC32'}": "checksum"} + stubber.add_response( + "upload_part_copy", + {"CopyPartResult": result}, + { + "Bucket": "bucket", + "Key": "dst", + "UploadId": "u", + "PartNumber": number, + "CopySource": {"Bucket": "bucket", "Key": "src"}, + "CopySourceRange": range_, + }, + ) + expected_parts.append( + { + "ETag": result["ETag"], + "PartNumber": number, + **({f"Checksum{algorithm}": "checksum"} if algorithm else {}), + } + ) + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"done"'}, + { + "Bucket": "bucket", + "Key": "dst", + "UploadId": "u", + "MultipartUpload": {"Parts": expected_parts}, + **({"ChecksumType": checksum_type} if checksum_type else {}), + }, + ) + kwargs = { + "source": S3Path("bucket", "src"), + "destination": S3Path("bucket", "dst"), + "block_size": block_size, + "MetadataDirective": "REPLACE", + "TaggingDirective": "REPLACE", + "AnnotationDirective": "EXCLUDE", + **checksum_kwargs, + } + await fs._copy_object_with_multipart_upload(**kwargs) + stubber.assert_no_pending_responses() + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): request.param = {} return AioS3FileSystem(connection=connect(), **request.param) + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_open_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_async_open_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + with fs.open(path, "wb", block_size=block_size, s3_additional_kwargs=kwargs) as file: + file.write(data) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_put_file_multipart_with_checksum(self, fs, tmp_path, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_async_put_file_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + lpath = tmp_path / "data" + lpath.write_bytes(data) + fs.put_file(str(lpath), path, block_size=block_size, s3_additional_kwargs=kwargs) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_pipe_file_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_async_pipe_file_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + fs.pipe_file(path, data, block_size=block_size, s3_additional_kwargs=kwargs) + assert fs.cat_file(path) == data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", None), ("CRC32", None), ("CRC32", "FULL_OBJECT")], + ) + def test_append_multipart_with_checksum(self, fs, algorithm, checksum_type): + block_size = 5 * 2**20 + data = b"x" * (block_size + 1) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_async_append_multipart_with_checksum/{uuid.uuid4()}" + ) + kwargs = {"ChecksumAlgorithm": algorithm} if algorithm else {} + if checksum_type: + kwargs["ChecksumType"] = checksum_type + try: + original = b"y" * block_size + fs.pipe_file(path, original) + with fs.open(path, "ab", block_size=block_size, s3_additional_kwargs=kwargs) as file: + file.write(data) + assert fs.cat_file(path) == original + data + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + if fs.exists(path): + fs.rm(path) + + def test_clear_multipart_uploads_after_listed_upload_is_aborted(self, fs): + sync_fs = fs._sync_fs + prefix = ( + f"{ENV.s3_staging_key}{ENV.schema}/filesystem/test_async_clear_multipart_race/" + f"{uuid.uuid4()}/" + ) + path = f"s3://{ENV.s3_staging_bucket}/{prefix}" + gone_path = S3Path(ENV.s3_staging_bucket, f"{prefix}gone") + gone = sync_fs.core.create_multipart_upload(gone_path) + try: + sync_fs.core.create_multipart_upload(S3Path(ENV.s3_staging_bucket, f"{prefix}pending")) + list_uploads = sync_fs.list_multipart_uploads + + def list_then_abort(path): + uploads = list_uploads(path) + assert len(uploads) == 2 + assert any(upload.upload_id == gone.upload_id for upload in uploads) + sync_fs.core.abort_multipart_upload(gone) + return uploads + + with mock.patch.object( + sync_fs, "list_multipart_uploads", side_effect=list_then_abort + ) as list_mock: + fs.clear_multipart_uploads(path) + list_mock.assert_called_once_with(path) + assert fs.list_multipart_uploads(path) == [] + finally: + fs.clear_multipart_uploads(path) + @pytest.mark.parametrize( ("fs", "start", "end", "target_data"), list( @@ -1944,7 +2231,9 @@ def call(**kwargs): sync_fs = fs._sync_fs sync_fs.core.create_multipart_upload = mock.MagicMock( - return_value=SimpleNamespace(upload_id="uploadid") + return_value=S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "uploadid"} + ) ) sync_fs.core.upload_part = track( lambda **kw: S3MultipartUploadPart(kw["part_number"], {"ETag": '"e"'}) diff --git a/tests/pyathena/filesystem/test_s3_core.py b/tests/pyathena/filesystem/test_s3_core.py index bf6425a0..01538087 100644 --- a/tests/pyathena/filesystem/test_s3_core.py +++ b/tests/pyathena/filesystem/test_s3_core.py @@ -28,7 +28,7 @@ S3MultipartCopyPlan, S3ObjectSummary, ) -from pyathena.filesystem.s3_object import S3MultipartUploadPart +from pyathena.filesystem.s3_object import S3MultipartUpload, S3MultipartUploadPart from pyathena.filesystem.s3_path import S3Path from pyathena.util import RetryConfig from tests.pyathena.util import ( @@ -453,8 +453,7 @@ def test_upload_part(self): ) with stubber: part = core.upload_part( - S3Path("bucket", "key"), - "u", + S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "u"}), 1, b"data", SSECustomerAlgorithm="AES256", @@ -462,6 +461,59 @@ def test_upload_part(self): ) assert (part.part_number, part.etag) == (1, '"e1"') + @pytest.mark.parametrize( + "bucket", ["myap-abc123-s3alias", "arn:aws:s3:us-east-1:123456789012:accesspoint/myap"] + ) + def test_multipart_upload_preserves_access_point_identity(self, bucket): + core, stubber = _make_core() + identity = {"Bucket": bucket, "Key": "key", "UploadId": "u"} + checksum = {"ChecksumAlgorithm": "SHA256", "ChecksumType": "COMPOSITE"} + first = {"ETag": '"first"', "ChecksumSHA256": "sha1"} + copied = {"ETag": '"copy"', "ChecksumSHA256": "sha2"} + stubber.add_response( + "create_multipart_upload", + {"Bucket": "underlying-bucket", "Key": "key", "UploadId": "u", **checksum}, + {"Bucket": bucket, "Key": "key", **checksum}, + ) + stubber.add_response( + "upload_part", + first, + {**identity, "PartNumber": 1, "Body": b"data", "ChecksumAlgorithm": "SHA256"}, + ) + stubber.add_response( + "upload_part_copy", + {"CopyPartResult": copied}, + {**identity, "PartNumber": 2, "CopySource": {"Bucket": "source", "Key": "object"}}, + ) + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"done"'}, + { + **identity, + "ChecksumType": "COMPOSITE", + "MultipartUpload": { + "Parts": [{**first, "PartNumber": 1}, {**copied, "PartNumber": 2}] + }, + }, + ) + stubber.add_client_error( + "abort_multipart_upload", + service_error_code="NoSuchUpload", + http_status_code=404, + expected_params=identity, + ) + with stubber: + upload = core.create_multipart_upload(S3Path(bucket, "key"), **checksum) + assert (upload.bucket, upload.key, upload.upload_id) == (bucket, "key", "u") + parts = [ + core.upload_part(upload, 1, b"data"), + core.upload_part_copy(upload, 2, S3Path("source", "object")), + ] + core.complete_multipart_upload(upload, parts) + with pytest.raises(FileNotFoundError): + core.abort_multipart_upload(upload) + stubber.assert_no_pending_responses() + def test_upload_part_copy(self): core, stubber = _make_core() stubber.add_response( @@ -491,14 +543,17 @@ def test_upload_part_copy(self): ) with stubber: part = core.upload_part_copy( - S3Path("bucket", "dst"), - "u", + S3MultipartUpload({"Bucket": "bucket", "Key": "dst", "UploadId": "u"}), 2, S3Path("src-bucket", "src", "v1"), range_=(10, 20), CopySourceIfMatch='"src"', ) - whole = core.upload_part_copy(S3Path("bucket", "dst"), "u", 1, S3Path("bucket", "dst")) + whole = core.upload_part_copy( + S3MultipartUpload({"Bucket": "bucket", "Key": "dst", "UploadId": "u"}), + 1, + S3Path("bucket", "dst"), + ) stubber.assert_no_pending_responses() assert (part.part_number, part.etag) == (2, '"p2"') assert (whole.part_number, whole.etag) == (1, '"p1"') @@ -521,7 +576,9 @@ def test_complete_multipart_upload(self): parts = [S3MultipartUploadPart(n, {"ETag": f'"e{n}"'}) for n in (1, 2)] with stubber: completed = core.complete_multipart_upload( - S3Path("bucket", "key"), "u", parts, RequestPayer="requester" + S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "u"}), + parts, + RequestPayer="requester", ) assert (completed.etag, completed.version_id) == ('"dst"', "v-dst") @@ -537,19 +594,27 @@ def test_abort_multipart_upload(self): "abort_multipart_upload", service_error_code="NoSuchUpload", http_status_code=404 ) with stubber: - core.abort_multipart_upload(S3Path("bucket", "key"), "u", RequestPayer="requester") + core.abort_multipart_upload( + S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "u"}), + RequestPayer="requester", + ) with pytest.raises(FileNotFoundError): - core.abort_multipart_upload(S3Path("bucket", "key"), "u") + core.abort_multipart_upload( + S3MultipartUpload({"Bucket": "bucket", "Key": "key", "UploadId": "u"}) + ) @pytest.mark.parametrize( ("method", "args"), [ ("create_multipart_upload", (S3Path("bucket"),)), - ("upload_part", (S3Path("bucket"), "u", 1, b"")), - ("upload_part_copy", (S3Path("bucket"), "u", 1, S3Path("bucket", "src"))), - ("upload_part_copy", (S3Path("bucket", "dst"), "u", 1, S3Path("bucket"))), - ("complete_multipart_upload", (S3Path("bucket"), "u", [])), - ("abort_multipart_upload", (S3Path("bucket"), "u")), + ( + "upload_part_copy", + ( + S3MultipartUpload({"Bucket": "bucket", "Key": "dst", "UploadId": "u"}), + 1, + S3Path("bucket"), + ), + ), ], ) def test_multipart_upload_requires_keys(self, method, args): @@ -557,6 +622,56 @@ def test_multipart_upload_requires_keys(self, method, args): with pytest.raises(ValueError, match="has no key"): getattr(core, method)(*args) + @pytest.mark.parametrize( + ("method", "args"), + [ + ("upload_part", (1, b"data")), + ("upload_part_copy", (1, S3Path("source", "key"))), + ("complete_multipart_upload", ([],)), + ("abort_multipart_upload", ()), + ], + ) + @pytest.mark.parametrize( + ("missing", "message"), + [("Bucket", "no bucket"), ("Key", "no key"), ("UploadId", "no upload ID")], + ) + def test_multipart_upload_requires_identity(self, method, args, missing, message): + core, _ = _make_core() + response = {"Bucket": "bucket", "Key": "key", "UploadId": "u"} + del response[missing] + with pytest.raises(ValueError, match=message): + getattr(core, method)(S3MultipartUpload(response), *args) + + def test_upload_part_uses_creation_algorithm_and_identity(self): + core, stubber = _make_core() + upload = S3MultipartUpload( + {"Bucket": "bucket", "Key": "key", "UploadId": "u", "ChecksumAlgorithm": "SHA256"} + ) + stubber.add_response( + "upload_part", + {"ETag": '"part"', "ChecksumSHA256": "sha"}, + { + "Bucket": "bucket", + "Key": "key", + "UploadId": "u", + "PartNumber": 1, + "Body": b"data", + "ChecksumAlgorithm": "SHA256", + }, + ) + with stubber: + part = core.upload_part( + upload, + 1, + b"data", + Bucket="other", + Key="other", + UploadId="other", + ChecksumAlgorithm="CRC32", + ) + stubber.assert_no_pending_responses() + assert part.to_api_repr()["ChecksumSHA256"] == "sha" + def test_create_multipart_upload_rejects_versions(self): # A write replaces the object at the key, not the named version. core, _ = _make_core() @@ -899,6 +1014,117 @@ def test_copy_object_annotation(self): ) stubber.assert_no_pending_responses() + @pytest.mark.parametrize( + "field", + [ + "ChecksumCRC32", + "ChecksumCRC32C", + "ChecksumCRC64NVME", + "ChecksumSHA1", + "ChecksumSHA256", + "ChecksumSHA512", + "ChecksumMD5", + "ChecksumXXHASH64", + "ChecksumXXHASH3", + "ChecksumXXHASH128", + ], + ) + @pytest.mark.parametrize("copy", [False, True]) + def test_complete_multipart_upload_preserves_part_checksums(self, field, copy): + core, stubber = _make_core() + identity = {"Bucket": "bucket", "Key": "key", "UploadId": "u"} + algorithm = field.removeprefix("Checksum") + stubber.add_response( + "create_multipart_upload", + {**identity, "ChecksumAlgorithm": algorithm}, + {"Bucket": "bucket", "Key": "key", "ChecksumAlgorithm": algorithm}, + ) + expected_parts = [] + for number in (1, 2): + result = {"ETag": f'"e{number}"', field: f"checksum{number}"} + request = {"Bucket": "bucket", "Key": "key", "UploadId": "u", "PartNumber": number} + if copy: + request["CopySource"] = {"Bucket": "bucket", "Key": "source"} + stubber.add_response("upload_part_copy", {"CopyPartResult": result}, request) + else: + request["Body"] = b"data" + request["ChecksumAlgorithm"] = algorithm + request[field] = f"checksum{number}" + stubber.add_response("upload_part", result, request) + expected_parts.append({**result, "PartNumber": number}) + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"done"'}, + { + "Bucket": "bucket", + "Key": "key", + "UploadId": "u", + "MultipartUpload": {"Parts": expected_parts}, + "RequestPayer": "requester", + }, + ) + with stubber: + upload = core.create_multipart_upload( + S3Path("bucket", "key"), ChecksumAlgorithm=algorithm + ) + if copy: + parts = [ + core.upload_part_copy(upload, n, S3Path("bucket", "source")) for n in (1, 2) + ] + else: + parts = [ + core.upload_part(upload, n, b"data", **{field: f"checksum{n}"}) for n in (1, 2) + ] + completed = core.complete_multipart_upload( + upload, + parts, + RequestPayer="requester", + MultipartUpload={"Parts": []}, + ) + stubber.assert_no_pending_responses() + assert completed.etag == '"done"' + + @pytest.mark.parametrize( + ("algorithm", "checksum_type"), + [(None, None), ("SHA256", "COMPOSITE"), ("CRC32", "FULL_OBJECT")], + ) + def test_complete_multipart_upload_uses_creation_algorithm(self, algorithm, checksum_type): + core, stubber = _make_core() + upload = S3MultipartUpload( + { + "Bucket": "bucket", + "Key": "key", + "UploadId": "u", + "ChecksumAlgorithm": algorithm, + "ChecksumType": checksum_type, + } + ) + part = S3MultipartUploadPart( + 1, + { + "ETag": '"part"', + "ChecksumCRC32": "sdk-crc", + "ChecksumSHA256": "upload-sha", + }, + ) + expected = {"ETag": '"part"', "PartNumber": 1} + if algorithm: + expected[f"Checksum{algorithm}"] = "upload-sha" if algorithm == "SHA256" else "sdk-crc" + stubber.add_response( + "complete_multipart_upload", + {"ETag": '"done"'}, + { + "Bucket": "bucket", + "Key": "key", + "UploadId": "u", + "MultipartUpload": {"Parts": [expected]}, + **({"ChecksumType": checksum_type} if checksum_type else {}), + }, + ) + with stubber: + core.complete_multipart_upload(upload, [part]) + stubber.assert_no_pending_responses() + class TestS3DeleteBatch: def test_from_paths(self): diff --git a/tests/pyathena/filesystem/test_s3_object.py b/tests/pyathena/filesystem/test_s3_object.py index b2f50fb9..05198746 100644 --- a/tests/pyathena/filesystem/test_s3_object.py +++ b/tests/pyathena/filesystem/test_s3_object.py @@ -344,6 +344,7 @@ def test_init(self): "BucketKeyEnabled": True, "RequestCharged": "requester", "ChecksumAlgorithm": "CRC32", + "ChecksumType": "FULL_OBJECT", "Initiated": datetime(2015, 1, 1, 0, 0, 0), "StorageClass": "STANDARD", "Owner": {"DisplayName": "test_owner", "ID": "test_owner_id"}, @@ -363,6 +364,7 @@ def test_init(self): assert actual.bucket_key_enabled is True assert actual.request_charged == "requester" assert actual.checksum_algorithm == "CRC32" + assert actual.checksum_type == "FULL_OBJECT" assert actual.initiated == datetime(2015, 1, 1, 0, 0, 0) assert actual.storage_class == "STANDARD" assert actual.owner @@ -381,6 +383,8 @@ def test_init_without_list_fields(self): } ) assert actual.initiated is None + assert actual.checksum_algorithm is None + assert actual.checksum_type is None assert actual.storage_class is None assert actual.owner is None assert actual.initiator is None @@ -453,6 +457,41 @@ def test_init(self): assert actual.bucket_key_enabled is False assert actual.request_charged is None + @pytest.mark.parametrize( + ("field", "property_name"), + [ + ("ChecksumCRC32", "checksum_crc32"), + ("ChecksumCRC32C", "checksum_crc32c"), + ("ChecksumCRC64NVME", "checksum_crc64nvme"), + ("ChecksumSHA1", "checksum_sha1"), + ("ChecksumSHA256", "checksum_sha256"), + ("ChecksumSHA512", "checksum_sha512"), + ("ChecksumMD5", "checksum_md5"), + ("ChecksumXXHASH64", "checksum_xxhash64"), + ("ChecksumXXHASH3", "checksum_xxhash3"), + ("ChecksumXXHASH128", "checksum_xxhash128"), + ], + ) + @pytest.mark.parametrize("copy", [False, True]) + def test_to_api_repr_with_checksum(self, field, property_name, copy): + result = {"ETag": '"part"', field: "checksum"} + response = {"CopyPartResult": result} if copy else result + part = S3MultipartUploadPart(2, response) + assert getattr(part, property_name) == "checksum" + assert part.to_api_repr() == {"ETag": '"part"', "PartNumber": 2, field: "checksum"} + + @pytest.mark.parametrize("copy", [False, True]) + def test_to_api_repr_omits_missing_checksums(self, copy): + result = {"ETag": '"part"', "ChecksumSHA256": None} + part = S3MultipartUploadPart(1, {"CopyPartResult": result} if copy else result) + assert part.to_api_repr() == {"ETag": '"part"', "PartNumber": 1} + assert part.checksum_crc64nvme is None + assert part.checksum_sha512 is None + assert part.checksum_md5 is None + assert part.checksum_xxhash64 is None + assert part.checksum_xxhash3 is None + assert part.checksum_xxhash128 is None + class TestS3CompleteMultipartUpload: def test_init(self):