diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index ba7abb984..57ff4b943 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -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 = { @@ -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 @@ -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"). @@ -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): @@ -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") ): diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2d0fa6cc0..b1b13fc2e 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -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 = [ @@ -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}/"