diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index ba7abb984..827a623e5 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1394,14 +1394,19 @@ def cat_file( from the end of the object. end: Byte offset to stop reading at (exclusive). A negative value counts from the end of the object. - **kwargs: Additional parameters passed to the GetObject API. + **kwargs: Additional parameters passed to the GetObject API, + except ``version_id``: the version ID to read when the path + has none. Returns: The bytes read from the object. """ - bucket, key, version_id = self.parse_path(path) + bucket, key, path_version_id = self.parse_path(path) + version_id = kwargs.pop("version_id", None) + if path_version_id: + version_id = path_version_id if start is not None or end is not None: - size = self.info(path).get("size", 0) + size = self.info(path, version_id=version_id).get("size", 0) if start is None: range_start = 0 elif start < 0: @@ -1972,7 +1977,6 @@ def _open( self, path, mode, - version_id=None, max_workers=max_workers, executor=self._create_executor(max_workers=max_workers), block_size=block_size, @@ -2158,7 +2162,8 @@ def __init__( path: S3 path (s3://bucket/key) of the file. mode: The file mode, such as ``rb``, ``wb`` or ``ab``. version_id: The version ID to read. Must match the version ID in - the path if both are given. + the path if both are given. A version cannot be given, in + either form, for writing or appending. max_workers: The number of parallel workers for range reads and part copies. executor: The executor for parallel operations. If None, a new @@ -2177,22 +2182,13 @@ def __init__( Raises: ValueError: If the path has no key, the version IDs do not match, - or the block size is too small for writing. + a version is given for writing, or the block size is too small + for writing. """ self.max_workers = max_workers self._executor: S3Executor = executor or S3ThreadPoolExecutor(max_workers=max_workers) self.s3_additional_kwargs = s3_additional_kwargs if s3_additional_kwargs else {} - super().__init__( - fs=fs, - path=path, - mode=mode, - block_size=block_size, - autocommit=autocommit, - cache_type=cache_type, - cache_options=cache_options, - size=size, - ) bucket, key, path_version_id = S3FileSystem.parse_path(path) self.bucket = bucket if not key: @@ -2209,16 +2205,19 @@ def __init__( self.version_id = path_version_id else: self.version_id = version_id - if "r" not in mode and block_size < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: - # When writing occurs, the block size should not be smaller - # than the minimum size of a part in a multipart upload. - raise ValueError(f"Block size must be >= {self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE}MB.") + if self.version_id and "r" not in mode: + raise ValueError("Cannot write to the file with the version specified.") + if self.version_id and not path_version_id: + # Carry the version in the path, as with the ?versionId= suffix, + # so that a reopened (e.g., unpickled) file reads the same version. + path = f"{path}?versionId={self.version_id}" - self.append_block = False - self._details: S3Object | dict[str, Any] + self._details: S3Object | dict[str, Any] = {} if "r" in mode: - info = self.fs.info(self.path, version_id=self.version_id) - if self.fs.version_aware and not self.version_id: + # Looked up before the base class initializer, which would + # otherwise take the size from the latest version of the object. + info = fs.info(path, version_id=self.version_id) + if fs.version_aware and not self.version_id: # Pin the version observed at open time so that reads are # consistent even if the object is overwritten. info() heads # the object when the cached entry carries no version. @@ -2226,7 +2225,26 @@ def __init__( if etag := info.get("etag"): self.s3_additional_kwargs.update({"IfMatch": etag}) self._details = info - elif "a" in mode and self.fs.exists(path): + if size is None: + size = info.get("size") + + super().__init__( + fs=fs, + path=path, + mode=mode, + block_size=block_size, + autocommit=autocommit, + cache_type=cache_type, + cache_options=cache_options, + size=size, + ) + if "r" not in mode and block_size < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: + # When writing occurs, the block size should not be smaller + # than the minimum size of a part in a multipart upload. + raise ValueError(f"Block size must be >= {self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE}MB.") + + self.append_block = False + if "a" in mode and self.fs.exists(path): info = self.fs.info(self.path, version_id=self.version_id) loc = info.get("size", 0) if loc < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: @@ -2239,8 +2257,6 @@ def __init__( self.loc = loc self.s3_additional_kwargs.update(info.to_api_repr()) self._details = info - else: - self._details = {} self.multipart_upload: S3MultipartUpload | None = None self.multipart_upload_parts: list[Future[S3MultipartUploadPart]] = [] diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index cb0592c25..2d6dd3e5a 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -366,7 +366,6 @@ def _open( self._sync_fs, path, mode, - version_id=None, max_workers=max_workers, executor=S3AioExecutor(loop=self._loop), block_size=block_size, diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 2d0fa6cc0..704ffc3e2 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -438,6 +438,73 @@ def test_open_max_workers(self): with fs.open("s3://bucket/key", "wb", max_workers=2) as f: assert f.max_workers == 2 + def test_open_version_id(self): + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs.info = mock.MagicMock(return_value=self._file_object("key")) + fs.info.return_value.size = 4 + + with fs.open("s3://bucket/key", "rb", version_id="v1") as f: + assert f.version_id == "v1" + # The size is that of the requested version, not the latest one. + assert f.size == 4 + fs.info.assert_called_once_with("bucket/key?versionId=v1", version_id="v1") + # The version is carried in the path, which fsspec reopens an + # unpickled file with. + assert f.path == "bucket/key?versionId=v1" + assert f.__reduce__()[1][1] == "bucket/key?versionId=v1" + # The argument must match the version in the path. + with pytest.raises(ValueError, match="do not match"): + fs.open("s3://bucket/key?versionId=v2", "rb", version_id="v1") + + @pytest.mark.parametrize("mode", ["wb", "ab", "xb"]) + @pytest.mark.parametrize( + ("path", "kwargs"), + [ + ("s3://bucket/key", {"version_id": "v1"}), + ("s3://bucket/key?versionId=v1", {}), + ], + ) + def test_open_version_id_for_writing(self, mode, path, kwargs): + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._call.side_effect = AssertionError("No request is expected.") + + with pytest.raises(ValueError, match="version specified"): + fs.open(path, mode, **kwargs) + + @pytest.mark.parametrize( + ("path", "expected"), + [ + ("s3://bucket/key", "v1"), + # The version in the path takes precedence, as in info(). + ("s3://bucket/key?versionId=v2", "v2"), + ], + ) + def test_cat_file_version_id(self, path, expected): + fs = self._make_fs() + fs.info = mock.MagicMock(return_value=self._file_object("key")) + fs.info.return_value.size = 10 + + fs._call.return_value = {"Body": io.BytesIO(b"data")} + assert fs.cat_file(path, version_id="v1") == b"data" + fs._call.assert_called_once_with( + fs._client.get_object, Bucket="bucket", Key="key", VersionId=expected + ) + + # A range is resolved against the size of the same version. + fs._call.reset_mock() + fs._call.return_value = {"Body": io.BytesIO(b"ta")} + assert fs.cat_file(path, start=2, end=4, version_id="v1") == b"ta" + fs.info.assert_called_once_with(path, version_id=expected) + fs._call.assert_called_once_with( + fs._client.get_object, + Bucket="bucket", + Key="key", + Range="bytes=2-3", + VersionId=expected, + ) + def test_finish_multipart_upload(self): fs = self._make_fs() fs._complete_multipart_upload = mock.MagicMock() @@ -1676,6 +1743,25 @@ def test_object_version_info(self, fs): # An unversioned bucket reports the "null" version. assert version.version_id + def test_read_version_id(self, fs): + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_read_version_id/{uuid.uuid4()}" + ) + data = b"0123456789" + fs.pipe(path, data) + # An unversioned bucket reports the "null" version, which can be + # read explicitly. + version_id = fs.object_version_info(path)[0].version_id + + assert fs.cat_file(path, version_id=version_id) == data + assert fs.cat_file(path, start=2, end=5, version_id=version_id) == data[2:5] + with fs.open(path, "rb", version_id=version_id) as f: + assert f.read() == data + # The version reaches S3, which rejects an unknown one. + with pytest.raises(OSError, match="Invalid version id"): + fs.cat_file(path, version_id="invalid") + @pytest.mark.parametrize("fs", [{"version_aware": True}], indirect=True) def test_version_aware_read(self, fs): # On an unversioned bucket, the version-aware mode is a no-op for diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index fc0c354f5..ad30d9e8f 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -15,7 +15,7 @@ from pyathena.filesystem.s3 import S3File, S3FileSystem from pyathena.filesystem.s3_async import AioS3File, AioS3FileSystem -from pyathena.filesystem.s3_object import S3ObjectType, S3StorageClass +from pyathena.filesystem.s3_object import S3Object, S3ObjectType, S3StorageClass from tests import ENV from tests.pyathena.conftest import connect @@ -901,6 +901,24 @@ def test_open_max_workers(self): assert isinstance(f, AioS3File) assert f.max_workers == 2 + def test_open_version_id(self): + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + fs._sync_fs.info = mock.MagicMock( + return_value=S3Object( + init={"Key": "key"}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="key", + ) + ) + fs._sync_fs.info.return_value.size = 4 + + with fs.open("s3://bucket/key", "rb", version_id="v1") as f: + assert isinstance(f, AioS3File) + assert f.version_id == "v1" + assert f.size == 4 + fs._sync_fs.info.assert_called_once_with("bucket/key?versionId=v1", version_id="v1") + @pytest.mark.parametrize( ("objects", "target"), [