diff --git a/docs/api/filesystem.rst b/docs/api/filesystem.rst index 170e1c2f..b7c78f56 100644 --- a/docs/api/filesystem.rst +++ b/docs/api/filesystem.rst @@ -49,6 +49,9 @@ S3 Paths .. autoclass:: pyathena.filesystem.s3_path.S3Path :members: +.. autoclass:: pyathena.filesystem.s3_path_pairing.S3PathPairing + :members: + S3 Core ------- diff --git a/docs/filesystem.md b/docs/filesystem.md index 765cf0c8..181b09b3 100644 --- a/docs/filesystem.md +++ b/docs/filesystem.md @@ -371,6 +371,31 @@ except Exception: raise ``` +## Path pairing + +`S3PathPairing` pairs the paths of one `copy()`, `get()` or `mv()`: `mv()` always, and +`copy()` and `get()` when a source has a version ID and the destination is one path +(fsspec pairs the others). Its `delete_paths()` splits the paths of an `rm()` into +those deleted as given and those expanded. The pairing is fsspec's, except that a +path with a version ID names that version, and its destination is named after its key. + +A pairing is a frozen dataclass of `path1`, `path2`, `recursive` and `maxdepth`, and +holds no filesystem. The filesystems read from it what to look up (`expands`, +`skips_directories`, `looks_up_destination` and `conflict_candidates()`), expand the +sources, leave out the directories when `skips_directories` is true, look up the +paths, and pass the results to `copy_pairs()` and `move_pairs()`. A `sources`, +`destination_is_dir` or `missing` that a rule needs and that is not passed raises +`ValueError`. + +```python +from pyathena.filesystem.s3_path_pairing import S3PathPairing + +pairing = S3PathPairing("s3://YOUR_S3_BUCKET/src/", "s3://YOUR_S3_BUCKET/dst/", recursive=True) +sources = fs.expand_path(pairing.path1, recursive=pairing.recursive) +for source, destination in pairing.copy_pairs(sources): + print(source, "->", destination) +``` + ## Async filesystem `AioS3FileSystem` provides the same functionality on top of fsspec's diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 387c814d..3e2f8a01 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -9,12 +9,10 @@ import mimetypes import os.path import time -from collections import Counter from collections.abc import Callable, Mapping from concurrent.futures import Future, as_completed, wait from copy import deepcopy from datetime import datetime -from glob import has_magic from io import BytesIO from multiprocessing import cpu_count from re import Pattern @@ -31,7 +29,7 @@ from fsspec.core import get_compression from fsspec.implementations.local import LocalFileSystem, make_path_posix, trailing_sep from fsspec.spec import AbstractBufferedFile -from fsspec.utils import check_contained, isfilelike, other_paths, tokenize +from fsspec.utils import check_contained, isfilelike, tokenize import pyathena from pyathena.connection import Connection @@ -50,6 +48,7 @@ S3StorageClass, ) from pyathena.filesystem.s3_path import S3Path +from pyathena.filesystem.s3_path_pairing import S3PathPairing from pyathena.util import RetryConfig, override _logger = logging.getLogger(__name__) @@ -1100,13 +1099,13 @@ def rm(self, path, recursive=False, maxdepth=None, **kwargs) -> None: ValueError: If a path is a bucket. OSError: If S3 could not delete some of the objects. """ - paths = self._expand_delete_paths(path, recursive=recursive, maxdepth=maxdepth) + paths = self._delete_paths(path, recursive=recursive, maxdepth=maxdepth) self._delete_objects(paths, **kwargs) - def _expand_delete_paths( + def _delete_paths( self, path: str | list[str], recursive: bool = False, maxdepth: int | None = None ) -> list[str]: - """Expand the paths that ``rm`` deletes. + """Expand the paths that ``rm()`` deletes (see :meth:`S3PathPairing.delete_paths`). Args: path: S3 path or list of paths. @@ -1120,25 +1119,13 @@ def _expand_delete_paths( Raises: ValueError: If a path is a bucket. """ - paths = [path] if isinstance(path, str) else list(path) - versioned_paths, unversioned_paths = [], [] - for p in paths: - s3_path = S3Path.parse(p) - # expand_path strips the slashes of "bucket//" to the bucket. - if s3_path.is_bucket: - raise ValueError("Cannot delete the bucket.") - if s3_path.version_id: - versioned_paths.append(p) - else: - unversioned_paths.append(p) - - if unversioned_paths: - # Versioned paths are deleted as given, without the lookup that - # expand_path makes for them with recursive. - unversioned_paths = self.expand_path( - unversioned_paths, recursive=recursive, maxdepth=maxdepth - ) - return versioned_paths + unversioned_paths + versioned_paths, unversioned_paths = S3PathPairing.delete_paths(path) + if not unversioned_paths: + # expand_path raises FileNotFoundError for no paths. + return versioned_paths + return versioned_paths + self.expand_path( + unversioned_paths, recursive=recursive, maxdepth=maxdepth + ) def _create_executor(self, max_workers: int) -> S3Executor: """Create an executor strategy for parallel operations. @@ -1449,15 +1436,19 @@ def mv(self, path1, path2, recursive=False, maxdepth=None, **kwargs) -> None: return copied = [ p1 - for p1, p2 in self._move_paths(path1, path2, recursive=recursive, maxdepth=maxdepth) + for p1, p2 in self._move_pairs(path1, path2, recursive=recursive, maxdepth=maxdepth) if self._copy_file(p1, p2, **kwargs) ] self._delete_objects(copied) - def _move_paths( - self, path1, path2, recursive: bool = False, maxdepth: int | None = None + def _move_pairs( + self, + path1: str | list[str], + path2: str | list[str], + recursive: bool = False, + maxdepth: int | None = None, ) -> list[tuple[str, str]]: - """Pair the sources and destinations of a move as fsspec's ``copy()`` does. + """Pair and check the paths of a move (see :meth:`S3PathPairing.move_pairs`). Args: path1: Source S3 path, glob pattern, or list of paths. @@ -1467,111 +1458,57 @@ def _move_paths( maxdepth: Maximum depth of the expansion. Returns: - The source and destination paths, except the sources whose - destination is the source itself or, for a ``null`` version, the - key of the source. + The sources and destinations that are moved. Raises: - ValueError: If two sources have the same destination, or a - destination is another source, including one left in place, - except for a directory with no object at its key, which is not - copied. - """ - paths1, paths2 = self._copy_paths(path1, path2, recursive=recursive, maxdepth=maxdepth) - # The paths are copied as given, and compared by what they name. - named = [ - (p1, p2, self._move_target(p1), self._move_target(p2)) - for p1, p2 in zip(paths1, paths2, strict=False) - ] - pairs = [(p1, p2) for p1, p2, source, dest in named if source != dest] - moved = [(p1, source, dest) for p1, _, source, dest in named if source != dest] - # The sources left in place count too; a copy onto one overwrites it. - sources = {source for _, _, source, _ in named} - counts = Counter(dest for _, _, dest in moved) - # A source with another source below it may be a directory. - directories: set[str] = set() - for source in sources: - parent = source.rpartition("/")[0] - while parent and parent not in directories: - directories.add(parent) - parent = parent.rpartition("/")[0] - # A directory without an object at its key is not copied, so it - # writes no destination and is left out of the checks. A path with a - # version always names an object. - writers = [ - (source, dest) - for p1, source, dest in moved - if not ( - (counts[dest] > 1 or dest in sources) - and source in directories - and not S3Path.parse(p1).version_id - and self._head_object(source) is None - ) - ] - counts = Counter(dest for _, dest in writers) - for _, dest in writers: - if counts[dest] > 1: - raise ValueError("Cannot move several paths to the same destination.") - if dest in sources: - raise ValueError("Cannot move a path onto another path that is moved.") - return pairs - - def _copy_paths( - self, - path1: str | list[str], - path2: str | list[str], - recursive: bool = False, - maxdepth: int | None = None, - isdir: Callable[[str], bool] | None = None, - ) -> tuple[list[str], list[str]]: - """Pair the sources of a copy with their destinations as fsspec's ``copy()`` does. + ValueError: If the move has conflicting paths. + """ + pairing = S3PathPairing(path1, path2, recursive=recursive, maxdepth=maxdepth) + pairs = self._copy_pairs(pairing) + missing = { + source + for source in pairing.conflict_candidates(pairs) + if self._head_object(source) is None + } + return pairing.move_pairs(pairs, missing=missing) + + def _copy_pairs( + self, pairing: S3PathPairing, isdir: Callable[[str], bool] | None = None + ) -> list[tuple[str, str]]: + """Expand and pair the paths of a copy (see :meth:`S3PathPairing.copy_pairs`). - A source with a version ID names that version of an object: it is not - a glob pattern, and its destination is named after its key without - the version. + The destination is looked up only when it decides the pairing. Args: - path1: Source S3 path, glob pattern, or list of them. - path2: Destination path, or list of paths when ``path1`` is a - list. - recursive: Whether to include the contents of the directories. - maxdepth: Maximum depth of the expansion. - isdir: Whether a destination path is a directory, by default + pairing: The paths of the copy. + isdir: Whether the destination is a directory, by default ``self.isdir``; ``get()`` passes the local filesystem's. Returns: - The sources and their destinations. Both are empty if ``path1`` - is a string that matches only directories without ``recursive``. - """ - if isinstance(path1, list) and isinstance(path2, list): - return path1, path2 - source_is_str = isinstance(path1, str) - paths1 = self.expand_path(path1, recursive=recursive, maxdepth=maxdepth) - if source_is_str and (not recursive or maxdepth is not None): - # Non-recursive glob does not copy directories. - paths1 = [p for p in paths1 if not (trailing_sep(p) or self.isdir(p))] - if not paths1: - return [], [] - glob = isinstance(path1, str) and has_magic(path1) and not S3Path.has_version_id(path1) - # The destination is looked up only when it decides the mapping. - exists = source_is_str and ( - (glob and len(paths1) == 1) - or ( - not glob - and not trailing_sep(path1) - and isinstance(path2, str) - and (trailing_sep(path2) or (isdir or self.isdir)(path2)) - ) + The sources and their destinations. + """ + if not pairing.expands: + return pairing.copy_pairs() + sources = self.expand_path( + pairing.path1, recursive=pairing.recursive, maxdepth=pairing.maxdepth + ) + if pairing.skips_directories: + sources = [p for p in sources if not (trailing_sep(p) or self.isdir(p))] + destination_is_dir = ( + # A string, as looks_up_destination checks. + (isdir or self.isdir)(cast(str, pairing.path2)) + if sources and pairing.looks_up_destination + else None ) - names = [S3Path.split_version_id(p)[0] for p in paths1] - return paths1, other_paths(names, path2, exists=exists, flatten=not source_is_str) + return pairing.copy_pairs(sources, destination_is_dir) def copy(self, path1, path2, recursive=False, maxdepth=None, on_error=None, **kwargs) -> None: """Copy files within S3. As fsspec's ``copy()``, except that a source with a version ID copies that version of the object to a destination named after its key, as - ``_copy_paths`` pairs them. + :meth:`~pyathena.filesystem.s3_path_pairing.S3PathPairing.copy_pairs` + pairs them. Args: path1: Source S3 path, glob pattern, or list of them. @@ -1586,9 +1523,12 @@ def copy(self, path1, path2, recursive=False, maxdepth=None, on_error=None, **kw """ sources = [path1] if isinstance(path1, (str, os.PathLike)) else path1 if isinstance(path2, str) and any(S3Path.has_version_id(p) for p in sources): - path1, path2 = self._copy_paths(path1, path2, recursive=recursive, maxdepth=maxdepth) - if not path1: + pairs = self._copy_pairs( + S3PathPairing(path1, path2, recursive=recursive, maxdepth=maxdepth) + ) + if not pairs: return + path1, path2 = [p1 for p1, _ in pairs], [p2 for _, p2 in pairs] super().copy( path1, path2, recursive=recursive, maxdepth=maxdepth, on_error=on_error, **kwargs ) @@ -1600,8 +1540,8 @@ def get( As fsspec's ``get()``, except that a source with a version ID downloads that version of the object to a local path named after its - key, as ``_copy_paths`` pairs them. Those destinations are checked to - lie under ``lpath``. + key, as :meth:`~pyathena.filesystem.s3_path_pairing.S3PathPairing.copy_pairs` + pairs them. Those destinations are checked to lie under ``lpath``. Args: rpath: Source S3 path, glob pattern, or list of them. @@ -1619,9 +1559,11 @@ def get( sources = [rpath] if isinstance(rpath, (str, os.PathLike)) else rpath if isinstance(lpath, (str, os.PathLike)) and any(S3Path.has_version_id(p) for p in sources): root = make_path_posix(lpath) - rpath, lpath = self._copy_paths( - rpath, root, recursive=recursive, maxdepth=maxdepth, isdir=LocalFileSystem().isdir + pairs = self._copy_pairs( + S3PathPairing(rpath, root, recursive=recursive, maxdepth=maxdepth), + isdir=LocalFileSystem().isdir, ) + rpath, lpath = [p1 for p1, _ in pairs], [p2 for _, p2 in pairs] check_contained(root, lpath) if not rpath: return @@ -1629,24 +1571,6 @@ def get( rpath, lpath, recursive=recursive, callback=callback, maxdepth=maxdepth, **kwargs ) - def _move_target(self, path: str) -> str: - """Return what a path of a move names, for comparing the paths. - - A write to a key replaces its ``null`` version, which the objects of a - bucket without versioning have, so that version names the key itself. - - Args: - path: S3 path, possibly with a version ID. - - Returns: - The path in ``bucket/key`` form, with the version ID unless it is - ``null``. - """ - 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 ): diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 75f95cf3..59256218 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -35,6 +35,7 @@ S3ObjectVersion, ) from pyathena.filesystem.s3_path import S3Path +from pyathena.filesystem.s3_path_pairing import S3PathPairing if TYPE_CHECKING: from datetime import datetime @@ -324,7 +325,7 @@ async def _rm( OSError: If S3 could not delete some of the objects. """ paths = await asyncio.to_thread( - self._sync_fs._expand_delete_paths, path, recursive=recursive, maxdepth=maxdepth + self._sync_fs._delete_paths, path, recursive=recursive, maxdepth=maxdepth ) await self._delete_objects(paths, **kwargs) @@ -372,7 +373,7 @@ async def _mv(self, path1, path2, recursive=False, maxdepth=None, **kwargs) -> N if path1 == path2: return pairs = await asyncio.to_thread( - self._sync_fs._move_paths, path1, path2, recursive=recursive, maxdepth=maxdepth + self._sync_fs._move_pairs, path1, path2, recursive=recursive, maxdepth=maxdepth ) # Every copy finishes before a failure is raised, as in fsspec's # _copy(), and nothing is deleted after a failure. @@ -429,11 +430,13 @@ async def _copy( """ sources = [path1] if isinstance(path1, (str, os.PathLike)) else path1 if isinstance(path2, str) and any(S3Path.has_version_id(p) for p in sources): - path1, path2 = await asyncio.to_thread( - self._sync_fs._copy_paths, path1, path2, recursive=recursive, maxdepth=maxdepth + pairs = await asyncio.to_thread( + self._sync_fs._copy_pairs, + S3PathPairing(path1, path2, recursive=recursive, maxdepth=maxdepth), ) - if not path1: + if not pairs: return + path1, path2 = [p1 for p1, _ in pairs], [p2 for _, p2 in pairs] await super()._copy( path1, path2, @@ -468,14 +471,12 @@ async def _get( sources = [rpath] if isinstance(rpath, (str, os.PathLike)) else rpath if isinstance(lpath, (str, os.PathLike)) and any(S3Path.has_version_id(p) for p in sources): root = make_path_posix(lpath) - rpath, lpath = await asyncio.to_thread( - self._sync_fs._copy_paths, - rpath, - root, - recursive=recursive, - maxdepth=maxdepth, + pairs = await asyncio.to_thread( + self._sync_fs._copy_pairs, + S3PathPairing(rpath, root, recursive=recursive, maxdepth=maxdepth), isdir=LocalFileSystem().isdir, ) + rpath, lpath = [p1 for p1, _ in pairs], [p2 for _, p2 in pairs] check_contained(root, lpath) if not rpath: return diff --git a/pyathena/filesystem/s3_path.py b/pyathena/filesystem/s3_path.py index 6b2af357..1f2a4f38 100644 --- a/pyathena/filesystem/s3_path.py +++ b/pyathena/filesystem/s3_path.py @@ -133,6 +133,18 @@ def uri(self) -> str: """The path as an ``s3://`` URI, with its version ID query, if any.""" return f"s3://{self}" + @property + def target(self) -> S3Path: + """What the path names when paths are compared: a ``null`` version names its key. + + A ``null`` version is taken to name the key itself: in a bucket without + versioning, or with versioning suspended, a write to the key replaces + its ``null`` version. With versioning enabled, a write adds a new + version instead, but the ``null`` version is still taken to name the + key. Any other path is its own target. + """ + return self.with_version_id(None) if self.version_id == "null" else self + def with_version_id(self, version_id: str | None) -> S3Path: """Return the path with another version ID. diff --git a/pyathena/filesystem/s3_path_pairing.py b/pyathena/filesystem/s3_path_pairing.py new file mode 100644 index 00000000..d4233416 --- /dev/null +++ b/pyathena/filesystem/s3_path_pairing.py @@ -0,0 +1,292 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""The pairing of the paths that the S3 filesystem copies, moves and deletes.""" + +from __future__ import annotations + +from collections import Counter +from collections.abc import Collection, Sequence +from dataclasses import dataclass +from glob import has_magic + +from fsspec.implementations.local import trailing_sep +from fsspec.utils import other_paths + +from pyathena.filesystem.s3_path import S3Path + + +@dataclass(frozen=True) +class S3PathPairing: + """The pairing of the paths of one ``copy()``, ``get()`` or ``mv()``. + + The sources are paired with their destinations as fsspec's ``copy()`` + pairs them, except that a path with a version ID names that version of an + object: it is not a glob pattern, and its destination is named after its + key without the version. A move compares the paths by what they name + (see :attr:`~pyathena.filesystem.s3_path.S3Path.target`). + :meth:`delete_paths` splits the paths of an ``rm()``. + + The pairing is a frozen dataclass of the paths as given, and holds no + filesystem. The caller makes the lookups that it asks for + (:attr:`expands`, :attr:`skips_directories`, :attr:`looks_up_destination` + and :meth:`conflict_candidates`) with its own requests and cache: it + expands the sources, leaves out the directories when + :attr:`skips_directories` says so, and passes the results in. A + ``sources``, ``destination_is_dir`` or ``missing`` that a rule needs and + that is not passed raises ``ValueError``. + + Attributes: + path1: The source path, glob pattern, or list of them, as given to + the copy. + path2: The destination path, or list of paths. + recursive: Whether the copy includes the contents of directories. + maxdepth: The maximum depth of the expansion. + + Example: + >>> pairing = S3PathPairing("s3://bucket/dir/", "s3://bucket/copy/", recursive=True) + >>> sources = fs.expand_path(pairing.path1, recursive=True) + >>> pairing.copy_pairs(sources) + """ + + path1: str | list[str] + path2: str | list[str] + recursive: bool = False + maxdepth: int | None = None + + @property + def expands(self) -> bool: + """Whether the sources are expanded. + + False if both paths are lists, which are paired as given, without + expansion or lookups. + """ + return not (isinstance(self.path1, list) and isinstance(self.path2, list)) + + @property + def skips_directories(self) -> bool: + """Whether the directories among the expanded sources are left out. + + True for a string source copied without ``recursive``, or with a + ``maxdepth``. A directory is a source that ends with a slash or that + the filesystem reports as a directory. + """ + return isinstance(self.path1, str) and (not self.recursive or self.maxdepth is not None) + + @property + def looks_up_destination(self) -> bool: + """Whether the pairing depends on the destination being a directory. + + True for a string source that is neither a glob pattern nor ends with + a slash, copied to a string destination that does not end with a + slash. :meth:`copy_pairs` then needs ``destination_is_dir``. + """ + return ( + isinstance(self.path1, str) + and not self._is_glob(self.path1) + and not trailing_sep(self.path1) + and isinstance(self.path2, str) + and not trailing_sep(self.path2) + ) + + def copy_pairs( + self, + sources: Sequence[str] | None = None, + destination_is_dir: bool | None = None, + ) -> list[tuple[str, str]]: + """Pair the sources with their destinations as fsspec's ``copy()`` does. + + Args: + sources: The expansion of ``path1``, without the directories that + :attr:`skips_directories` leaves out; needed when + :attr:`expands`. + destination_is_dir: Whether ``path2`` is a directory; needed when + :attr:`looks_up_destination` is true and there are sources. + + Returns: + The sources and their destinations, which keep the form of + ``path2``. When both paths are lists, they are paired as given, up + to the end of the shorter list, as in fsspec's ``copy()``. Empty + if there are no sources. + + Raises: + ValueError: If ``sources`` or ``destination_is_dir`` is needed and + None. + """ + path1, path2 = self.path1, self.path2 + if not self.expands: + return list(zip(path1, path2, strict=False)) + if sources is None: + raise ValueError("sources is needed to pair the paths.") + if not sources: + return [] + if destination_is_dir is None and self.looks_up_destination: + raise ValueError("destination_is_dir is needed to pair the paths.") + source_is_str = isinstance(path1, str) + glob = isinstance(path1, str) and self._is_glob(path1) + # As in fsspec's copy(); destination_is_dir is None only when + # looks_up_destination is false, where the other terms decide. + dest_is_dir = isinstance(path2, str) and (trailing_sep(path2) or bool(destination_is_dir)) + exists = source_is_str and ( + (glob and len(sources) == 1) or (not glob and dest_is_dir and not trailing_sep(path1)) + ) + names = [S3Path.split_version_id(p)[0] for p in sources] + destinations = other_paths(names, path2, exists=exists, flatten=not source_is_str) + return list(zip(sources, destinations, strict=True)) + + def conflict_candidates(self, pairs: Sequence[tuple[str, str]]) -> list[str]: + """Return the sources of a move whose conflicts depend on an object at their key. + + A source with another source below it may be a directory. If no object + exists at its key, it is not copied, so it writes no destination and + is left out of the conflict checks of :meth:`move_pairs`. A path with + a version always names an object. + + Args: + pairs: The sources and destinations of the move, as + :meth:`copy_pairs` of this pairing returns them. + + Returns: + The sources, in ``bucket/key`` form (their + :attr:`~pyathena.filesystem.s3_path.S3Path.target`), that have + another source below them and a destination that conflicts, one + per pair in the order of the pairs; empty if nothing needs to be + looked up. + """ + return self._moves(pairs)[2] + + def move_pairs( + self, pairs: Sequence[tuple[str, str]], missing: Collection[str] | None = None + ) -> list[tuple[str, str]]: + """Check the pairs of a move and leave out the sources that stay in place. + + Args: + pairs: The sources and destinations of the move, as + :meth:`copy_pairs` of this pairing returns them. + missing: The :meth:`conflict_candidates` without an object at their + key, in any form that names them; needed when there are + candidates. + + Returns: + The pairs, except those whose destination is the source itself or, + for a ``null`` version, the key of the source. + + Raises: + ValueError: If two sources have the same destination, or a + destination is another source, including one left in place, + except for a directory with no object at its key, which is not + copied. Also if ``missing`` is needed and None. + TypeError: If ``missing`` is a string instead of a collection. + """ + if isinstance(missing, str): + raise TypeError("missing is a collection of paths, not a path.") + named, sources, candidates = self._moves(pairs) + if missing is None and candidates: + raise ValueError("missing is needed to check the pairs.") + # A directory without an object at its key writes no destination. + skipped = set(candidates).intersection( + str(S3Path.parse(path).target) for path in missing or () + ) + # A path with a version always names an object, so it writes its + # destination even when the key has no current object. + writers = Counter( + dest + for _, _, versioned, source, dest in named + if source != dest and (versioned or source not in skipped) + ) + for dest in writers: + if writers[dest] > 1: + raise ValueError("Cannot move several paths to the same destination.") + if dest in sources: + raise ValueError("Cannot move a path onto another path that is moved.") + return [(p1, p2) for p1, p2, _, source, dest in named if source != dest] + + @classmethod + def delete_paths(cls, path: str | list[str]) -> tuple[list[str], list[str]]: + """Split the paths that ``rm()`` deletes into those deleted as given and those expanded. + + Args: + path: S3 path or list of paths. + + Returns: + The paths with a version ID, which are deleted as given, and the + other paths, which the filesystem expands. + + Raises: + ValueError: If a path is a bucket. + """ + paths = [path] if isinstance(path, str) else list(path) + versioned_paths, unversioned_paths = [], [] + for p in paths: + s3_path = S3Path.parse(p) + # expand_path strips the slashes of "bucket//" to the bucket. + if s3_path.is_bucket: + raise ValueError("Cannot delete the bucket.") + if s3_path.version_id: + versioned_paths.append(p) + else: + unversioned_paths.append(p) + return versioned_paths, unversioned_paths + + @staticmethod + def _is_glob(path: str) -> bool: + """Return whether a path is a glob pattern; a version ID query is not one. + + Args: + path: The path. + + Returns: + True if the path has glob characters and no version ID query. + """ + return has_magic(path) and not S3Path.has_version_id(path) + + @staticmethod + def _moves( + pairs: Sequence[tuple[str, str]], + ) -> tuple[list[tuple[str, str, bool, str, str]], set[str], list[str]]: + """Compare the paths of a move by what they name. + + Args: + pairs: The sources and destinations of the move. + + Returns: + Each pair with whether its source has a version and the targets + of its source and destination; the targets of all sources, + including those left in place; and the conflict candidates (see + :meth:`conflict_candidates`). + """ + named = [] + for p1, p2 in pairs: + source_path = S3Path.parse(p1) + named.append( + ( + p1, + p2, + bool(source_path.version_id), + str(source_path.target), + str(S3Path.parse(p2).target), + ) + ) + # The sources left in place count too; a copy onto one overwrites it. + sources = {source for _, _, _, source, _ in named} + # A source with another source below it may be a directory. + directories: set[str] = set() + for source in sources: + parent = source.rpartition("/")[0] + while parent and parent not in directories: + directories.add(parent) + parent = parent.rpartition("/")[0] + counts = Counter(dest for _, _, _, source, dest in named if source != dest) + candidates = [ + source + for _, _, versioned, source, dest in named + if source != dest + and (counts[dest] > 1 or dest in sources) + and source in directories + and not versioned + ] + return named, sources, candidates diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 9bdd00fd..6f99e53a 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -16,6 +16,7 @@ import urllib.parse import urllib.request import uuid +import weakref from base64 import b64encode from concurrent.futures import Future, ThreadPoolExecutor, wait from datetime import UTC, datetime @@ -43,6 +44,7 @@ from pyathena.filesystem.s3_executor import S3AioExecutor, S3ThreadPoolExecutor from pyathena.filesystem.s3_object import S3MultipartUpload, S3Object, S3ObjectType, S3StorageClass from pyathena.filesystem.s3_path import S3Path +from pyathena.filesystem.s3_path_pairing import S3PathPairing from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -1815,6 +1817,55 @@ def test_get_version_outside_destination(self, tmp_path, rpath): fs.get(rpath, f"{tmp_path}/d/") assert sorted(p.name for p in tmp_path.rglob("*")) == ["d"] + @pytest.mark.parametrize( + ("path2", "lookups", "expected"), + [ + # The source is checked for a directory; the destination decides + # the pairing, so it is looked up too. + ("s3://bucket/d", ["bucket/b?versionId=v1", "s3://bucket/d"], "s3://bucket/d/b"), + # A trailing slash decides it without a lookup. + ("s3://bucket/d/", ["bucket/b?versionId=v1"], "s3://bucket/d/b"), + ], + ) + def test_copy_pairs_destination_lookup(self, path2, lookups, expected): + fs = self._make_fs() + self._serve_keys(fs, {"b"}) + # Only the destination is a directory. + fs.isdir = mock.MagicMock(side_effect=lambda p: p.rstrip("/").endswith("/d")) + + pairs = fs._copy_pairs(S3PathPairing("s3://bucket/b?versionId=v1", path2)) + + assert pairs == [("bucket/b?versionId=v1", expected)] + assert [c.args[0] for c in fs.isdir.call_args_list] == lookups + + def test_move_pairs_looks_up_only_conflict_candidates(self): + # Only a source that may be a directory and whose destination + # conflicts is looked up with HeadObject. + fs = self._make_fs() + self._serve_keys(fs, {"d/x", "e/y", "f"}) + fs._head_object = mock.MagicMock(return_value=None) + + pairs = fs._move_pairs( + ["s3://bucket/d", "s3://bucket/d/x", "s3://bucket/e/y", "s3://bucket/f"], + ["s3://bucket/e", "s3://bucket/e", "s3://bucket/out", "s3://bucket/g"], + ) + + assert len(pairs) == 4 + fs._head_object.assert_called_once_with("bucket/d") + + def test_freed_by_reference_counting(self): + # The filesystem holds no reference cycle, so a filesystem that is + # not cached, such as the internal one of a cursor (GH-978), is freed + # as soon as it is unused. + fs = self._stubbed_fs() + ref = weakref.ref(fs) + gc.disable() + try: + del fs + assert ref() is None + finally: + gc.enable() + def test_mv_nothing_within_maxdepth(self): # Only directories within maxdepth: nothing is moved, as with copy(). fs = self._make_fs() diff --git a/tests/pyathena/filesystem/test_s3_path.py b/tests/pyathena/filesystem/test_s3_path.py index cf143f33..0dda470d 100644 --- a/tests/pyathena/filesystem/test_s3_path.py +++ b/tests/pyathena/filesystem/test_s3_path.py @@ -123,3 +123,17 @@ def test_with_version_id(self): assert path == S3Path("bucket", "key", "v1") with pytest.raises(AttributeError): path.key = "other" # type: ignore[misc] + + @pytest.mark.parametrize( + ("path", "expected"), + [ + # A write to the key replaces its "null" version. + (S3Path("bucket", "key", "null"), S3Path("bucket", "key")), + # Any other path is its own target. + (S3Path("bucket", "key", "v1"), S3Path("bucket", "key", "v1")), + (S3Path("bucket", "key"), S3Path("bucket", "key")), + (S3Path("bucket"), S3Path("bucket")), + ], + ) + def test_target(self, path, expected): + assert path.target == expected diff --git a/tests/pyathena/filesystem/test_s3_path_pairing.py b/tests/pyathena/filesystem/test_s3_path_pairing.py new file mode 100644 index 00000000..f57fbacf --- /dev/null +++ b/tests/pyathena/filesystem/test_s3_path_pairing.py @@ -0,0 +1,228 @@ +# 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_pairing import S3PathPairing + + +def _pairing(pairs): + # The pairing of a move of the paths as lists. + return S3PathPairing([p1 for p1, _ in pairs], [p2 for _, p2 in pairs]) + + +class TestS3PathPairing: + @pytest.mark.parametrize( + ("path1", "path2", "expected"), + [ + ("s3://bucket/a", "s3://bucket/b", True), + (["s3://bucket/a"], "s3://bucket/b", True), + ("s3://bucket/a", ["s3://bucket/b"], True), + # Lists are paired as given. + (["s3://bucket/a"], ["s3://bucket/b"], False), + ], + ) + def test_expands(self, path1, path2, expected): + assert S3PathPairing(path1, path2).expands is expected + + @pytest.mark.parametrize( + ("path1", "recursive", "maxdepth", "expected"), + [ + ("s3://bucket/d", False, None, True), + ("s3://bucket/d", True, None, False), + ("s3://bucket/d", True, 1, True), + (["s3://bucket/d"], False, None, False), + ], + ) + def test_skips_directories(self, path1, recursive, maxdepth, expected): + assert ( + S3PathPairing(path1, "s3://bucket/out", recursive, maxdepth).skips_directories + is expected + ) + + @pytest.mark.parametrize( + ("path1", "path2", "expected"), + [ + ("s3://bucket/a", "s3://bucket/b", True), + # A version ID query is not a glob pattern. + ("s3://bucket/a?versionId=v1", "s3://bucket/b", True), + # A trailing slash decides the pairing without a lookup. + ("s3://bucket/a/", "s3://bucket/b", False), + ("s3://bucket/a", "s3://bucket/b/", False), + # So does a glob pattern. + ("s3://bucket/a*", "s3://bucket/b", False), + (["s3://bucket/a"], "s3://bucket/b", False), + ("s3://bucket/a", ["s3://bucket/b"], False), + ], + ) + def test_looks_up_destination(self, path1, path2, expected): + assert S3PathPairing(path1, path2).looks_up_destination is expected + + @pytest.mark.parametrize( + ("path1", "path2", "sources", "destination_is_dir", "expected"), + [ + # The directory and the objects below it, as fsspec's copy() + # pairs them. + ( + "s3://bucket/d", + "s3://bucket/out", + ["bucket/d", "bucket/d/a", "bucket/d/b"], + False, + [ + ("bucket/d", "s3://bucket/out"), + ("bucket/d/a", "s3://bucket/out/a"), + ("bucket/d/b", "s3://bucket/out/b"), + ], + ), + # A version is copied into a directory under its key, + ( + "s3://bucket/b?versionId=v1", + "s3://bucket/d", + ["bucket/b?versionId=v1"], + True, + [("bucket/b?versionId=v1", "s3://bucket/d/b")], + ), + # or to the destination itself. + ( + "s3://bucket/b?versionId=v1", + "s3://bucket/c", + ["bucket/b?versionId=v1"], + False, + [("bucket/b?versionId=v1", "s3://bucket/c")], + ), + # A trailing slash needs no lookup. + ( + "s3://bucket/b?versionId=v1", + "s3://bucket/d/", + ["bucket/b?versionId=v1"], + None, + [("bucket/b?versionId=v1", "s3://bucket/d/b")], + ), + # A list of sources is copied into the destination by name. + ( + ["s3://bucket/b?versionId=v1"], + "s3://bucket/d", + ["bucket/b?versionId=v1"], + None, + [("bucket/b?versionId=v1", "s3://bucket/d/b")], + ), + # Lists are paired as given, up to the end of the shorter one. + ( + ["s3://bucket/a", "s3://bucket/b", "s3://bucket/c"], + ["s3://bucket/x", "s3://bucket/y"], + (), + None, + [("s3://bucket/a", "s3://bucket/x"), ("s3://bucket/b", "s3://bucket/y")], + ), + # Nothing to copy. + ("s3://bucket/d", "s3://bucket/out", [], None, []), + ], + ) + def test_copy_pairs(self, path1, path2, sources, destination_is_dir, expected): + assert S3PathPairing(path1, path2).copy_pairs(sources, destination_is_dir) == expected + + def test_copy_pairs_needs_destination_lookup(self): + with pytest.raises(ValueError, match="destination_is_dir"): + S3PathPairing("s3://bucket/a", "s3://bucket/b").copy_pairs(["bucket/a"]) + + def test_copy_pairs_needs_sources(self): + # Not passing the expansion is not the same as expanding to nothing. + with pytest.raises(ValueError, match="sources"): + S3PathPairing("s3://bucket/a/", "s3://bucket/b/").copy_pairs() + + def test_move_pairs(self): + # The "null" version of a key moved onto the key stays in place. + pairs = [("s3://bucket/b?versionId=null", "s3://bucket/b"), ("bucket/d/a", "s3://bucket/z")] + + assert _pairing(pairs).conflict_candidates(pairs) == [] + assert _pairing(pairs).move_pairs(pairs) == [("bucket/d/a", "s3://bucket/z")] + + @pytest.mark.parametrize( + ("pairs", "match"), + [ + ( + [("s3://bucket/b", "s3://bucket/o"), ("s3://bucket/z", "s3://bucket/o")], + "same destination", + ), + ( + [("s3://bucket/b", "s3://bucket/z"), ("s3://bucket/z", "s3://bucket/n")], + "another path that is moved", + ), + ], + ) + def test_move_pairs_conflicts(self, pairs, match): + with pytest.raises(ValueError, match=match): + _pairing(pairs).move_pairs(pairs) + + def test_move_pairs_directory_without_object(self): + # A source with another source below it may be a directory; one + # without an object at its key is not copied, so its destination + # does not conflict. + pairs = [ + ("s3://bucket/d", "s3://bucket/e"), + ("s3://bucket/d/x", "s3://bucket/e"), + ("s3://bucket/e/y", "s3://bucket/out"), + ] + + assert _pairing(pairs).conflict_candidates(pairs) == ["bucket/d"] + with pytest.raises(ValueError, match="missing"): + _pairing(pairs).move_pairs(pairs) + with pytest.raises(ValueError, match="same destination"): + _pairing(pairs).move_pairs(pairs, missing=set()) + assert _pairing(pairs).move_pairs(pairs, missing={"bucket/d"}) == pairs + # The missing sources can be given in any form that names them. + assert _pairing(pairs).move_pairs(pairs, missing={"s3://bucket/d"}) == pairs + + def test_move_pairs_missing_is_a_collection(self): + pairs = [("s3://bucket/d", "s3://bucket/e"), ("s3://bucket/d/x", "s3://bucket/e")] + with pytest.raises(TypeError, match="collection"): + _pairing(pairs).move_pairs(pairs, missing="bucket/d") + + def test_conflict_candidates_one_per_pair(self): + # A source given twice is looked up twice, as each pair is checked. + pairs = [ + ("s3://bucket/d", "s3://bucket/out"), + ("s3://bucket/d", "s3://bucket/out"), + ("s3://bucket/d/x", "s3://bucket/x"), + ] + assert _pairing(pairs).conflict_candidates(pairs) == ["bucket/d", "bucket/d"] + + def test_move_pairs_version_names_an_object(self): + # A version is never taken for a directory. + pairs = [ + ("s3://bucket/d?versionId=null", "s3://bucket/out"), + ("s3://bucket/d/x", "s3://bucket/x"), + ("s3://bucket/a", "s3://bucket/out"), + ] + + assert _pairing(pairs).conflict_candidates(pairs) == [] + with pytest.raises(ValueError, match="same destination"): + _pairing(pairs).move_pairs(pairs) + + def test_move_pairs_version_of_missing_directory_key(self): + # The "null" version of a key without a current object still names an + # object, so it writes its destination although the key is missing. + pairs = [ + ("src/d", "dst/out"), + ("src/d?versionId=null", "dst/out"), + ("src/a", "dst/out"), + ("src/d/x", "dst/x"), + ] + + assert _pairing(pairs).conflict_candidates(pairs) == ["src/d"] + with pytest.raises(ValueError, match="same destination"): + _pairing(pairs).move_pairs(pairs, missing={"src/d"}) + + def test_delete_paths(self): + assert S3PathPairing.delete_paths( + ["s3://bucket/d", "s3://bucket/b?versionId=v1", "s3://bucket/c"] + ) == (["s3://bucket/b?versionId=v1"], ["s3://bucket/d", "s3://bucket/c"]) + + @pytest.mark.parametrize("path", ["s3://bucket", "s3://bucket/", "s3://bucket//"]) + def test_delete_paths_bucket(self, path): + with pytest.raises(ValueError, match="Cannot delete the bucket"): + S3PathPairing.delete_paths(path)