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
22 changes: 18 additions & 4 deletions pyathena/filesystem/s3.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,6 +340,14 @@ def _head_object(
) -> S3Object | None:
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:
# 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}"
# 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"
if path not in self.dircache or refresh:
try:
request = {
Expand All @@ -366,7 +374,8 @@ def _head_object(
key=key,
version_id=version_id,
)
self.dircache[path] = file
if cacheable:
self.dircache[path] = file
else:
file = self.dircache[path]
return file
Expand Down Expand Up @@ -585,7 +594,11 @@ def info(self, path: str, **kwargs) -> S3Object:
exists, with a ListObjectsV2 request (``Delimiter="/"``,
``MaxKeys=1``) that checks whether it is a key prefix; a bucket path
is looked up with HeadBucket. With ``version_aware``, a cached file
entry without a version ID is looked up again.
entry without a version ID is looked up again. With an explicit
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.

Args:
path: S3 path (e.g., "s3://bucket" or "s3://bucket/key").
Expand Down Expand Up @@ -617,7 +630,9 @@ def info(self, path: str, **kwargs) -> S3Object:
key=None,
version_id=None,
)
if not refresh:
# 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:
caches: list[S3Object] | S3Object | None = self._ls_from_cache(path)
if caches is not None:
if isinstance(caches, list):
Expand All @@ -630,7 +645,6 @@ def info(self, path: str, **kwargs) -> S3Object:
if cache:
if (
self.version_aware
and not version_id
and cache.get("type") == S3ObjectType.S3_OBJECT_TYPE_FILE
and not cache.get("version_id")
):
Expand Down
42 changes: 42 additions & 0 deletions tests/pyathena/filesystem/test_s3.py
Original file line number Diff line number Diff line change
Expand Up @@ -504,6 +504,35 @@ def test_head_object_version_aware(self):
fs._call.return_value = {"ContentLength": 4, "ETag": '"etag"', "VersionId": "v1"}
assert fs._head_object("bucket/key").version_id == "v1"

@pytest.mark.parametrize("version_aware", [False, True])
def test_info_caches_each_version_separately(self, version_aware):
fs = self._make_fs()
fs.version_aware = version_aware
responses = {
None: {"ContentLength": 3, "ETag": '"e3"', "VersionId": "v3"},
"v1": {"ContentLength": 1, "ETag": '"e1"', "VersionId": "v1"},
"v2": {"ContentLength": 2, "ETag": '"e2"', "VersionId": "v2"},
}
fs._call.side_effect = lambda _, **kwargs: responses[kwargs.get("VersionId")]

for _ in range(2):
assert fs.info("s3://bucket/key", version_id="v1").size == 1
assert fs.info("s3://bucket/key", version_id="v2").size == 2
assert fs.info("s3://bucket/key?versionId=v1").size == 1
assert fs.info("s3://bucket/key").size == 3
# The second round is served from the cache.
assert fs._call.call_count == 3

def test_info_does_not_cache_null_version(self):
fs = self._make_fs()
fs._call.return_value = {"ContentLength": 4, "ETag": '"etag"', "VersionId": "null"}

for _ in range(2):
assert fs.info("s3://bucket/key", version_id="null").size == 4
assert fs.info("s3://bucket/key?versionId=null").size == 4
# An overwrite can replace the null version, so it is looked up every time.
assert fs._call.call_count == 4

def test_object_version_info_paginates(self):
fs = self._make_fs()
fs._call.side_effect = [
Expand Down Expand Up @@ -1684,6 +1713,19 @@ def test_version_aware_read(self, fs):
with fs.open(path, "rb") as f:
assert f.read() == b"0123456789"

def test_read_null_version(self, fs):
path = (
f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/"
f"filesystem/test_read_null_version/{uuid.uuid4()}"
)
# An unversioned bucket stores each object as the "null" version,
# which an overwrite replaces.
for data in (b"1", b"22"):
fs.pipe(path, data)
for _ in range(2):
with fs.open(f"{path}?versionId=null", "rb") as f:
assert f.read() == data

def test_file_url_metadata_getxattr_setxattr(self, fs):
path = (
f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/"
Expand Down
Loading