"""存储抽象与内存实现。 依据 docs/01-design/database-design.md 中 fault_log / root_cause / fault_type / fault_filter_rule / event 表结构,提供统一的进程内存储接口。 生产环境通过 integrations.py 对接 MySQL / Elasticsearch / Redis; 本模块保证零第三方依赖,测试与离线运行使用 MemoryStorage。 """ from __future__ import annotations import threading from datetime import datetime from typing import Dict, List, Optional from .models import Event, FaultFilterRule, FaultLog, FaultType, RootCause class Storage: """存储接口(抽象基类)。 设计上把五类实体(故障日志 / 根因 / 故障类型 / 过滤规则 / 事件)收敛到 一个 Storage 门面,便于 pipeline / workers / api 复用同一份存储。 """ # ---- 故障日志 ------------------------------------------------------ def save_fault_log(self, log: FaultLog) -> None: raise NotImplementedError def get_fault_log(self, fault_log_id: str) -> Optional[FaultLog]: raise NotImplementedError def list_fault_logs(self) -> List[FaultLog]: raise NotImplementedError def query_fault_logs( self, host_id: Optional[str] = None, fault_type: Optional[str] = None, level: Optional[str] = None, keyword: Optional[str] = None, page: int = 1, page_size: int = 20, ) -> Dict: raise NotImplementedError # ---- 根因 ---------------------------------------------------------- def save_root_cause(self, root_cause: RootCause) -> None: raise NotImplementedError def get_root_cause(self, fault_log_id: str) -> Optional[RootCause]: raise NotImplementedError # ---- 故障类型 ------------------------------------------------------ def list_fault_types(self) -> List[FaultType]: raise NotImplementedError def get_fault_type(self, fault_type: str) -> Optional[FaultType]: raise NotImplementedError def add_fault_type(self, fault_type: FaultType) -> FaultType: raise NotImplementedError # ---- 过滤规则 ------------------------------------------------------ def list_filter_rules(self) -> List[FaultFilterRule]: raise NotImplementedError def add_filter_rule(self, rule: FaultFilterRule) -> FaultFilterRule: raise NotImplementedError # ---- 事件(根因分析关联输入) -------------------------------------- def save_event(self, event: Event) -> None: raise NotImplementedError def list_events_by_host( self, host_id: str, start: datetime, end: datetime ) -> List[Event]: raise NotImplementedError class MemoryStorage(Storage): """进程内内存实现,用于单元测试与离线运行。""" def __init__(self) -> None: self._lock = threading.RLock() self._fault_logs: Dict[str, FaultLog] = {} self._root_causes: Dict[str, RootCause] = {} self._fault_types: Dict[str, FaultType] = {} self._filter_rules: Dict[str, FaultFilterRule] = {} self._events: List[Event] = [] # ---- 故障日志 ------------------------------------------------------ def save_fault_log(self, log: FaultLog) -> None: with self._lock: self._fault_logs[log.fault_log_id] = log def get_fault_log(self, fault_log_id: str) -> Optional[FaultLog]: with self._lock: return self._fault_logs.get(fault_log_id) def list_fault_logs(self) -> List[FaultLog]: with self._lock: return sorted( self._fault_logs.values(), key=lambda l: (l.occurred_at, l.fault_log_id), ) def query_fault_logs( self, host_id: Optional[str] = None, fault_type: Optional[str] = None, level: Optional[str] = None, keyword: Optional[str] = None, page: int = 1, page_size: int = 20, ) -> Dict: with self._lock: items = list(self._fault_logs.values()) if host_id: items = [l for l in items if l.host_id == host_id] if fault_type: items = [l for l in items if l.fault_type == fault_type] if level: items = [l for l in items if l.level.upper() == level.upper()] if keyword: needle = keyword.lower() items = [l for l in items if needle in (l.message or "").lower()] items.sort(key=lambda l: (l.occurred_at, l.fault_log_id), reverse=True) total = len(items) page = max(1, page) page_size = max(1, min(page_size, 200)) start = (page - 1) * page_size return {"total": total, "items": items[start : start + page_size]} # ---- 根因 ---------------------------------------------------------- def save_root_cause(self, root_cause: RootCause) -> None: with self._lock: self._root_causes[root_cause.fault_log_id] = root_cause def get_root_cause(self, fault_log_id: str) -> Optional[RootCause]: with self._lock: return self._root_causes.get(fault_log_id) # ---- 故障类型 ------------------------------------------------------ def list_fault_types(self) -> List[FaultType]: with self._lock: return list(self._fault_types.values()) def get_fault_type(self, fault_type: str) -> Optional[FaultType]: with self._lock: return self._fault_types.get(fault_type) def add_fault_type(self, fault_type: FaultType) -> FaultType: with self._lock: self._fault_types[fault_type.fault_type] = fault_type return fault_type # ---- 过滤规则 ------------------------------------------------------ def list_filter_rules(self) -> List[FaultFilterRule]: with self._lock: return list(self._filter_rules.values()) def add_filter_rule(self, rule: FaultFilterRule) -> FaultFilterRule: with self._lock: self._filter_rules[rule.name] = rule return rule # ---- 事件 ---------------------------------------------------------- def save_event(self, event: Event) -> None: with self._lock: self._events.append(event) def list_events_by_host( self, host_id: str, start: datetime, end: datetime ) -> List[Event]: with self._lock: return [ e for e in self._events if e.host_id == host_id and start <= e.fired_at <= end ]