diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 4104f9354..940a2649e 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -315,6 +315,11 @@ def _head_bucket(self, bucket, refresh: bool = False) -> S3Object | None: Bucket=bucket, ) except FileNotFoundError: + self.dircache.pop(bucket, None) + # 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): + self.dircache.pop("", None) return None file = S3Object( init={ @@ -352,6 +357,7 @@ def _head_object( **request, ) except FileNotFoundError: + self.dircache.pop(path, None) return None if self.version_aware and not version_id: # Pin the version of the object so that subsequent reads see @@ -404,13 +410,34 @@ def _ls_dirs( max_keys: int | None = None, refresh: bool = False, ) -> list[S3Object]: + """List the objects and common prefixes under a path. + + A complete, non-empty listing of the path is cached under + ``(path, delimiter)``, and an empty one evicts it. + ``invalidate_cache`` drops it when the path or a path under it is + invalidated. + + Args: + path: The bucket or directory path to list. + prefix: Key prefix to filter by, relative to the path. A prefixed + listing is neither read from nor written to the cache. + delimiter: Delimiter to group keys by; ``""`` lists recursively. + next_token: Continuation token to start listing from. A listing + that starts from a token is neither read from nor written to + the cache. + max_keys: Maximum number of keys per ListObjectsV2 request. + refresh: If True, bypass the cache and list from S3. + + Returns: + The listed directories and files. + """ bucket, key, version_id = self.parse_path(path) + use_cache = not prefix and not next_token if key: prefix = f"{key}/{prefix if prefix else ''}" - # Create a cache key that includes the delimiter cache_key = (path, delimiter) - if cache_key in self.dircache and not refresh: + if use_cache and cache_key in self.dircache and not refresh: return cast(list[S3Object], self.dircache[cache_key]) files: list[S3Object] = [] @@ -444,8 +471,11 @@ def _ls_dirs( next_token = response.get("NextContinuationToken") if not next_token: break - if files: - self.dircache[cache_key] = files + if use_cache: + if files: + self.dircache[cache_key] = files + else: + self.dircache.pop(cache_key, None) return files def ls( @@ -696,13 +726,15 @@ def _find( raise ValueError("Cannot traverse all files in S3.") bucket, key, _ = self.parse_path(path) prefix = kwargs.pop("prefix", "") + # Keep refresh in kwargs so that the recursive calls also refresh. + refresh = kwargs.get("refresh", False) # When maxdepth is specified, use a recursive approach with delimiter if maxdepth is not None: result: list[S3Object] = [] # List files and directories at current level - current_items = self._ls_dirs(path, prefix=prefix, delimiter="/") + current_items = self._ls_dirs(path, prefix=prefix, delimiter="/", refresh=refresh) for item in current_items: if item.type == S3ObjectType.S3_OBJECT_TYPE_FILE: @@ -724,16 +756,17 @@ def _find( return result # For unlimited depth, use the original approach (get all files at once) - files = self._ls_dirs(path, prefix=prefix, delimiter="") + files = self._ls_dirs(path, prefix=prefix, delimiter="", refresh=refresh) if not files and key: try: - files = [self.info(path)] + files = [self.info(path, refresh=refresh)] except FileNotFoundError: files = [] # If withdirs is True, we need to derive directories from file paths if withdirs: - files.extend(self._extract_parent_directories(files, bucket, key)) + # Build a new list; files may be the cached listing. + files = files + self._extract_parent_directories(files, bucket, key) # Filter directories if withdirs is False (default) if withdirs is False or withdirs is None: @@ -760,7 +793,11 @@ def find( maxdepth: Maximum depth to recurse (None for unlimited). withdirs: Whether to include directories in results (None = default behavior). detail: If True, return dict of {path: S3Object}; if False, return list of paths. - **kwargs: Additional arguments. + **kwargs: Additional arguments including: + prefix: Key prefix, relative to the path, to filter the listed keys + by. Without maxdepth, if nothing is listed and the path itself is + an object, that object is returned regardless of the prefix. + refresh: If True, bypass the cache and list from S3. Returns: Dictionary mapping paths to S3Objects (if detail=True) or @@ -784,7 +821,8 @@ def exists(self, path: str, **kwargs) -> bool: Args: path: S3 path to check (e.g., "s3://bucket" or "s3://bucket/key"). - **kwargs: Additional arguments (unused). + **kwargs: Additional arguments including: + refresh: If True, bypass the cache and query S3. Returns: True if the path exists, False otherwise. @@ -794,6 +832,7 @@ def exists(self, path: str, **kwargs) -> bool: >>> fs.exists("s3://my-bucket/file.txt") >>> fs.exists("s3://my-bucket/") """ + refresh = kwargs.pop("refresh", False) path = self._strip_protocol(path) if path in ["", "/"]: # The root always exists. @@ -801,22 +840,22 @@ def exists(self, path: str, **kwargs) -> bool: bucket, key, _ = self.parse_path(path) if key: try: - if self._ls_from_cache(path): + if not refresh and self._ls_from_cache(path): return True - info = self.info(path) + info = self.info(path, refresh=refresh) return bool(info) except FileNotFoundError: return False - elif self.dircache.get(bucket, False): - return True - else: + if not refresh: + if self.dircache.get(bucket, False): + return True try: if self._ls_from_cache(bucket): return True except FileNotFoundError: pass - file = self._head_bucket(bucket) - return bool(file) + file = self._head_bucket(bucket, refresh=refresh) + return bool(file) def rm_file(self, path: str, **kwargs) -> None: """Delete an S3 object with DeleteObject. @@ -1882,6 +1921,9 @@ def invalidate_cache(self, path: str | None = None) -> None: path = self._strip_protocol(path) while path: self.dircache.pop(path, None) + # _ls_dirs caches listings under (path, delimiter). + for delimiter in ("/", ""): + self.dircache.pop((path, delimiter), None) path = self._parent(path) def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None: diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index a2ead0af2..ed6335019 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -148,6 +148,16 @@ def _make_fs(): fs.version_aware = False return fs + @staticmethod + def _file_object(key): + # Build a listed file entry in the bucket named "bucket". + return S3Object( + init={"Key": key}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key=key, + ) + def test_get_client_compatible_with_s3fs(self): # Only constructs a boto3 client; no AWS access. fs = S3FileSystem( @@ -192,6 +202,137 @@ def test_ls_from_cache_with_cached_object(self): # every cache value is a listing and raises TypeError here). assert fs._ls_from_cache("bucket/key/child") is None + def test_invalidate_cache_drops_listings_of_path_and_parents(self): + fs = self._make_fs() + invalidated = [ + "bucket/a/b/c.txt", + ("bucket/a/b", "/"), + ("bucket/a/b", ""), + ("bucket/a", "/"), + ("bucket/a", ""), + ("bucket", "/"), + ("bucket", ""), + ] + kept = ["", ("bucket/a/x", "/")] + for cache_key in invalidated + kept: + fs.dircache[cache_key] = [] + + fs.invalidate_cache("s3://bucket/a/b/c.txt") + assert list(fs.dircache) == kept + + @pytest.mark.parametrize( + ("prefix", "next_token"), + [ + ("test_", None), + ("", "token"), + ], + ) + def test_ls_dirs_partial_listing_bypasses_cache(self, prefix, next_token): + fs = self._make_fs() + cached = self._file_object("dir/cached") + fs.dircache[("bucket/dir", "")] = [cached] + fs._call.return_value = {"Contents": [{"Key": "dir/test_1"}]} + + files = fs._ls_dirs("bucket/dir", prefix=prefix, delimiter="", next_token=next_token) + assert [f.name for f in files] == ["bucket/dir/test_1"] + assert fs.dircache[("bucket/dir", "")] == [cached] + + # A complete listing of the path is still served from the cache. + fs._call.reset_mock() + assert fs._ls_dirs("bucket/dir", delimiter="") == [cached] + fs._call.assert_not_called() + + def test_ls_dirs_empty_refresh_evicts_cached_listing(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "/")] = [self._file_object("dir/deleted")] + fs._call.return_value = {} + + assert fs._ls_dirs("bucket/dir", refresh=True) == [] + # The next listing must not return the deleted object from the cache. + fs._call.reset_mock() + assert fs._ls_dirs("bucket/dir") == [] + fs._call.assert_called_once() + + def test_find_withdirs_does_not_modify_cached_listing(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "")] = [self._file_object("dir/sub/file")] + + expected = ["bucket/dir/sub", "bucket/dir/sub/file"] + assert sorted(fs.find("s3://bucket/dir", withdirs=True)) == expected + assert sorted(fs.find("s3://bucket/dir", withdirs=True)) == expected + assert fs.find("s3://bucket/dir") == ["bucket/dir/sub/file"] + fs._call.assert_not_called() + + def test_find_refresh_bypasses_cached_listings(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "")] = [self._file_object("dir/old")] + fs.dircache[("bucket/dir", "/")] = [self._file_object("dir/old")] + fs.dircache[("bucket/dir/sub", "/")] = [self._file_object("dir/sub/old")] + responses = { + ("dir/", ""): {"Contents": [{"Key": "dir/sub/new"}]}, + ("dir/", "/"): {"CommonPrefixes": [{"Prefix": "dir/sub/"}]}, + ("dir/sub/", "/"): {"Contents": [{"Key": "dir/sub/new"}]}, + } + fs._call.side_effect = lambda method, **kwargs: responses[ + (kwargs["Prefix"], kwargs["Delimiter"]) + ] + + assert fs.find("s3://bucket/dir", refresh=True) == ["bucket/dir/sub/new"] + # The subdirectory listings of maxdepth are refreshed as well. + assert fs.find("s3://bucket/dir", maxdepth=1, refresh=True) == ["bucket/dir/sub/new"] + + def test_refresh_evicts_cached_object_and_bucket_not_found(self): + fs = self._make_fs() + fs.dircache["bucket/key"] = self._file_object("key") + fs.dircache["bucket"] = fs._directory_object("bucket", None) + fs.dircache[""] = [fs._directory_object("bucket", None)] + + def call(method, **kwargs): + if method in (fs._client.head_object, fs._client.head_bucket): + raise FileNotFoundError + return {} + + fs._call.side_effect = call + + assert fs.ls("s3://bucket/key", refresh=True) == [] + # The next lookups must not return the deleted object and bucket from the cache. + assert fs.ls("s3://bucket/key") == [] + assert not fs.exists("s3://bucket/key") + with pytest.raises(FileNotFoundError): + fs.info("s3://bucket", refresh=True) + assert not fs.exists("s3://bucket") + + def test_exists_refresh_bypasses_cache(self): + fs = self._make_fs() + fs.dircache["bucket/key"] = self._file_object("key") + fs.dircache["bucket"] = fs._directory_object("bucket", None) + fs.dircache[""] = [fs._directory_object("bucket", None)] + + def call(method, **kwargs): + if method in (fs._client.head_object, fs._client.head_bucket): + raise FileNotFoundError + return {} + + fs._call.side_effect = call + + assert fs.exists("s3://bucket/key") + assert fs.exists("s3://bucket") + fs._call.assert_not_called() + + assert not fs.exists("s3://bucket/key", refresh=True) + assert not fs.exists("s3://bucket", refresh=True) + + def test_missing_bucket_keeps_bucket_listing_without_it(self): + fs = self._make_fs() + fs.dircache[""] = [fs._directory_object("bucket", None)] + fs._call.side_effect = FileNotFoundError + + assert not fs.exists("s3://missing") + # Other buckets are still answered from the cached bucket listing. + fs._call.reset_mock() + assert fs.exists("s3://bucket") + fs._call.assert_not_called() + def test_mkdir_creates_bucket(self): fs = self._make_fs() fs.allow_bucket_creation = True @@ -719,6 +860,29 @@ def test_ls_dirs(self, fs): assert test_1_detail[0].name == fs._strip_protocol(f"{dir_}/prefix/test_1") assert test_1_detail[0].size == 1 + def test_ls_and_find_reflect_changes_through_the_filesystem(self, fs): + dir_ = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_ls_and_find_reflect_changes/{uuid.uuid4()}" + ) + path = fs._strip_protocol(dir_) + fs.touch(f"{dir_}/a.txt") + fs.touch(f"{dir_}/b.txt") + assert sorted(fs.ls(dir_)) == [f"{path}/a.txt", f"{path}/b.txt"] + assert sorted(fs.find(dir_)) == [f"{path}/a.txt", f"{path}/b.txt"] + + fs.rm(f"{dir_}/a.txt") + assert fs.ls(dir_) == [f"{path}/b.txt"] + assert fs.find(dir_) == [f"{path}/b.txt"] + + fs.touch(f"{dir_}/c.txt") + assert sorted(fs.ls(dir_)) == [f"{path}/b.txt", f"{path}/c.txt"] + assert sorted(fs.find(dir_)) == [f"{path}/b.txt", f"{path}/c.txt"] + # A prefixed find must not be served from the unprefixed listing. + assert fs.find(dir_, prefix="c") == [f"{path}/c.txt"] + + fs.rm(dir_, recursive=True) + def test_info_bucket(self, fs): dir_ = f"s3://{ENV.s3_staging_bucket}" bucket, key, version_id = fs.parse_path(dir_)