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