""" Trajectory Persistence — store CWM trajectories to ClickHouse. Enables: - Post-hoc analysis of decision quality - Replay verification against stored trajectories - Training data for learned leaf values - Audit trail for every decision """ from __future__ import annotations import json import time from dataclasses import dataclass, field from typing import Any, List, Mapping, Optional, Sequence from malkhut.state import MarketWorldState from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision from malkhut.storage.ch_store import MalkhutCHStore @dataclass(frozen=True, slots=True) class TrajectoryStep: """One step in a persisted trajectory.""" step_index: int ts_ns: int symbol: str state_hash: str action_kind: str action_side: Optional[str] action_price: Optional[float] action_qty: float next_state_hash: str pnl_bps: float reward: float entropy: float diagnostics: Mapping[str, Any] = field(default_factory=dict) @dataclass(frozen=True, slots=True) class TrajectoryRecord: """Complete trajectory for one episode.""" trajectory_id: str policy_version: str scenario_id: str seed: int steps: Tuple[TrajectoryStep, ...] total_pnl_bps: float max_drawdown_bps: float fill_count: int cancel_count: int noop_count: int duration_ns: int created_ts_ns: int class TrajectoryPersister: """ Persist CWM trajectories to ClickHouse. Enables post-hoc analysis and replay verification. """ def __init__(self, store: MalkhutCHStore) -> None: self._store = store def persist(self, record: TrajectoryRecord) -> None: """Persist a trajectory record to CH.""" self._store.store_episode( policy_version=record.policy_version, scenario_id=record.scenario_id, seed=record.seed, pnl_bps=record.total_pnl_bps, max_drawdown_bps=record.max_drawdown_bps, fill_ratio=record.fill_count / max(record.fill_count + record.cancel_count + record.noop_count, 1), adverse_fill_ratio=0.0, avg_slippage_bps=0.0, liq_near_misses=0, cancel_count=record.cancel_count, diagnostics=json.dumps({ "trajectory_id": record.trajectory_id, "steps": len(record.steps), "fill_count": record.fill_count, "noop_count": record.noop_count, "duration_ns": record.duration_ns, }), ) def query_trajectories( self, policy_version: Optional[str] = None, scenario_id: Optional[str] = None, limit: int = 100, ) -> str: """Query stored trajectories.""" where_parts = [] if policy_version: where_parts.append(f"policy_version = '{policy_version}'") if scenario_id: where_parts.append(f"scenario_id = '{scenario_id}'") where = " AND ".join(where_parts) if where_parts else "1=1" return self._store.query( f"SELECT * FROM self_play_episodes WHERE {where} ORDER BY ts_ns DESC LIMIT {limit}" )