diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index eb48eccee..e05abad65 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1367,7 +1367,7 @@ def _finish_multipart_upload( upload_id=upload_id, parts=parts, ) - except Exception: + except Exception as error: for future in futures: future.cancel() try: @@ -1381,6 +1381,9 @@ def _finish_multipart_upload( _logger.exception( f"Failed to abort multipart upload {upload_id} to s3://{bucket}/{key}." ) + # The upload still exists, so the caller must keep its upload ID + # for a later discard() to retry the abort. + error.multipart_abort_failed = True # type: ignore[attr-defined] raise def cat_file( @@ -2391,11 +2394,13 @@ def commit(self) -> None: upload_id=cast(str, self.multipart_upload.upload_id), futures=self.multipart_upload_parts, ) - except Exception: - # The multipart upload has been aborted by the helper; - # prevent discard() from aborting it again. - self.multipart_upload = None - self.multipart_upload_parts = [] + except Exception as error: + # Unless the abort failed, the helper has aborted the upload; + # clear it so discard() does not abort it again. If the abort + # failed, keep the upload ID so discard() can retry the abort. + if not getattr(error, "multipart_abort_failed", False): + self.multipart_upload = None + self.multipart_upload_parts = [] raise self.fs.invalidate_cache(self.path) diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 7e4f36e5b..0515bed0c 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1955,6 +1955,45 @@ def test_discard(self, multipart): assert file.multipart_upload is None assert file.multipart_upload_parts == [] + @pytest.mark.parametrize( + ("abort_fails", "keeps_upload_id"), + [(True, True), (False, False)], + ) + def test_commit_failed_completion_upload_id(self, abort_fails, keeps_upload_id): + # When CompleteMultipartUpload fails, the abort runs. If the abort + # also fails, the upload still exists, so commit() must keep the upload + # ID for discard() to retry. Otherwise the upload ID is cleared, so a + # later discard() does not abort an upload that is already gone. + file = self._make_multipart_write_file(b"x" * 8, autocommit=False) + part = Future() + part.set_result(SimpleNamespace(etag='"e1"', part_number=1)) + file.multipart_upload_parts = [part] + real_fs = S3FileSystem.__new__(S3FileSystem) + real_fs._client = mock.MagicMock() + real_fs._call = mock.MagicMock( + side_effect=RuntimeError("abort failed") if abort_fails else None + ) + real_fs._complete_multipart_upload = mock.MagicMock( + side_effect=RuntimeError("complete failed") + ) + file.fs._finish_multipart_upload.side_effect = lambda **kw: ( + S3FileSystem._finish_multipart_upload(real_fs, **kw) + ) + file.fs.invalidate_cache = mock.MagicMock() + + with pytest.raises(RuntimeError, match="complete failed"): + file.commit() + + assert (file.multipart_upload is not None) is keeps_upload_id + assert (file.multipart_upload_parts != []) is keeps_upload_id + + file.discard() + if keeps_upload_id: + assert file.fs._call.call_args.args[0] == "abort_multipart_upload" + assert file.fs._call.call_args.kwargs["UploadId"] == "uploadid" + else: + file.fs._call.assert_not_called() + @pytest.mark.parametrize("autocommit", [True, False]) def test_upload_chunk_multipart(self, autocommit): # Multipart upload (CompleteMultipartUpload), completed from the uploaded