""" Replay verification — mandatory before trusting CWM. Three verification modes: 1. Historical replay: venue data → ReplayStep → CWM transition → compare 2. Self-play replay: persist trajectory → deterministic re-run → exact match 3. hftbacktest comparison: CWM vs known replay engine for queue/fill validation Design rule from spec: "Replay correctness before search depth. A wrong CWM plus deep search creates confident nonsense." """ from __future__ import annotations import hashlib import json import time from dataclasses import dataclass, field from typing import Any, Callable, List, Optional, Protocol, Sequence, Tuple from malkhut.state import ( AccountState, MarketWorldState, Mode, OpenOrderState, OrderBookState, PositionState, PriceLevel, Side, VenueRules, ) from malkhut.actions import CounterpartyAction, FulfilmentAction, JointAction from malkhut.cwm.core import CodeWorldModel # ============================================================================== # Data types # ============================================================================== @dataclass(frozen=True, slots=True) class ReplayStep: """One step in a replay trajectory.""" before: MarketWorldState joint_action: JointAction after_ground_truth: MarketWorldState step_index: int = 0 metadata: Mapping[str, Any] = field(default_factory=dict) @dataclass(frozen=True, slots=True) class ReplayMismatch: """One field mismatch between predicted and ground truth.""" index: int field: str expected: Any actual: Any severity: str # "critical", "warning", "info" tolerance: float = 0.0 @property def is_critical(self) -> bool: return self.severity == "critical" @dataclass(frozen=True, slots=True) class ReplayResult: """Complete result of a replay verification run.""" passed: bool mismatches: List[ReplayMismatch] steps_verified: int total_steps: int first_mismatch_index: Optional[int] duration_ns: int trajectory_hash: str @property def match_rate(self) -> float: return self.steps_verified / max(self.total_steps, 1) @property def critical_count(self) -> int: return sum(1 for m in self.mismatches if m.is_critical) @property def warning_count(self) -> int: return sum(1 for m in self.mismatches if m.severity == "warning") @dataclass(frozen=True, slots=True) class TrajectoryRecord: """One step in a persisted trajectory for self-play verification.""" step_index: int before_hash: str action_hash: str after_hash: str ts_ns: int symbol: str # ============================================================================== # Deep state comparison # ============================================================================== def _hash_state(state: MarketWorldState) -> str: """Deterministic hash of a MarketWorldState for trajectory recording.""" parts = [ str(state.ts_ns), state.venue.symbol, str(state.book.best_bid) if state.book.bids else "0", str(state.book.best_ask) if state.book.asks else "0", str(state.account.equity), str(len(state.open_orders)), ] return hashlib.sha256(":".join(parts).encode()).hexdigest()[:16] def _hash_action(action: Any) -> str: """Deterministic hash of an action.""" return hashlib.sha256(str(action).encode()).hexdigest()[:16] def _compare_deep( i: int, expected: MarketWorldState, actual: MarketWorldState, tolerances: Optional[Mapping[str, float]] = None, ) -> List[ReplayMismatch]: """ Deep comparison of two MarketWorldStates. Compares all fields with appropriate tolerances: - ts_ns: exact match - venue: exact match - book prices: float tolerance (default 1e-6) - book quantities: float tolerance - account equity: float tolerance - open orders: count + individual comparison - positions: per-symbol comparison """ tol = tolerances or {} diffs: List[ReplayMismatch] = [] def _cmp(field: str, exp_val: Any, act_val: Any, tolerance: float = 1e-9) -> None: if isinstance(exp_val, float): if abs(exp_val - act_val) > tolerance: diffs.append(ReplayMismatch(i, field, exp_val, act_val, "warning", tolerance)) elif isinstance(exp_val, int): if exp_val != act_val: diffs.append(ReplayMismatch(i, field, exp_val, act_val, "info")) elif exp_val != act_val: diffs.append(ReplayMismatch(i, field, str(exp_val), str(act_val), "info")) # Timestamp _cmp("ts_ns", expected.ts_ns, actual.ts_ns) # Venue _cmp("venue.symbol", expected.venue.symbol, actual.venue.symbol) _cmp("venue.exchange", expected.venue.exchange, actual.venue.exchange) _cmp("venue.tick_size", expected.venue.tick_size, actual.venue.tick_size) # Book if expected.book and actual.book: book_tol = tol.get("book_price", 1e-6) _cmp("book.best_bid", expected.book.best_bid, actual.book.best_bid, book_tol) _cmp("book.best_ask", expected.book.best_ask, actual.book.best_ask, book_tol) _cmp("book.bid_depth", len(expected.book.bids), len(actual.book.bids)) _cmp("book.ask_depth", len(expected.book.asks), len(actual.book.asks)) # Compare top N levels for n in range(min(5, len(expected.book.bids), len(actual.book.bids))): _cmp(f"book.bid[{n}].price", expected.book.bids[n].price, actual.book.bids[n].price, book_tol) _cmp(f"book.bid[{n}].qty", expected.book.bids[n].qty, actual.book.bids[n].qty, book_tol) for n in range(min(5, len(expected.book.asks), len(actual.book.asks))): _cmp(f"book.ask[{n}].price", expected.book.asks[n].price, actual.book.asks[n].price, book_tol) _cmp(f"book.ask[{n}].qty", expected.book.asks[n].qty, actual.book.asks[n].qty, book_tol) # Account if expected.account and actual.account: acct_tol = tol.get("account_equity", 1e-6) _cmp("account.equity", expected.account.equity, actual.account.equity, acct_tol) _cmp("account.wallet_balance", expected.account.wallet_balance, actual.account.wallet_balance, acct_tol) _cmp("account.available_balance", expected.account.available_balance, actual.account.available_balance, acct_tol) _cmp("account.total_notional", expected.account.total_notional, actual.account.total_notional, acct_tol) # Open orders _cmp("open_orders.count", len(expected.open_orders), len(actual.open_orders)) for n in range(min(len(expected.open_orders), len(actual.open_orders))): eo = expected.open_orders[n] ao = actual.open_orders[n] _cmp(f"open_orders[{n}].price", eo.price, ao.price, tol.get("order_price", 1e-6)) _cmp(f"open_orders[{n}].qty", eo.qty, ao.qty, tol.get("order_qty", 1e-9)) _cmp(f"open_orders[{n}].side", eo.side.value, ao.side.value) # Positions exp_pos = expected.account.positions if expected.account else {} act_pos = actual.account.positions if actual.account else {} _cmp("positions.count", len(exp_pos), len(act_pos)) for sym in set(list(exp_pos.keys()) + list(act_pos.keys())): ep = exp_pos.get(sym) ap = act_pos.get(sym) if ep and ap: pos_tol = tol.get("position_qty", 1e-9) _cmp(f"positions[{sym}].qty", ep.qty, ap.qty, pos_tol) _cmp(f"positions[{sym}].avg_entry", ep.avg_entry, ap.avg_entry, pos_tol) _cmp(f"positions[{sym}].side", ep.side.value if ep.side else None, ap.side.value if ap.side else None) elif ep and not ap: diffs.append(ReplayMismatch(i, f"positions[{sym}]", "present", "missing", "critical")) elif not ep and ap: diffs.append(ReplayMismatch(i, f"positions[{sym}]", "missing", "present", "critical")) # Trade path if expected.trade_path and actual.trade_path: ep = expected.trade_path ap = actual.trade_path _cmp("trade_path.pnl_bps", ep.pnl_bps, ap.pnl_bps, tol.get("pnl_bps", 0.1)) _cmp("trade_path.mae_bps", ep.mae_bps, ap.mae_bps, tol.get("mae_bps", 0.1)) _cmp("trade_path.mfe_bps", ep.mfe_bps, ap.mfe_bps, tol.get("mfe_bps", 0.1)) return diffs # ============================================================================== # Binary search for first mismatch # ============================================================================== def bisect_first_mismatch( cwm: CodeWorldModel, replay: Sequence[ReplayStep], lo: int = 0, hi: Optional[int] = None, tolerances: Optional[Mapping[str, float]] = None, ) -> Optional[ReplayMismatch]: """ Binary search for the first mismatch in a replay trajectory. Uses the CWM to re-simulate from known-good prefix, narrowing to the first divergence point. Much faster than linear scan for long trajectories. """ if hi is None: hi = len(replay) - 1 if lo > hi: return None # Find any mismatch in the range mid = (lo + hi) // 2 mismatches = _compare_deep( mid, replay[mid].after_ground_truth, cwm.transition(replay[mid].before, replay[mid].joint_action), tolerances, ) if mismatches: # Check if earlier steps also mismatch if mid > lo: earlier = bisect_first_mismatch(cwm, replay, lo, mid - 1, tolerances) if earlier: return earlier return mismatches[0] # No mismatch at mid, check right half return bisect_first_mismatch(cwm, replay, mid + 1, hi, tolerances) # ============================================================================== # Trajectory recording for self-play verification # ============================================================================== class TrajectoryRecorder: """ Records every state/action/next_state for deterministic re-run verification. For self-play: persist trajectory → re-run must produce exact same states. For historical: persist trajectory → CWM prediction must match ground truth. """ def __init__(self, max_steps: int = 10_000) -> None: self._max_steps = max_steps self._steps: list[TrajectoryRecord] = [] self._full_states: list[Tuple[MarketWorldState, Any, MarketWorldState]] = [] def record( self, step_index: int, before: MarketWorldState, action: Any, after: MarketWorldState, ) -> None: """Record one step. Keeps full states for detailed comparison.""" if len(self._steps) >= self._max_steps: return self._steps.append(TrajectoryRecord( step_index=step_index, before_hash=_hash_state(before), action_hash=_hash_action(action), after_hash=_hash_state(after), ts_ns=after.ts_ns, symbol=before.venue.symbol, )) self._full_states.append((before, action, after)) def verify_deterministic( self, cwm: CodeWorldModel, ) -> Tuple[bool, List[ReplayMismatch]]: """ Re-run the trajectory through CWM and verify exact match. Must produce identical states for same inputs. """ mismatches: List[ReplayMismatch] = [] for idx, (before, action, expected_after) in enumerate(self._full_states): actual_after = cwm.transition(before, action if isinstance(action, tuple) else (action,)) step_mismatches = _compare_deep(idx, expected_after, actual_after) mismatches.extend(step_mismatches) if any(m.is_critical for m in step_mismatches): break return (len(mismatches) == 0, mismatches) def trajectory_hash(self) -> str: """Hash of the entire trajectory for quick comparison.""" parts = [s.before_hash + s.action_hash + s.after_hash for s in self._steps] return hashlib.sha256("".join(parts).encode()).hexdigest()[:16] @property def step_count(self) -> int: return len(self._steps) @property def steps(self) -> List[TrajectoryRecord]: return list(self._steps) def to_replay_steps(self) -> List[ReplayStep]: """Convert recorded trajectory to ReplayStep list.""" return [ ReplayStep( before=before, joint_action=action if isinstance(action, tuple) else (action,), after_ground_truth=after, step_index=idx, ) for idx, (before, action, after) in enumerate(self._full_states) ] # ============================================================================== # ReplayVerifier — main interface # ============================================================================== class ReplayVerifier: """ Replay matching is mandatory. A fast wrong CWM is worse than a slow correct one. Three verification modes: 1. verify(): compare CWM predictions against ground truth steps 2. verify_deterministic(): re-run trajectory, check exact match 3. bisect(): binary search for first mismatch Tolerances: - Historical replay: exchange-data tolerances (feeds can drop) - Self-play replay: tight tolerances (deterministic) """ HISTORICAL_TOLERANCES = { "book_price": 0.01, # 1 cent "book_qty": 0.001, "account_equity": 0.01, "position_qty": 0.0001, "pnl_bps": 0.5, } SELF_PLAY_TOLERANCES = { "book_price": 1e-9, "book_qty": 1e-12, "account_equity": 1e-9, "position_qty": 1e-12, "pnl_bps": 1e-6, } def verify( self, cwm: CodeWorldModel, replay: Sequence[ReplayStep], tolerances: Optional[Mapping[str, float]] = None, ) -> ReplayResult: """Verify CWM predictions against ground truth steps.""" t0 = time.perf_counter_ns() tol = tolerances or self.HISTORICAL_TOLERANCES all_mismatches: List[ReplayMismatch] = [] steps_verified = 0 first_mismatch_idx = None for i, step in enumerate(replay): pred = cwm.transition(step.before, step.joint_action) mismatches = _compare_deep(i, step.after_ground_truth, pred, tol) steps_verified += 1 if mismatches: all_mismatches.extend(mismatches) if first_mismatch_idx is None: first_mismatch_idx = i break # Hash trajectory for caching traj_hash = hashlib.sha256( "".join(_hash_state(s.before) for s in replay).encode() ).hexdigest()[:16] return ReplayResult( passed=len(all_mismatches) == 0, mismatches=all_mismatches, steps_verified=steps_verified, total_steps=len(replay), first_mismatch_index=first_mismatch_idx, duration_ns=time.perf_counter_ns() - t0, trajectory_hash=traj_hash, ) def verify_determinism( self, cwm: CodeWorldModel, replay: Sequence[ReplayStep], ) -> ReplayResult: """Verify that re-running produces identical results.""" t0 = time.perf_counter_ns() mismatches: List[ReplayMismatch] = [] # Run once, collect results first_results: list = [] for step in replay: first_results.append(cwm.transition(step.before, step.joint_action)) # Run again, compare for i, step in enumerate(replay): second = cwm.transition(step.before, step.joint_action) step_mismatches = _compare_deep(i, first_results[i], second, self.SELF_PLAY_TOLERANCES) mismatches.extend(step_mismatches) if any(m.is_critical for m in step_mismatches): break traj_hash = hashlib.sha256( "".join(_hash_state(s.before) for s in replay).encode() ).hexdigest()[:16] return ReplayResult( passed=len(mismatches) == 0, mismatches=mismatches, steps_verified=len(replay), total_steps=len(replay), first_mismatch_index=mismatches[0].index if mismatches else None, duration_ns=time.perf_counter_ns() - t0, trajectory_hash=traj_hash, ) def bisect( self, cwm: CodeWorldModel, replay: Sequence[ReplayStep], tolerances: Optional[Mapping[str, float]] = None, ) -> Optional[ReplayMismatch]: """Binary search for first mismatch.""" return bisect_first_mismatch(cwm, replay, tolerances=tolerances)