diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index fc573173..7fc0c82d 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1879,49 +1879,6 @@ def _check_multipart_upload_size(self, path: str, size: int, block_size: int) -> f"{min_block_size} bytes." ) - @staticmethod - def _write_and_close(f: S3File, value: bytes | bytearray | memoryview) -> None: - """Write the whole value to a file opened for writing and close it. - - Unlike a ``with`` block, a failed write closes the file without - committing it, so the existing object is left unchanged. - - Args: - f: The file to write to. - value: The bytes to write. - """ - try: - if isinstance(value, memoryview) and not value.c_contiguous: - # The buffer of the file cannot write a non-contiguous memoryview. - value = value.tobytes() - f.write(value) - except BaseException: - f._close_without_commit() - raise - f.close() - - @staticmethod - def _write_file_and_close(f: S3File, local: BinaryIO, callback: Callback) -> None: - """Write the rest of a local file to a file opened for writing and close it. - - Unlike a ``with`` block, a failed read, write, or progress update - closes the file without committing it, so the existing object is - left unchanged. - - Args: - f: The file to write to. - local: The local file to read from. - callback: Progress callback, updated with the size of each block. - """ - try: - while data := local.read(f.blocksize): - f.write(data) - callback.relative_update(len(data)) - except BaseException: - f._close_without_commit() - raise - f.close() - def pipe_file( self, path: str, value: bytes | bytearray | memoryview, mode: str = "overwrite", **kwargs ) -> None: @@ -1975,9 +1932,7 @@ def pipe_file( # Defer to the buffered open() path, which keeps the # deferred-commit semantics of fsspec transactions and uploads # large data as a parallel multipart upload. - self._write_and_close( - self.open(path, "xb" if mode == "create" else "wb", **kwargs), value - ) + self.open(path, "xb" if mode == "create" else "wb", **kwargs)._write_and_close(value) return bucket, key, version_id = self.parse_path(path) if version_id: @@ -2210,17 +2165,13 @@ def put_file( # The local file is opened first, so that an unreadable one fails # before the remote file is opened. with open(lpath, "rb") as local: - self._write_file_and_close( - self.open( - rpath, - "xb" if mode == "create" else "wb", - block_size=block_size, - max_workers=max_workers, - s3_additional_kwargs=s3_additional_kwargs, - ), - local, - callback, - ) + self.open( + rpath, + "xb" if mode == "create" else "wb", + block_size=block_size, + max_workers=max_workers, + s3_additional_kwargs=s3_additional_kwargs, + )._write_file_and_close(local, callback) self.invalidate_cache(rpath) @@ -3403,6 +3354,45 @@ def _close_without_commit(self) -> None: self.multipart_upload_parts = [] self._executor.shutdown() + def _write_and_close(self, value: bytes | bytearray | memoryview) -> None: + """Write the whole value and close the file. + + Unlike a ``with`` block, a failed write closes the file without + committing it, so the existing object is left unchanged. + + Args: + value: The bytes to write. + """ + try: + if isinstance(value, memoryview) and not value.c_contiguous: + # The buffer of the file cannot write a non-contiguous memoryview. + value = value.tobytes() + self.write(value) + except BaseException: + self._close_without_commit() + raise + self.close() + + def _write_file_and_close(self, local: BinaryIO, callback: Callback) -> None: + """Write the rest of a local file and close the file. + + Unlike a ``with`` block, a failed read, write, or progress update + closes the file without committing it, so the existing object is + left unchanged. + + Args: + local: The local file to read from. + callback: Progress callback, updated with the size of each block. + """ + try: + while data := local.read(self.blocksize): + self.write(data) + callback.relative_update(len(data)) + except BaseException: + self._close_without_commit() + raise + self.close() + def _initiate_upload(self) -> None: if not self.append_block and self.tell() < self.blocksize: # Files smaller than block size in size cannot be multipart uploaded. diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index fdbd0121..85165eba 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -205,9 +205,7 @@ def _pipe_file_in_transaction( block_size = kwargs.get("block_size") or self._sync_fs.default_block_size # The size in bytes; the length of a memoryview counts its items. self._sync_fs._check_multipart_upload_size(path, memoryview(value).nbytes, block_size) - self._sync_fs._write_and_close( - self.open(path, "xb" if mode == "create" else "wb", **kwargs), value - ) + self.open(path, "xb" if mode == "create" else "wb", **kwargs)._write_and_close(value) async def _put_file( self, @@ -273,17 +271,13 @@ def _put_file_in_transaction( # See S3FileSystem.put_file. with open(lpath, "rb") as local: - self._sync_fs._write_file_and_close( - self.open( - rpath, - "xb" if mode == "create" else "wb", - block_size=block_size, - max_workers=max_workers, - s3_additional_kwargs=s3_additional_kwargs, - ), - local, - callback, - ) + self.open( + rpath, + "xb" if mode == "create" else "wb", + block_size=block_size, + max_workers=max_workers, + s3_additional_kwargs=s3_additional_kwargs, + )._write_file_and_close(local, callback) self.invalidate_cache(rpath) async def _get_file(self, rpath: str, lpath: str, callback=_DEFAULT_CALLBACK, **kwargs) -> None: diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 600372aa..910af255 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -2079,7 +2079,6 @@ def test_put_file_block_size(self, tmp_path): # API. fs = self._make_fs() fs.open = mock.MagicMock() - fs.open.return_value.blocksize = 8 lpath = tmp_path / "data" lpath.write_bytes(b"a" * 13) @@ -2107,7 +2106,6 @@ def test_put_file_content_type(self, tmp_path, filesystem_kwargs, kwargs, expect fs = self._make_fs() fs.s3_additional_kwargs = filesystem_kwargs fs.open = mock.MagicMock() - fs.open.return_value.blocksize = 8 lpath = tmp_path / "data.csv" lpath.write_bytes(b"a") diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index c8d3db98..55109dc0 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -368,7 +368,6 @@ def test_transaction_put_file_block_size(self, tmp_path): # open() instead of the S3 API, as outside one. fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) fs.open = mock.MagicMock() - fs.open.return_value.blocksize = 8 local = tmp_path / "local" local.write_bytes(b"a" * 13) @@ -517,7 +516,6 @@ def test_put_file_in_transaction_open_parameters(self, tmp_path, mode, open_mode # GH-972: fsspec's mode argument selects the mode of the file. fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) fs.open = mock.MagicMock() - fs.open.return_value.blocksize = 4 lpath = tmp_path / "data.csv" lpath.write_bytes(b"a")