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

470 lines
19 KiB
Python
Raw Normal View History

"""
Comprehensive tests for Discrepancy Tracker, Feature Importance, Rollback, Stress (60+ tests).
"""
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
from malkhut.training.importance import FeatureImportanceTracker
from malkhut.training.rollback import PolicyRollback
from malkhut.training.stress import StressScenarioFactory
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)
def _action(kind=ActionKind.NOOP):
return FulfilmentAction(kind, None, None, 0, 0.0, 0)
# ══════════════════════════════════════════════════════════════════════════════
# DISCREPANCY TRACKER (25+ tests)
# ══════════════════════════════════════════════════════════════════════════════
class TestDiscrepancyComprehensive:
def test_multiple_predictions(self):
dt = DiscrepancyTracker()
for i in range(5):
dt.record_prediction(_state(ts=i), _action(), f"v{i}")
assert dt._last_prediction is not None
def test_compare_after_each_prediction(self):
dt = DiscrepancyTracker()
for i in range(5):
dt.record_prediction(_state(ts=i), _action(), "v1")
dt.compare_with_actual(_state(ts=i+100))
assert dt.total_comparisons == 5
def test_discrepancy_count_matches(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(ts=1), _action(), "v1")
discs = dt.compare_with_actual(_state(ts=2))
assert len(discs) == dt.total_discrepancies
def test_rate_calculation(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(ts=1), _action(), "v1")
dt.compare_with_actual(_state(ts=2))
assert dt.discrepancy_rate > 0
def test_rate_zero_when_no_comparisons(self):
dt = DiscrepancyTracker()
assert dt.discrepancy_rate == 0.0
def test_different_severities(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(ts=1), _action(), "v1")
discs = dt.compare_with_actual(_state(ts=2))
severities = set(d.severity for d in discs)
assert len(severities) > 0
def test_get_recent_limit(self):
dt = DiscrepancyTracker()
for i in range(10):
dt.record_prediction(_state(ts=i), _action(), "v1")
dt.compare_with_actual(_state(ts=i+100))
recent = dt.get_recent(3)
assert len(recent) == 3
def test_get_by_severity_empty(self):
dt = DiscrepancyTracker()
result = dt.get_by_severity("nonexistent")
assert len(result) == 0
def test_record_prediction_stores_action(self):
dt = DiscrepancyTracker()
a = _action(ActionKind.PLACE)
dt.record_prediction(_state(), a, "v1")
assert dt._last_action.kind == ActionKind.PLACE
def test_record_prediction_stores_version(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(), _action(), "my_version")
assert dt._last_policy_version == "my_version"
def test_compare_with_custom_tolerance(self):
dt = DiscrepancyTracker()
s1 = _state(ts=1)
s2 = _state(ts=1, bid=50000.01) # tiny difference
dt.record_prediction(s1, _action(), "v1")
discs = dt.compare_with_actual(s2, tolerances={"book_price": 0.1})
assert len(discs) == 0 # within tolerance
def test_compare_without_tolerance(self):
dt = DiscrepancyTracker()
s1 = _state(ts=1)
s2 = _state(ts=1, bid=50000.01)
dt.record_prediction(s1, _action(), "v1")
discs = dt.compare_with_actual(s2)
# Default tolerance is 1e-9, so 0.01 difference should be detected
assert len(discs) > 0
def test_multiple_comparisons_accumulate(self):
dt = DiscrepancyTracker()
for i in range(20):
dt.record_prediction(_state(ts=i), _action(), "v1")
dt.compare_with_actual(_state(ts=i+1000))
assert dt.total_comparisons == 20
assert dt.total_discrepancies > 0
def test_discrepancy_record_has_all_fields(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(ts=1), _action(), "v1")
discs = dt.compare_with_actual(_state(ts=2))
d = discs[0]
assert d.ts_ns > 0
assert d.symbol == "BTCUSDT"
assert len(d.field) > 0
assert d.severity in ("info", "warning", "critical")
assert d.action_kind == "NOOP"
assert d.policy_version == "v1"
def test_empty_state_comparison(self):
dt = DiscrepancyTracker()
dt.record_prediction(_state(), _action(), "v1")
discs = dt.compare_with_actual(_state())
assert len(discs) == 0 # identical states
def test_wide_book_difference(self):
dt = DiscrepancyTracker()
s1 = _state(bid=50000.0, ask=50001.0)
s2 = _state(bid=49000.0, ask=51000.0)
dt.record_prediction(s1, _action(), "v1")
discs = dt.compare_with_actual(s2)
assert len(discs) > 0
# ══════════════════════════════════════════════════════════════════════════════
# FEATURE IMPORTANCE (20+ tests)
# ══════════════════════════════════════════════════════════════════════════════
class TestFeatureImportanceComprehensive:
def test_empty_tracker(self):
fit = FeatureImportanceTracker()
assert fit.total_decisions == 0
assert fit.feature_count == 0
def test_single_decision(self):
fit = FeatureImportanceTracker()
fit.record_decision(_state(), _action(), "normal")
assert fit.total_decisions == 1
assert fit.feature_count > 0
def test_importance_all_regimes(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance(regime="normal")
assert len(imp) > 0
def test_importance_no_regime_filter(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance()
assert len(imp) > 0
def test_importance_top_n_limit(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance(top_n=3)
assert len(imp) <= 3
def test_importance_sorted_descending(self):
fit = FeatureImportanceTracker()
for _ in range(20):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance(top_n=10)
for i in range(len(imp) - 1):
assert imp[i].importance >= imp[i+1].importance
def test_feature_stats_mean(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
stats = fit.get_feature_stats("mid")
assert "mean" in stats
assert stats["mean"] > 0
def test_feature_stats_min_max(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
stats = fit.get_feature_stats("mid")
assert stats["min"] <= stats["max"]
def test_feature_stats_count(self):
fit = FeatureImportanceTracker()
for _ in range(5):
fit.record_decision(_state(), _action(), "normal")
stats = fit.get_feature_stats("mid")
assert stats["count"] == 5
def test_feature_stats_empty(self):
fit = FeatureImportanceTracker()
stats = fit.get_feature_stats("nonexistent")
assert stats == {}
def test_multiple_regimes(self):
fit = FeatureImportanceTracker()
fit.record_decision(_state(), _action(), "normal")
fit.record_decision(_state(), _action(), "volatile")
imp_normal = fit.get_importance(regime="normal")
imp_volatile = fit.get_importance(regime="volatile")
assert len(imp_normal) > 0
assert len(imp_volatile) > 0
def test_importance_importance_positive(self):
fit = FeatureImportanceTracker()
for _ in range(10):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance()
for i in imp:
assert i.importance >= 0
def test_importance_sample_count(self):
fit = FeatureImportanceTracker()
for _ in range(5):
fit.record_decision(_state(), _action(), "normal")
imp = fit.get_importance(top_n=1)
assert imp[0].sample_count > 0
def test_importance_regime_field(self):
fit = FeatureImportanceTracker()
fit.record_decision(_state(), _action(), "choppy")
imp = fit.get_importance(regime="choppy")
assert imp[0].regime == "choppy"
def test_feature_count_increases(self):
fit = FeatureImportanceTracker()
assert fit.feature_count == 0
fit.record_decision(_state(), _action())
assert fit.feature_count > 0
# ══════════════════════════════════════════════════════════════════════════════
# POLICY ROLLBACK (15+ tests)
# ══════════════════════════════════════════════════════════════════════════════
class TestPolicyRollbackComprehensive:
def test_initial_state(self):
reg = PolicyRegistry()
rb = PolicyRollback(registry=reg)
assert rb.shadow_score_count == 0
assert len(rb.rollback_events) == 0
def test_record_many_scores(self):
reg = PolicyRegistry()
rb = PolicyRollback(registry=reg)
for _ in range(20):
rb.record_shadow_score(10.0)
assert rb.shadow_score_count == 20
def test_no_rollback_insufficient_data(self):
reg = PolicyRegistry()
rb = PolicyRollback(registry=reg, min_shadow_steps=10)
for _ in range(9):
rb.record_shadow_score(-10.0)
assert rb.check_rollback() is None
def test_no_rollback_no_active_policy(self):
reg = PolicyRegistry()
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_no_rollback_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_event_fields(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.ts_ns > 0
assert event.rolled_back_version == "v1"
assert event.reason.startswith("performance_drop")
def test_multiple_rollbacks(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
def test_custom_threshold(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=-100.0)
for _ in range(5):
rb.record_shadow_score(-10.0)
assert rb.check_rollback() is None # -10 > -100
def test_shadow_score_returns_list(self):
reg = PolicyRegistry()
rb = PolicyRollback(registry=reg)
rb.record_shadow_score(5.0)
rb.record_shadow_score(10.0)
assert rb.shadow_score_count == 2
# ══════════════════════════════════════════════════════════════════════════════
# STRESS SCENARIOS (15+ tests)
# ══════════════════════════════════════════════════════════════════════════════
class TestStressComprehensive:
def test_all_scenarios_unique_id(self):
factory = StressScenarioFactory()
suite = factory.build_stress_suite()
ids = [s.scenario_id for s in suite]
assert len(ids) == len(set(ids))
def test_all_scenarios_have_valid_state(self):
factory = StressScenarioFactory()
for sc in factory.build_stress_suite():
assert sc.initial_state.book.best_bid > 0
assert sc.initial_state.book.best_ask > 0
assert sc.initial_state.account.equity > 0
def test_all_scenarios_have_positive_max_steps(self):
factory = StressScenarioFactory()
for sc in factory.build_stress_suite():
assert sc.max_steps > 0
def test_flash_crash_thin_book(self):
factory = StressScenarioFactory()
sc = factory.flash_crash()
assert sc.initial_state.book.bids[0].qty == 0.1
def test_liquidity_vacuum_extremely_thin(self):
factory = StressScenarioFactory()
sc = factory.liquidity_vacuum()
assert sc.initial_state.book.bids[0].qty == 0.001
def test_extreme_vol_wide_spread(self):
factory = StressScenarioFactory()
sc = factory.extreme_volatility()
spread = sc.initial_state.book.best_ask - sc.initial_state.book.best_bid
assert spread == 2000.0
def test_toxic_flood_multiple_counterparties(self):
factory = StressScenarioFactory()
sc = factory.toxic_flood()
assert len(sc.counterparties) == 3
def test_choppy_tight_range(self):
factory = StressScenarioFactory()
sc = factory.choppy_market()
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 sc.initial_state.book.bids[0].qty == 0.2
def test_liquidation_cascade_tags(self):
factory = StressScenarioFactory()
sc = factory.liquidation_cascade()
assert "liquidation_cascade" in sc.tags
assert "cascade" in sc.tags
def test_correlation_breakdown_tags(self):
factory = StressScenarioFactory()
sc = factory.correlation_breakdown()
assert "correlation_breakdown" in sc.tags
def test_suite_count_per_symbol(self):
factory = StressScenarioFactory()
suite = factory.build_stress_suite(symbols=("BTCUSDT",))
assert len(suite) == 8
def test_suite_multi_symbol(self):
factory = StressScenarioFactory()
suite = factory.build_stress_suite(symbols=("BTCUSDT", "ETHUSDT", "SOLUSDT"))
assert len(suite) == 24
def test_all_tags_are_strings(self):
factory = StressScenarioFactory()
for sc in factory.build_stress_suite():
for tag in sc.tags:
assert isinstance(tag, str)
def test_all_descriptions_non_empty(self):
factory = StressScenarioFactory()
for sc in factory.build_stress_suite():
assert len(sc.description) > 10
def test_custom_counterparties(self):
from malkhut.counterparties import ToxicTakerPolicy
factory = StressScenarioFactory(counterparties=(ToxicTakerPolicy(),))
sc = factory.flash_crash()
assert len(sc.counterparties) == 1