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

339 lines
14 KiB
Python
Raw Normal View History

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