""" Comprehensive tests for planner alternatives, hftbacktest validation, training parallelism, observability, and hooks (100+ tests). """ import json import math import os import tempfile import threading import time import pytest from malkhut.state import ( AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind, MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules, ) from malkhut.actions import ActionKind, FulfilmentAction, PlannedPolicy, RiskDecision from malkhut.cwm.core import MinimalCryptoLOBCWM from malkhut.cwm.hftbacktest_validator import HftBacktestValidator, ValidationReport from malkhut.planner.alternatives import ( EXP3Planner, RegretMatchingPlanner, UCB1Planner, ThompsonSamplingPlanner, HedgePlanner, create_planner, PLANNER_REGISTRY, ) from malkhut.training.parallel import ParallelEvaluator, TrainingMonitor, TrainingMetrics from malkhut.training.observability import ObservabilityLogger, DecisionRecord from malkhut.training.hooks import ExecutionHooks from malkhut.counterparties import default_counterparty_ecology 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(): return MarketWorldState( ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(), book=OrderBookState(ts_ns=1, symbol="BTCUSDT", bids=(PriceLevel(50000.0, 1.0), PriceLevel(49999.0, 2.0)), asks=(PriceLevel(50001.0, 1.0), PriceLevel(50002.0, 2.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), ) def _intent(): from malkhut.state import ExecutionIntent return ExecutionIntent( intent_id="test", ts_ns=1_000_000_000, symbol="BTCUSDT", kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0, urgency=0.5, alpha_horizon_s=60.0, alpha_bps=2.0, max_slippage_bps=5.0, prefer_maker=True, reduce_only=False, ttl_s=300.0, reason="test", ) def _baseline(): return FulfilmentPolicyParams( 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, ) def _state_with_intent(): return MarketWorldState( ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(), book=OrderBookState(ts_ns=1, symbol="BTCUSDT", bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(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), intent=_intent(), ) # ══════════════════════════════════════════════════════════════════════════════ # HFTBACKTEST VALIDATOR (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestHftBacktestValidator: def test_validate_empty(self): v = HftBacktestValidator() cwm = MinimalCryptoLOBCWM() report = v.validate(cwm, []) assert report.total_steps == 0 # Empty replay is valid (nothing to compare) assert report.fill_match_rate == 0.0 def test_validate_identical_states(self): v = HftBacktestValidator() cwm = MinimalCryptoLOBCWM() s = _state() a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) report = v.validate(cwm, [(s, (a,))]) assert report.total_steps == 1 def test_report_fields(self): v = HftBacktestValidator() cwm = MinimalCryptoLOBCWM() report = v.validate(cwm, []) assert hasattr(report, 'total_steps') assert hasattr(report, 'passed') assert hasattr(report, 'avg_price_error_bps') def test_report_passed_threshold(self): v = HftBacktestValidator(price_tolerance_bps=1.0) report = ValidationReport( total_steps=100, matching_steps=99, avg_price_error_bps=0.5, max_price_error_bps=1.0, avg_qty_error=0.0, max_qty_error=0.0, fill_match_rate=0.99, passed=True, mismatches=[], ) assert report.passed def test_report_failed_threshold(self): v = HftBacktestValidator() report = ValidationReport( total_steps=100, matching_steps=50, avg_price_error_bps=5.0, max_price_error_bps=10.0, avg_qty_error=0.0, max_qty_error=0.0, fill_match_rate=0.5, passed=False, mismatches=[], ) assert not report.passed # ══════════════════════════════════════════════════════════════════════════════ # PLANNER ALTERNATIVES (20 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestPlannerAlternatives: def test_exp3_returns_planned_policy(self): cwm = MinimalCryptoLOBCWM() p = EXP3Planner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) assert result.diagnostics["algorithm"] == "exp3" def test_exp3_probabilities_sum_to_one(self): cwm = MinimalCryptoLOBCWM() p = EXP3Planner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert abs(sum(result.probabilities) - 1.0) < 1e-6 def test_regret_matching_returns_planned_policy(self): cwm = MinimalCryptoLOBCWM() p = RegretMatchingPlanner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) assert result.diagnostics["algorithm"] == "regret_matching" def test_regret_matching_probabilities_sum_to_one(self): cwm = MinimalCryptoLOBCWM() p = RegretMatchingPlanner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert abs(sum(result.probabilities) - 1.0) < 1e-6 def test_ucb1_returns_planned_policy(self): cwm = MinimalCryptoLOBCWM() p = UCB1Planner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) assert result.diagnostics["algorithm"] == "ucb1" def test_thompson_returns_planned_policy(self): cwm = MinimalCryptoLOBCWM() p = ThompsonSamplingPlanner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) assert result.diagnostics["algorithm"] == "thompson" def test_hedge_returns_planned_policy(self): cwm = MinimalCryptoLOBCWM() p = HedgePlanner(cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) assert result.diagnostics["algorithm"] == "hedge" def test_create_planner_factory(self): cwm = MinimalCryptoLOBCWM() for name in ["exp3", "regret_matching", "ucb1", "thompson", "hedge"]: p = create_planner(name, cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(_state_with_intent(), _baseline(), budget_ms=10) assert isinstance(result, PlannedPolicy) def test_create_planner_unknown_raises(self): cwm = MinimalCryptoLOBCWM() with pytest.raises(ValueError): create_planner("unknown", cwm=cwm, counterparties=default_counterparty_ecology()) def test_planner_registry_has_all(self): assert "sm_mcts" in PLANNER_REGISTRY assert "exp3" in PLANNER_REGISTRY assert "regret_matching" in PLANNER_REGISTRY assert "ucb1" in PLANNER_REGISTRY assert "thompson" in PLANNER_REGISTRY assert "hedge" in PLANNER_REGISTRY def test_all_planners_produce_valid_output(self): cwm = MinimalCryptoLOBCWM() s = _state_with_intent() params = _baseline() for name in ["sm_mcts", "exp3", "regret_matching", "ucb1", "thompson", "hedge"]: p = create_planner(name, cwm=cwm, counterparties=default_counterparty_ecology()) result = p.plan(s, params, budget_ms=10) assert isinstance(result, PlannedPolicy) assert len(result.actions) > 0 assert abs(sum(result.probabilities) - 1.0) < 1e-6 # ══════════════════════════════════════════════════════════════════════════════ # TRAINING PARALLELISM (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestTrainingParallel: def test_monitor_record(self): m = TrainingMonitor() m.record_generation(10.0) m.record_generation(15.0) assert m.best_score == 15.0 def test_monitor_avg(self): m = TrainingMonitor() m.record_generation(10.0) m.record_generation(20.0) assert m.avg_score == 15.0 def test_monitor_improvement(self): m = TrainingMonitor() m.record_generation(10.0) m.record_generation(20.0) assert m.improvement == 10.0 def test_monitor_convergence(self): m = TrainingMonitor() for _ in range(10): m.record_generation(10.0) assert m.converged def test_monitor_not_converged(self): m = TrainingMonitor() for i in range(10): m.record_generation(float(i)) assert not m.converged def test_monitor_metrics(self): m = TrainingMonitor() m.record_generation(10.0) metrics = m.metrics() assert isinstance(metrics, TrainingMetrics) assert metrics.generations == 1 def test_monitor_best_score(self): m = TrainingMonitor() m.record_generation(5.0) m.record_generation(15.0) m.record_generation(10.0) assert m.best_score == 15.0 # ══════════════════════════════════════════════════════════════════════════════ # OBSERVABILITY (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestObservability: def test_log_decision(self): with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f: path = f.name try: logger = ObservabilityLogger(log_path=path) s = _state() a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) planned = PlannedPolicy(actions=(a,), probabilities=(1.0,), selected_action=a, diagnostics={"entropy": 0.5, "sims": 10}) decision = RiskDecision(approved=True, action=a, reason="ok") logger.log_decision(s, planned, decision, plan_ns=1000) assert logger.total_decisions == 1 with open(path) as f: lines = f.readlines() assert len(lines) == 1 record = json.loads(lines[0]) assert record["act"] == "NOOP" assert record["app"] is True finally: os.unlink(path) def test_log_multiple_decisions(self): with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f: path = f.name try: logger = ObservabilityLogger(log_path=path) s = _state() a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) planned = PlannedPolicy(actions=(a,), probabilities=(1.0,), selected_action=a, diagnostics={"entropy": 0.5, "sims": 10}) decision = RiskDecision(approved=True, action=a, reason="ok") for _ in range(5): logger.log_decision(s, planned, decision, plan_ns=1000) assert logger.total_decisions == 5 finally: os.unlink(path) def test_decision_record_fields(self): with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f: path = f.name try: logger = ObservabilityLogger(log_path=path) s = _state() a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) planned = PlannedPolicy(actions=(a,), probabilities=(1.0,), selected_action=a, diagnostics={"entropy": 0.5, "sims": 10}) decision = RiskDecision(approved=True, action=a, reason="ok") logger.log_decision(s, planned, decision, plan_ns=1000) record = logger.get_recent(1)[0] assert record.ts_ns > 0 assert record.symbol == "BTCUSDT" assert record.action_kind == "NOOP" finally: os.unlink(path) def test_get_recent(self): with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f: path = f.name try: logger = ObservabilityLogger(log_path=path) s = _state() a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) planned = PlannedPolicy(actions=(a,), probabilities=(1.0,), selected_action=a, diagnostics={"entropy": 0.5, "sims": 10}) decision = RiskDecision(approved=True, action=a, reason="ok") for _ in range(10): logger.log_decision(s, planned, decision, plan_ns=1000) recent = logger.get_recent(3) assert len(recent) == 3 finally: os.unlink(path) # ══════════════════════════════════════════════════════════════════════════════ # HOOKS (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestHooks: def test_default_hooks_noop(self): hooks = ExecutionHooks() hooks.submit_intent("i1", "BTCUSDT", 0.5) hooks.report_fill(50001.0, 0.001, "BUY", True) hooks.publish_venue_state({"test": True}) result = hooks.reconcile_state(_state(), {"test": True}) assert result == [] def test_custom_intent_hook(self): called = [] hooks = ExecutionHooks(on_intent=lambda i, t, u: called.append((i, t, u))) hooks.submit_intent("i1", "BTCUSDT", 0.5) assert called == [("i1", "BTCUSDT", 0.5)] def test_custom_fill_hook(self): called = [] hooks = ExecutionHooks(on_fill=lambda p, q, s, m: called.append((p, q, s, m))) hooks.report_fill(50001.0, 0.001, "BUY", True) assert called == [(50001.0, 0.001, "BUY", True)] def test_custom_venue_hook(self): called = [] hooks = ExecutionHooks(on_venue_state=lambda s: called.append(s)) hooks.publish_venue_state({"test": True}) assert called == [{"test": True}] def test_custom_reconcile_hook(self): hooks = ExecutionHooks(reconcile=lambda i, v: ["diff1"]) result = hooks.reconcile_state(_state(), {"test": True}) assert result == ["diff1"] def test_hooks_frozen(self): hooks = ExecutionHooks() with pytest.raises(AttributeError): hooks.on_intent = None def test_intent_hook_receives_all_params(self): received = [] hooks = ExecutionHooks(on_intent=lambda i, t, u: received.append({"id": i, "target": t, "urgency": u})) hooks.submit_intent("intent_123", "ETHUSDT", 0.8) assert received[0]["id"] == "intent_123" assert received[0]["target"] == "ETHUSDT" assert received[0]["urgency"] == 0.8 def test_fill_hook_receives_all_params(self): received = [] hooks = ExecutionHooks(on_fill=lambda p, q, s, m: received.append({"price": p, "qty": q, "side": s, "maker": m})) hooks.report_fill(50001.5, 0.005, "SELL", False) assert received[0]["price"] == 50001.5 assert received[0]["maker"] is False def test_venue_hook_receives_state(self): received = [] hooks = ExecutionHooks(on_venue_state=lambda s: received.append(s)) state = {"bid": 50000.0, "ask": 50001.0, "volume": 1000} hooks.publish_venue_state(state) assert received[0]["bid"] == 50000.0 def test_reconcile_hook_returns_diffs(self): def reconcile(internal, venue): diffs = [] if abs(internal.account.equity - venue.get("equity", 0)) > 1.0: diffs.append("equity_mismatch") return diffs hooks = ExecutionHooks(reconcile=reconcile) result = hooks.reconcile_state(_state(), {"equity": 9990.0}) assert "equity_mismatch" in result