Files
sentiment-engine/MALKHUT/malkhut/tests/test_new_features.py
Codex 4c239f7774 malkhut(tests): 1140 test functions across 46 test files
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
2026-07-11 10:46:12 +02:00

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