From 2a68d3481253291aef5089c1d1c8e56f59d9da83 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 19:58:57 +0900 Subject: [PATCH 1/2] Follow fsspec's get_file() contract for directories and file objects get_file() wrote a directory rpath as an empty local file, did not create the parent directories of lpath, rejected a file-like lpath and ignored outfile, so a recursive get() failed. A directory or bucket rpath now creates the local directory, the parent directories of a local file are created, and a file-like lpath or outfile is written to and left open. A requested version, as an argument or in the path, always names an object, so it skips the directory check. Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 37 ++++++++++---- tests/pyathena/filesystem/test_s3.py | 76 +++++++++++++++++++++++++--- 2 files changed, 95 insertions(+), 18 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index a7a60877..76d6e5e1 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -24,7 +24,7 @@ from fsspec import AbstractFileSystem from fsspec.callbacks import _DEFAULT_CALLBACK from fsspec.spec import AbstractBufferedFile -from fsspec.utils import tokenize +from fsspec.utils import isfilelike, tokenize import pyathena from pyathena.connection import Connection @@ -1814,32 +1814,47 @@ def put_file(self, lpath: str, rpath: str, callback=_DEFAULT_CALLBACK, **kwargs) self.invalidate_cache(rpath) - def get_file(self, rpath: str, lpath: str, callback=_DEFAULT_CALLBACK, outfile=None, **kwargs): + def get_file(self, rpath: str, lpath, callback=_DEFAULT_CALLBACK, outfile=None, **kwargs): """Download an S3 file to local filesystem. Downloads a file from S3 to the local filesystem with progress tracking. Reads the file in chunks to handle large files efficiently. + As with fsspec's ``AbstractFileSystem.get_file()``, a directory + ``rpath`` creates the local directory ``lpath``, and the parent + directories of a local file ``lpath`` are created as needed. + Args: rpath: S3 source path (s3://bucket/key). - lpath: Local destination file path. + lpath: Local destination path, or a file-like object to write to. callback: Progress callback for tracking download progress. - outfile: Unused parameter for fsspec compatibility. + outfile: A file-like object to write to instead of ``lpath``. **kwargs: Additional S3 parameters passed to open(). - - Note: - If lpath is a directory, the method returns without performing - any operation. """ - if os.path.isdir(lpath): + _, _, path_version_id = self.parse_path(self._strip_protocol(rpath)) + if outfile is None and isfilelike(lpath): + outfile = lpath + elif ( + outfile is None + and not (path_version_id or kwargs.get("version_id")) + and self.isdir(rpath) + ): + # A requested version always names an object, while isdir() + # would look up the latest version, or the prefix of the same + # name when the version does not exist. + os.makedirs(lpath, exist_ok=True) return # The remote file is opened first so that no local file is created # when open() finds no object at the path. - with self.open(rpath, "rb", **kwargs) as remote, open(lpath, "wb") as local: + with contextlib.ExitStack() as stack: + remote = stack.enter_context(self.open(rpath, "rb", **kwargs)) + if outfile is None: + os.makedirs(os.path.dirname(os.path.abspath(lpath)), exist_ok=True) + outfile = stack.enter_context(open(lpath, "wb")) callback.set_size(remote.size) while data := remote.read(remote.blocksize): - local.write(data) + outfile.write(data) callback.relative_update(len(data)) def checksum(self, path: str, **kwargs): diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index ad185930..00e237f7 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1472,14 +1472,64 @@ def test_cat_file_range_key_ending_in_slash(self, start, end): fs.cat_file("s3://bucket/dir/", start=start, end=end) assert ranges == [] - def test_get_file_directory(self, tmp_path): - fs = self._make_fs() - fs.default_cache_type = "bytes" - fs.info = mock.MagicMock(return_value=S3FileSystem._directory_object("bucket", "dir")) + @pytest.mark.parametrize("key", ["dir", None]) + def test_get_file_directory(self, tmp_path, key): + # GH-974: recursive get() passes directories, including the bucket, + # which become local directories as with fsspec's get_file(). + fs = self._make_fs() + fs.info = mock.MagicMock(return_value=S3FileSystem._directory_object("bucket", key)) + rpath = f"s3://bucket/{key}" if key else "s3://bucket" + lpath = tmp_path / "out" / "dir" + + fs.get_file(rpath, str(lpath)) + assert lpath.is_dir() + # Existing directories are kept. + fs.get_file(rpath, str(lpath)) + assert lpath.is_dir() + fs._call.assert_not_called() - with pytest.raises(FileNotFoundError): - fs.get_file("s3://bucket/dir", str(tmp_path / "dir")) - assert list(tmp_path.iterdir()) == [] + def test_get_file_creates_parent_directories(self, tmp_path): + # GH-974: the parent directories used to raise FileNotFoundError. + fs, _ = self._make_object_fs(b"data") + lpath = tmp_path / "new" / "dir" / "key" + callback = Callback() + + fs.get_file("s3://bucket/key", str(lpath), callback=callback) + assert lpath.read_bytes() == b"data" + assert callback.size == callback.value == 4 + + @pytest.mark.parametrize( + ("rpath", "kwargs"), + [ + ("s3://bucket/key", {"version_id": "v1"}), + ("s3://bucket/key?versionId=v1", {}), + ], + ) + def test_get_file_version_id(self, tmp_path, rpath, kwargs): + # A requested version is looked up only by open(), not as a possible + # directory: isdir() would look up the latest version, or the prefix + # of the same name when the version does not exist. + fs, _ = self._make_object_fs(b"data") + lpath = tmp_path / "key" + + fs.get_file(rpath, str(lpath), **kwargs) + assert lpath.read_bytes() == b"data" + assert fs.info.call_count == 1 + + def test_get_file_file_like(self, tmp_path): + # GH-974: a file-like lpath used to raise TypeError, and outfile was + # ignored in favor of the local file lpath. + fs, _ = self._make_object_fs(b"data") + lpath = io.BytesIO() + fs.get_file("s3://bucket/key", lpath) + assert lpath.getvalue() == b"data" + assert not lpath.closed + + outfile = io.BytesIO() + fs.get_file("s3://bucket/key", str(tmp_path / "key"), outfile=outfile) + assert outfile.getvalue() == b"data" + assert not outfile.closed + assert not (tmp_path / "key").exists() def test_cat_ranges_range(self): fs, ranges = self._make_object_fs(b"0123456789") @@ -2799,6 +2849,18 @@ def test_move(self, fs): # assert fs.cat(f"{dir2}test_{i}") == bytes(i) # assert not fs.exists(f"{dir1}test_{i}") + def test_get_recursive(self, fs, tmp_path): + # GH-974: the directory entries used to be written as empty files. + base = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_get_recursive/{uuid.uuid4()}" + ) + fs.pipe(f"{base}/a", b"a") + fs.pipe(f"{base}/sub/b", b"b") + fs.get(base, str(tmp_path / "out"), recursive=True) + assert (tmp_path / "out" / "a").read_bytes() == b"a" + assert (tmp_path / "out" / "sub" / "b").read_bytes() == b"b" + def test_get_file(self, fs): with tempfile.TemporaryDirectory() as tmp: rpath = f"s3://{ENV.s3_staging_bucket}/{ENV.s3_filesystem_test_file_key}" From a87a59e2cd784c0ac4f1493da56781ffebca155d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 20:12:12 +0900 Subject: [PATCH 2/2] Create the parent of lpath as open() resolves it, and allow outfile alone os.path.abspath() resolves ".." before symlinks, so a path such as "link/../out/key" created the parent next to the symlink instead of the one that open() writes to. Use the parent of lpath as given. lpath now defaults to None as in fsspec, so get_file() accepts outfile alone. Test that a missing object leaves nothing behind, and give the directory test the cache type that open() needs. Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/s3.py | 7 +++++-- tests/pyathena/filesystem/test_s3.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 33 insertions(+), 2 deletions(-) diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 76d6e5e1..3f17c4f9 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1814,7 +1814,7 @@ def put_file(self, lpath: str, rpath: str, callback=_DEFAULT_CALLBACK, **kwargs) self.invalidate_cache(rpath) - def get_file(self, rpath: str, lpath, callback=_DEFAULT_CALLBACK, outfile=None, **kwargs): + def get_file(self, rpath: str, lpath=None, callback=_DEFAULT_CALLBACK, outfile=None, **kwargs): """Download an S3 file to local filesystem. Downloads a file from S3 to the local filesystem with progress tracking. @@ -1827,6 +1827,7 @@ def get_file(self, rpath: str, lpath, callback=_DEFAULT_CALLBACK, outfile=None, Args: rpath: S3 source path (s3://bucket/key). lpath: Local destination path, or a file-like object to write to. + Not needed when ``outfile`` is given. callback: Progress callback for tracking download progress. outfile: A file-like object to write to instead of ``lpath``. **kwargs: Additional S3 parameters passed to open(). @@ -1850,7 +1851,9 @@ def get_file(self, rpath: str, lpath, callback=_DEFAULT_CALLBACK, outfile=None, with contextlib.ExitStack() as stack: remote = stack.enter_context(self.open(rpath, "rb", **kwargs)) if outfile is None: - os.makedirs(os.path.dirname(os.path.abspath(lpath)), exist_ok=True) + # Not abspath(), which would resolve ".." before symlinks. + if parent := os.path.dirname(lpath): + os.makedirs(parent, exist_ok=True) outfile = stack.enter_context(open(lpath, "wb")) callback.set_size(remote.size) while data := remote.read(remote.blocksize): diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 00e237f7..f08d1b02 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1477,6 +1477,7 @@ def test_get_file_directory(self, tmp_path, key): # GH-974: recursive get() passes directories, including the bucket, # which become local directories as with fsspec's get_file(). fs = self._make_fs() + fs.default_cache_type = "bytes" fs.info = mock.MagicMock(return_value=S3FileSystem._directory_object("bucket", key)) rpath = f"s3://bucket/{key}" if key else "s3://bucket" lpath = tmp_path / "out" / "dir" @@ -1488,6 +1489,16 @@ def test_get_file_directory(self, tmp_path, key): assert lpath.is_dir() fs._call.assert_not_called() + def test_get_file_missing(self, tmp_path): + # A missing object leaves no local file or parent directory. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs.info = mock.MagicMock(side_effect=FileNotFoundError("bucket/key")) + + with pytest.raises(FileNotFoundError): + fs.get_file("s3://bucket/key", str(tmp_path / "new" / "key")) + assert list(tmp_path.iterdir()) == [] + def test_get_file_creates_parent_directories(self, tmp_path): # GH-974: the parent directories used to raise FileNotFoundError. fs, _ = self._make_object_fs(b"data") @@ -1498,6 +1509,18 @@ def test_get_file_creates_parent_directories(self, tmp_path): assert lpath.read_bytes() == b"data" assert callback.size == callback.value == 4 + def test_get_file_parent_through_symlink(self, tmp_path): + # The parent is created as open() resolves it: "link/.." is the + # parent of the symlink's target, not tmp_path / "a". + fs, _ = self._make_object_fs(b"data") + (tmp_path / "b" / "sub").mkdir(parents=True) + (tmp_path / "a").mkdir() + (tmp_path / "a" / "link").symlink_to(tmp_path / "b" / "sub") + (tmp_path / "a" / "out").touch() + + fs.get_file("s3://bucket/key", str(tmp_path / "a" / "link" / ".." / "out" / "key")) + assert (tmp_path / "b" / "out" / "key").read_bytes() == b"data" + @pytest.mark.parametrize( ("rpath", "kwargs"), [ @@ -1531,6 +1554,11 @@ def test_get_file_file_like(self, tmp_path): assert not outfile.closed assert not (tmp_path / "key").exists() + # As with fsspec, lpath may be omitted when outfile is given. + outfile = io.BytesIO() + fs.get_file("s3://bucket/key", outfile=outfile) + assert outfile.getvalue() == b"data" + def test_cat_ranges_range(self): fs, ranges = self._make_object_fs(b"0123456789")