diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index c9fa16868..a3cc388d2 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -25,7 +25,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 @@ -1839,32 +1839,50 @@ def put_file( 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=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. 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. + Not needed when ``outfile`` is given. 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: + # 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): - 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 0a90461fc..cf302f1b3 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1665,15 +1665,93 @@ 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): + @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.default_cache_type = "bytes" - fs.info = mock.MagicMock(return_value=S3FileSystem._directory_object("bucket", "dir")) + 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() + + 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/dir", str(tmp_path / "dir")) + 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") + 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 + + 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"), + [ + ("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() + + # 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") @@ -3127,6 +3205,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}"