Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
2 changes: 2 additions & 0 deletions src/dotenv/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,6 +204,8 @@ def rewrite(
try:
with source:
yield (source, dest)
dest.flush()
os.fsync(dest.fileno())
except BaseException as err:
error = err

Expand Down
88 changes: 88 additions & 0 deletions tests/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down
Loading