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
This commit is contained in:
421
MALKHUT/malkhut/tests/test_planner_alternatives_and_hooks.py
Normal file
421
MALKHUT/malkhut/tests/test_planner_alternatives_and_hooks.py
Normal file
@@ -0,0 +1,421 @@
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user