Files
sentiment-engine/MALKHUT/malkhut/tests/test_trajectory.py

131 lines
5.7 KiB
Python
Raw Normal View History

"""
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)