diff --git a/CHANGELOG.md b/CHANGELOG.md index e5f25414..f4ce52d6 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 +- Flush and sync temporary file contents before `set_key` and `unset_key` replace a `.env` file. If syncing fails, preserve the original file and remove the temporary file. - 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..ed84b44d 100644 --- a/src/dotenv/main.py +++ b/src/dotenv/main.py @@ -204,6 +204,8 @@ def rewrite( try: with source: yield (source, dest) + dest.flush() + os.fsync(dest.fileno()) except BaseException as err: error = err diff --git a/tests/test_main.py b/tests/test_main.py index 930ab171..fb705815 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -126,6 +126,94 @@ def tracking_open(*args, **kwargs): assert all(handle.closed for handle in opened_handles) +@pytest.mark.parametrize( + "before,rewrite,after", + [ + (None, lambda path: dotenv.set_key(path, "a", "é"), "a='é'\n"), + ("a=x\n", lambda path: dotenv.set_key(path, "a", "é"), "a='é'\n"), + ("a=x\nb=y\n", lambda path: dotenv.unset_key(path, "a"), "b=y\n"), + ], + ids=["set_key_new_file", "set_key", "unset_key"], +) +def test_rewrite_syncs_flushed_contents_before_replace( + tmp_path, before, rewrite, after +): + dotenv_path = tmp_path / ".env" + if before is not None: + dotenv_path.write_text(before, encoding="utf-8") + real_fsync = os.fsync + real_replace = os.replace + calls = [] + + def fsync(fd): + [temp_path] = tmp_path.glob(".tmp_*") + assert os.path.samestat(os.fstat(fd), temp_path.stat()) + # A separate reader sees the complete contents only after the flush. + assert temp_path.read_text(encoding="utf-8") == after + calls.append("fsync") + real_fsync(fd) + + def replace(src, dst): + assert calls == ["fsync"] + calls.append("replace") + real_replace(src, dst) + + with mock.patch("dotenv.main.os.fsync", side_effect=fsync): + with mock.patch("dotenv.main.os.replace", side_effect=replace): + rewrite(dotenv_path) + + assert calls == ["fsync", "replace"] + assert dotenv_path.read_text(encoding="utf-8") == after + assert list(tmp_path.iterdir()) == [dotenv_path] + + +@pytest.mark.parametrize( + "before,rewrite", + [ + (None, lambda path: dotenv.set_key(path, "a", "y")), + ("a=x\n", lambda path: dotenv.set_key(path, "a", "y")), + ("a=x\n", lambda path: dotenv.unset_key(path, "a")), + ], + ids=["set_key_new_file", "set_key", "unset_key"], +) +def test_rewrite_fsync_failure_preserves_target(tmp_path, before, rewrite): + dotenv_path = tmp_path / ".env" + if before is not None: + dotenv_path.write_text(before) + sync_error = OSError("fsync failed") + + with mock.patch("dotenv.main.os.fsync", side_effect=sync_error): + with mock.patch("dotenv.main.os.replace") as replace: + with pytest.raises(OSError) as exc_info: + rewrite(dotenv_path) + + assert exc_info.value is sync_error + 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_*")) == [] + + +def test_rewrite_does_not_sync_failed_writes(dotenv_path): + dotenv_path.write_text("a=x\n") + write_error = OSError("write failed") + + with mock.patch("dotenv.main.os.fsync") as fsync: + with mock.patch("dotenv.main.os.replace") as replace: + with pytest.raises(OSError) as exc_info: + with dotenv.main.rewrite(dotenv_path, encoding="utf-8") as (_, dest): + dest.write("a=y\n") + raise write_error + + assert exc_info.value is write_error + fsync.assert_not_called() + replace.assert_not_called() + assert dotenv_path.read_text() == "a=x\n" + assert list(dotenv_path.parent.glob(".tmp_*")) == [] + + @pytest.mark.skipif( sys.platform == "win32", reason="symlinks require elevated privileges on Windows" )