diff --git a/docs/filesystem.md b/docs/filesystem.md index f48be6c5..cd755b19 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -300,13 +300,15 @@ directories below the bucket level) and is always a no-op. ## Typed S3 operations `S3FileSystem.core` is an `S3Core`, the typed operations that the filesystem sends -its listing, lookup, read, write, delete, multipart upload and copy requests with. It can also be -built on a boto3 S3 client. Each operation sends one request (one per page for the -iterators and `list_object_annotations()`); `plan_multipart_copy()` and -`copy_object_annotation()`, described below, send several. The requests are sent with -the retry policy. An operation raises `FileNotFoundError` for a missing bucket or -multipart upload, or for a missing object or version that it reads, and caches -nothing. Requests sent through `fs.core` do not invalidate the filesystem's cache: call +its listing, lookup, read, write, delete, multipart upload, copy, tagging, ACL and +metadata replacement requests with. It can also be built on a boto3 S3 client. Each +operation sends one request (one per page for the iterators and +`list_object_annotations()`); `plan_multipart_copy()` and `copy_object_annotation()`, +described below, send several, and `generate_presigned_url()` signs a URL locally +without a request. The requests are sent with the retry policy. An operation raises +`FileNotFoundError` for a missing bucket or multipart upload, or for a missing object or +version that it reads, and caches nothing. Requests sent through `fs.core` do not +invalidate the filesystem's cache: call `fs.invalidate_cache()` after a change, or make it through the filesystem. ```python @@ -357,6 +359,42 @@ for batch in S3DeleteBatch.from_paths(paths): print(error) # path (code: message) ``` +`get_object_tagging()` returns the tags of an object as a dictionary, and +`put_object_tagging()` replaces all of its tags with the given ones. Both act on the +version ID of the path, if any, including `null`. `put_object_acl()` and +`put_bucket_acl()` apply a canned ACL, which must be in `S3Core.OBJECT_ACLS` or +`S3Core.BUCKET_ACLS`, respectively; another value raises `ValueError`, and an object +ACL also applies to the version ID of the path. + +`replace_object_metadata()` replaces the user-defined metadata of an object, given the +`head_object()` result of the same path, by copying the object onto itself. The copy +rewrites the object, or creates a new version in a versioned bucket, so the path must +not have a version ID. The copy retains the content headers (`CacheControl`, +`ContentDisposition`, `ContentEncoding`, `ContentLanguage`, `ContentType`), `Expires`, +`WebsiteRedirectLocation`, `StorageClass`, and the `ServerSideEncryption`, +`SSEKMSKeyId` and `BucketKeyEnabled` of the object. Additional CopyObject parameters +take precedence over the retained fields, and any encryption parameter replaces all of +the retained encryption fields. + +```python +path = S3Path.parse("s3://YOUR_S3_BUCKET/path/to/object") +core.put_object_tagging(path, {**core.get_object_tagging(path), "tag2": "value2"}) + +head = core.head_object(path) +core.replace_object_metadata(path, head, {**head, "attr1": "value1"}) +``` + +`generate_presigned_url()` signs a URL locally, for `get_object` unless another client +method is given, and sends no request. Its parameters take precedence over the +`Bucket`, `Key` and `VersionId` of the path. The `request_kwargs` of the core, such as +`RequestPayer`, are not included in the signed parameters; pass them as parameters +of `generate_presigned_url()` to sign them. + +```python +url = core.generate_presigned_url(path, expires_in=600) +upload_url = core.generate_presigned_url(path, "put_object", ContentType="text/csv") +``` + ### Multipart writer `S3MultipartWriter` provides synchronous multipart requests and part planning without fsspec. diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 3bcc8d0f..dc326fc2 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -161,34 +161,6 @@ class S3FileSystem(AbstractFileSystem): """ DEFAULT_BLOCK_SIZE: int = 5 * 2**20 # 5MiB - # https://docs.aws.amazon.com/AmazonS3/latest/userguide/acl-overview.html#canned-acl - OBJECT_ACLS: frozenset[str] = frozenset( - { - "private", - "public-read", - "public-read-write", - "authenticated-read", - "aws-exec-read", - "bucket-owner-read", - "bucket-owner-full-control", - } - ) - BUCKET_ACLS: frozenset[str] = frozenset( - {"private", "public-read", "public-read-write", "authenticated-read"} - ) - # https://docs.aws.amazon.com/AmazonS3/latest/API/API_CopyObject.html - # The CopyObject parameters that set the encryption of the copy. - _SSE_COPY_PARAMS: frozenset[str] = frozenset( - { - "ServerSideEncryption", - "SSEKMSKeyId", - "SSEKMSEncryptionContext", - "BucketKeyEnabled", - "SSECustomerAlgorithm", - "SSECustomerKey", - "SSECustomerKeyMD5", - } - ) PATTERN_PATH: Pattern[str] = S3Path.PATTERN protocol = ("s3", "s3a") @@ -1282,8 +1254,8 @@ def mkdir(self, path: str, create_parents: bool = True, **kwargs) -> None: "Set allow_bucket_creation=True on the filesystem to enable it." ) acl = kwargs.pop("acl", "") - if acl and acl not in self.BUCKET_ACLS: - raise ValueError(f"ACL not in {self.BUCKET_ACLS}.") + if acl and acl not in self.core.BUCKET_ACLS: + raise ValueError(f"ACL not in {self.core.BUCKET_ACLS}.") request: dict[str, Any] = {"Bucket": s3_path.bucket} if acl: request.update({"ACL": acl}) @@ -2197,23 +2169,9 @@ def sign(self, path: str, expiration: int = 3600, **kwargs): ... client_method="put_object" ... ) """ - s3_path = S3Path.parse(path) client_method = kwargs.pop("client_method", "get_object") - params = {"Bucket": s3_path.bucket, "Key": s3_path.key} - if s3_path.version_id: - params.update({"VersionId": s3_path.version_id}) - if kwargs: - params.update(kwargs) - request = { - "ClientMethod": client_method, - "Params": params, - "ExpiresIn": expiration, - } - - _logger.debug(f"Generate signed url: {s3_path.uri}") - return self._call( - self._client.generate_presigned_url, - **request, + return self.core.generate_presigned_url( + S3Path.parse(path), client_method, expiration, **kwargs ) def metadata(self, path: str, **kwargs) -> S3Metadata: @@ -2295,43 +2253,7 @@ def setxattr(self, path: str, copy_kwargs: dict[str, Any] | None = None, **kw_ar metadata.pop(k, None) else: metadata[k] = v - - # With the REPLACE directive, S3 does not copy what the request - # omits: the system-defined metadata is dropped, and the copy is - # written as STANDARD with the default encryption of the bucket. - kept: dict[str, Any] = { - "CacheControl": head.cache_control, - "ContentDisposition": head.content_disposition, - "ContentEncoding": head.content_encoding, - "ContentLanguage": head.content_language, - "ContentType": head.content_type, - "Expires": head.expires, - "WebsiteRedirectLocation": head.website_redirect_location, - "StorageClass": head.storage_class, - } - copy_kwargs = copy_kwargs if copy_kwargs else {} - if not self._SSE_COPY_PARAMS.intersection(copy_kwargs): - kept.update( - { - "ServerSideEncryption": head.server_side_encryption, - "SSEKMSKeyId": head.sse_kms_key_id, - "BucketKeyEnabled": head.bucket_key_enabled, - } - ) - - _logger.debug(f"Set object metadata: {s3_path.uri}") - self._call( - self._client.copy_object, - CopySource={"Bucket": s3_path.bucket, "Key": s3_path.key}, - Bucket=s3_path.bucket, - Key=s3_path.key, - Metadata=metadata, - MetadataDirective="REPLACE", - **{ - **{k: v for k, v in kept.items() if v is not None}, - **copy_kwargs, - }, - ) + self.core.replace_object_metadata(s3_path, head, metadata, **(copy_kwargs or {})) self.invalidate_cache(path) def get_tags(self, path: str) -> dict[str, str]: @@ -2346,16 +2268,7 @@ def get_tags(self, path: str) -> dict[str, str]: s3_path = S3Path.parse(path) if not s3_path.key: raise ValueError("Cannot get tags of a bucket.") - request: dict[str, Any] = {"Bucket": s3_path.bucket, "Key": s3_path.key} - if s3_path.version_id: - request.update({"VersionId": s3_path.version_id}) - - _logger.debug(f"Get object tagging: {s3_path.uri}") - response = self._call( - self._client.get_object_tagging, - **request, - ) - return {v["Key"]: v["Value"] for v in response["TagSet"]} + return self.core.get_object_tagging(s3_path) def put_tags(self, path: str, tags: dict[str, str], mode: str = "o") -> None: """Set the tags for the given existing key. @@ -2377,26 +2290,12 @@ def put_tags(self, path: str, tags: dict[str, str], mode: str = "o") -> None: if not s3_path.key: raise ValueError("Cannot put tags of a bucket.") if mode == "m": - existing_tags = self.get_tags(path) - existing_tags.update(tags) - new_tags = [{"Key": k, "Value": v} for k, v in existing_tags.items()] + new_tags = {**self.core.get_object_tagging(s3_path), **tags} elif mode == "o": - new_tags = [{"Key": k, "Value": v} for k, v in tags.items()] + new_tags = tags else: raise ValueError(f"Mode must be {{'o', 'm'}}, not {mode}.") - request: dict[str, Any] = { - "Bucket": s3_path.bucket, - "Key": s3_path.key, - "Tagging": {"TagSet": new_tags}, - } - if s3_path.version_id: - request.update({"VersionId": s3_path.version_id}) - - _logger.debug(f"Put object tagging: {s3_path.uri}") - self._call( - self._client.put_object_tagging, - **request, - ) + self.core.put_object_tagging(s3_path, new_tags) def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: """Set the Access Control on a bucket/key. @@ -2414,10 +2313,10 @@ def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: s3_path = S3Path.parse(path) # Validate before any ACL is applied so that a recursive call cannot # partially apply object ACLs and then fail on the bucket ACL. - if not s3_path.key and acl not in self.BUCKET_ACLS: - raise ValueError(f"ACL not in {self.BUCKET_ACLS}.") - if s3_path.key and acl not in self.OBJECT_ACLS: - raise ValueError(f"ACL not in {self.OBJECT_ACLS}.") + if not s3_path.key and acl not in self.core.BUCKET_ACLS: + raise ValueError(f"ACL not in {self.core.BUCKET_ACLS}.") + if s3_path.key and acl not in self.core.OBJECT_ACLS: + raise ValueError(f"ACL not in {self.core.OBJECT_ACLS}.") if recursive: with self._create_executor(max_workers=self.max_workers) as executor: futures = [ @@ -2431,24 +2330,9 @@ def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: # below it have ACLs. return if s3_path.key: - request: dict[str, Any] = {"Bucket": s3_path.bucket, "Key": s3_path.key, "ACL": acl} - if s3_path.version_id: - request.update({"VersionId": s3_path.version_id}) - - _logger.debug(f"Put object acl: {s3_path.uri}") - self._call( - self._client.put_object_acl, - **request, - **kwargs, - ) + self.core.put_object_acl(s3_path, acl, **kwargs) else: - _logger.debug(f"Put bucket acl: {s3_path.uri}") - self._call( - self._client.put_bucket_acl, - Bucket=s3_path.bucket, - ACL=acl, - **kwargs, - ) + self.core.put_bucket_acl(s3_path.bucket, acl, **kwargs) def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: """List in-progress (incomplete) multipart uploads in a bucket. diff --git a/pyathena/filesystem/s3_core.py b/pyathena/filesystem/s3_core.py index 51d83455..408fb6ef 100644 --- a/pyathena/filesystem/s3_core.py +++ b/pyathena/filesystem/s3_core.py @@ -451,12 +451,13 @@ class S3Core: """Typed S3 operations on a boto3 S3 client. Each operation sends one request, or one per page for the iterators, - except those that say otherwise, such as :meth:`plan_multipart_copy` - and :meth:`copy_object_annotation`. The requests are sent - with the retry policy, and S3 errors are translated into ``OSError`` - subclasses (see :class:`~pyathena.filesystem.s3_errors.S3ClientError`): - a missing bucket or multipart upload, or a missing object or version that - an operation reads, raises ``FileNotFoundError``, and a denied request + except those that say otherwise, such as :meth:`plan_multipart_copy`, + :meth:`copy_object_annotation` and :meth:`generate_presigned_url`. The + requests are sent with the retry policy, and S3 errors are translated + into ``OSError`` subclasses (see + :class:`~pyathena.filesystem.s3_errors.S3ClientError`): a missing bucket + or multipart upload, or a missing object or version that an operation + reads, raises ``FileNotFoundError``, and a denied request ``PermissionError``. As in S3, deleting a missing key is not an error. Nothing is cached. @@ -486,6 +487,36 @@ class S3Core: "Expires", "Metadata", ) + # https://docs.aws.amazon.com/AmazonS3/latest/API/API_CopyObject.html + # The CopyObject parameters that set the encryption of the copy. + _SSE_COPY_PARAMS: ClassVar[frozenset[str]] = frozenset( + { + "ServerSideEncryption", + "SSEKMSKeyId", + "SSEKMSEncryptionContext", + "BucketKeyEnabled", + "SSECustomerAlgorithm", + "SSECustomerKey", + "SSECustomerKeyMD5", + } + ) + # https://docs.aws.amazon.com/AmazonS3/latest/userguide/acl-overview.html#canned-acl + # The canned ACLs that an object accepts. + OBJECT_ACLS: ClassVar[frozenset[str]] = frozenset( + { + "private", + "public-read", + "public-read-write", + "authenticated-read", + "aws-exec-read", + "bucket-owner-read", + "bucket-owner-full-control", + } + ) + # The canned ACLs that a bucket accepts. + BUCKET_ACLS: ClassVar[frozenset[str]] = frozenset( + {"private", "public-read", "public-read-write", "authenticated-read"} + ) def __init__( self, @@ -1100,16 +1131,9 @@ def plan_multipart_copy( # Directory buckets do not support GetObjectTagging, and their # objects have no tags. if not self._is_directory_bucket(source.bucket): - _logger.debug(f"Get tags to copy: {source.uri}") - tagging_request: dict[str, Any] = {"Bucket": source.bucket, "Key": source.key} - if source.version_id: - tagging_request.update({"VersionId": source.version_id}) - response = self.call( - self._client.get_object_tagging, - **self.operation_params("get_object_tagging", source_params), - **tagging_request, + tags = self.get_object_tagging( + source, **self.operation_params("get_object_tagging", source_params) ) - tags = [(t["Key"], t["Value"]) for t in response["TagSet"]] if tags: request.update({"Tagging": urlencode(tags)}) copy_params = self.operation_params("copy_object", request) @@ -1244,6 +1268,206 @@ def copy_object_annotation( **{**self.operation_params("put_object_annotation", params), **put_request}, ) + def get_object_tagging(self, path: S3Path, **params) -> dict[str, str]: + """Get the tags of an object, or of a version of it, with GetObjectTagging. + + Args: + path: The path of the object, with the version ID to read, if + any, including ``null``. + **params: Additional request parameters, sent as given. + + Returns: + The tags, mapping each key to its value, in the order of the + response. + + Raises: + ValueError: If the path has no key. + FileNotFoundError: If the object or version does not exist. + """ + if not path.key: + raise ValueError(f"The path has no key: {path.uri}.") + request: dict[str, Any] = {"Bucket": path.bucket, "Key": path.key} + if path.version_id: + request.update({"VersionId": path.version_id}) + _logger.debug(f"Get object tagging: {path.uri}") + response = self.call(self._client.get_object_tagging, **request, **params) + return {t["Key"]: t["Value"] for t in response["TagSet"]} + + def put_object_tagging(self, path: S3Path, tags: Mapping[str, str], **params) -> None: + """Replace the tags of an object, or of a version of it, with PutObjectTagging. + + Args: + path: The path of the object, with the version ID to tag, if any, + including ``null``. + tags: The tags, mapping each key to its value. They replace all + the existing tags. + **params: Additional request parameters, sent as given. + + Raises: + ValueError: If the path has no key. + FileNotFoundError: If the object or version does not exist. + """ + if not path.key: + raise ValueError(f"The path has no key: {path.uri}.") + request: dict[str, Any] = { + "Bucket": path.bucket, + "Key": path.key, + "Tagging": {"TagSet": [{"Key": k, "Value": v} for k, v in tags.items()]}, + } + if path.version_id: + request.update({"VersionId": path.version_id}) + _logger.debug(f"Put object tagging: {path.uri}") + self.call(self._client.put_object_tagging, **request, **params) + + def put_object_acl(self, path: S3Path, acl: str, **params) -> None: + """Apply a canned ACL to an object, or to a version of it, with PutObjectAcl. + + Args: + path: The path of the object, with the version ID to apply the + ACL to, if any, including ``null``. + acl: The canned ACL, one of ``OBJECT_ACLS``. + **params: Additional request parameters, sent as given. + + Raises: + ValueError: If the path has no key, or the ACL is not in + ``OBJECT_ACLS``. + FileNotFoundError: If the object or version does not exist. + """ + if not path.key: + raise ValueError(f"The path has no key: {path.uri}.") + if acl not in self.OBJECT_ACLS: + raise ValueError(f"ACL not in {self.OBJECT_ACLS}.") + request: dict[str, Any] = {"Bucket": path.bucket, "Key": path.key, "ACL": acl} + if path.version_id: + request.update({"VersionId": path.version_id}) + _logger.debug(f"Put object acl: {path.uri}") + self.call(self._client.put_object_acl, **request, **params) + + def put_bucket_acl(self, bucket: str, acl: str, **params) -> None: + """Apply a canned ACL to a bucket with PutBucketAcl. + + Args: + bucket: The name of the bucket. + acl: The canned ACL, one of ``BUCKET_ACLS``. + **params: Additional request parameters, sent as given. + + Raises: + ValueError: If the ACL is not in ``BUCKET_ACLS``. + FileNotFoundError: If the bucket does not exist. + """ + if acl not in self.BUCKET_ACLS: + raise ValueError(f"ACL not in {self.BUCKET_ACLS}.") + _logger.debug(f"Put bucket acl: s3://{bucket}") + self.call(self._client.put_bucket_acl, Bucket=bucket, ACL=acl, **params) + + def replace_object_metadata( + self, path: S3Path, head: S3Metadata, metadata: Mapping[str, str], **params + ) -> None: + """Replace the user-defined metadata of an object by copying it onto itself. + + S3 does not update the metadata of an object in place, so the object + is copied onto itself with CopyObject and the REPLACE metadata + directive, which writes a new object, or a new version in a + versioned bucket. With that directive, S3 does not copy what the + request omits, so the request also sends the content headers, + ``Expires``, ``WebsiteRedirectLocation`` and ``StorageClass`` of + ``head``, and, unless ``params`` set an encryption parameter, its + ``ServerSideEncryption``, ``SSEKMSKeyId`` and ``BucketKeyEnabled``. + Fields that are None in ``head`` are not sent; note that + :class:`~pyathena.filesystem.s3_object.S3Metadata` reports + ``STANDARD`` when HeadObject omits the storage class. HeadObject does + not return the KMS encryption context, so it is not retained. + + Args: + path: The path of the object, without a version ID. + head: The HeadObject result of ``path`` (see :meth:`head_object`), + whose fields are retained. + metadata: The user-defined metadata of the copy, which replaces + all the existing user-defined metadata. + **params: Additional CopyObject parameters. They take precedence + over the retained fields; one that the copy itself sets, + such as ``Metadata``, ``MetadataDirective`` or ``Key``, + raises ``TypeError``. + + Raises: + ValueError: If the path has no key or has a version ID, which a + write cannot replace. + """ + if not path.key: + raise ValueError(f"The path has no key: {path.uri}.") + if path.version_id: + raise ValueError(f"Cannot write to a version: {path.uri}.") + _logger.debug(f"Replace object metadata: {path.uri}") + retained: dict[str, Any] = { + "CacheControl": head.cache_control, + "ContentDisposition": head.content_disposition, + "ContentEncoding": head.content_encoding, + "ContentLanguage": head.content_language, + "ContentType": head.content_type, + "Expires": head.expires, + "WebsiteRedirectLocation": head.website_redirect_location, + "StorageClass": head.storage_class, + } + if not self._SSE_COPY_PARAMS.intersection(params): + retained.update( + { + "ServerSideEncryption": head.server_side_encryption, + "SSEKMSKeyId": head.sse_kms_key_id, + "BucketKeyEnabled": head.bucket_key_enabled, + } + ) + self.copy_object( + path, + path, + # botocore accepts only a dict, not another mapping such as an + # S3Metadata. + Metadata=dict(metadata), + MetadataDirective="REPLACE", + **{ + **{k: v for k, v in retained.items() if v is not None}, + **params, + }, + ) + + def generate_presigned_url( + self, + path: S3Path, + client_method: str = "get_object", + expires_in: int = 3600, + **params, + ) -> str: + """Generate a presigned URL for a request on an object. + + The URL is signed locally with the client's credentials; no request + is sent. ``request_kwargs`` are not added to the signed parameters. + + Args: + path: The path of the object, with the version ID to sign for, + if any. A path without a key is not rejected; botocore + validates the parameters of the method. + client_method: The name of the client method to sign, such as + ``get_object`` or ``put_object``. + expires_in: The number of seconds for which the URL is valid. + **params: Parameters of the method to sign. They take precedence + over the bucket, key and version ID of the path. + + Returns: + The presigned URL. + """ + request: dict[str, Any] = {"Bucket": path.bucket, "Key": path.key} + if path.version_id: + request.update({"VersionId": path.version_id}) + _logger.debug(f"Generate signed url: {path.uri}") + return cast( + str, + self.call( + self._client.generate_presigned_url, + ClientMethod=client_method, + Params={**request, **params}, + ExpiresIn=expires_in, + ), + ) + @staticmethod def _is_directory_bucket(bucket: str) -> bool: """Return whether the bucket is a directory bucket (S3 Express One Zone). diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2fc116c3..4aac93b5 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -4192,6 +4192,37 @@ def test_chmod_recursive(self): fs._client.put_object_acl, Bucket="bucket", Key="key2", ACL="private" ) + def test_put_tags_merge(self): + fs = self._make_fs() + fs._call.side_effect = [ + {"TagSet": [{"Key": "a", "Value": "1"}, {"Key": "b", "Value": "2"}]}, + {}, + ] + + fs.put_tags("s3://bucket/key?versionId=v1", {"b": "3", "c": "4"}, mode="m") + request = {"Bucket": "bucket", "Key": "key", "VersionId": "v1"} + assert fs._call.call_args_list == [ + mock.call(fs._client.get_object_tagging, **request), + mock.call( + fs._client.put_object_tagging, + **request, + Tagging={ + "TagSet": [ + {"Key": "a", "Value": "1"}, + {"Key": "b", "Value": "3"}, + {"Key": "c", "Value": "4"}, + ] + }, + ), + ] + + def test_put_tags_invalid_mode(self): + fs = self._make_fs() + + with pytest.raises(ValueError, match="Mode must be"): + fs.put_tags("s3://bucket/key", {"a": "1"}, mode="x") + fs._call.assert_not_called() + def test_list_multipart_uploads_paginates(self): fs = self._make_fs() fs._call.side_effect = [ diff --git a/tests/pyathena/filesystem/test_s3_core.py b/tests/pyathena/filesystem/test_s3_core.py index db637a0d..073711f0 100644 --- a/tests/pyathena/filesystem/test_s3_core.py +++ b/tests/pyathena/filesystem/test_s3_core.py @@ -8,6 +8,8 @@ import io from datetime import UTC, datetime from itertools import pairwise +from unittest import mock +from urllib.parse import parse_qs, urlsplit import boto3 import botocore.exceptions @@ -28,7 +30,7 @@ S3MultipartCopyPlan, S3ObjectSummary, ) -from pyathena.filesystem.s3_object import S3MultipartUpload, S3MultipartUploadPart +from pyathena.filesystem.s3_object import S3Metadata, S3MultipartUpload, S3MultipartUploadPart from pyathena.filesystem.s3_path import S3Path from pyathena.util import RetryConfig from tests.pyathena.util import ( @@ -1249,6 +1251,270 @@ def test_complete_multipart_upload_uses_creation_algorithm(self, algorithm, chec core.complete_multipart_upload(upload, [part]) stubber.assert_no_pending_responses() + @pytest.mark.parametrize( + ("version_id", "request_version"), + [(None, {}), ("v1", {"VersionId": "v1"}), ("null", {"VersionId": "null"})], + ) + def test_object_tagging(self, version_id, request_version): + core, stubber = _make_core(request_kwargs={"ExpectedBucketOwner": "111122223333"}) + expected = { + "Bucket": "bucket", + "Key": "key", + **request_version, + "ExpectedBucketOwner": "111122223333", + } + stubber.add_response( + "get_object_tagging", + {"TagSet": [{"Key": "b", "Value": "2"}, {"Key": "a", "Value": "1"}]}, + expected, + ) + stubber.add_response( + "put_object_tagging", + {}, + { + **expected, + "Tagging": {"TagSet": [{"Key": "c", "Value": "3"}, {"Key": "d", "Value": "4"}]}, + "RequestPayer": "requester", + }, + ) + path = S3Path("bucket", "key", version_id) + with stubber: + tags = core.get_object_tagging(path) + # In the order of the response. + assert list(tags.items()) == [("b", "2"), ("a", "1")] + core.put_object_tagging(path, {"c": "3", "d": "4"}, RequestPayer="requester") + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + ("version_id", "request_version"), + [(None, {}), ("v1", {"VersionId": "v1"}), ("null", {"VersionId": "null"})], + ) + def test_put_object_acl(self, version_id, request_version): + core, stubber = _make_core() + stubber.add_response( + "put_object_acl", + {}, + { + "Bucket": "bucket", + "Key": "key", + **request_version, + "ACL": "bucket-owner-full-control", + "ExpectedBucketOwner": "111122223333", + }, + ) + with stubber: + core.put_object_acl( + S3Path("bucket", "key", version_id), + "bucket-owner-full-control", + ExpectedBucketOwner="111122223333", + ) + stubber.assert_no_pending_responses() + + def test_put_bucket_acl(self): + core, stubber = _make_core() + stubber.add_response( + "put_bucket_acl", + {}, + {"Bucket": "bucket", "ACL": "private", "ExpectedBucketOwner": "111122223333"}, + ) + with stubber: + core.put_bucket_acl("bucket", "private", ExpectedBucketOwner="111122223333") + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + ("method", "args"), + [ + ("put_object_acl", (S3Path("bucket", "key"), "invalid")), + # An object ACL that a bucket does not accept. + ("put_bucket_acl", ("bucket", "bucket-owner-full-control")), + ], + ) + def test_acl_validation(self, method, args): + core, stubber = _make_core() + with stubber, pytest.raises(ValueError, match="ACL not in"): + getattr(core, method)(*args) + + @pytest.mark.parametrize( + ("method", "args", "match"), + [ + ("get_object_tagging", (S3Path("bucket"),), "has no key"), + ("put_object_tagging", (S3Path("bucket"), {}), "has no key"), + ("put_object_acl", (S3Path("bucket"), "private"), "has no key"), + ( + "replace_object_metadata", + (S3Path("bucket"), S3Metadata({}), {}), + "has no key", + ), + ( + "replace_object_metadata", + (S3Path("bucket", "key", "v1"), S3Metadata({}), {}), + "Cannot write to a version", + ), + ( + "replace_object_metadata", + (S3Path("bucket", "key", "null"), S3Metadata({}), {}), + "Cannot write to a version", + ), + ], + ) + def test_object_operations_reject_paths(self, method, args, match): + core, stubber = _make_core() + with stubber, pytest.raises(ValueError, match=match): + getattr(core, method)(*args) + + def test_replace_object_metadata(self): + core, stubber = _make_core() + expires = datetime(2030, 1, 1, tzinfo=UTC) + head = S3Metadata( + { + "CacheControl": "max-age=60", + "ContentDisposition": "attachment", + "ContentEncoding": "gzip", + "ContentLanguage": "en", + "ContentType": "text/csv", + "Expires": expires, + "WebsiteRedirectLocation": "/other", + "StorageClass": "STANDARD_IA", + "ServerSideEncryption": "aws:kms", + "SSEKMSKeyId": "k", + "BucketKeyEnabled": True, + "Metadata": {"old": "1"}, + } + ) + request = { + "CopySource": {"Bucket": "bucket", "Key": "key"}, + "Bucket": "bucket", + "Key": "key", + "Metadata": {"new": "2"}, + "MetadataDirective": "REPLACE", + "CacheControl": "max-age=60", + "ContentDisposition": "attachment", + "ContentEncoding": "gzip", + "ContentLanguage": "en", + "Expires": expires, + "WebsiteRedirectLocation": "/other", + "StorageClass": "STANDARD_IA", + } + stubber.add_response( + "copy_object", + {}, + { + **request, + "ContentType": "text/csv", + "ServerSideEncryption": "aws:kms", + "SSEKMSKeyId": "k", + "BucketKeyEnabled": True, + }, + ) + # A parameter takes precedence over the retained field, and an + # encryption parameter replaces all the retained encryption. + stubber.add_response( + "copy_object", + {}, + {**request, "ContentType": "text/plain", "SSECustomerAlgorithm": "AES256"}, + ) + with stubber: + core.replace_object_metadata(S3Path("bucket", "key"), head, {"new": "2"}) + core.replace_object_metadata( + S3Path("bucket", "key"), + head, + {"new": "2"}, + ContentType="text/plain", + SSECustomerAlgorithm="AES256", + ) + stubber.assert_no_pending_responses() + + def test_replace_object_metadata_omits_unset_fields(self): + core, stubber = _make_core() + stubber.add_response( + "copy_object", + {}, + { + "CopySource": {"Bucket": "bucket", "Key": "key"}, + "Bucket": "bucket", + "Key": "key", + "Metadata": {}, + "MetadataDirective": "REPLACE", + # S3Metadata reports STANDARD when HeadObject omits it. + "StorageClass": "STANDARD", + }, + ) + # The user-defined metadata can be any mapping, such as an S3Metadata. + stubber.add_response( + "copy_object", + {}, + { + "CopySource": {"Bucket": "bucket", "Key": "key"}, + "Bucket": "bucket", + "Key": "key", + "Metadata": {"a": "1"}, + "MetadataDirective": "REPLACE", + "StorageClass": "STANDARD", + }, + ) + with stubber: + core.replace_object_metadata(S3Path("bucket", "key"), S3Metadata({}), {}) + head = S3Metadata({"Metadata": {"a": "1"}}) + core.replace_object_metadata(S3Path("bucket", "key"), head, head) + stubber.assert_no_pending_responses() + + @pytest.mark.parametrize( + "name", ["Metadata", "MetadataDirective", "CopySource", "Bucket", "Key"] + ) + def test_replace_object_metadata_rejects_fields_of_the_copy(self, name): + core, stubber = _make_core() + with stubber, pytest.raises(TypeError, match=name): + core.replace_object_metadata(S3Path("bucket", "key"), S3Metadata({}), {}, **{name: "x"}) + + @pytest.mark.parametrize( + ("path", "client_method", "params", "expected"), + [ + (S3Path("bucket", "key"), "get_object", {}, {"Bucket": "bucket", "Key": "key"}), + ( + S3Path("bucket", "key", "null"), + "get_object", + {"ResponseContentType": "text/csv"}, + { + "Bucket": "bucket", + "Key": "key", + "VersionId": "null", + "ResponseContentType": "text/csv", + }, + ), + # The parameters take precedence over the path. + ( + S3Path("bucket", "key", "v1"), + "put_object", + {"Key": "other", "VersionId": None}, + {"Bucket": "bucket", "Key": "other", "VersionId": None}, + ), + ], + ) + def test_generate_presigned_url(self, path, client_method, params, expected): + core, _ = _make_core(request_kwargs={"RequestPayer": "requester"}) + core.call = mock.MagicMock(return_value="https://signed") + assert ( + core.generate_presigned_url(path, client_method, expires_in=60, **params) + == "https://signed" + ) + # request_kwargs are not signed. + core.call.assert_called_once_with( + core.client.generate_presigned_url, + ClientMethod=client_method, + Params=expected, + ExpiresIn=60, + ) + + def test_generate_presigned_url_signs_locally(self): + # No request is sent, so the stubber has no responses to return. + core, stubber = _make_core(request_kwargs={"RequestPayer": "requester"}) + with stubber: + url = core.generate_presigned_url(S3Path("bucket", "key", "v1")) + query = parse_qs(urlsplit(url).query) + assert query["versionId"] == ["v1"] + # request_kwargs are not signed. + assert not any(k.lower() == "x-amz-request-payer" for k in query) + class TestS3DeleteBatch: def test_from_paths(self):