diff --git a/docs/filesystem.md b/docs/filesystem.md index d69a230a6..c65b77850 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": YOUR_32_BYTE_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..f4ba43eda 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -3,11 +3,13 @@ from __future__ import annotations import contextlib +import hashlib import logging import math 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 @@ -46,6 +48,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 +332,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,39 +385,50 @@ 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 (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 + 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. """ 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. 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 +439,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 +464,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 +701,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 +744,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 +781,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 +986,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 +1007,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 +1718,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 +1811,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 +1848,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 +2466,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 +2490,108 @@ 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 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. + + 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. 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 + 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)) + 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. + + 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) + lookups = self.dircache.get(key) + if lookups is None: + 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)] = (time.time(), 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], ...]: + """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 +2990,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 +3014,19 @@ 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) + # 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. - if fs.exists(path): + if fs.exists(path, **lookup_kwargs): raise FileExistsError(path) self.s3_additional_kwargs.update({"IfNoneMatch": "*"}) @@ -2937,7 +3111,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..31d974fac 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 @@ -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( { @@ -1685,7 +1692,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 +1727,211 @@ def info(path, **kwargs): assert unraisable == [] fs._call.assert_not_called() + LOOKUP_KWARGS = { + "ExpectedBucketOwner": "111122223333", + "RequestPayer": "requester", + "SSECustomerAlgorithm": "AES256", + "SSECustomerKey": "k" * 32, + } + + @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_requests(fs, exists=mode != "xb") + + 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: + 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: + 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 InterleavedDirCache(DirCache): + 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 + + 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_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] + monkeypatch.setattr("fsspec.dircache.time.time", lambda: now[0]) + 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. + 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(dict(fs.dircache)) + assert "j" * 32 not in repr(dict(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_requests(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_requests(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], @@ -2323,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"} @@ -3806,6 +4042,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 +4097,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"), [