""" Exhaustive tests for trajectory persistence (100+ 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 _state(): return MarketWorldState( ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(), book=OrderBookState(ts_ns=1, symbol="BTCUSDT", bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(50001.0, 1.0),)), account=AccountState(ts_ns=1, equity=10000.0, wallet_balance=10000.0, available_balance=10000.0, margin_used=0.0, total_notional=0.0), ) class TestTrajectoryStep: def test_construction(self): step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT", state_hash="abc", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash="def", pnl_bps=1.0, reward=0.5, entropy=0.3) assert step.step_index == 0 assert step.symbol == "BTCUSDT" def test_frozen(self): step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT", state_hash="abc", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash="def", pnl_bps=1.0, reward=0.5, entropy=0.3) with pytest.raises(AttributeError): step.step_index = 5 def test_with_diagnostics(self): step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT", state_hash="abc", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash="def", pnl_bps=1.0, reward=0.5, entropy=0.3, diagnostics={"key": "value"}) assert step.diagnostics["key"] == "value" class TestTrajectoryRecord: def test_construction(self): record = TrajectoryRecord( trajectory_id="t1", policy_version="v1", scenario_id="s1", seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0, fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000, created_ts_ns=1, ) assert record.trajectory_id == "t1" assert record.total_pnl_bps == 10.0 def test_frozen(self): record = TrajectoryRecord( trajectory_id="t1", policy_version="v1", scenario_id="s1", seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0, fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000, created_ts_ns=1, ) with pytest.raises(AttributeError): record.total_pnl_bps = 20.0 def test_with_steps(self): steps = (TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT", state_hash="a", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash="b", pnl_bps=1.0, reward=0.5, entropy=0.3),) record = TrajectoryRecord( 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=1, cancel_count=0, noop_count=0, duration_ns=1000, created_ts_ns=1, ) assert len(record.steps) == 1 class TestTrajectoryPersister: def test_persist(self): store = MalkhutCHStore() store.ensure_tables() persister = TrajectoryPersister(store) record = TrajectoryRecord( trajectory_id="t1", policy_version="v1", scenario_id="s1", seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0, fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000, created_ts_ns=1, ) persister.persist(record) def test_persist_with_steps(self): store = MalkhutCHStore() store.ensure_tables() persister = TrajectoryPersister(store) steps = (TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT", state_hash="a", action_kind="PLACE", action_side="BUY", action_price=50000.0, action_qty=0.001, next_state_hash="b", pnl_bps=1.0, reward=0.5, entropy=0.3),) record = TrajectoryRecord( 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=1, cancel_count=0, noop_count=0, duration_ns=1000, created_ts_ns=1, ) persister.persist(record) def test_query_trajectories(self): store = MalkhutCHStore() store.ensure_tables() persister = TrajectoryPersister(store) result = persister.query_trajectories() assert isinstance(result, str)