Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 23 additions & 5 deletions pyathena/filesystem/s3.py
Original file line number Diff line number Diff line change
Expand Up @@ -1908,6 +1908,11 @@ def modified(self, path: str) -> datetime:
def invalidate_cache(self, path: str | None = None) -> None:
"""Remove the cached entries of the path and its parent paths.

A version-qualified path invalidates the version under every query
spelling that ``parse_path`` accepts, and also the object path without
the version, because deleting or copying a version can change the
current version of the object.

Args:
path: The path to invalidate. If None, clear the whole cache.
"""
Expand All @@ -1916,11 +1921,24 @@ def invalidate_cache(self, path: str | None = None) -> None:
else:
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)
# parse_path does not accept "?" in keys, so it starts the
# versionId query.
base, _, query = path.partition("?")
cache_paths = [path]
if query:
version_id = query.partition("=")[2]
cache_paths.extend(
f"{base}?{name}={version_id}"
for name in ("versionId", "versionID", "versionid", "version_id")
)
for cache_path in cache_paths:
self.dircache.pop(cache_path, None)
# _ls_dirs caches listings under (path, delimiter).
for delimiter in ("/", ""):
self.dircache.pop((cache_path, delimiter), None)
# A version-qualified path continues with the path without
# the version.
path = self._strip_protocol(base) if query else self._parent(path)

def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None:
"""Check the dircache for a cached entry of the path.
Expand Down
43 changes: 43 additions & 0 deletions tests/pyathena/filesystem/test_s3.py
Original file line number Diff line number Diff line change
Expand Up @@ -220,6 +220,49 @@ def test_invalidate_cache_drops_listings_of_path_and_parents(self):
fs.invalidate_cache("s3://bucket/a/b/c.txt")
assert list(fs.dircache) == kept

@pytest.mark.parametrize(
("path", "cache_key"),
[
("s3://bucket/a/c.txt?versionId=v1", "bucket/a/c.txt?versionId=v1"),
(Path("bucket/a/c.txt?versionId=v1"), "bucket/a/c.txt?versionId=v1"),
# A directory marker object keeps the trailing slash before the query.
("s3://bucket/a/c.txt/?versionId=v1", "bucket/a/c.txt/?versionId=v1"),
# parse_path accepts other spellings of the query.
("s3://bucket/a/c.txt?version_id=v1", "bucket/a/c.txt?versionId=v1"),
("s3://bucket/a/c.txt?versionId=v1", "bucket/a/c.txt?versionid=v1"),
],
)
def test_invalidate_cache_version_drops_object_path(self, path, cache_key):
fs = self._make_fs()
invalidated = [
cache_key,
(cache_key, "/"),
"bucket/a/c.txt",
("bucket/a", "/"),
("bucket", "/"),
]
# Other versions of the object do not change.
kept = ["bucket/a/c.txt?versionId=v2"]
for key in invalidated + kept:
fs.dircache[key] = []

fs.invalidate_cache(path)
assert list(fs.dircache) == kept

def test_rm_file_version_invalidates_object_path(self):
fs = self._make_fs()
fs.dircache["bucket/a/c.txt"] = self._file_object("a/c.txt")

fs.rm_file("s3://bucket/a/c.txt?versionId=v1")
fs._call.assert_called_once_with(
fs._client.delete_object, Bucket="bucket", Key="a/c.txt", VersionId="v1"
)

# The deleted version was the only one: HeadObject and the prefix
# listing find nothing, instead of the cached object answering.
fs._call.side_effect = [FileNotFoundError("bucket/a/c.txt"), {}]
assert not fs.exists("s3://bucket/a/c.txt")

@pytest.mark.parametrize(
("prefix", "next_token"),
[
Expand Down
Loading