128 lines
4.0 KiB
Python
128 lines
4.0 KiB
Python
|
|
"""
|
||
|
|
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]
|