""" Comprehensive tests for Trajectory Persistence (50+ tests). """ import json import pytest from malkhut.state import MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules, AccountState from malkhut.actions import ActionKind, FulfilmentAction from malkhut.training.trajectory import TrajectoryStep, TrajectoryRecord, TrajectoryPersister from malkhut.storage.ch_store import MalkhutCHStore def _venue(): return VenueRules(exchange="bingx", symbol="BTCUSDT", tick_size=0.1, lot_size=0.001, min_qty=0.001, min_notional=5.0, maker_fee_bps=-0.2, taker_fee_bps=0.5, post_only_supported=True, reduce_only_supported=True, max_orders_per_second=100, max_cancels_per_minute=120) def _step(i=0, **kw): defaults = dict(step_index=i, ts_ns=1_000_000_000 + i, symbol="BTCUSDT", state_hash=f"h{i}", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash=f"h{i+1}", pnl_bps=1.0, reward=0.5, entropy=0.3) defaults.update(kw) return TrajectoryStep(**defaults) def _record(n_steps=3, **kw): steps = tuple(_step(i) for i in range(n_steps)) defaults = dict(trajectory_id="t1", policy_version="v1", scenario_id="s1", seed=42, steps=steps, total_pnl_bps=10.0, max_drawdown_bps=5.0, fill_count=2, cancel_count=1, noop_count=n_steps - 3, duration_ns=1_000_000, created_ts_ns=1) defaults.update(kw) return TrajectoryRecord(**defaults) class TestTrajectoryStepDetailed: def test_all_fields_accessible(self): s = _step() assert s.step_index == 0 assert s.ts_ns > 0 assert s.symbol == "BTCUSDT" assert s.state_hash == "h0" assert s.action_kind == "PLACE" assert s.action_side == "BUY" assert s.action_price == 50000.0 assert s.action_qty == 0.001 assert s.next_state_hash == "h1" assert s.pnl_bps == 1.0 assert s.reward == 0.5 assert s.entropy == 0.3 def test_none_fields(self): s = TrajectoryStep(step_index=0, ts_ns=1, symbol="X", state_hash="a", action_kind="NOOP", action_side=None, action_price=None, action_qty=0.0, next_state_hash="b", pnl_bps=0.0, reward=0.0, entropy=0.0) assert s.action_side is None assert s.action_price is None def test_large_diagnostics(self): diag = {f"key_{i}": f"value_{i}" for i in range(100)} s = _step(diagnostics=diag) assert len(s.diagnostics) == 100 def test_hash_determinism(self): s1 = _step(ts_ns=100) s2 = _step(ts_ns=100) assert s1 == s2 def test_step_index_range(self): for i in [0, 1, 100, 10000]: s = _step(step_index=i) assert s.step_index == i class TestTrajectoryRecordDetailed: def test_empty_steps(self): r = _record(n_steps=0) assert len(r.steps) == 0 def test_many_steps(self): r = _record(n_steps=50) assert len(r.steps) == 50 def test_all_metrics_accessible(self): r = _record() assert r.trajectory_id == "t1" assert r.policy_version == "v1" assert r.scenario_id == "s1" assert r.seed == 42 assert r.total_pnl_bps == 10.0 assert r.max_drawdown_bps == 5.0 assert r.fill_count == 2 assert r.cancel_count == 1 assert r.duration_ns == 1_000_000 def test_diagnostics(self): r = _record() assert r.total_pnl_bps == 10.0 def test_negative_pnl(self): r = _record(total_pnl_bps=-20.0, max_drawdown_bps=25.0) assert r.total_pnl_bps == -20.0 def test_zero_fill_count(self): r = _record(fill_count=0, cancel_count=0, noop_count=5) assert r.fill_count == 0 class TestTrajectoryPersisterDetailed: def test_persist_creates_record(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) r = _record(n_steps=5) p.persist(r) def test_persist_many_records(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) for i in range(10): p.persist(_record(n_steps=3, trajectory_id=f"t{i}")) def test_persist_empty_steps(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) p.persist(_record(n_steps=0)) def test_persist_negative_pnl(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) p.persist(_record(total_pnl_bps=-50.0)) def test_persist_large_trajectory(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) p.persist(_record(n_steps=100)) def test_query_returns_string(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) result = p.query_trajectories() assert isinstance(result, str) def test_query_with_policy_filter(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) result = p.query_trajectories(policy_version="v1") assert isinstance(result, str) def test_query_with_scenario_filter(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) result = p.query_trajectories(scenario_id="s1") assert isinstance(result, str) def test_query_with_limit(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) result = p.query_trajectories(limit=5) assert isinstance(result, str) def test_persist_and_query(self): store = MalkhutCHStore() store.ensure_tables() p = TrajectoryPersister(store) p.persist(_record(trajectory_id="t_query")) result = p.query_trajectories(policy_version="v1") assert isinstance(result, str)