From 19d79c4628f41041621aba9a3a800152e6e140e3 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 22:44:38 +0900 Subject: [PATCH 1/8] Send the lookup parameters of a file with its lookups S3File looked up the object while opening without the file's request parameters, so an object encrypted with a customer-provided key could not be opened, and RequestPayer and ExpectedBucketOwner were not sent with the lookups. info() and exists() now accept ExpectedBucketOwner, RequestPayer, and the SSECustomer* parameters, send them with the HeadObject, HeadBucket, and ListObjectsV2 requests that accept them, and use only the cached results of lookups with the same values. S3File, pipe_file(mode="create"), and cat_file() pass them to their lookups, and the append reads the existing object with the file's GetObject parameters. Closes #1004 Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 13 ++ pyathena/filesystem/s3.py | 258 +++++++++++++++++---- tests/pyathena/filesystem/test_s3.py | 166 ++++++++++++- tests/pyathena/filesystem/test_s3_async.py | 17 ++ 4 files changed, 404 insertions(+), 50 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index d69a230a6..daa9b6bc5 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -128,6 +128,19 @@ with fs.open("s3://YOUR_S3_BUCKET/path/to/data.csv", "wb", ContentType="text/csv f.write(b"col1\n1\n") ``` +`info` and `exists` accept `ExpectedBucketOwner`, `RequestPayer`, and the +`SSECustomer*` parameters of an object encrypted with a customer-provided key (SSE-C) +as keyword arguments, and ignore other request parameters. The lookups of the object +by `open`, and the existence check of `pipe` with `mode="create"`, send these +parameters of the file or the write. A lookup with them uses only the cached results of +lookups with the same values, not cached listings. + +```python +sse_c = {"SSECustomerAlgorithm": "AES256", "SSECustomerKey": key} +with fs.open("s3://YOUR_S3_BUCKET/path/to/encrypted.csv", "rb", **sse_c) as f: + data = f.read() +``` + ## 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 9943121f6..7c333bf8d 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -3,6 +3,7 @@ from __future__ import annotations import contextlib +import hashlib import logging import math import mimetypes @@ -46,6 +47,21 @@ _logger = logging.getLogger(__name__) +# The request parameters on which the authorization of a lookup (HeadObject, +# HeadBucket, or ListObjectsV2) depends. +_LOOKUP_REQUEST_PARAMETERS = frozenset( + { + "ExpectedBucketOwner", + "RequestPayer", + "SSECustomerAlgorithm", + "SSECustomerKey", + "SSECustomerKeyMD5", + } +) +# The second element of the dircache key, ``(path, _LOOKUPS_CACHE_KEY)``, of +# the lookup results of a path made with lookup request parameters. +_LOOKUPS_CACHE_KEY = "lookups" + class S3FileSystem(AbstractFileSystem): """A filesystem interface for Amazon S3 that implements the fsspec protocol. @@ -315,28 +331,41 @@ def _versioned_file_object(bucket: str, version: dict[str, Any]) -> S3Object: is_latest=version.get("IsLatest", False), ) - def _head_bucket(self, bucket, refresh: bool = False) -> S3Object | None: + def _head_bucket( + self, + bucket, + refresh: bool = False, + lookup_kwargs: Mapping[str, Any] | None = None, + ) -> S3Object | None: """Get the bucket as a directory object with HeadBucket. - The result is cached under the bucket name. A missing bucket evicts - its entry and the cached bucket listing that still lists it. + The result is cached under the bucket name, apart for each set of + lookup parameters (see ``_get_cached_lookup``). A missing bucket + evicts its entries and the cached bucket listing that still lists it. Args: bucket: The bucket name. refresh: If True, bypass the cache and call HeadBucket. + lookup_kwargs: The lookup parameters (see ``_get_lookup_kwargs``) + to send with the request. Returns: The bucket object, or None if the bucket does not exist. """ - file = None if refresh else self.dircache.get(bucket) + lookup_kwargs = lookup_kwargs or {} + file = None if refresh else self._get_cached_lookup(bucket, lookup_kwargs) if file is None: try: self._call( self._client.head_bucket, - Bucket=bucket, + **{ + **self._get_operation_kwargs("head_bucket", lookup_kwargs), + "Bucket": bucket, + }, ) except FileNotFoundError: self._evict_cache(bucket) + self._evict_cache((bucket, _LOOKUPS_CACHE_KEY)) # Evict the cached bucket listing only if it still lists the bucket. buckets = self.dircache.get("") if buckets and any(b.name == bucket for b in buckets): @@ -355,24 +384,31 @@ def _head_bucket(self, bucket, refresh: bool = False) -> S3Object | None: key=None, version_id=None, ) - self.dircache[bucket] = file + self._cache_lookup(bucket, lookup_kwargs, file) return file def _head_object( - self, path: str, version_id: str | None = None, refresh: bool = False + self, + path: str, + version_id: str | None = None, + refresh: bool = False, + lookup_kwargs: Mapping[str, Any] | None = None, ) -> S3Object | None: """Get the object with HeadObject. The result is cached under the path, or under the version-qualified - path for an explicit version. An explicitly requested ``"null"`` - version is not cached. A missing object evicts its entry and, unless - a version was requested, the cached listing of its parent that still - lists it. + path for an explicit version, apart for each set of lookup parameters + (see ``_get_cached_lookup``). An explicitly requested ``"null"`` + version is not cached. A missing object evicts its entries and, + unless a version was requested, the cached listing of its parent that + still lists it. Args: path: The object path, optionally with a versionId query. version_id: The version to get when the path has no version. refresh: If True, bypass the cache and call HeadObject. + lookup_kwargs: The lookup parameters (see ``_get_lookup_kwargs``) + to send with the request. Returns: The object, or None if it does not exist. @@ -387,7 +423,8 @@ def _head_object( # overwrite replaces the "null" version of a bucket without # versioning, so that version is looked up every time. cacheable = version_id != "null" - file = None if refresh else self.dircache.get(path) + lookup_kwargs = lookup_kwargs or {} + file = None if refresh else self._get_cached_lookup(path, lookup_kwargs) if file is None: try: request = { @@ -398,10 +435,11 @@ def _head_object( request.update({"VersionId": version_id}) response = self._call( self._client.head_object, - **request, + **{**self._get_operation_kwargs("head_object", lookup_kwargs), **request}, ) except FileNotFoundError: self._evict_cache(path) + self._evict_cache((path, _LOOKUPS_CACHE_KEY)) if not version_id: # Evict the cached listing of the parent only if it still # lists the path. @@ -422,7 +460,7 @@ def _head_object( version_id=version_id, ) if cacheable: - self.dircache[path] = file + self._cache_lookup(path, lookup_kwargs, file) return file def _ls_buckets(self, refresh: bool = False) -> list[S3Object]: @@ -659,13 +697,22 @@ def info(self, path: str, **kwargs) -> S3Object: version, the cached entries of the path are skipped, and the HeadObject result is cached under the version-qualified path apart from other versions, except for the ``null`` version, which an - overwrite replaces. + overwrite replaces. With request parameters on which the + authorization of the requests depends (``ExpectedBucketOwner``, + ``RequestPayer``, and the ``SSECustomer*`` parameters of an object + encrypted with a customer-provided key), each request receives those + that its operation accepts, and only the cached HeadObject or + HeadBucket results of lookups with the same values of these + parameters are used. Args: path: S3 path (e.g., "s3://bucket" or "s3://bucket/key"). **kwargs: Additional arguments including: refresh: If True, bypass the cache and query S3. version_id: The version ID to look up when the path has none. + ExpectedBucketOwner, RequestPayer, SSECustomerAlgorithm, + SSECustomerKey, SSECustomerKeyMD5: The request parameters + described above. Other request parameters are ignored. Returns: S3Object describing the bucket, directory, or file. The root path @@ -693,9 +740,12 @@ def info(self, path: str, **kwargs) -> S3Object: ) bucket, key, path_version_id = self.parse_path(path) version_id = path_version_id if path_version_id else kwargs.pop("version_id", None) - # Cached entries describe the current version of a path, so an - # explicit version uses only the HeadObject cache of that version. - if not refresh and not version_id: + lookup_kwargs = self._get_lookup_kwargs(kwargs) + # Cached entries describe the current version of a path as looked up + # without lookup parameters, so an explicit version or lookup + # parameters use only the HeadObject cache of that version and those + # parameters. + if not refresh and not version_id and not lookup_kwargs: caches: list[S3Object] | S3Object | None = self._ls_from_cache(path) if caches is not None: if isinstance(caches, list): @@ -727,21 +777,26 @@ def info(self, path: str, **kwargs) -> S3Object: bucket, key.rstrip("/") if key else None, version_id ) if key: - object_info = self._head_object(path, refresh=refresh, version_id=version_id) + object_info = self._head_object( + path, refresh=refresh, version_id=version_id, lookup_kwargs=lookup_kwargs + ) if object_info: return object_info else: - bucket_info = self._head_bucket(path, refresh=refresh) + bucket_info = self._head_bucket(path, refresh=refresh, lookup_kwargs=lookup_kwargs) if bucket_info: return bucket_info raise FileNotFoundError(path) response = self._call( self._client.list_objects_v2, - Bucket=bucket, - Prefix=f"{key.rstrip('/')}/" if key else "", - Delimiter="/", - MaxKeys=1, + **{ + **self._get_operation_kwargs("list_objects_v2", lookup_kwargs), + "Bucket": bucket, + "Prefix": f"{key.rstrip('/')}/" if key else "", + "Delimiter": "/", + "MaxKeys": 1, + }, ) if ( response.get("KeyCount", 0) > 0 @@ -927,6 +982,11 @@ def exists(self, path: str, **kwargs) -> bool: path: S3 path to check (e.g., "s3://bucket" or "s3://bucket/key"). **kwargs: Additional arguments including: refresh: If True, bypass the cache and query S3. + ExpectedBucketOwner, RequestPayer, SSECustomerAlgorithm, + SSECustomerKey, SSECustomerKeyMD5: The request parameters + on which the authorization of the requests depends, as + described in :meth:`info`. With them, cached listings are + not used. Other request parameters are ignored. Returns: True if the path exists, False otherwise. A bucket that HeadBucket @@ -943,18 +1003,21 @@ def exists(self, path: str, **kwargs) -> bool: # The root always exists. return True bucket, key, _ = self.parse_path(path) + lookup_kwargs = self._get_lookup_kwargs(kwargs) + # The cached listings are made without lookup parameters. + use_listings = not refresh and not lookup_kwargs if key: try: - if not refresh and self._ls_from_cache(path): + if use_listings and self._ls_from_cache(path): return True - info = self.info(path, refresh=refresh) + info = self.info(path, refresh=refresh, **lookup_kwargs) return bool(info) except FileNotFoundError: return False - if not refresh and self._ls_from_cache(bucket): + if use_listings and self._ls_from_cache(bucket): return True try: - file = self._head_bucket(bucket, refresh=refresh) + file = self._head_bucket(bucket, refresh=refresh, lookup_kwargs=lookup_kwargs) except PermissionError: # HeadBucket answers 403 for a bucket that exists but that the # caller may not access. @@ -1651,24 +1714,24 @@ def pipe_file( raise ValueError("Cannot write to the file with the version specified.") if not key: raise ValueError("Cannot write to a bucket.") + kwargs.pop("block_size", None) + kwargs.pop("max_workers", None) + request_kwargs = { + **self._get_operation_kwargs("put_object", self.s3_additional_kwargs), + **kwargs.pop("s3_additional_kwargs", {}), + **kwargs, + } if mode == "create": # Checked up front, as open() does in "xb" mode, and with # IfNoneMatch for an object created since. - if self.exists(path): + if self.exists(path, **self._get_lookup_kwargs(request_kwargs)): raise FileExistsError(path) - kwargs["IfNoneMatch"] = "*" + request_kwargs["IfNoneMatch"] = "*" if not isinstance(value, bytes): # Accept bytes-like values (bytearray, memoryview) as the # buffered path does. value = bytes(value) - kwargs.pop("block_size", None) - kwargs.pop("max_workers", None) - request_kwargs = { - **self._get_operation_kwargs("put_object", self.s3_additional_kwargs), - **kwargs.pop("s3_additional_kwargs", {}), - **kwargs, - } self._put_object(bucket=bucket, key=key, body=value, **request_kwargs) self.invalidate_cache(path) @@ -1744,7 +1807,7 @@ def cat_file( ``start`` without an ``end``, as a suffix range of the last bytes. Other negative offsets are resolved against the size from :meth:`info`, which also checks that the object exists for an empty - range. + range and receives the lookup parameters among ``kwargs``. Args: path: S3 path (s3://bucket/key) of the object. @@ -1781,7 +1844,7 @@ def cat_file( # A negative offset needs the size of the object, and an # empty range sends no GetObject request that would report a # missing object. - info = self.info(path, version_id=version_id) + info = self.info(path, version_id=version_id, **self._get_lookup_kwargs(kwargs)) if info.get("type") == S3ObjectType.S3_OBJECT_TYPE_DIRECTORY or info.key != key: # There is no object to read, as GetObject reports for # the other ranges, or info() describes the key without @@ -2399,8 +2462,15 @@ def invalidate_cache(self, path: str | None = None) -> None: for name in ("versionId", "versionID", "versionid", "version_id") ) for cache_path in cache_paths: - # _ls_dirs caches listings under (path, delimiter). - for cache_key in (cache_path, (cache_path, "/"), (cache_path, "")): + # _ls_dirs caches listings under (path, delimiter), and + # lookups with parameters are cached under + # (path, _LOOKUPS_CACHE_KEY). + for cache_key in ( + cache_path, + (cache_path, "/"), + (cache_path, ""), + (cache_path, _LOOKUPS_CACHE_KEY), + ): self._evict_cache(cache_key) # A version-qualified path continues with the path without # the version. @@ -2416,12 +2486,95 @@ def _evict_cache(self, key: str | tuple[str, str]) -> None: already gone, which this ignores. Args: - key: The dircache key, a path or a ``(path, delimiter)`` listing - key. + key: The dircache key, a path, a ``(path, delimiter)`` listing + key, or a ``(path, _LOOKUPS_CACHE_KEY)`` key. """ with contextlib.suppress(KeyError): del self.dircache[key] + @staticmethod + def _get_lookup_kwargs(kwargs: Mapping[str, Any]) -> dict[str, Any]: + """Select the request parameters that lookups send. + + Lookups (``info()`` and ``exists()``) send only the parameters on + which their authorization depends, which also select their cached + results. Other parameters, such as ``IfMatch`` or + ``ResponseContentType``, are not sent: they would change the result + that is cached for the path. + + Args: + kwargs: The parameters to select from. + + Returns: + The lookup parameters. + """ + return {k: v for k, v in kwargs.items() if k in _LOOKUP_REQUEST_PARAMETERS} + + def _get_cached_lookup(self, path: str, lookup_kwargs: Mapping[str, Any]) -> S3Object | None: + """Get the cached HeadObject or HeadBucket result of a lookup. + + A lookup without lookup parameters uses the entry under the path. A + lookup with them uses only the result of a lookup with the same + values, cached under ``(path, _LOOKUPS_CACHE_KEY)``, because the + authorization of the request depends on them. + + Args: + path: The path of the object, optionally version-qualified, or + the bucket name. + lookup_kwargs: The lookup parameters (see ``_get_lookup_kwargs``). + + Returns: + The cached object, or None if there is none. + """ + if not lookup_kwargs: + return cast("S3Object | None", self.dircache.get(path)) + lookups = self.dircache.get((path, _LOOKUPS_CACHE_KEY)) + return lookups.get(self._get_lookup_cache_id(lookup_kwargs)) if lookups else None + + def _cache_lookup(self, path: str, lookup_kwargs: Mapping[str, Any], file: S3Object) -> None: + """Cache the HeadObject or HeadBucket result of a lookup. + + Args: + path: The path of the object, optionally version-qualified, or + the bucket name. + lookup_kwargs: The lookup parameters (see ``_get_lookup_kwargs``) + that the request was made with. + file: The object to cache. + """ + if not lookup_kwargs: + self.dircache[path] = file + return + key = (path, _LOOKUPS_CACHE_KEY) + # A new dictionary, so that a concurrent reader does not see it + # change. + self.dircache[key] = { + **(self.dircache.get(key) or {}), + self._get_lookup_cache_id(lookup_kwargs): file, + } + + @staticmethod + def _get_lookup_cache_id(lookup_kwargs: Mapping[str, Any]) -> tuple[tuple[str, Any], ...]: + """Identify a set of lookup parameters in the cache. + + Args: + lookup_kwargs: The lookup parameters (see ``_get_lookup_kwargs``). + + Returns: + The sorted parameters, with ``SSECustomerKey`` replaced by its + SHA-256 digest so that the cache does not keep the key. + """ + return tuple( + sorted( + ( + k, + hashlib.sha256(v if isinstance(v, bytes) else str(v).encode()).hexdigest() + if k == "SSECustomerKey" + else v, + ) + for k, v in lookup_kwargs.items() + ) + ) + def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None: """Check the dircache for a cached entry of the path. @@ -2820,10 +2973,13 @@ def __init__( self._details: S3Object | dict[str, Any] = {} append_info: S3Object | None = None append_data: bytes | None = None + # The lookups need the parameters of the file on which their + # authorization depends, such as the customer-provided key. + lookup_kwargs = fs._get_lookup_kwargs(self.s3_additional_kwargs) if "r" in mode: # Looked up before the base class initializer, which would # otherwise take the size from the latest version of the object. - info = fs.info(path, version_id=self.version_id) + info = fs.info(path, version_id=self.version_id, **lookup_kwargs) if info.get("type") == S3ObjectType.S3_OBJECT_TYPE_DIRECTORY: # A prefix has no object to read. raise FileNotFoundError(path) @@ -2841,18 +2997,20 @@ def __init__( # The rewritten object keeps the metadata of the existing one, # which a cached listing entry lacks, so look up the object. with contextlib.suppress(FileNotFoundError): - append_info = fs.info(path, refresh=True) + append_info = fs.info(path, refresh=True, **lookup_kwargs) if ( append_info is not None and append_info.get("size", 0) < fs.MULTIPART_UPLOAD_MIN_PART_SIZE ): # Too small to be a part of a multipart upload: rewritten # from the buffer. - append_data = fs.cat(path) + append_data = fs.cat_file( + path, **fs._get_operation_kwargs("get_object", self.s3_additional_kwargs) + ) elif "x" in mode: # Checked up front so that no data is uploaded for an existing # object, and on commit with IfNoneMatch for one created since. - if fs.exists(path): + if fs.exists(path, **lookup_kwargs): raise FileExistsError(path) self.s3_additional_kwargs.update({"IfNoneMatch": "*"}) @@ -2937,7 +3095,11 @@ def _initiate_upload(self) -> None: ) if self.append_block: if self.tell() > self.fs.MULTIPART_UPLOAD_MAX_PART_SIZE: - info = self.fs.info(self.path, version_id=self.version_id) + info = self.fs.info( + self.path, + version_id=self.version_id, + **self.fs._get_lookup_kwargs(self.s3_additional_kwargs), + ) ranges = self.fs._get_copy_ranges( # Set copy source file byte size info.get("size", 0), diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 5f2efbefd..33f8cc1d0 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1685,7 +1685,7 @@ def test_open_append_keeps_metadata_of_listed_object(self): "ContentType": "text/plain", "Metadata": {"k": "v"}, } - fs.cat = mock.MagicMock(return_value=b"aa") + fs.cat_file = mock.MagicMock(return_value=b"aa") fs._put_object = mock.MagicMock() with fs.open("s3://bucket/key", "ab") as f: @@ -1720,6 +1720,167 @@ def info(path, **kwargs): assert unraisable == [] fs._call.assert_not_called() + LOOKUP_KWARGS = { + "ExpectedBucketOwner": "111122223333", + "RequestPayer": "requester", + "SSECustomerAlgorithm": "AES256", + "SSECustomerKey": "k" * 32, + } + + @staticmethod + def _record_lookups(fs, exists=True): + # Record the S3 requests of the filesystem by operation name, with + # an object of 2 bytes at every key if it exists. + requests = [] + + def call(method, **request): + name = method._extract_mock_name().split(".")[-1] + requests.append((name, request)) + if name == "head_object": + if not exists: + raise FileNotFoundError(request["Key"]) + return {"ContentLength": 2, "ETag": '"e"'} + if name == "get_object": + return {"Body": io.BytesIO(b"aa")} + return {"UploadId": "uploadid", "ETag": '"e"'} + + fs._call.side_effect = call + return requests + + @pytest.mark.parametrize("mode", ["rb", "ab", "xb"]) + def test_open_lookup_parameters(self, mode): + # GH-1004: the lookups made while opening a file did not send its + # parameters, so an object encrypted with a customer-provided key, or + # in a requester-pays bucket, could not be opened. + fs = self._make_fs() + fs.default_cache_type = "bytes" + requests = self._record_lookups(fs, exists=mode != "xb") + + with fs.open("s3://bucket/key", mode, ContentType="text/plain", **self.LOOKUP_KWARGS) as f: + if mode == "rb": + assert f.read() == b"aa" + else: + f.write(b"bb") + + lookups = {name: request for name, request in requests if name != "put_object"} + # Only the parameters on which the authorization of a lookup depends. + assert lookups.pop("head_object") == { + "Bucket": "bucket", + "Key": "key", + **self.LOOKUP_KWARGS, + } + if mode == "xb": + assert lookups.pop("list_objects_v2") == { + "Bucket": "bucket", + "Prefix": "key/", + "Delimiter": "/", + "MaxKeys": 1, + "ExpectedBucketOwner": "111122223333", + "RequestPayer": "requester", + } + else: + assert self.LOOKUP_KWARGS.items() <= lookups.pop("get_object").items() + assert lookups == {} + + def test_info_lookup_parameters_cache(self): + # GH-1004: a cached result serves only lookups with the same lookup + # parameters, on which the authorization of the requests depends. + fs = self._make_fs() + fs._call.return_value = {"ContentLength": 2, "ETag": '"e"'} + fs.dircache[("bucket", "/")] = [self._file_object("key")] + path = "s3://bucket/key" + other_key = {**self.LOOKUP_KWARGS, "SSECustomerKey": "j" * 32} + + for _ in range(2): + assert fs.info(path).size == 0 + assert fs.info(path, **self.LOOKUP_KWARGS).size == 2 + assert fs.info(path, IfMatch='"x"', **other_key).size == 2 + assert fs.exists(path, **self.LOOKUP_KWARGS) + # The listing serves only the lookups without the parameters, and the + # other parameters are not sent. + assert [c.kwargs for c in fs._call.call_args_list] == [ + {"Bucket": "bucket", "Key": "key", **self.LOOKUP_KWARGS}, + {"Bucket": "bucket", "Key": "key", **other_key}, + ] + # The cache does not keep the customer-provided keys. + assert "k" * 32 not in repr(fs.dircache) + assert "j" * 32 not in repr(fs.dircache) + + fs.invalidate_cache(path) + fs.info(path, **self.LOOKUP_KWARGS) + assert fs._call.call_count == 3 + + def test_info_lookup_parameters_missing_object(self): + # GH-1004: a missing object evicts the cached results of the lookups + # with parameters, and the request that checks for a key prefix + # receives those that ListObjectsV2 accepts. + fs = self._make_fs() + fs._call.side_effect = [ + {"ContentLength": 2, "ETag": '"e"'}, + FileNotFoundError("key"), + {"KeyCount": 0}, + ] + path = "s3://bucket/key" + + fs.info(path, **self.LOOKUP_KWARGS) + with pytest.raises(FileNotFoundError): + fs.info(path, refresh=True, **self.LOOKUP_KWARGS) + + assert fs.dircache == {} + assert fs._call.call_args.kwargs == { + "Bucket": "bucket", + "Prefix": "key/", + "Delimiter": "/", + "MaxKeys": 1, + "ExpectedBucketOwner": "111122223333", + "RequestPayer": "requester", + } + + def test_exists_bucket_lookup_parameters(self): + # GH-1004: a bucket lookup with parameters uses neither the cached + # bucket listing nor the result of a lookup without them. + fs = self._make_fs() + fs._call.return_value = {} + fs.dircache[""] = [fs._directory_object("bucket", None)] + + for _ in range(2): + assert fs.exists("s3://bucket") + assert fs.exists("s3://bucket", **self.LOOKUP_KWARGS) + assert fs.info("s3://bucket", **self.LOOKUP_KWARGS).type == ( + S3ObjectType.S3_OBJECT_TYPE_DIRECTORY + ) + fs._call.assert_called_once_with( + fs._client.head_bucket, Bucket="bucket", ExpectedBucketOwner="111122223333" + ) + + def test_pipe_file_create_lookup_parameters(self): + # GH-1004: the existence check of pipe_file(mode="create") sends the + # lookup parameters of the write, as open() does in "xb" mode. + fs = self._make_fs() + requests = self._record_lookups(fs, exists=False) + + fs.pipe_file("s3://bucket/key", b"a", mode="create", **self.LOOKUP_KWARGS) + + assert requests[0] == ( + "head_object", + {"Bucket": "bucket", "Key": "key", **self.LOOKUP_KWARGS}, + ) + assert requests[-1][0] == "put_object" + assert requests[-1][1]["IfNoneMatch"] == "*" + + def test_cat_file_range_lookup_parameters(self): + # GH-1004: the lookup that resolves negative offsets sends the lookup + # parameters of the read. + fs = self._make_fs() + requests = self._record_lookups(fs) + + assert fs.cat_file("s3://bucket/key", start=-2, end=-1, **self.LOOKUP_KWARGS) == b"aa" + + assert [(name, self.LOOKUP_KWARGS.items() <= r.items()) for name, r in requests] == [ + ("head_object", True), + ("get_object", True), + ] + @pytest.mark.parametrize( "block_size", [S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE, S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE], @@ -3806,6 +3967,7 @@ def _make_mock_fs(): fs._get_operation_kwargs.side_effect = functools.partial( S3FileSystem._get_operation_kwargs, fs ) + fs._get_lookup_kwargs.side_effect = S3FileSystem._get_lookup_kwargs return fs @staticmethod @@ -3860,7 +4022,7 @@ def _make_append_fs(existing: bytes): bucket="bucket", key="key.txt", ) - fs.cat.return_value = existing + fs.cat_file.return_value = existing fs._create_multipart_upload.return_value = SimpleNamespace(upload_id="uploadid") def part(**kw): diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index e9990dc2f..cf7b4f811 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -1437,6 +1437,23 @@ def test_open_version_id(self): assert f.size == 4 fs._sync_fs.info.assert_called_once_with("bucket/key?versionId=v1", version_id="v1") + def test_open_lookup_parameters(self): + # GH-1004: the lookup of the file sends its lookup parameters. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._sync_fs.info = mock.MagicMock( + return_value=S3Object( + init={"ContentLength": 4}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="key", + ) + ) + sse_c = {"SSECustomerAlgorithm": "AES256", "SSECustomerKey": "k" * 32} + + with fs.open("s3://bucket/key", "rb", ContentType="text/plain", **sse_c) as f: + assert isinstance(f, AioS3File) + fs._sync_fs.info.assert_called_once_with("bucket/key", version_id=None, **sse_c) + @pytest.mark.parametrize( ("objects", "target"), [ From fa36dc124ad2ea9f81fcfb7b4b713993ff13a163 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 22:55:12 +0900 Subject: [PATCH 2/8] Reuse the request recorder for the lookup tests Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/filesystem/test_s3.py | 37 +++++++++------------------- 1 file changed, 12 insertions(+), 25 deletions(-) diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 33f8cc1d0..c4c93320a 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1091,8 +1091,9 @@ def test_put_file_open_parameters(self, tmp_path): ) @staticmethod - def _record_requests(fs, precondition_failed=False): - # Record the S3 requests of the filesystem by operation name. With + def _record_requests(fs, precondition_failed=False, exists=True): + # Record the S3 requests of the filesystem by operation name, with an + # object of 2 bytes at every key if it exists. With # precondition_failed, the conditional writes fail as S3 fails them # when an object exists. requests = [] @@ -1101,6 +1102,12 @@ 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 == "head_object": + if not exists: + raise FileNotFoundError(request["Key"]) + return {"ContentLength": 2, "ETag": '"e"'} + if name == "get_object": + return {"Body": io.BytesIO(b"aa")} if precondition_failed and name in {"put_object", "complete_multipart_upload"}: error = botocore.exceptions.ClientError( { @@ -1727,26 +1734,6 @@ def info(path, **kwargs): "SSECustomerKey": "k" * 32, } - @staticmethod - def _record_lookups(fs, exists=True): - # Record the S3 requests of the filesystem by operation name, with - # an object of 2 bytes at every key if it exists. - requests = [] - - def call(method, **request): - name = method._extract_mock_name().split(".")[-1] - requests.append((name, request)) - if name == "head_object": - if not exists: - raise FileNotFoundError(request["Key"]) - return {"ContentLength": 2, "ETag": '"e"'} - if name == "get_object": - return {"Body": io.BytesIO(b"aa")} - return {"UploadId": "uploadid", "ETag": '"e"'} - - fs._call.side_effect = call - return requests - @pytest.mark.parametrize("mode", ["rb", "ab", "xb"]) def test_open_lookup_parameters(self, mode): # GH-1004: the lookups made while opening a file did not send its @@ -1754,7 +1741,7 @@ def test_open_lookup_parameters(self, mode): # in a requester-pays bucket, could not be opened. fs = self._make_fs() fs.default_cache_type = "bytes" - requests = self._record_lookups(fs, exists=mode != "xb") + requests = self._record_requests(fs, exists=mode != "xb") with fs.open("s3://bucket/key", mode, ContentType="text/plain", **self.LOOKUP_KWARGS) as f: if mode == "rb": @@ -1857,7 +1844,7 @@ def test_pipe_file_create_lookup_parameters(self): # GH-1004: the existence check of pipe_file(mode="create") sends the # lookup parameters of the write, as open() does in "xb" mode. fs = self._make_fs() - requests = self._record_lookups(fs, exists=False) + requests = self._record_requests(fs, exists=False) fs.pipe_file("s3://bucket/key", b"a", mode="create", **self.LOOKUP_KWARGS) @@ -1872,7 +1859,7 @@ def test_cat_file_range_lookup_parameters(self): # GH-1004: the lookup that resolves negative offsets sends the lookup # parameters of the read. fs = self._make_fs() - requests = self._record_lookups(fs) + requests = self._record_requests(fs) assert fs.cat_file("s3://bucket/key", start=-2, end=-1, **self.LOOKUP_KWARGS) == b"aa" From 04e1f7e624ec65703a9288791dcbe1ad4470950a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 22:56:22 +0900 Subject: [PATCH 3/8] Narrow the lookup parameter notes Co-Authored-By: Claude Opus 5.5 --- docs/filesystem.md | 2 +- pyathena/filesystem/s3.py | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/docs/filesystem.md b/docs/filesystem.md index daa9b6bc5..c65b77850 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -136,7 +136,7 @@ parameters of the file or the write. A lookup with them uses only the cached res lookups with the same values, not cached listings. ```python -sse_c = {"SSECustomerAlgorithm": "AES256", "SSECustomerKey": key} +sse_c = {"SSECustomerAlgorithm": "AES256", "SSECustomerKey": YOUR_32_BYTE_KEY} with fs.open("s3://YOUR_S3_BUCKET/path/to/encrypted.csv", "rb", **sse_c) as f: data = f.read() ``` diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 7c333bf8d..bc61ac02f 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -2498,9 +2498,9 @@ def _get_lookup_kwargs(kwargs: Mapping[str, Any]) -> dict[str, Any]: Lookups (``info()`` and ``exists()``) send only the parameters on which their authorization depends, which also select their cached - results. Other parameters, such as ``IfMatch`` or - ``ResponseContentType``, are not sent: they would change the result - that is cached for the path. + results. Other parameters that HeadObject accepts, such as + ``IfMatch``, ``PartNumber`` or ``ResponseContentType``, are not sent, + because the result would depend on them. Args: kwargs: The parameters to select from. From 5d970c3a8abe73b8226243cf554275481d96d88b Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 23:04:13 +0900 Subject: [PATCH 4/8] Read the whole object for an append and update the lookups in place The append read the existing object with every GetObject parameter of the file, so a Range among them truncated the object. It now sends only the lookup parameters. A lookup cached with some parameters wrote back a copy of the dictionary, which could replace a result cached in the meantime for other parameters; the dictionary is now updated in place. Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 20 +++++++------- tests/pyathena/filesystem/test_s3.py | 41 ++++++++++++++++++++++++++-- 2 files changed, 49 insertions(+), 12 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index bc61ac02f..5b3a013fb 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -2545,12 +2545,13 @@ def _cache_lookup(self, path: str, lookup_kwargs: Mapping[str, Any], file: S3Obj self.dircache[path] = file return key = (path, _LOOKUPS_CACHE_KEY) - # A new dictionary, so that a concurrent reader does not see it - # change. - self.dircache[key] = { - **(self.dircache.get(key) or {}), - self._get_lookup_cache_id(lookup_kwargs): file, - } + lookups = self.dircache.get(key) + if lookups is None: + lookups = {} + self.dircache[key] = lookups + # Updated in place: a copy written back could replace a result that + # another thread has cached since for other parameters. + lookups[self._get_lookup_cache_id(lookup_kwargs)] = file @staticmethod def _get_lookup_cache_id(lookup_kwargs: Mapping[str, Any]) -> tuple[tuple[str, Any], ...]: @@ -3003,10 +3004,9 @@ def __init__( and append_info.get("size", 0) < fs.MULTIPART_UPLOAD_MIN_PART_SIZE ): # Too small to be a part of a multipart upload: rewritten - # from the buffer. - append_data = fs.cat_file( - path, **fs._get_operation_kwargs("get_object", self.s3_additional_kwargs) - ) + # from the buffer. Only the lookup parameters are sent, so + # that the whole object is read. + append_data = fs.cat_file(path, **lookup_kwargs) elif "x" in mode: # Checked up front so that no data is uploaded for an existing # object, and on commit with IfNoneMatch for one created since. diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index c4c93320a..b3250c48c 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1743,7 +1743,13 @@ def test_open_lookup_parameters(self, mode): fs.default_cache_type = "bytes" requests = self._record_requests(fs, exists=mode != "xb") - with fs.open("s3://bucket/key", mode, ContentType="text/plain", **self.LOOKUP_KWARGS) as f: + with fs.open( + "s3://bucket/key", + mode, + ContentType="text/plain", + Range="bytes=0-0", + **self.LOOKUP_KWARGS, + ) as f: if mode == "rb": assert f.read() == b"aa" else: @@ -1766,9 +1772,40 @@ def test_open_lookup_parameters(self, mode): "RequestPayer": "requester", } else: - assert self.LOOKUP_KWARGS.items() <= lookups.pop("get_object").items() + get_object = lookups.pop("get_object") + assert self.LOOKUP_KWARGS.items() <= get_object.items() + if mode == "ab": + # The whole existing object is read. + assert "Range" not in get_object assert lookups == {} + def test_cache_lookup_concurrent_parameters(self): + # GH-1004: caching a lookup with some parameters does not replace a + # result cached in the meantime for other parameters. + fs = self._make_fs() + path = "bucket/key" + lookup_kwargs = self.LOOKUP_KWARGS + other_key = {**lookup_kwargs, "SSECustomerKey": "j" * 32} + stale, fresh, other = (fs._directory_object("bucket", "key") for _ in range(3)) + fs._cache_lookup(path, lookup_kwargs, stale) + + class InterleavedDict(dict): + interleaved = False + + def get(self, key, default=None): + value = super().get(key, default) + if not self.interleaved: + # Another thread refreshes the lookup in between. + self.interleaved = True + fs._cache_lookup(path, lookup_kwargs, fresh) + return value + + fs.dircache = InterleavedDict(fs.dircache) + fs._cache_lookup(path, other_key, other) + + assert fs._get_cached_lookup(path, lookup_kwargs) is fresh + assert fs._get_cached_lookup(path, other_key) is other + def test_info_lookup_parameters_cache(self): # GH-1004: a cached result serves only lookups with the same lookup # parameters, on which the authorization of the requests depends. From 01b8b6a88d109b9ef80ff49da7de131b2d3e63dc Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 23:12:51 +0900 Subject: [PATCH 5/8] Renew the expiry time of the cached lookups when one is cached Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 3 ++- tests/pyathena/filesystem/test_s3.py | 17 +++++++++++++++++ 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 5b3a013fb..c46782840 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -2548,10 +2548,11 @@ def _cache_lookup(self, path: str, lookup_kwargs: Mapping[str, Any], file: S3Obj lookups = self.dircache.get(key) if lookups is None: lookups = {} - self.dircache[key] = lookups # Updated in place: a copy written back could replace a result that # another thread has cached since for other parameters. lookups[self._get_lookup_cache_id(lookup_kwargs)] = file + # Set again so that the expiry time of the entry is renewed. + self.dircache[key] = lookups @staticmethod def _get_lookup_cache_id(lookup_kwargs: Mapping[str, Any]) -> tuple[tuple[str, Any], ...]: diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index b3250c48c..2a0a5b861 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1806,6 +1806,23 @@ def get(self, key, default=None): assert fs._get_cached_lookup(path, lookup_kwargs) is fresh assert fs._get_cached_lookup(path, other_key) is other + def test_cache_lookup_renews_expiry(self, monkeypatch): + # GH-1004: caching a lookup renews the expiry time of the cached + # lookups of the path, as caching other entries does. + fs = self._make_fs() + fs.dircache = DirCache(listings_expiry_time=60) + now = [0.0] + monkeypatch.setattr("fsspec.dircache.time.time", lambda: now[0]) + path = "bucket/key" + stale, fresh = (fs._directory_object("bucket", "key") for _ in range(2)) + + fs._cache_lookup(path, self.LOOKUP_KWARGS, stale) + now[0] = 59.0 + fs._cache_lookup(path, self.LOOKUP_KWARGS, fresh) + now[0] = 61.0 + + assert fs._get_cached_lookup(path, self.LOOKUP_KWARGS) is fresh + def test_info_lookup_parameters_cache(self): # GH-1004: a cached result serves only lookups with the same lookup # parameters, on which the authorization of the requests depends. From cc9d5bdb9878d4c279133289e5c0b7f81d0c84dc Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 23:20:28 +0900 Subject: [PATCH 6/8] Expire each cached lookup with parameters on its own Caching the lookup of one set of parameters renews the dircache entry of the path, which extended the other sets' results beyond listings_expiry_time. Each result now keeps the time it was cached and expires on its own. Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 18 +++++++++++++++--- tests/pyathena/filesystem/test_s3.py | 24 +++++++++++++++++------- 2 files changed, 32 insertions(+), 10 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index c46782840..e311c7edf 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -9,6 +9,7 @@ import mimetypes import os.path import re +import time from collections.abc import Callable, Iterator, Mapping from concurrent.futures import Future, as_completed, wait from copy import deepcopy @@ -2516,7 +2517,9 @@ def _get_cached_lookup(self, path: str, lookup_kwargs: Mapping[str, Any]) -> S3O A lookup without lookup parameters uses the entry under the path. A lookup with them uses only the result of a lookup with the same values, cached under ``(path, _LOOKUPS_CACHE_KEY)``, because the - authorization of the request depends on them. + authorization of the request depends on them. Each of these results + expires on its own after the ``listings_expiry_time`` of the + dircache. Args: path: The path of the object, optionally version-qualified, or @@ -2529,7 +2532,16 @@ def _get_cached_lookup(self, path: str, lookup_kwargs: Mapping[str, Any]) -> S3O if not lookup_kwargs: return cast("S3Object | None", self.dircache.get(path)) lookups = self.dircache.get((path, _LOOKUPS_CACHE_KEY)) - return lookups.get(self._get_lookup_cache_id(lookup_kwargs)) if lookups else None + cached = lookups.get(self._get_lookup_cache_id(lookup_kwargs)) if lookups else None + if cached is None: + return None + cached_at, file = cached + expiry_time = self.dircache.listings_expiry_time + if expiry_time and time.time() - cached_at > expiry_time: + # Caching the result of any parameters renews the dircache entry + # of the path, so each result expires on its own. + return None + return cast(S3Object, file) def _cache_lookup(self, path: str, lookup_kwargs: Mapping[str, Any], file: S3Object) -> None: """Cache the HeadObject or HeadBucket result of a lookup. @@ -2550,7 +2562,7 @@ def _cache_lookup(self, path: str, lookup_kwargs: Mapping[str, Any], file: S3Obj lookups = {} # Updated in place: a copy written back could replace a result that # another thread has cached since for other parameters. - lookups[self._get_lookup_cache_id(lookup_kwargs)] = file + lookups[self._get_lookup_cache_id(lookup_kwargs)] = (time.time(), file) # Set again so that the expiry time of the entry is renewed. self.dircache[key] = lookups diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2a0a5b861..e33b82001 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -154,7 +154,7 @@ def _make_fs(): # Build a minimal S3FileSystem without touching AWS, bypassing # __init__ which would require a boto3 client. fs = S3FileSystem.__new__(S3FileSystem) - fs.dircache = {} + fs.dircache = 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 @@ -1789,7 +1789,7 @@ def test_cache_lookup_concurrent_parameters(self): stale, fresh, other = (fs._directory_object("bucket", "key") for _ in range(3)) fs._cache_lookup(path, lookup_kwargs, stale) - class InterleavedDict(dict): + class InterleavedDirCache(DirCache): interleaved = False def get(self, key, default=None): @@ -1800,15 +1800,17 @@ def get(self, key, default=None): fs._cache_lookup(path, lookup_kwargs, fresh) return value - fs.dircache = InterleavedDict(fs.dircache) + cache = InterleavedDirCache() + cache.update(fs.dircache) + fs.dircache = cache fs._cache_lookup(path, other_key, other) assert fs._get_cached_lookup(path, lookup_kwargs) is fresh assert fs._get_cached_lookup(path, other_key) is other - def test_cache_lookup_renews_expiry(self, monkeypatch): - # GH-1004: caching a lookup renews the expiry time of the cached - # lookups of the path, as caching other entries does. + def test_cache_lookup_expiry(self, monkeypatch): + # GH-1004: the cached lookups with parameters expire after the + # listings_expiry_time of the dircache, each on its own. fs = self._make_fs() fs.dircache = DirCache(listings_expiry_time=60) now = [0.0] @@ -1816,13 +1818,21 @@ def test_cache_lookup_renews_expiry(self, monkeypatch): path = "bucket/key" stale, fresh = (fs._directory_object("bucket", "key") for _ in range(2)) + other_key = {**self.LOOKUP_KWARGS, "SSECustomerKey": "j" * 32} + fs._cache_lookup(path, self.LOOKUP_KWARGS, stale) now[0] = 59.0 fs._cache_lookup(path, self.LOOKUP_KWARGS, fresh) now[0] = 61.0 - assert fs._get_cached_lookup(path, self.LOOKUP_KWARGS) is fresh + # Each result expires on its own, also while the results of other + # parameters keep renewing the entry of the path. + fs._cache_lookup(path, other_key, stale) + now[0] = 120.0 + assert fs._get_cached_lookup(path, other_key) is stale + assert fs._get_cached_lookup(path, self.LOOKUP_KWARGS) is None + def test_info_lookup_parameters_cache(self): # GH-1004: a cached result serves only lookups with the same lookup # parameters, on which the authorization of the requests depends. From b175d8bf0e86be2bc66f3125eb2d972f128d6f5e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 23:27:08 +0900 Subject: [PATCH 7/8] Inspect the cached entries for customer-provided keys Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/filesystem/test_s3.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index e33b82001..ba3c7d781 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1854,8 +1854,8 @@ def test_info_lookup_parameters_cache(self): {"Bucket": "bucket", "Key": "key", **other_key}, ] # The cache does not keep the customer-provided keys. - assert "k" * 32 not in repr(fs.dircache) - assert "j" * 32 not in repr(fs.dircache) + assert "k" * 32 not in repr(dict(fs.dircache)) + assert "j" * 32 not in repr(dict(fs.dircache)) fs.invalidate_cache(path) fs.info(path, **self.LOOKUP_KWARGS) From 07c1a1418e5ccd915dec5b5db332771f6446de23 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 00:37:30 +0900 Subject: [PATCH 8/8] Cache each version under a single spelling of its query A version looked up with ?version_id= and with ?versionId= was cached twice, and a missing version evicted only the spelling that looked it up, so the other spelling kept returning the deleted version. Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 11 +++++++---- tests/pyathena/filesystem/test_s3.py | 24 ++++++++++++++++++++++++ 2 files changed, 31 insertions(+), 4 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index e311c7edf..f4ba43eda 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -398,7 +398,8 @@ def _head_object( """Get the object with HeadObject. The result is cached under the path, or under the version-qualified - path for an explicit version, apart for each set of lookup parameters + path (spelled ``?versionId=``) for an explicit version, apart for each + set of lookup parameters (see ``_get_cached_lookup``). An explicitly requested ``"null"`` version is not cached. A missing object evicts its entries and, unless a version was requested, the cached listing of its parent that @@ -416,10 +417,12 @@ def _head_object( """ bucket, key, path_version_id = self.parse_path(path) version_id = path_version_id if path_version_id else version_id - if version_id and not path_version_id: + if version_id: # Cache an explicit version under its version-qualified path so - # that it neither reuses nor replaces the entry of another version. - path = f"{path}?versionId={version_id}" + # that it neither reuses nor replaces the entry of another version, + # with a single spelling of the query so that a missing version + # evicts the entry whatever spelling looked it up. + path = f"{path.partition('?')[0]}?versionId={version_id}" # Writes invalidate only the path without the version, and an # overwrite replaces the "null" version of a bucket without # versioning, so that version is looked up every time. diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index ba3c7d781..31d974fac 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -2535,6 +2535,30 @@ def test_info_caches_each_version_separately(self, version_aware): # The second round is served from the cache. assert fs._call.call_count == 3 + @pytest.mark.parametrize("lookup", [False, True]) + def test_info_version_spellings_share_cache(self, lookup): + # A version looked up with any spelling of the query is cached once, + # so a missing version evicts it for every spelling. + fs = self._make_fs() + kwargs = self.LOOKUP_KWARGS if lookup else {} + fs._call.side_effect = [ + {"ContentLength": 4, "ETag": '"etag"', "VersionId": "v1"}, + FileNotFoundError("key"), + {"KeyCount": 0}, + FileNotFoundError("key"), + {"KeyCount": 0}, + ] + + assert fs.info("s3://bucket/key?versionId=v1", **kwargs).size == 4 + assert fs.info("s3://bucket/key?version_id=v1", **kwargs).size == 4 + assert fs.info("s3://bucket/key", version_id="v1", **kwargs).size == 4 + assert fs._call.call_count == 1 + with pytest.raises(FileNotFoundError): + fs.info("s3://bucket/key?version_id=v1", refresh=True, **kwargs) + with pytest.raises(FileNotFoundError): + fs.info("s3://bucket/key?versionId=v1", **kwargs) + assert fs._call.call_count == 5 + def test_info_does_not_cache_null_version(self): fs = self._make_fs() fs._call.return_value = {"ContentLength": 4, "ETag": '"etag"', "VersionId": "null"}