422 lines
19 KiB
Python
422 lines
19 KiB
Python
|
|
"""
|
||
|
|
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
|