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