diff --git a/src/base_cli_demo/cli.py b/src/base_cli_demo/cli.py index 8f625c1..ade8417 100644 --- a/src/base_cli_demo/cli.py +++ b/src/base_cli_demo/cli.py @@ -3,8 +3,12 @@ from __future__ import annotations import json +import os +import stat +import tempfile from collections.abc import Mapping from importlib.resources import files +from pathlib import Path from typing import Any import base_cli @@ -92,18 +96,77 @@ def _persist_reconciliation( return state_path = context.state_dir / "last-reconciliation.json" - state_path.parent.mkdir(parents=True, exist_ok=True) - state_path.write_text( - json.dumps(dict(record), sort_keys=True) + "\n", - encoding="utf-8", - ) + had_previous_snapshot = state_path.exists() + staged_path: Path | None = None + try: + serialized = _serialize_reconciliation(record) + state_path.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile( + mode="w", + encoding="utf-8", + dir=state_path.parent, + prefix=f".{state_path.name}.", + suffix=".tmp", + delete=False, + ) as staged: + staged_path = Path(staged.name) + staged.write(serialized) + staged.flush() + os.fsync(staged.fileno()) + + _preserve_state_mode(staged_path, state_path) + temporary_input = context.temp_dir / "reconciliation-input.json" + temporary_input.write_text(serialized, encoding="utf-8") + context.on_cleanup(lambda: temporary_input.unlink(missing_ok=True)) + _replace_state(staged_path, state_path) + staged_path = None + except (OSError, TypeError, ValueError) as exc: + snapshot_message = ( + "the previous snapshot was left unchanged." + if had_previous_snapshot + else "no reconciliation snapshot was published." + ) + raise click.ClickException( + f"Could not persist the reconciliation snapshot; {snapshot_message}" + ) from exc + finally: + if staged_path is not None: + try: + staged_path.unlink(missing_ok=True) + except OSError: + pass - temporary_input = context.temp_dir / "reconciliation-input.json" - temporary_input.write_text( - json.dumps(dict(record), sort_keys=True) + "\n", - encoding="utf-8", - ) - context.on_cleanup(lambda: temporary_input.unlink(missing_ok=True)) + +def _serialize_reconciliation(record: Mapping[str, Any]) -> str: + """Serialize once so both local artifacts describe the same snapshot.""" + + return json.dumps(dict(record), sort_keys=True) + "\n" + + +def _replace_state(staged_path: Path, state_path: Path) -> None: + """Atomically publish a complete snapshot from the same filesystem.""" + + os.replace(staged_path, state_path) + try: + directory_fd = os.open(state_path.parent, os.O_RDONLY) + except OSError: + if os.name != "nt": + raise + return + try: + os.fsync(directory_fd) + finally: + os.close(directory_fd) + + +def _preserve_state_mode(staged_path: Path, state_path: Path) -> None: + """Keep an existing snapshot's permissions across atomic replacement.""" + + try: + mode = stat.S_IMODE(state_path.stat().st_mode) + except FileNotFoundError: + return + os.chmod(staged_path, mode) def _service_option(function: Any) -> Any: diff --git a/tests/test_cli.py b/tests/test_cli.py index ecdba6f..42969b4 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1,11 +1,15 @@ from __future__ import annotations import json +import stat import tempfile +import threading +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Any import base_cli +import base_cli_demo.cli as cli_module from base_cli_demo.cli import command @@ -130,6 +134,99 @@ def test_reconcile_persists_state_and_cleans_temporary_input() -> None: assert list(root.rglob("reconciliation-input.json")) == [] +def test_reconcile_preserves_existing_snapshot_permissions() -> None: + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + first = invoke(["release", "reconcile"], root) + assert first.exit_code == 0, first.output + state_path = next(root.rglob("last-reconciliation.json")) + os_mode = 0o640 + state_path.chmod(os_mode) + + second = invoke(["release", "reconcile", "--version", "2.6.0"], root) + + assert second.exit_code == 0, second.output + assert stat.S_IMODE(state_path.stat().st_mode) == os_mode + + +def test_failed_atomic_replace_preserves_the_previous_snapshot( + monkeypatch: Any, +) -> None: + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + first = invoke(["release", "reconcile", "--version", "2.6.0"], root) + assert first.exit_code == 0, first.output + state_path = next(root.rglob("last-reconciliation.json")) + previous = state_path.read_bytes() + + def fail_replace(_staged_path: Path, _state_path: Path) -> None: + raise OSError("simulated replace failure") + + monkeypatch.setattr(cli_module, "_replace_state", fail_replace) + failed = invoke(["release", "reconcile", "--version", "2.7.0"], root) + + assert failed.exit_code == 1 + assert "previous snapshot was left unchanged" in failed.output + assert state_path.read_bytes() == previous + assert not list(state_path.parent.glob(f".{state_path.name}.*.tmp")) + + +def test_serialization_failure_preserves_the_previous_snapshot(monkeypatch: Any) -> None: + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + first = invoke(["release", "reconcile"], root) + assert first.exit_code == 0, first.output + state_path = next(root.rglob("last-reconciliation.json")) + previous = state_path.read_bytes() + + def fail_serialization(_record: Any) -> str: + raise TypeError("simulated serialization failure") + + monkeypatch.setattr(cli_module, "_serialize_reconciliation", fail_serialization) + failed = invoke(["release", "reconcile", "--version", "2.7.0"], root) + + assert failed.exit_code == 1 + assert state_path.read_bytes() == previous + + +def test_concurrent_public_reconciliations_leave_a_complete_snapshot() -> None: + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + initial = invoke(["release", "reconcile"], root) + assert initial.exit_code == 0, initial.output + state_path = next(root.rglob("last-reconciliation.json")) + stop_reader = threading.Event() + reader_errors: list[BaseException] = [] + + def read_while_writing() -> None: + while not stop_reader.is_set(): + try: + json.loads(state_path.read_text(encoding="utf-8")) + except BaseException as exc: # captured for assertion in the test thread + reader_errors.append(exc) + stop_reader.set() + + with ThreadPoolExecutor(max_workers=5) as pool: + reader = pool.submit(read_while_writing) + futures = [ + pool.submit( + invoke, + ["release", "reconcile", "--version", f"2.{minor}.0"], + root, + ) + for minor in range(8, 12) + ] + results = [future.result() for future in futures] + stop_reader.set() + reader.result() + + assert not reader_errors + assert all(result.exit_code == 0 for result in results) + final = json.loads(state_path.read_text(encoding="utf-8")) + assert final["action"] == "reconciled" + assert final["target_version"] in {f"2.{minor}.0" for minor in range(8, 12)} + + def test_json_error_envelope_preserves_nonzero_exit_status() -> None: with tempfile.TemporaryDirectory() as directory: result = invoke(