from __future__ import annotations import json import os import tempfile from pathlib import Path from typing import Any from blackout_notifier.models import Outage, OutageDataError SCHEMA_VERSION = 1 class StateError(RuntimeError): """Raised when persisted state is unreadable or unsafe to overwrite.""" def empty_state() -> dict[str, Any]: return {"schema_version": SCHEMA_VERSION, "updated_at": None, "bills": {}} class StateRepository: def __init__(self, path: Path) -> None: self.path = path def load(self) -> dict[str, Any]: if not self.path.exists(): return empty_state() try: payload = json.loads(self.path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as error: raise StateError(f"Cannot read state file {self.path}: {error}") from error self._validate(payload) return payload def save(self, state: dict[str, Any]) -> None: self._validate(state) self.path.parent.mkdir(parents=True, exist_ok=True) temporary_path: Path | None = None try: with tempfile.NamedTemporaryFile( "w", encoding="utf-8", dir=self.path.parent, prefix=f".{self.path.name}.", suffix=".tmp", delete=False, ) as temporary_file: temporary_path = Path(temporary_file.name) json.dump(state, temporary_file, ensure_ascii=False, indent=2, sort_keys=True) temporary_file.write("\n") temporary_file.flush() os.fsync(temporary_file.fileno()) os.replace(temporary_path, self.path) except OSError as error: raise StateError(f"Cannot save state file {self.path}: {error}") from error finally: if temporary_path and temporary_path.exists(): temporary_path.unlink() @staticmethod def _validate(state: Any) -> None: if not isinstance(state, dict) or state.get("schema_version") != SCHEMA_VERSION: raise StateError(f"State must use schema_version {SCHEMA_VERSION}") if state.get("updated_at") is not None and not isinstance(state.get("updated_at"), str): raise StateError("State updated_at must be a string or null") bills = state.get("bills") if not isinstance(bills, dict): raise StateError("State bills must be an object") for bill_id, bill_state in bills.items(): if not isinstance(bill_id, str) or not isinstance(bill_state, dict): raise StateError("Invalid bill state") events = bill_state.get("events") if not isinstance(events, dict): raise StateError(f"State events for bill {bill_id} must be an object") for event_id, record in events.items(): _validate_record(bill_id, event_id, record) def _validate_record(bill_id: str, event_id: Any, record: Any) -> None: if not isinstance(event_id, str) or not isinstance(record, dict): raise StateError(f"Invalid event record for bill {bill_id}") try: outage = Outage.from_dict(record.get("snapshot")) except OutageDataError as error: raise StateError( f"Invalid snapshot for bill {bill_id}, event {event_id}: {error}" ) from error if outage.outage_number != event_id: raise StateError(f"Event key does not match outage_number for bill {bill_id}") if record.get("status") not in {"active", "cancelled"}: raise StateError(f"Invalid event status for bill {bill_id}, event {event_id}") missing_checks = record.get("missing_checks") if ( not isinstance(missing_checks, int) or isinstance(missing_checks, bool) or missing_checks < 0 ): raise StateError(f"Invalid missing_checks for bill {bill_id}, event {event_id}") for name in ("notified_hash", "last_seen_at", "notified_at"): if not isinstance(record.get(name), str): raise StateError(f"Invalid {name} for bill {bill_id}, event {event_id}")