""" Live Discrepancy Tracker — compare CWM predictions vs actual market fills. Enables: - Detecting CWM model drift - Alerting when predictions diverge from reality - Feeding discrepancies back for CWM improvement """ from __future__ import annotations import json import time from dataclasses import dataclass, field from typing import Any, List, Mapping, Optional, Tuple from malkhut.state import MarketWorldState from malkhut.actions import FulfilmentAction from malkhut.cwm.core import MinimalCryptoLOBCWM from malkhut.cwm.replay_verify import _compare_deep, ReplayMismatch from malkhut.storage.ch_store import MalkhutCHStore @dataclass(frozen=True, slots=True) class DiscrepancyRecord: """One discrepancy between predicted and actual state.""" ts_ns: int symbol: str field: str predicted: Any actual: Any severity: str # "info", "warning", "critical" action_kind: str policy_version: str class DiscrepancyTracker: """ Track discrepancies between CWM predictions and actual market state. Runs in shadow mode: CWM predicts next state, actual state arrives later, discrepancy is logged and analyzed. """ def __init__(self, store: Optional[MalkhutCHStore] = None) -> None: self._store = store self._discrepancies: list[DiscrepancyRecord] = [] self._total_comparisons = 0 self._total_discrepancies = 0 def record_prediction( self, predicted_state: MarketWorldState, action: FulfilmentAction, policy_version: str, ) -> None: """Record a CWM prediction for later comparison.""" # Store for comparison when actual state arrives self._last_prediction = predicted_state self._last_action = action self._last_policy_version = policy_version def compare_with_actual( self, actual_state: MarketWorldState, tolerances: Optional[Mapping[str, float]] = None, ) -> List[DiscrepancyRecord]: """ Compare last prediction with actual state. Returns list of discrepancies found. """ if not hasattr(self, '_last_prediction') or self._last_prediction is None: return [] self._total_comparisons += 1 mismatches = _compare_deep( 0, self._last_prediction, actual_state, tolerances, ) discrepancies = [] for m in mismatches: disc = DiscrepancyRecord( ts_ns=actual_state.ts_ns, symbol=actual_state.venue.symbol, field=m.field, predicted=m.expected, actual=m.actual, severity=m.severity, action_kind=self._last_action.kind.value if self._last_action else "unknown", policy_version=self._last_policy_version, ) discrepancies.append(disc) self._discrepancies.append(disc) self._total_discrepancies += 1 # Persist to CH if self._store: self._store.store_discrepancy( ts_ns=actual_state.ts_ns, exchange=actual_state.venue.exchange, symbol=actual_state.venue.symbol, predicted=str(m.expected), actual=str(m.actual), severity=m.severity, ) return discrepancies @property def discrepancy_rate(self) -> float: if self._total_comparisons == 0: return 0.0 return self._total_discrepancies / self._total_comparisons @property def total_comparisons(self) -> int: return self._total_comparisons @property def total_discrepancies(self) -> int: return self._total_discrepancies def get_recent(self, n: int = 10) -> List[DiscrepancyRecord]: return self._discrepancies[-n:] def get_by_severity(self, severity: str) -> List[DiscrepancyRecord]: return [d for d in self._discrepancies if d.severity == severity]