diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index ba7abb984..3bd5fcf2a 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -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. """ @@ -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. diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2d0fa6cc0..7a26569b4 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -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"), [