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

- 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]
Expand Down
35 changes: 17 additions & 18 deletions src/dotenv/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
63 changes: 63 additions & 0 deletions tests/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down