from __future__ import annotations import json from dataclasses import dataclass from datetime import datetime, timezone from typing import Iterable from redis.asyncio import Redis from app.correlation import CorrelationEventRecord from app.models import NormalizedEvent, NotificationDecision @dataclass class FingerprintState: fingerprint: str count: int first_seen: str last_seen: str last_severity: str | None last_event_id: str | None last_correlation_id: str | None @dataclass class SuppressState: active: bool last_notified_at: str | None ttl_seconds: int @dataclass class OpenIncidentState: active: bool fingerprint: str severity: str | None routing_class: str | None channels: list[str] opened_at: str | None last_problem_event_id: str | None last_correlation_id: str | None @dataclass class FlapState: active: bool event_count: int phases: list[str] window_seconds: int @dataclass class TriageCacheState: active: bool verdict: str classification: str | None reason: str | None source: str | None ttl_seconds: int class RedisStateRepository: def __init__( self, client: Redis, key_prefix: str, fingerprint_ttl_seconds: int, event_ttl_seconds: int, suppress_window_seconds: int, flap_window_seconds: int, flap_threshold: int, triage_cache_ttl_seconds: int, ) -> None: self.client = client self.key_prefix = key_prefix self.fingerprint_ttl_seconds = fingerprint_ttl_seconds self.event_ttl_seconds = event_ttl_seconds self.suppress_window_seconds = suppress_window_seconds self.flap_window_seconds = flap_window_seconds self.flap_threshold = flap_threshold self.triage_cache_ttl_seconds = triage_cache_ttl_seconds def _fingerprint_key(self, fingerprint: str) -> str: return f"{self.key_prefix}:fingerprint:{fingerprint}" def _event_key(self, event_id: str) -> str: return f"{self.key_prefix}:event:{event_id}" def _suppress_key(self, fingerprint: str) -> str: return f"{self.key_prefix}:suppress:{fingerprint}" def _incident_key(self, fingerprint: str) -> str: return f"{self.key_prefix}:incident:{fingerprint}" def _flap_key(self, fingerprint: str) -> str: return f"{self.key_prefix}:flap:{fingerprint}" def _triage_key(self, fingerprint: str, severity: str | None) -> str: sev = (severity or "unknown").strip().lower().replace(" ", "_") return f"{self.key_prefix}:triage:{fingerprint}:{sev}" def _correlation_key(self, scope_key: str) -> str: safe_scope = scope_key.strip().lower().replace(" ", "_") return f"{self.key_prefix}:correlation:{safe_scope}" async def ping(self) -> bool: result = await self.client.ping() return bool(result) async def update_fingerprint_state( self, fingerprint: str, event: NormalizedEvent, ) -> FingerprintState: key = self._fingerprint_key(fingerprint) now = datetime.now(timezone.utc).isoformat() exists = await self.client.exists(key) if not exists: await self.client.hset( key, mapping={ "fingerprint": fingerprint, "count": 1, "first_seen": now, "last_seen": now, "last_severity": event.severity or "", "last_event_id": event.event_id or "", "last_correlation_id": event.correlation_id, }, ) else: await self.client.hincrby(key, "count", 1) await self.client.hset( key, mapping={ "last_seen": now, "last_severity": event.severity or "", "last_event_id": event.event_id or "", "last_correlation_id": event.correlation_id, }, ) await self.client.expire(key, self.fingerprint_ttl_seconds) raw = await self.client.hgetall(key) return FingerprintState( fingerprint=(raw.get(b"fingerprint") or b"").decode(), count=int((raw.get(b"count") or b"0").decode()), first_seen=(raw.get(b"first_seen") or b"").decode(), last_seen=(raw.get(b"last_seen") or b"").decode(), last_severity=((raw.get(b"last_severity") or b"").decode() or None), last_event_id=((raw.get(b"last_event_id") or b"").decode() or None), last_correlation_id=((raw.get(b"last_correlation_id") or b"").decode() or None), ) async def save_event_snapshot( self, event: NormalizedEvent, fingerprint: str, repeat_count: int, ) -> None: if not event.event_id: return key = self._event_key(event.event_id) await self.client.hset( key, mapping={ "event_id": event.event_id, "problem_id": event.problem_id or "", "correlation_id": event.correlation_id, "host": event.host or "", "service": event.service or "", "trigger_name": event.trigger_name or "", "severity": event.severity or "", "event_type": event.event_type or "", "fingerprint": fingerprint, "repeat_count": repeat_count, "received_at": event.received_at.isoformat(), }, ) await self.client.expire(key, self.event_ttl_seconds) async def get_suppress_state(self, fingerprint: str) -> SuppressState: key = self._suppress_key(fingerprint) value = await self.client.get(key) ttl = await self.client.ttl(key) if value is None: return SuppressState( active=False, last_notified_at=None, ttl_seconds=0, ) last_notified_at = value.decode() if isinstance(value, bytes) else str(value) return SuppressState( active=True, last_notified_at=last_notified_at, ttl_seconds=max(ttl, 0), ) async def activate_suppress_window(self, fingerprint: str) -> str: key = self._suppress_key(fingerprint) now = datetime.now(timezone.utc).isoformat() await self.client.set( key, now, ex=self.suppress_window_seconds, ) return now async def clear_suppress_window(self, fingerprint: str) -> bool: key = self._suppress_key(fingerprint) deleted = await self.client.delete(key) return bool(deleted) async def upsert_open_incident( self, fingerprint: str, event: NormalizedEvent, decision: NotificationDecision, ) -> None: key = self._incident_key(fingerprint) now = datetime.now(timezone.utc).isoformat() await self.client.hset( key, mapping={ "fingerprint": fingerprint, "severity": event.severity or "", "routing_class": decision.routing_class, "channels": ",".join(decision.channels), "opened_at": now, "last_problem_event_id": event.event_id or "", "last_correlation_id": event.correlation_id, }, ) await self.client.expire(key, self.fingerprint_ttl_seconds) async def get_open_incident(self, fingerprint: str) -> OpenIncidentState | None: key = self._incident_key(fingerprint) raw = await self.client.hgetall(key) if not raw: return None channels_raw = (raw.get(b"channels") or b"").decode() channels = [item for item in channels_raw.split(",") if item] return OpenIncidentState( active=True, fingerprint=(raw.get(b"fingerprint") or b"").decode(), severity=((raw.get(b"severity") or b"").decode() or None), routing_class=((raw.get(b"routing_class") or b"").decode() or None), channels=channels, opened_at=((raw.get(b"opened_at") or b"").decode() or None), last_problem_event_id=((raw.get(b"last_problem_event_id") or b"").decode() or None), last_correlation_id=((raw.get(b"last_correlation_id") or b"").decode() or None), ) async def clear_open_incident(self, fingerprint: str) -> bool: key = self._incident_key(fingerprint) deleted = await self.client.delete(key) return bool(deleted) async def record_phase_transition( self, fingerprint: str, event_phase: str, correlation_id: str, event_id: str | None, ) -> FlapState: key = self._flap_key(fingerprint) now_dt = datetime.now(timezone.utc) now_ts = now_dt.timestamp() member = f"{int(now_ts * 1000)}|{event_phase}|{correlation_id}|{event_id or ''}" min_score = now_ts - self.flap_window_seconds await self.client.zadd(key, {member: now_ts}) await self.client.zremrangebyscore(key, 0, min_score) await self.client.expire(key, max(self.flap_window_seconds * 2, 300)) raw_members = await self.client.zrange(key, 0, -1) phases = self._extract_phases(raw_members) event_count = len(phases) active = ( event_count >= self.flap_threshold and "problem" in phases and "recovery" in phases ) return FlapState( active=active, event_count=event_count, phases=phases, window_seconds=self.flap_window_seconds, ) async def get_triage_cache( self, fingerprint: str, severity: str | None, ) -> TriageCacheState | None: key = self._triage_key(fingerprint, severity) raw = await self.client.hgetall(key) if not raw: return None ttl = await self.client.ttl(key) return TriageCacheState( active=True, verdict=((raw.get(b"verdict") or b"").decode() or "hold"), classification=((raw.get(b"classification") or b"").decode() or None), reason=((raw.get(b"reason") or b"").decode() or None), source=((raw.get(b"source") or b"").decode() or None), ttl_seconds=max(ttl, 0), ) async def save_triage_cache( self, fingerprint: str, severity: str | None, verdict: str, classification: str | None, reason: str | None, source: str = "llm", ) -> None: key = self._triage_key(fingerprint, severity) await self.client.hset( key, mapping={ "verdict": verdict, "classification": classification or "", "reason": reason or "", "source": source, }, ) await self.client.expire(key, self.triage_cache_ttl_seconds) async def get_recent_correlation_events( self, host: str | None, window_seconds: int, scope_keys: list[str] | None = None, ) -> list[CorrelationEventRecord]: now_ts = datetime.now(timezone.utc).timestamp() min_score = now_ts - window_seconds all_scope_keys = list(scope_keys or []) if host: host_key = f"host:{host.strip().lower().replace(' ', '_')}" if host_key not in all_scope_keys: all_scope_keys.append(host_key) result: list[CorrelationEventRecord] = [] seen: set[str] = set() for scope_key in all_scope_keys: key = self._correlation_key(scope_key) await self.client.zremrangebyscore(key, 0, min_score) raw_items = await self.client.zrangebyscore( key, min=min_score, max=now_ts, withscores=True, ) for raw_member, score in raw_items: member = raw_member.decode() if isinstance(raw_member, bytes) else str(raw_member) payload = json.loads(member) record = CorrelationEventRecord( host=payload.get("host", ""), event_id=payload.get("event_id"), correlation_id=payload.get("correlation_id", ""), kind=payload.get("kind", "unknown"), severity=payload.get("severity"), routing_class=payload.get("routing_class"), fingerprint=payload.get("fingerprint"), root_candidate=bool(payload.get("root_candidate", False)), role=payload.get("role", "standalone"), group_id=payload.get("group_id"), parent_event_id=payload.get("parent_event_id"), parent_correlation_id=payload.get("parent_correlation_id"), timestamp=float(score), tags=payload.get("tags") or {}, service=payload.get("service"), scope=payload.get("scope"), component=payload.get("component"), domain=payload.get("domain"), ) key_id = record.correlation_id or record.event_id or member if key_id in seen: continue seen.add(key_id) result.append(record) result.sort(key=lambda item: item.timestamp, reverse=True) return result async def save_correlation_event( self, host: str, event_id: str | None, correlation_id: str, kind: str, severity: str | None, routing_class: str | None, fingerprint: str | None, root_candidate: bool, role: str, group_id: str | None, parent_event_id: str | None, parent_correlation_id: str | None, scope_keys: list[str] | None = None, tags: dict[str, str] | None = None, service: str | None = None, scope: str | None = None, component: str | None = None, domain: str | None = None, ) -> None: now_ts = datetime.now(timezone.utc).timestamp() payload = { "host": host, "event_id": event_id, "correlation_id": correlation_id, "kind": kind, "severity": severity, "routing_class": routing_class, "fingerprint": fingerprint, "root_candidate": root_candidate, "role": role, "group_id": group_id, "parent_event_id": parent_event_id, "parent_correlation_id": parent_correlation_id, "tags": tags or {}, "service": service, "scope": scope, "component": component, "domain": domain, } raw = json.dumps(payload, ensure_ascii=False, sort_keys=True) all_scope_keys = list(scope_keys or []) host_key = f"host:{host.strip().lower().replace(' ', '_')}" if host_key not in all_scope_keys: all_scope_keys.append(host_key) for scope_key in all_scope_keys: key = self._correlation_key(scope_key) await self.client.zadd(key, {raw: now_ts}) await self.client.expire(key, self.fingerprint_ttl_seconds) @staticmethod def _extract_phases(raw_members: Iterable[bytes | str]) -> list[str]: phases: list[str] = [] for member in raw_members: value = member.decode() if isinstance(member, bytes) else str(member) parts = value.split("|", 3) if len(parts) >= 2: phases.append(parts[1]) return phases