From ff6f1d8b1b61247dc8170555eabadd833c0e971e Mon Sep 17 00:00:00 2001 From: fan wu <336798967+FanWu-ai@users.noreply.github.com> Date: Sun, 4 Oct 2026 21:22:44 -0700 Subject: [PATCH] Clean up temporary files when closing a rewrite fails --- CHANGELOG.md | 1 + src/dotenv/main.py | 35 +++++++++++++------------- tests/test_main.py | 63 ++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 81 insertions(+), 18 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e5f25414..5d1f5aeb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ### Fixed +- Remove temporary files when a buffered write fails while closing the file during `set_key` or `unset_key`. - Fix a package build deprecation warning caused by a non-string `license` value in `pyproject.toml` by [@kurtmckee] in [#648] - `set_key`, `unset_key` and the `dotenv set`/`unset` commands now name the `.env` path instead of an internal temporary file when its directory is missing or not writable, and the CLI prints a short error and exits with code 2 instead of a traceback by [@jamalkamaladdin] in [#711] - `set_key` and `unset_key` no longer leave a `.tmp_*` file behind on Windows when writing a read-only `.env` fails, and the error raised is the one from the failed write rather than from cleaning up the temporary file by [@MohammedAlkindi] in [#686] diff --git a/src/dotenv/main.py b/src/dotenv/main.py index 5faa7f0e..10226017 100644 --- a/src/dotenv/main.py +++ b/src/dotenv/main.py @@ -197,28 +197,27 @@ def rewrite( err.filename = os.fspath(path) raise - with temp_file as dest: - dest_path = pathlib.Path(dest.name) - error = None + dest_path = pathlib.Path(temp_file.name) + try: + with temp_file as dest: + error = None - try: - with source: - yield (source, dest) - except BaseException as err: - error = err + try: + with source: + yield (source, dest) + except BaseException as err: + error = err - if error is None: - try: - if original_mode is not None: - os.chmod(dest_path, original_mode) + if error is not None: + raise error from None - os.replace(dest_path, path) - except BaseException: - _discard_temp_file(dest_path) - raise - else: + if original_mode is not None: + os.chmod(dest_path, original_mode) + + os.replace(dest_path, path) + except BaseException: _discard_temp_file(dest_path) - raise error from None + raise def set_key( diff --git a/tests/test_main.py b/tests/test_main.py index 930ab171..49772553 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -238,6 +238,69 @@ def test_rewrite_reports_original_error_when_cleanup_fails(dotenv_path, caplog): ] +@pytest.mark.parametrize( + "before,rewrite", + [ + (None, lambda path: dotenv.set_key(path, "a", "new")), + ("a=old\n", lambda path: dotenv.set_key(path, "a", "new")), + ("a=old\nb=keep\n", lambda path: dotenv.unset_key(path, "a")), + ], + ids=["set_key_new_file", "set_key", "unset_key"], +) +def test_rewrite_close_failure_leaves_no_temp_file(tmp_path, before, rewrite): + dotenv_path = tmp_path / ".env" + if before is not None: + dotenv_path.write_text(before) + close_error = OSError("buffered write failed while closing") + real_temp_file = dotenv.main.tempfile.NamedTemporaryFile + opened_files = [] + + def failing_temp_file(*args, **kwargs): + temp_file = real_temp_file(*args, **kwargs) + opened_files.append(temp_file) + wrapper = mock.MagicMock(wraps=temp_file) + wrapper.name = temp_file.name + wrapper.__enter__.return_value = temp_file + + def close(*exc): + temp_file.close() + raise close_error + + wrapper.__exit__.side_effect = close + return wrapper + + with mock.patch("dotenv.main.tempfile.NamedTemporaryFile", failing_temp_file): + with mock.patch("dotenv.main.os.replace") as replace: + with pytest.raises(OSError) as exc_info: + rewrite(dotenv_path) + + assert exc_info.value is close_error + assert all(temp_file.closed for temp_file in opened_files) + replace.assert_not_called() + if before is None: + assert not dotenv_path.exists() + else: + assert dotenv_path.read_text() == before + assert list(tmp_path.glob(".tmp_*")) == [] + + +@pytest.mark.parametrize("error_type", [ValueError, KeyboardInterrupt]) +def test_rewrite_body_failure_preserves_error_and_target(dotenv_path, error_type): + dotenv_path.write_text("a=old\n") + error = error_type("rewrite interrupted") + + with pytest.raises(error_type) as exc_info: + with dotenv.main.rewrite(dotenv_path, encoding="utf-8") as (source, dest): + dest.write("a=new\n") + raise error + + assert exc_info.value is error + assert source.closed + assert dest.closed + assert dotenv_path.read_text() == "a=old\n" + assert list(dotenv_path.parent.glob(".tmp_*")) == [] + + def test_set_key_missing_directory(tmp_path): dotenv_path = tmp_path / "nx_dir" / ".env"