diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index ba7abb984..9211d6fde 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -717,6 +717,24 @@ def _find( withdirs: bool | None = None, **kwargs, ) -> list[S3Object]: + """List the objects below a path, as described in ``find``. + + Args: + path: S3 path to search under. + maxdepth: Maximum number of levels to descend, at least 1 + (None for unlimited). + withdirs: Whether to include directories in the result. + **kwargs: Additional arguments including ``prefix`` and + ``refresh``, as described in ``find``. + + Returns: + The objects found, and the directories if ``withdirs`` is True. + + Raises: + ValueError: If ``maxdepth`` is less than 1 or the path is the root. + """ + if maxdepth is not None and maxdepth < 1: + raise ValueError("maxdepth must be at least 1") path = self._strip_protocol(path) if path in ["", "/"]: raise ValueError("Cannot traverse all files in S3.") @@ -742,7 +760,7 @@ def _find( result.append(item) # Recursively explore subdirectory if depth allows - if maxdepth > 0: + if maxdepth > 1: sub_path = f"s3://{bucket}/{item.key}" sub_results = self._find( sub_path, maxdepth=maxdepth - 1, withdirs=withdirs, **kwargs @@ -786,7 +804,9 @@ def find( Args: path: S3 path to search under (e.g., "s3://bucket/prefix"). - maxdepth: Maximum depth to recurse (None for unlimited). + maxdepth: Maximum number of levels to descend, at least 1 + (None for unlimited). With 1, only the entries directly under + the path are listed. withdirs: Whether to include directories in results (None = default behavior). detail: If True, return dict of {path: S3Object}; if False, return list of paths. **kwargs: Additional arguments including: @@ -799,6 +819,9 @@ def find( Dictionary mapping paths to S3Objects (if detail=True) or list of paths (if detail=False). + Raises: + ValueError: If ``maxdepth`` is less than 1 or the path is the root. + Example: >>> fs = S3FileSystem() >>> fs.find("s3://bucket/data/", maxdepth=2) # Limit depth diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2d0fa6cc0..e0ea0b342 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -279,7 +279,41 @@ def test_find_refresh_bypasses_cached_listings(self): assert fs.find("s3://bucket/dir", refresh=True) == ["bucket/dir/sub/new"] # The subdirectory listings of maxdepth are refreshed as well. - assert fs.find("s3://bucket/dir", maxdepth=1, refresh=True) == ["bucket/dir/sub/new"] + assert fs.find("s3://bucket/dir", maxdepth=2, refresh=True) == ["bucket/dir/sub/new"] + + def test_find_maxdepth_counts_levels_like_fsspec(self): + fs = self._make_fs() + responses = { + "dir/": { + "Contents": [{"Key": "dir/direct"}], + "CommonPrefixes": [{"Prefix": "dir/sub/"}], + }, + "dir/sub/": { + "Contents": [{"Key": "dir/sub/nested"}], + "CommonPrefixes": [{"Prefix": "dir/sub/deep/"}], + }, + "dir/sub/deep/": {"Contents": [{"Key": "dir/sub/deep/file"}]}, + } + fs._call.side_effect = lambda method, **kwargs: responses[kwargs["Prefix"]] + + with pytest.raises(ValueError, match="maxdepth must be at least 1"): + fs.find("s3://bucket/dir", maxdepth=0) + fs._call.assert_not_called() + + assert fs.find("s3://bucket/dir", maxdepth=1) == ["bucket/dir/direct"] + assert sorted(fs.find("s3://bucket/dir", maxdepth=1, withdirs=True)) == [ + "bucket/dir/direct", + "bucket/dir/sub", + ] + assert sorted(fs.find("s3://bucket/dir", maxdepth=2)) == [ + "bucket/dir/direct", + "bucket/dir/sub/nested", + ] + assert sorted(fs.find("s3://bucket/dir", maxdepth=3)) == [ + "bucket/dir/direct", + "bucket/dir/sub/deep/file", + "bucket/dir/sub/nested", + ] def test_refresh_evicts_cached_object_and_bucket_not_found(self): fs = self._make_fs() @@ -1089,19 +1123,23 @@ def test_find_maxdepth(self, fs): fs.touch(f"{dir_}/level1/level2/file2.txt") fs.touch(f"{dir_}/level1/level2/level3/file3.txt") - # Test maxdepth=0 (only files in the root) - result = fs.find(dir_, maxdepth=0) + # maxdepth must be at least 1, as in fsspec + with pytest.raises(ValueError, match="maxdepth must be at least 1"): + fs.find(dir_, maxdepth=0) + + # Test maxdepth=1 (only files in the root) + result = fs.find(dir_, maxdepth=1) assert len(result) == 1 assert fs._strip_protocol(f"{dir_}/file0.txt") in result - # Test maxdepth=1 (files in root and level1) - result = fs.find(dir_, maxdepth=1) + # Test maxdepth=2 (files in root and level1) + result = fs.find(dir_, maxdepth=2) assert len(result) == 2 assert fs._strip_protocol(f"{dir_}/file0.txt") in result assert fs._strip_protocol(f"{dir_}/level1/file1.txt") in result - # Test maxdepth=2 (files in root, level1, and level2) - result = fs.find(dir_, maxdepth=2) + # Test maxdepth=3 (files in root, level1, and level2) + result = fs.find(dir_, maxdepth=3) assert len(result) == 3 assert fs._strip_protocol(f"{dir_}/level1/level2/file2.txt") in result diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index fc0c354f5..cbc383f1f 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -442,19 +442,23 @@ async def test_find_maxdepth(self, fs): await fs._touch(f"{dir_}/level1/level2/file2.txt") await fs._touch(f"{dir_}/level1/level2/level3/file3.txt") - # Test maxdepth=0 (only files in the root) - result = await fs._find(dir_, maxdepth=0) + # maxdepth must be at least 1, as in fsspec + with pytest.raises(ValueError, match="maxdepth must be at least 1"): + await fs._find(dir_, maxdepth=0) + + # Test maxdepth=1 (only files in the root) + result = await fs._find(dir_, maxdepth=1) assert len(result) == 1 assert fs._strip_protocol(f"{dir_}/file0.txt") in result - # Test maxdepth=1 (files in root and level1) - result = await fs._find(dir_, maxdepth=1) + # Test maxdepth=2 (files in root and level1) + result = await fs._find(dir_, maxdepth=2) assert len(result) == 2 assert fs._strip_protocol(f"{dir_}/file0.txt") in result assert fs._strip_protocol(f"{dir_}/level1/file1.txt") in result - # Test maxdepth=2 (files in root, level1, and level2) - result = await fs._find(dir_, maxdepth=2) + # Test maxdepth=3 (files in root, level1, and level2) + result = await fs._find(dir_, maxdepth=3) assert len(result) == 3 assert fs._strip_protocol(f"{dir_}/level1/level2/file2.txt") in result