103 lines
4.1 KiB
Python
103 lines
4.1 KiB
Python
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}")
|