diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index c9fa1686..ed034861 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1557,6 +1557,27 @@ 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() + def pipe_file( self, path: str, value: bytes | bytearray | memoryview, mode: str = "overwrite", **kwargs ) -> None: @@ -1566,7 +1587,8 @@ def pipe_file( instead of the inherited ``open()`` + ``write()`` path. Larger data and writes inside an fsspec transaction go through the buffered path, which uploads the data as a parallel multipart upload and - keeps the deferred-commit semantics of transactions. + keeps the deferred-commit semantics of transactions. A write that + fails on that path leaves the existing object unchanged. Args: path: S3 path (s3://bucket/key) to write to. @@ -1595,8 +1617,9 @@ 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. - with self.open(path, "xb" if mode == "create" else "wb", **kwargs) as f: - f.write(value) + self._write_and_close( + self.open(path, "xb" if mode == "create" else "wb", **kwargs), value + ) return bucket, key, version_id = self.parse_path(path) if version_id: @@ -2801,6 +2824,28 @@ def close(self) -> None: # The executor is shut down even if the final flush fails. self._executor.shutdown() + def _close_without_commit(self) -> None: + """Close the file without uploading the written data. + + Drops the buffered data, so that neither close() nor a deferred + commit() uploads it, and aborts the multipart upload, if any. An + abort failure is logged instead of raised, so it does not mask the + error that the caller is handling. Even if the abort fails or is + interrupted, commit() does not complete the upload afterwards. The + executor is shut down here, as fsspec does not close a closed file + again when it is garbage collected. + """ + self.buffer = None + self.closed = True + try: + self.discard() + except Exception: + _logger.exception(f"Failed to abort multipart upload to s3://{self.bucket}/{self.key}.") + finally: + self.multipart_upload = None + self.multipart_upload_parts = [] + self._executor.shutdown() + 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. @@ -2891,23 +2936,7 @@ def _upload_chunk(self, final: bool = False) -> bool: for upload in uploads: if part_number >= self.fs.MULTIPART_UPLOAD_MAX_PARTS: - # Close the file without the buffered data, so that - # neither close() nor commit() uploads it, and abort the - # upload. An abort failure does not mask this error, and - # commit() does not complete the upload afterwards. The - # executor is shut down here, as fsspec does not close a - # closed file again when it is garbage collected. - self.buffer = None - self.closed = True - try: - self.discard() - except Exception: - _logger.exception( - f"Failed to abort multipart upload to s3://{self.bucket}/{self.key}." - ) - self.multipart_upload = None - self.multipart_upload_parts = [] - self._executor.shutdown() + self._close_without_commit() raise ValueError( f"Cannot upload more than {self.fs.MULTIPART_UPLOAD_MAX_PARTS} " f"parts to s3://{self.bucket}/{self.key} with a block size of " diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index c1ff4cf9..c582f6c5 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -195,8 +195,9 @@ 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) - with self.open(path, "xb" if mode == "create" else "wb", **kwargs) as f: - f.write(value) + self._sync_fs._write_and_close( + self.open(path, "xb" if mode == "create" else "wb", **kwargs), value + ) async def _put_file( self, diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 0a90461f..d9f95e3b 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1249,8 +1249,7 @@ def test_pipe_file_invalid_path_raises(self): def test_pipe_file_non_contiguous_memoryview(self): # A non-contiguous memoryview within the block size in items, 4 items - # of 8 bytes here, is uploaded with PutObject, as the buffered path - # cannot write it. + # of 8 bytes here, is uploaded with PutObject. fs = self._make_fs() fs._put_object = mock.MagicMock() value = memoryview(b"ab" * 8).cast("H")[::2] @@ -1267,6 +1266,70 @@ def test_pipe_file_small_drops_max_workers(self): fs.pipe_file("s3://bucket/key", b"data", max_workers=2) fs._put_object.assert_called_once_with(bucket="bucket", key="key", body=b"data") + def test_pipe_file_buffered_non_contiguous_memoryview(self): + # GH-997: a non-contiguous memoryview larger than the block size is + # uploaded as a multipart upload; the buffer of the file used to + # raise BufferError for it. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._upload_part = mock.MagicMock( + side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + ) + fs._finish_multipart_upload = mock.MagicMock() + size = S3FileSystem.DEFAULT_BLOCK_SIZE + 1 + + fs.pipe_file("s3://bucket/key", memoryview(b"ab" * size)[::2]) + + assert b"".join(c.kwargs["body"] for c in fs._upload_part.call_args_list) == b"a" * size + fs._finish_multipart_upload.assert_called_once() + fs._call.assert_not_called() + + @pytest.mark.parametrize("intrans", [False, True]) + def test_pipe_file_failed_write(self, intrans): + # GH-997: a write that fails on the buffered path leaves the existing + # object unchanged. The file used to be committed when the failure + # left the with block of fsspec's pipe_file(), which replaced the + # object with an empty one, also later in a transaction. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._transaction = None + fs._put_object = mock.MagicMock() + + with ( + mock.patch.object(S3File, "write", side_effect=RuntimeError("write failed")), + fs.transaction if intrans else contextlib.nullcontext(), + pytest.raises(RuntimeError, match="write failed"), + ): + fs.pipe_file("s3://bucket/key", b"a" * (S3FileSystem.DEFAULT_BLOCK_SIZE + 1)) + + fs._put_object.assert_not_called() + fs._call.assert_not_called() + + def test_pipe_file_failed_write_aborts_multipart_upload(self): + # GH-997: a write that fails after its multipart upload has started + # aborts the upload instead of completing it. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._finish_multipart_upload = mock.MagicMock() + executor = mock.MagicMock() + executor.submit.side_effect = [Future(), RuntimeError("submit failed")] + fs._create_executor = mock.MagicMock(return_value=executor) + + with pytest.raises(RuntimeError, match="submit failed"): + fs.pipe_file("s3://bucket/key", b"a" * (3 * S3FileSystem.DEFAULT_BLOCK_SIZE)) + + fs._finish_multipart_upload.assert_not_called() + fs._call.assert_called_once_with( + "abort_multipart_upload", Bucket="bucket", Key="key", UploadId="uploadid" + ) + executor.shutdown.assert_called_once() + @pytest.mark.parametrize( ("size", "block_size", "min_block_size"), [ @@ -3726,6 +3789,32 @@ def test_write_exceeding_max_parts_abort_failure(self, autocommit): fs._finish_multipart_upload.assert_not_called() fs._put_object.assert_not_called() + def test_write_exceeding_max_parts_abort_interrupted(self): + # GH-997: an interrupted abort propagates, and a deferred commit + # still does not complete the upload; the executor is shut down. + fs = self._make_append_fs(b"") + fs.MULTIPART_UPLOAD_MAX_PARTS = 3 + fs._call.side_effect = KeyboardInterrupt + + executor = mock.MagicMock(wraps=S3ThreadPoolExecutor(max_workers=1)) + f = S3File( + fs, + "s3://bucket/key.txt", + mode="wb", + block_size=4, + autocommit=False, + executor=executor, + ) + with pytest.raises(KeyboardInterrupt): + self._write_and_close(f, [b"a" * 4] * 4) + f.commit() + + assert f.closed + executor.shutdown.assert_called() + fs._call.assert_called_once() + fs._finish_multipart_upload.assert_not_called() + fs._put_object.assert_not_called() + def test_write_exceeding_max_parts_without_close(self): # The executor of the closed file is shut down, as fsspec does not # close it again when it is garbage collected. diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 50b9e185..91348b43 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -268,6 +268,25 @@ def test_put_file_mode(self, tmp_path, mode): assert "mode" not in call.kwargs assert call.kwargs.get("IfNoneMatch") == ("*" if mode == "create" else None) + def test_transaction_pipe_file_write(self): + # GH-997: in a transaction, a non-contiguous memoryview is written, + # and a failed write does not replace the object with an empty one + # when the transaction commits. + fs = AioS3FileSystem(connection=mock.MagicMock(), skip_instance_cache=True) + put_object = fs._sync_fs._put_object = mock.MagicMock() + + with fs.transaction: + with ( + mock.patch.object(AioS3File, "write", side_effect=RuntimeError("write failed")), + pytest.raises(RuntimeError, match="write failed"), + ): + fs.pipe_file("s3://bucket/k1", b"data") + fs.pipe_file("s3://bucket/k2", memoryview(b"ab" * 4)[::2]) + + assert [(c.kwargs["key"], c.kwargs["body"]) for c in put_object.call_args_list] == [ + ("k2", b"aaaa") + ] + @pytest.mark.parametrize("kwargs", [{"block_size": 4}, {}]) def test_transaction_pipe_put_file_exceeding_max_parts(self, tmp_path, kwargs): # GH-953: in a transaction, as outside one, pipe_file() and put_file()