diff --git a/docs/api/filesystem.rst b/docs/api/filesystem.rst index 0a2627b35..958c8d05d 100644 --- a/docs/api/filesystem.rst +++ b/docs/api/filesystem.rst @@ -43,6 +43,12 @@ S3 Executor .. autoclass:: pyathena.filesystem.s3_executor.S3AioExecutor :members: +S3 Paths +-------- + +.. autoclass:: pyathena.filesystem.s3_path.S3Path + :members: + S3 Objects ---------- diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 0ecc0d0dd..042488dbf 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -8,7 +8,6 @@ import math import mimetypes import os.path -import re import time from collections import Counter from collections.abc import Callable, Iterator, Mapping @@ -49,6 +48,7 @@ S3PutObject, S3StorageClass, ) +from pyathena.filesystem.s3_path import S3Path from pyathena.util import RetryConfig, override, retry_api_call _logger = logging.getLogger(__name__) @@ -212,10 +212,7 @@ class S3FileSystem(AbstractFileSystem): "Expires", "Metadata", ) - PATTERN_PATH: Pattern[str] = re.compile( - r"(^s3://|^s3a://|^)(?P[a-zA-Z0-9.\-_]+)(/(?P[^?]+)|/)?" - r"($|\?version(Id|ID|id|_id)=(?P.+)$)" - ) + PATTERN_PATH: Pattern[str] = S3Path.PATTERN protocol = ("s3", "s3a") _extra_tokenize_attributes = ("default_block_size",) @@ -350,9 +347,7 @@ def _get_client_compatible_with_s3fs(self, **kwargs) -> BaseClient: def parse_path(path: str) -> tuple[str, str | None, str | None]: """Parse an S3 path into its bucket, key and version ID. - The path may have an ``s3://`` or ``s3a://`` scheme and a version ID - query (``?versionId=``, ``?versionID=``, ``?versionid=`` or - ``?version_id=``). + See :meth:`S3Path.parse`, which returns them as an :class:`S3Path`. Args: path: The S3 path (e.g., "s3://bucket/key?versionId=..."). @@ -364,10 +359,8 @@ def parse_path(path: str) -> tuple[str, str | None, str | None]: Raises: ValueError: If the path is not a valid S3 path. """ - match = S3FileSystem.PATTERN_PATH.search(path) - if match: - return match.group("bucket"), match.group("key"), match.group("version_id") - raise ValueError(f"Invalid S3 path format {path}.") + s3_path = S3Path.parse(path) + return s3_path.bucket, s3_path.key, s3_path.version_id @staticmethod def _directory_object(bucket: str, key: str | None, version_id: str | None = None) -> S3Object: @@ -481,14 +474,14 @@ def _head_object( Returns: The object, or None if it does not exist. """ - bucket, key, path_version_id = self.parse_path(path) - version_id = path_version_id if path_version_id else version_id + s3_path = S3Path.parse(path) + version_id = s3_path.version_id if s3_path.version_id else version_id if version_id: # Cache an explicit version under its version-qualified path so # that it neither reuses nor replaces the entry of another version, # with a single spelling of the query so that a missing version # evicts the entry whatever spelling looked it up. - path = f"{path.partition('?')[0]}?versionId={version_id}" + path = str(s3_path.with_version_id(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. @@ -498,8 +491,8 @@ def _head_object( if file is None: try: request = { - "Bucket": bucket, - "Key": key, + "Bucket": s3_path.bucket, + "Key": s3_path.key, } if version_id: request.update({"VersionId": version_id}) @@ -525,8 +518,8 @@ def _head_object( file = S3Object( init=response, type=S3ObjectType.S3_OBJECT_TYPE_FILE, - bucket=bucket, - key=key, + bucket=s3_path.bucket, + key=s3_path.key, version_id=version_id, ) if cacheable: @@ -598,10 +591,10 @@ def _ls_dirs( Returns: The listed directories and files. """ - bucket, key, version_id = self.parse_path(path) + s3_path = S3Path.parse(path) use_cache = not prefix and not next_token - if key: - prefix = f"{key}/{prefix if prefix else ''}" + if s3_path.key: + prefix = f"{s3_path.key}/{prefix if prefix else ''}" cache_key = (path, delimiter) cached = self.dircache.get(cache_key) if use_cache and not refresh else None @@ -611,7 +604,7 @@ def _ls_dirs( files: list[S3Object] = [] while True: request: dict[Any, Any] = { - "Bucket": bucket, + "Bucket": s3_path.bucket, "Prefix": prefix, "Delimiter": delimiter, } @@ -624,14 +617,16 @@ def _ls_dirs( **request, ) files.extend( - self._directory_object(bucket, c["Prefix"][:-1].rstrip("/"), version_id) + self._directory_object( + s3_path.bucket, c["Prefix"][:-1].rstrip("/"), s3_path.version_id + ) for c in response.get("CommonPrefixes", []) ) files.extend( S3Object( init=c, type=S3ObjectType.S3_OBJECT_TYPE_FILE, - bucket=bucket, + bucket=s3_path.bucket, key=c["Key"], ) for c in response.get("Contents", []) @@ -694,26 +689,30 @@ def _ls_object_versions(self, path: str) -> list[S3Object]: The listing is always fetched from S3 and is not cached, because the dircache stores the current view of a path. """ - bucket, key, _ = self.parse_path(path) - prefix = f"{key}/" if key else "" + s3_path = S3Path.parse(path) + prefix = f"{s3_path.key}/" if s3_path.key else "" files: list[S3Object] = [] - for response in self._list_object_versions_pages(bucket, prefix=prefix, delimiter="/"): + for response in self._list_object_versions_pages( + s3_path.bucket, prefix=prefix, delimiter="/" + ): files.extend( - self._directory_object(bucket, c["Prefix"][:-1].rstrip("/")) + self._directory_object(s3_path.bucket, c["Prefix"][:-1].rstrip("/")) for c in response.get("CommonPrefixes", []) ) files.extend( - self._versioned_file_object(bucket, v) for v in response.get("Versions", []) + self._versioned_file_object(s3_path.bucket, v) for v in response.get("Versions", []) ) - if not files and key: + if not files and s3_path.key: # The path may point at an object rather than a key prefix. files = [ - self._versioned_file_object(bucket, v) - for response in self._list_object_versions_pages(bucket, prefix=key, delimiter="/") + self._versioned_file_object(s3_path.bucket, v) + for response in self._list_object_versions_pages( + s3_path.bucket, prefix=s3_path.key, delimiter="/" + ) for v in response.get("Versions", []) - if v["Key"] == key + if v["Key"] == s3_path.key ] return files @@ -808,8 +807,8 @@ def info(self, path: str, **kwargs) -> S3Object: key=None, version_id=None, ) - bucket, key, path_version_id = self.parse_path(path) - version_id = path_version_id if path_version_id else kwargs.pop("version_id", None) + s3_path = S3Path.parse(path) + version_id = s3_path.version_id if s3_path.version_id else kwargs.pop("version_id", None) lookup_kwargs = self._get_lookup_kwargs(kwargs) # Cached entries describe the current version of a path as looked up # without lookup parameters, so an explicit version or lookup @@ -844,9 +843,9 @@ def info(self, path: str, **kwargs) -> S3Object: return cache else: return self._directory_object( - bucket, key.rstrip("/") if key else None, version_id + s3_path.bucket, s3_path.key.rstrip("/") if s3_path.key else None, version_id ) - if key: + if s3_path.key: object_info = self._head_object( path, refresh=refresh, version_id=version_id, lookup_kwargs=lookup_kwargs ) @@ -862,8 +861,8 @@ def info(self, path: str, **kwargs) -> S3Object: self._client.list_objects_v2, **{ **self._get_operation_kwargs("list_objects_v2", lookup_kwargs), - "Bucket": bucket, - "Prefix": f"{key.rstrip('/')}/" if key else "", + "Bucket": s3_path.bucket, + "Prefix": f"{s3_path.key.rstrip('/')}/" if s3_path.key else "", "Delimiter": "/", "MaxKeys": 1, }, @@ -876,7 +875,9 @@ def info(self, path: str, **kwargs) -> S3Object: # Nothing caches the key prefix, and the cached listing of the # parent may predate it. self._evict_cache((self._parent(path), "/")) - return self._directory_object(bucket, key.rstrip("/") if key else None, version_id) + return self._directory_object( + s3_path.bucket, s3_path.key.rstrip("/") if s3_path.key else None, version_id + ) raise FileNotFoundError(path) def _extract_parent_directories( @@ -948,7 +949,7 @@ def _find( path = self._strip_protocol(path) if path in ["", "/"]: raise ValueError("Cannot traverse all files in S3.") - bucket, key, _ = self.parse_path(path) + s3_path = S3Path.parse(path) prefix = kwargs.pop("prefix", "") refresh = kwargs.pop("refresh", False) @@ -964,17 +965,17 @@ def _find( # directories are derived from the listed keys, below the last # slash of the prefix, as with maxdepth. if withdirs: - base_key = "/".join(k for k in (key, prefix.rpartition("/")[0]) if k) + base_key = "/".join(k for k in (s3_path.key, prefix.rpartition("/")[0]) if k) # Build a new list; files may be the cached listing. - files = files + self._extract_parent_directories(files, bucket, base_key) + files = files + self._extract_parent_directories(files, s3_path.bucket, base_key) if files: # Something is listed below the path, so the path is a directory, # which fsspec includes with the directories. A bucket is not # included, since cp_file() cannot copy it in a recursive copy. - if withdirs and key: - files = [self._directory_object(bucket, key), *files] - elif key: + if withdirs and s3_path.key: + files = [self._directory_object(s3_path.bucket, s3_path.key), *files] + elif s3_path.key: # As in fsspec, the path itself is returned if it is an object, # or with withdirs if it is a directory. try: @@ -1088,11 +1089,11 @@ def exists(self, path: str, **kwargs) -> bool: if path in ["", "/"]: # The root always exists. return True - bucket, key, _ = self.parse_path(path) + s3_path = S3Path.parse(path) lookup_kwargs = self._get_lookup_kwargs(kwargs) # The cached listings are made without lookup parameters. use_listings = not refresh and not lookup_kwargs - if key: + if s3_path.key: try: if use_listings and self._ls_from_cache(path): return True @@ -1100,10 +1101,10 @@ def exists(self, path: str, **kwargs) -> bool: return bool(info) except FileNotFoundError: return False - if use_listings and self._ls_from_cache(bucket): + if use_listings and self._ls_from_cache(s3_path.bucket): return True try: - file = self._head_bucket(bucket, refresh=refresh, lookup_kwargs=lookup_kwargs) + file = self._head_bucket(s3_path.bucket, refresh=refresh, lookup_kwargs=lookup_kwargs) except PermissionError: # HeadBucket answers 403 for a bucket that exists but that the # caller may not access. @@ -1121,10 +1122,12 @@ def rm_file(self, path: str, **kwargs) -> None: **kwargs: Accepted for fsspec compatibility; not used in the request. """ - bucket, key, version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: return - self._delete_object(bucket=bucket, key=key, version_id=version_id, **kwargs) + self._delete_object( + bucket=s3_path.bucket, key=s3_path.key, version_id=s3_path.version_id, **kwargs + ) self.invalidate_cache(path) def rm(self, path, recursive=False, maxdepth=None, **kwargs) -> None: @@ -1169,11 +1172,11 @@ def _expand_delete_paths( paths = [path] if isinstance(path, str) else list(path) versioned_paths, unversioned_paths = [], [] for p in paths: - _, key, version_id = self.parse_path(p) + s3_path = S3Path.parse(p) # expand_path strips the slashes of "bucket//" to the bucket. - if not key or not key.strip("/"): + if s3_path.is_bucket: raise ValueError("Cannot delete the bucket.") - if version_id: + if s3_path.version_id: versioned_paths.append(p) else: unversioned_paths.append(p) @@ -1195,7 +1198,7 @@ def _delete_object( if version_id: request.update({"VersionId": version_id}) - _logger.debug(f"Delete object: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Delete object: {S3Path(bucket, key, version_id).uri}") self._call( self._client.delete_object, **request, @@ -1258,12 +1261,12 @@ def _delete_objects_requests(self, paths: list[str], **kwargs) -> list[dict[str, quiet = kwargs.pop("Quiet", True) delete_objects: dict[str, list[dict[str, str]]] = {} for p in paths: - bucket, key, version_id = self.parse_path(p) - if key: - object_ = {"Key": key} - if version_id: - object_.update({"VersionId": version_id}) - delete_objects.setdefault(bucket, []).append(object_) + s3_path = S3Path.parse(p) + if s3_path.key: + object_ = {"Key": s3_path.key} + if s3_path.version_id: + object_.update({"VersionId": s3_path.version_id}) + delete_objects.setdefault(s3_path.bucket, []).append(object_) return [ { "Bucket": bucket, @@ -1306,10 +1309,7 @@ def _delete_objects_path(bucket: str, object_: dict[str, Any]) -> str: Returns: The path, with a ``?versionId=`` query if the entry has a version. """ - path = f"{bucket}/{object_['Key']}" - if object_.get("VersionId"): - path += f"?versionId={object_['VersionId']}" - return path + return str(S3Path(bucket, object_["Key"], object_.get("VersionId"))) @staticmethod def _raise_delete_objects_errors( @@ -1380,13 +1380,13 @@ def mkdir(self, path: str, create_parents: bool = True, **kwargs) -> None: path = self._strip_protocol(path).rstrip("/") if not path: raise ValueError("Cannot create the root directory.") - bucket, key, _ = self.parse_path(path) - if self.exists(bucket): - if not key: + s3_path = S3Path.parse(path) + if self.exists(s3_path.bucket): + if not s3_path.key: # Requested to create a bucket, but the bucket already exists. - raise FileExistsError(bucket) + raise FileExistsError(s3_path.bucket) # Do nothing as the bucket already exists. - elif not key or create_parents: + elif not s3_path.key or create_parents: if not self.allow_bucket_creation: raise PermissionError( "Bucket creation is disabled. " @@ -1395,7 +1395,7 @@ def mkdir(self, path: str, create_parents: bool = True, **kwargs) -> None: acl = kwargs.pop("acl", "") if acl and acl not in self.BUCKET_ACLS: raise ValueError(f"ACL not in {self.BUCKET_ACLS}.") - request: dict[str, Any] = {"Bucket": bucket} + request: dict[str, Any] = {"Bucket": s3_path.bucket} if acl: request.update({"ACL": acl}) region_name = kwargs.pop("region_name", None) or self._client.meta.region_name @@ -1403,22 +1403,22 @@ def mkdir(self, path: str, create_parents: bool = True, **kwargs) -> None: # us-east-1 does not accept a location constraint. request.update({"CreateBucketConfiguration": {"LocationConstraint": region_name}}) - _logger.debug(f"Create bucket: s3://{bucket}") + _logger.debug(f"Create bucket: s3://{s3_path.bucket}") try: self._call( self._client.create_bucket, **request, ) except botocore.exceptions.ParamValidationError as e: - raise ValueError(f"Bucket create failed {bucket!r}: {e}") from e + raise ValueError(f"Bucket create failed {s3_path.bucket!r}: {e}") from e # invalidate_cache of the bucket keeps the cached bucket # listing, so evict it directly. self._evict_cache("") - self.invalidate_cache(bucket) + self.invalidate_cache(s3_path.bucket) else: # exists() has already confirmed the bucket does not exist, # and it is not requested to be created. - raise FileNotFoundError(bucket) + raise FileNotFoundError(s3_path.bucket) def makedirs(self, path: str, exist_ok: bool = False) -> None: """Recursively create a directory, creating the bucket if necessary. @@ -1466,8 +1466,8 @@ def rmdir(self, path: str) -> None: OSError: If the bucket is not empty. """ path = self._strip_protocol(path).rstrip("/") - bucket, key, _ = self.parse_path(path) - if key: + s3_path = S3Path.parse(path) + if s3_path.key: if self.exists(path): # The user may have meant rm(path, recursive=True). raise FileExistsError(path) @@ -1478,12 +1478,12 @@ def rmdir(self, path: str) -> None: "Set allow_bucket_deletion=True on the filesystem to enable it." ) - _logger.debug(f"Delete bucket: s3://{bucket}") + _logger.debug(f"Delete bucket: s3://{s3_path.bucket}") self._call( self._client.delete_bucket, - Bucket=bucket, + Bucket=s3_path.bucket, ) - self.invalidate_cache(bucket) + self.invalidate_cache(s3_path.bucket) # invalidate_cache of the bucket keeps the cached bucket listing, # so evict it directly. self._evict_cache("") @@ -1505,15 +1505,15 @@ def touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]: ValueError: If the path has a version ID, is a bucket, or exists while ``truncate`` is False. """ - bucket, key, version_id = self.parse_path(path) - if version_id: + s3_path = S3Path.parse(path) + if s3_path.version_id: raise ValueError("Cannot touch the file with the version specified.") if not truncate and self.exists(path): raise ValueError("Cannot touch the existing file without specifying truncate.") - if not key: + if not s3_path.key: raise ValueError("Cannot touch the bucket.") - object_ = self._put_object(bucket=bucket, key=key, body=None, **kwargs) + object_ = self._put_object(bucket=s3_path.bucket, key=s3_path.key, body=None, **kwargs) self.invalidate_cache(path) return object_.to_dict() @@ -1622,7 +1622,7 @@ def _move_paths( if not ( (counts[dest] > 1 or dest in sources) and source in directories - and not self.parse_path(p1)[2] + and not S3Path.parse(p1).version_id and self._head_object(source) is None ) ] @@ -1647,11 +1647,10 @@ def _move_target(self, path: str) -> str: The path in ``bucket/key`` form, with the version ID unless it is ``null``. """ - bucket, key, version_id = self.parse_path(path) - target = f"{bucket}/{key}" if key else bucket - if version_id and version_id != "null": - return f"{target}?versionId={version_id}" - return target + s3_path = S3Path.parse(path) + if s3_path.version_id == "null": + return s3_path.name + return str(s3_path) def cp_file( self, path1: str, path2: str, recursive=False, maxdepth=None, on_error=None, **kwargs @@ -1712,11 +1711,11 @@ def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: # Parameters of the multipart copy, not of the S3 requests. block_size = kwargs.pop("block_size", None) max_workers = kwargs.pop("max_workers", None) - bucket1, key1, version_id1 = self.parse_path(path1) - bucket2, key2, version_id2 = self.parse_path(path2) - if version_id2: + source = S3Path.parse(path1) + destination = S3Path.parse(path2) + if destination.version_id: raise ValueError("Cannot copy to a versioned file.") - if not key1 or not key2: + if not source.key or not destination.key: raise ValueError("Cannot copy buckets.") info1 = self.info(path1) @@ -1728,21 +1727,21 @@ def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: try: if size1 <= self.MULTIPART_UPLOAD_MAX_PART_SIZE: self._copy_object( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, + bucket1=source.bucket, + key1=source.key, + version_id1=source.version_id, + bucket2=destination.bucket, + key2=destination.key, **kwargs, ) else: self._copy_object_with_multipart_upload( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, + bucket1=source.bucket, + key1=source.key, + version_id1=source.version_id, size1=size1, - bucket2=bucket2, - key2=key2, + bucket2=destination.bucket, + key2=destination.key, max_workers=max_workers, block_size=block_size, **kwargs, @@ -1775,8 +1774,8 @@ def _copy_object( } _logger.debug( - f"Copy object from s3://{bucket1}/{key1}?versionId={version_id1} " - f"to s3://{bucket2}/{key2}." + f"Copy object from {S3Path(bucket1, key1, version_id1).uri} " + f"to {S3Path(bucket2, key2).uri}." ) self._call(self._client.copy_object, **request, **kwargs) @@ -1983,7 +1982,7 @@ def _get_multipart_copy_kwargs( source = {"Bucket": bucket, "Key": key} if version_id: source.update({"VersionId": version_id}) - _logger.debug(f"Head object to copy: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Head object to copy: {S3Path(bucket, key, version_id).uri}") head = S3Metadata( self._call( self._client.head_object, @@ -2018,7 +2017,7 @@ def _get_multipart_copy_kwargs( # Directory buckets do not support GetObjectTagging, and their # objects have no tags. if not self._is_directory_bucket(bucket): - _logger.debug(f"Get tags to copy: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Get tags to copy: {S3Path(bucket, key, version_id).uri}") response = self._call( self._client.get_object_tagging, **self._get_operation_kwargs("get_object_tagging", source_kwargs), @@ -2084,7 +2083,7 @@ def _list_object_annotations( request.update({"VersionId": version_id}) names: list[str] = [] while True: - _logger.debug(f"List object annotations: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"List object annotations: {S3Path(bucket, key, version_id).uri}") response = self._call(self._client.list_object_annotations, **request) names.extend(a["AnnotationName"] for a in response.get("Annotations", [])) token = response.get("NextContinuationToken") @@ -2124,7 +2123,7 @@ def _copy_object_annotation( if version_id1: source.update({"VersionId": version_id1}) _logger.debug( - f"Copy object annotation {name} from s3://{bucket1}/{key1}?versionId={version_id1} " + f"Copy object annotation {name} from {S3Path(bucket1, key1, version_id1).uri} " f"to s3://{bucket2}/{key2}." ) response = self._call( @@ -2265,10 +2264,10 @@ def pipe_file( # large data as a parallel multipart upload. self.open(path, "xb" if mode == "create" else "wb", **kwargs)._write_and_close(value) return - bucket, key, version_id = self.parse_path(path) - if version_id: + s3_path = S3Path.parse(path) + if s3_path.version_id: raise ValueError("Cannot write to the file with the version specified.") - if not key: + if not s3_path.key: raise ValueError("Cannot write to a bucket.") kwargs.pop("block_size", None) kwargs.pop("max_workers", None) @@ -2288,7 +2287,7 @@ def pipe_file( # buffered path does. value = bytes(value) - self._put_object(bucket=bucket, key=key, body=value, **request_kwargs) + self._put_object(bucket=s3_path.bucket, key=s3_path.key, body=value, **request_kwargs) self.invalidate_cache(path) def _finish_multipart_upload( @@ -2400,12 +2399,12 @@ def cat_file( FileNotFoundError: If the path has no key or the key does not exist. """ - bucket, key, path_version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: raise FileNotFoundError(path) version_id = kwargs.pop("version_id", None) - if path_version_id: - version_id = path_version_id + if s3_path.version_id: + version_id = s3_path.version_id ranges: tuple[int, int | None] | None = None if start is not None and start < 0 and end is None: # S3 returns the last bytes, or the whole object when it is @@ -2419,7 +2418,10 @@ def cat_file( # empty range sends no GetObject request that would report a # missing object. info = self.info(path, version_id=version_id, **self._get_lookup_kwargs(kwargs)) - if info.get("type") == S3ObjectType.S3_OBJECT_TYPE_DIRECTORY or info.key != key: + if ( + info.get("type") == S3ObjectType.S3_OBJECT_TYPE_DIRECTORY + or info.key != s3_path.key + ): # There is no object to read, as GetObject reports for # the other ranges, or info() describes the key without # the trailing slash of this one. @@ -2433,8 +2435,8 @@ def cat_file( ranges = (start, end) try: return self._get_object( - bucket=bucket, - key=key, + bucket=s3_path.bucket, + key=s3_path.key, ranges=ranges, version_id=version_id, **kwargs, @@ -2493,8 +2495,7 @@ def put_file( # No support for directory uploads. return - bucket, key, _ = self.parse_path(rpath) - if not key: + if not S3Path.parse(rpath).key: # No support for bucket copy. return @@ -2541,7 +2542,7 @@ def get_file(self, rpath: str, lpath=None, callback=_DEFAULT_CALLBACK, outfile=N outfile: A file-like object to write to instead of ``lpath``. **kwargs: Additional S3 parameters passed to open(). """ - _, _, path_version_id = self.parse_path(self._strip_protocol(rpath)) + path_version_id = S3Path.parse(self._strip_protocol(rpath)).version_id if outfile is None and isfilelike(lpath): outfile = lpath elif ( @@ -2623,11 +2624,11 @@ def sign(self, path: str, expiration: int = 3600, **kwargs): ... client_method="put_object" ... ) """ - bucket, key, version_id = self.parse_path(path) + s3_path = S3Path.parse(path) client_method = kwargs.pop("client_method", "get_object") - params = {"Bucket": bucket, "Key": key} - if version_id: - params.update({"VersionId": version_id}) + params = {"Bucket": s3_path.bucket, "Key": s3_path.key} + if s3_path.version_id: + params.update({"VersionId": s3_path.version_id}) if kwargs: params.update(kwargs) request = { @@ -2636,7 +2637,7 @@ def sign(self, path: str, expiration: int = 3600, **kwargs): "ExpiresIn": expiration, } - _logger.debug(f"Generate signed url: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Generate signed url: {s3_path.uri}") return self._call( self._client.generate_presigned_url, **request, @@ -2655,14 +2656,14 @@ def metadata(self, path: str, **kwargs) -> S3Metadata: system-defined metadata (content type, encryption settings, etc.) as typed properties. """ - bucket, key, version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: raise ValueError("Cannot get metadata of a bucket.") - request: dict[str, Any] = {"Bucket": bucket, "Key": key} - if version_id: - request.update({"VersionId": version_id}) + request: dict[str, Any] = {"Bucket": s3_path.bucket, "Key": s3_path.key} + if s3_path.version_id: + request.update({"VersionId": s3_path.version_id}) - _logger.debug(f"Head object metadata: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Head object metadata: {s3_path.uri}") response = self._call( self._client.head_object, **request, @@ -2718,10 +2719,10 @@ def setxattr(self, path: str, copy_kwargs: dict[str, Any] | None = None, **kw_ar Raises: ValueError: If the path is a bucket or has a version ID. """ - bucket, key, version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: raise ValueError("Cannot set metadata of a bucket.") - if version_id: + if s3_path.version_id: raise ValueError("Cannot set metadata of a version.") head = self.metadata(path) metadata = dict(head) @@ -2754,12 +2755,12 @@ def setxattr(self, path: str, copy_kwargs: dict[str, Any] | None = None, **kw_ar } ) - _logger.debug(f"Set object metadata: s3://{bucket}/{key}") + _logger.debug(f"Set object metadata: {s3_path.uri}") self._call( self._client.copy_object, - CopySource={"Bucket": bucket, "Key": key}, - Bucket=bucket, - Key=key, + CopySource={"Bucket": s3_path.bucket, "Key": s3_path.key}, + Bucket=s3_path.bucket, + Key=s3_path.key, Metadata=metadata, MetadataDirective="REPLACE", **{ @@ -2778,14 +2779,14 @@ def get_tags(self, path: str) -> dict[str, str]: Returns: Dictionary mapping tag keys to tag values. """ - bucket, key, version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: raise ValueError("Cannot get tags of a bucket.") - request: dict[str, Any] = {"Bucket": bucket, "Key": key} - if version_id: - request.update({"VersionId": version_id}) + request: dict[str, Any] = {"Bucket": s3_path.bucket, "Key": s3_path.key} + if s3_path.version_id: + request.update({"VersionId": s3_path.version_id}) - _logger.debug(f"Get object tagging: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Get object tagging: {s3_path.uri}") response = self._call( self._client.get_object_tagging, **request, @@ -2808,8 +2809,8 @@ def put_tags(self, path: str, tags: dict[str, str], mode: str = "o") -> None: 'm' will merge in new tags with existing tags, which incurs two remote calls. """ - bucket, key, version_id = self.parse_path(path) - if not key: + s3_path = S3Path.parse(path) + if not s3_path.key: raise ValueError("Cannot put tags of a bucket.") if mode == "m": existing_tags = self.get_tags(path) @@ -2820,14 +2821,14 @@ def put_tags(self, path: str, tags: dict[str, str], mode: str = "o") -> None: else: raise ValueError(f"Mode must be {{'o', 'm'}}, not {mode}.") request: dict[str, Any] = { - "Bucket": bucket, - "Key": key, + "Bucket": s3_path.bucket, + "Key": s3_path.key, "Tagging": {"TagSet": new_tags}, } - if version_id: - request.update({"VersionId": version_id}) + if s3_path.version_id: + request.update({"VersionId": s3_path.version_id}) - _logger.debug(f"Put object tagging: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Put object tagging: {s3_path.uri}") self._call( self._client.put_object_tagging, **request, @@ -2846,12 +2847,12 @@ def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: **kwargs: Additional parameters passed to the PutObjectAcl or PutBucketAcl API. """ - bucket, key, version_id = self.parse_path(path) + s3_path = S3Path.parse(path) # Validate before any ACL is applied so that a recursive call cannot # partially apply object ACLs and then fail on the bucket ACL. - if not key and acl not in self.BUCKET_ACLS: + if not s3_path.key and acl not in self.BUCKET_ACLS: raise ValueError(f"ACL not in {self.BUCKET_ACLS}.") - if key and acl not in self.OBJECT_ACLS: + if s3_path.key and acl not in self.OBJECT_ACLS: raise ValueError(f"ACL not in {self.OBJECT_ACLS}.") if recursive: with self._create_executor(max_workers=self.max_workers) as executor: @@ -2861,26 +2862,26 @@ def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: ] for future in as_completed(futures): future.result() - if key: + if s3_path.key: # A key prefix is not an object itself; only the objects # below it have ACLs. return - if key: - request: dict[str, Any] = {"Bucket": bucket, "Key": key, "ACL": acl} - if version_id: - request.update({"VersionId": version_id}) + if s3_path.key: + request: dict[str, Any] = {"Bucket": s3_path.bucket, "Key": s3_path.key, "ACL": acl} + if s3_path.version_id: + request.update({"VersionId": s3_path.version_id}) - _logger.debug(f"Put object acl: s3://{bucket}/{key}?versionId={version_id}") + _logger.debug(f"Put object acl: {s3_path.uri}") self._call( self._client.put_object_acl, **request, **kwargs, ) else: - _logger.debug(f"Put bucket acl: s3://{bucket}") + _logger.debug(f"Put bucket acl: {s3_path.uri}") self._call( self._client.put_bucket_acl, - Bucket=bucket, + Bucket=s3_path.bucket, ACL=acl, **kwargs, ) @@ -2903,19 +2904,19 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: List of S3MultipartUpload instances describing the in-progress multipart uploads. """ - bucket, key, _ = self.parse_path(path) + s3_path = S3Path.parse(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 "" + prefix = f"{s3_path.key.rstrip('/')}/" if s3_path.key else "" - _logger.debug(f"List multipart uploads: s3://{bucket}/{key}") + _logger.debug(f"List multipart uploads: {s3_path.uri}") uploads: list[S3MultipartUpload] = [] next_key_marker: str | None = None next_upload_id_marker: str | None = None while True: - request: dict[str, Any] = {"Bucket": bucket} - if key: - request.update({"Prefix": key}) + request: dict[str, Any] = {"Bucket": s3_path.bucket} + if s3_path.key: + request.update({"Prefix": s3_path.key}) if next_key_marker: request.update( {"KeyMarker": next_key_marker, "UploadIdMarker": next_upload_id_marker} @@ -2925,9 +2926,9 @@ def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: **request, ) uploads.extend( - S3MultipartUpload({**u, "Bucket": bucket}) + S3MultipartUpload({**u, "Bucket": s3_path.bucket}) for u in response.get("Uploads", []) - if u["Key"] == key or u["Key"].startswith(prefix) + if u["Key"] == s3_path.key or u["Key"].startswith(prefix) ) if not response.get("IsTruncated"): break @@ -2960,30 +2961,32 @@ def object_version_info( Returns: List of S3ObjectVersion instances describing the versions. """ - bucket, key, _ = self.parse_path(path) + s3_path = S3Path.parse(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 "" + prefix = f"{s3_path.key.rstrip('/')}/" if s3_path.key else "" - _logger.debug(f"List object versions: s3://{bucket}/{key}") + _logger.debug(f"List object versions: {s3_path.uri}") versions: list[S3ObjectVersion] = [] - for response in self._list_object_versions_pages(bucket, prefix=key or "", **kwargs): + for response in self._list_object_versions_pages( + s3_path.bucket, prefix=s3_path.key or "", **kwargs + ): versions.extend( - S3ObjectVersion(bucket=bucket, is_delete_marker=False, response=v) + S3ObjectVersion(bucket=s3_path.bucket, is_delete_marker=False, response=v) for v in response.get("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) + S3ObjectVersion(bucket=s3_path.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] + if s3_path.key and not s3_path.key.endswith("/") and s3_path.key in keys: + selected = [v for v, k in zip(versions, keys, strict=True) if k == s3_path.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] @@ -3224,10 +3227,10 @@ def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None: cache = self.dircache.get(path.rstrip("/")) if cache is not None: return cast("list[S3Object] | S3Object", cache) - _, key, version_id = self.parse_path(path) - if version_id: + s3_path = S3Path.parse(path) + if s3_path.version_id: return None - if key: + if s3_path.key: parent_cache = self.dircache.get((self._parent(path), "/")) else: parent_cache = self.dircache.get("") @@ -3243,7 +3246,7 @@ def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None: ] if files: return files - if key: + if s3_path.key: raise FileNotFoundError(path) # The bucket listing holds only the buckets that the caller owns, # so a bucket missing from it is looked up with HeadBucket. @@ -3325,7 +3328,7 @@ def _get_object( if version_id: request.update({"VersionId": version_id}) - _logger.debug(f"Get object: s3://{bucket}/{key}?versionId={version_id}&range={range_}") + _logger.debug(f"Get object: {S3Path(bucket, key, version_id).uri} range={range_}") response = self._call( self._client.get_object, # The fields of the request take precedence over inherited @@ -3559,11 +3562,12 @@ def __init__( # The arguments are validated, and the objects looked up, before the # base class initializer: a file that fails here is never opened, so # its garbage collection does not close (flush and commit) it. - bucket, key, path_version_id = S3FileSystem.parse_path(path) - self.bucket = bucket - if not key: + s3_path = S3Path.parse(path) + path_version_id = s3_path.version_id + self.bucket = s3_path.bucket + if not s3_path.key: raise ValueError("The path does not contain a key.") - self.key = key + self.key = s3_path.key if version_id and path_version_id: if version_id != path_version_id: raise ValueError( diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index cb8f9f133..0d98a28a4 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -29,6 +29,7 @@ S3ObjectType, S3ObjectVersion, ) +from pyathena.filesystem.s3_path import S3Path if TYPE_CHECKING: from datetime import datetime @@ -253,8 +254,7 @@ def _put_file_in_transaction( """ if os.path.isdir(lpath): return - _, key, _ = self.parse_path(rpath) - if not key: + if not S3Path.parse(rpath).key: return size = os.path.getsize(lpath) @@ -429,11 +429,11 @@ async def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: # Parameters of the multipart copy, not of the S3 requests. block_size = kwargs.pop("block_size", None) max_workers = kwargs.pop("max_workers", None) - bucket1, key1, version_id1 = self.parse_path(path1) - bucket2, key2, version_id2 = self.parse_path(path2) - if version_id2: + source = S3Path.parse(path1) + destination = S3Path.parse(path2) + if destination.version_id: raise ValueError("Cannot copy to a versioned file.") - if not key1 or not key2: + if not source.key or not destination.key: raise ValueError("Cannot copy buckets.") info1 = await self._info(path1) @@ -446,21 +446,21 @@ async def _copy_file(self, path1: str, path2: str, **kwargs) -> bool: if size1 <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: await asyncio.to_thread( self._sync_fs._copy_object, - bucket1=bucket1, - key1=key1, - version_id1=version_id1, - bucket2=bucket2, - key2=key2, + bucket1=source.bucket, + key1=source.key, + version_id1=source.version_id, + bucket2=destination.bucket, + key2=destination.key, **kwargs, ) else: await self._copy_object_with_multipart_upload( - bucket1=bucket1, - key1=key1, - version_id1=version_id1, + bucket1=source.bucket, + key1=source.key, + version_id1=source.version_id, size1=size1, - bucket2=bucket2, - key2=key2, + bucket2=destination.bucket, + key2=destination.key, max_workers=max_workers, block_size=block_size, **kwargs, diff --git a/pyathena/filesystem/s3_path.py b/pyathena/filesystem/s3_path.py new file mode 100644 index 000000000..3013092c1 --- /dev/null +++ b/pyathena/filesystem/s3_path.py @@ -0,0 +1,107 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""The S3 path model of the S3 filesystem.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass, replace +from re import Pattern +from typing import ClassVar + +from pyathena.util import override + + +@dataclass(frozen=True) +class S3Path: + """An S3 path: a bucket, optionally a key, and optionally a version ID. + + Paths are parsed from strings such as ``s3://bucket/key``, + ``s3a://bucket/key`` or ``bucket/key``, optionally followed by a version + ID query (``?versionId=``, ``?versionID=``, ``?versionid=`` or + ``?version_id=``). The root path, which names no bucket, is not an + ``S3Path``. + + Attributes: + bucket: The name of the bucket. + key: The key, or None for a bucket path. A trailing slash is kept, + and a key of only slashes (``bucket//``) names the bucket (see + ``is_bucket``). + version_id: The version ID, or None for a path without a version. + + Example: + >>> path = S3Path.parse("s3://bucket/dir/key?versionId=v1") + >>> path.bucket, path.key, path.version_id + ('bucket', 'dir/key', 'v1') + >>> path.name + 'bucket/dir/key' + >>> str(path) # The name, with a version ID query if there is one. + 'bucket/dir/key?versionId=v1' + >>> path.uri + 's3://bucket/dir/key?versionId=v1' + """ + + PATTERN: ClassVar[Pattern[str]] = re.compile( + r"(^s3://|^s3a://|^)(?P[a-zA-Z0-9.\-_]+)(/(?P[^?]+)|/)?" + r"($|\?version(Id|ID|id|_id)=(?P.+)$)" + ) + + bucket: str + key: str | None = None + version_id: str | None = None + + @classmethod + def parse(cls, path: str) -> S3Path: + """Parse a string into an S3 path. + + Args: + path: The S3 path (e.g., "s3://bucket/key?versionId=..."). + + Returns: + The path. The key is None for a bucket path (``bucket`` or + ``bucket/``), and keeps a trailing slash otherwise. + + Raises: + ValueError: If the string is not a valid S3 path. + """ + match = cls.PATTERN.search(path) + if not match: + raise ValueError(f"Invalid S3 path format {path}.") + return cls(match.group("bucket"), match.group("key"), match.group("version_id")) + + @property + def is_bucket(self) -> bool: + """Whether the path names the bucket: it has no key, or a key of only slashes.""" + return not self.key or not self.key.strip("/") + + @property + def name(self) -> str: + """The path without a scheme or version, in ``bucket/key`` form, or the bucket.""" + return f"{self.bucket}/{self.key}" if self.key else self.bucket + + @property + def uri(self) -> str: + """The path as an ``s3://`` URI, with its version ID query, if any.""" + return f"s3://{self}" + + def with_version_id(self, version_id: str | None) -> S3Path: + """Return the path with another version ID. + + Args: + version_id: The version ID, or None for the path without a version. + + Returns: + The path with the version ID. + """ + return replace(self, version_id=version_id) + + @override + def __str__(self) -> str: + if self.version_id: + return f"{self.name}?versionId={self.version_id}" + return self.name diff --git a/tests/pyathena/filesystem/test_s3_path.py b/tests/pyathena/filesystem/test_s3_path.py new file mode 100644 index 000000000..56908b2f0 --- /dev/null +++ b/tests/pyathena/filesystem/test_s3_path.py @@ -0,0 +1,80 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +import pytest + +from pyathena.filesystem.s3_path import S3Path + + +class TestS3Path: + @pytest.mark.parametrize( + ("path", "expected"), + [ + ("s3://bucket", S3Path("bucket")), + ("s3a://bucket/", S3Path("bucket")), + ("bucket", S3Path("bucket")), + ("s3://bucket/path/to/obj", S3Path("bucket", "path/to/obj")), + ("s3a://bucket/path/to/obj", S3Path("bucket", "path/to/obj")), + ("bucket/path/to/dir/", S3Path("bucket", "path/to/dir/")), + ("bucket//", S3Path("bucket", "/")), + ("s3://bucket/obj?versionId=v1", S3Path("bucket", "obj", "v1")), + ("bucket/obj?versionID=v1", S3Path("bucket", "obj", "v1")), + ("bucket/obj?versionid=v1", S3Path("bucket", "obj", "v1")), + ("bucket/obj?version_id=v1", S3Path("bucket", "obj", "v1")), + ("bucket?versionId=v1", S3Path("bucket", None, "v1")), + ], + ) + def test_parse(self, path, expected): + assert S3Path.parse(path) == expected + + @pytest.mark.parametrize("path", ["", "s3://", "http://bucket", "bucket/obj?x=1"]) + def test_parse_invalid(self, path): + with pytest.raises(ValueError, match="Invalid S3 path format"): + S3Path.parse(path) + + @pytest.mark.parametrize( + ("path", "expected"), + [ + (S3Path("bucket"), True), + (S3Path("bucket", "/"), True), + (S3Path("bucket", "//"), True), + (S3Path("bucket", "key"), False), + (S3Path("bucket", "dir/"), False), + ], + ) + def test_is_bucket(self, path, expected): + assert path.is_bucket is expected + + @pytest.mark.parametrize( + ("path", "name", "string", "uri"), + [ + (S3Path("bucket"), "bucket", "bucket", "s3://bucket"), + (S3Path("bucket", "dir/"), "bucket/dir/", "bucket/dir/", "s3://bucket/dir/"), + ( + S3Path("bucket", "key", "v1"), + "bucket/key", + "bucket/key?versionId=v1", + "s3://bucket/key?versionId=v1", + ), + ], + ) + def test_names(self, path, name, string, uri): + assert path.name == name + assert str(path) == string + assert path.uri == uri + + def test_str_spells_the_version_query_as_version_id(self): + assert str(S3Path.parse("s3a://bucket/key?version_id=v1")) == "bucket/key?versionId=v1" + + def test_with_version_id(self): + path = S3Path("bucket", "key", "v1") + assert path.with_version_id("v2") == S3Path("bucket", "key", "v2") + assert path.with_version_id(None) == S3Path("bucket", "key") + # The path is a frozen value. + assert path == S3Path("bucket", "key", "v1") + with pytest.raises(AttributeError): + path.key = "other" # type: ignore[misc]