diff --git a/docs/filesystem.md b/docs/filesystem.md index fd2439ead..9159fd6ef 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -109,6 +109,22 @@ entry. `cat_file` uses the key as written. Without a `?versionId=` suffix, it re such an object without a range, with a non-empty range of non-negative offsets, or with a negative `start` and no `end`, and raises `FileNotFoundError` for other ranges. +S3 request parameters, such as `ContentType`, `ServerSideEncryption`, or `RequestPayer`, +can be given to `open`, `pipe`, and `put` as keyword arguments or in +`s3_additional_kwargs`, and to all of them through the `s3_additional_kwargs` argument +of `S3FileSystem`. The parameters of a call take precedence over those of the +filesystem. A file sends each of its requests only the parameters that the S3 operation +accepts, so, for example, `ServerSideEncryption` for writes is not sent with reads. A +`pipe` of data up to the block size sends its parameters with a single PutObject +request as given. `put` sets `ContentType` from the file extension unless the call or +the filesystem gives one. + +```python +fs = S3FileSystem(s3_additional_kwargs={"ServerSideEncryption": "AES256"}) +with fs.open("s3://YOUR_S3_BUCKET/path/to/data.csv", "wb", ContentType="text/csv") as f: + f.write(b"col1\n1\n") +``` + ## Error translation S3 error responses are translated into standard Python exceptions, so filesystem diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 7be763aa6..5e6b1a1a8 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -8,7 +8,7 @@ import mimetypes import os.path import re -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterator, Mapping from concurrent.futures import Future, as_completed, wait from copy import deepcopy from datetime import datetime @@ -162,7 +162,8 @@ def __init__( max_workers: The number of threads for parallel transfers. s3_additional_kwargs: Extra arguments for the object requests of ``open()`` and ``pipe_file()``; listings and other requests do - not use them. + not use them. Each request receives those that its operation + accepts, and the parameters of a call take precedence. allow_bucket_creation: Whether ``mkdir``/``makedirs`` may create a bucket. allow_bucket_deletion: Whether ``rmdir`` may delete a bucket. @@ -171,7 +172,8 @@ def __init__( *args: Passed to ``fsspec.AbstractFileSystem``. **kwargs: Passed to ``fsspec.AbstractFileSystem``; without a ``connection``, also s3fs-compatible client arguments. - ``requester_pays=True`` sends requester-pays requests. + ``requester_pays=True`` sends requester-pays requests with + the operations that accept ``RequestPayer``. """ super().__init__(*args, **kwargs) if connection: @@ -1318,7 +1320,9 @@ def cp_file( recursive: Unused parameter for fsspec compatibility. maxdepth: Unused parameter for fsspec compatibility. on_error: Unused parameter for fsspec compatibility. - **kwargs: Additional S3 copy parameters (e.g., metadata, storage class). + **kwargs: Additional S3 copy parameters (e.g., metadata, storage + class). The ``block_size`` and ``max_workers`` parameters + control a multipart copy and are not sent to S3. Raises: ValueError: If trying to copy to a versioned file or copy buckets. @@ -1335,6 +1339,9 @@ def cp_file( # >= 2026.6.0, where mv() passes on_error correctly. # https://github.com/fsspec/filesystem_spec/commit/346a589fef9308550ffa3d0d510f2db67281bb05 kwargs.pop("onerror", None) + # Parameters of the multipart copy, not of the S3 requests. + block_size = kwargs.pop("block_size", None) + max_workers = kwargs.pop("max_workers", None) bucket1, key1, version_id1 = self.parse_path(path1) bucket2, key2, version_id2 = self.parse_path(path2) if version_id2: @@ -1361,6 +1368,8 @@ def cp_file( size1=size1, bucket2=bucket2, key2=key2, + max_workers=max_workers, + block_size=block_size, **kwargs, ) self.invalidate_cache(path2) @@ -1439,6 +1448,7 @@ def _copy_object_with_multipart_upload( upload_id=cast(str, multipart_upload.upload_id), part_number=i + 1, copy_source_ranges=range_, + **self._get_operation_kwargs("upload_part_copy", kwargs), ) for i, range_ in enumerate(ranges) ] @@ -1447,6 +1457,7 @@ def _copy_object_with_multipart_upload( key=key2, upload_id=cast(str, multipart_upload.upload_id), futures=futures, + request_kwargs=kwargs, ) def _get_copy_ranges(self, size: int, block_size: int) -> list[tuple[int, int]]: @@ -1561,7 +1572,7 @@ def pipe_file( kwargs.pop("block_size", None) kwargs.pop("max_workers", None) request_kwargs = { - **self.s3_additional_kwargs, + **self._get_operation_kwargs("put_object", self.s3_additional_kwargs), **kwargs.pop("s3_additional_kwargs", {}), **kwargs, } @@ -1574,6 +1585,7 @@ def _finish_multipart_upload( key: str, upload_id: str, futures: list[Future[S3MultipartUploadPart]], + request_kwargs: Mapping[str, Any] | None = None, ) -> S3CompleteMultipartUpload: """Collect the uploaded parts and complete the multipart upload. @@ -1587,10 +1599,14 @@ def _finish_multipart_upload( key: Object key being uploaded. upload_id: Unique identifier for the multipart upload. 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 + the abort receive those that they accept. Returns: S3CompleteMultipartUpload of the completed upload. """ + request_kwargs = request_kwargs or {} try: # The futures are in part-number order. results = [future.result() for future in futures] @@ -1600,6 +1616,7 @@ def _finish_multipart_upload( key=key, upload_id=upload_id, parts=parts, + **self._get_operation_kwargs("complete_multipart_upload", request_kwargs), ) except Exception: # A part that is still uploading when the upload is aborted may @@ -1609,9 +1626,12 @@ def _finish_multipart_upload( try: self._call( self._client.abort_multipart_upload, - Bucket=bucket, - Key=key, - UploadId=upload_id, + **{ + **self._get_operation_kwargs("abort_multipart_upload", request_kwargs), + "Bucket": bucket, + "Key": key, + "UploadId": upload_id, + }, ) except Exception: _logger.exception( @@ -1711,7 +1731,8 @@ def put_file(self, lpath: str, rpath: str, callback=_DEFAULT_CALLBACK, **kwargs) rpath: S3 destination path (s3://bucket/key). callback: Progress callback for tracking upload progress. **kwargs: Additional S3 parameters (e.g., ContentType, StorageClass). - The ``block_size`` parameter of ``open()`` is also accepted. + The ``block_size``, ``max_workers``, and ``s3_additional_kwargs`` + parameters of ``open()`` are also accepted. Raises: ValueError: If the file takes more than @@ -1733,15 +1754,24 @@ def put_file(self, lpath: str, rpath: str, callback=_DEFAULT_CALLBACK, **kwargs) size = os.path.getsize(lpath) block_size = kwargs.pop("block_size", None) or self.default_block_size + max_workers = kwargs.pop("max_workers", self.max_workers) + # The other parameters are S3 request parameters, as in pipe_file(). + s3_additional_kwargs = {**kwargs.pop("s3_additional_kwargs", {}), **kwargs} self._check_multipart_upload_size(rpath, size, block_size) callback.set_size(size) - if "ContentType" not in kwargs: + if "ContentType" not in {**self.s3_additional_kwargs, **s3_additional_kwargs}: content_type, _ = mimetypes.guess_type(lpath) if content_type is not None: - kwargs["ContentType"] = content_type + s3_additional_kwargs["ContentType"] = content_type with ( - self.open(rpath, "wb", block_size=block_size, s3_additional_kwargs=kwargs) as remote, + self.open( + rpath, + "wb", + block_size=block_size, + max_workers=max_workers, + s3_additional_kwargs=s3_additional_kwargs, + ) as remote, open(lpath, "rb") as local, ): while data := local.read(remote.blocksize): @@ -2266,8 +2296,12 @@ def _open( if cache_type is None: cache_type = self.default_cache_type max_workers = kwargs.pop("max_workers", self.max_workers) - s3_additional_kwargs = kwargs.pop("s3_additional_kwargs", {}) - s3_additional_kwargs.update(self.s3_additional_kwargs) + # The parameters of the call take precedence over those of the + # filesystem; the caller's dictionary is not modified. + s3_additional_kwargs = { + **self.s3_additional_kwargs, + **kwargs.pop("s3_additional_kwargs", {}), + } return S3File( self, @@ -2326,8 +2360,9 @@ def _get_object( _logger.debug(f"Get object: s3://{bucket}/{key}?versionId={version_id}&range={range_}") response = self._call( self._client.get_object, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return ranges[0], cast(bytes, response["Body"].read()) @@ -2339,8 +2374,9 @@ def _put_object(self, bucket: str, key: str, body: bytes | None, **kwargs) -> S3 _logger.debug(f"Put object: s3://{bucket}/{key}") response = self._call( self._client.put_object, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return S3PutObject(response) @@ -2353,8 +2389,9 @@ def _create_multipart_upload(self, bucket: str, key: str, **kwargs) -> S3Multipa _logger.debug(f"Create multipart upload to s3://{bucket}/{key}.") response = self._call( self._client.create_multipart_upload, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return S3MultipartUpload(response) @@ -2383,8 +2420,9 @@ def _upload_part_copy( ) response = self._call( self._client.upload_part_copy, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return S3MultipartUploadPart(part_number, response) @@ -2408,8 +2446,9 @@ def _upload_part( _logger.debug(f"Upload part of {upload_id} to s3://{bucket}/{key} as part {part_number}.") response = self._call( self._client.upload_part, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return S3MultipartUploadPart(part_number, response) @@ -2426,17 +2465,47 @@ def _complete_multipart_upload( _logger.debug(f"Complete multipart upload {upload_id} to s3://{bucket}/{key}.") response = self._call( self._client.complete_multipart_upload, - **request, - **kwargs, + # The fields of the request take precedence over inherited + # parameters of the same name. + **{**kwargs, **request}, ) return S3CompleteMultipartUpload(response) + def _get_operation_kwargs(self, method: str, kwargs: Mapping[str, Any]) -> dict[str, Any]: + """Select the parameters that an S3 operation accepts. + + Parameters that are inherited by several requests (the + ``requester_pays`` parameter, ``s3_additional_kwargs``, or the + parameters of a file or a multipart copy) are filtered by the input + shape of each operation, so that, e.g., ``ServerSideEncryption`` for + writes is not sent with GetObject. + + Args: + method: The name of the client method, such as ``get_object``. + kwargs: The parameters to select from. + + Returns: + The parameters that the operation accepts. Empty for a method + that is not an S3 API operation, such as + ``generate_presigned_url``. + """ + operation = self._client.meta.method_to_api_mapping.get(method) + if not kwargs or operation is None: + return {} + members = self._client.meta.service_model.operation_model(operation).input_shape.members + return {k: v for k, v in kwargs.items() if k in members} + def _call(self, method: str | Callable[..., Any], **kwargs) -> dict[str, Any]: func = getattr(self._client, method) if isinstance(method, str) else method + # The requester_pays parameter goes only to the operations that + # accept it, and a parameter of the call takes precedence. + request = ( + {**self._get_operation_kwargs(func.__name__, self.request_kwargs), **kwargs} + if self.request_kwargs + else kwargs + ) try: - response = retry_api_call( - func, config=self._retry_config, logger=_logger, **kwargs, **self.request_kwargs - ) + response = retry_api_call(func, config=self._retry_config, logger=_logger, **request) except botocore.exceptions.ClientError as e: raise S3ClientError(e).os_error from e return cast(dict[str, Any], response) @@ -2496,9 +2565,11 @@ def __init__( cache_options: Options for the fsspec cache. size: The size of the object, if known. Passed to ``fsspec.spec.AbstractBufferedFile``. - s3_additional_kwargs: Additional parameters for the object requests - of the file. - **kwargs: Accepted for compatibility; not used. + s3_additional_kwargs: Additional parameters for the S3 requests of + the file, such as ``ContentType`` or ``RequestPayer``. Each + request receives those that its operation accepts. + **kwargs: Additional parameters for the S3 requests of the file, + which take precedence over ``s3_additional_kwargs``. Raises: FileNotFoundError: If no object exists at the path when reading, @@ -2509,7 +2580,8 @@ def __init__( ``MULTIPART_UPLOAD_MAX_PART_SIZE`` for writing. """ self.max_workers = max_workers - self.s3_additional_kwargs = s3_additional_kwargs if s3_additional_kwargs else {} + # A new dictionary, so that the caller's is not modified. + self.s3_additional_kwargs: dict[str, Any] = {**(s3_additional_kwargs or {}), **kwargs} # The arguments are validated, and the objects looked up, before the # base class initializer: a file that fails here is never opened, so @@ -2599,6 +2671,18 @@ def __init__( self.s3_additional_kwargs.update(append_info.to_api_repr()) self._details = append_info + def _get_request_kwargs(self, method: str) -> dict[str, Any]: + """Select the parameters of the file that an S3 operation accepts. + + Args: + method: The name of the client method, such as ``upload_part``. + + Returns: + The parameters in ``s3_additional_kwargs`` that the operation + accepts. + """ + return self.fs._get_operation_kwargs(method, self.s3_additional_kwargs) + def close(self) -> None: """Close the file, flushing any written data, and shut down its executor.""" try: @@ -2617,7 +2701,7 @@ def _initiate_upload(self) -> None: self.multipart_upload = self.fs._create_multipart_upload( bucket=self.bucket, key=self.key, - **self.s3_additional_kwargs, + **self._get_request_kwargs("create_multipart_upload"), ) if self.append_block: if self.tell() > self.fs.MULTIPART_UPLOAD_MAX_PART_SIZE: @@ -2637,6 +2721,7 @@ def _initiate_upload(self) -> None: upload_id=cast(str, self.multipart_upload.upload_id), part_number=i + 1, copy_source_ranges=range_, + **self._get_request_kwargs("upload_part_copy"), ) ) else: @@ -2648,6 +2733,7 @@ def _initiate_upload(self) -> None: copy_source=self.path, upload_id=cast(str, self.multipart_upload.upload_id), part_number=1, + **self._get_request_kwargs("upload_part_copy"), ) ) @@ -2729,6 +2815,7 @@ def _upload_chunk(self, final: bool = False) -> bool: upload_id=cast(str, self.multipart_upload.upload_id), part_number=part_number, body=upload, + **self._get_request_kwargs("upload_part"), ) ) @@ -2753,7 +2840,7 @@ def commit(self) -> None: if self.tell() == 0: if self.buffer is not None: self.discard() - self.fs.touch(self.path, **self.s3_additional_kwargs) + self.fs.touch(self.path, **self._get_request_kwargs("put_object")) elif not self.multipart_upload_parts: if self.buffer is not None: # Upload files smaller than block size. @@ -2763,7 +2850,7 @@ def commit(self) -> None: bucket=self.bucket, key=self.key, body=data, - **self.s3_additional_kwargs, + **self._get_request_kwargs("put_object"), ) else: if not self.multipart_upload: @@ -2775,6 +2862,7 @@ def commit(self) -> None: key=self.key, upload_id=cast(str, self.multipart_upload.upload_id), futures=self.multipart_upload_parts, + request_kwargs=self.s3_additional_kwargs, ) except Exception: # The multipart upload has been aborted by the helper; @@ -2796,18 +2884,13 @@ def discard(self) -> None: # be stored after the abort, so wait for the parts that could not # be cancelled first. wait([f for f in self.multipart_upload_parts if not f.cancel()]) - # s3_additional_kwargs also holds object parameters (e.g., the - # existing object's metadata in append mode) that - # AbortMultipartUpload rejects. self.fs._call( "abort_multipart_upload", - Bucket=self.bucket, - Key=self.key, - UploadId=self.multipart_upload.upload_id, **{ - k: v - for k, v in self.s3_additional_kwargs.items() - if k in ("RequestPayer", "ExpectedBucketOwner") + **self._get_request_kwargs("abort_multipart_upload"), + "Bucket": self.bucket, + "Key": self.key, + "UploadId": self.multipart_upload.upload_id, }, ) @@ -2898,7 +2981,7 @@ def _fetch_range(self, start: int, end: int) -> bytes: key=self.key, ranges=r, version_id=self.version_id, - **self.s3_additional_kwargs, + **self._get_request_kwargs("get_object"), ) for r in ranges ] @@ -2909,7 +2992,7 @@ def _fetch_range(self, start: int, end: int) -> bytes: self.key, ranges[0], self.version_id, - **self.s3_additional_kwargs, + **self._get_request_kwargs("get_object"), )[1] return object_ diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 441154674..1993d2c3d 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -216,7 +216,8 @@ def _put_file_in_transaction(self, lpath: str, rpath: str, callback, **kwargs) - rpath: S3 destination path (s3://bucket/key). callback: Progress callback for tracking upload progress. **kwargs: Additional S3 parameters (e.g., ContentType, StorageClass). - The ``block_size`` parameter of ``open()`` is also accepted. + The ``block_size``, ``max_workers``, and ``s3_additional_kwargs`` + parameters of ``open()`` are also accepted. Raises: ValueError: If the file takes more than @@ -230,15 +231,24 @@ def _put_file_in_transaction(self, lpath: str, rpath: str, callback, **kwargs) - size = os.path.getsize(lpath) block_size = kwargs.pop("block_size", None) or self._sync_fs.default_block_size + max_workers = kwargs.pop("max_workers", self._sync_fs.max_workers) + # The other parameters are S3 request parameters, as in pipe_file(). + s3_additional_kwargs = {**kwargs.pop("s3_additional_kwargs", {}), **kwargs} self._sync_fs._check_multipart_upload_size(rpath, size, block_size) callback.set_size(size) - if "ContentType" not in kwargs: + if "ContentType" not in {**self._sync_fs.s3_additional_kwargs, **s3_additional_kwargs}: content_type, _ = mimetypes.guess_type(lpath) if content_type is not None: - kwargs["ContentType"] = content_type + s3_additional_kwargs["ContentType"] = content_type with ( - self.open(rpath, "wb", block_size=block_size, s3_additional_kwargs=kwargs) as remote, + self.open( + rpath, + "wb", + block_size=block_size, + max_workers=max_workers, + s3_additional_kwargs=s3_additional_kwargs, + ) as remote, open(lpath, "rb") as local, ): while data := local.read(remote.blocksize): @@ -298,6 +308,9 @@ async def _cp_file(self, path1: str, path2: str, **kwargs) -> None: # fsspec < 2026.6.0 leaks the typo'd "onerror" keyword from mv(); # see S3FileSystem.cp_file. kwargs.pop("onerror", None) + # Parameters of the multipart copy, not of the S3 requests. + block_size = kwargs.pop("block_size", None) + max_workers = kwargs.pop("max_workers", None) bucket1, key1, version_id1 = self.parse_path(path1) bucket2, key2, version_id2 = self.parse_path(path2) if version_id2: @@ -325,6 +338,8 @@ async def _cp_file(self, path1: str, path2: str, **kwargs) -> None: size1=size1, bucket2=bucket2, key2=key2, + max_workers=max_workers, + block_size=block_size, **kwargs, ) self._sync_fs.invalidate_cache(path2) @@ -336,10 +351,12 @@ async def _copy_object_with_multipart_upload( size1: int, bucket2: str, key2: str, + max_workers: int | None = None, block_size: int | None = None, version_id1: str | None = None, **kwargs, ) -> None: + max_workers = max_workers if max_workers else self._sync_fs.max_workers block_size = block_size if block_size else S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE if ( block_size < S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE @@ -367,16 +384,21 @@ async def _copy_object_with_multipart_upload( **kwargs, ) + semaphore = asyncio.Semaphore(max_workers) + part_kwargs = self._sync_fs._get_operation_kwargs("upload_part_copy", kwargs) + async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any]: - result = await asyncio.to_thread( - self._sync_fs._upload_part_copy, - bucket=bucket2, - key=key2, - copy_source=copy_source, - upload_id=cast(str, multipart_upload.upload_id), - part_number=i + 1, - copy_source_ranges=range_, - ) + async with semaphore: + result = await asyncio.to_thread( + self._sync_fs._upload_part_copy, + bucket=bucket2, + key=key2, + copy_source=copy_source, + upload_id=cast(str, multipart_upload.upload_id), + part_number=i + 1, + copy_source_ranges=range_, + **part_kwargs, + ) return { "ETag": result.etag, "PartNumber": result.part_number, @@ -391,6 +413,7 @@ async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any]: key=key2, upload_id=cast(str, multipart_upload.upload_id), parts=parts_list, + **self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs), ) async def _find( @@ -440,8 +463,12 @@ def _open( if cache_type is None: cache_type = self._sync_fs.default_cache_type max_workers = kwargs.pop("max_workers", self._sync_fs.max_workers) - s3_additional_kwargs = kwargs.pop("s3_additional_kwargs", {}) - s3_additional_kwargs.update(self._sync_fs.s3_additional_kwargs) + # The parameters of the call take precedence over those of the + # filesystem; the caller's dictionary is not modified. + s3_additional_kwargs = { + **self._sync_fs.s3_additional_kwargs, + **kwargs.pop("s3_additional_kwargs", {}), + } return AioS3File( self._sync_fs, diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index dbfefc361..ca47605d1 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1,4 +1,5 @@ import asyncio +import contextlib import functools import gc import io @@ -18,8 +19,10 @@ from types import SimpleNamespace from unittest import mock +import boto3 import botocore.exceptions import pytest +from botocore.stub import Stubber from fsspec import Callback from fsspec.dircache import DirCache @@ -32,6 +35,12 @@ from tests import ENV from tests.pyathena.conftest import connect +# A client that sends no requests; its service model selects the parameters +# that each S3 operation accepts. +S3_CLIENT = boto3.client( + "s3", region_name="us-east-1", aws_access_key_id="dummy", aws_secret_access_key="dummy" +) + @pytest.fixture(scope="class") def register_filesystem(): @@ -145,6 +154,8 @@ def _make_fs(): fs = S3FileSystem.__new__(S3FileSystem) fs.dircache = {} fs._client = mock.MagicMock() + fs._client.meta.method_to_api_mapping = S3_CLIENT.meta.method_to_api_mapping + fs._client.meta.service_model = S3_CLIENT.meta.service_model fs._call = mock.MagicMock() fs._retry_config = RetryConfig() fs.request_kwargs = {} @@ -646,6 +657,328 @@ def test_pipe_file_small_uses_put_object(self): ContentType="text/plain", ) + @pytest.mark.parametrize( + ("method", "kwargs", "expected"), + [ + ( + "get_object", + {"ServerSideEncryption": "AES256", "RequestPayer": "requester", "IfMatch": '"e"'}, + {"RequestPayer": "requester", "IfMatch": '"e"'}, + ), + ("head_bucket", {"RequestPayer": "requester"}, {}), + ( + "upload_part", + {"ContentType": "text/csv", "SSECustomerAlgorithm": "AES256"}, + {"SSECustomerAlgorithm": "AES256"}, + ), + # Not an S3 API operation. + ("generate_presigned_url", {"RequestPayer": "requester"}, {}), + ], + ) + def test_get_operation_kwargs(self, method, kwargs, expected): + assert self._make_fs()._get_operation_kwargs(method, kwargs) == expected + + def test_requester_pays(self): + # GH-969: RequestPayer is sent only with the operations that accept + # it, and one given to a call does not conflict with it (GH-946). + fs = S3FileSystem( + key="dummy", + secret="dummy", + region_name="us-east-1", + requester_pays=True, + skip_instance_cache=True, + ) + head_object = {"ContentLength": 1, "ETag": '"e"'} + with Stubber(fs._client) as stubber: + stubber.add_response("head_bucket", {}, {"Bucket": "bucket"}) + stubber.add_response( + "head_object", + head_object, + {"Bucket": "bucket", "Key": "key", "RequestPayer": "requester"}, + ) + stubber.add_response( + "head_object", + head_object, + {"Bucket": "bucket", "Key": "key2", "RequestPayer": "requester"}, + ) + fs.info("s3://bucket") + fs.metadata("s3://bucket/key") + fs.metadata("s3://bucket/key2", RequestPayer="requester") + stubber.assert_no_pending_responses() + assert fs.sign("s3://bucket/key").startswith("https://") + + def test_open_s3_additional_kwargs(self): + # GH-969: the parameters of the call take precedence over those of + # the filesystem, keyword parameters are added to them, each request + # receives those that it accepts, and the caller's dictionary is not + # modified. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs.s3_additional_kwargs = {"ServerSideEncryption": "AES256", "StorageClass": "STANDARD"} + fs.info = mock.MagicMock( + return_value=S3Object( + init={"ContentLength": 3, "ETag": '"e"'}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="key", + ) + ) + fs._get_object = mock.MagicMock(return_value=(0, b"abc")) + fs._put_object = mock.MagicMock() + kwargs = {"StorageClass": "GLACIER_IR", "ExpectedBucketOwner": "111122223333"} + + with fs.open("s3://bucket/key", "rb", s3_additional_kwargs=kwargs) as f: + assert f.read() == b"abc" + with fs.open( + "s3://bucket/key", "wb", s3_additional_kwargs=kwargs, ContentType="text/csv" + ) as f: + f.write(b"x") + + assert kwargs == {"StorageClass": "GLACIER_IR", "ExpectedBucketOwner": "111122223333"} + fs._get_object.assert_called_once_with( + "bucket", "key", (0, 3), None, ExpectedBucketOwner="111122223333", IfMatch='"e"' + ) + fs._put_object.assert_called_once_with( + bucket="bucket", + key="key", + body=b"x", + ServerSideEncryption="AES256", + StorageClass="GLACIER_IR", + ExpectedBucketOwner="111122223333", + ContentType="text/csv", + ) + + @pytest.mark.parametrize("transaction", [False, True]) + def test_pipe_file_buffered_s3_parameters(self, transaction): + # GH-969: the parameters of the call reach the upload when the data + # goes through the buffered path, as on the single-request path. + fs = S3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._upload_part = mock.MagicMock( + side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + ) + fs._complete_multipart_upload = mock.MagicMock() + fs._put_object = mock.MagicMock() + data = b"x" * (fs.MULTIPART_UPLOAD_MIN_PART_SIZE + 1) + + if transaction: + with fs.transaction: + fs.pipe_file("s3://bucket/key", b"x", ContentType="text/csv") + fs._put_object.assert_called_once_with( + bucket="bucket", key="key", body=b"x", ContentType="text/csv" + ) + else: + fs.pipe_file("s3://bucket/key", data, ContentType="text/csv") + fs._create_multipart_upload.assert_called_once_with( + bucket="bucket", key="key", ContentType="text/csv" + ) + + def test_put_file_open_parameters(self, tmp_path): + # GH-969: the open() parameters of put_file() go to open(), and the + # other parameters, also in s3_additional_kwargs, to S3. + fs = S3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) + lpath = tmp_path / "data.csv" + lpath.write_bytes(b"a") + block_size = fs.MULTIPART_UPLOAD_MIN_PART_SIZE + + with ( + mock.patch.object(fs, "open", wraps=fs.open) as open_, + Stubber(fs._client) as stubber, + ): + stubber.add_response( + "put_object", + {"ETag": '"e"'}, + { + "Bucket": "bucket", + "Key": "key", + "Body": b"a", + "ContentType": "text/csv", + "StorageClass": "STANDARD_IA", + }, + ) + fs.put_file( + str(lpath), + "s3://bucket/key", + block_size=block_size, + max_workers=2, + s3_additional_kwargs={"StorageClass": "STANDARD_IA"}, + ) + stubber.assert_no_pending_responses() + + open_.assert_called_once_with( + "s3://bucket/key", + "wb", + block_size=block_size, + max_workers=2, + s3_additional_kwargs={"StorageClass": "STANDARD_IA", "ContentType": "text/csv"}, + ) + + @pytest.mark.parametrize("fail", [False, True]) + def test_open_parameters_named_as_request_fields(self, fail): + # Parameters of a file named like the fields that a request sets + # itself do not replace them, so the parts, the completion and the + # abort use the upload of the file. + fs = self._make_fs() + fs.default_cache_type = "bytes" + requests = [] + + def call(method, **request): + name = method if isinstance(method, str) else method._extract_mock_name() + name = name.split(".")[-1] + requests.append((name, request)) + if name == "upload_part" and fail: + raise OSError("upload failed") + return {"UploadId": "uploadid", "ETag": '"e"'} + + fs._call.side_effect = call + block_size = fs.MULTIPART_UPLOAD_MIN_PART_SIZE + + with ( + pytest.raises(OSError, match="upload failed") if fail else contextlib.nullcontext(), + fs.open( + "s3://bucket/key", + "wb", + block_size=block_size, + Key="other", + UploadId="other", + PartNumber=99, + ) as f, + ): + f.write(b"x" * (block_size + 1)) + + names = [name for name, _ in requests] + expected = "abort_multipart_upload" if fail else "complete_multipart_upload" + assert names == ["create_multipart_upload", "upload_part", expected] + for name, request in requests: + assert request["Key"] == "key" + if name != "create_multipart_upload": + assert request["UploadId"] == "uploadid" + assert requests[1][1]["PartNumber"] == 1 + + 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() + fs._complete_multipart_upload = mock.MagicMock() + kwargs = { + "ContentType": "text/csv", + "RequestPayer": "requester", + "SSECustomerAlgorithm": "AES256", + } + 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 + ) + 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", + futures=[failed], + request_kwargs=kwargs, + ) + + fs._complete_multipart_upload.assert_called_once_with( + bucket="bucket", + key="key", + upload_id="uploadid", + parts=[{"ETag": '"e1"', "PartNumber": 1}], + RequestPayer="requester", + SSECustomerAlgorithm="AES256", + ) + fs._call.assert_called_once_with( + fs._client.abort_multipart_upload, + Bucket="bucket", + Key="key", + UploadId="uploadid", + RequestPayer="requester", + ) + + @pytest.mark.parametrize("size", [10, 5 * 2**30 + 1]) + def test_cp_file_multipart_parameters(self, size): + # GH-967: block_size and max_workers control a multipart copy and are + # not sent to S3, whatever the size of the object. + fs = self._make_fs() + fs.info = mock.MagicMock( + return_value=S3Object( + init={"ContentLength": size}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="src", + ) + ) + fs._copy_object = mock.MagicMock() + fs._copy_object_with_multipart_upload = mock.MagicMock() + + fs.cp_file( + "s3://bucket/src", + "s3://bucket/dst", + block_size=fs.MULTIPART_UPLOAD_MIN_PART_SIZE, + max_workers=2, + RequestPayer="requester", + ) + + if size <= fs.MULTIPART_UPLOAD_MAX_PART_SIZE: + fs._copy_object.assert_called_once_with( + bucket1="bucket", + key1="src", + version_id1=None, + bucket2="bucket", + key2="dst", + RequestPayer="requester", + ) + else: + fs._copy_object_with_multipart_upload.assert_called_once_with( + bucket1="bucket", + key1="src", + version_id1=None, + size1=size, + bucket2="bucket", + key2="dst", + max_workers=2, + block_size=fs.MULTIPART_UPLOAD_MIN_PART_SIZE, + RequestPayer="requester", + ) + + def test_copy_object_with_multipart_upload_request_parameters(self): + # GH-946: the part copies receive the parameters of the copy that + # they accept, and the completion and the abort get them all. + fs = self._make_fs() + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._upload_part_copy = mock.MagicMock( + side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + ) + fs._finish_multipart_upload = mock.MagicMock() + kwargs = {"ContentType": "text/csv", "RequestPayer": "requester"} + + fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=5 * 2**30 + 2**20, + bucket2="bucket", + key2="dst", + **kwargs, + ) + + fs._create_multipart_upload.assert_called_once_with(bucket="bucket", key="dst", **kwargs) + assert all( + c.kwargs["RequestPayer"] == "requester" and "ContentType" not in c.kwargs + for c in fs._upload_part_copy.call_args_list + ) + assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == kwargs + def test_pipe_file_invalid_path_raises(self): fs = self._make_fs() with pytest.raises(ValueError, match="Cannot write to a bucket"): @@ -712,7 +1045,8 @@ def test_put_file_exceeding_max_parts(self, tmp_path, kwargs): fs._call.assert_not_called() def test_put_file_block_size(self, tmp_path): - # block_size is passed to open() instead of the S3 API. + # block_size and max_workers are passed to open() instead of the S3 + # API. fs = self._make_fs() fs.open = mock.MagicMock() fs.open.return_value.__enter__.return_value.blocksize = 8 @@ -722,9 +1056,35 @@ def test_put_file_block_size(self, tmp_path): fs.put_file(str(lpath), "s3://bucket/key", block_size=8) fs.open.assert_called_once_with( - "s3://bucket/key", "wb", block_size=8, s3_additional_kwargs={} + "s3://bucket/key", + "wb", + block_size=8, + max_workers=fs.max_workers, + s3_additional_kwargs={}, ) + @pytest.mark.parametrize( + ("filesystem_kwargs", "kwargs", "expected"), + [ + ({}, {}, {"ContentType": "text/csv"}), + ({}, {"ContentType": "text/plain"}, {"ContentType": "text/plain"}), + # An explicit ContentType of the filesystem takes precedence over + # the one guessed from the file extension. + ({"ContentType": "application/octet-stream"}, {}, {}), + ], + ) + def test_put_file_content_type(self, tmp_path, filesystem_kwargs, kwargs, expected): + fs = self._make_fs() + fs.s3_additional_kwargs = filesystem_kwargs + fs.open = mock.MagicMock() + fs.open.return_value.__enter__.return_value.blocksize = 8 + lpath = tmp_path / "data.csv" + lpath.write_bytes(b"a") + + fs.put_file(str(lpath), "s3://bucket/key", **kwargs) + + assert fs.open.call_args.kwargs["s3_additional_kwargs"] == expected + @pytest.mark.parametrize( ("value", "kwargs"), [ @@ -2661,12 +3021,23 @@ def test_pandas_write_csv(self, line_count): class TestS3File: + @staticmethod + def _make_mock_fs(): + # A mocked filesystem that selects the request parameters of each + # operation as the real one does. + fs = mock.MagicMock(spec=S3FileSystem) + fs._client = S3_CLIENT + fs._get_operation_kwargs.side_effect = functools.partial( + S3FileSystem._get_operation_kwargs, fs + ) + return fs + @staticmethod def _make_write_file(data: bytes, autocommit: bool): # Build a minimal write-mode S3File without touching AWS, bypassing # __init__ which would require a real connection. file = S3File.__new__(S3File) - file.fs = mock.MagicMock(spec=S3FileSystem) + file.fs = TestS3File._make_mock_fs() file.path = "s3://bucket/key.txt" file.bucket = "bucket" file.key = "key.txt" @@ -2703,7 +3074,7 @@ def _make_append_fs(existing: bytes): # A mocked filesystem holding an existing object, with a minimum part # size of 4 bytes so that the write and append paths can be exercised # with tiny data and no AWS access. - fs = mock.MagicMock(spec=S3FileSystem) + fs = TestS3File._make_mock_fs() fs.MULTIPART_UPLOAD_MIN_PART_SIZE = 4 fs.MULTIPART_UPLOAD_MAX_PART_SIZE = 64 fs.MULTIPART_UPLOAD_MAX_PARTS = S3FileSystem.MULTIPART_UPLOAD_MAX_PARTS @@ -2939,6 +3310,46 @@ def test_write_exceeding_max_parts_without_close(self): assert f.closed executor.shutdown.assert_called_once() + def test_multipart_write_request_parameters(self): + # GH-946: the parts receive the parameters of the file that they + # accept, such as RequestPayer and SSE-C, and the completion receives + # them all. + fs = self._make_append_fs(b"") + kwargs = { + "ContentType": "text/csv", + "RequestPayer": "requester", + "SSECustomerAlgorithm": "AES256", + "SSECustomerKey": "key", + } + + with S3File( + fs, "s3://bucket/key.txt", mode="wb", block_size=4, s3_additional_kwargs=kwargs + ) as f: + f.write(b"x" * 8) + + fs._create_multipart_upload.assert_called_once_with( + bucket="bucket", key="key.txt", **kwargs + ) + assert fs._upload_part.call_count == 2 + for c in fs._upload_part.call_args_list: + assert {k: v for k, v in c.kwargs.items() if k[0].isupper()} == { + "RequestPayer": "requester", + "SSECustomerAlgorithm": "AES256", + "SSECustomerKey": "key", + } + assert fs._finish_multipart_upload.call_args.kwargs["request_kwargs"] == kwargs + + def test_multipart_write_keyword_named_as_argument(self): + # 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: + 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"} + def test_append_discard(self): # Rolling back an append aborts its multipart upload without the # existing object's metadata, which AbortMultipartUpload rejects, diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 933edd17a..59604bb97 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -205,7 +205,10 @@ def test_transaction_pipe_put_file(self, tmp_path, commit): # GH-977: pipe_file() and put_file() join the transaction of this # filesystem; they used to write through the internal S3FileSystem, # which is not in the transaction, and were not rolled back. - fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + # A real client selects the request parameters of each operation. + fs = AioS3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) put_object = fs._sync_fs._put_object = mock.MagicMock() local = tmp_path / "local.txt" local.write_bytes(b"local") @@ -258,8 +261,8 @@ def test_transaction_pipe_put_file_exceeding_max_parts(self, tmp_path, kwargs): fs._sync_fs._call.assert_not_called() def test_transaction_put_file_block_size(self, tmp_path): - # In a transaction, put_file() passes block_size to open() instead of - # the S3 API, as outside one. + # In a transaction, put_file() passes block_size and max_workers to + # open() instead of the S3 API, as outside one. fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) fs.open = mock.MagicMock() fs.open.return_value.__enter__.return_value.blocksize = 8 @@ -270,7 +273,11 @@ def test_transaction_put_file_block_size(self, tmp_path): fs.put_file(str(local), "s3://bucket/key", block_size=8) fs.open.assert_called_once_with( - "s3://bucket/key", "wb", block_size=8, s3_additional_kwargs={} + "s3://bucket/key", + "wb", + block_size=8, + max_workers=fs._sync_fs.max_workers, + s3_additional_kwargs={}, ) def test_touch_sync_wrapper(self): @@ -400,6 +407,102 @@ async def test_rm_maxdepth(self): (call,) = sync_fs._call.call_args_list assert call.kwargs["Delete"]["Objects"] == [{"Key": "dir"}, {"Key": "dir/a"}] + def test_put_file_in_transaction_open_parameters(self, tmp_path): + # GH-969: the open() parameters of put_file() go to open(), and the + # other parameters, also in s3_additional_kwargs, to S3. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs.open = mock.MagicMock() + fs.open.return_value.__enter__.return_value.blocksize = 4 + lpath = tmp_path / "data.csv" + lpath.write_bytes(b"a") + + fs._put_file_in_transaction( + str(lpath), + "s3://bucket/key", + Callback(), + block_size=S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + max_workers=2, + StorageClass="STANDARD_IA", + ) + + fs.open.assert_called_once_with( + "s3://bucket/key", + "wb", + block_size=S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, + max_workers=2, + s3_additional_kwargs={"StorageClass": "STANDARD_IA", "ContentType": "text/csv"}, + ) + + @pytest.mark.parametrize("size", [10, 5 * 2**30 + 1]) + @pytest.mark.asyncio + async def test_cp_file_multipart_parameters(self, size): + # GH-967: block_size and max_workers control a multipart copy and are + # not sent to S3, whatever the size of the object. + fs = AioS3FileSystem( + key="dummy", secret="dummy", region_name="us-east-1", skip_instance_cache=True + ) + fs._info = mock.AsyncMock( + return_value=S3Object( + init={"ContentLength": size}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="src", + ) + ) + sync_fs = fs._sync_fs + sync_fs._copy_object = mock.MagicMock() + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + running = [] + concurrency = [] + + def upload_part_copy(**kw): + running.append(kw["part_number"]) + concurrency.append(len(running)) + time.sleep(0.01) + running.remove(kw["part_number"]) + return SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + + sync_fs._upload_part_copy = mock.MagicMock(side_effect=upload_part_copy) + sync_fs._complete_multipart_upload = mock.MagicMock() + + await fs._cp_file( + "s3://bucket/src", + "s3://bucket/dst", + block_size=S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE // 2, + max_workers=1, + RequestPayer="requester", + ContentType="text/csv", + ) + + if size <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: + sync_fs._copy_object.assert_called_once_with( + bucket1="bucket", + key1="src", + version_id1=None, + bucket2="bucket", + key2="dst", + RequestPayer="requester", + ContentType="text/csv", + ) + else: + sync_fs._create_multipart_upload.assert_called_once_with( + bucket="bucket", key="dst", RequestPayer="requester", ContentType="text/csv" + ) + # The part copies receive the parameters that they accept, and + # max_workers limits how many run at once. + # Two parts, the second with the 1-byte tail. + assert sync_fs._upload_part_copy.call_count == 2 + assert all( + c.kwargs["RequestPayer"] == "requester" and "ContentType" not in c.kwargs + for c in sync_fs._upload_part_copy.call_args_list + ) + assert max(concurrency) == 1 + assert ( + sync_fs._complete_multipart_upload.call_args.kwargs["RequestPayer"] == "requester" + ) + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"):