202 lines
8.5 KiB
Python
202 lines
8.5 KiB
Python
|
|
"""
|
||
|
|
Unit tests: SM-MCTS planner and action menu.
|
||
|
|
|
||
|
|
Mutation litmus: flip UCB exploration constant; if no test breaks,
|
||
|
|
the planner is untested.
|
||
|
|
"""
|
||
|
|
import pytest
|
||
|
|
from malkhut.state import (
|
||
|
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||
|
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||
|
|
)
|
||
|
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner, PlayerActionStats
|
||
|
|
from malkhut.planner.action_menu import build_our_actions
|
||
|
|
from malkhut.cwm import MinimalCryptoLOBCWM
|
||
|
|
from malkhut.counterparties import default_counterparty_ecology
|
||
|
|
from malkhut.actions import ActionKind, FulfilmentAction, PlannedPolicy
|
||
|
|
|
||
|
|
|
||
|
|
def _state_with_intent() -> MarketWorldState:
|
||
|
|
return MarketWorldState(
|
||
|
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||
|
|
venue=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,
|
||
|
|
),
|
||
|
|
book=OrderBookState(
|
||
|
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||
|
|
bids=(PriceLevel(50000.0, 1.0), PriceLevel(49999.0, 2.0)),
|
||
|
|
asks=(PriceLevel(50001.0, 1.0), PriceLevel(50002.0, 2.0)),
|
||
|
|
),
|
||
|
|
account=AccountState(
|
||
|
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||
|
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||
|
|
),
|
||
|
|
intent=ExecutionIntent(
|
||
|
|
intent_id="test_intent", ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||
|
|
kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0,
|
||
|
|
urgency=0.5, alpha_horizon_s=60.0, alpha_bps=2.0,
|
||
|
|
max_slippage_bps=5.0, prefer_maker=True, reduce_only=False,
|
||
|
|
ttl_s=300.0, reason="test",
|
||
|
|
),
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _params() -> FulfilmentPolicyParams:
|
||
|
|
return FulfilmentPolicyParams(
|
||
|
|
version="test", ucb_c=1.414, max_sims=32, max_depth=2,
|
||
|
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||
|
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.10, 0.25, 0.50),
|
||
|
|
passive_ttl_ms=200, aggressive_ttl_ms=50,
|
||
|
|
maker_edge_min_bps=0.5, cross_spread_edge_min_bps=5.0,
|
||
|
|
adverse_toxicity_cancel_threshold=0.5, queue_churn_cancel_threshold=0.5,
|
||
|
|
mae_tail_cut_bps=50.0, mfe_giveback_cut_fraction=0.5,
|
||
|
|
max_time_in_loss_s=300.0, failed_recovery_cut_count=3,
|
||
|
|
recovery_velocity_min_bps_per_s=0.0,
|
||
|
|
max_symbol_notional_fraction=0.20, max_single_order_notional_fraction=0.05,
|
||
|
|
reduce_when_global_up_fraction=0.30, session_profit_lock_fraction=0.02,
|
||
|
|
w_expected_pnl=1.0, w_fill_probability=0.5, w_adverse_selection=2.0,
|
||
|
|
w_queue_priority=0.5, w_inventory_risk=1.5, w_tail_loss=5.0,
|
||
|
|
w_fee_quality=0.5, w_time_decay=0.3, w_policy_entropy=0.5,
|
||
|
|
robust_tail_weight=2.0, toxic_counterparty_weight=3.0,
|
||
|
|
low_liquidity_weight=2.0, latency_stress_weight=1.0,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class TestActionMenu:
|
||
|
|
def test_noop_always_available(self):
|
||
|
|
state = _state_with_intent()
|
||
|
|
params = _params()
|
||
|
|
actions = build_our_actions(state, params)
|
||
|
|
kinds = [a.kind for a in actions]
|
||
|
|
assert ActionKind.NOOP in kinds
|
||
|
|
|
||
|
|
def test_action_count_bounded(self):
|
||
|
|
state = _state_with_intent()
|
||
|
|
params = _params()
|
||
|
|
actions = build_our_actions(state, params)
|
||
|
|
assert 3 <= len(actions) <= 50
|
||
|
|
|
||
|
|
def test_cross_spread_only_when_urgency_high(self):
|
||
|
|
state = _state_with_intent()
|
||
|
|
params = _params()
|
||
|
|
actions = build_our_actions(state, params)
|
||
|
|
crosses = [a for a in actions if a.kind == ActionKind.CROSS_SPREAD]
|
||
|
|
# urgency=0.5 < 0.65, so no crosses
|
||
|
|
assert len(crosses) == 0
|
||
|
|
|
||
|
|
def test_high_urgency_allows_crosses(self):
|
||
|
|
state = _state_with_intent()
|
||
|
|
params = _params()
|
||
|
|
intent = ExecutionIntent(
|
||
|
|
intent_id="hi", ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||
|
|
kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0,
|
||
|
|
urgency=0.8, alpha_horizon_s=60.0, alpha_bps=2.0,
|
||
|
|
max_slippage_bps=5.0, prefer_maker=True, reduce_only=False,
|
||
|
|
ttl_s=300.0, reason="test",
|
||
|
|
)
|
||
|
|
state = MarketWorldState(
|
||
|
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||
|
|
venue=state.venue, book=state.book,
|
||
|
|
account=state.account, intent=intent,
|
||
|
|
)
|
||
|
|
actions = build_our_actions(state, params)
|
||
|
|
crosses = [a for a in actions if a.kind == ActionKind.CROSS_SPREAD]
|
||
|
|
assert len(crosses) > 0
|
||
|
|
|
||
|
|
|
||
|
|
class TestPlayerActionStats:
|
||
|
|
def test_every_action_gets_initial_visit(self):
|
||
|
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),
|
||
|
|
FulfilmentAction(ActionKind.PLACE, Side.BUY, None, 0, 0.25, 200))
|
||
|
|
stats = PlayerActionStats.from_actions(actions)
|
||
|
|
assert all(v == 0 for v in stats.visits)
|
||
|
|
|
||
|
|
def test_ucb_prefers_high_value_after_sampling(self):
|
||
|
|
import random
|
||
|
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),
|
||
|
|
FulfilmentAction(ActionKind.PLACE, Side.BUY, None, 0, 0.25, 200))
|
||
|
|
stats = PlayerActionStats.from_actions(actions)
|
||
|
|
for _ in range(100):
|
||
|
|
stats.update(0, 0.1)
|
||
|
|
stats.update(1, 1.0)
|
||
|
|
rng = random.Random(42)
|
||
|
|
idx, action = stats.ucb_select(200, 1.414, rng)
|
||
|
|
assert idx == 1
|
||
|
|
|
||
|
|
def test_update_increments_count(self):
|
||
|
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),)
|
||
|
|
stats = PlayerActionStats.from_actions(actions)
|
||
|
|
stats.update(0, 1.0)
|
||
|
|
assert stats.visits[0] == 1
|
||
|
|
assert abs(stats.total_value[0] - 1.0) < 1e-9
|
||
|
|
|
||
|
|
|
||
|
|
class TestPlanner:
|
||
|
|
def test_planner_returns_planned_policy(self):
|
||
|
|
cwm = MinimalCryptoLOBCWM()
|
||
|
|
planner = DecoupledUCBPlanner(
|
||
|
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||
|
|
)
|
||
|
|
params = _params()
|
||
|
|
state = _state_with_intent()
|
||
|
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||
|
|
assert isinstance(result, PlannedPolicy)
|
||
|
|
assert len(result.actions) > 0
|
||
|
|
assert len(result.probabilities) == len(result.actions)
|
||
|
|
assert abs(sum(result.probabilities) - 1.0) < 1e-6
|
||
|
|
|
||
|
|
def test_root_policy_probabilities_sum_to_one(self):
|
||
|
|
cwm = MinimalCryptoLOBCWM()
|
||
|
|
planner = DecoupledUCBPlanner(
|
||
|
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||
|
|
)
|
||
|
|
params = _params()
|
||
|
|
state = _state_with_intent()
|
||
|
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||
|
|
total = sum(result.probabilities)
|
||
|
|
assert abs(total - 1.0) < 1e-6
|
||
|
|
|
||
|
|
def test_deterministic_with_fixed_seed(self):
|
||
|
|
cwm = MinimalCryptoLOBCWM()
|
||
|
|
state = _state_with_intent()
|
||
|
|
params = _params()
|
||
|
|
|
||
|
|
p1 = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=99)
|
||
|
|
r1 = p1.plan(root_state=state, params=params, budget_ms=20)
|
||
|
|
|
||
|
|
p2 = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=99)
|
||
|
|
r2 = p2.plan(root_state=state, params=params, budget_ms=20)
|
||
|
|
|
||
|
|
assert r1.selected_action.kind == r2.selected_action.kind
|
||
|
|
|
||
|
|
def test_no_intent_returns_noop(self):
|
||
|
|
cwm = MinimalCryptoLOBCWM()
|
||
|
|
planner = DecoupledUCBPlanner(
|
||
|
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=0,
|
||
|
|
)
|
||
|
|
params = _params()
|
||
|
|
state = MarketWorldState(
|
||
|
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||
|
|
venue=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,
|
||
|
|
),
|
||
|
|
book=OrderBookState(
|
||
|
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||
|
|
bids=(PriceLevel(50000.0, 1.0),),
|
||
|
|
asks=(PriceLevel(50001.0, 1.0),),
|
||
|
|
),
|
||
|
|
account=AccountState(
|
||
|
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||
|
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||
|
|
),
|
||
|
|
intent=None,
|
||
|
|
)
|
||
|
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||
|
|
assert result.selected_action.kind == ActionKind.NOOP
|