diff --git a/docs/filesystem.md b/docs/filesystem.md index fd2439ea..3c0373a6 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -156,7 +156,7 @@ uploads = fs.list_multipart_uploads("s3://YOUR_S3_BUCKET") for upload in uploads: print(upload.key, upload.upload_id, upload.initiated) -# Abort all incomplete uploads under a bucket or key prefix. +# Abort all incomplete uploads to a key and the keys under it. fs.clear_multipart_uploads("s3://YOUR_S3_BUCKET/path/to/") ``` diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 2819a99d..f8674fd8 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -16,6 +16,7 @@ from multiprocessing import cpu_count from re import Pattern from typing import Any, cast +from urllib.parse import unquote_plus import botocore.exceptions from boto3 import Session @@ -2100,15 +2101,20 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: to abort all of them. Args: - path: S3 bucket or prefix path (e.g., "bucket", "s3://bucket" or - "s3://bucket/prefix"). If the path contains a key prefix, - only the uploads under that prefix are listed. + path: S3 bucket or key path (e.g., "bucket", "s3://bucket" or + "s3://bucket/prefix"). If the path contains a key, only the + uploads to that key and to the keys under ``key/`` are + listed, not those to sibling keys that merely start with the + same characters (e.g., ``prefix2/a``). Returns: List of S3MultipartUpload instances describing the in-progress multipart uploads. """ bucket, key, _ = self.parse_path(path) + # S3 matches Prefix as a plain string, so the uploads are filtered to + # the key itself and the keys under it. + prefix = f"{key.rstrip('/')}/" if key else "" _logger.debug(f"List multipart uploads: s3://{bucket}/{key}") uploads: list[S3MultipartUpload] = [] @@ -2127,7 +2133,9 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: **request, ) uploads.extend( - S3MultipartUpload({**u, "Bucket": bucket}) for u in response.get("Uploads", []) + S3MultipartUpload({**u, "Bucket": bucket}) + for u in response.get("Uploads", []) + if u["Key"] == key or u["Key"].startswith(prefix) ) if not response.get("IsTruncated"): break @@ -2140,7 +2148,15 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: def object_version_info( self, path: str, delete_markers: bool = False, **kwargs ) -> list[S3ObjectVersion]: - """List the versions of the objects under the path. + """List the versions of the object or of the objects under the path. + + A key path without a trailing slash selects that key if it has any + versions or delete markers, and otherwise the keys under ``key/``. + The choice does not depend on ``delete_markers``, so a key that has + only delete markers yields no versions without them. A key path with + a trailing slash selects the keys under it, and a bucket path selects + all the keys in the bucket. Sibling keys that merely start with the + same characters (e.g., ``key.bak``) are never included. Args: path: S3 path (s3://bucket/key or a key prefix) to list the @@ -2153,6 +2169,9 @@ def object_version_info( List of S3ObjectVersion instances describing the versions. """ bucket, key, _ = self.parse_path(path) + # S3 matches Prefix as a plain string, so the versions are filtered to + # the key itself or the keys under it. + prefix = f"{key.rstrip('/')}/" if key else "" _logger.debug(f"List object versions: s3://{bucket}/{key}") versions: list[S3ObjectVersion] = [] @@ -2161,20 +2180,30 @@ def object_version_info( S3ObjectVersion(bucket=bucket, is_delete_marker=False, response=v) for v in response.get("Versions", []) ) - if delete_markers: - versions.extend( - S3ObjectVersion(bucket=bucket, is_delete_marker=True, response=m) - for m in response.get("DeleteMarkers", []) - ) - return versions + # Delete markers are kept until the key is chosen, so that the + # choice is the same with and without them. + versions.extend( + S3ObjectVersion(bucket=bucket, is_delete_marker=True, response=m) + for m in response.get("DeleteMarkers", []) + ) + # botocore decodes the keys only when it sets EncodingType itself, so + # the keys of an explicit EncodingType="url" are decoded for matching. + url_encoded = kwargs.get("EncodingType") == "url" + keys = [unquote_plus(v.key) if url_encoded else v.key for v in versions] + if key and not key.endswith("/") and key in keys: + selected = [v for v, k in zip(versions, keys, strict=True) if k == key] + else: + selected = [v for v, k in zip(versions, keys, strict=True) if k.startswith(prefix)] + return [v for v in selected if delete_markers or not v.is_delete_marker] def clear_multipart_uploads(self, path: str) -> None: """Abort any incomplete multipart uploads in the bucket. Args: - path: S3 bucket or prefix path (e.g., "bucket", "s3://bucket" or - "s3://bucket/prefix"). If the path contains a key prefix, - only the uploads under that prefix are aborted. + path: S3 bucket or key path (e.g., "bucket", "s3://bucket" or + "s3://bucket/prefix"). If the path contains a key, only the + uploads to that key and to the keys under ``key/`` are + aborted, as listed by :meth:`list_multipart_uploads`. """ uploads = self.list_multipart_uploads(path) if not uploads: diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 44115467..e8586d4f 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -570,7 +570,7 @@ def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: def object_version_info( self, path: str, delete_markers: bool = False, **kwargs ) -> list[S3ObjectVersion]: - """List the versions of the objects under the path. + """List the versions of the object or of the objects under the path. See :meth:`S3FileSystem.object_version_info`. @@ -590,7 +590,7 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: See :meth:`S3FileSystem.list_multipart_uploads`. Args: - path: S3 bucket or prefix path (e.g., "s3://bucket" or "s3://bucket/prefix"). + path: S3 bucket or key path (e.g., "s3://bucket" or "s3://bucket/prefix"). Returns: List of S3MultipartUpload instances describing the uploads. @@ -603,7 +603,7 @@ def clear_multipart_uploads(self, path: str) -> None: See :meth:`S3FileSystem.clear_multipart_uploads`. Args: - path: S3 bucket or prefix path (e.g., "s3://bucket" or "s3://bucket/prefix"). + path: S3 bucket or key path (e.g., "s3://bucket" or "s3://bucket/prefix"). """ self._sync_fs.clear_multipart_uploads(path) diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 76253929..e4d1f606 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1497,6 +1497,80 @@ def test_object_version_info_with_delete_markers(self): ("m1", True), ] + @pytest.mark.parametrize( + ("path", "expected"), + [ + # The key itself wins over the keys under it. + ("s3://bucket/a.csv", ["a.csv"]), + # Without an object, the keys under the path are returned. + ("s3://bucket/dir", ["dir/", "dir/x"]), + ("s3://bucket/dir/", ["dir/", "dir/x"]), + # A trailing slash selects the keys under the path even if the + # key without it exists. + ("s3://bucket/a.csv/", ["a.csv/x"]), + ("s3://bucket", ["a.csv", "a.csv.bak", "a.csv/x", "dir/", "dir/x", "dir2/y"]), + ], + ) + def test_object_version_info_excludes_sibling_keys(self, path, expected): + fs = self._make_fs() + keys = ["a.csv", "a.csv.bak", "a.csv/x", "dir/", "dir/x", "dir2/y"] + # S3 matches Prefix as a plain string prefix. + fs._call.side_effect = lambda _, **request: { + "Versions": [ + {"Key": k, "VersionId": f"v-{k}", "IsLatest": True} + for k in keys + if k.startswith(request["Prefix"]) + ], + "DeleteMarkers": [ + {"Key": k, "VersionId": f"m-{k}", "IsLatest": False} + for k in keys + if k.startswith(request["Prefix"]) + ], + "IsTruncated": False, + } + + actual = fs.object_version_info(path) + assert [v.key for v in actual] == expected + actual = fs.object_version_info(path, delete_markers=True) + assert sorted(v.key for v in actual if not v.is_delete_marker) == expected + assert sorted(v.key for v in actual if v.is_delete_marker) == expected + + @pytest.mark.parametrize( + ("delete_markers", "expected"), + [ + (False, []), + (True, [("dir", "m1", True)]), + ], + ) + def test_object_version_info_chooses_key_with_only_delete_markers( + self, delete_markers, expected + ): + fs = self._make_fs() + # "dir" is both a deleted object, with only a delete marker left, and + # a folder. + fs._call.return_value = { + "Versions": [{"Key": "dir/x", "VersionId": "v1", "IsLatest": True}], + "DeleteMarkers": [{"Key": "dir", "VersionId": "m1", "IsLatest": True}], + "IsTruncated": False, + } + + actual = fs.object_version_info("s3://bucket/dir", delete_markers=delete_markers) + assert [(v.key, v.version_id, v.is_delete_marker) for v in actual] == expected + + def test_object_version_info_matches_url_encoded_keys(self): + fs = self._make_fs() + # With an explicit EncodingType="url", botocore leaves the keys encoded. + fs._call.return_value = { + "Versions": [ + {"Key": "a+b", "VersionId": "v1", "IsLatest": True}, + {"Key": "a+b.bak", "VersionId": "v2", "IsLatest": True}, + ], + "IsTruncated": False, + } + + actual = fs.object_version_info("s3://bucket/a b", EncodingType="url") + assert [(v.key, v.version_id) for v in actual] == [("a+b", "v1")] + def test_ls_versions_requires_version_aware(self): fs = self._make_fs() with pytest.raises(ValueError, match="version aware"): @@ -1646,6 +1720,39 @@ def test_list_multipart_uploads_paginates(self): UploadIdMarker="upload1", ) + @pytest.mark.parametrize( + ("path", "expected"), + [ + ("s3://bucket/data", ["data", "data/part.csv"]), + ("s3://bucket/data/", ["data/part.csv"]), + ("s3://bucket", ["data", "data.csv", "data/part.csv", "data2/other.csv"]), + ], + ) + def test_list_and_clear_multipart_uploads_exclude_sibling_keys(self, path, expected): + fs = self._make_fs() + keys = ["data", "data.csv", "data/part.csv", "data2/other.csv"] + aborted = [] + + def call(method, **request): + if method is fs._client.abort_multipart_upload: + aborted.append(request["Key"]) + return {} + # S3 matches Prefix as a plain string prefix. + return { + "Uploads": [ + {"Key": k, "UploadId": f"u-{k}"} + for k in keys + if k.startswith(request.get("Prefix", "")) + ], + "IsTruncated": False, + } + + fs._call.side_effect = call + + assert [u.key for u in fs.list_multipart_uploads(path)] == expected + fs.clear_multipart_uploads(path) + assert sorted(aborted) == expected + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): @@ -2617,17 +2724,24 @@ def test_list_and_clear_multipart_uploads(self, fs): prefix_path = f"s3://{bucket}/{prefix}" key = f"{prefix}/file" upload = fs._create_multipart_upload(bucket=bucket, key=key) - - uploads = fs.list_multipart_uploads(prefix_path) - listed = next((u for u in uploads if u.upload_id == upload.upload_id), None) - assert listed - assert listed.bucket == bucket - assert listed.key == key - assert listed.initiated - - fs.clear_multipart_uploads(prefix_path) - uploads = fs.list_multipart_uploads(prefix_path) - assert not any(u.upload_id == upload.upload_id for u in uploads) + # A sibling key that starts with the same characters as the prefix. + sibling = fs._create_multipart_upload(bucket=bucket, key=f"{prefix}2/file") + try: + uploads = fs.list_multipart_uploads(prefix_path) + listed = next((u for u in uploads if u.upload_id == upload.upload_id), None) + assert listed + assert listed.bucket == bucket + assert listed.key == key + assert listed.initiated + assert not any(u.upload_id == sibling.upload_id for u in uploads) + + fs.clear_multipart_uploads(prefix_path) + uploads = fs.list_multipart_uploads(prefix_path) + assert not any(u.upload_id == upload.upload_id for u in uploads) + uploads = fs.list_multipart_uploads(f"{prefix_path}2") + assert any(u.upload_id == sibling.upload_id for u in uploads) + finally: + fs.clear_multipart_uploads(f"{prefix_path}2") def test_object_version_info(self, fs): path = ( @@ -2635,6 +2749,8 @@ def test_object_version_info(self, fs): f"filesystem/test_object_version_info/{uuid.uuid4()}" ) fs.pipe(path, b"data") + # A sibling key that starts with the same characters as the path. + fs.pipe(f"{path}.bak", b"backup") versions = fs.object_version_info(path) assert len(versions) == 1