117 lines
3.9 KiB
Python
117 lines
3.9 KiB
Python
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)
|