from __future__ import annotations import logging import os from pathlib import Path from typing import Iterable from uuid import uuid4 logger = logging.getLogger(__name__) def atomic_write_bytes(path: Path, content: bytes) -> None: """Durably replace *path* without exposing a partially written file.""" atomic_write_many_bytes(((Path(path), content),)) def atomic_write_many_bytes(writes: Iterable[tuple[Path, bytes]]) -> None: """Durably replace one or more files and roll back a failed transaction.""" entries = [(Path(path), bytes(content)) for path, content in writes] if not entries: return targets = [path for path, _content in entries] if len(set(targets)) != len(targets): raise ValueError("atomic write transaction contains duplicate paths") originals: dict[Path, bytes | None] = {} temporaries: dict[Path, Path] = {} replaced: list[Path] = [] parents = sorted({target.parent for target in targets}, key=str) try: for target, _content in entries: target.parent.mkdir(parents=True, exist_ok=True) originals[target] = target.read_bytes() if target.exists() else None for target, content in entries: temporary = _write_temporary(target, content) temporaries[target] = temporary for parent in parents: _fsync_directory(parent) for target, _content in entries: os.replace(temporaries[target], target) replaced.append(target) for parent in parents: _fsync_directory(parent) except BaseException as original_error: _cleanup_temporaries(temporaries.values()) if replaced: try: _restore_originals(replaced, originals) except BaseException as recovery_error: logger.critical( "Unable to roll back failed durable write transaction", exc_info=recovery_error, ) raise RuntimeError( "durable write failed and its previous files could not be restored" ) from original_error raise def _write_temporary(target: Path, content: bytes) -> Path: temporary = target.with_name(f".{target.name}.{uuid4().hex}.tmp") try: with temporary.open("xb") as handle: handle.write(content) handle.flush() os.fsync(handle.fileno()) except BaseException: temporary.unlink(missing_ok=True) raise return temporary def _cleanup_temporaries(paths: Iterable[Path]) -> None: for path in paths: try: path.unlink(missing_ok=True) except OSError: logger.warning("Unable to clean temporary persistence file %s", path, exc_info=True) def _restore_originals( replaced: Iterable[Path], originals: dict[Path, bytes | None], ) -> None: targets = list(replaced) recovery_files: dict[Path, Path] = {} try: for target in targets: original = originals[target] if original is not None: recovery_files[target] = _write_temporary(target, original) for target in targets: recovery = recovery_files.get(target) if recovery is None: target.unlink(missing_ok=True) else: os.replace(recovery, target) for parent in sorted({target.parent for target in targets}, key=str): _fsync_directory(parent) finally: _cleanup_temporaries(recovery_files.values()) def _fsync_directory(directory: Path) -> None: """Sync a directory entry where the platform supports directory handles.""" if os.name == "nt": return flags = os.O_RDONLY if hasattr(os, "O_DIRECTORY"): flags |= os.O_DIRECTORY descriptor = os.open(directory, flags) try: os.fsync(descriptor) finally: os.close(descriptor)