dheraingoud's picture
feat: sync upstream commits up to f17c92bc
a1bab2d
Raw
History Blame Contribute Delete
5.42 kB
"""Atomic JSON persistence for messaging session state."""
from __future__ import annotations
import contextlib
import json
import os
import tempfile
import threading
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
from loguru import logger
@dataclass(frozen=True)
class _PendingWrite:
generation: int
snapshot: dict[str, Any]
class DebouncedJsonPersistence:
"""Thread-safe debounced JSON writer with atomic replace semantics."""
def __init__(
self,
storage_path: str,
*,
snapshot: Callable[[], dict[str, Any]],
on_dirty: Callable[[bool], None],
debounce_secs: float = 0.5,
) -> None:
self.storage_path = storage_path
self._snapshot = snapshot
self._on_dirty = on_dirty
self._debounce_secs = debounce_secs
self._save_timer: threading.Timer | None = None
self._timer_lock = threading.Lock()
self._writer_lock = threading.Lock()
self._save_generation = 0
def load_json(self) -> dict[str, Any]:
if not os.path.exists(self.storage_path):
return {}
with open(self.storage_path, encoding="utf-8") as file:
data = json.load(file)
return data if isinstance(data, dict) else {}
def schedule_save(self) -> None:
self._on_dirty(True)
with self._timer_lock:
if self._save_timer is not None:
self._save_timer.cancel()
self._save_generation += 1
generation = self._save_generation
timer = threading.Timer(
self._debounce_secs,
self._save_from_timer,
args=(generation,),
)
timer.daemon = True
self._save_timer = timer
timer.start()
def flush(self) -> None:
self._on_dirty(True)
pending = self._snapshot_for_write()
if pending is None:
return
self._write_pending(pending)
def _save_from_timer(self, generation: int) -> None:
try:
pending = self._snapshot_for_write(expected_generation=generation)
if pending is None:
return
self._write_pending(pending)
except Exception as e:
self._on_dirty(True)
logger.error(
"Failed to save sessions: exc_type={}",
type(e).__name__,
)
def _write_pending(self, pending: _PendingWrite) -> None:
try:
written = self._write_if_current(pending)
except Exception:
self._on_dirty(True)
raise
if written:
self._mark_clean_if_current(pending.generation)
def _write_if_current(self, pending: _PendingWrite) -> bool:
"""Serialize writers and reject a snapshot superseded before replace."""
with self._writer_lock:
with self._timer_lock:
if pending.generation != self._save_generation:
return False
self._write_file(pending.snapshot)
return True
def _snapshot_for_write(
self, *, expected_generation: int | None = None
) -> _PendingWrite | None:
generation = self._claim_timer(expected_generation)
if generation is None:
return None
snapshot = self._snapshot()
return _PendingWrite(generation=generation, snapshot=snapshot)
def _claim_timer(self, expected_generation: int | None) -> int | None:
with self._timer_lock:
if expected_generation is not None and (
expected_generation != self._save_generation or self._save_timer is None
):
return None
if self._save_timer is not None:
self._save_timer.cancel()
self._save_timer = None
return self._save_generation
def _mark_clean_if_current(self, generation: int) -> None:
with self._timer_lock:
is_current = (
self._save_timer is None and generation == self._save_generation
)
if is_current:
self._on_dirty(False)
def write_data(self, data: dict[str, Any]) -> None:
"""Write authoritative state after invalidating older pending snapshots."""
self._on_dirty(True)
with self._timer_lock:
if self._save_timer is not None:
self._save_timer.cancel()
self._save_timer = None
self._save_generation += 1
pending = _PendingWrite(
generation=self._save_generation,
snapshot=data,
)
self._write_pending(pending)
def _write_file(self, data: dict[str, Any]) -> None:
abs_target = os.path.abspath(self.storage_path)
dir_name = os.path.dirname(abs_target) or "."
fd, tmp_path = tempfile.mkstemp(
dir=dir_name,
prefix=".sessions.",
suffix=".tmp.json",
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as file:
json.dump(data, file, indent=2)
file.flush()
os.fsync(file.fileno())
os.replace(tmp_path, abs_target)
except BaseException:
with contextlib.suppress(OSError):
os.unlink(tmp_path)
raise