diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index d52998fb..5434a297 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -16,7 +16,7 @@ from io import BytesIO from multiprocessing import cpu_count from re import Pattern -from typing import Any, cast +from typing import Any, BinaryIO, cast from urllib.parse import unquote_plus import botocore.exceptions @@ -24,7 +24,7 @@ from botocore import UNSIGNED from botocore.client import BaseClient, Config from fsspec import AbstractFileSystem -from fsspec.callbacks import _DEFAULT_CALLBACK +from fsspec.callbacks import _DEFAULT_CALLBACK, Callback from fsspec.implementations.local import trailing_sep from fsspec.spec import AbstractBufferedFile from fsspec.utils import isfilelike, other_paths, tokenize @@ -1716,6 +1716,28 @@ def _write_and_close(f: S3File, value: bytes | bytearray | memoryview) -> None: 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: @@ -1795,10 +1817,11 @@ def _finish_multipart_upload( ) -> S3CompleteMultipartUpload: """Collect the uploaded parts and complete the multipart upload. - When any part or the completion fails, the parts that have not - started are cancelled, the running ones are waited for, and the - multipart upload is aborted so that no incomplete upload or part is - left behind. The original error is then re-raised. + When any part or the completion fails, or the wait for them is + interrupted, the parts that have not started are cancelled, the + running ones are waited for, and the multipart upload is aborted so + that no incomplete upload or part is left behind. The original error + is then re-raised. Args: bucket: S3 bucket name. @@ -1824,7 +1847,7 @@ def _finish_multipart_upload( parts=parts, **self._get_operation_kwargs("complete_multipart_upload", request_kwargs), ) - except Exception: + except BaseException: # A part that is still uploading when the upload is aborted may # be stored after the abort, so wait for the parts that could not # be cancelled first. @@ -1937,7 +1960,9 @@ def put_file( Uploads a file from the local filesystem to an S3 location. Supports automatic content type detection based on file extension and provides - progress callback functionality. + progress callback functionality. An upload that fails before it is + completed, including one of a local file that cannot be read, leaves + the existing object unchanged. Args: lpath: Local file path to upload. @@ -1984,19 +2009,20 @@ def put_file( if content_type is not None: s3_additional_kwargs["ContentType"] = content_type - with ( - self.open( - rpath, - "xb" if mode == "create" else "wb", - block_size=block_size, - max_workers=max_workers, - s3_additional_kwargs=s3_additional_kwargs, - ) as remote, - open(lpath, "rb") as local, - ): - while data := local.read(remote.blocksize): - remote.write(data) - callback.relative_update(len(data)) + # 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.invalidate_cache(rpath) @@ -3160,7 +3186,8 @@ def commit(self) -> None: Creates an empty object if nothing was written, uploads the buffered data with PutObject if no multipart upload part was submitted, and otherwise completes the multipart upload, which is aborted if the - completion fails. Invalidates the cache of the path afterwards. + completion fails or is interrupted. Invalidates the cache of the path + afterwards. Raises: FileExistsError: If an object was created at the path after the @@ -3197,7 +3224,9 @@ def commit(self) -> None: ) except Exception: # The multipart upload has been aborted by the helper; - # prevent discard() from aborting it again. + # prevent discard() from aborting it again. An interrupt may + # have stopped the helper before the abort, so the upload is + # kept for discard() then. self.multipart_upload = None self.multipart_upload_parts = [] raise diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 541e660b..a7d933bc 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -264,19 +264,19 @@ def _put_file_in_transaction( if content_type is not None: s3_additional_kwargs["ContentType"] = content_type - with ( - self.open( - rpath, - "xb" if mode == "create" else "wb", - block_size=block_size, - max_workers=max_workers, - s3_additional_kwargs=s3_additional_kwargs, - ) as remote, - open(lpath, "rb") as local, - ): - while data := local.read(remote.blocksize): - remote.write(data) - callback.relative_update(len(data)) + # 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.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 0b14f799..4eed9dd0 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1558,6 +1558,67 @@ def test_pipe_file_failed_write_aborts_multipart_upload(self): ) executor.shutdown.assert_called_once() + @pytest.mark.parametrize("intrans", [False, True]) + @pytest.mark.parametrize("error", [RuntimeError, KeyboardInterrupt, PermissionError]) + def test_put_file_failed_write(self, tmp_path, intrans, error): + # GH-1014: a failure inside the write loop, or a local file that + # cannot be read, leaves the existing object unchanged. The remote + # file used to be committed when the failure left the with block, + # which replaced the object with the data written so far or 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() + lpath = tmp_path / "data" + lpath.write_bytes(b"a") + callback = Callback() + if error is PermissionError: + if os.geteuid() == 0: + pytest.skip("root can read a file without read permission.") + lpath.chmod(0) + else: + callback.relative_update = mock.MagicMock(side_effect=error("callback failed")) + + with ( + fs.transaction if intrans else contextlib.nullcontext(), + pytest.raises(error), + ): + fs.put_file(str(lpath), "s3://bucket/key", callback=callback) + + fs._put_object.assert_not_called() + fs._call.assert_not_called() + + @pytest.mark.parametrize("intrans", [False, True]) + def test_put_file_failed_write_aborts_multipart_upload(self, tmp_path, intrans): + # GH-1014: a failure after the first block was uploaded aborts the + # multipart upload instead of completing it with that block only. + fs = self._make_fs() + fs.default_cache_type = "bytes" + fs._transaction = None + 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() + callback = Callback() + callback.relative_update = mock.MagicMock(side_effect=RuntimeError("callback failed")) + lpath = tmp_path / "data" + lpath.write_bytes(b"a" * (2 * S3FileSystem.DEFAULT_BLOCK_SIZE + 1)) + + with ( + fs.transaction if intrans else contextlib.nullcontext(), + pytest.raises(RuntimeError, match="callback failed"), + ): + fs.put_file(str(lpath), "s3://bucket/key", callback=callback) + + fs._finish_multipart_upload.assert_not_called() + fs._call.assert_called_once_with( + "abort_multipart_upload", Bucket="bucket", Key="key", UploadId="uploadid" + ) + @pytest.mark.parametrize( ("size", "block_size", "min_block_size"), [ @@ -1601,7 +1662,7 @@ def test_put_file_block_size(self, tmp_path): # API. fs = self._make_fs() fs.open = mock.MagicMock() - fs.open.return_value.__enter__.return_value.blocksize = 8 + fs.open.return_value.blocksize = 8 lpath = tmp_path / "data" lpath.write_bytes(b"a" * 13) @@ -1629,7 +1690,7 @@ 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.__enter__.return_value.blocksize = 8 + fs.open.return_value.blocksize = 8 lpath = tmp_path / "data.csv" lpath.write_bytes(b"a") @@ -2138,13 +2199,16 @@ def test_finish_multipart_upload(self): ) fs._call.assert_not_called() - def test_finish_multipart_upload_aborts_on_failure(self): + # GH-1014: an interrupt while waiting for the parts used to leave the + # multipart upload behind. + @pytest.mark.parametrize("error", [RuntimeError, KeyboardInterrupt]) + def test_finish_multipart_upload_aborts_on_failure(self, error): fs = self._make_fs() fs._complete_multipart_upload = mock.MagicMock() future: Future[SimpleNamespace] = Future() - future.set_exception(RuntimeError("upload failed")) + future.set_exception(error("upload failed")) - with pytest.raises(RuntimeError, match="upload failed"): + with pytest.raises(error, match="upload failed"): fs._finish_multipart_upload( bucket="bucket", key="key", upload_id="uploadid", futures=[future] ) @@ -4415,6 +4479,22 @@ def wait_parts(futures): waited.assert_called_once_with([running]) assert pending.cancelled() + @pytest.mark.parametrize(("error", "aborts"), [(RuntimeError, 0), (KeyboardInterrupt, 1)]) + def test_commit_failure_and_discard(self, error, aborts): + # GH-1014: an error from _finish_multipart_upload() follows its + # abort, so a later discard(), as a transaction calls after a failed + # commit(), does not abort the upload again. An interrupt may have + # stopped it before the abort, so the upload is kept for discard(). + file = self._make_multipart_write_file(b"x" * 16, autocommit=False) + file._upload_chunk(final=True) + file.fs._finish_multipart_upload.side_effect = error("failed") + + with pytest.raises(error): + file.commit() + file.discard() + + assert file.fs._call.call_count == aborts + def test_discard_on_event_loop_thread(self): # GH-976: the parts that have not started are cancelled and not # waited for, so a rollback on the thread of the event loop that diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 944ab797..f65e72de 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -1,4 +1,5 @@ import asyncio +import contextlib import os import tempfile import threading @@ -287,6 +288,36 @@ def test_transaction_pipe_file_write(self): ("k2", b"aaaa") ] + @pytest.mark.parametrize("error", [RuntimeError, PermissionError]) + def test_transaction_put_file_failed_write(self, tmp_path, error): + # GH-1014: in a transaction, a failed write or a local file that + # cannot be read does not replace the object with the data written + # so far, or 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() + local = tmp_path / "local" + local.write_bytes(b"a") + failing = tmp_path / "failing" + failing.write_bytes(b"b") + if error is PermissionError: + if os.geteuid() == 0: + pytest.skip("root can read a file without read permission.") + failing.chmod(0) + + with fs.transaction: + with ( + mock.patch.object(AioS3File, "write", side_effect=RuntimeError("write failed")) + if error is RuntimeError + else contextlib.nullcontext(), + pytest.raises(error), + ): + fs.put_file(str(failing), "s3://bucket/k1") + fs.put_file(str(local), "s3://bucket/k2") + + assert [(c.kwargs["key"], c.kwargs["body"]) for c in put_object.call_args_list] == [ + ("k2", b"a") + ] + @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() @@ -313,7 +344,7 @@ 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.__enter__.return_value.blocksize = 8 + fs.open.return_value.blocksize = 8 local = tmp_path / "local" local.write_bytes(b"a" * 13) @@ -462,7 +493,7 @@ 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.__enter__.return_value.blocksize = 4 + fs.open.return_value.blocksize = 4 lpath = tmp_path / "data.csv" lpath.write_bytes(b"a")