CWM (103): core mechanics, exhaustive edge cases, numba, exchange mechanics Replay (118): exhaustive verification, microstructure, trajectory Training (190): asset classification, phase0 extensive, pipeline, exhaustive DSL (102): v2 syntax, expanded, new features ASEx (33): validate-before-mutate, single-writer Planner (48): MCTS, alternatives, hooks Counterparties (19): 9 adversarial agent policies Clock (30): event-driven reactor BingX (28): venue adapter IPC (8): Zinc SHM Storage (9): ClickHouse Risk (4): hard invariants State (17): frozen dataclass invariants Integration: E2E, concurrency, sync/async seams, hypothesis, fuzz, adversarial
339 lines
14 KiB
Python
339 lines
14 KiB
Python
"""
|
|
Exhaustive tests for discrepancy tracker, feature importance, rollback, stress scenarios.
|
|
"""
|
|
import pytest
|
|
from malkhut.state import (
|
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
|
OrderBookState, PriceLevel, Side, TradePathState, VenueRules,
|
|
)
|
|
from malkhut.actions import ActionKind, FulfilmentAction
|
|
from malkhut.training.discrepancy import DiscrepancyTracker, DiscrepancyRecord
|
|
from malkhut.training.importance import FeatureImportanceTracker, FeatureImportance
|
|
from malkhut.training.rollback import PolicyRollback, RollbackEvent
|
|
from malkhut.training.stress import StressScenarioFactory, StressScenario
|
|
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
|
|
|
|
|
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(**kw):
|
|
tp = kw.get("trade_path")
|
|
return MarketWorldState(
|
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
|
book=OrderBookState(ts_ns=1, symbol="BTCUSDT",
|
|
bids=(PriceLevel(kw.get("bid", 50000.0), 1.0),),
|
|
asks=(PriceLevel(kw.get("ask", 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),
|
|
trade_path=tp,
|
|
)
|
|
|
|
|
|
def _baseline(**kw):
|
|
d = dict(version="baseline", ucb_c=1.414, max_sims=256, max_depth=3,
|
|
rollout_depth=3, root_temperature=0.5, min_root_entropy=0.25,
|
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.1, 0.25, 0.5),
|
|
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)
|
|
d.update(kw)
|
|
return FulfilmentPolicyParams(**d)
|
|
|
|
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
# DISCREPANCY TRACKER
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
|
|
class TestDiscrepancyTracker:
|
|
def test_record_prediction(self):
|
|
dt = DiscrepancyTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s, a, "v1")
|
|
assert dt._last_prediction is not None
|
|
|
|
def test_compare_identical_states(self):
|
|
dt = DiscrepancyTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s, a, "v1")
|
|
discs = dt.compare_with_actual(s)
|
|
assert len(discs) == 0
|
|
|
|
def test_compare_different_states(self):
|
|
dt = DiscrepancyTracker()
|
|
s1 = _state(ts=1)
|
|
s2 = _state(ts=2)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s1, a, "v1")
|
|
discs = dt.compare_with_actual(s2)
|
|
assert len(discs) > 0
|
|
|
|
def test_discrepancy_record_fields(self):
|
|
dt = DiscrepancyTracker()
|
|
s1 = _state(ts=1)
|
|
s2 = _state(ts=2)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s1, a, "v1")
|
|
discs = dt.compare_with_actual(s2)
|
|
assert discs[0].field == "ts_ns"
|
|
assert discs[0].predicted == 1
|
|
assert discs[0].actual == 2
|
|
|
|
def test_discrepancy_rate(self):
|
|
dt = DiscrepancyTracker()
|
|
s1 = _state(ts=1)
|
|
s2 = _state(ts=2)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s1, a, "v1")
|
|
dt.compare_with_actual(s2)
|
|
assert dt.discrepancy_rate > 0
|
|
|
|
def test_total_comparisons(self):
|
|
dt = DiscrepancyTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s, a, "v1")
|
|
dt.compare_with_actual(s)
|
|
dt.record_prediction(s, a, "v1")
|
|
dt.compare_with_actual(s)
|
|
assert dt.total_comparisons == 2
|
|
|
|
def test_get_recent(self):
|
|
dt = DiscrepancyTracker()
|
|
s1 = _state(ts=1)
|
|
s2 = _state(ts=2)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s1, a, "v1")
|
|
dt.compare_with_actual(s2)
|
|
recent = dt.get_recent(5)
|
|
assert len(recent) >= 1
|
|
|
|
def test_get_by_severity(self):
|
|
dt = DiscrepancyTracker()
|
|
s1 = _state(ts=1)
|
|
s2 = _state(ts=2)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
dt.record_prediction(s1, a, "v1")
|
|
dt.compare_with_actual(s2)
|
|
# ts_ns mismatch is "info" severity
|
|
info_discs = dt.get_by_severity("info")
|
|
assert len(info_discs) >= 1
|
|
|
|
def test_no_prediction_returns_empty(self):
|
|
dt = DiscrepancyTracker()
|
|
s = _state()
|
|
discs = dt.compare_with_actual(s)
|
|
assert len(discs) == 0
|
|
|
|
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
# FEATURE IMPORTANCE TRACKER
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
|
|
class TestFeatureImportance:
|
|
def test_record_decision(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
fit.record_decision(s, a, "normal")
|
|
assert fit.total_decisions == 1
|
|
|
|
def test_get_importance(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
for _ in range(10):
|
|
fit.record_decision(s, a, "normal")
|
|
importance = fit.get_importance(top_n=5)
|
|
assert len(importance) > 0
|
|
assert all(isinstance(i, FeatureImportance) for i in importance)
|
|
|
|
def test_importance_sorted(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
for _ in range(10):
|
|
fit.record_decision(s, a, "normal")
|
|
importance = fit.get_importance(top_n=10)
|
|
for i in range(len(importance) - 1):
|
|
assert importance[i].importance >= importance[i+1].importance
|
|
|
|
def test_get_feature_stats(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
for _ in range(10):
|
|
fit.record_decision(s, a, "normal")
|
|
stats = fit.get_feature_stats("mid")
|
|
assert "mean" in stats
|
|
assert "min" in stats
|
|
assert "max" in stats
|
|
assert stats["count"] == 10
|
|
|
|
def test_feature_count(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
fit.record_decision(s, a, "normal")
|
|
assert fit.feature_count > 0
|
|
|
|
def test_regime_filter(self):
|
|
fit = FeatureImportanceTracker()
|
|
s = _state()
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
fit.record_decision(s, a, "normal")
|
|
fit.record_decision(s, a, "volatile")
|
|
importance_normal = fit.get_importance(top_n=5, regime="normal")
|
|
importance_volatile = fit.get_importance(top_n=5, regime="volatile")
|
|
assert len(importance_normal) > 0
|
|
assert len(importance_volatile) > 0
|
|
|
|
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
# POLICY ROLLBACK
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
|
|
class TestPolicyRollback:
|
|
def test_record_shadow_score(self):
|
|
reg = PolicyRegistry()
|
|
rb = PolicyRollback(registry=reg)
|
|
rb.record_shadow_score(10.0)
|
|
assert rb.shadow_score_count == 1
|
|
|
|
def test_no_rollback_when_insufficient_data(self):
|
|
reg = PolicyRegistry()
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=10)
|
|
for _ in range(5):
|
|
rb.record_shadow_score(-10.0)
|
|
assert rb.check_rollback() is None
|
|
|
|
def test_no_rollback_when_performance_ok(self):
|
|
reg = PolicyRegistry()
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3)
|
|
for _ in range(5):
|
|
rb.record_shadow_score(10.0)
|
|
assert rb.check_rollback() is None
|
|
|
|
def test_rollback_on_degradation(self):
|
|
reg = PolicyRegistry()
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-5.0)
|
|
for _ in range(5):
|
|
rb.record_shadow_score(-10.0)
|
|
event = rb.check_rollback()
|
|
assert event is not None
|
|
assert event.rolled_back_version == "v1"
|
|
|
|
def test_rollback_events_tracked(self):
|
|
reg = PolicyRegistry()
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-5.0)
|
|
for _ in range(5):
|
|
rb.record_shadow_score(-10.0)
|
|
rb.check_rollback()
|
|
assert len(rb.rollback_events) == 1
|
|
|
|
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
# STRESS SCENARIOS
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
|
|
|
class TestStressScenarios:
|
|
def test_flash_crash(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.flash_crash()
|
|
assert isinstance(sc, StressScenario)
|
|
assert "flash_crash" in sc.tags
|
|
assert sc.max_steps == 10
|
|
|
|
def test_liquidity_vacuum(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.liquidity_vacuum()
|
|
assert "liquidity_vacuum" in sc.tags
|
|
assert sc.initial_state.book.bids[0].qty == 0.001
|
|
|
|
def test_extreme_volatility(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.extreme_volatility()
|
|
assert "extreme_volatility" in sc.tags
|
|
spread = sc.initial_state.book.best_ask - sc.initial_state.book.best_bid
|
|
assert spread == 2000.0
|
|
|
|
def test_toxic_flood(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.toxic_flood()
|
|
assert "toxic_flood" in sc.tags
|
|
assert len(sc.counterparties) == 3
|
|
|
|
def test_choppy_market(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.choppy_market()
|
|
assert "choppy" in sc.tags
|
|
spread = sc.initial_state.book.best_ask - sc.initial_state.book.best_bid
|
|
assert spread == 0.5
|
|
|
|
def test_weekend_low_participation(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.weekend_low_participation()
|
|
assert "weekend" in sc.tags
|
|
|
|
def test_liquidation_cascade(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.liquidation_cascade()
|
|
assert "liquidation_cascade" in sc.tags
|
|
|
|
def test_correlation_breakdown(self):
|
|
factory = StressScenarioFactory()
|
|
sc = factory.correlation_breakdown()
|
|
assert "correlation_breakdown" in sc.tags
|
|
|
|
def test_build_stress_suite(self):
|
|
factory = StressScenarioFactory()
|
|
suite = factory.build_stress_suite(symbols=("BTCUSDT",))
|
|
assert len(suite) == 8
|
|
|
|
def test_multi_symbol_suite(self):
|
|
factory = StressScenarioFactory()
|
|
suite = factory.build_stress_suite(symbols=("BTCUSDT", "ETHUSDT"))
|
|
assert len(suite) == 16
|
|
|
|
def test_all_scenarios_have_state(self):
|
|
factory = StressScenarioFactory()
|
|
for sc in factory.build_stress_suite():
|
|
assert sc.initial_state is not None
|
|
assert sc.initial_state.account.equity > 0
|
|
|
|
def test_all_scenarios_have_counterparties(self):
|
|
factory = StressScenarioFactory()
|
|
for sc in factory.build_stress_suite():
|
|
assert len(sc.counterparties) > 0
|
|
|
|
def test_all_scenarios_have_tags(self):
|
|
factory = StressScenarioFactory()
|
|
for sc in factory.build_stress_suite():
|
|
assert len(sc.tags) > 0
|
|
|
|
def test_all_scenarios_have_descriptions(self):
|
|
factory = StressScenarioFactory()
|
|
for sc in factory.build_stress_suite():
|
|
assert len(sc.description) > 0
|