from __future__ import annotations import logging from collections.abc import Mapping from dataclasses import dataclass, field from datetime import datetime, timedelta from typing import Any from blackout_notifier.clients import ExternalServiceError, Notifier, SaapaClient from blackout_notifier.messages import ( cancelled_outage_message, new_outage_message, reactivated_outage_message, updated_outage_message, ) from blackout_notifier.models import Outage, OutageDataError, jalali_query_window from blackout_notifier.state import StateRepository LOGGER = logging.getLogger(__name__) @dataclass(slots=True) class RunResult: sent: int = 0 state_changed: bool = False errors: list[str] = field(default_factory=list) class OutageService: def __init__( self, *, bill_chat_ids: Mapping[str, tuple[str, ...]], lookahead_days: int, saapa: SaapaClient, notifier: Notifier, state_repository: StateRepository, ) -> None: self._bill_chat_ids = dict(bill_chat_ids) self._lookahead_days = lookahead_days self._saapa = saapa self._notifier = notifier self._state_repository = state_repository def run(self, now: datetime, *, persist: bool = True) -> RunResult: state = self._state_repository.load() result = RunResult() from_date, to_date = jalali_query_window(now, self._lookahead_days) for bill_id, chat_ids in self._bill_chat_ids.items(): try: outages = self._saapa.fetch(bill_id, from_date, to_date) except ExternalServiceError as error: result.errors.append(str(error)) LOGGER.error("%s", error) continue self._process_bill(state, bill_id, chat_ids, outages, now, result) if self._prune_expired(state, now): result.state_changed = True if result.state_changed and persist: state["updated_at"] = now.isoformat() self._state_repository.save(state) return result def _process_bill( self, state: dict[str, Any], bill_id: str, chat_ids: tuple[str, ...], outages: list[Outage], now: datetime, result: RunResult, ) -> None: bills = state["bills"] bill_state = bills.get(bill_id) if bill_state is None: bill_state = {"events": {}} records: dict[str, dict[str, Any]] = bill_state["events"] current = { outage.outage_number: outage for outage in outages if outage.stop_datetime() > now } for event_id, outage in current.items(): record = records.get(event_id) if record is None: if self._notify(chat_ids, new_outage_message(bill_id, outage, now), result): records[event_id] = _record(outage, now) result.state_changed = True continue previous = Outage.from_dict(record["snapshot"]) if record["status"] == "cancelled": if self._notify(chat_ids, reactivated_outage_message(bill_id, outage, now), result): records[event_id] = _record(outage, now) result.state_changed = True continue if outage.fingerprint() != record["notified_hash"]: if self._notify( chat_ids, updated_outage_message(bill_id, previous, outage, now), result, ): records[event_id] = _record(outage, now) result.state_changed = True continue if record["missing_checks"]: record["missing_checks"] = 0 record["last_seen_at"] = now.isoformat() result.state_changed = True for event_id, record in list(records.items()): if event_id in current or record["status"] != "active": continue try: previous = Outage.from_dict(record["snapshot"]) except OutageDataError as error: result.errors.append(f"Invalid stored outage for bill {bill_id}: {error}") continue if previous.start_datetime() <= now: continue if record["missing_checks"] == 0: record["missing_checks"] = 1 result.state_changed = True elif self._notify(chat_ids, cancelled_outage_message(bill_id, previous, now), result): record["status"] = "cancelled" record["missing_checks"] = 2 record["notified_at"] = now.isoformat() result.state_changed = True if records and bill_id not in bills: bills[bill_id] = bill_state def _notify(self, chat_ids: tuple[str, ...], message: str, result: RunResult) -> bool: delivered_to_all = True for chat_id in chat_ids: try: self._notifier.send(chat_id, message) except ExternalServiceError as error: result.errors.append(str(error)) LOGGER.error("%s", error) delivered_to_all = False continue result.sent += 1 return delivered_to_all @staticmethod def _prune_expired(state: dict[str, Any], now: datetime) -> bool: changed = False cutoff = now - timedelta(days=30) for bill_id, bill_state in list(state["bills"].items()): records = bill_state["events"] for event_id, record in list(records.items()): outage = Outage.from_dict(record["snapshot"]) if outage.stop_datetime() < cutoff: del records[event_id] changed = True if not records: del state["bills"][bill_id] changed = True return changed def _record(outage: Outage, now: datetime) -> dict[str, Any]: return { "snapshot": outage.to_dict(), "notified_hash": outage.fingerprint(), "status": "active", "missing_checks": 0, "last_seen_at": now.isoformat(), "notified_at": now.isoformat(), }