180 lines
6.1 KiB
Python
180 lines
6.1 KiB
Python
|
|
"""
|
||
|
|
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)
|