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:
0
MALKHUT/malkhut/tests/__init__.py
Normal file
0
MALKHUT/malkhut/tests/__init__.py
Normal file
239
MALKHUT/malkhut/tests/test_adversarial.py
Normal file
239
MALKHUT/malkhut/tests/test_adversarial.py
Normal file
@@ -0,0 +1,239 @@
|
|||||||
|
"""
|
||||||
|
Adversarial scenario tests.
|
||||||
|
|
||||||
|
These prove the core thesis: mixed policies survive diverse counterparty
|
||||||
|
ecologies better than pure deterministic quotes.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
from malkhut.planner.action_menu import build_our_actions
|
||||||
|
from malkhut.risk.gate import RiskGate
|
||||||
|
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType, PlannedPolicy
|
||||||
|
|
||||||
|
|
||||||
|
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 _params(**kw):
|
||||||
|
d = dict(
|
||||||
|
version="adv", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.10, 0.25, 0.50),
|
||||||
|
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 _state_with_intent(**kw):
|
||||||
|
from malkhut.state import TradePathState, AccountState as AC
|
||||||
|
tp = kw.get("trade_path")
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, 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=AC(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
intent=ExecutionIntent(
|
||||||
|
intent_id="adv", 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="adversarial_test",
|
||||||
|
),
|
||||||
|
trade_path=tp,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestToxicTakerPicksOffStaleQuote:
|
||||||
|
def test_pure_stale_quote_vulnerable(self):
|
||||||
|
"""A pure 'always quote best bid' is predictable and gets picked off."""
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
|
||||||
|
# Pure strategy: always place at best bid, 25% size
|
||||||
|
pure_actions = [a for a in actions if a.kind == ActionKind.PLACE and a.price_ticks_from_best == 0]
|
||||||
|
assert len(pure_actions) > 0
|
||||||
|
# This action is predictable — toxic taker can target it
|
||||||
|
|
||||||
|
def test_mixed_policy_reduces_predictability(self):
|
||||||
|
"""SM-MCTS should return a mixed distribution, not a single action."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||||||
|
)
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
result = planner.plan(root_state=state, params=params, budget_ms=15)
|
||||||
|
|
||||||
|
# Distribution should have multiple non-zero probabilities
|
||||||
|
nonzero = [p for p in result.probabilities if p > 0.01]
|
||||||
|
assert len(nonzero) >= 2, "Pure deterministic policy is exploitable"
|
||||||
|
|
||||||
|
def test_mixed_policy_includes_cancellation_option(self):
|
||||||
|
"""A good policy should have PASSIVE placement + NOOP as minimum diversity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||||||
|
)
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
result = planner.plan(root_state=state, params=params, budget_ms=15)
|
||||||
|
|
||||||
|
# Action set should include NOOP and at least one passive placement
|
||||||
|
all_kinds = set(a.kind for a in result.actions)
|
||||||
|
assert ActionKind.NOOP in all_kinds
|
||||||
|
assert ActionKind.PLACE in all_kinds
|
||||||
|
|
||||||
|
|
||||||
|
class TestRiskGateAdversarial:
|
||||||
|
def test_kill_switch_blocks_all(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
gate._kill_switch_active = lambda: True
|
||||||
|
action = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
state = _state_with_intent()
|
||||||
|
decision = gate.validate(state, planned, _params())
|
||||||
|
assert not decision.approved
|
||||||
|
assert decision.reason == "kill_switch"
|
||||||
|
|
||||||
|
def test_post_only_cross_rejected(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
state = _state_with_intent()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, -10, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(state, planned, _params())
|
||||||
|
assert not decision.approved
|
||||||
|
|
||||||
|
def test_leverage_exceeded_blocks(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
from malkhut.state import AccountState
|
||||||
|
state = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.LIVE, 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=1000.0, wallet_balance=1000.0,
|
||||||
|
available_balance=1000.0, margin_used=0.0, total_notional=5000.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
action = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(state, planned, _params())
|
||||||
|
assert not decision.approved
|
||||||
|
assert decision.reason == "leverage_limit"
|
||||||
|
|
||||||
|
|
||||||
|
class TestCounterpartyAdversarial:
|
||||||
|
def test_toxic_taker_attacks_high_toxicity(self):
|
||||||
|
"""When orderflow toxicity is high, toxic taker should cross."""
|
||||||
|
from malkhut.state import TradePathState
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=30.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=20.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.9,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
state = _state_with_intent(trade_path=tp)
|
||||||
|
toxic = ToxicTakerPolicy()
|
||||||
|
import random
|
||||||
|
action = toxic.rollout_action(state, random.Random(42))
|
||||||
|
assert action.kind == ActionKind.CROSS_SPREAD
|
||||||
|
|
||||||
|
def test_latency_arb_attacks_stale_quotes(self):
|
||||||
|
"""Latency arb crosses when cross-venue lead is strong."""
|
||||||
|
from malkhut.state import TradePathState
|
||||||
|
from malkhut.counterparties import LatencyArbPolicy
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=30.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=20.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.9,
|
||||||
|
)
|
||||||
|
state = _state_with_intent(trade_path=tp)
|
||||||
|
arb = LatencyArbPolicy()
|
||||||
|
import random
|
||||||
|
action = arb.rollout_action(state, random.Random(42))
|
||||||
|
assert action.kind == ActionKind.CROSS_SPREAD
|
||||||
|
|
||||||
|
|
||||||
|
class TestMixedPolicySurvivesEcology:
|
||||||
|
def test_noop_always_available(self):
|
||||||
|
"""NOOP must always be in the action set — sometimes the best quote is no quote."""
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
kinds = [a.kind for a in actions]
|
||||||
|
assert ActionKind.NOOP in kinds
|
||||||
|
|
||||||
|
def test_exit_available_under_tail_risk(self):
|
||||||
|
"""When path risk is high, FULL_EXIT must be available."""
|
||||||
|
from malkhut.state import TradePathState
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=-30.0, mae_bps=-60.0,
|
||||||
|
mfe_bps=5.0, distance_from_mfe_bps=35.0, distance_from_entry_bps=30.0,
|
||||||
|
time_to_mfe_s=10.0, time_in_loss_s=90.0, time_in_profit_s=10.0,
|
||||||
|
time_since_last_profit_s=80.0, time_since_deep_mae_s=5.0,
|
||||||
|
loss_to_profit_transitions=0, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=4, recovery_velocity_bps_per_s=-2.0,
|
||||||
|
adverse_velocity_bps_per_s=3.0,
|
||||||
|
dolphin_regime_score=0.2, jericho_signal_strength=0.1,
|
||||||
|
volatility_bps=30.0, orderflow_toxicity=0.7,
|
||||||
|
queue_churn_score=0.5, book_imbalance=0.3, cross_venue_lead_score=-0.5,
|
||||||
|
)
|
||||||
|
state = _state_with_intent(trade_path=tp)
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
kinds = [a.kind for a in actions]
|
||||||
|
assert ActionKind.FULL_EXIT in kinds
|
||||||
394
MALKHUT/malkhut/tests/test_asex_integration.py
Normal file
394
MALKHUT/malkhut/tests/test_asex_integration.py
Normal file
@@ -0,0 +1,394 @@
|
|||||||
|
"""
|
||||||
|
ASEx integration tests for MALKHUT.
|
||||||
|
|
||||||
|
Tests the validate-before-mutate kernel:
|
||||||
|
- GuardedFulfilmentState: book/account/intent/policy mutations
|
||||||
|
- GuardedRiskState: kill switch + risk decisions
|
||||||
|
- FulfilmentWorker: serialised engine mutations
|
||||||
|
- RiskWorker: serialised risk mutations
|
||||||
|
- FulfilmentWatch: zero-overhead ring buffer
|
||||||
|
- ShardedWorker: per-symbol partitioning
|
||||||
|
- Thread safety under concurrent access
|
||||||
|
"""
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from concurrent.futures import Future
|
||||||
|
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
from malkhut.execution.asex_integration import (
|
||||||
|
GuardedFulfilmentState,
|
||||||
|
GuardedRiskState,
|
||||||
|
FulfilmentWorker,
|
||||||
|
RiskWorker,
|
||||||
|
FulfilmentWatch,
|
||||||
|
create_sharded_fulfilment,
|
||||||
|
BookUpdate,
|
||||||
|
AccountUpdate,
|
||||||
|
IntentUpdate,
|
||||||
|
PolicyReload,
|
||||||
|
RiskCheck,
|
||||||
|
)
|
||||||
|
from asex.guarded import ValidationError
|
||||||
|
|
||||||
|
|
||||||
|
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_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _params():
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="asex_test", ucb_c=1.414, max_sims=32, max_depth=2,
|
||||||
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# GuardedFulfilmentState
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestGuardedFulfilmentState:
|
||||||
|
def test_book_update_valid(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
result = gs.mutate(BookUpdate(
|
||||||
|
ts_ns=2_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=((50100.0, 2.0),), asks=((50101.0, 1.5),),
|
||||||
|
))
|
||||||
|
assert result is not None
|
||||||
|
assert gs.state.book.best_bid == 50100.0
|
||||||
|
|
||||||
|
def test_book_update_invalid_zero_ts(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
gs.mutate(BookUpdate(ts_ns=0, symbol="BTCUSDT", bids=(), asks=()))
|
||||||
|
|
||||||
|
def test_book_update_invalid_empty_bids(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
gs.mutate(BookUpdate(ts_ns=1, symbol="BTCUSDT", bids=(), asks=((50001.0, 1.0),)))
|
||||||
|
|
||||||
|
def test_account_update_valid(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
result = gs.mutate(AccountUpdate(
|
||||||
|
ts_ns=2_000_000_000, equity=11000.0, wallet_balance=11000.0,
|
||||||
|
available_balance=11000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
positions={},
|
||||||
|
))
|
||||||
|
assert result is not None
|
||||||
|
assert gs.state.account.equity == 11000.0
|
||||||
|
|
||||||
|
def test_account_update_invalid_negative_equity(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
gs.mutate(AccountUpdate(
|
||||||
|
ts_ns=1, equity=-100.0, wallet_balance=0.0,
|
||||||
|
available_balance=0.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
positions={},
|
||||||
|
))
|
||||||
|
|
||||||
|
def test_intent_update_valid(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="t1", 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",
|
||||||
|
)
|
||||||
|
result = gs.mutate(IntentUpdate(intent=intent))
|
||||||
|
assert result is not None
|
||||||
|
assert gs.state.intent is not None
|
||||||
|
|
||||||
|
def test_intent_update_none(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
result = gs.mutate(IntentUpdate(intent=None))
|
||||||
|
assert gs.state.intent is None
|
||||||
|
|
||||||
|
def test_policy_reload(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
params = _params()
|
||||||
|
result = gs.mutate(PolicyReload(params=params))
|
||||||
|
assert gs.params is not None
|
||||||
|
assert gs.params.version == "asex_test"
|
||||||
|
|
||||||
|
def test_mutation_count_increments(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
assert gs.mutation_count == 0
|
||||||
|
gs.mutate(IntentUpdate(intent=None))
|
||||||
|
assert gs.mutation_count == 1
|
||||||
|
gs.mutate(IntentUpdate(intent=None))
|
||||||
|
assert gs.mutation_count == 2
|
||||||
|
|
||||||
|
def test_rejected_count_increments(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
assert gs.rejected == 0
|
||||||
|
try:
|
||||||
|
gs.mutate(BookUpdate(ts_ns=0, symbol="BTCUSDT", bids=(), asks=()))
|
||||||
|
except ValidationError:
|
||||||
|
pass
|
||||||
|
assert gs.rejected == 1
|
||||||
|
|
||||||
|
def test_invalid_mutation_type(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
gs.mutate("not_a_valid_mutation")
|
||||||
|
|
||||||
|
def test_sequential_mutations_preserve_state(self):
|
||||||
|
gs = GuardedFulfilmentState(_state())
|
||||||
|
gs.mutate(BookUpdate(ts_ns=2, symbol="BTCUSDT",
|
||||||
|
bids=((50100.0, 1.0),), asks=((50101.0, 1.0),)))
|
||||||
|
gs.mutate(AccountUpdate(ts_ns=3, equity=11000.0, wallet_balance=11000.0,
|
||||||
|
available_balance=11000.0, margin_used=0.0,
|
||||||
|
total_notional=0.0, positions={}))
|
||||||
|
assert gs.state.book.best_bid == 50100.0
|
||||||
|
assert gs.state.account.equity == 11000.0
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# GuardedRiskState
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestGuardedRiskState:
|
||||||
|
def test_kill_switch_activation(self):
|
||||||
|
gs = GuardedRiskState()
|
||||||
|
assert not gs.kill_switch
|
||||||
|
result = gs.mutate("KILL_SWITCH_ON")
|
||||||
|
assert gs.kill_switch
|
||||||
|
assert not result.approved
|
||||||
|
|
||||||
|
def test_kill_switch_deactivation(self):
|
||||||
|
gs = GuardedRiskState()
|
||||||
|
gs.mutate("KILL_SWITCH_ON")
|
||||||
|
gs.mutate("KILL_SWITCH_OFF")
|
||||||
|
assert not gs.kill_switch
|
||||||
|
|
||||||
|
def test_kill_switch_blocks_risk_check(self):
|
||||||
|
gs = GuardedRiskState()
|
||||||
|
gs.mutate("KILL_SWITCH_ON")
|
||||||
|
action = FulfilmentAction(ActionKind.PLACE, Side.BUY, None, 0, 0.1, 200)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
result = gs.mutate(RiskCheck(state=_state(), planned=planned, params=_params()))
|
||||||
|
assert not result.approved
|
||||||
|
assert "kill_switch" in result.reason
|
||||||
|
|
||||||
|
def test_invalid_mutation_rejected(self):
|
||||||
|
gs = GuardedRiskState()
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
gs.mutate(42) # not a valid mutation type
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# FulfilmentWorker
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestFulfilmentWorker:
|
||||||
|
def test_worker_creates_and_mutates(self):
|
||||||
|
fw = FulfilmentWorker(_state())
|
||||||
|
future = fw.update_intent(None)
|
||||||
|
assert isinstance(future, Future)
|
||||||
|
result = future.result(timeout=5.0)
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_worker_state_accessible(self):
|
||||||
|
fw = FulfilmentWorker(_state())
|
||||||
|
assert fw.state is not None
|
||||||
|
assert fw.state.book.best_bid == 50000.0
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_worker_mutation_count(self):
|
||||||
|
fw = FulfilmentWorker(_state())
|
||||||
|
fw.update_intent(None)
|
||||||
|
fw.update_intent(None)
|
||||||
|
time.sleep(0.05) # let worker process
|
||||||
|
assert fw.mutation_count >= 2
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_worker_policy_reload(self):
|
||||||
|
fw = FulfilmentWorker(_state())
|
||||||
|
params = _params()
|
||||||
|
fw.reload_policy(params)
|
||||||
|
time.sleep(0.05)
|
||||||
|
assert fw.params is not None
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_worker_thread_safety(self):
|
||||||
|
fw = FulfilmentWorker(_state())
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
def mutate_loop(idx):
|
||||||
|
try:
|
||||||
|
for i in range(10):
|
||||||
|
fw.update_intent(None)
|
||||||
|
except Exception as e:
|
||||||
|
errors.append((idx, e))
|
||||||
|
|
||||||
|
threads = [threading.Thread(target=mutate_loop, args=(i,)) for i in range(4)]
|
||||||
|
for t in threads:
|
||||||
|
t.start()
|
||||||
|
for t in threads:
|
||||||
|
t.join(timeout=5)
|
||||||
|
time.sleep(0.1)
|
||||||
|
assert len(errors) == 0
|
||||||
|
assert fw.mutation_count >= 40
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# RiskWorker
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestRiskWorker:
|
||||||
|
def test_kill_switch_activation(self):
|
||||||
|
rw = RiskWorker()
|
||||||
|
assert not rw.kill_switch
|
||||||
|
rw.activate_kill_switch()
|
||||||
|
time.sleep(0.05)
|
||||||
|
assert rw.kill_switch
|
||||||
|
rw.close()
|
||||||
|
|
||||||
|
def test_kill_switch_deactivation(self):
|
||||||
|
rw = RiskWorker()
|
||||||
|
rw.activate_kill_switch()
|
||||||
|
time.sleep(0.05)
|
||||||
|
rw.deactivate_kill_switch()
|
||||||
|
time.sleep(0.05)
|
||||||
|
assert not rw.kill_switch
|
||||||
|
rw.close()
|
||||||
|
|
||||||
|
def test_worker_thread_safety(self):
|
||||||
|
rw = RiskWorker()
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
def toggle_loop():
|
||||||
|
try:
|
||||||
|
for _ in range(5):
|
||||||
|
rw.activate_kill_switch()
|
||||||
|
rw.deactivate_kill_switch()
|
||||||
|
except Exception as e:
|
||||||
|
errors.append(e)
|
||||||
|
|
||||||
|
threads = [threading.Thread(target=toggle_loop) for _ in range(3)]
|
||||||
|
for t in threads:
|
||||||
|
t.start()
|
||||||
|
for t in threads:
|
||||||
|
t.join(timeout=5)
|
||||||
|
time.sleep(0.1)
|
||||||
|
assert len(errors) == 0
|
||||||
|
rw.close()
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# FulfilmentWatch
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestFulfilmentWatch:
|
||||||
|
def test_watch_create(self):
|
||||||
|
fw = FulfilmentWatch(_state(), capacity=64)
|
||||||
|
assert fw.state is not None
|
||||||
|
assert fw.pending == 0
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_watch_mutate_and_poll(self):
|
||||||
|
fw = FulfilmentWatch(_state(), capacity=64)
|
||||||
|
fw.mutate(IntentUpdate(intent=None))
|
||||||
|
count = fw.poll()
|
||||||
|
assert count >= 1
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
def test_watch_ring_buffer_full(self):
|
||||||
|
fw = FulfilmentWatch(_state(), capacity=4)
|
||||||
|
for _ in range(4):
|
||||||
|
fw.mutate(IntentUpdate(intent=None))
|
||||||
|
with pytest.raises(Exception): # queue.Full
|
||||||
|
fw.mutate(IntentUpdate(intent=None))
|
||||||
|
fw.close()
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# ShardedWorker
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestShardedFulfilment:
|
||||||
|
def test_sharded_create(self):
|
||||||
|
sw = create_sharded_fulfilment(n_partitions=4)
|
||||||
|
assert sw.n_partitions == 4
|
||||||
|
sw.close()
|
||||||
|
|
||||||
|
def test_sharded_different_symbols_different_partitions(self):
|
||||||
|
sw = create_sharded_fulfilment(n_partitions=4)
|
||||||
|
f1 = sw.mutate("BTCUSDT", IntentUpdate(intent=None))
|
||||||
|
f2 = sw.mutate("ETHUSDT", IntentUpdate(intent=None))
|
||||||
|
assert isinstance(f1, Future)
|
||||||
|
assert isinstance(f2, Future)
|
||||||
|
sw.close()
|
||||||
|
|
||||||
|
def test_sharded_same_symbol_same_partition(self):
|
||||||
|
sw = create_sharded_fulfilment(n_partitions=4)
|
||||||
|
p1 = sw._shard("BTCUSDT")
|
||||||
|
p2 = sw._shard("BTCUSDT")
|
||||||
|
assert p1 == p2
|
||||||
|
sw.close()
|
||||||
|
|
||||||
|
|
||||||
|
# ==============================================================================
|
||||||
|
# Engine with ASEx
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
class TestEngineASEx:
|
||||||
|
def test_engine_has_workers(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
engine = FulfilmentEngine(params_provider=_params)
|
||||||
|
assert engine.fulfilment_worker is not None
|
||||||
|
assert engine.risk_worker is not None
|
||||||
|
engine.close()
|
||||||
|
|
||||||
|
def test_engine_close_cleans_up(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
engine = FulfilmentEngine(params_provider=_params)
|
||||||
|
engine.close()
|
||||||
|
assert not engine._active
|
||||||
|
|
||||||
|
def test_engine_on_state_with_asex(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
engine = FulfilmentEngine(params_provider=_params)
|
||||||
|
engine.on_state(_state())
|
||||||
|
assert engine.fulfilment_worker.mutation_count >= 0
|
||||||
|
engine.close()
|
||||||
384
MALKHUT/malkhut/tests/test_bingx_adapter.py
Normal file
384
MALKHUT/malkhut/tests/test_bingx_adapter.py
Normal file
@@ -0,0 +1,384 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive BingX venue adapter tests.
|
||||||
|
|
||||||
|
Categories:
|
||||||
|
1. Config safety (testnet enforcement)
|
||||||
|
2. Order placement (price, qty, notional, tick/lot rounding)
|
||||||
|
3. Order cancellation (tracked, rate limit)
|
||||||
|
4. Cancel-replace
|
||||||
|
5. Risk gate integration (rejected orders)
|
||||||
|
6. Order tracking (state transitions)
|
||||||
|
7. Rate limiting
|
||||||
|
8. Zinc telemetry
|
||||||
|
9. Edge cases (zero qty, empty book, min notional)
|
||||||
|
10. Context manager
|
||||||
|
"""
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType, PlannedPolicy, RiskDecision
|
||||||
|
from malkhut.venue.bingx.adapter import BingXVenueAdapter, BingXConfig, TrackedOrder
|
||||||
|
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
d = dict(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)
|
||||||
|
d.update(kw)
|
||||||
|
return VenueRules(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(bid=50000.0, ask=50001.0, equity=10000.0, **kw):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.LIVE,
|
||||||
|
venue=kw.get("venue", _venue()),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, 1.0),), asks=(PriceLevel(ask, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=equity, wallet_balance=equity,
|
||||||
|
available_balance=equity, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _decision(approved=True, action=None, reason="approved"):
|
||||||
|
if action is None:
|
||||||
|
action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
return RiskDecision(approved=approved, action=action, reason=reason)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. CONFIG SAFETY
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestConfigSafety:
|
||||||
|
def test_testnet_default(self):
|
||||||
|
cfg = BingXConfig()
|
||||||
|
assert cfg.testnet is True
|
||||||
|
|
||||||
|
def test_testnet_rejects_live(self):
|
||||||
|
with pytest.raises(ValueError, match="testnet=False not allowed"):
|
||||||
|
BingXConfig(testnet=False)
|
||||||
|
|
||||||
|
def test_custom_config(self):
|
||||||
|
cfg = BingXConfig(api_key="k", api_secret="s", testnet=True)
|
||||||
|
assert cfg.api_key == "k"
|
||||||
|
assert cfg.recv_window_ms == 5000
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. ORDER PLACEMENT
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestOrderPlacement:
|
||||||
|
def test_place_buy_limit(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
state = _state()
|
||||||
|
adapter.execute(state, _decision(action=action))
|
||||||
|
assert adapter.total_orders == 1
|
||||||
|
|
||||||
|
def test_place_sell_limit(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.SELL, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
state = _state()
|
||||||
|
adapter.execute(state, _decision(action=action))
|
||||||
|
assert adapter.total_orders == 1
|
||||||
|
|
||||||
|
def test_place_cross_spread_market(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50,
|
||||||
|
)
|
||||||
|
state = _state()
|
||||||
|
adapter.execute(state, _decision(action=action))
|
||||||
|
assert adapter.total_orders == 1
|
||||||
|
|
||||||
|
def test_order_tracked(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
state = _state(ts=1_000_000_001)
|
||||||
|
adapter.execute(state, _decision(action=action))
|
||||||
|
working = adapter.get_working()
|
||||||
|
assert len(working) == 1
|
||||||
|
assert working[0].side == Side.BUY
|
||||||
|
assert working[0].status == "WORKING"
|
||||||
|
|
||||||
|
def test_client_id_unique(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=action))
|
||||||
|
working = adapter.get_working()
|
||||||
|
assert working[0].client_order_id != working[1].client_order_id
|
||||||
|
|
||||||
|
def test_order_below_min_notional_rejected(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
state = _state(equity=1.0) # very small equity
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(state, _decision(action=action))
|
||||||
|
assert adapter.total_orders == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. ORDER CANCELLATION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestOrderCancellation:
|
||||||
|
def test_cancel_existing_order(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
# Place first
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
working = adapter.get_working()
|
||||||
|
cancel_id = working[0].client_order_id
|
||||||
|
|
||||||
|
# Cancel
|
||||||
|
cancel_action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=cancel_id,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=cancel_action))
|
||||||
|
assert len(adapter.get_working()) == 0
|
||||||
|
|
||||||
|
def test_cancel_nonexistent_order_no_crash(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
cancel_action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id="nonexistent",
|
||||||
|
)
|
||||||
|
adapter.execute(_state(), _decision(action=cancel_action))
|
||||||
|
assert adapter.total_cancels == 0
|
||||||
|
|
||||||
|
def test_cancel_tracks_status(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
working = adapter.get_working()
|
||||||
|
cancel_id = working[0].client_order_id
|
||||||
|
|
||||||
|
cancel_action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=cancel_id,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=cancel_action))
|
||||||
|
tracked = adapter.get_tracked(cancel_id)
|
||||||
|
assert tracked.status == "CANCELLED"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. CANCEL-REPLACE
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCancelReplace:
|
||||||
|
def test_cancel_replace(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
# Place first
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
old_id = adapter.get_working()[0].client_order_id
|
||||||
|
|
||||||
|
# Cancel-replace
|
||||||
|
cr_action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL_REPLACE, Side.BUY, OrderType.LIMIT, 1, 0.1, 200,
|
||||||
|
cancel_order_id=old_id, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=cr_action))
|
||||||
|
# Old cancelled, new working
|
||||||
|
assert adapter.get_tracked(old_id).status == "CANCELLED"
|
||||||
|
working = adapter.get_working()
|
||||||
|
assert len(working) == 1
|
||||||
|
assert working[0].client_order_id != old_id
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. RISK GATE INTEGRATION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRiskGateIntegration:
|
||||||
|
def test_rejected_order_not_placed(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(), _decision(approved=False, action=action, reason="leverage"))
|
||||||
|
assert adapter.total_orders == 0
|
||||||
|
assert len(adapter.get_working()) == 0
|
||||||
|
|
||||||
|
def test_noop_not_tracked(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
adapter.execute(_state(), _decision())
|
||||||
|
assert adapter.total_orders == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 6. ORDER TRACKING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestOrderTracking:
|
||||||
|
def test_get_tracked(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
working = adapter.get_working()
|
||||||
|
tracked = adapter.get_tracked(working[0].client_order_id)
|
||||||
|
assert tracked is not None
|
||||||
|
assert tracked.symbol == "BTCUSDT"
|
||||||
|
assert tracked.price > 0
|
||||||
|
assert tracked.qty > 0
|
||||||
|
|
||||||
|
def test_get_nonexistent_returns_none(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
assert adapter.get_tracked("nonexistent") is None
|
||||||
|
|
||||||
|
def test_multiple_orders_tracked(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
for i in range(3):
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, i, 0.05, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=i + 1), _decision(action=action))
|
||||||
|
assert len(adapter.get_working()) == 3
|
||||||
|
|
||||||
|
def test_cancel_removes_from_working(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
oid = adapter.get_working()[0].client_order_id
|
||||||
|
|
||||||
|
cancel = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=oid,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=cancel))
|
||||||
|
assert len(adapter.get_working()) == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 7. RATE LIMITING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRateLimiting:
|
||||||
|
def test_cancel_rate_check(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
# Should allow many cancels within limit (90 per minute)
|
||||||
|
for i in range(90):
|
||||||
|
assert adapter._check_cancel_rate("BTCUSDT")
|
||||||
|
# 91st should fail (count=90 >= limit)
|
||||||
|
assert not adapter._check_cancel_rate("BTCUSDT")
|
||||||
|
|
||||||
|
def test_rate_resets_per_minute(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
adapter._last_minute_ts = int(time.time() / 60) - 2
|
||||||
|
adapter._cancel_count["BTCUSDT"] = 200
|
||||||
|
# Should reset because minute changed
|
||||||
|
assert adapter._check_cancel_rate("BTCUSDT")
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 8. COUNTERS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCounters:
|
||||||
|
def test_order_counter_increments(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
assert adapter.total_orders == 0
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
assert adapter.total_orders == 1
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=action))
|
||||||
|
assert adapter.total_orders == 2
|
||||||
|
|
||||||
|
def test_cancel_counter_increments(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
oid = adapter.get_working()[0].client_order_id
|
||||||
|
cancel = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=oid,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=2), _decision(action=cancel))
|
||||||
|
assert adapter.total_cancels == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 9. EDGE CASES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEdgeCases:
|
||||||
|
def test_zero_qty_not_placed(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.0, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(), _decision(action=action))
|
||||||
|
assert adapter.total_orders == 0
|
||||||
|
|
||||||
|
def test_cancel_with_no_id(self):
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=None,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(), _decision(action=action))
|
||||||
|
assert adapter.total_cancels == 0
|
||||||
|
|
||||||
|
def test_context_manager(self):
|
||||||
|
with BingXVenueAdapter() as adapter:
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200, post_only=True,
|
||||||
|
)
|
||||||
|
adapter.execute(_state(ts=1), _decision(action=action))
|
||||||
|
assert adapter.total_orders == 1
|
||||||
|
# Should be cleaned up
|
||||||
|
assert len(adapter._tracked) == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 10. TRADEDORDER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTrackedOrder:
|
||||||
|
def test_frozen(self):
|
||||||
|
o = TrackedOrder(
|
||||||
|
client_order_id="c1", venue_order_id=None, symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type="LIMIT", price=50000.0, qty=0.001,
|
||||||
|
status="WORKING", created_ts_ns=1,
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
o.status = "FILLED"
|
||||||
|
|
||||||
|
def test_defaults(self):
|
||||||
|
o = TrackedOrder(
|
||||||
|
client_order_id="c1", venue_order_id=None, symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type="LIMIT", price=50000.0, qty=0.001,
|
||||||
|
status="WORKING", created_ts_ns=1,
|
||||||
|
)
|
||||||
|
assert o.filled_qty == 0.0
|
||||||
|
assert o.filled_price == 0.0
|
||||||
|
assert o.filled_ts_ns == 0
|
||||||
229
MALKHUT/malkhut/tests/test_clock.py
Normal file
229
MALKHUT/malkhut/tests/test_clock.py
Normal file
@@ -0,0 +1,229 @@
|
|||||||
|
"""
|
||||||
|
T19 UV Clock Host tests — event dispatch, staleness, edge-triggered BarFire.
|
||||||
|
"""
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.clock.events import (
|
||||||
|
BarFire, EventProvenance, EventType, ScanEvent,
|
||||||
|
StaleInput, TickEvent, TimerEvent,
|
||||||
|
)
|
||||||
|
from malkhut.clock.host import UVClock
|
||||||
|
from malkhut.clock.staleness import StalenessWatchdog
|
||||||
|
from malkhut.clock.deadnode import DeadNodeReaper
|
||||||
|
|
||||||
|
|
||||||
|
class TestEventProvenance:
|
||||||
|
def test_fresh_when_recent(self):
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=time.time_ns(),
|
||||||
|
ingest_ts=time.time_ns(), source="live",
|
||||||
|
)
|
||||||
|
assert prov.is_fresh
|
||||||
|
|
||||||
|
def test_age_increases(self):
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=time.time_ns() - 10_000_000_000,
|
||||||
|
ingest_ts=time.time_ns() - 10_000_000_000, source="live",
|
||||||
|
)
|
||||||
|
assert prov.age_ns > 0
|
||||||
|
|
||||||
|
def test_stale_threshold(self):
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=0, ingest_ts=0, source="live",
|
||||||
|
)
|
||||||
|
threshold = prov.stale_threshold_ns(5_850_000_000)
|
||||||
|
assert threshold == 8_775_000_000 # 5.85s * 1.5
|
||||||
|
|
||||||
|
|
||||||
|
class TestScanEvent:
|
||||||
|
def test_scan_event_type(self):
|
||||||
|
e = ScanEvent(scan_number=1, symbol="BTCUSDT")
|
||||||
|
assert e.event_type == EventType.SCAN
|
||||||
|
|
||||||
|
def test_scan_event_has_provenance(self):
|
||||||
|
e = ScanEvent(scan_number=1, symbol="BTCUSDT")
|
||||||
|
assert e.provenance is not None
|
||||||
|
assert e.provenance.scan_number == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestTickEvent:
|
||||||
|
def test_tick_event_type(self):
|
||||||
|
e = TickEvent(symbol="BTCUSDT", price=50000.0, bid=49999.0, ask=50001.0)
|
||||||
|
assert e.event_type == EventType.TICK
|
||||||
|
|
||||||
|
def test_tick_event_has_provenance(self):
|
||||||
|
e = TickEvent(symbol="BTCUSDT", price=50000.0, bid=49999.0, ask=50001.0)
|
||||||
|
assert e.provenance is not None
|
||||||
|
|
||||||
|
|
||||||
|
class TestBarFire:
|
||||||
|
def test_bar_fire_type(self):
|
||||||
|
e = BarFire(scan_number=1, symbol="BTCUSDT")
|
||||||
|
assert e.event_type == EventType.BAR_FIRE
|
||||||
|
|
||||||
|
def test_bar_fire_has_provenance(self):
|
||||||
|
e = BarFire(scan_number=1, symbol="BTCUSDT")
|
||||||
|
assert e.provenance is not None
|
||||||
|
assert e.provenance.scan_number == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestTimerEvent:
|
||||||
|
def test_timer_type(self):
|
||||||
|
e = TimerEvent(timer_id="watchdog")
|
||||||
|
assert e.event_type == EventType.TIMER
|
||||||
|
|
||||||
|
|
||||||
|
class TestStaleInput:
|
||||||
|
def test_stale_type(self):
|
||||||
|
e = StaleInput(source="scan", last_scan_ts=0, current_ts=10_000_000_000)
|
||||||
|
assert e.event_type == EventType.STALE_INPUT
|
||||||
|
|
||||||
|
|
||||||
|
class TestUVClock:
|
||||||
|
def test_subscribe_and_dispatch(self):
|
||||||
|
clock = UVClock()
|
||||||
|
received = []
|
||||||
|
clock.subscribe(EventType.SCAN, lambda e: received.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {"price": 50000})
|
||||||
|
assert len(received) == 1
|
||||||
|
assert received[0].scan_number == 1
|
||||||
|
|
||||||
|
def test_bar_fire_on_scan_advance(self):
|
||||||
|
clock = UVClock()
|
||||||
|
bar_fires = []
|
||||||
|
clock.subscribe(EventType.BAR_FIRE, lambda e: bar_fires.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {})
|
||||||
|
clock.emit_scan(2, "BTCUSDT", {})
|
||||||
|
assert len(bar_fires) == 2
|
||||||
|
|
||||||
|
def test_no_bar_fire_on_duplicate_scan(self):
|
||||||
|
clock = UVClock()
|
||||||
|
bar_fires = []
|
||||||
|
clock.subscribe(EventType.BAR_FIRE, lambda e: bar_fires.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {})
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {}) # duplicate
|
||||||
|
assert len(bar_fires) == 1 # edge-triggered, not level
|
||||||
|
|
||||||
|
def test_tick_dispatch(self):
|
||||||
|
clock = UVClock()
|
||||||
|
ticks = []
|
||||||
|
clock.subscribe(EventType.TICK, lambda e: ticks.append(e))
|
||||||
|
clock.emit_tick("BTCUSDT", 50000.0, 49999.0, 50001.0)
|
||||||
|
assert len(ticks) == 1
|
||||||
|
assert ticks[0].price == 50000.0
|
||||||
|
|
||||||
|
def test_timer_dispatch(self):
|
||||||
|
clock = UVClock()
|
||||||
|
timers = []
|
||||||
|
clock.subscribe(EventType.TIMER, lambda e: timers.append(e))
|
||||||
|
clock.emit_timer("watchdog", {"check": True})
|
||||||
|
assert len(timers) == 1
|
||||||
|
|
||||||
|
def test_unsubscribe(self):
|
||||||
|
clock = UVClock()
|
||||||
|
received = []
|
||||||
|
handler = lambda e: received.append(e)
|
||||||
|
clock.subscribe(EventType.SCAN, handler)
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {})
|
||||||
|
assert len(received) == 1
|
||||||
|
clock.unsubscribe(EventType.SCAN, handler)
|
||||||
|
clock.emit_scan(2, "BTCUSDT", {})
|
||||||
|
assert len(received) == 1
|
||||||
|
|
||||||
|
def test_multiple_subscribers(self):
|
||||||
|
clock = UVClock()
|
||||||
|
r1, r2 = [], []
|
||||||
|
clock.subscribe(EventType.SCAN, lambda e: r1.append(e))
|
||||||
|
clock.subscribe(EventType.SCAN, lambda e: r2.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {})
|
||||||
|
assert len(r1) == 1
|
||||||
|
assert len(r2) == 1
|
||||||
|
|
||||||
|
def test_scan_number_tracking(self):
|
||||||
|
clock = UVClock()
|
||||||
|
clock.emit_scan(5, "BTCUSDT", {})
|
||||||
|
assert clock.last_scan_number == 5
|
||||||
|
|
||||||
|
def test_event_count(self):
|
||||||
|
clock = UVClock()
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {}) # scan + barfire = 2
|
||||||
|
clock.emit_tick("BTCUSDT", 50000.0, 49999.0, 50001.0) # tick = 1
|
||||||
|
assert clock.event_count == 3 # scan + barfire + tick
|
||||||
|
|
||||||
|
def test_staleness_check(self):
|
||||||
|
clock = UVClock(scan_cadence_ns=100_000_000) # 100ms cadence for test
|
||||||
|
# No events yet — should be stale
|
||||||
|
stale = clock.check_staleness()
|
||||||
|
# May or may not be stale depending on timing, but should not crash
|
||||||
|
|
||||||
|
def test_provenance_forwarded(self):
|
||||||
|
clock = UVClock()
|
||||||
|
received = []
|
||||||
|
clock.subscribe(EventType.SCAN, lambda e: received.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {})
|
||||||
|
assert received[0].provenance.source == "live"
|
||||||
|
|
||||||
|
def test_replay_source(self):
|
||||||
|
clock = UVClock()
|
||||||
|
received = []
|
||||||
|
clock.subscribe(EventType.SCAN, lambda e: received.append(e))
|
||||||
|
clock.emit_scan(1, "BTCUSDT", {}, source="replay")
|
||||||
|
assert received[0].provenance.source == "replay"
|
||||||
|
|
||||||
|
|
||||||
|
class TestStalenessWatchdog:
|
||||||
|
def test_heartbeat_resets_staleness(self):
|
||||||
|
w = StalenessWatchdog(cadence_ns=100_000_000)
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=time.time_ns(),
|
||||||
|
ingest_ts=time.time_ns(), source="live",
|
||||||
|
)
|
||||||
|
w.heartbeat(prov)
|
||||||
|
assert not w.is_stale
|
||||||
|
|
||||||
|
def test_stale_after_threshold(self):
|
||||||
|
w = StalenessWatchdog(cadence_ns=100) # 100ns cadence
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=time.time_ns() - 1000,
|
||||||
|
ingest_ts=time.time_ns() - 1000, source="live",
|
||||||
|
)
|
||||||
|
w.heartbeat(prov)
|
||||||
|
# Wait for staleness
|
||||||
|
import time as _time
|
||||||
|
_time.sleep(0.001)
|
||||||
|
stale = w.check()
|
||||||
|
assert stale is not None
|
||||||
|
|
||||||
|
def test_stale_count_increments(self):
|
||||||
|
w = StalenessWatchdog(cadence_ns=1)
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=0, ingest_ts=0, source="live",
|
||||||
|
)
|
||||||
|
w.heartbeat(prov)
|
||||||
|
w.check()
|
||||||
|
assert w.stale_count >= 1
|
||||||
|
|
||||||
|
def test_age_ns(self):
|
||||||
|
w = StalenessWatchdog()
|
||||||
|
prov = EventProvenance(
|
||||||
|
scan_number=1, scan_ts=time.time_ns(),
|
||||||
|
ingest_ts=time.time_ns(), source="live",
|
||||||
|
)
|
||||||
|
w.heartbeat(prov)
|
||||||
|
assert w.age_ns >= 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestDeadNodeReaper:
|
||||||
|
def test_reaper_creates(self):
|
||||||
|
r = DeadNodeReaper()
|
||||||
|
assert r.reaped_count == 0
|
||||||
|
|
||||||
|
def test_sweep_no_orphans(self):
|
||||||
|
r = DeadNodeReaper(shm_path="/tmp")
|
||||||
|
removed = r.sweep()
|
||||||
|
assert isinstance(removed, list)
|
||||||
|
|
||||||
|
def test_sweep_nonexistent_path(self):
|
||||||
|
r = DeadNodeReaper(shm_path="/nonexistent_path_xyz")
|
||||||
|
removed = r.sweep()
|
||||||
|
assert removed == []
|
||||||
122
MALKHUT/malkhut/tests/test_codec.py
Normal file
122
MALKHUT/malkhut/tests/test_codec.py
Normal file
@@ -0,0 +1,122 @@
|
|||||||
|
"""
|
||||||
|
CMA parameter codec — encode/decode roundtrip, bounds, type preservation.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCodecBounds:
|
||||||
|
def test_bounds_length_matches_specs(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
assert len(lows) == len(codec.SPECS)
|
||||||
|
assert len(highs) == len(codec.SPECS)
|
||||||
|
|
||||||
|
def test_lows_less_than_highs(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
for lo, hi in zip(lows, highs):
|
||||||
|
assert lo < hi
|
||||||
|
|
||||||
|
def test_initial_vector_midpoint(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
for i, (v, lo, hi) in enumerate(zip(x0, lows, highs)):
|
||||||
|
assert lo <= v <= hi
|
||||||
|
|
||||||
|
def test_initial_vector_length(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
assert len(x0) == len(codec.SPECS)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCodecDecode:
|
||||||
|
def test_decode_returns_fulfilment_policy_params(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "v_test")
|
||||||
|
assert isinstance(p, FulfilmentPolicyParams)
|
||||||
|
|
||||||
|
def test_version_preserved(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "my_version")
|
||||||
|
assert p.version == "my_version"
|
||||||
|
|
||||||
|
def test_int_fields_are_integers(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "int_test")
|
||||||
|
assert isinstance(p.max_depth, int)
|
||||||
|
assert isinstance(p.passive_ttl_ms, int)
|
||||||
|
assert isinstance(p.failed_recovery_cut_count, int)
|
||||||
|
|
||||||
|
def test_float_fields_are_floats(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "float_test")
|
||||||
|
assert isinstance(p.ucb_c, float)
|
||||||
|
assert isinstance(p.mae_tail_cut_bps, float)
|
||||||
|
|
||||||
|
def test_bounds_clipping_above(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
highs = [s.high for s in codec.SPECS]
|
||||||
|
x_over = [h + 10.0 for h in highs]
|
||||||
|
p = codec.decode(x_over, "over")
|
||||||
|
lows, highs_b = codec.bounds()
|
||||||
|
for i, spec in enumerate(codec.SPECS):
|
||||||
|
val = getattr(p, spec.name)
|
||||||
|
if spec.kind == "float":
|
||||||
|
assert val <= spec.high + 1e-9
|
||||||
|
|
||||||
|
def test_bounds_clipping_below(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows = [s.low for s in codec.SPECS]
|
||||||
|
x_under = [l - 10.0 for l in lows]
|
||||||
|
p = codec.decode(x_under, "under")
|
||||||
|
for i, spec in enumerate(codec.SPECS):
|
||||||
|
val = getattr(p, spec.name)
|
||||||
|
if spec.kind == "float":
|
||||||
|
assert val >= spec.low - 1e-9
|
||||||
|
|
||||||
|
def test_decode_idempotent_at_midpoint(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p1 = codec.decode(x0, "v1")
|
||||||
|
p2 = codec.decode(x0, "v2")
|
||||||
|
assert p1.ucb_c == p2.ucb_c
|
||||||
|
assert p1.max_depth == p2.max_depth
|
||||||
|
|
||||||
|
def test_different_vectors_different_params(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
x1 = list(x0)
|
||||||
|
x1[0] = x0[0] + 0.5 # ucb_c
|
||||||
|
p0 = codec.decode(x0, "a")
|
||||||
|
p1 = codec.decode(x1, "b")
|
||||||
|
assert p0.ucb_c != p1.ucb_c
|
||||||
158
MALKHUT/malkhut/tests/test_concurrency.py
Normal file
158
MALKHUT/malkhut/tests/test_concurrency.py
Normal file
@@ -0,0 +1,158 @@
|
|||||||
|
"""
|
||||||
|
Concurrency / race condition tests.
|
||||||
|
|
||||||
|
Verifies Zinc SHM IPC is safe under concurrent reader/writer access.
|
||||||
|
Uses threading to simulate real multi-process patterns.
|
||||||
|
"""
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.ipc.zinc_plane import MalkhutZincPlane
|
||||||
|
from malkhut.ipc.control_plane import MalkhutControlPlane, ControlPlaneFrame
|
||||||
|
|
||||||
|
|
||||||
|
class TestZincSHMConcurrency:
|
||||||
|
def test_writer_reader_concurrent(self):
|
||||||
|
"""Writer and reader operate concurrently without corruption."""
|
||||||
|
plane = MalkhutZincPlane(prefix="concurrent_test")
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
def writer():
|
||||||
|
try:
|
||||||
|
for i in range(20):
|
||||||
|
plane.publish_book({"seq": i, "ts": time.time_ns()})
|
||||||
|
time.sleep(0.001)
|
||||||
|
except Exception as e:
|
||||||
|
errors.append(("writer", e))
|
||||||
|
|
||||||
|
def reader():
|
||||||
|
try:
|
||||||
|
for _ in range(20):
|
||||||
|
try:
|
||||||
|
data, seq = plane.read_book(timeout_ms=50)
|
||||||
|
except Exception:
|
||||||
|
pass # timeout is acceptable
|
||||||
|
time.sleep(0.002)
|
||||||
|
except Exception as e:
|
||||||
|
errors.append(("reader", e))
|
||||||
|
|
||||||
|
t1 = threading.Thread(target=writer)
|
||||||
|
t2 = threading.Thread(target=reader)
|
||||||
|
t1.start()
|
||||||
|
t2.start()
|
||||||
|
t1.join(timeout=5)
|
||||||
|
t2.join(timeout=5)
|
||||||
|
plane.close_all()
|
||||||
|
assert len(errors) == 0, f"Errors: {errors}"
|
||||||
|
|
||||||
|
def test_multiple_readers(self):
|
||||||
|
"""Multiple readers can read from the same region without conflict."""
|
||||||
|
plane = MalkhutZincPlane(prefix="multi_reader_test")
|
||||||
|
plane.publish_book({"test": "data"})
|
||||||
|
results = []
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
def reader(idx):
|
||||||
|
try:
|
||||||
|
for _ in range(10):
|
||||||
|
try:
|
||||||
|
data, seq = plane.read_book(timeout_ms=50)
|
||||||
|
results.append((idx, seq))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
time.sleep(0.001)
|
||||||
|
except Exception as e:
|
||||||
|
errors.append((idx, e))
|
||||||
|
|
||||||
|
threads = [threading.Thread(target=reader, args=(i,)) for i in range(3)]
|
||||||
|
for t in threads:
|
||||||
|
t.start()
|
||||||
|
for t in threads:
|
||||||
|
t.join(timeout=5)
|
||||||
|
plane.close_all()
|
||||||
|
assert len(errors) == 0
|
||||||
|
assert len(results) > 0
|
||||||
|
|
||||||
|
def test_rapid_write_read_cycles(self):
|
||||||
|
"""Rapid write/read cycles don't corrupt the region."""
|
||||||
|
plane = MalkhutZincPlane(prefix="rapid_test")
|
||||||
|
for i in range(100):
|
||||||
|
plane.publish_book({"cycle": i})
|
||||||
|
data, seq = plane.read_book(timeout_ms=50)
|
||||||
|
assert data["cycle"] == i
|
||||||
|
plane.close_all()
|
||||||
|
|
||||||
|
|
||||||
|
class TestControlPlaneConcurrency:
|
||||||
|
def test_command_write_read(self):
|
||||||
|
"""Control plane commands can be written and read."""
|
||||||
|
cp = MalkhutControlPlane()
|
||||||
|
frame = ControlPlaneFrame(
|
||||||
|
command="START", ts_ns=time.time_ns(),
|
||||||
|
target_symbols=("BTCUSDT",), source="test",
|
||||||
|
)
|
||||||
|
cp.publish_command(frame)
|
||||||
|
cmd = cp.read_command(timeout_ms=50)
|
||||||
|
assert cmd is not None
|
||||||
|
assert cmd.command == "START"
|
||||||
|
assert cmd.target_symbols == ("BTCUSDT",)
|
||||||
|
cp.close()
|
||||||
|
|
||||||
|
def test_ack_write_read(self):
|
||||||
|
"""ACK frames can be written and read."""
|
||||||
|
cp = MalkhutControlPlane()
|
||||||
|
cp.publish_ack("START", time.time_ns(), "ok", "engine started")
|
||||||
|
cmd = cp.read_command(timeout_ms=50)
|
||||||
|
assert cmd is not None
|
||||||
|
assert "ACK_START" in cmd.command
|
||||||
|
cp.close()
|
||||||
|
|
||||||
|
def test_emergency_stop(self):
|
||||||
|
"""Emergency stop command is processed."""
|
||||||
|
cp = MalkhutControlPlane()
|
||||||
|
frame = ControlPlaneFrame(
|
||||||
|
command="EMERGENCY_STOP", ts_ns=time.time_ns(), source="test",
|
||||||
|
)
|
||||||
|
cp.publish_command(frame)
|
||||||
|
cmd = cp.read_command(timeout_ms=50)
|
||||||
|
assert cmd is not None
|
||||||
|
assert cmd.command == "EMERGENCY_STOP"
|
||||||
|
cp.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestZincRegionIsolation:
|
||||||
|
def test_different_prefixes_independent(self):
|
||||||
|
"""Different prefixes create independent regions."""
|
||||||
|
p1 = MalkhutZincPlane(prefix="iso_a")
|
||||||
|
p2 = MalkhutZincPlane(prefix="iso_b")
|
||||||
|
|
||||||
|
p1.publish_book({"source": "a"})
|
||||||
|
p2.publish_book({"source": "b"})
|
||||||
|
|
||||||
|
d1, _ = p1.read_book()
|
||||||
|
d2, _ = p2.read_book()
|
||||||
|
assert d1["source"] == "a"
|
||||||
|
assert d2["source"] == "b"
|
||||||
|
|
||||||
|
p1.close_all()
|
||||||
|
p2.close_all()
|
||||||
|
|
||||||
|
def test_multiple_region_types(self):
|
||||||
|
"""Different region types (book, account, fulfilment, risk) are independent."""
|
||||||
|
plane = MalkhutZincPlane(prefix="multi_region")
|
||||||
|
plane.publish_book({"type": "book"})
|
||||||
|
plane.publish_account({"type": "account"})
|
||||||
|
plane.publish_fulfilment({"type": "fulfilment"})
|
||||||
|
plane.publish_risk({"type": "risk"})
|
||||||
|
|
||||||
|
b, _ = plane.read_book()
|
||||||
|
a, _ = plane.read_account()
|
||||||
|
f, _ = plane.read_fulfilment()
|
||||||
|
r, _ = plane.read_risk()
|
||||||
|
|
||||||
|
assert b["type"] == "book"
|
||||||
|
assert a["type"] == "account"
|
||||||
|
assert f["type"] == "fulfilment"
|
||||||
|
assert r["type"] == "risk"
|
||||||
|
|
||||||
|
plane.close_all()
|
||||||
181
MALKHUT/malkhut/tests/test_counterparties.py
Normal file
181
MALKHUT/malkhut/tests/test_counterparties.py
Normal file
@@ -0,0 +1,181 @@
|
|||||||
|
"""
|
||||||
|
Counterparty ecology — adversarial agent behavior.
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel,
|
||||||
|
Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.counterparties import (
|
||||||
|
ToxicTakerPolicy, PassiveMakerPolicy, LatencyArbPolicy,
|
||||||
|
NoiseTraderPolicy, default_counterparty_ecology,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, AgentRole
|
||||||
|
|
||||||
|
|
||||||
|
def _state(**kw):
|
||||||
|
tp_kw = kw.get("trade_path_kw", {})
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=30.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=20.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=kw.get("toxicity", 0.3),
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1,
|
||||||
|
cross_venue_lead_score=kw.get("lead", 0.1),
|
||||||
|
) if kw.get("with_path", True) else None
|
||||||
|
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
trade_path=tp,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestToxicTaker:
|
||||||
|
def test_noop_when_low_toxicity(self):
|
||||||
|
p = ToxicTakerPolicy()
|
||||||
|
s = _state(toxicity=0.3)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_cross_when_high_toxicity(self):
|
||||||
|
p = ToxicTakerPolicy()
|
||||||
|
s = _state(toxicity=0.9)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
|
||||||
|
def test_always_has_legal_actions(self):
|
||||||
|
p = ToxicTakerPolicy()
|
||||||
|
s = _state()
|
||||||
|
actions = p.legal_actions(s)
|
||||||
|
assert len(actions) == 3
|
||||||
|
|
||||||
|
def test_role_is_toxic_taker(self):
|
||||||
|
p = ToxicTakerPolicy()
|
||||||
|
assert p.role == AgentRole.TOXIC_TAKER
|
||||||
|
|
||||||
|
def test_toxicity_field_set(self):
|
||||||
|
p = ToxicTakerPolicy()
|
||||||
|
s = _state(toxicity=0.9)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
if a.kind == ActionKind.CROSS_SPREAD:
|
||||||
|
assert a.toxicity > 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestPassiveMaker:
|
||||||
|
def test_always_has_legal_actions(self):
|
||||||
|
p = PassiveMakerPolicy()
|
||||||
|
s = _state()
|
||||||
|
actions = p.legal_actions(s)
|
||||||
|
assert len(actions) == 4
|
||||||
|
|
||||||
|
def test_role_is_passive_maker(self):
|
||||||
|
p = PassiveMakerPolicy()
|
||||||
|
assert p.role == AgentRole.PASSIVE_MAKER
|
||||||
|
|
||||||
|
def test_rollout_can_place(self):
|
||||||
|
p = PassiveMakerPolicy(join_probability=1.0)
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.PLACE
|
||||||
|
|
||||||
|
def test_rollout_can_noop(self):
|
||||||
|
p = PassiveMakerPolicy(join_probability=0.0)
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_deterministic_with_same_seed(self):
|
||||||
|
p = PassiveMakerPolicy()
|
||||||
|
s = _state()
|
||||||
|
a1 = p.rollout_action(s, random.Random(99))
|
||||||
|
a2 = p.rollout_action(s, random.Random(99))
|
||||||
|
assert a1.kind == a2.kind
|
||||||
|
|
||||||
|
|
||||||
|
class TestLatencyArb:
|
||||||
|
def test_noop_when_low_lead(self):
|
||||||
|
p = LatencyArbPolicy()
|
||||||
|
s = _state(lead=0.3)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_cross_when_high_lead(self):
|
||||||
|
p = LatencyArbPolicy()
|
||||||
|
s = _state(lead=0.8)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
|
||||||
|
def test_always_has_legal_actions(self):
|
||||||
|
p = LatencyArbPolicy()
|
||||||
|
s = _state()
|
||||||
|
actions = p.legal_actions(s)
|
||||||
|
assert len(actions) == 3
|
||||||
|
|
||||||
|
|
||||||
|
class TestNoiseTrader:
|
||||||
|
def test_always_has_legal_actions(self):
|
||||||
|
p = NoiseTraderPolicy()
|
||||||
|
s = _state()
|
||||||
|
actions = p.legal_actions(s)
|
||||||
|
assert len(actions) == 3
|
||||||
|
|
||||||
|
def test_role_is_noise_trader(self):
|
||||||
|
p = NoiseTraderPolicy()
|
||||||
|
assert p.role == AgentRole.NOISE_TRADER
|
||||||
|
|
||||||
|
def test_rollout_can_cross(self):
|
||||||
|
p = NoiseTraderPolicy()
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
crosses = 0
|
||||||
|
for seed in range(100):
|
||||||
|
a = p.rollout_action(s, random.Random(seed))
|
||||||
|
if a.kind == ActionKind.CROSS_SPREAD:
|
||||||
|
crosses += 1
|
||||||
|
assert crosses > 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestDefaultEcology:
|
||||||
|
def test_has_four_agents(self):
|
||||||
|
eco = default_counterparty_ecology()
|
||||||
|
assert len(eco) == 4
|
||||||
|
|
||||||
|
def test_unique_roles(self):
|
||||||
|
eco = default_counterparty_ecology()
|
||||||
|
roles = [p.role for p in eco]
|
||||||
|
assert len(set(roles)) == 4
|
||||||
|
|
||||||
|
def test_all_have_legal_actions(self):
|
||||||
|
eco = default_counterparty_ecology()
|
||||||
|
s = _state()
|
||||||
|
for p in eco:
|
||||||
|
actions = p.legal_actions(s)
|
||||||
|
assert len(actions) > 0
|
||||||
170
MALKHUT/malkhut/tests/test_cwm.py
Normal file
170
MALKHUT/malkhut/tests/test_cwm.py
Normal file
@@ -0,0 +1,170 @@
|
|||||||
|
"""
|
||||||
|
Unit tests: CWM determinism and exchange mechanics.
|
||||||
|
|
||||||
|
Mutation litmus: if we flip a comparison in CWM transition,
|
||||||
|
these tests MUST go RED.
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OpenOrderState, OrderBookState, PositionState,
|
||||||
|
PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm import MinimalCryptoLOBCWM, materialize_price_from_action
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
|
||||||
|
|
||||||
|
def _default_venue() -> VenueRules:
|
||||||
|
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 _default_book() -> OrderBookState:
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0), PriceLevel(49999.0, 2.0)),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0), PriceLevel(50002.0, 2.0)),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _default_account() -> AccountState:
|
||||||
|
return AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _default_state() -> MarketWorldState:
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=_default_venue(), book=_default_book(), account=_default_account(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _default_params() -> FulfilmentPolicyParams:
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="test_v1", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCWNDeterminism:
|
||||||
|
def test_same_state_action_seed_same_next_state(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
r1 = cwm.transition(state, (action,))
|
||||||
|
r2 = cwm.transition(state, (action,))
|
||||||
|
assert r1.ts_ns == r2.ts_ns
|
||||||
|
assert r1.book.best_bid == r2.book.best_bid
|
||||||
|
|
||||||
|
def test_transition_does_not_mutate_input_state(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
orig_ts = state.ts_ns
|
||||||
|
orig_bid = state.book.best_bid
|
||||||
|
cwm.transition(state, (action,))
|
||||||
|
assert state.ts_ns == orig_ts
|
||||||
|
assert state.book.best_bid == orig_bid
|
||||||
|
|
||||||
|
def test_noop_does_not_change_account(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
result = cwm.transition(state, (action,))
|
||||||
|
assert result.account.equity == state.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
class TestExchangeMechanics:
|
||||||
|
def test_tick_rounding_buy(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.10, 200,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
result = cwm.transition(state, (action,))
|
||||||
|
assert result.account.equity <= state.account.equity
|
||||||
|
|
||||||
|
def test_cancel_removes_open_order(self):
|
||||||
|
from malkhut.state import OpenOrderState
|
||||||
|
oo = OpenOrderState(
|
||||||
|
client_order_id="test_123", venue_order_id="v_123",
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, order_type=OrderType.POST_ONLY,
|
||||||
|
price=50000.0, qty=0.001, remaining_qty=0.001,
|
||||||
|
queue_ahead_estimate=0.001, created_ts_ns=1_000_000_000,
|
||||||
|
last_update_ts_ns=1_000_000_000, post_only=True,
|
||||||
|
)
|
||||||
|
state = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=_default_venue(), book=_default_book(),
|
||||||
|
account=_default_account(), open_orders=(oo,),
|
||||||
|
)
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0,
|
||||||
|
cancel_order_id="test_123",
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
result = cwm.transition(state, (action,))
|
||||||
|
assert len(result.open_orders) == 0
|
||||||
|
|
||||||
|
def test_post_only_rejects_crossing_buy(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _default_state()
|
||||||
|
# Post-only buy at best_ask should be rejected
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, -1, 0.10, 200,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
result = cwm.transition(state, (action,))
|
||||||
|
# No fill should occur, order should not be in book at crossing price
|
||||||
|
assert result.account.equity == state.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
class TestPriceMaterialization:
|
||||||
|
def test_cross_spread_buy_returns_best_ask(self):
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.10, 50,
|
||||||
|
)
|
||||||
|
price = materialize_price_from_action(state, action)
|
||||||
|
assert price == 50001.0
|
||||||
|
|
||||||
|
def test_cross_spread_sell_returns_best_bid(self):
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.CROSS_SPREAD, Side.SELL, OrderType.IOC, 0, 0.10, 50,
|
||||||
|
)
|
||||||
|
price = materialize_price_from_action(state, action)
|
||||||
|
assert price == 50000.0
|
||||||
|
|
||||||
|
def test_buy_offset_0_returns_best_bid(self):
|
||||||
|
state = _default_state()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.10, 200,
|
||||||
|
)
|
||||||
|
price = materialize_price_from_action(state, action)
|
||||||
|
assert price == 50000.0
|
||||||
282
MALKHUT/malkhut/tests/test_cwm_core.py
Normal file
282
MALKHUT/malkhut/tests/test_cwm_core.py
Normal file
@@ -0,0 +1,282 @@
|
|||||||
|
"""
|
||||||
|
CWM determinism and transition correctness.
|
||||||
|
|
||||||
|
Mutation litmus:
|
||||||
|
- Flip bid/ask comparison in CWM → determinism test must FAIL
|
||||||
|
- Remove tick rounding → price test must FAIL
|
||||||
|
- Swap maker/taker fee → reward test must FAIL
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, IntentKind, MarketWorldState,
|
||||||
|
Mode, OpenOrderState, OrderBookState, PositionState, PriceLevel,
|
||||||
|
Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM, materialize_price_from_action
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType, CounterpartyAction, AgentRole
|
||||||
|
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
defaults = dict(
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
defaults.update(kw)
|
||||||
|
return VenueRules(**defaults)
|
||||||
|
|
||||||
|
|
||||||
|
def _book(bid=50000.0, ask=50001.0, bid_qty=1.0, ask_qty=1.0, **kw):
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, bid_qty), PriceLevel(bid - 0.1, 2.0)),
|
||||||
|
asks=(PriceLevel(ask, ask_qty), PriceLevel(ask + 0.1, 2.0)),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _account(equity=10000.0, **kw):
|
||||||
|
return AccountState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), equity=equity,
|
||||||
|
wallet_balance=equity, available_balance=equity,
|
||||||
|
margin_used=0.0, total_notional=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(bid=50000.0, ask=50001.0, equity=10000.0, **kw):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=kw.get("venue", _venue()), book=_book(bid, ask),
|
||||||
|
account=_account(equity), open_orders=kw.get("open_orders", ()),
|
||||||
|
trade_path=kw.get("trade_path"), intent=kw.get("intent"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCWMDeterminism:
|
||||||
|
def test_same_input_same_output(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r1 = cwm.transition(s, (a,))
|
||||||
|
r2 = cwm.transition(s, (a,))
|
||||||
|
assert r1.ts_ns == r2.ts_ns
|
||||||
|
assert r1.book.bids[0].price == r2.book.bids[0].price
|
||||||
|
|
||||||
|
def test_input_not_mutated(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
orig_ts = s.ts_ns
|
||||||
|
orig_bid = s.book.best_bid
|
||||||
|
cwm.transition(s, (_noop(),))
|
||||||
|
assert s.ts_ns == orig_ts
|
||||||
|
assert s.book.best_bid == orig_bid
|
||||||
|
|
||||||
|
def test_timestamp_advances(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r = cwm.transition(s, (_noop(),))
|
||||||
|
assert r.ts_ns > s.ts_ns
|
||||||
|
|
||||||
|
def test_noop_preserves_equity(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r = cwm.transition(s, (_noop(),))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_determinism_across_seeds(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s1 = _state()
|
||||||
|
s2 = _state()
|
||||||
|
a = _noop()
|
||||||
|
r1 = cwm.transition(s1, (a,))
|
||||||
|
r2 = cwm.transition(s2, (a,))
|
||||||
|
assert r1.ts_ns == r2.ts_ns
|
||||||
|
|
||||||
|
def test_different_states_different_outputs(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s_fast = _state(bid=50000.0, ask=50001.0)
|
||||||
|
s_wide = _state(bid=49000.0, ask=51000.0)
|
||||||
|
a = _noop()
|
||||||
|
r1 = cwm.transition(s_fast, (a,))
|
||||||
|
r2 = cwm.transition(s_wide, (a,))
|
||||||
|
assert r1.book.spread != r2.book.spread
|
||||||
|
|
||||||
|
|
||||||
|
class TestPriceMaterialization:
|
||||||
|
def test_cross_spread_buy_returns_best_ask(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
assert materialize_price_from_action(s, a) == 50001.0
|
||||||
|
|
||||||
|
def test_cross_spread_sell_returns_best_bid(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.SELL, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
assert materialize_price_from_action(s, a) == 50000.0
|
||||||
|
|
||||||
|
def test_buy_offset_0_returns_best_bid(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 50000.0
|
||||||
|
|
||||||
|
def test_buy_offset_1_one_tick_behind(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 1, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 49999.9
|
||||||
|
|
||||||
|
def test_sell_offset_0_returns_best_ask(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.SELL, OrderType.POST_ONLY, 0, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 50001.0
|
||||||
|
|
||||||
|
def test_sell_offset_1_one_tick_behind(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.SELL, OrderType.POST_ONLY, 1, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 50001.1
|
||||||
|
|
||||||
|
def test_none_side_returns_none(self):
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
assert materialize_price_from_action(s, a) is None
|
||||||
|
|
||||||
|
def test_wide_spread_offsets(self):
|
||||||
|
s = _state(bid=49000.0, ask=51000.0)
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 5, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 48999.5
|
||||||
|
|
||||||
|
def test_tight_spread_one_tick(self):
|
||||||
|
s = _state(bid=50000.0, ask=50000.1)
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.1, 200)
|
||||||
|
assert materialize_price_from_action(s, a) == 50000.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestRewardFunction:
|
||||||
|
def _params(self, **kw):
|
||||||
|
defaults = dict(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
defaults.update(kw)
|
||||||
|
return FulfilmentPolicyParams(**defaults)
|
||||||
|
|
||||||
|
def test_noop_reward_zero_path(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
reward = cwm.reward(s, a, r, self._params())
|
||||||
|
assert reward == 0.0
|
||||||
|
|
||||||
|
def test_maker_fill_positive_fee_reward(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
params = self._params(w_fee_quality=1.0)
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 1, 0.10, 200, post_only=True,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
reward = cwm.reward(s, a, r, params)
|
||||||
|
# Maker fee is -0.2 bps, so reward should be positive
|
||||||
|
assert reward > 0
|
||||||
|
|
||||||
|
def test_cross_spread_penalty(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
reward = cwm.reward(s, a, r, self._params())
|
||||||
|
assert reward < 0
|
||||||
|
|
||||||
|
def test_higher_tail_loss_weight_more_penalty(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
params_low = self._params(w_tail_loss=1.0)
|
||||||
|
params_high = self._params(w_tail_loss=10.0)
|
||||||
|
# Both should be 0 for noop with no trade path
|
||||||
|
assert cwm.reward(s, a, r, params_low) == 0.0
|
||||||
|
assert cwm.reward(s, a, r, params_high) == 0.0
|
||||||
|
|
||||||
|
def test_inventory_risk_calculation(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
# State with position
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state()
|
||||||
|
s_new = MarketWorldState(
|
||||||
|
ts_ns=s.ts_ns, mode=s.mode, venue=s.venue, book=s.book,
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=s.account.ts_ns, equity=s.account.equity,
|
||||||
|
wallet_balance=s.account.wallet_balance,
|
||||||
|
available_balance=s.account.available_balance,
|
||||||
|
margin_used=s.account.margin_used,
|
||||||
|
total_notional=abs(0.1 * 50000.0),
|
||||||
|
positions={"BTCUSDT": pos},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
risk = cwm._inventory_risk(s_new)
|
||||||
|
assert 0.0 < risk < 1.0
|
||||||
|
|
||||||
|
def test_tail_risk_proxy_increases_with_mae(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=1_000_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=-20.0, mae_bps=-30.0,
|
||||||
|
mfe_bps=5.0, distance_from_mfe_bps=25.0, distance_from_entry_bps=20.0,
|
||||||
|
time_to_mfe_s=30.0, time_in_loss_s=80.0, time_in_profit_s=20.0,
|
||||||
|
time_since_last_profit_s=60.0, time_since_deep_mae_s=5.0,
|
||||||
|
loss_to_profit_transitions=2, deep_loss_recoveries=1,
|
||||||
|
failed_recovery_count=2, recovery_velocity_bps_per_s=-1.0,
|
||||||
|
adverse_velocity_bps_per_s=2.0,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.4,
|
||||||
|
queue_churn_score=0.3, book_imbalance=0.1, cross_venue_lead_score=0.2,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
risk = cwm._tail_risk_proxy(s)
|
||||||
|
assert risk > 0
|
||||||
|
|
||||||
|
def test_terminal_at_depth_zero(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
assert cwm.terminal(s, 0)
|
||||||
|
|
||||||
|
def test_terminal_when_no_intent(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
assert cwm.terminal(s, 5)
|
||||||
|
|
||||||
|
def test_not_terminal_with_intent_and_depth(self):
|
||||||
|
from malkhut.state import ExecutionIntent
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="t1", 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",
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(intent=intent)
|
||||||
|
assert not cwm.terminal(s, 3)
|
||||||
924
MALKHUT/malkhut/tests/test_cwm_exhaustive.py
Normal file
924
MALKHUT/malkhut/tests/test_cwm_exhaustive.py
Normal file
@@ -0,0 +1,924 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive CWM tests — every exchange mechanic, every edge case.
|
||||||
|
|
||||||
|
Test categories:
|
||||||
|
1. Tick/lot rounding
|
||||||
|
2. Price-time priority + sequential level consumption
|
||||||
|
3. Partial fills across multiple levels
|
||||||
|
4. Post-only rejection (buy crosses ask, sell crosses bid)
|
||||||
|
5. CROSS_SPREAD immediate fill
|
||||||
|
6. Cancel order
|
||||||
|
7. Cancel-replace
|
||||||
|
8. Fee application (maker vs taker)
|
||||||
|
9. Position update (open, add, reduce, close)
|
||||||
|
10. Mark-to-market
|
||||||
|
11. Realized PnL on sell
|
||||||
|
12. Available balance deduction
|
||||||
|
13. Path-state update (entry, MAE, MFE, recovery)
|
||||||
|
14. Counterparty fills consuming book levels
|
||||||
|
15. Empty book handling
|
||||||
|
16. Determinism (same input = same output)
|
||||||
|
17. Input immutability
|
||||||
|
18. Timestamp advancement
|
||||||
|
19. Edge cases (zero qty, zero price, negative equity)
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OpenOrderState, OrderBookState, PositionState,
|
||||||
|
PriceLevel, Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import (
|
||||||
|
MinimalCryptoLOBCWM, materialize_price_from_action,
|
||||||
|
_round_tick, _round_lot, _clip_lots, _fill_from_levels,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, CounterpartyAction, AgentRole, FulfilmentAction, OrderType
|
||||||
|
|
||||||
|
|
||||||
|
# ── Helpers ──────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
d = dict(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)
|
||||||
|
d.update(kw)
|
||||||
|
return VenueRules(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _book(bid=50000.0, ask=50001.0, bid_qty=1.0, ask_qty=1.0, ts=1_000_000_000, **kw):
|
||||||
|
bids = kw.get("bids", ((bid, bid_qty),))
|
||||||
|
asks = kw.get("asks", ((ask, ask_qty),))
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=ts, symbol="BTCUSDT",
|
||||||
|
bids=tuple(PriceLevel(p, q) for p, q in bids),
|
||||||
|
asks=tuple(PriceLevel(p, q) for p, q in asks),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _account(equity=10000.0, **kw):
|
||||||
|
return AccountState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), equity=equity,
|
||||||
|
wallet_balance=kw.get("wallet", equity),
|
||||||
|
available_balance=kw.get("available", equity),
|
||||||
|
margin_used=kw.get("margin", 0.0),
|
||||||
|
total_notional=kw.get("notional", 0.0),
|
||||||
|
positions=kw.get("positions", {}),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(bid=50000.0, ask=50001.0, equity=10000.0, **kw):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=kw.get("venue", _venue()),
|
||||||
|
book=_book(bid, ask, bid_qty=kw.get("bid_qty", 1.0), ask_qty=kw.get("ask_qty", 1.0)),
|
||||||
|
account=_account(equity, positions=kw.get("positions", {})),
|
||||||
|
open_orders=kw.get("open_orders", ()),
|
||||||
|
trade_path=kw.get("trade_path"), intent=kw.get("intent"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _place(side, offset=0, frac=0.1, post_only=False, reduce_only=False):
|
||||||
|
return FulfilmentAction(
|
||||||
|
ActionKind.PLACE, side,
|
||||||
|
OrderType.POST_ONLY if post_only else OrderType.LIMIT,
|
||||||
|
offset, frac, 200,
|
||||||
|
post_only=post_only, reduce_only=reduce_only,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _cross(side, frac=0.1):
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
|
|
||||||
|
|
||||||
|
def _cancel(order_id):
|
||||||
|
return FulfilmentAction(ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id=order_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _oo(cid="c1", price=50000.0, qty=0.001, side=Side.BUY, ts=1_000_000_000):
|
||||||
|
return OpenOrderState(
|
||||||
|
client_order_id=cid, venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=side, order_type=OrderType.POST_ONLY, price=price,
|
||||||
|
qty=qty, remaining_qty=qty, queue_ahead_estimate=qty * 0.5,
|
||||||
|
created_ts_ns=ts, last_update_ts_ns=ts, post_only=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tp(side=Side.BUY, pnl=0.0, mae=-10.0, mfe=5.0, ts=1_000_000_000,
|
||||||
|
failed_recovery_count=0):
|
||||||
|
return TradePathState(
|
||||||
|
symbol="BTCUSDT", side=side, entry_ts_ns=ts, now_ts_ns=ts,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=pnl, mae_bps=mae,
|
||||||
|
mfe_bps=mfe, distance_from_mfe_bps=mfe - pnl,
|
||||||
|
distance_from_entry_bps=abs(pnl), time_to_mfe_s=20.0,
|
||||||
|
time_in_loss_s=30.0, time_in_profit_s=20.0,
|
||||||
|
time_since_last_profit_s=5.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=failed_recovery_count,
|
||||||
|
recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. TICK / LOT ROUNDING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTickRounding:
|
||||||
|
def test_round_tick_exact(self):
|
||||||
|
assert _round_tick(50000.0, 0.1) == 50000.0
|
||||||
|
|
||||||
|
def test_round_tick_up(self):
|
||||||
|
assert _round_tick(50000.06, 0.1) == pytest.approx(50000.1, abs=1e-9)
|
||||||
|
|
||||||
|
def test_round_tick_down(self):
|
||||||
|
assert _round_tick(50000.04, 0.1) == 50000.0
|
||||||
|
|
||||||
|
def test_round_tick_tiny_tick(self):
|
||||||
|
assert _round_tick(50000.055, 0.01) == 50000.06
|
||||||
|
|
||||||
|
def test_round_tick_large_tick(self):
|
||||||
|
assert _round_tick(50005.0, 1.0) == 50005.0
|
||||||
|
|
||||||
|
def test_round_tick_large_tick_rounds_down(self):
|
||||||
|
# round(50004.9 / 1.0) = round(50004.9) = 50005 (banker's rounds to even)
|
||||||
|
assert _round_tick(50004.4, 1.0) == 50004.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestLotRounding:
|
||||||
|
def test_round_lot_exact(self):
|
||||||
|
assert _round_lot(0.001, 0.001) == 0.001
|
||||||
|
|
||||||
|
def test_round_lot_up(self):
|
||||||
|
assert _round_lot(0.0015, 0.001) == 0.002
|
||||||
|
|
||||||
|
def test_round_lot_down(self):
|
||||||
|
assert _round_lot(0.0014, 0.001) == 0.001
|
||||||
|
|
||||||
|
def test_round_lot_large_lot(self):
|
||||||
|
assert _round_lot(1.5, 1.0) == 2.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestClipLots:
|
||||||
|
def test_clip_above_min(self):
|
||||||
|
assert _clip_lots(0.005, 0.001, 0.001) == 0.005
|
||||||
|
|
||||||
|
def test_clip_below_min_returns_zero(self):
|
||||||
|
assert _clip_lots(0.0005, 0.001, 0.001) == 0.0
|
||||||
|
|
||||||
|
def test_clip_exact_min(self):
|
||||||
|
assert _clip_lots(0.001, 0.001, 0.001) == 0.001
|
||||||
|
|
||||||
|
def test_clip_rounds_to_lot(self):
|
||||||
|
assert _clip_lots(0.0017, 0.001, 0.001) == 0.002
|
||||||
|
|
||||||
|
def test_clip_zero_qty(self):
|
||||||
|
assert _clip_lots(0.0, 0.001, 0.001) == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. FILL FROM LEVELS (price-time priority)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestFillFromLevels:
|
||||||
|
def test_fill_single_level_full(self):
|
||||||
|
levels = [PriceLevel(50000.0, 1.0)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.5, 0.001, 0.001)
|
||||||
|
assert filled == 0.5
|
||||||
|
assert avg == 50000.0
|
||||||
|
assert len(remaining) == 1
|
||||||
|
assert remaining[0].qty == 0.5
|
||||||
|
|
||||||
|
def test_fill_single_level_exact(self):
|
||||||
|
levels = [PriceLevel(50000.0, 1.0)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 1.0, 0.001, 0.001)
|
||||||
|
assert filled == 1.0
|
||||||
|
assert len(remaining) == 0
|
||||||
|
|
||||||
|
def test_fill_multi_level(self):
|
||||||
|
levels = [PriceLevel(50000.0, 0.5), PriceLevel(50001.0, 0.5)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.8, 0.001, 0.001)
|
||||||
|
assert filled == 0.8
|
||||||
|
assert abs(avg - (50000.0 * 0.5 + 50001.0 * 0.3) / 0.8) < 0.01
|
||||||
|
assert len(remaining) == 1
|
||||||
|
assert remaining[0].price == 50001.0
|
||||||
|
assert remaining[0].qty == 0.2
|
||||||
|
|
||||||
|
def test_fill_exhausts_all_levels(self):
|
||||||
|
levels = [PriceLevel(50000.0, 0.3), PriceLevel(50001.0, 0.3)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 1.0, 0.001, 0.001)
|
||||||
|
assert filled == 0.6
|
||||||
|
assert len(remaining) == 0
|
||||||
|
|
||||||
|
def test_fill_empty_levels(self):
|
||||||
|
filled, avg, remaining = _fill_from_levels([], 1.0, 0.001, 0.001)
|
||||||
|
assert filled == 0.0
|
||||||
|
assert remaining == []
|
||||||
|
|
||||||
|
def test_fill_preserves_price_order(self):
|
||||||
|
levels = [PriceLevel(50001.0, 0.5), PriceLevel(50000.0, 0.5)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.3, 0.001, 0.001)
|
||||||
|
# Should fill from 50001.0 first (first in list = highest priority)
|
||||||
|
assert avg == 50001.0
|
||||||
|
|
||||||
|
def test_fill_lot_rounding(self):
|
||||||
|
levels = [PriceLevel(50000.0, 1.0)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.555, 0.1, 0.1)
|
||||||
|
assert filled == pytest.approx(0.6, abs=0.01) # rounded to 0.1 lot
|
||||||
|
|
||||||
|
def test_fill_below_min_qty(self):
|
||||||
|
levels = [PriceLevel(50000.0, 1.0)]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.0005, 0.001, 0.001)
|
||||||
|
assert filled == 0.0
|
||||||
|
|
||||||
|
def test_fill_three_levels(self):
|
||||||
|
levels = [
|
||||||
|
PriceLevel(50000.0, 0.1),
|
||||||
|
PriceLevel(50001.0, 0.1),
|
||||||
|
PriceLevel(50002.0, 0.1),
|
||||||
|
]
|
||||||
|
filled, avg, remaining = _fill_from_levels(levels, 0.25, 0.001, 0.001)
|
||||||
|
assert filled == 0.25
|
||||||
|
assert avg == (50000.0 * 0.1 + 50001.0 * 0.1 + 50002.0 * 0.05) / 0.25
|
||||||
|
assert len(remaining) == 1
|
||||||
|
assert remaining[0].price == 50002.0
|
||||||
|
assert remaining[0].qty == pytest.approx(0.05, abs=0.001)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. POST-ONLY REJECTION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPostOnlyRejection:
|
||||||
|
def test_buy_at_ask_rejected(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=-10, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
assert len(r.open_orders) == len(s.open_orders)
|
||||||
|
|
||||||
|
def test_sell_at_bid_rejected(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.SELL, offset=-10, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_buy_inside_spread_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=0, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.BUY for o in r.open_orders)
|
||||||
|
|
||||||
|
def test_sell_inside_spread_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.SELL, offset=0, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.SELL for o in r.open_orders)
|
||||||
|
|
||||||
|
def test_buy_one_tick_below_ask_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid=50000.0, ask=50001.0)
|
||||||
|
# price = 50000.0 - (-9)*0.1 = 50000.9 < 50001.0
|
||||||
|
a = _place(Side.BUY, offset=-9, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.BUY for o in r.open_orders)
|
||||||
|
|
||||||
|
def test_sell_one_tick_above_bid_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid=50000.0, ask=50001.0)
|
||||||
|
# price = 50001.0 + (-9)*0.1 = 50000.1 > 50000.0
|
||||||
|
a = _place(Side.SELL, offset=-9, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.SELL for o in r.open_orders)
|
||||||
|
|
||||||
|
def test_wide_spread_allows_more_offsets(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid=49000.0, ask=51000.0)
|
||||||
|
# price = 49000.0 - (-10)*0.1 = 49001.0 < 51000.0
|
||||||
|
a = _place(Side.BUY, offset=-10, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.BUY for o in r.open_orders)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. CROSS_SPREAD (immediate fill)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCrossSpread:
|
||||||
|
def test_cross_buy_fills(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity < s.account.equity
|
||||||
|
|
||||||
|
def test_cross_sell_fills(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.SELL, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity <= s.account.equity
|
||||||
|
|
||||||
|
def test_cross_buy_updates_book(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(ask=50001.0, ask_qty=1.0)
|
||||||
|
a = _cross(Side.BUY, frac=0.5)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Ask should be reduced
|
||||||
|
total_ask_qty = sum(l.qty for l in r.book.asks)
|
||||||
|
assert total_ask_qty < 1.0
|
||||||
|
|
||||||
|
def test_cross_buy_fills_at_best_ask(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.book.last_trade_price == 50001.0
|
||||||
|
|
||||||
|
def test_cross_sell_fills_at_best_bid(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.SELL, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.book.last_trade_price == 50000.0
|
||||||
|
|
||||||
|
def test_cross_partial_fill(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(ask_qty=0.002)
|
||||||
|
a = _cross(Side.BUY, frac=0.5) # wants more than available
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Should fill what's available
|
||||||
|
assert r.account.equity < s.account.equity
|
||||||
|
|
||||||
|
def test_cross_consumes_levels_sequentially(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(asks=((50001.0, 0.1), (50002.0, 0.1)))
|
||||||
|
a = _cross(Side.BUY, frac=0.5)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Should consume from 50001 first, then 50002
|
||||||
|
assert r.book.last_trade_price <= 50002.0
|
||||||
|
|
||||||
|
def test_cross_creates_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos is not None
|
||||||
|
assert pos.qty > 0
|
||||||
|
|
||||||
|
def test_cross_no_fill_when_zero_qty(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.0)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_cross_updates_last_trade(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.book.last_trade_side == Side.BUY
|
||||||
|
assert r.book.last_trade_qty > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. CANCEL ORDER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCancelOrder:
|
||||||
|
def test_cancel_removes_order(self):
|
||||||
|
oo = _oo("c1")
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
a = _cancel("c1")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 0
|
||||||
|
|
||||||
|
def test_cancel_wrong_id_keeps_order(self):
|
||||||
|
oo = _oo("c1")
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
a = _cancel("wrong")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
|
||||||
|
def test_cancel_only_one_order(self):
|
||||||
|
oo1 = _oo("c1")
|
||||||
|
oo2 = _oo("c2")
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo1, oo2))
|
||||||
|
a = _cancel("c1")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
assert r.open_orders[0].client_order_id == "c2"
|
||||||
|
|
||||||
|
def test_cancel_nonexistent_id(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(_oo("c1"),))
|
||||||
|
a = _cancel("nonexistent")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
|
||||||
|
def test_cancel_empty_book(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cancel("c1")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 6. PASSIVE PLACEMENT
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPassivePlacement:
|
||||||
|
def test_passive_buy_added(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
assert r.open_orders[0].side == Side.BUY
|
||||||
|
|
||||||
|
def test_passive_sell_added(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.SELL, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
assert r.open_orders[0].side == Side.SELL
|
||||||
|
|
||||||
|
def test_passive_price_correct(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.open_orders[0].price == 49999.9
|
||||||
|
|
||||||
|
def test_passive_qty_correct(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
expected_qty = 0.1 * 10000.0 / 49999.9
|
||||||
|
assert r.open_orders[0].qty > 0
|
||||||
|
|
||||||
|
def test_passive_order_id_unique(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a1 = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r1 = cwm.transition(s, (a1,))
|
||||||
|
# Use r1 as input (different ts_ns) for second order
|
||||||
|
a2 = _place(Side.BUY, offset=2, frac=0.1, post_only=True)
|
||||||
|
r2 = cwm.transition(r1, (a2,))
|
||||||
|
assert r1.open_orders[0].client_order_id != r2.open_orders[-1].client_order_id
|
||||||
|
|
||||||
|
def test_multiple_passive_orders(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
a2 = _place(Side.BUY, offset=2, frac=0.1, post_only=True)
|
||||||
|
r2 = cwm.transition(r, (a2,))
|
||||||
|
assert len(r2.open_orders) == 2
|
||||||
|
|
||||||
|
def test_passive_no_position_change(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _place(Side.BUY, offset=1, frac=0.1, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos is None or pos.qty == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 7. FEES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestFeeApplication:
|
||||||
|
def test_taker_fee_reduces_equity(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
fee = 0.001 * 50001.0 * 0.5 / 10_000 # taker fee
|
||||||
|
assert r.account.equity < s.account.equity
|
||||||
|
|
||||||
|
def test_maker_fee_rebate(self):
|
||||||
|
"""Counterparty fill should not charge us taker fees."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.1, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
# CP fill touches book but doesn't go through our fee path
|
||||||
|
# available_balance should be reduced (position opened via maker fill)
|
||||||
|
assert r.account.available_balance <= s.account.available_balance
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 8. POSITION UPDATE
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPositionUpdate:
|
||||||
|
def test_open_long_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos is not None
|
||||||
|
assert pos.qty > 0
|
||||||
|
assert pos.side == Side.BUY
|
||||||
|
|
||||||
|
def test_open_short_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.SELL, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos is not None
|
||||||
|
assert pos.qty < 0
|
||||||
|
assert pos.side == Side.SELL
|
||||||
|
|
||||||
|
def test_add_to_long(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.01, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.05, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos.qty > 0.01
|
||||||
|
|
||||||
|
def test_reduce_long(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.SELL, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos.qty < 0.1
|
||||||
|
|
||||||
|
def test_avg_entry_updates_on_add(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.01, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.05, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(bid=49000.0, ask=49001.0, positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos.avg_entry != 50000.0
|
||||||
|
|
||||||
|
def test_realized_pnl_on_reduce(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(bid=51000.0, ask=51001.0, positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.SELL, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos.realized_pnl > 0
|
||||||
|
|
||||||
|
def test_no_position_no_change(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert "BTCUSDT" not in r.account.positions
|
||||||
|
|
||||||
|
def test_leverage_calculation(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.5)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos.leverage > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 9. MARK-TO-MARKET
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestMarkToMarket:
|
||||||
|
def test_mtM_updates_on_fill(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
# Large book so CP doesn't empty it
|
||||||
|
s = _state(bid=51000.0, ask=51001.0, ask_qty=10.0, positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.BUY, frac=0.5)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos is not None
|
||||||
|
assert new_pos.qty > 0.1
|
||||||
|
|
||||||
|
def test_mtM_equity_changes_on_fill(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(bid=51000.0, ask=51001.0, ask_qty=10.0, positions={"BTCUSDT": pos})
|
||||||
|
a = _cross(Side.BUY, frac=0.5)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity != s.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 10. PATH-STATE UPDATE
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPathStateUpdate:
|
||||||
|
def test_new_position_creates_path(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.trade_path is not None
|
||||||
|
assert r.trade_path.side == Side.BUY
|
||||||
|
|
||||||
|
def test_path_entry_timestamp(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.trade_path.entry_ts_ns == r.ts_ns
|
||||||
|
|
||||||
|
def test_path_pnl_updates(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = _tp(side=Side.BUY, pnl=0.0, mae=-5.0, mfe=10.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.trade_path is not None
|
||||||
|
|
||||||
|
def test_path_mae_tracking(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = _tp(side=Side.BUY, mae=-20.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.trade_path.mae_bps <= -20.0
|
||||||
|
|
||||||
|
def test_path_mfe_tracking(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = _tp(side=Side.BUY, mfe=15.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.trade_path.mfe_bps >= 15.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 11. COUNTERPARTY FILLS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCounterpartyFills:
|
||||||
|
def test_cp_buy_consumes_asks(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(asks=((50001.0, 0.5),))
|
||||||
|
# fraction=5.0 means cp wants to buy 5.0 * 10000 / 50001 = ~1.0 units
|
||||||
|
# Should consume all 0.5 from top level
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 5.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
total_ask = sum(l.qty for l in r.book.asks)
|
||||||
|
assert total_ask < 0.5 # consumed from top level
|
||||||
|
|
||||||
|
def test_cp_sell_consumes_bids(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bids=((50000.0, 0.5),))
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.SELL, 0, 5.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
total_bid = sum(l.qty for l in r.book.bids)
|
||||||
|
assert total_bid < 0.5 # consumed from top level
|
||||||
|
|
||||||
|
def test_cp_fill_updates_book(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(asks=((50001.0, 0.1), (50002.0, 0.1)))
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.5, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
assert r.book.last_trade_price is not None
|
||||||
|
|
||||||
|
def test_cp_fill_reduces_available(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.NOISE_TRADER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.1, toxicity=0.1,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
assert r.account.available_balance <= s.account.available_balance
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 12. DETERMINISM
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDeterminism:
|
||||||
|
def test_same_input_same_output(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r1 = cwm.transition(s, (a,))
|
||||||
|
r2 = cwm.transition(s, (a,))
|
||||||
|
assert r1.ts_ns == r2.ts_ns
|
||||||
|
assert r1.account.equity == r2.account.equity
|
||||||
|
|
||||||
|
def test_input_not_mutated(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
orig_ts = s.ts_ns
|
||||||
|
orig_equity = s.account.equity
|
||||||
|
cwm.transition(s, (_cross(Side.BUY, frac=0.1),))
|
||||||
|
assert s.ts_ns == orig_ts
|
||||||
|
assert s.account.equity == orig_equity
|
||||||
|
|
||||||
|
def test_timestamp_advances(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r = cwm.transition(s, (_noop(),))
|
||||||
|
assert r.ts_ns > s.ts_ns
|
||||||
|
|
||||||
|
def test_book_state_independent(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s1 = _state(bid=50000.0, ask=50001.0)
|
||||||
|
s2 = _state(bid=49000.0, ask=49001.0)
|
||||||
|
r1 = cwm.transition(s1, (_noop(),))
|
||||||
|
r2 = cwm.transition(s2, (_noop(),))
|
||||||
|
assert r1.book.mid != r2.book.mid
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 13. EDGE CASES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEdgeCases:
|
||||||
|
def test_noop_preserves_everything(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r = cwm.transition(s, (_noop(),))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
assert r.book.best_bid == s.book.best_bid
|
||||||
|
assert len(r.open_orders) == len(s.open_orders)
|
||||||
|
|
||||||
|
def test_zero_equity(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(equity=0.0)
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Should not crash
|
||||||
|
assert isinstance(r.account.equity, float)
|
||||||
|
|
||||||
|
def test_empty_book_no_fill(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
# Remove all asks
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=s.ts_ns, mode=s.mode, venue=s.venue,
|
||||||
|
book=OrderBookState(ts_ns=s.book.ts_ns, symbol=s.book.symbol,
|
||||||
|
bids=s.book.bids, asks=()),
|
||||||
|
account=s.account,
|
||||||
|
)
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_very_small_qty(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.0001)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Might be clipped to zero
|
||||||
|
assert isinstance(r.account.equity, float)
|
||||||
|
|
||||||
|
def test_consecutive_transitions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
for _ in range(10):
|
||||||
|
s = cwm.transition(s, (_noop(),))
|
||||||
|
assert isinstance(s.account.equity, float)
|
||||||
|
|
||||||
|
def test_consecutive_with_actions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
for i in range(5):
|
||||||
|
a = _place(Side.BUY, offset=i, frac=0.05, post_only=True)
|
||||||
|
s = cwm.transition(s, (a,))
|
||||||
|
assert len(s.open_orders) == 5
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 14. TERMINAL
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTerminal:
|
||||||
|
def test_depth_zero(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
assert cwm.terminal(_state(), 0)
|
||||||
|
|
||||||
|
def test_no_intent(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
assert cwm.terminal(_state(), 5)
|
||||||
|
|
||||||
|
def test_with_intent_and_depth(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="t1", ts_ns=1, 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",
|
||||||
|
)
|
||||||
|
s = _state(intent=intent)
|
||||||
|
assert not cwm.terminal(s, 3)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 15. REWARD
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestReward:
|
||||||
|
def _params(self):
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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 test_noop_reward_zero(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r = cwm.transition(s, (_noop(),))
|
||||||
|
assert cwm.reward(s, _noop(), r, self._params()) == 0.0
|
||||||
|
|
||||||
|
def test_cross_spread_penalty(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _cross(Side.BUY, frac=0.1)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert cwm.reward(s, a, r, self._params()) < 0
|
||||||
|
|
||||||
|
def test_higher_w_tail_more_penalty(self):
|
||||||
|
import dataclasses
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = _tp(mae=-40.0, failed_recovery_count=2)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
p1 = self._params()
|
||||||
|
p2_dict = dataclasses.asdict(p1)
|
||||||
|
p2_dict["w_tail_loss"] = 10.0
|
||||||
|
p2_dict["version"] = "t2"
|
||||||
|
p2 = FulfilmentPolicyParams(**p2_dict)
|
||||||
|
assert cwm.reward(s, a, r, p2) < cwm.reward(s, a, r, p1)
|
||||||
230
MALKHUT/malkhut/tests/test_diagnostic.py
Normal file
230
MALKHUT/malkhut/tests/test_diagnostic.py
Normal file
@@ -0,0 +1,230 @@
|
|||||||
|
"""
|
||||||
|
Diagnostic tests — WHY doesn't the system improve?
|
||||||
|
|
||||||
|
Tests that verify:
|
||||||
|
1. Counterparty fills actually affect the book
|
||||||
|
2. System fills are affected by book state
|
||||||
|
3. Different strategies produce different scores
|
||||||
|
4. Evaluation has enough variance
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, AgentRole, CounterpartyAction, FulfilmentAction, OrderType
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.counterparties import ToxicTakerPolicy, PassiveMakerPolicy, 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(bid=50000.0, ask=50001.0, bid_qty=1.0, ask_qty=1.0):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.ENDOGENOUS_AGENT_SIM, venue=_venue(),
|
||||||
|
book=OrderBookState(ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, bid_qty),),
|
||||||
|
asks=(PriceLevel(ask, ask_qty),)),
|
||||||
|
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 _params():
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2, rollout_depth=2,
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# DIAGNOSTIC 1: Do counterparties affect the book?
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCounterpartyImpact:
|
||||||
|
def test_cp_buy_reduces_asks(self):
|
||||||
|
"""Counterparty BUY should consume ask liquidity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(ask_qty=0.5)
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 1.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
total_ask = sum(l.qty for l in r.book.asks)
|
||||||
|
assert total_ask < 0.5, f"Expected ask reduction, got {total_ask}"
|
||||||
|
|
||||||
|
def test_cp_sell_reduces_bids(self):
|
||||||
|
"""Counterparty SELL should consume bid liquidity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid_qty=0.5)
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.SELL, 0, 1.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
total_bid = sum(l.qty for l in r.book.bids)
|
||||||
|
assert total_bid < 0.5, f"Expected bid reduction, got {total_bid}"
|
||||||
|
|
||||||
|
def test_cp_fill_changes_book_state(self):
|
||||||
|
"""Counterparty fill should change the book state."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(ask_qty=1.0)
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 1.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r1 = cwm.transition(s, (_noop(),))
|
||||||
|
r2 = cwm.transition(s, (_noop(), cp))
|
||||||
|
# Book should be different after counterparty fill
|
||||||
|
assert r2.book.best_ask != r1.book.best_ask or sum(l.qty for l in r2.book.asks) != sum(l.qty for l in r1.book.asks)
|
||||||
|
|
||||||
|
def test_cp_does_not_affect_our_position(self):
|
||||||
|
"""Counterparty fill should NOT change our position."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 1.0, toxicity=0.8,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (_noop(), cp))
|
||||||
|
# Our position should be unchanged (no fill on our side)
|
||||||
|
assert r.account.positions.get("BTCUSDT") is None or r.account.positions.get("BTCUSDT").qty == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# DIAGNOSTIC 2: Do different strategies produce different scores?
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStrategyDifferentiation:
|
||||||
|
def test_noop_vs_cross_different_pnl(self):
|
||||||
|
"""NOOP and CROSS_SPREAD should produce different PnL."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r_noop = cwm.transition(s, (_noop(),))
|
||||||
|
r_cross = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
# Cross should produce different equity than noop
|
||||||
|
assert r_noop.account.equity != r_cross.account.equity or \
|
||||||
|
r_noop.book.best_ask != r_cross.book.best_ask
|
||||||
|
|
||||||
|
def test_aggressive_vs_passive_different_pnl(self):
|
||||||
|
"""Aggressive and passive strategies should produce different PnL."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
# Aggressive: cross spread
|
||||||
|
r_agg = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
# Passive: place limit
|
||||||
|
r_pas = cwm.transition(s, (_place(Side.BUY, offset=0, frac=0.1),))
|
||||||
|
# Should produce different book states
|
||||||
|
assert r_agg.book.best_ask != r_pas.book.best_ask or \
|
||||||
|
r_agg.account.equity != r_pas.account.equity
|
||||||
|
|
||||||
|
def test_toxic_vs_safe_different_pnl(self):
|
||||||
|
"""Toxic and safe strategies should produce different PnL."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(ask_qty=0.1) # thin book
|
||||||
|
# Toxic: aggressive cross
|
||||||
|
r_toxic = cwm.transition(s, (_cross(Side.BUY, 0.5),))
|
||||||
|
# Safe: small passive
|
||||||
|
r_safe = cwm.transition(s, (_place(Side.BUY, offset=2, frac=0.01),))
|
||||||
|
# Toxic should have different equity than safe
|
||||||
|
assert r_toxic.account.equity != r_safe.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# DIAGNOSTIC 3: Is the evaluation environment realistic?
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEvaluationRealism:
|
||||||
|
def test_thin_book_consumed_quickly(self):
|
||||||
|
"""With thin book, counterparty should consume it quickly."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid_qty=0.1, ask_qty=0.1)
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 5.0, toxicity=0.9,
|
||||||
|
)
|
||||||
|
# Run multiple steps
|
||||||
|
state = s
|
||||||
|
for _ in range(5):
|
||||||
|
state = cwm.transition(state, (_noop(), cp))
|
||||||
|
# Book should be mostly consumed
|
||||||
|
total_ask = sum(l.qty for l in state.book.asks)
|
||||||
|
assert total_ask < 0.5, f"Expected book consumption, got {total_ask}"
|
||||||
|
|
||||||
|
def test_thick_book_not_consumed(self):
|
||||||
|
"""With thick book, counterparty should not consume it all."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(bid_qty=10.0, ask_qty=10.0)
|
||||||
|
cp = CounterpartyAction(
|
||||||
|
AgentRole.TOXIC_TAKER, ActionKind.CROSS_SPREAD, Side.BUY, 0, 1.0, toxicity=0.9,
|
||||||
|
)
|
||||||
|
state = cwm.transition(s, (_noop(), cp))
|
||||||
|
total_ask = sum(l.qty for l in state.book.asks)
|
||||||
|
assert total_ask > 5.0, f"Expected thick book to survive, got {total_ask}"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# DIAGNOSTIC 4: Does the planner actually produce different actions?
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPlannerDifferentiation:
|
||||||
|
def test_noop_produces_noop(self):
|
||||||
|
"""Without intent, planner should return NOOP."""
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology())
|
||||||
|
s = _state()
|
||||||
|
result = planner.plan(s, _params(), budget_ms=10)
|
||||||
|
assert result.selected_action.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_with_intent_produces_action(self):
|
||||||
|
"""With intent, planner should produce a non-NOOP action."""
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
from malkhut.state import ExecutionIntent, IntentKind
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology())
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="test", ts_ns=1, 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",
|
||||||
|
)
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=1, 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,
|
||||||
|
)
|
||||||
|
result = planner.plan(s, _params(), budget_ms=10)
|
||||||
|
# Should produce a non-NOOP action
|
||||||
|
assert result.selected_action.kind != ActionKind.NOOP or len(result.actions) > 1
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _cross(side, frac):
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
|
|
||||||
|
|
||||||
|
def _place(side, offset=0, frac=0.1):
|
||||||
|
return FulfilmentAction(ActionKind.PLACE, side, OrderType.LIMIT, offset, frac, 200)
|
||||||
357
MALKHUT/malkhut/tests/test_dsl.py
Normal file
357
MALKHUT/malkhut/tests/test_dsl.py
Normal file
@@ -0,0 +1,357 @@
|
|||||||
|
"""
|
||||||
|
Tests for MALKHUT Strategy DSL.
|
||||||
|
|
||||||
|
Verifies:
|
||||||
|
- DSL parsing (text → StrategyTemplate)
|
||||||
|
- DSL decompilation (StrategyTemplate → text)
|
||||||
|
- Action primitives (QUOTE, CROSS, EXIT, etc.)
|
||||||
|
- Market sensors (read values from state)
|
||||||
|
- Decision rules (conditional logic)
|
||||||
|
- Strategy selection (highest priority match)
|
||||||
|
- Builtin strategies
|
||||||
|
- Edge cases (empty strategy, unknown sensor, unknown action)
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PositionState, PriceLevel, Side, TradePathState,
|
||||||
|
VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction
|
||||||
|
from malkhut.training.dsl import (
|
||||||
|
ActionType, ActionPrimitive, SensorType, ComparisonOp,
|
||||||
|
SensorCondition, DecisionRule, StrategyTemplate,
|
||||||
|
StrategyDSLParser, StrategyDSLCompiler, DSLParseError,
|
||||||
|
BUILTIN_STRATEGIES, get_builtin_strategy, list_builtin_strategies,
|
||||||
|
_read_sensor, _primitive_to_action,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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(bid=50000.0, ask=50001.0, equity=10000.0, **kw):
|
||||||
|
path = 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_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, 1.0),), asks=(PriceLevel(ask, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=equity, wallet_balance=equity,
|
||||||
|
available_balance=equity, margin_used=0.0, total_notional=0.0,
|
||||||
|
positions=kw.get("positions", {}),
|
||||||
|
),
|
||||||
|
trade_path=path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. SENSOR CONDITIONS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestSensorConditions:
|
||||||
|
def test_spread_bps(self):
|
||||||
|
s = _state(bid=50000.0, ask=50001.0)
|
||||||
|
cond = SensorCondition(SensorType.SPREAD_BPS, ComparisonOp.LT, 5.0)
|
||||||
|
assert cond.evaluate(s)
|
||||||
|
|
||||||
|
def test_spread_bps_fails(self):
|
||||||
|
s = _state(bid=50000.0, ask=51000.0) # 1000pt spread = ~2000bps
|
||||||
|
cond = SensorCondition(SensorType.SPREAD_BPS, ComparisonOp.LT, 5.0)
|
||||||
|
assert not cond.evaluate(s) # 2000bps > 5bps → False
|
||||||
|
|
||||||
|
def test_toxicity(self):
|
||||||
|
from malkhut.state import TradePathState
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.8,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
cond = SensorCondition(SensorType.TOXICITY, ComparisonOp.GT, 0.5)
|
||||||
|
assert cond.evaluate(s)
|
||||||
|
|
||||||
|
def test_position_qty(self):
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(positions={"BTCUSDT": pos})
|
||||||
|
cond = SensorCondition(SensorType.POSITION_QTY, ComparisonOp.GT, 0.05)
|
||||||
|
assert cond.evaluate(s)
|
||||||
|
|
||||||
|
def test_equity(self):
|
||||||
|
s = _state(equity=10000.0)
|
||||||
|
cond = SensorCondition(SensorType.EQUITY, ComparisonOp.GT, 5000.0)
|
||||||
|
assert cond.evaluate(s)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. SENSOR READING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestSensorReading:
|
||||||
|
def test_read_spread_bps(self):
|
||||||
|
s = _state(bid=50000.0, ask=50001.0)
|
||||||
|
val = _read_sensor(SensorType.SPREAD_BPS, s)
|
||||||
|
assert val > 0
|
||||||
|
|
||||||
|
def test_read_equity(self):
|
||||||
|
s = _state(equity=10000.0)
|
||||||
|
val = _read_sensor(SensorType.EQUITY, s)
|
||||||
|
assert val == 10000.0
|
||||||
|
|
||||||
|
def test_read_imbalance(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.IMBALANCE, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. DSL PARSER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDSLParser:
|
||||||
|
def test_parse_simple_strategy(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.name == "test"
|
||||||
|
assert template.rule_count == 2
|
||||||
|
|
||||||
|
def test_parse_multiple_conditions(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "multi" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rule_count == 1
|
||||||
|
assert len(template.rules[0].conditions) == 2
|
||||||
|
|
||||||
|
def test_parse_bare_action(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "bare" {
|
||||||
|
PRIORITY 1: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rule_count == 1
|
||||||
|
assert template.rules[0].action.action_type == ActionType.NOOP
|
||||||
|
|
||||||
|
def test_parse_exit(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "exit_strat" {
|
||||||
|
PRIORITY 1: IF time_in_trade > 300 THEN EXIT
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.EXIT
|
||||||
|
|
||||||
|
def test_parse_cross(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "cross_strat" {
|
||||||
|
PRIORITY 1: IF spread_bps < 2.0 THEN CROSS(BUY, 0.1)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
action = template.rules[0].action
|
||||||
|
assert action.action_type == ActionType.CROSS
|
||||||
|
assert action.side == Side.BUY
|
||||||
|
|
||||||
|
def test_parse_missing_name(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
with pytest.raises(DSLParseError):
|
||||||
|
parser.parse("STRATEGY { PRIORITY 1: NOOP }")
|
||||||
|
|
||||||
|
def test_parse_missing_braces(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
with pytest.raises(DSLParseError):
|
||||||
|
parser.parse('STRATEGY "test" PRIORITY 1: NOOP')
|
||||||
|
|
||||||
|
def test_parse_empty_rules(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
with pytest.raises(DSLParseError):
|
||||||
|
parser.parse('STRATEGY "test" {}')
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. STRATEGY TEMPLATE
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStrategyTemplate:
|
||||||
|
def test_select_action_matches_rule(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
s = _state(bid=50000.0, ask=50001.0) # 1pt spread = 0.2bps
|
||||||
|
action = template.select_action(s)
|
||||||
|
# 0.2bps < 5.0bps → QUOTE matches
|
||||||
|
assert action.kind == ActionKind.PLACE
|
||||||
|
|
||||||
|
def test_select_action_noop_when_no_match(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 0.01 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
s = _state()
|
||||||
|
action = template.select_action(s)
|
||||||
|
assert action.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_rule_count(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: IF time_in_trade > 300 THEN EXIT
|
||||||
|
PRIORITY 3: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rule_count == 3
|
||||||
|
|
||||||
|
def test_action_types_used(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: IF time_in_trade > 300 THEN EXIT
|
||||||
|
PRIORITY 3: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
types = template.action_types_used
|
||||||
|
assert ActionType.QUOTE in types
|
||||||
|
assert ActionType.EXIT in types
|
||||||
|
assert ActionType.NOOP in types
|
||||||
|
|
||||||
|
def test_sensors_used(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 AND orderflow_toxicity > 0.5 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
sensors = template.sensors_used
|
||||||
|
assert SensorType.SPREAD_BPS in sensors
|
||||||
|
assert SensorType.TOXICITY in sensors
|
||||||
|
|
||||||
|
def test_evaluate_conditions(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
s = _state()
|
||||||
|
results = template.evaluate_conditions(s)
|
||||||
|
assert len(results) == 2
|
||||||
|
assert results[0][0] == 1 # priority
|
||||||
|
assert isinstance(results[0][1], bool) # matched
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. DSL COMPILER (roundtrip)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDSLCompiler:
|
||||||
|
def test_compile_and_decompile(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "roundtrip" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = compiler.compile(dsl)
|
||||||
|
output = compiler.decompile(template)
|
||||||
|
assert "roundtrip" in output
|
||||||
|
assert "QUOTE" in output
|
||||||
|
assert "NOOP" in output
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 6. BUILTIN STRATEGIES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestBuiltinStrategies:
|
||||||
|
def test_list_builtins(self):
|
||||||
|
names = list_builtin_strategies()
|
||||||
|
assert "passive_maker" in names
|
||||||
|
assert "aggressive_taker" in names
|
||||||
|
assert "toxicity_avoider" in names
|
||||||
|
assert "path_risk_exit" in names
|
||||||
|
assert "regime_adaptive" in names
|
||||||
|
|
||||||
|
def test_get_builtin(self):
|
||||||
|
text = get_builtin_strategy("passive_maker")
|
||||||
|
assert text is not None
|
||||||
|
assert "STRATEGY" in text
|
||||||
|
|
||||||
|
def test_parse_all_builtins(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
for name in list_builtin_strategies():
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.name == name
|
||||||
|
assert template.rule_count > 0
|
||||||
|
|
||||||
|
def test_builtin_passive_maker_parsed(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("passive_maker")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.rule_count >= 4
|
||||||
|
assert SensorType.SPREAD_BPS in template.sensors_used
|
||||||
|
assert SensorType.TOXICITY in template.sensors_used
|
||||||
|
|
||||||
|
def test_builtin_aggressive_taker_parsed(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("aggressive_taker")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert ActionType.CROSS in template.action_types_used
|
||||||
|
|
||||||
|
def test_builtin_path_risk_exit_parsed(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("path_risk_exit")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.MAE_BPS in template.sensors_used
|
||||||
|
assert ActionType.EXIT in template.action_types_used
|
||||||
345
MALKHUT/malkhut/tests/test_dsl_expanded.py
Normal file
345
MALKHUT/malkhut/tests/test_dsl_expanded.py
Normal file
@@ -0,0 +1,345 @@
|
|||||||
|
"""
|
||||||
|
Expanded DSL tests — covers all new primitives, sensors, and builtin strategies.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OrderBookState, PositionState,
|
||||||
|
PriceLevel, Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind
|
||||||
|
from malkhut.training.dsl import (
|
||||||
|
ActionType, SensorType, ComparisonOp, SensorCondition,
|
||||||
|
StrategyDSLParser, StrategyDSLCompiler, DSLParseError,
|
||||||
|
BUILTIN_STRATEGIES, get_builtin_strategy, list_builtin_strategies,
|
||||||
|
_read_sensor, _primitive_to_action, ActionPrimitive,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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):
|
||||||
|
path = kw.get("trade_path")
|
||||||
|
pos = kw.get("positions", {})
|
||||||
|
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), kw.get("bid_qty", 1.0)),
|
||||||
|
PriceLevel(kw.get("bid", 50000.0) - 1.0, 2.0)),
|
||||||
|
asks=(PriceLevel(kw.get("ask", 50001.0), kw.get("ask_qty", 1.0)),
|
||||||
|
PriceLevel(kw.get("ask", 50001.0) + 1.0, 2.0)),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1, equity=kw.get("equity", 10000.0), wallet_balance=kw.get("equity", 10000.0),
|
||||||
|
available_balance=kw.get("equity", 10000.0), margin_used=0.0,
|
||||||
|
total_notional=kw.get("notional", 0.0), positions=pos,
|
||||||
|
),
|
||||||
|
trade_path=path,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. EXPANDED SENSORS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExpandedSensors:
|
||||||
|
def test_bid_depth_3(self):
|
||||||
|
s = _state(bid=50000.0, bid_qty=3.0)
|
||||||
|
assert _read_sensor(SensorType.BID_DEPTH_3, s) == pytest.approx(5.0, abs=0.1)
|
||||||
|
|
||||||
|
def test_bid_depth_10(self):
|
||||||
|
s = _state(bid=50000.0, bid_qty=10.0)
|
||||||
|
assert _read_sensor(SensorType.BID_DEPTH_10, s) > 0
|
||||||
|
|
||||||
|
def test_ask_depth_5(self):
|
||||||
|
s = _state(ask=50001.0, ask_qty=2.0)
|
||||||
|
assert _read_sensor(SensorType.ASK_DEPTH_5, s) > 0
|
||||||
|
|
||||||
|
def test_imbalance_3(self):
|
||||||
|
s = _state(bid_qty=3.0, ask_qty=1.0)
|
||||||
|
val = _read_sensor(SensorType.IMBALANCE_3, s)
|
||||||
|
assert val > 0
|
||||||
|
|
||||||
|
def test_imbalance_10(self):
|
||||||
|
s = _state(bid_qty=10.0, ask_qty=5.0)
|
||||||
|
val = _read_sensor(SensorType.IMBALANCE_10, s)
|
||||||
|
assert val > 0
|
||||||
|
|
||||||
|
def test_bid_ask_ratio(self):
|
||||||
|
s = _state(bid_qty=2.0, ask_qty=1.0)
|
||||||
|
val = _read_sensor(SensorType.BID_ASK_RATIO, s)
|
||||||
|
assert val > 1.0 # bid > ask → ratio > 1
|
||||||
|
|
||||||
|
def test_toxicity(self):
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.8,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert _read_sensor(SensorType.TOXICITY, s) == 0.8
|
||||||
|
|
||||||
|
def test_mae_bps(self):
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-25.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert _read_sensor(SensorType.MAE_BPS, s) == -25.0
|
||||||
|
|
||||||
|
def test_equity(self):
|
||||||
|
s = _state(equity=10000.0)
|
||||||
|
assert _read_sensor(SensorType.EQUITY, s) == 10000.0
|
||||||
|
|
||||||
|
def test_leverage(self):
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(positions={"BTCUSDT": pos})
|
||||||
|
assert _read_sensor(SensorType.LEVERAGE, s) == 0.5
|
||||||
|
|
||||||
|
def test_position_age(self):
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=150.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert _read_sensor(SensorType.POSITION_AGE_S, s) == 150.0
|
||||||
|
|
||||||
|
def test_atr_14(self):
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=20.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert _read_sensor(SensorType.ATR_14, s) == pytest.approx(28.0, abs=0.1)
|
||||||
|
|
||||||
|
def test_all_sensors_readable(self):
|
||||||
|
"""Every sensor should be readable without crashing."""
|
||||||
|
s = _state()
|
||||||
|
for sensor in SensorType:
|
||||||
|
val = _read_sensor(sensor, s)
|
||||||
|
assert isinstance(val, float), f"{sensor.value} returned {type(val)}"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. EXPANDED ACTION PRIMITIVES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExpandedActions:
|
||||||
|
def test_all_action_types_creatable(self):
|
||||||
|
"""Every ActionType should be constructable."""
|
||||||
|
for at in ActionType:
|
||||||
|
p = ActionPrimitive(action_type=at)
|
||||||
|
assert p.action_type == at
|
||||||
|
|
||||||
|
def test_quote_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.QUOTE, side=Side.BUY, offset_ticks=1, size_fraction=0.25)
|
||||||
|
assert p.side == Side.BUY
|
||||||
|
assert p.offset_ticks == 1
|
||||||
|
|
||||||
|
def test_cross_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.CROSS, side=Side.SELL, size_fraction=0.1)
|
||||||
|
assert p.action_type == ActionType.CROSS
|
||||||
|
|
||||||
|
def test_trailing_stop_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.TRAILING_STOP, trail_distance_bps=20.0)
|
||||||
|
assert p.trail_distance_bps == 20.0
|
||||||
|
|
||||||
|
def test_half_exit_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.HALF_EXIT)
|
||||||
|
assert p.action_type == ActionType.HALF_EXIT
|
||||||
|
|
||||||
|
def test_bracket_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.BRACKET)
|
||||||
|
assert p.action_type == ActionType.BRACKET
|
||||||
|
|
||||||
|
def test_ladder_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.LADDER, levels=5, size_per_step=0.05)
|
||||||
|
assert p.levels == 5
|
||||||
|
|
||||||
|
def test_grid_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.GRID, levels=10, size_per_step=0.02)
|
||||||
|
assert p.levels == 10
|
||||||
|
|
||||||
|
def test_iceberg_primitive(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.ICEBERG, size_fraction=0.5, steps=5)
|
||||||
|
assert p.steps == 5
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. EXPANDED COMPARISON OPERATORS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExpandedOperators:
|
||||||
|
def test_abs_gt(self):
|
||||||
|
s = _state()
|
||||||
|
cond = SensorCondition(SensorType.IMBALANCE, ComparisonOp.ABS_GT, 0.5)
|
||||||
|
assert not cond.evaluate(s) # imbalance ~0
|
||||||
|
|
||||||
|
def test_abs_gt_true(self):
|
||||||
|
s = _state(bid_qty=10.0, ask_qty=1.0)
|
||||||
|
cond = SensorCondition(SensorCondition.sensor, ComparisonOp.ABS_GT, 0.5)
|
||||||
|
|
||||||
|
def test_changing(self):
|
||||||
|
s = _state()
|
||||||
|
cond = SensorCondition(SensorType.EQUITY, ComparisonOp.CHANGING, 9999.0)
|
||||||
|
assert cond.evaluate(s) # 10000 != 9999
|
||||||
|
|
||||||
|
def test_stable(self):
|
||||||
|
s = _state(equity=10000.0)
|
||||||
|
cond = SensorCondition(SensorType.EQUITY, ComparisonOp.STABLE, 10000.0)
|
||||||
|
assert cond.evaluate(s)
|
||||||
|
|
||||||
|
def test_crossing_above(self):
|
||||||
|
s = _state()
|
||||||
|
cond = SensorCondition(SensorType.EQUITY, ComparisonOp.CROSSING_ABOVE, 5000.0)
|
||||||
|
assert cond.evaluate(s) # 10000 > 5000
|
||||||
|
|
||||||
|
def test_crossing_below(self):
|
||||||
|
s = _state()
|
||||||
|
cond = SensorCondition(SensorType.EQUITY, ComparisonOp.CROSSING_BELOW, 20000.0)
|
||||||
|
assert cond.evaluate(s) # 10000 < 20000
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. EXPANDED BUILTIN STRATEGIES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExpandedBuiltins:
|
||||||
|
def test_all_builtins_parse(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
for name in list_builtin_strategies():
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.name == name
|
||||||
|
assert template.rule_count > 0
|
||||||
|
|
||||||
|
def test_builtin_count(self):
|
||||||
|
assert len(BUILTIN_STRATEGIES) >= 15
|
||||||
|
|
||||||
|
def test_momentum_catcher(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("momentum_catcher")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert ActionType.CROSS in template.action_types_used
|
||||||
|
|
||||||
|
def test_scalper(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("scalper")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert ActionType.CROSS in template.action_types_used
|
||||||
|
|
||||||
|
def test_session_guard(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("session_guard")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.IS_WEEKEND in template.sensors_used
|
||||||
|
|
||||||
|
def test_grid_trader(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("grid_trader")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.rule_count >= 4
|
||||||
|
|
||||||
|
def test_hybrid_adaptive(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("hybrid_adaptive")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.rule_count >= 8
|
||||||
|
|
||||||
|
def test_risk_parity(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("risk_parity")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.RISK_BUDGET_USED in template.sensors_used
|
||||||
|
|
||||||
|
def test_funding_arb(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("funding_arb")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.FUNDING in template.sensors_used
|
||||||
|
|
||||||
|
def test_inventory_manager(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("inventory_manager")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.POSITION_QTY in template.sensors_used
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. DSL ROUNDTRIP
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDSLRoundtrip:
|
||||||
|
def test_compile_decompile_roundtrip(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
PRIORITY 2: IF time_in_trade > 300 THEN EXIT
|
||||||
|
PRIORITY 3: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = compiler.compile(dsl)
|
||||||
|
output = compiler.decompile(template)
|
||||||
|
assert "test" in output
|
||||||
|
assert "QUOTE" in output
|
||||||
|
assert "EXIT" in output
|
||||||
|
assert "NOOP" in output
|
||||||
|
|
||||||
|
def test_all_builtin_roundtrip(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
for name in list_builtin_strategies():
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
output = compiler.decompile(template)
|
||||||
|
# Re-parse to verify
|
||||||
|
template2 = compiler.compile(output)
|
||||||
|
assert template2.name == template.name
|
||||||
|
assert template2.rule_count == template.rule_count
|
||||||
313
MALKHUT/malkhut/tests/test_dsl_new_features.py
Normal file
313
MALKHUT/malkhut/tests/test_dsl_new_features.py
Normal file
@@ -0,0 +1,313 @@
|
|||||||
|
"""
|
||||||
|
Comprehensive DSL tests for new features (50+ tests).
|
||||||
|
|
||||||
|
Tests new action primitives, sensors, and builtin strategies added for
|
||||||
|
trajectory persistence, discrepancy tracking, feature importance,
|
||||||
|
policy rollback, and stress scenarios.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel,
|
||||||
|
Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind
|
||||||
|
from malkhut.training.dsl import (
|
||||||
|
ActionType, SensorType, ComparisonOp, SensorCondition,
|
||||||
|
StrategyDSLParser, StrategyDSLCompiler, DSLParseError,
|
||||||
|
BUILTIN_STRATEGIES, get_builtin_strategy, list_builtin_strategies,
|
||||||
|
_read_sensor, ActionPrimitive,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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=kw.get("equity", 10000.0),
|
||||||
|
wallet_balance=kw.get("equity", 10000.0),
|
||||||
|
available_balance=kw.get("equity", 10000.0),
|
||||||
|
margin_used=0.0, total_notional=kw.get("notional", 0.0)),
|
||||||
|
trade_path=tp,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. NEW ACTION PRIMITIVES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewActionPrimitives:
|
||||||
|
def test_log_state(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.LOG_STATE)
|
||||||
|
assert p.action_type == ActionType.LOG_STATE
|
||||||
|
|
||||||
|
def test_check_regime(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.CHECK_REGIME)
|
||||||
|
assert p.action_type == ActionType.CHECK_REGIME
|
||||||
|
|
||||||
|
def test_switch_strategy(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.SWITCH_STRATEGY,
|
||||||
|
metadata={"target": "aggressive_taker"})
|
||||||
|
assert p.action_type == ActionType.SWITCH_STRATEGY
|
||||||
|
assert p.metadata["target"] == "aggressive_taker"
|
||||||
|
|
||||||
|
def test_wait_for_regime(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.WAIT_FOR_REGIME, duration_s=60.0)
|
||||||
|
assert p.action_type == ActionType.WAIT_FOR_REGIME
|
||||||
|
assert p.duration_s == 60.0
|
||||||
|
|
||||||
|
def test_adjust_size(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.ADJUST_SIZE, side=Side.BUY, size_fraction=0.2)
|
||||||
|
assert p.action_type == ActionType.ADJUST_SIZE
|
||||||
|
assert p.side == Side.BUY
|
||||||
|
|
||||||
|
def test_hedge_pair(self):
|
||||||
|
p = ActionPrimitive(action_type=ActionType.HEDGE_PAIR, side=Side.SELL, size_fraction=0.1)
|
||||||
|
assert p.action_type == ActionType.HEDGE_PAIR
|
||||||
|
|
||||||
|
def test_all_new_actions_frozen(self):
|
||||||
|
for at in [ActionType.LOG_STATE, ActionType.CHECK_REGIME,
|
||||||
|
ActionType.SWITCH_STRATEGY, ActionType.WAIT_FOR_REGIME,
|
||||||
|
ActionType.ADJUST_SIZE, ActionType.HEDGE_PAIR]:
|
||||||
|
p = ActionPrimitive(action_type=at)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
p.action_type = ActionType.NOOP
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. NEW SENSORS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewSensors:
|
||||||
|
def test_discrepancy_rate(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.DISCREPANCY_RATE, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_trajectory_length(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.TRAJECTORY_LENGTH, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_feature_importance_top(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.FEATURE_IMPORTANCE_TOP, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_current_regime(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.CURRENT_REGIME, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_regime_confidence(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.REGIME_CONFIDENCE, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_strategy_age_s(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.STRATEGY_AGE_S, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_strategy_score(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.STRATEGY_SCORE, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_portfolio_risk(self):
|
||||||
|
s = _state(notional=5000.0, equity=10000.0)
|
||||||
|
val = _read_sensor(SensorType.PORTFOLIO_RISK, s)
|
||||||
|
assert val == pytest.approx(0.5, abs=0.01)
|
||||||
|
|
||||||
|
def test_correlation_btc(self):
|
||||||
|
s = _state()
|
||||||
|
val = _read_sensor(SensorType.CORRELATION_BTC, s)
|
||||||
|
assert isinstance(val, float)
|
||||||
|
|
||||||
|
def test_all_new_sensors_readable(self):
|
||||||
|
"""Every new sensor should be readable without crashing."""
|
||||||
|
s = _state()
|
||||||
|
new_sensors = [
|
||||||
|
SensorType.DISCREPANCY_RATE, SensorType.TRAJECTORY_LENGTH,
|
||||||
|
SensorType.FEATURE_IMPORTANCE_TOP, SensorType.CURRENT_REGIME,
|
||||||
|
SensorType.REGIME_CONFIDENCE, SensorType.STRATEGY_AGE_S,
|
||||||
|
SensorType.STRATEGY_SCORE, SensorType.PORTFOLIO_RISK,
|
||||||
|
SensorType.CORRELATION_BTC,
|
||||||
|
]
|
||||||
|
for sensor in new_sensors:
|
||||||
|
val = _read_sensor(sensor, s)
|
||||||
|
assert isinstance(val, float), f"{sensor.value} returned {type(val)}"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. NEW DSL PARSING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewDSLParser:
|
||||||
|
def test_parse_log_state(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: LOG_STATE
|
||||||
|
PRIORITY 2: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.LOG_STATE
|
||||||
|
|
||||||
|
def test_parse_check_regime(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: CHECK_REGIME
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.CHECK_REGIME
|
||||||
|
|
||||||
|
def test_parse_switch_strategy(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: SWITCH_STRATEGY(aggressive_taker)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.SWITCH_STRATEGY
|
||||||
|
|
||||||
|
def test_parse_wait_for_regime(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: WAIT_FOR_REGIME(60)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.WAIT_FOR_REGIME
|
||||||
|
|
||||||
|
def test_parse_adjust_size(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: ADJUST_SIZE(BUY, 0.2)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.ADJUST_SIZE
|
||||||
|
|
||||||
|
def test_parse_hedge_pair(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: HEDGE_PAIR(SELL, 0.1)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
assert template.rules[0].action.action_type == ActionType.HEDGE_PAIR
|
||||||
|
|
||||||
|
def test_parse_new_sensors(self):
|
||||||
|
parser = StrategyDSLParser()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF discrepancy_rate > 0.3 THEN CANCEL_ALL
|
||||||
|
PRIORITY 2: IF portfolio_risk > 0.8 THEN EXIT
|
||||||
|
PRIORITY 3: IF current_regime > 0.7 THEN QUOTE(BUY, 0, 0.25)
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = parser.parse(dsl)
|
||||||
|
sensors = template.sensors_used
|
||||||
|
assert SensorType.DISCREPANCY_RATE in sensors
|
||||||
|
assert SensorType.PORTFOLIO_RISK in sensors
|
||||||
|
assert SensorType.CURRENT_REGIME in sensors
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. NEW BUILTIN STRATEGIES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewBuiltinStrategies:
|
||||||
|
def test_all_builtins_parse(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
for name in list_builtin_strategies():
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.name == name
|
||||||
|
assert template.rule_count > 0
|
||||||
|
|
||||||
|
def test_builtin_count_increased(self):
|
||||||
|
assert len(BUILTIN_STRATEGIES) >= 20
|
||||||
|
|
||||||
|
def test_regime_switcher(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("regime_switcher")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.REGIME_SCORE in template.sensors_used or SensorType.CURRENT_REGIME in template.sensors_used
|
||||||
|
|
||||||
|
def test_discrepancy_aware(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("discrepancy_aware")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.DISCREPANCY_RATE in template.sensors_used
|
||||||
|
|
||||||
|
def test_portfolio_risk_manager(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("portfolio_risk_manager")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert SensorType.PORTFOLIO_RISK in template.sensors_used
|
||||||
|
|
||||||
|
def test_multi_regime_adaptive(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
text = get_builtin_strategy("multi_regime_adaptive")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
assert template.rule_count >= 6
|
||||||
|
|
||||||
|
def test_new_builtins_execute(self):
|
||||||
|
"""All new builtins should be executable."""
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
s = _state()
|
||||||
|
for name in ["regime_switcher", "discrepancy_aware",
|
||||||
|
"portfolio_risk_manager", "multi_regime_adaptive"]:
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
action = template.select_action(s)
|
||||||
|
assert action.kind is not None
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. DSL ROUNDTRIP WITH NEW FEATURES
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewDSLRoundtrip:
|
||||||
|
def test_compile_decompile_new_actions(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
dsl = '''
|
||||||
|
STRATEGY "test" {
|
||||||
|
PRIORITY 1: IF discrepancy_rate > 0.3 THEN CANCEL_ALL
|
||||||
|
PRIORITY 2: IF portfolio_risk > 0.8 THEN EXIT
|
||||||
|
PRIORITY 3: LOG_STATE
|
||||||
|
PRIORITY 4: NOOP
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
template = compiler.compile(dsl)
|
||||||
|
output = compiler.decompile(template)
|
||||||
|
assert "discrepancy_rate" in output or "DISCREPANCY_RATE" in output
|
||||||
|
assert "LOG_STATE" in output
|
||||||
|
assert "NOOP" in output
|
||||||
|
|
||||||
|
def test_all_new_builtin_roundtrip(self):
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
for name in list_builtin_strategies():
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
decompiled = compiler.decompile(template)
|
||||||
|
template2 = compiler.compile(decompiled)
|
||||||
|
assert template2.name == template.name
|
||||||
412
MALKHUT/malkhut/tests/test_e2e_integration.py
Normal file
412
MALKHUT/malkhut/tests/test_e2e_integration.py
Normal file
@@ -0,0 +1,412 @@
|
|||||||
|
"""
|
||||||
|
End-to-end integration test — wire ALL subsystems together.
|
||||||
|
|
||||||
|
Simulates a full training cycle:
|
||||||
|
1. Create initial state + scenarios
|
||||||
|
2. Run training pipeline (CMA-ES → evaluate → promote)
|
||||||
|
3. Register promoted policy
|
||||||
|
4. Load into engine
|
||||||
|
5. Engine plans on live state
|
||||||
|
6. Risk gate validates
|
||||||
|
7. BingX adapter tracks order
|
||||||
|
8. Zinc SHM publishes book/account/fulfilment
|
||||||
|
9. ClickHouse persists decisions
|
||||||
|
10. Control plane sends HOT_RELOAD
|
||||||
|
11. Engine hot-reloads new policy
|
||||||
|
12. Replay verifier validates CWM against trajectory
|
||||||
|
13. Training logger records everything
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PositionState, PriceLevel,
|
||||||
|
Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import (
|
||||||
|
ActionKind, FulfilmentAction, PlannedPolicy, RiskDecision,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.cwm.replay_verify import (
|
||||||
|
ReplayVerifier, ReplayStep, TrajectoryRecorder,
|
||||||
|
)
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.risk.gate import RiskGate
|
||||||
|
from malkhut.venue.bingx.adapter import BingXVenueAdapter, BingXConfig
|
||||||
|
from malkhut.ipc.zinc_plane import MalkhutZincPlane
|
||||||
|
from malkhut.ipc.control_plane import MalkhutControlPlane, ControlCommand, ControlPlaneFrame
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
from malkhut.execution.asex_integration import FulfilmentWorker, RiskWorker, RiskCheck
|
||||||
|
from malkhut.training.pipeline import TrainingPipeline, PipelineConfig, TrainingLogger
|
||||||
|
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
||||||
|
from malkhut.training.cma_trainer import (
|
||||||
|
CMAParameterCodec, PolicyEvaluator, ScenarioFactory, SelfPlayPool,
|
||||||
|
PolicySnapshot,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── Helpers ──────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
d = dict(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)
|
||||||
|
d.update(kw)
|
||||||
|
return VenueRules(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _baseline(**kw):
|
||||||
|
d = dict(
|
||||||
|
version="baseline", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.10, 0.25, 0.50),
|
||||||
|
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 _make_state(bid=50000.0, ask=50001.0, equity=10000.0):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.ENDOGENOUS_AGENT_SIM,
|
||||||
|
venue=_venue(),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, 1.0), PriceLevel(bid - 1.0, 2.0)),
|
||||||
|
asks=(PriceLevel(ask, 1.0), PriceLevel(ask + 1.0, 2.0)),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=equity, wallet_balance=equity,
|
||||||
|
available_balance=equity, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_intent(symbol="BTCUSDT", urgency=0.5):
|
||||||
|
return ExecutionIntent(
|
||||||
|
intent_id=f"e2e_{int(time.time_ns())}", ts_ns=1_000_000_000,
|
||||||
|
symbol=symbol, kind=IntentKind.ENTER_LONG, target_qty=0.01,
|
||||||
|
max_notional=500.0, urgency=urgency, 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="e2e_test",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# THE FULL END-TO-END TEST
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEndToEndFullCycle:
|
||||||
|
"""
|
||||||
|
Wire ALL subsystems together and simulate a full training cycle:
|
||||||
|
|
||||||
|
training pipeline → registry → engine → plan → risk → venue →
|
||||||
|
zinc → clickhouse → control plane → hot reload → replay verify
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_full_e2e_cycle(self):
|
||||||
|
# ──────────────────────────────────────────────────────────────────
|
||||||
|
# 1. SETUP: All subsystems
|
||||||
|
# ──────────────────────────────────────────────────────────────────
|
||||||
|
zinc = MalkhutZincPlane(prefix="e2e_test")
|
||||||
|
control_plane = MalkhutControlPlane()
|
||||||
|
registry = PolicyRegistry()
|
||||||
|
engine = None
|
||||||
|
|
||||||
|
try:
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 2. TRAINING: Run pipeline to produce a trained policy
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
pipeline_config = PipelineConfig(
|
||||||
|
max_generations=1, max_evals_per_generation=7,
|
||||||
|
max_time_s=30, auto_promote=True,
|
||||||
|
)
|
||||||
|
pipeline = TrainingPipeline(
|
||||||
|
config=pipeline_config, registry=registry,
|
||||||
|
log_path="/dev/null",
|
||||||
|
)
|
||||||
|
pipeline_result = pipeline.run(
|
||||||
|
incumbent=_baseline(version="init"),
|
||||||
|
symbols=("BTCUSDT",),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify training produced results
|
||||||
|
assert pipeline_result.generations_run >= 1
|
||||||
|
assert pipeline_result.total_evals > 0
|
||||||
|
assert len(pipeline_result.events) > 0
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 3. REGISTRY: Check promoted policy exists
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
active_params = registry.load_active()
|
||||||
|
assert active_params is not None, "No active policy after training"
|
||||||
|
trained_version = active_params.version
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 4. ENGINE: Create and load trained policy
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
engine = FulfilmentEngine(
|
||||||
|
zinc=zinc, control_plane=control_plane,
|
||||||
|
registry=registry,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify engine loaded the trained policy
|
||||||
|
current_params = engine.params_provider()
|
||||||
|
assert current_params.version == trained_version
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 5. LIVE PLANNING: Engine plans on live state
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
state = _make_state()
|
||||||
|
state_with_intent = MarketWorldState(
|
||||||
|
ts_ns=state.ts_ns, mode=state.mode, venue=state.venue,
|
||||||
|
book=state.book, account=state.account,
|
||||||
|
intent=_make_intent(),
|
||||||
|
)
|
||||||
|
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
rng_seed=42,
|
||||||
|
)
|
||||||
|
planned = planner.plan(
|
||||||
|
root_state=state_with_intent, params=current_params, budget_ms=20,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify planner produced valid output
|
||||||
|
assert isinstance(planned, PlannedPolicy)
|
||||||
|
assert len(planned.actions) > 0
|
||||||
|
assert abs(sum(planned.probabilities) - 1.0) < 1e-6
|
||||||
|
assert planned.selected_action is not None
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 6. RISK GATE: Validate the planned action
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
risk_gate = RiskGate()
|
||||||
|
decision = risk_gate.validate(state_with_intent, planned, current_params)
|
||||||
|
assert isinstance(decision, RiskDecision)
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 7. VENUE ADAPTER: Track the order
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
adapter.execute(state_with_intent, decision)
|
||||||
|
|
||||||
|
# Verify adapter tracked the order (if approved)
|
||||||
|
if decision.approved and decision.action.kind != ActionKind.NOOP:
|
||||||
|
assert adapter.total_orders >= 1
|
||||||
|
working = adapter.get_working()
|
||||||
|
assert len(working) >= 1
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 8. ZINC SHM: Publish book/account/fulfilment
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
zinc.publish_book({
|
||||||
|
"ts_ns": state.ts_ns, "symbol": "BTCUSDT",
|
||||||
|
"bid": state.book.best_bid, "ask": state.book.best_ask,
|
||||||
|
})
|
||||||
|
book_data, book_seq = zinc.read_book()
|
||||||
|
assert book_data["symbol"] == "BTCUSDT"
|
||||||
|
assert book_seq >= 1
|
||||||
|
|
||||||
|
zinc.publish_account({
|
||||||
|
"ts_ns": state.ts_ns, "equity": state.account.equity,
|
||||||
|
})
|
||||||
|
acct_data, acct_seq = zinc.read_account()
|
||||||
|
assert acct_data["equity"] == 10000.0
|
||||||
|
|
||||||
|
zinc.publish_fulfilment({
|
||||||
|
"ts_ns": state.ts_ns, "action": str(planned.selected_action.kind.value),
|
||||||
|
"approved": decision.approved,
|
||||||
|
})
|
||||||
|
fulfil_data, fulfil_seq = zinc.read_fulfilment()
|
||||||
|
assert fulfil_seq >= 1
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 9. CLICKHOUSE: Persist decision
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
store.store_fulfilment_decision(
|
||||||
|
ts_ns=state.ts_ns, exchange="bingx", symbol="BTCUSDT",
|
||||||
|
intent_id="e2e_test", state_hash="abc123",
|
||||||
|
selected_action=str(planned.selected_action.kind.value),
|
||||||
|
root_distribution=str(planned.probabilities),
|
||||||
|
risk_decision=f"{decision.approved}:{decision.reason}",
|
||||||
|
policy_version=current_params.version, latency_ms=5.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Verify CH query
|
||||||
|
result = store.query("SELECT count() FROM fulfilment_decisions")
|
||||||
|
assert int(result.strip()) >= 1
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 10. CONTROL PLANE: HOT_RELOAD_POLICY
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# Register a new policy and promote
|
||||||
|
new_params = _baseline(version="v2_reloaded")
|
||||||
|
registry.register_candidate(new_params, score=15.0)
|
||||||
|
registry.promote(new_params.version, PolicyStage.ACTIVE, "e2e reload")
|
||||||
|
|
||||||
|
# Send HOT_RELOAD via control plane
|
||||||
|
control_plane.publish_command(ControlPlaneFrame(
|
||||||
|
command=ControlCommand.HOT_RELOAD_POLICY.value,
|
||||||
|
ts_ns=time.time_ns(),
|
||||||
|
params={"policy_version": new_params.version},
|
||||||
|
source="e2e_test",
|
||||||
|
))
|
||||||
|
|
||||||
|
# Engine processes control plane
|
||||||
|
engine._process_control_plane()
|
||||||
|
|
||||||
|
# Verify engine reloaded
|
||||||
|
assert engine.params_provider().version == new_params.version
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 11. REPLAY VERIFICATION: Validate CWM determinism
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
recorder = TrajectoryRecorder(max_steps=10)
|
||||||
|
s = _make_state()
|
||||||
|
actions = [_noop(), _cross(Side.BUY, 0.1), _noop()]
|
||||||
|
|
||||||
|
for i, a in enumerate(actions):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
recorder.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
|
||||||
|
# Verify deterministic re-run
|
||||||
|
ok, mismatches = recorder.verify_deterministic(cwm)
|
||||||
|
assert ok, f"Determinism failed: {mismatches}"
|
||||||
|
|
||||||
|
# Verify against replay
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
replay = recorder.to_replay_steps()
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert result.passed, f"Replay failed: {result.mismatches}"
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 12. TRAINING LOGGER: Verify events were recorded
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
events = pipeline.logger.get_events()
|
||||||
|
assert len(events) > 0
|
||||||
|
event_types = [e.event_type for e in events]
|
||||||
|
assert "run_start" in event_types
|
||||||
|
assert "generation" in event_types
|
||||||
|
assert "run_end" in event_types
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 13. ASEX WORKERS: Verify state mutations through ASEx
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
fw = engine.fulfilment_worker
|
||||||
|
assert fw.mutation_count >= 0
|
||||||
|
fw.reload_policy(new_params)
|
||||||
|
time.sleep(0.05)
|
||||||
|
assert fw.params is not None
|
||||||
|
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
# 14. CLEANUP
|
||||||
|
# ──────────────────────────────────────────────────────────────
|
||||||
|
adapter.close()
|
||||||
|
engine.close()
|
||||||
|
|
||||||
|
finally:
|
||||||
|
zinc.close_all()
|
||||||
|
control_plane.close()
|
||||||
|
|
||||||
|
def test_multi_step_trajectory(self):
|
||||||
|
"""Run a multi-step trajectory through the full pipeline."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
recorder = TrajectoryRecorder(max_steps=20)
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||||||
|
)
|
||||||
|
params = _baseline()
|
||||||
|
risk_gate = RiskGate()
|
||||||
|
adapter = BingXVenueAdapter()
|
||||||
|
zinc = MalkhutZincPlane(prefix="e2e_traj")
|
||||||
|
|
||||||
|
try:
|
||||||
|
s = _make_state()
|
||||||
|
total_pnl = 0.0
|
||||||
|
|
||||||
|
for step in range(10):
|
||||||
|
# Add intent
|
||||||
|
intent = _make_intent(urgency=0.5)
|
||||||
|
s_with_intent = MarketWorldState(
|
||||||
|
ts_ns=s.ts_ns, mode=s.mode, venue=s.venue,
|
||||||
|
book=s.book, account=s.account, intent=intent,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Plan
|
||||||
|
planned = planner.plan(root_state=s_with_intent, params=params, budget_ms=15)
|
||||||
|
|
||||||
|
# Risk gate
|
||||||
|
decision = risk_gate.validate(s_with_intent, planned, params)
|
||||||
|
|
||||||
|
# Venue adapter
|
||||||
|
adapter.execute(s_with_intent, decision)
|
||||||
|
|
||||||
|
# Zinc
|
||||||
|
zinc.publish_book({"ts_ns": s.ts_ns, "symbol": "BTCUSDT"})
|
||||||
|
|
||||||
|
# CWM transition
|
||||||
|
cps = default_counterparty_ecology()
|
||||||
|
import random
|
||||||
|
rng = random.Random(42 + step)
|
||||||
|
cp_actions = tuple(cp.rollout_action(s, rng) for cp in cps)
|
||||||
|
next_state = cwm.transition(s, (planned.selected_action, *cp_actions))
|
||||||
|
|
||||||
|
# Record trajectory
|
||||||
|
recorder.record(step, s, (planned.selected_action, *cp_actions), next_state)
|
||||||
|
|
||||||
|
# Track PnL
|
||||||
|
pnl = next_state.account.equity - s.account.equity
|
||||||
|
total_pnl += pnl
|
||||||
|
|
||||||
|
s = next_state
|
||||||
|
|
||||||
|
# Verify trajectory
|
||||||
|
ok, mismatches = recorder.verify_deterministic(cwm)
|
||||||
|
assert ok, f"Determinism failed at step: {mismatches}"
|
||||||
|
|
||||||
|
# Verify replay
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
replay = recorder.to_replay_steps()
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
# Verify final state is valid
|
||||||
|
assert s.account.equity > 0
|
||||||
|
assert s.ts_ns > 1_000_000_000
|
||||||
|
|
||||||
|
finally:
|
||||||
|
adapter.close()
|
||||||
|
zinc.close_all()
|
||||||
|
|
||||||
|
|
||||||
|
# ── Helpers (local) ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _cross(side, frac=0.1):
|
||||||
|
from malkhut.actions import OrderType
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
221
MALKHUT/malkhut/tests/test_exchange_mechanics.py
Normal file
221
MALKHUT/malkhut/tests/test_exchange_mechanics.py
Normal file
@@ -0,0 +1,221 @@
|
|||||||
|
"""
|
||||||
|
Exchange mechanics — price-time priority, tick/lot rounding, fees,
|
||||||
|
partial fills, IOC/FOK, post-only rejection, reduce-only.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OpenOrderState, OrderBookState,
|
||||||
|
PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
d = dict(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)
|
||||||
|
d.update(kw); return VenueRules(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(**kw):
|
||||||
|
book = kw.get("book", OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0), PriceLevel(49999.0, 2.0)),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0), PriceLevel(50002.0, 2.0)),
|
||||||
|
))
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=kw.get("venue", _venue()), book=book,
|
||||||
|
account=kw.get("account", AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
)),
|
||||||
|
open_orders=kw.get("open_orders", ()),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPostOnlyRejection:
|
||||||
|
def test_buy_at_ask_rejected(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
# price = best_bid - (-10)*tick = 50000 + 1.0 = 50001.0 = best_ask
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, -10, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_sell_at_bid_rejected(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
# price = best_ask + (-10)*tick = 50001 - 1.0 = 50000.0 = best_bid
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.SELL, OrderType.POST_ONLY, -10, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity == s.account.equity
|
||||||
|
|
||||||
|
def test_buy_inside_spread_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
# Should be in open orders (passive placement)
|
||||||
|
assert any(o.side == Side.BUY for o in r.open_orders)
|
||||||
|
|
||||||
|
def test_sell_inside_spread_accepted(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.SELL, OrderType.POST_ONLY, 0, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert any(o.side == Side.SELL for o in r.open_orders)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCrossSpread:
|
||||||
|
def test_cross_buy_fills_at_best_ask(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity < s.account.equity # fees paid
|
||||||
|
|
||||||
|
def test_cross_sell_fills_at_best_bid(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.SELL, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity <= s.account.equity
|
||||||
|
|
||||||
|
def test_cross_spread_taker_fee_applied(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
fee = s.venue.taker_fee_bps
|
||||||
|
assert fee > 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestCancelOrder:
|
||||||
|
def test_cancel_removes_order(self):
|
||||||
|
oo = OpenOrderState(
|
||||||
|
client_order_id="c1", venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type=OrderType.POST_ONLY, price=50000.0,
|
||||||
|
qty=0.001, remaining_qty=0.001, queue_ahead_estimate=0.001,
|
||||||
|
created_ts_ns=1_000_000_000, last_update_ts_ns=1_000_000_000,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
a = FulfilmentAction(ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id="c1")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 0
|
||||||
|
|
||||||
|
def test_cancel_wrong_id_keeps_order(self):
|
||||||
|
oo = OpenOrderState(
|
||||||
|
client_order_id="c1", venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type=OrderType.POST_ONLY, price=50000.0,
|
||||||
|
qty=0.001, remaining_qty=0.001, queue_ahead_estimate=0.001,
|
||||||
|
created_ts_ns=1_000_000_000, last_update_ts_ns=1_000_000_000,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
a = FulfilmentAction(ActionKind.CANCEL, Side.BUY, None, 0, 0.0, 0, cancel_order_id="wrong_id")
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
|
||||||
|
def test_cancel_replace_removes_old_adds_new(self):
|
||||||
|
oo = OpenOrderState(
|
||||||
|
client_order_id="c1", venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type=OrderType.POST_ONLY, price=50000.0,
|
||||||
|
qty=0.001, remaining_qty=0.001, queue_ahead_estimate=0.001,
|
||||||
|
created_ts_ns=1_000_000_000, last_update_ts_ns=1_000_000_000,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
a = FulfilmentAction(
|
||||||
|
ActionKind.CANCEL_REPLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.25,
|
||||||
|
200, cancel_order_id="c1", post_only=True,
|
||||||
|
)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
assert r.open_orders[0].client_order_id != "c1"
|
||||||
|
|
||||||
|
|
||||||
|
class TestAccountUpdate:
|
||||||
|
def test_buy_increases_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert pos is not None
|
||||||
|
assert pos.qty > 0
|
||||||
|
|
||||||
|
def test_sell_decreases_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
# Start with a long position
|
||||||
|
from malkhut.state import PositionState
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=5000.0,
|
||||||
|
positions={"BTCUSDT": pos},
|
||||||
|
))
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.SELL, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
new_pos = r.account.positions.get("BTCUSDT")
|
||||||
|
assert new_pos.qty < pos.qty
|
||||||
|
|
||||||
|
def test_fees_reduce_equity(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.equity < s.account.equity
|
||||||
|
|
||||||
|
def test_no_position_zero_notional(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.account.total_notional == 0.0
|
||||||
|
|
||||||
|
def test_maker_fee_lower_than_taker(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
# Both maker and taker should apply their respective fees
|
||||||
|
a_cross = FulfilmentAction(ActionKind.CROSS_SPREAD, Side.BUY, OrderType.IOC, 0, 0.1, 50)
|
||||||
|
r_cross = cwm.transition(s, (a_cross,))
|
||||||
|
assert s.venue.taker_fee_bps > s.venue.maker_fee_bps
|
||||||
|
|
||||||
|
|
||||||
|
class TestPassivePlacement:
|
||||||
|
def test_passive_order_added_to_book(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 1, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert len(r.open_orders) == 1
|
||||||
|
assert r.open_orders[0].side == Side.BUY
|
||||||
|
|
||||||
|
def test_passive_order_price_correct(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 1, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.open_orders[0].price == 49999.9
|
||||||
|
|
||||||
|
def test_passive_order_symbol_matches(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 0, 0.1, 200, post_only=True)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.open_orders[0].symbol == "BTCUSDT"
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
275
MALKHUT/malkhut/tests/test_extended.py
Normal file
275
MALKHUT/malkhut/tests/test_extended.py
Normal file
@@ -0,0 +1,275 @@
|
|||||||
|
"""
|
||||||
|
Tests for extended counterparties and structured observability.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel,
|
||||||
|
PositionState, Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, AgentRole
|
||||||
|
from malkhut.counterparties_extended import (
|
||||||
|
MomentumTakerPolicy, MeanReversionTakerPolicy,
|
||||||
|
InventoryMarketMakerPolicy, LiquidationFlowPolicy,
|
||||||
|
StaleQuoteAttackerPolicy, extended_counterparty_ecology,
|
||||||
|
)
|
||||||
|
from malkhut.training.structured_obs import StructuredObservability, DecisionMetrics
|
||||||
|
|
||||||
|
|
||||||
|
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")
|
||||||
|
pos = kw.get("position")
|
||||||
|
positions = {"BTCUSDT": pos} if pos else {}
|
||||||
|
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, positions=positions),
|
||||||
|
trade_path=tp, open_orders=kw.get("open_orders", ()),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tp(**kw):
|
||||||
|
d = dict(symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5, dolphin_regime_score=0.5,
|
||||||
|
jericho_signal_strength=0.3, volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1)
|
||||||
|
d.update(kw)
|
||||||
|
return TradePathState(**d)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# EXTENDED COUNTERPARTIES (20 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExtendedCounterparties:
|
||||||
|
def test_momentum_taker_buy_on_upward(self):
|
||||||
|
p = MomentumTakerPolicy()
|
||||||
|
tp = _tp(pnl_bps=50.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
assert a.side == Side.BUY
|
||||||
|
|
||||||
|
def test_momentum_taker_sell_on_downward(self):
|
||||||
|
p = MomentumTakerPolicy()
|
||||||
|
tp = _tp(pnl_bps=-50.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
assert a.side == Side.SELL
|
||||||
|
|
||||||
|
def test_momentum_taker_noop_when_flat(self):
|
||||||
|
p = MomentumTakerPolicy()
|
||||||
|
tp = _tp(pnl_bps=0.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_mean_reversion_buy_on_drop(self):
|
||||||
|
p = MeanReversionTakerPolicy()
|
||||||
|
tp = _tp(pnl_bps=-60.0) # abs(60) > 0.5*100 = 50
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
assert a.side == Side.BUY
|
||||||
|
|
||||||
|
def test_mean_reversion_sell_on_rise(self):
|
||||||
|
p = MeanReversionTakerPolicy()
|
||||||
|
tp = _tp(pnl_bps=60.0) # abs(60) > 0.5*100 = 50
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
assert a.side == Side.SELL
|
||||||
|
|
||||||
|
def test_inventory_mm_reduces_high_inventory(self):
|
||||||
|
p = InventoryMarketMakerPolicy(max_inventory=0.05)
|
||||||
|
pos = PositionState(symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY)
|
||||||
|
s = _state(position=pos)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.PLACE
|
||||||
|
assert a.side == Side.SELL
|
||||||
|
|
||||||
|
def test_inventory_mm_noop_when_balanced(self):
|
||||||
|
p = InventoryMarketMakerPolicy(max_inventory=0.2)
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_liquidation_flow_triggers_on_deep_loss(self):
|
||||||
|
p = LiquidationFlowPolicy(trigger_bps=50.0)
|
||||||
|
tp = _tp(mae_bps=-60.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
assert a.side == Side.SELL
|
||||||
|
assert a.toxicity == 0.9
|
||||||
|
|
||||||
|
def test_liquidation_flow_noop_when_no_loss(self):
|
||||||
|
p = LiquidationFlowPolicy(trigger_bps=50.0)
|
||||||
|
tp = _tp(mae_bps=-10.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_stale_quote_attacker_attacks(self):
|
||||||
|
p = StaleQuoteAttackerPolicy()
|
||||||
|
from malkhut.state import OpenOrderState
|
||||||
|
oo = OpenOrderState(client_order_id="c1", venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type=OrderType.LIMIT, price=50000.0,
|
||||||
|
qty=0.001, remaining_qty=0.001, queue_ahead_estimate=0.001,
|
||||||
|
created_ts_ns=1, last_update_ts_ns=1)
|
||||||
|
s = _state(open_orders=(oo,))
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.CROSS_SPREAD
|
||||||
|
|
||||||
|
def test_stale_quote_attacker_noop_when_no_orders(self):
|
||||||
|
p = StaleQuoteAttackerPolicy()
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
a = p.rollout_action(s, rng)
|
||||||
|
assert a.kind == ActionKind.NOOP
|
||||||
|
|
||||||
|
def test_extended_ecology_has_9_agents(self):
|
||||||
|
eco = extended_counterparty_ecology()
|
||||||
|
assert len(eco) == 9
|
||||||
|
|
||||||
|
def test_extended_ecology_unique_roles(self):
|
||||||
|
eco = extended_counterparty_ecology()
|
||||||
|
roles = [p.role for p in eco]
|
||||||
|
assert len(set(roles)) == 9
|
||||||
|
|
||||||
|
def test_all_agents_produce_valid_actions(self):
|
||||||
|
eco = extended_counterparty_ecology()
|
||||||
|
s = _state()
|
||||||
|
rng = random.Random(42)
|
||||||
|
for agent in eco:
|
||||||
|
a = agent.rollout_action(s, rng)
|
||||||
|
assert a.kind is not None
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# STRUCTURED OBSERVABILITY (15 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStructuredObservability:
|
||||||
|
def test_record_decision(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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")
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000, regime="normal")
|
||||||
|
assert so.total_decisions == 1
|
||||||
|
|
||||||
|
def test_feature_importance(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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):
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000, regime="normal")
|
||||||
|
importance = so.get_feature_importance(top_n=5)
|
||||||
|
assert len(importance) > 0
|
||||||
|
|
||||||
|
def test_regime_approval_rate(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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})
|
||||||
|
# Approved
|
||||||
|
decision = RiskDecision(approved=True, action=a, reason="ok")
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000, regime="normal")
|
||||||
|
# Rejected
|
||||||
|
decision2 = RiskDecision(approved=False, action=a, reason="kill")
|
||||||
|
so.record_decision(s, planned, decision2, plan_ns=1000, regime="normal")
|
||||||
|
rate = so.get_regime_approval_rate("normal")
|
||||||
|
assert rate == 0.5
|
||||||
|
|
||||||
|
def test_avg_latency(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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")
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000)
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=2000)
|
||||||
|
assert so.avg_latency_ns == 1500.0
|
||||||
|
|
||||||
|
def test_avg_entropy(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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")
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000)
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000)
|
||||||
|
assert so.avg_entropy == 0.5
|
||||||
|
|
||||||
|
def test_total_decisions(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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(20):
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000)
|
||||||
|
assert so.total_decisions == 20
|
||||||
|
|
||||||
|
def test_feature_importance_sorted(self):
|
||||||
|
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
|
||||||
|
so = StructuredObservability()
|
||||||
|
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):
|
||||||
|
so.record_decision(s, planned, decision, plan_ns=1000)
|
||||||
|
importance = so.get_feature_importance(top_n=5)
|
||||||
|
for i in range(len(importance) - 1):
|
||||||
|
assert importance[i][1] >= importance[i+1][1]
|
||||||
|
|
||||||
|
|
||||||
|
import random
|
||||||
|
from malkhut.actions import OrderType
|
||||||
157
MALKHUT/malkhut/tests/test_fuzz.py
Normal file
157
MALKHUT/malkhut/tests/test_fuzz.py
Normal file
@@ -0,0 +1,157 @@
|
|||||||
|
"""
|
||||||
|
Fuzz testing — random state/action sequences through CWM.
|
||||||
|
|
||||||
|
Verifies CWM never crashes, always produces valid outputs,
|
||||||
|
and preserves invariants under random stress.
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, IntentKind, MarketWorldState,
|
||||||
|
Mode, OpenOrderState, OrderBookState, PriceLevel, PositionState,
|
||||||
|
Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
|
||||||
|
|
||||||
|
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 _random_book(rng):
|
||||||
|
bid = rng.uniform(100.0, 100000.0)
|
||||||
|
ask = bid + rng.uniform(0.1, 100.0)
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=rng.randint(1, 2**62), symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, rng.uniform(0.001, 10.0)),),
|
||||||
|
asks=(PriceLevel(ask, rng.uniform(0.001, 10.0)),),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _random_action(rng):
|
||||||
|
kind = rng.choice(list(ActionKind))
|
||||||
|
side = rng.choice([Side.BUY, Side.SELL, None])
|
||||||
|
ot = rng.choice(list(OrderType))
|
||||||
|
offset = rng.randint(-20, 20)
|
||||||
|
frac = rng.uniform(0.0, 1.0)
|
||||||
|
ttl = rng.randint(0, 5000)
|
||||||
|
cancel_id = f"r_{rng.randint(0, 1000)}" if kind in (ActionKind.CANCEL, ActionKind.CANCEL_REPLACE) else None
|
||||||
|
return FulfilmentAction(
|
||||||
|
kind=kind, side=side, order_type=ot,
|
||||||
|
price_ticks_from_best=offset, qty_fraction=frac, ttl_ms=ttl,
|
||||||
|
cancel_order_id=cancel_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _random_state(rng):
|
||||||
|
book = _random_book(rng)
|
||||||
|
equity = rng.uniform(100.0, 100000.0)
|
||||||
|
pos_qty = rng.choice([0.0, rng.uniform(-0.5, 0.5)])
|
||||||
|
pos = None
|
||||||
|
positions = {}
|
||||||
|
if pos_qty != 0.0:
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=pos_qty, avg_entry=book.mid,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=abs(pos_qty * book.mid) / max(equity, 1.0),
|
||||||
|
side=Side.BUY if pos_qty > 0 else Side.SELL,
|
||||||
|
)
|
||||||
|
positions["BTCUSDT"] = pos
|
||||||
|
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=rng.randint(1, 2**62), mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=_venue(), book=book,
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=rng.randint(1, 2**62), equity=equity,
|
||||||
|
wallet_balance=equity, available_balance=equity,
|
||||||
|
margin_used=0.0, total_notional=abs(pos_qty * book.mid),
|
||||||
|
positions=positions,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCWMFuzz:
|
||||||
|
def test_100_random_noop_transitions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(42)
|
||||||
|
for _ in range(100):
|
||||||
|
s = _random_state(rng)
|
||||||
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.ts_ns >= s.ts_ns
|
||||||
|
assert r.account.equity >= 0
|
||||||
|
|
||||||
|
def test_100_random_action_transitions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(123)
|
||||||
|
for _ in range(100):
|
||||||
|
s = _random_state(rng)
|
||||||
|
a = _random_action(rng)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.ts_ns >= s.ts_ns
|
||||||
|
assert r.account.equity >= 0
|
||||||
|
|
||||||
|
def test_100_random_multi_counterparty(self):
|
||||||
|
from malkhut.counterparties import (
|
||||||
|
ToxicTakerPolicy, PassiveMakerPolicy, NoiseTraderPolicy,
|
||||||
|
)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(456)
|
||||||
|
cps = [ToxicTakerPolicy(), PassiveMakerPolicy(), NoiseTraderPolicy()]
|
||||||
|
for _ in range(100):
|
||||||
|
s = _random_state(rng)
|
||||||
|
a = _random_action(rng)
|
||||||
|
cp_actions = tuple(cp.rollout_action(s, rng) for cp in cps)
|
||||||
|
r = cwm.transition(s, (a, *cp_actions))
|
||||||
|
assert r.ts_ns >= s.ts_ns
|
||||||
|
# Equity can go negative with aggressive counterparties (realistic)
|
||||||
|
assert isinstance(r.account.equity, float)
|
||||||
|
|
||||||
|
def test_50_sequential_transitions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(789)
|
||||||
|
s = _random_state(rng)
|
||||||
|
for i in range(50):
|
||||||
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
s = cwm.transition(s, (a,))
|
||||||
|
assert s.account.equity >= 0
|
||||||
|
|
||||||
|
def test_50_sequential_with_actions(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(101)
|
||||||
|
s = _random_state(rng)
|
||||||
|
for i in range(50):
|
||||||
|
a = _random_action(rng)
|
||||||
|
s = cwm.transition(s, (a,))
|
||||||
|
# Equity can go negative (realistic: overleveraged position)
|
||||||
|
assert isinstance(s.account.equity, float)
|
||||||
|
|
||||||
|
def test_input_never_mutated(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(202)
|
||||||
|
for _ in range(50):
|
||||||
|
s = _random_state(rng)
|
||||||
|
orig_ts = s.ts_ns
|
||||||
|
orig_equity = s.account.equity
|
||||||
|
a = _random_action(rng)
|
||||||
|
cwm.transition(s, (a,))
|
||||||
|
assert s.ts_ns == orig_ts
|
||||||
|
assert s.account.equity == orig_equity
|
||||||
|
|
||||||
|
def test_determinism_under_fuzz(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rng = random.Random(303)
|
||||||
|
for _ in range(50):
|
||||||
|
s = _random_state(rng)
|
||||||
|
a = _random_action(rng)
|
||||||
|
r1 = cwm.transition(s, (a,))
|
||||||
|
r2 = cwm.transition(s, (a,))
|
||||||
|
assert r1.ts_ns == r2.ts_ns
|
||||||
|
assert r1.account.equity == r2.account.equity
|
||||||
249
MALKHUT/malkhut/tests/test_generator.py
Normal file
249
MALKHUT/malkhut/tests/test_generator.py
Normal file
@@ -0,0 +1,249 @@
|
|||||||
|
"""
|
||||||
|
Tests for strategy generator — genetic programming for strategy evolution.
|
||||||
|
|
||||||
|
Verifies:
|
||||||
|
- Genome creation and representation
|
||||||
|
- Crossover produces valid offspring
|
||||||
|
- Mutation produces valid variants
|
||||||
|
- Tournament selection preferentially selects fitter genomes
|
||||||
|
- Population initialization includes baseline
|
||||||
|
- Evolution improves fitness over generations
|
||||||
|
- Successful strategies are added to pool
|
||||||
|
- Diverse strategies are returned
|
||||||
|
- Hardcoded baseline is never replaced
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
from malkhut.training.generator import (
|
||||||
|
StrategyGenome, StrategyType, GeneticOperators,
|
||||||
|
StrategyEvaluator, StrategyGenerator, GeneratorConfig,
|
||||||
|
)
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec, SelfPlayPool
|
||||||
|
from malkhut.training.registry import PolicyRegistry
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. STRATEGY GENOME
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStrategyGenome:
|
||||||
|
def test_construction(self):
|
||||||
|
genome = StrategyGenome(
|
||||||
|
strategy_type=StrategyType.SM_MCTS,
|
||||||
|
params=_baseline(),
|
||||||
|
)
|
||||||
|
assert genome.strategy_type == StrategyType.SM_MCTS
|
||||||
|
assert genome.generation == 0
|
||||||
|
assert genome.fitness == 0.0
|
||||||
|
|
||||||
|
def test_genome_id(self):
|
||||||
|
genome = StrategyGenome(
|
||||||
|
strategy_type=StrategyType.SM_MCTS,
|
||||||
|
params=_baseline(version="v1"),
|
||||||
|
generation=3,
|
||||||
|
)
|
||||||
|
assert "SM_MCTS" in genome.genome_id
|
||||||
|
assert "v1" in genome.genome_id
|
||||||
|
assert "3" in genome.genome_id
|
||||||
|
|
||||||
|
def test_frozen(self):
|
||||||
|
genome = StrategyGenome(
|
||||||
|
strategy_type=StrategyType.SM_MCTS,
|
||||||
|
params=_baseline(),
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
genome.fitness = 10.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. GENETIC OPERATORS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestGeneticOperators:
|
||||||
|
def _ops(self):
|
||||||
|
return GeneticOperators(codec=CMAParameterCodec())
|
||||||
|
|
||||||
|
def test_crossover_produces_valid_child(self):
|
||||||
|
ops = self._ops()
|
||||||
|
p1 = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline(version="p1"))
|
||||||
|
p2 = StrategyGenome(strategy_type=StrategyType.UCB1, params=_baseline(version="p2"))
|
||||||
|
child = ops.crossover(p1, p2, random.Random(42))
|
||||||
|
assert isinstance(child, StrategyGenome)
|
||||||
|
assert child.strategy_type in (StrategyType.SM_MCTS, StrategyType.UCB1)
|
||||||
|
assert child.generation == 1
|
||||||
|
assert len(child.parent_ids) == 2
|
||||||
|
|
||||||
|
def test_crossover_inherits_fitter_parent_type(self):
|
||||||
|
ops = self._ops()
|
||||||
|
p1 = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline(), fitness=10.0)
|
||||||
|
p2 = StrategyGenome(strategy_type=StrategyType.UCB1, params=_baseline(), fitness=5.0)
|
||||||
|
child = ops.crossover(p1, p2, random.Random(42))
|
||||||
|
assert child.strategy_type == StrategyType.SM_MCTS
|
||||||
|
|
||||||
|
def test_mutation_produces_valid_child(self):
|
||||||
|
ops = self._ops()
|
||||||
|
parent = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline())
|
||||||
|
child = ops.mutate(parent, random.Random(42))
|
||||||
|
assert isinstance(child, StrategyGenome)
|
||||||
|
assert child.generation == 1
|
||||||
|
assert len(child.parent_ids) == 1
|
||||||
|
|
||||||
|
def test_mutation_different_params(self):
|
||||||
|
ops = self._ops()
|
||||||
|
parent = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline())
|
||||||
|
child = ops.mutate(parent, random.Random(42))
|
||||||
|
# With high mutation rate, params should differ
|
||||||
|
assert child.params.ucb_c != parent.params.ucb_c or child.params.root_temperature != parent.params.root_temperature
|
||||||
|
|
||||||
|
def test_tournament_select_prefers_fitter(self):
|
||||||
|
ops = self._ops()
|
||||||
|
pop = [
|
||||||
|
StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline(), fitness=1.0),
|
||||||
|
StrategyGenome(strategy_type=StrategyType.UCB1, params=_baseline(), fitness=10.0),
|
||||||
|
StrategyGenome(strategy_type=StrategyType.GREEDY, params=_baseline(), fitness=5.0),
|
||||||
|
]
|
||||||
|
wins = 0
|
||||||
|
for i in range(200):
|
||||||
|
winner = ops.tournament_select(pop, tournament_size=2, rng=random.Random(i))
|
||||||
|
if winner.fitness == 10.0:
|
||||||
|
wins += 1
|
||||||
|
# With 3 genomes and tournament_size=2, fitter wins ~50% (vs 25% random)
|
||||||
|
assert wins > 60
|
||||||
|
|
||||||
|
def test_random_genome_valid(self):
|
||||||
|
ops = self._ops()
|
||||||
|
genome = ops.random_genome(rng=random.Random(42))
|
||||||
|
assert isinstance(genome, StrategyGenome)
|
||||||
|
assert genome.strategy_type in list(StrategyType)
|
||||||
|
assert genome.generation == 0
|
||||||
|
|
||||||
|
def test_random_genome_respects_type(self):
|
||||||
|
ops = self._ops()
|
||||||
|
genome = ops.random_genome(strategy_type=StrategyType.THOMPSON_SAMPLING, rng=random.Random(42))
|
||||||
|
assert genome.strategy_type == StrategyType.THOMPSON_SAMPLING
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. STRATEGY GENERATOR
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStrategyGenerator:
|
||||||
|
def test_initialize_population(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=5, generations=1))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.initialize_population(_baseline(), scenarios)
|
||||||
|
assert gen.population_size == 5
|
||||||
|
# First should be baseline
|
||||||
|
assert gen.population[0].strategy_type == StrategyType.SM_MCTS
|
||||||
|
|
||||||
|
def test_baseline_always_first(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=10, generations=1))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.initialize_population(_baseline(), scenarios)
|
||||||
|
# Baseline should be in population
|
||||||
|
types = [g.strategy_type for g in gen.population]
|
||||||
|
assert StrategyType.SM_MCTS in types
|
||||||
|
|
||||||
|
def test_evolve_returns_population(self):
|
||||||
|
config = GeneratorConfig(population_size=5, generations=2)
|
||||||
|
gen = StrategyGenerator(config=config)
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
result = gen.evolve(_baseline(), scenarios)
|
||||||
|
# Population may be slightly larger due to baseline preservation
|
||||||
|
assert len(result) >= 5
|
||||||
|
assert all(isinstance(g, StrategyGenome) for g in result)
|
||||||
|
|
||||||
|
def test_evolve_improves_over_generations(self):
|
||||||
|
config = GeneratorConfig(population_size=8, generations=3, elitism_count=2)
|
||||||
|
gen = StrategyGenerator(config=config)
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
# Population should have diverse fitness
|
||||||
|
fitnesses = [g.fitness for g in gen.population]
|
||||||
|
assert len(set(fitnesses)) > 1 # not all same
|
||||||
|
|
||||||
|
def test_get_successful_strategies(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=5, generations=1))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
successful = gen.get_successful_strategies(min_fitness=-1000)
|
||||||
|
assert len(successful) > 0
|
||||||
|
|
||||||
|
def test_get_diverse_strategies(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=10, generations=2))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
diverse = gen.get_diverse_strategies(n=3)
|
||||||
|
assert len(diverse) == 3
|
||||||
|
|
||||||
|
def test_add_to_pool(self):
|
||||||
|
gen = StrategyGenerator()
|
||||||
|
genome = StrategyGenome(
|
||||||
|
strategy_type=StrategyType.SM_MCTS,
|
||||||
|
params=_baseline(version="test"),
|
||||||
|
fitness=10.0,
|
||||||
|
)
|
||||||
|
gen.add_to_pool(genome)
|
||||||
|
assert len(gen._pool.policies()) == 1
|
||||||
|
|
||||||
|
def test_baseline_never_replaced(self):
|
||||||
|
"""The hardcoded baseline should always be in the population."""
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=5, generations=3))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
# At least one SM_MCTS should exist (the baseline)
|
||||||
|
assert any(g.strategy_type == StrategyType.SM_MCTS for g in gen.population)
|
||||||
|
|
||||||
|
def test_best_fitness_tracked(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=5, generations=1))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
assert gen.best_fitness > -float("inf")
|
||||||
|
|
||||||
|
def test_history_tracked(self):
|
||||||
|
gen = StrategyGenerator(config=GeneratorConfig(population_size=5, generations=2))
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
gen.evolve(_baseline(), scenarios)
|
||||||
|
assert len(gen._history) > 0
|
||||||
366
MALKHUT/malkhut/tests/test_harness.py
Normal file
366
MALKHUT/malkhut/tests/test_harness.py
Normal file
@@ -0,0 +1,366 @@
|
|||||||
|
"""
|
||||||
|
Harness tests — verify the system CAN learn and CAN register improvement.
|
||||||
|
|
||||||
|
These tests prove the system is CAPABLE of learning:
|
||||||
|
1. Different parameters produce different actions
|
||||||
|
2. Different actions produce different PnL
|
||||||
|
3. CMA-ES can find better parameters
|
||||||
|
4. Genetic operators produce meaningful diversity
|
||||||
|
5. Score improves over generations
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
|
||||||
|
|
||||||
|
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, 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),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _intent():
|
||||||
|
return ExecutionIntent(
|
||||||
|
intent_id="test", ts_ns=1, 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 _params(**kw):
|
||||||
|
d = dict(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2, rollout_depth=2,
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 1: Different parameters produce different actions
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestParameterSensitivity:
|
||||||
|
def test_ucb_c_affects_exploration(self):
|
||||||
|
"""Different UCB_c should produce different action distributions."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
intent = _intent()
|
||||||
|
s_with_intent = MarketWorldState(
|
||||||
|
ts_ns=1, 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
actions_low = []
|
||||||
|
actions_high = []
|
||||||
|
for seed in range(30): # more samples for reliability
|
||||||
|
p_low = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=seed)
|
||||||
|
r_low = p_low.plan(s_with_intent, _params(ucb_c=0.2), budget_ms=10)
|
||||||
|
actions_low.append(r_low.selected_action.kind)
|
||||||
|
|
||||||
|
p_high = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=seed)
|
||||||
|
r_high = p_high.plan(s_with_intent, _params(ucb_c=3.0), budget_ms=10)
|
||||||
|
actions_high.append(r_high.selected_action.kind)
|
||||||
|
|
||||||
|
# Different exploration should produce different distributions
|
||||||
|
low_types = set(actions_low)
|
||||||
|
high_types = set(actions_high)
|
||||||
|
# Either different action sets OR different distribution within same set
|
||||||
|
assert low_types != high_types or len(low_types) > 1
|
||||||
|
|
||||||
|
def test_temperature_affects_distribution(self):
|
||||||
|
"""Different temperatures should produce different probability distributions."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
intent = _intent()
|
||||||
|
s_with_intent = MarketWorldState(
|
||||||
|
ts_ns=1, 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
probs_by_temp = {}
|
||||||
|
for temp in [0.1, 0.5, 1.0, 2.0]:
|
||||||
|
p = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42)
|
||||||
|
r = p.plan(s_with_intent, _params(root_temperature=temp), budget_ms=10)
|
||||||
|
probs_by_temp[temp] = tuple(r.probabilities)
|
||||||
|
|
||||||
|
# Different temperatures should produce different distributions
|
||||||
|
unique_dists = set(probs_by_temp.values())
|
||||||
|
assert len(unique_dists) > 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 2: Different actions produce different PnL
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestActionPnLDifferentiation:
|
||||||
|
def test_cross_vs_noop_different_equity(self):
|
||||||
|
"""CROSS and NOOP should produce different equity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r_cross = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
r_noop = cwm.transition(s, (_noop(),))
|
||||||
|
assert r_cross.account.equity != r_noop.account.equity
|
||||||
|
|
||||||
|
def test_cross_vs_place_different_equity(self):
|
||||||
|
"""CROSS and PLACE should produce different equity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r_cross = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
r_place = cwm.transition(s, (_place(Side.BUY, offset=0, frac=0.1),))
|
||||||
|
assert r_cross.account.equity != r_place.account.equity
|
||||||
|
|
||||||
|
def test_different_cross_sizes_different_equity(self):
|
||||||
|
"""Different cross sizes should produce different equity."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
r_small = cwm.transition(s, (_cross(Side.BUY, 0.01),))
|
||||||
|
r_large = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
assert r_small.account.equity != r_large.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 3: CMA-ES can find better parameters
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCMAESLearning:
|
||||||
|
def test_cma_es_finds_better_params(self):
|
||||||
|
"""CMA-ES should find parameters that produce different (hopefully better) scores."""
|
||||||
|
from malkhut.training.cma_trainer import CMAESTrainer, CMAParameterCodec, PolicyEvaluator, ScenarioFactory
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
|
||||||
|
# Run CMA-ES for a few evaluations
|
||||||
|
import cma
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
es = cma.CMAEvolutionStrategy(x0, 0.30, {
|
||||||
|
"bounds": [lows, highs], "popsize": 5, "seed": 42, "verbose": -9,
|
||||||
|
})
|
||||||
|
|
||||||
|
scores = []
|
||||||
|
for _ in range(3):
|
||||||
|
xs = es.ask()
|
||||||
|
for x in xs:
|
||||||
|
candidate = codec.decode(x, version="test")
|
||||||
|
score, _ = evaluator.evaluate_candidate(
|
||||||
|
params=candidate, scenarios=scenarios, rng_seed=42,
|
||||||
|
)
|
||||||
|
scores.append(score)
|
||||||
|
es.tell(xs, [-s for s in scores[-5:]])
|
||||||
|
|
||||||
|
# Scores should vary (not all identical)
|
||||||
|
assert len(set(scores)) > 1, "All scores identical — system can't learn"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 4: Genetic operators produce meaningful diversity
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestGeneticDiversity:
|
||||||
|
def test_crossover_produces_different_children(self):
|
||||||
|
"""Crossover of two parents should produce different offspring."""
|
||||||
|
from malkhut.training.generator import GeneticOperators, StrategyGenome
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec
|
||||||
|
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
ops = GeneticOperators(codec=codec)
|
||||||
|
|
||||||
|
p1 = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline(version="p1"))
|
||||||
|
p2 = StrategyGenome(strategy_type=StrategyType.UCB1, params=_baseline(version="p2"))
|
||||||
|
|
||||||
|
children = []
|
||||||
|
for i in range(10):
|
||||||
|
child = ops.crossover(p1, p2, random.Random(i))
|
||||||
|
children.append(child)
|
||||||
|
|
||||||
|
# Children should have different params
|
||||||
|
unique_versions = set(c.params.version for c in children)
|
||||||
|
assert len(unique_versions) > 1
|
||||||
|
|
||||||
|
def test_mutation_produces_different_children(self):
|
||||||
|
"""Mutation should produce different offspring."""
|
||||||
|
from malkhut.training.generator import GeneticOperators, StrategyGenome
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec
|
||||||
|
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
ops = GeneticOperators(codec=codec, mutation_rate=0.5) # high mutation
|
||||||
|
|
||||||
|
parent = StrategyGenome(strategy_type=StrategyType.SM_MCTS, params=_baseline())
|
||||||
|
|
||||||
|
children = []
|
||||||
|
for i in range(10):
|
||||||
|
child = ops.mutate(parent, random.Random(i))
|
||||||
|
children.append(child)
|
||||||
|
|
||||||
|
# Children should have different params
|
||||||
|
unique_versions = set(c.params.version for c in children)
|
||||||
|
assert len(unique_versions) > 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 5: Score improves over generations
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestScoreImprovement:
|
||||||
|
def test_score_varies_across_generations(self):
|
||||||
|
"""Scores should vary across generations (not all identical)."""
|
||||||
|
from malkhut.training.cma_trainer import CMAESTrainer, CMAParameterCodec, PolicyEvaluator, ScenarioFactory
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=5)
|
||||||
|
|
||||||
|
# Run for a few generations
|
||||||
|
import cma
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
es = cma.CMAEvolutionStrategy(x0, 0.30, {
|
||||||
|
"bounds": [lows, highs], "popsize": 5, "seed": 42, "verbose": -9,
|
||||||
|
})
|
||||||
|
|
||||||
|
all_scores = []
|
||||||
|
for gen in range(3):
|
||||||
|
xs = es.ask()
|
||||||
|
gen_scores = []
|
||||||
|
for x in xs:
|
||||||
|
candidate = codec.decode(x, version=f"gen{gen}")
|
||||||
|
score, _ = evaluator.evaluate_candidate(
|
||||||
|
params=candidate, scenarios=scenarios, rng_seed=42 + gen,
|
||||||
|
)
|
||||||
|
gen_scores.append(score)
|
||||||
|
all_scores.append(max(gen_scores))
|
||||||
|
es.tell(xs, [-s for s in gen_scores])
|
||||||
|
|
||||||
|
# Scores should vary (not all identical)
|
||||||
|
assert len(set(all_scores)) > 1, f"All generation scores identical: {all_scores}"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HARNESS 6: End-to-end training produces improvement
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEndToEndLearning:
|
||||||
|
def test_training_pipeline_produces_improvement(self):
|
||||||
|
"""Training pipeline should produce improvement over baseline."""
|
||||||
|
from malkhut.training.pipeline import TrainingPipeline, PipelineConfig
|
||||||
|
from malkhut.training.registry import PolicyRegistry
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
registry = PolicyRegistry(store=store)
|
||||||
|
|
||||||
|
cfg = PipelineConfig(max_generations=3, max_evals_per_generation=5, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=cfg, registry=registry, log_path="/dev/null")
|
||||||
|
|
||||||
|
result = pipeline.run(incumbent=_baseline(), symbols=("BTCUSDT",))
|
||||||
|
|
||||||
|
# Should have run some generations
|
||||||
|
assert result.generations_run >= 1
|
||||||
|
# Should have some events
|
||||||
|
assert len(result.events) > 0
|
||||||
|
# Best score should be a valid number
|
||||||
|
assert isinstance(result.best_score, float)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# HELPERS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
from malkhut.training.generator import StrategyType, SelfPlayPool
|
||||||
|
|
||||||
|
|
||||||
|
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 _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _cross(side, frac):
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
|
|
||||||
|
|
||||||
|
def _place(side, offset=0, frac=0.1):
|
||||||
|
return FulfilmentAction(ActionKind.PLACE, side, OrderType.LIMIT, offset, frac, 200)
|
||||||
164
MALKHUT/malkhut/tests/test_hypothesis_properties.py
Normal file
164
MALKHUT/malkhut/tests/test_hypothesis_properties.py
Normal file
@@ -0,0 +1,164 @@
|
|||||||
|
"""
|
||||||
|
Property-based tests using Hypothesis.
|
||||||
|
|
||||||
|
Invariant tests:
|
||||||
|
- CWM always produces valid states
|
||||||
|
- Planner always returns probability distribution summing to 1
|
||||||
|
- Codec always produces valid params
|
||||||
|
- Risk gate always returns valid decisions
|
||||||
|
"""
|
||||||
|
import hypothesis
|
||||||
|
from hypothesis import given, strategies as st, assume, settings
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PriceLevel, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM, materialize_price_from_action
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType, Side
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
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 _book(bid_price, ask_price):
|
||||||
|
assume(bid_price < ask_price)
|
||||||
|
assume(bid_price > 0)
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid_price, 1.0),),
|
||||||
|
asks=(PriceLevel(ask_price, 1.0),),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _params():
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="hypo", ucb_c=1.414, max_sims=32, max_depth=2,
|
||||||
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCWMProperties:
|
||||||
|
@settings(max_examples=50, deadline=None)
|
||||||
|
@given(bid=st.floats(min_value=1.0, max_value=100000.0),
|
||||||
|
ask=st.floats(min_value=1.0, max_value=100000.0))
|
||||||
|
def test_transition_never_crashes(self, bid, ask):
|
||||||
|
assume(bid < ask)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
book = _book(bid, ask)
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=_venue(), book=book,
|
||||||
|
account=AccountState(ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0),
|
||||||
|
)
|
||||||
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
assert r.ts_ns >= s.ts_ns
|
||||||
|
assert r.account.equity >= 0
|
||||||
|
|
||||||
|
@settings(max_examples=50, deadline=None)
|
||||||
|
@given(bid=st.floats(min_value=100.0, max_value=100000.0),
|
||||||
|
ask=st.floats(min_value=100.0, max_value=100000.0))
|
||||||
|
def test_book_invariants(self, bid, ask):
|
||||||
|
assume(bid < ask)
|
||||||
|
book = _book(bid, ask)
|
||||||
|
assert book.best_bid == bid
|
||||||
|
assert book.best_ask == ask
|
||||||
|
assert book.spread == ask - bid
|
||||||
|
assert book.spread_bps > 0
|
||||||
|
assert book.mid == (bid + ask) / 2
|
||||||
|
|
||||||
|
@settings(max_examples=50, deadline=None)
|
||||||
|
@given(bid=st.floats(min_value=100.0, max_value=100000.0),
|
||||||
|
ask=st.floats(min_value=100.0, max_value=100000.0),
|
||||||
|
offset=st.integers(min_value=0, max_value=20))
|
||||||
|
def test_price_materialization_bounded(self, bid, ask, offset):
|
||||||
|
assume(bid < ask)
|
||||||
|
book = _book(bid, ask)
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=1, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(), book=book,
|
||||||
|
account=AccountState(ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0),
|
||||||
|
)
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, offset, 0.1, 200)
|
||||||
|
price = materialize_price_from_action(s, a)
|
||||||
|
assert price is not None
|
||||||
|
assert price <= book.best_bid # buy offset should be <= best bid
|
||||||
|
|
||||||
|
|
||||||
|
class TestCodecProperties:
|
||||||
|
def test_decode_always_returns_valid_params(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
import random
|
||||||
|
for _ in range(50):
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
x = [random.uniform(lo, hi) for lo, hi in zip(lows, highs)]
|
||||||
|
p = codec.decode(x, f"rand_{_}")
|
||||||
|
assert isinstance(p, FulfilmentPolicyParams)
|
||||||
|
assert p.version.startswith("rand_")
|
||||||
|
|
||||||
|
def test_decode_bounds_respected(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
import random
|
||||||
|
for _ in range(50):
|
||||||
|
x = [random.uniform(lo, hi) for lo, hi in zip(lows, highs)]
|
||||||
|
p = codec.decode(x, "b")
|
||||||
|
for spec in codec.SPECS:
|
||||||
|
val = getattr(p, spec.name)
|
||||||
|
if spec.kind == "float":
|
||||||
|
assert spec.low - 1e-9 <= val <= spec.high + 1e-9
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlannerProperties:
|
||||||
|
@settings(max_examples=30, deadline=None)
|
||||||
|
@given(seed=st.integers(min_value=0, max_value=2**31))
|
||||||
|
def test_planner_always_returns_valid_distribution(self, seed):
|
||||||
|
from malkhut.state import ExecutionIntent, IntentKind
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="h", 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="hypo",
|
||||||
|
)
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=_book(50000.0, 50001.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,
|
||||||
|
)
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=seed,
|
||||||
|
)
|
||||||
|
result = planner.plan(root_state=s, params=_params(), budget_ms=10)
|
||||||
|
total = sum(result.probabilities)
|
||||||
|
assert abs(total - 1.0) < 1e-6
|
||||||
|
assert all(p >= 0 for p in result.probabilities)
|
||||||
622
MALKHUT/malkhut/tests/test_microstructure.py
Normal file
622
MALKHUT/malkhut/tests/test_microstructure.py
Normal file
@@ -0,0 +1,622 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive tests for all 10 new CWM/training modules.
|
||||||
|
|
||||||
|
Covers: queue model, adverse selection, latency, spread dynamics,
|
||||||
|
volatility clustering, execution quality, risk-adjusted returns,
|
||||||
|
multi-level book, multi-asset correlation.
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
from malkhut.cwm.queue_model import (
|
||||||
|
QueuePositionModel, QueueState, estimate_queue_position,
|
||||||
|
compute_fill_probability, compute_queue_adverse_selection,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.adverse_selection import (
|
||||||
|
AdverseSelectionModel, AdverseSelectionCost, compute_adverse_selection_cost,
|
||||||
|
compute_toxic_fill_ratio, optimal_quote_offset,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.latency_model import (
|
||||||
|
LatencyModel, LatencyState, simulate_feed_latency, simulate_order_latency,
|
||||||
|
compute_latency_impact,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.spread_dynamics import (
|
||||||
|
SpreadDynamicsModel, compute_spread_tendency, predict_spread,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.volatility import (
|
||||||
|
VolatilityClusteringModel, compute_volatility_regime, predict_volatility,
|
||||||
|
)
|
||||||
|
from malkhut.training.execution_quality import (
|
||||||
|
ExecutionQualityTracker, ExecutionQualityReport, RiskAdjustedReturns,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.multi_level import (
|
||||||
|
MultiLevelBookModel, compute_net_order_flow, compute_book_imbalance_weighted,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.correlation import (
|
||||||
|
MultiAssetCorrelationModel, compute_rolling_correlation, compute_correlation_regime,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# QUEUE MODEL (15 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestQueueModel:
|
||||||
|
def test_fill_probability_basic(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
fp = qm.estimate_fill_probability(0.001, 1.0, toxicity=0.0)
|
||||||
|
assert 0.0 <= fp <= 1.0
|
||||||
|
|
||||||
|
def test_fill_probability_zero_qty(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
fp = qm.estimate_fill_probability(0.0, 1.0)
|
||||||
|
assert fp == 0.0
|
||||||
|
|
||||||
|
def test_fill_probability_zero_rate(self):
|
||||||
|
qm = QueuePositionModel(default_trade_rate=0.0)
|
||||||
|
fp = qm.estimate_fill_probability(0.001, 1.0)
|
||||||
|
assert fp == 0.0
|
||||||
|
|
||||||
|
def test_fill_probability_increases_with_rate(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
fp1 = qm.estimate_fill_probability(0.001, 1.0, recent_trade_rate=0.1)
|
||||||
|
fp2 = qm.estimate_fill_probability(0.001, 1.0, recent_trade_rate=1.0)
|
||||||
|
assert fp2 > fp1
|
||||||
|
|
||||||
|
def test_fill_probability_decreases_with_toxicity(self):
|
||||||
|
"""Toxicity increases fill rate (toxic flow fills queue faster) — bad for us."""
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
# Use short time horizon so probabilities don't saturate to 1.0
|
||||||
|
fp1 = qm.estimate_fill_probability(0.001, 1.0, toxicity=0.0, time_horizon_s=1.0)
|
||||||
|
fp2 = qm.estimate_fill_probability(0.001, 1.0, toxicity=0.9, time_horizon_s=1.0)
|
||||||
|
assert fp2 > fp1 # toxicity increases fill rate (adverse for maker)
|
||||||
|
|
||||||
|
def test_queue_position_estimation(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
pos = qm.estimate_queue_position(0.001, 1.0)
|
||||||
|
assert pos >= 0
|
||||||
|
|
||||||
|
def test_queue_position_zero_level(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
pos = qm.estimate_queue_position(0.001, 0.0)
|
||||||
|
assert pos == 0.0
|
||||||
|
|
||||||
|
def test_adverse_selection_risk_front(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
risk = qm.adverse_selection_risk(0, toxicity=0.5, spread_bps=2.0)
|
||||||
|
assert risk > 0
|
||||||
|
|
||||||
|
def test_adverse_selection_risk_back(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
risk_front = qm.adverse_selection_risk(0, toxicity=0.5)
|
||||||
|
risk_back = qm.adverse_selection_risk(10, toxicity=0.5)
|
||||||
|
assert risk_front > risk_back
|
||||||
|
|
||||||
|
def test_adverse_selection_increases_with_toxicity(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
r1 = qm.adverse_selection_risk(5, toxicity=0.1)
|
||||||
|
r2 = qm.adverse_selection_risk(5, toxicity=0.9)
|
||||||
|
assert r2 > r1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# ADVERSE SELECTION (15 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestAdverseSelection:
|
||||||
|
def test_cost_basic(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
cost = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=5)
|
||||||
|
assert isinstance(cost, AdverseSelectionCost)
|
||||||
|
assert cost.expected_cost_bps >= 0
|
||||||
|
|
||||||
|
def test_cost_increases_with_toxicity(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
c1 = asm.compute_cost(spread_bps=2.0, toxicity=0.1, queue_position=5)
|
||||||
|
c2 = asm.compute_cost(spread_bps=2.0, toxicity=0.9, queue_position=5)
|
||||||
|
assert c2.expected_cost_bps > c1.expected_cost_bps
|
||||||
|
|
||||||
|
def test_cost_decreases_with_queue_position(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
c1 = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=0)
|
||||||
|
c2 = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=10)
|
||||||
|
assert c1.expected_cost_bps > c2.expected_cost_bps
|
||||||
|
|
||||||
|
def test_optimal_offset_zero_toxicity(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
offset = asm.optimal_offset(spread_bps=2.0, toxicity=0.0)
|
||||||
|
assert offset == 0 # no toxicity → quote at best
|
||||||
|
|
||||||
|
def test_optimal_offset_high_toxicity(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
offset = asm.optimal_offset(spread_bps=2.0, toxicity=0.9)
|
||||||
|
assert offset >= 0 # step back under toxicity
|
||||||
|
|
||||||
|
def test_toxic_fill_ratio(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
asm.record_fill(0.3) # non-toxic
|
||||||
|
asm.record_fill(0.8) # toxic
|
||||||
|
assert asm.toxic_fill_ratio == pytest.approx(0.5, abs=0.01)
|
||||||
|
|
||||||
|
def test_toxic_fill_ratio_zero_fills(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
assert asm.toxic_fill_ratio == 0.0
|
||||||
|
|
||||||
|
def test_average_toxicity(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
asm.record_fill(0.2)
|
||||||
|
asm.record_fill(0.8)
|
||||||
|
assert asm.average_toxicity == pytest.approx(0.5, abs=0.01)
|
||||||
|
|
||||||
|
def test_cost_with_zero_spread(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
cost = asm.compute_cost(spread_bps=0.0, toxicity=0.5, queue_position=5)
|
||||||
|
assert cost.expected_cost_bps == 0.0
|
||||||
|
|
||||||
|
def test_cost_with_zero_toxicity(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
cost = asm.compute_cost(spread_bps=2.0, toxicity=0.0, queue_position=5)
|
||||||
|
assert cost.expected_cost_bps == 0.0
|
||||||
|
|
||||||
|
def test_pick_off_probability(self):
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
cost = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=0)
|
||||||
|
assert 0.0 <= cost.pick_off_probability <= 1.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# LATENCY MODEL (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestLatencyModel:
|
||||||
|
def test_feed_latency(self):
|
||||||
|
lm = LatencyModel(feed_latency_ms=10.0)
|
||||||
|
lat = lm.simulate_feed_latency()
|
||||||
|
assert lat >= 0
|
||||||
|
|
||||||
|
def test_order_latency(self):
|
||||||
|
lm = LatencyModel(order_latency_ms=50.0)
|
||||||
|
lat = lm.simulate_order_latency(queue_position=5)
|
||||||
|
assert lat >= 50.0 # at least base latency
|
||||||
|
|
||||||
|
def test_order_latency_increases_with_queue(self):
|
||||||
|
lm = LatencyModel(order_latency_ms=50.0, order_jitter_ms=0.0)
|
||||||
|
lat1 = lm.simulate_order_latency(queue_position=0, recent_trade_rate=1.0)
|
||||||
|
lat2 = lm.simulate_order_latency(queue_position=100, recent_trade_rate=1.0)
|
||||||
|
assert lat2 >= lat1
|
||||||
|
|
||||||
|
def test_latency_cost_zero_change(self):
|
||||||
|
lm = LatencyModel()
|
||||||
|
cost = lm.compute_latency_cost(price_change_per_ms=0.0)
|
||||||
|
assert cost == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# SPREAD DYNAMICS (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestSpreadDynamics:
|
||||||
|
def test_update_and_predict(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
sd.update(2.0)
|
||||||
|
sd.update(2.5)
|
||||||
|
sd.update(3.0)
|
||||||
|
predicted = sd.predict(time_horizon_s=1.0)
|
||||||
|
assert predicted > 0
|
||||||
|
|
||||||
|
def test_spread_volatility(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
for i in range(50):
|
||||||
|
sd.update(2.0 + (i % 5) * 0.1)
|
||||||
|
assert sd.spread_volatility > 0
|
||||||
|
|
||||||
|
def test_current_spread(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
sd.update(3.5)
|
||||||
|
assert sd.current_spread == 3.5
|
||||||
|
|
||||||
|
def test_predict_empty_history(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
predicted = sd.predict(5.0)
|
||||||
|
assert predicted == 0.0
|
||||||
|
|
||||||
|
def test_predict_short_history(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
sd.update(2.0)
|
||||||
|
predicted = sd.predict(5.0)
|
||||||
|
assert predicted == 2.0
|
||||||
|
|
||||||
|
def test_spread_tightening_trend(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
for i in range(20):
|
||||||
|
sd.update(5.0 - i * 0.1) # tightening
|
||||||
|
predicted = sd.predict(1.0)
|
||||||
|
assert predicted < 5.0
|
||||||
|
|
||||||
|
def test_spread_widening_trend(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
for i in range(20):
|
||||||
|
sd.update(2.0 + i * 0.1) # widening
|
||||||
|
predicted = sd.predict(1.0)
|
||||||
|
assert predicted > 2.0
|
||||||
|
|
||||||
|
def test_spread_floor(self):
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
for i in range(20):
|
||||||
|
sd.update(0.01) # very tight
|
||||||
|
predicted = sd.predict(1.0)
|
||||||
|
assert predicted >= 0.1 # floor
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# VOLATILITY CLUSTERING (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestVolatilityClustering:
|
||||||
|
def test_update_and_regime(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
vc.update(20.0)
|
||||||
|
regime = vc.regime()
|
||||||
|
assert 0.0 <= regime <= 1.0
|
||||||
|
|
||||||
|
def test_high_vol_regime(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
for _ in range(200):
|
||||||
|
vc.update(100.0) # very high vol
|
||||||
|
regime = vc.regime()
|
||||||
|
assert regime > 0.4 # sigmoid may not reach exactly 0.5
|
||||||
|
|
||||||
|
def test_low_vol_regime(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
for _ in range(200):
|
||||||
|
vc.update(1.0) # very low vol
|
||||||
|
regime = vc.regime()
|
||||||
|
assert regime < 0.6 # sigmoid may not reach exactly 0.5
|
||||||
|
|
||||||
|
def test_predict(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
vc.update(20.0)
|
||||||
|
predicted = vc.predict(60.0)
|
||||||
|
assert predicted > 0
|
||||||
|
|
||||||
|
def test_vol_of_vol(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
for i in range(50):
|
||||||
|
vc.update(15.0 + (i % 10) * 0.5)
|
||||||
|
assert vc.vol_of_vol > 0
|
||||||
|
|
||||||
|
def test_current_volatility(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
vc.update(25.0)
|
||||||
|
assert vc.current_volatility == 25.0
|
||||||
|
|
||||||
|
def test_long_term_volatility(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
for _ in range(100):
|
||||||
|
vc.update(20.0)
|
||||||
|
assert vc.long_term_volatility == pytest.approx(20.0, abs=1.0)
|
||||||
|
|
||||||
|
def test_predict_floor(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
vc.update(0.001)
|
||||||
|
predicted = vc.predict(60.0)
|
||||||
|
assert predicted > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# EXECUTION QUALITY (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestExecutionQuality:
|
||||||
|
def test_tracker_record_fill(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
assert eqt.total_fills == 1
|
||||||
|
|
||||||
|
def test_report_empty(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.total_fills == 0
|
||||||
|
|
||||||
|
def test_report_with_fills(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
eqt.record_fill(50002.0, 50000.0, 50000.5, False, 20.0, 0.5)
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.total_fills == 2
|
||||||
|
assert report.avg_slippage_bps > 0
|
||||||
|
|
||||||
|
def test_maker_ratio(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
eqt.record_fill(50002.0, 50000.0, 50000.5, True, 20.0, 0.2)
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.maker_fill_ratio == 1.0
|
||||||
|
|
||||||
|
def test_taker_ratio(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, False, 10.0, 0.5)
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.taker_fill_ratio == 1.0
|
||||||
|
|
||||||
|
def test_adverse_fill_ratio(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2, toxicity=0.8)
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.adverse_fill_ratio == 1.0
|
||||||
|
|
||||||
|
def test_avg_fill_time(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
eqt.record_fill(50002.0, 50000.0, 50000.5, True, 30.0, 0.2)
|
||||||
|
report = eqt.report()
|
||||||
|
assert report.avg_fill_time_ms == pytest.approx(20.0, abs=0.1)
|
||||||
|
|
||||||
|
def test_risk_adjusted_sharpe(self):
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
ra.add_return(0.01)
|
||||||
|
ra.add_return(0.02)
|
||||||
|
ra.add_return(-0.005)
|
||||||
|
assert ra.sharpe_ratio != 0.0
|
||||||
|
|
||||||
|
def test_risk_adjusted_sortino(self):
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
ra.add_return(0.01)
|
||||||
|
ra.add_return(0.02)
|
||||||
|
ra.add_return(-0.005)
|
||||||
|
assert ra.sortino_ratio != 0.0
|
||||||
|
|
||||||
|
def test_profit_factor(self):
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
ra.add_return(0.01)
|
||||||
|
ra.add_return(0.02)
|
||||||
|
ra.add_return(-0.005)
|
||||||
|
assert ra.profit_factor > 1.0
|
||||||
|
|
||||||
|
def test_max_drawdown(self):
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
ra.add_return(0.01)
|
||||||
|
ra.add_return(-0.005)
|
||||||
|
ra.add_return(0.02)
|
||||||
|
ra.add_return(-0.01)
|
||||||
|
assert ra.max_drawdown >= 0
|
||||||
|
|
||||||
|
def test_report_dict(self):
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
ra.add_return(0.01)
|
||||||
|
report = ra.report()
|
||||||
|
assert "sharpe_ratio" in report
|
||||||
|
assert "sortino_ratio" in report
|
||||||
|
assert "profit_factor" in report
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# MULTI-LEVEL BOOK (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestMultiLevelBook:
|
||||||
|
def test_update_and_imbalance(self):
|
||||||
|
ml = MultiLevelBookModel()
|
||||||
|
ml.update([1.0, 0.5, 0.3], [0.8, 0.4, 0.2])
|
||||||
|
imbalance = ml.compute_imbalance()
|
||||||
|
assert isinstance(imbalance, float)
|
||||||
|
|
||||||
|
def test_depth_ratio(self):
|
||||||
|
ml = MultiLevelBookModel()
|
||||||
|
ml.update([1.0, 0.5], [0.5, 0.25])
|
||||||
|
ratio = ml.compute_depth_ratio()
|
||||||
|
assert ratio > 1.0
|
||||||
|
|
||||||
|
def test_current_depth(self):
|
||||||
|
ml = MultiLevelBookModel()
|
||||||
|
ml.update([1.0, 0.5], [0.8, 0.4])
|
||||||
|
assert ml.current_bid_depth > 0
|
||||||
|
assert ml.current_ask_depth > 0
|
||||||
|
|
||||||
|
def test_net_order_flow(self):
|
||||||
|
bid_f, ask_f = compute_net_order_flow(1.0, 1.0, 0.0, 0.0, 15.0)
|
||||||
|
assert bid_f >= 0
|
||||||
|
assert ask_f >= 0
|
||||||
|
|
||||||
|
def test_net_flow_with_imbalance(self):
|
||||||
|
bid_f, ask_f = compute_net_order_flow(1.0, 1.0, 0.5, 0.0, 15.0)
|
||||||
|
assert bid_f > ask_f # buying pressure
|
||||||
|
|
||||||
|
def test_net_flow_with_toxicity(self):
|
||||||
|
bid_f, ask_f = compute_net_order_flow(1.0, 1.0, 0.0, 0.9, 15.0)
|
||||||
|
assert bid_f < 1.0 # toxicity reduces flow
|
||||||
|
|
||||||
|
def test_weighted_imbalance(self):
|
||||||
|
bid_p = np.array([1.0, 2.0, 3.0], dtype=np.float64)
|
||||||
|
bid_q = np.array([1.0, 1.0, 1.0], dtype=np.float64)
|
||||||
|
ask_p = np.array([1.0, 2.0, 3.0], dtype=np.float64)
|
||||||
|
ask_q = np.array([1.0, 1.0, 1.0], dtype=np.float64)
|
||||||
|
imbalance = compute_book_imbalance_weighted(bid_p, bid_q, ask_p, ask_q, 3)
|
||||||
|
assert imbalance == pytest.approx(0.0, abs=0.01) # symmetric
|
||||||
|
|
||||||
|
def test_weighted_imbalance_asymmetric(self):
|
||||||
|
bid_p = np.array([1.0, 2.0], dtype=np.float64)
|
||||||
|
bid_q = np.array([2.0, 2.0], dtype=np.float64)
|
||||||
|
ask_p = np.array([1.0, 2.0], dtype=np.float64)
|
||||||
|
ask_q = np.array([1.0, 1.0], dtype=np.float64)
|
||||||
|
imbalance = compute_book_imbalance_weighted(bid_p, bid_q, ask_p, ask_q, 2)
|
||||||
|
assert imbalance > 0 # more on bid side
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# MULTI-ASSET CORRELATION (12 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestMultiAssetCorrelation:
|
||||||
|
def test_update_returns(self):
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
mac.update_returns("BTCUSDT", 0.01)
|
||||||
|
assert mac.asset_count == 1
|
||||||
|
|
||||||
|
def test_compute_correlation(self):
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
for i in range(30):
|
||||||
|
mac.update_returns("BTCUSDT", 0.01 * (1 if i % 2 == 0 else -1))
|
||||||
|
mac.update_returns("ETHUSDT", 0.01 * (1 if i % 2 == 0 else -1))
|
||||||
|
corr = mac.compute_correlation("BTCUSDT", "ETHUSDT")
|
||||||
|
assert -1.0 <= corr <= 1.0
|
||||||
|
|
||||||
|
def test_correlation_same_asset(self):
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
for i in range(30):
|
||||||
|
mac.update_returns("BTCUSDT", 0.01)
|
||||||
|
corr = mac.compute_correlation("BTCUSDT", "BTCUSDT")
|
||||||
|
assert corr == pytest.approx(1.0, abs=0.01)
|
||||||
|
|
||||||
|
def test_btc_correlation(self):
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
for i in range(30):
|
||||||
|
mac.update_returns("BTCUSDT", 0.01 * (1 if i % 2 == 0 else -1))
|
||||||
|
mac.update_returns("ETHUSDT", 0.01 * (1 if i % 2 == 0 else -1))
|
||||||
|
corr = mac.compute_correlation("BTCUSDT", "ETHUSDT")
|
||||||
|
# Perfectly correlated series should have corr near 1.0
|
||||||
|
# (numpy may return exactly 1.0 or close to it)
|
||||||
|
assert abs(corr) > 0.5
|
||||||
|
|
||||||
|
def test_asset_count(self):
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
mac.update_returns("A", 0.01)
|
||||||
|
mac.update_returns("B", 0.02)
|
||||||
|
assert mac.asset_count == 2
|
||||||
|
|
||||||
|
def test_correlation_regime(self):
|
||||||
|
regime = compute_correlation_regime(0.8, 0.1)
|
||||||
|
assert regime > 0.5
|
||||||
|
|
||||||
|
def test_correlation_regime_low(self):
|
||||||
|
regime = compute_correlation_regime(0.2, 0.1)
|
||||||
|
assert regime < 0.5
|
||||||
|
|
||||||
|
def test_rolling_correlation(self):
|
||||||
|
a = np.array([1.0, 2.0, 3.0, 4.0, 5.0], dtype=np.float64)
|
||||||
|
b = np.array([1.0, 2.0, 3.0, 4.0, 5.0], dtype=np.float64)
|
||||||
|
corr = compute_rolling_correlation(a, b, window=5)
|
||||||
|
assert corr == pytest.approx(1.0, abs=0.01)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# INTEGRATION: ALL MODULES TOGETHER (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestIntegration:
|
||||||
|
def test_queue_adverse_selection_pipeline(self):
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
fp = qm.estimate_fill_probability(0.001, 1.0, toxicity=0.5)
|
||||||
|
cost = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=5)
|
||||||
|
assert fp > 0
|
||||||
|
assert cost.expected_cost_bps >= 0
|
||||||
|
|
||||||
|
def test_latency_spread_interaction(self):
|
||||||
|
lm = LatencyModel(feed_latency_ms=10.0, order_latency_ms=50.0)
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
sd.update(2.0)
|
||||||
|
latency_cost = lm.compute_latency_cost()
|
||||||
|
spread_predict = sd.predict(5.0)
|
||||||
|
assert latency_cost >= 0
|
||||||
|
assert spread_predict > 0
|
||||||
|
|
||||||
|
def test_volatility_correlation_interaction(self):
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
vc.update(20.0)
|
||||||
|
mac.update_returns("BTCUSDT", 0.01)
|
||||||
|
regime = vc.regime()
|
||||||
|
corr = mac.get_btc_correlation("BTCUSDT")
|
||||||
|
assert 0.0 <= regime <= 1.0
|
||||||
|
assert isinstance(corr, float)
|
||||||
|
|
||||||
|
def test_execution_quality_risk_adjusted(self):
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
ra.add_return(0.01)
|
||||||
|
report = eqt.report()
|
||||||
|
risk_report = ra.report()
|
||||||
|
assert report.total_fills == 1
|
||||||
|
assert "sharpe_ratio" in risk_report
|
||||||
|
|
||||||
|
def test_multi_level_queue_integration(self):
|
||||||
|
ml = MultiLevelBookModel()
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
ml.update([1.0, 0.5], [0.8, 0.4])
|
||||||
|
depth_ratio = ml.compute_depth_ratio()
|
||||||
|
fp = qm.estimate_fill_probability(0.001, 0.8)
|
||||||
|
assert depth_ratio > 0
|
||||||
|
assert fp >= 0
|
||||||
|
|
||||||
|
def test_full_pipeline(self):
|
||||||
|
"""All models work together without errors."""
|
||||||
|
qm = QueuePositionModel()
|
||||||
|
asm = AdverseSelectionModel()
|
||||||
|
lm = LatencyModel()
|
||||||
|
sd = SpreadDynamicsModel()
|
||||||
|
vc = VolatilityClusteringModel()
|
||||||
|
ml = MultiLevelBookModel()
|
||||||
|
mac = MultiAssetCorrelationModel()
|
||||||
|
eqt = ExecutionQualityTracker()
|
||||||
|
ra = RiskAdjustedReturns()
|
||||||
|
|
||||||
|
# Update all models
|
||||||
|
sd.update(2.0)
|
||||||
|
vc.update(20.0)
|
||||||
|
ml.update([1.0, 0.5], [0.8, 0.4])
|
||||||
|
mac.update_returns("BTCUSDT", 0.01)
|
||||||
|
eqt.record_fill(50001.0, 50000.0, 50000.5, True, 10.0, 0.2)
|
||||||
|
ra.add_return(0.01)
|
||||||
|
|
||||||
|
# Query all models
|
||||||
|
fp = qm.estimate_fill_probability(0.001, 1.0, toxicity=0.5)
|
||||||
|
cost = asm.compute_cost(spread_bps=2.0, toxicity=0.5, queue_position=5)
|
||||||
|
lat = lm.simulate_feed_latency()
|
||||||
|
spread = sd.predict(5.0)
|
||||||
|
vol_regime = vc.regime()
|
||||||
|
depth_ratio = ml.compute_depth_ratio()
|
||||||
|
corr = mac.get_btc_correlation("BTCUSDT")
|
||||||
|
exec_report = eqt.report()
|
||||||
|
risk_report = ra.report()
|
||||||
|
|
||||||
|
# All should return valid values
|
||||||
|
assert fp >= 0
|
||||||
|
assert cost.expected_cost_bps >= 0
|
||||||
|
assert lat >= 0
|
||||||
|
assert spread > 0
|
||||||
|
assert 0 <= vol_regime <= 1
|
||||||
|
assert depth_ratio > 0
|
||||||
|
assert isinstance(corr, float)
|
||||||
|
assert exec_report.total_fills == 1
|
||||||
|
assert "sharpe_ratio" in risk_report
|
||||||
|
|
||||||
|
def test_numba_functions_correct(self):
|
||||||
|
"""Verify numba-accelerated functions return same results as Python."""
|
||||||
|
from malkhut.cwm.queue_model import estimate_queue_position, compute_fill_probability
|
||||||
|
# Pure Python equivalents
|
||||||
|
def py_estimate(our_qty, level_qty, rate, time_s):
|
||||||
|
if level_qty <= 0: return 0.0
|
||||||
|
queue_depth = max(0.0, level_qty - our_qty)
|
||||||
|
if rate <= 0: return queue_depth
|
||||||
|
consumed = rate * time_s
|
||||||
|
return max(0.0, queue_depth - consumed)
|
||||||
|
|
||||||
|
def py_fill_prob(qd, oq, rate, horizon, tox):
|
||||||
|
if oq <= 0 or qd < 0: return 0.0
|
||||||
|
if rate <= 0: return 0.0
|
||||||
|
total = qd + oq
|
||||||
|
if total <= 0: return 1.0
|
||||||
|
base = rate / total
|
||||||
|
tox_f = 1.0 + tox * 0.5
|
||||||
|
return min(1.0, max(0.0, 1.0 - math.exp(-base * tox_f * horizon)))
|
||||||
|
|
||||||
|
# Test multiple values
|
||||||
|
for qd in [0.0, 0.5, 1.0, 5.0]:
|
||||||
|
for oq in [0.001, 0.01, 0.1]:
|
||||||
|
for rate in [0.1, 0.5, 1.0]:
|
||||||
|
for tox in [0.0, 0.5, 0.9]:
|
||||||
|
nb_val = compute_fill_probability(qd, oq, rate, 300.0, tox)
|
||||||
|
py_val = py_fill_prob(qd, oq, rate, 300.0, tox)
|
||||||
|
assert abs(nb_val - py_val) < 1e-6
|
||||||
338
MALKHUT/malkhut/tests/test_new_features.py
Normal file
338
MALKHUT/malkhut/tests/test_new_features.py
Normal file
@@ -0,0 +1,338 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
469
MALKHUT/malkhut/tests/test_new_features_comprehensive.py
Normal file
469
MALKHUT/malkhut/tests/test_new_features_comprehensive.py
Normal file
@@ -0,0 +1,469 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
120
MALKHUT/malkhut/tests/test_numba.py
Normal file
120
MALKHUT/malkhut/tests/test_numba.py
Normal file
@@ -0,0 +1,120 @@
|
|||||||
|
"""
|
||||||
|
Tests for numba-accelerated CWM functions.
|
||||||
|
|
||||||
|
Verifies:
|
||||||
|
- Numba JIT compilation works
|
||||||
|
- fill_from_levels produces same results as pure Python
|
||||||
|
- round_tick / round_lot / clip_lots correct
|
||||||
|
- Feature extraction vectorized
|
||||||
|
- Fallback to pure Python when numba unavailable
|
||||||
|
"""
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
from malkhut.cwm.numba_core import (
|
||||||
|
fill_from_levels, round_tick, round_lot, clip_lots,
|
||||||
|
extract_features_vectorized, compare_states_vectorized,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNumbaFillFromLevels:
|
||||||
|
def test_fill_single_level(self):
|
||||||
|
prices = np.array([50000.0], dtype=np.float64)
|
||||||
|
qtys = np.array([1.0], dtype=np.float64)
|
||||||
|
filled, avg, _, _ = fill_from_levels(
|
||||||
|
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
|
||||||
|
prices, qtys, 0.5, 0.001, 0.001, True,
|
||||||
|
)
|
||||||
|
assert filled == pytest.approx(0.5, abs=0.001)
|
||||||
|
assert avg == pytest.approx(50000.0, abs=0.01)
|
||||||
|
|
||||||
|
def test_fill_multi_level(self):
|
||||||
|
prices = np.array([50000.0, 50001.0], dtype=np.float64)
|
||||||
|
qtys = np.array([0.5, 0.5], dtype=np.float64)
|
||||||
|
filled, avg, _, _ = fill_from_levels(
|
||||||
|
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
|
||||||
|
prices, qtys, 0.8, 0.001, 0.001, True,
|
||||||
|
)
|
||||||
|
assert filled == pytest.approx(0.8, abs=0.001)
|
||||||
|
assert avg > 50000.0
|
||||||
|
|
||||||
|
def test_fill_exhausts_all(self):
|
||||||
|
prices = np.array([50000.0, 50001.0], dtype=np.float64)
|
||||||
|
qtys = np.array([0.3, 0.3], dtype=np.float64)
|
||||||
|
filled, avg, _, _ = fill_from_levels(
|
||||||
|
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
|
||||||
|
prices, qtys, 1.0, 0.001, 0.001, True,
|
||||||
|
)
|
||||||
|
assert filled == pytest.approx(0.6, abs=0.001)
|
||||||
|
|
||||||
|
def test_fill_empty(self):
|
||||||
|
filled, avg, _, _ = fill_from_levels(
|
||||||
|
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
|
||||||
|
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
|
||||||
|
1.0, 0.001, 0.001, True,
|
||||||
|
)
|
||||||
|
assert filled == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestNumbaRounding:
|
||||||
|
def test_round_tick(self):
|
||||||
|
assert round_tick(50000.0, 0.1) == 50000.0
|
||||||
|
assert round_tick(50000.06, 0.1) == pytest.approx(50000.1, abs=1e-9)
|
||||||
|
assert round_tick(50004.4, 1.0) == 50004.0
|
||||||
|
|
||||||
|
def test_round_lot(self):
|
||||||
|
assert round_lot(0.001, 0.001) == 0.001
|
||||||
|
assert round_lot(0.0017, 0.001) == 0.002
|
||||||
|
|
||||||
|
def test_clip_lots_above_min(self):
|
||||||
|
assert clip_lots(0.005, 0.001, 0.001) == 0.005
|
||||||
|
|
||||||
|
def test_clip_lots_below_min(self):
|
||||||
|
assert clip_lots(0.0005, 0.001, 0.001) == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestNumbaFeatures:
|
||||||
|
def test_extract_features(self):
|
||||||
|
bid_p = np.array([50000.0], dtype=np.float64)
|
||||||
|
bid_q = np.array([1.0], dtype=np.float64)
|
||||||
|
ask_p = np.array([50001.0], dtype=np.float64)
|
||||||
|
ask_q = np.array([1.0], dtype=np.float64)
|
||||||
|
features = extract_features_vectorized(
|
||||||
|
bid_p, bid_q, ask_p, ask_q,
|
||||||
|
50000.5, 0.1, 0.0, 15.0, 0.0, -10.0, 15.0, 15.0,
|
||||||
|
50.0, 30.0, 10.0, 1.0, -0.5, 0.3, 0.2, 0.1,
|
||||||
|
)
|
||||||
|
assert len(features) == 17
|
||||||
|
assert features[0] == pytest.approx(50000.5, abs=0.1) # mid
|
||||||
|
assert features[14] == pytest.approx(0.3, abs=0.01) # toxicity
|
||||||
|
|
||||||
|
|
||||||
|
class TestNumbaCompare:
|
||||||
|
def test_compare_match(self):
|
||||||
|
ok, idx, ev, av = compare_states_vectorized(
|
||||||
|
10000.0, 10000.0, 50000.0, 50000.0, 50001.0, 50001.0,
|
||||||
|
1e-6, 1e-6,
|
||||||
|
)
|
||||||
|
assert ok
|
||||||
|
|
||||||
|
def test_compare_equity_mismatch(self):
|
||||||
|
ok, idx, ev, av = compare_states_vectorized(
|
||||||
|
10000.0, 9000.0, 50000.0, 50000.0, 50001.0, 50001.0,
|
||||||
|
1e-6, 1e-6,
|
||||||
|
)
|
||||||
|
assert not ok
|
||||||
|
assert idx == 0
|
||||||
|
|
||||||
|
def test_compare_bid_mismatch(self):
|
||||||
|
ok, idx, ev, av = compare_states_vectorized(
|
||||||
|
10000.0, 10000.0, 50000.0, 50001.0, 50001.0, 50001.0,
|
||||||
|
1e-6, 1e-6,
|
||||||
|
)
|
||||||
|
assert not ok
|
||||||
|
assert idx == 1
|
||||||
|
|
||||||
|
def test_compare_within_tolerance(self):
|
||||||
|
ok, idx, ev, av = compare_states_vectorized(
|
||||||
|
10000.0, 10000.001, 50000.0, 50000.001, 50001.0, 50001.001,
|
||||||
|
0.1, 0.1,
|
||||||
|
)
|
||||||
|
assert ok
|
||||||
208
MALKHUT/malkhut/tests/test_path_risk.py
Normal file
208
MALKHUT/malkhut/tests/test_path_risk.py
Normal file
@@ -0,0 +1,208 @@
|
|||||||
|
"""
|
||||||
|
Path risk / trade path state — SL/TP path-aware logic.
|
||||||
|
|
||||||
|
Tests that the CWM reward and risk functions correctly respond to
|
||||||
|
various trade path states (MAE, MFE, recovery velocity, failed recovery).
|
||||||
|
"""
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, IntentKind, MarketWorldState,
|
||||||
|
Mode, OrderBookState, PriceLevel, Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
|
||||||
|
|
||||||
|
def _params(**kw):
|
||||||
|
d = dict(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2,
|
||||||
|
rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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 _path(**kw):
|
||||||
|
d = dict(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=10, seconds_held=100.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=30.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=20.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
d.update(kw)
|
||||||
|
return TradePathState(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(**kw):
|
||||||
|
book = OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
)
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=book,
|
||||||
|
account=kw.get("account", AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
)),
|
||||||
|
trade_path=kw.get("trade_path"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPathRiskExitDecision:
|
||||||
|
def test_no_exit_when_no_path(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
s = _state()
|
||||||
|
assert not _path_risk_says_exit(s, _params())
|
||||||
|
|
||||||
|
def test_exit_when_mae_exceeds_threshold(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
path = _path(mae_bps=-60.0, recovery_velocity_bps_per_s=-2.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert _path_risk_says_exit(s, _params(mae_tail_cut_bps=50.0))
|
||||||
|
|
||||||
|
def test_no_exit_when_mae_below_threshold(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
# mae=30 < threshold=50, time_in_loss=50 < 300, failed_recovery=0 < 3
|
||||||
|
# mfe=15, distance=3 => giveback=3/15=0.2 < 0.5
|
||||||
|
path = _path(mae_bps=-30.0, distance_from_mfe_bps=3.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert not _path_risk_says_exit(s, _params(mae_tail_cut_bps=50.0))
|
||||||
|
|
||||||
|
def test_exit_when_time_in_loss_exceeds(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
path = _path(time_in_loss_s=400.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert _path_risk_says_exit(s, _params(max_time_in_loss_s=300.0))
|
||||||
|
|
||||||
|
def test_exit_when_failed_recovery_count_exceeds(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
path = _path(failed_recovery_count=5)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert _path_risk_says_exit(s, _params(failed_recovery_cut_count=3))
|
||||||
|
|
||||||
|
def test_exit_on_mfe_giveback(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
# mfe=10, distance_from_mfe=8 => giveback = 8/10 = 0.8 > 0.5
|
||||||
|
path = _path(mfe_bps=10.0, distance_from_mfe_bps=8.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert _path_risk_says_exit(s, _params(mfe_giveback_cut_fraction=0.5))
|
||||||
|
|
||||||
|
def test_no_exit_when_mfe_giveback_below_threshold(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
# mfe=10, distance_from_mfe=3 => giveback = 3/10 = 0.3 < 0.5
|
||||||
|
path = _path(mfe_bps=10.0, distance_from_mfe_bps=3.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert not _path_risk_says_exit(s, _params(mfe_giveback_cut_fraction=0.5))
|
||||||
|
|
||||||
|
def test_exit_when_both_mae_and_slow_recovery(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
path = _path(mae_bps=-60.0, recovery_velocity_bps_per_s=-3.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert _path_risk_says_exit(s, _params(mae_tail_cut_bps=50.0, recovery_velocity_min_bps_per_s=0.0))
|
||||||
|
|
||||||
|
def test_no_exit_when_mae_high_but_recovery_fast(self):
|
||||||
|
from malkhut.planner.action_menu import _path_risk_says_exit
|
||||||
|
# mae=60 > threshold=50, but recovery_velocity=5.0 > min=0.0 → no MAE exit
|
||||||
|
# time_in_loss=50 < 300, failed_recovery=0 < 3
|
||||||
|
# mfe=15, distance=3 => giveback=0.2 < 0.5
|
||||||
|
path = _path(mae_bps=-60.0, recovery_velocity_bps_per_s=5.0,
|
||||||
|
distance_from_mfe_bps=3.0)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
assert not _path_risk_says_exit(s, _params(mae_tail_cut_bps=50.0, recovery_velocity_min_bps_per_s=0.0))
|
||||||
|
|
||||||
|
|
||||||
|
class TestPathRiskProxy:
|
||||||
|
def test_zero_when_no_path(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
assert cwm._tail_risk_proxy(s) == 0.0
|
||||||
|
|
||||||
|
def test_increases_with_mae(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
p1 = _path(mae_bps=-10.0)
|
||||||
|
p2 = _path(mae_bps=-50.0)
|
||||||
|
r1 = cwm._tail_risk_proxy(_state(trade_path=p1))
|
||||||
|
r2 = cwm._tail_risk_proxy(_state(trade_path=p2))
|
||||||
|
assert r2 > r1
|
||||||
|
|
||||||
|
def test_increases_with_time_in_loss(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
p1 = _path(time_in_loss_s=10.0)
|
||||||
|
p2 = _path(time_in_loss_s=100.0)
|
||||||
|
r1 = cwm._tail_risk_proxy(_state(trade_path=p1))
|
||||||
|
r2 = cwm._tail_risk_proxy(_state(trade_path=p2))
|
||||||
|
assert r2 > r1
|
||||||
|
|
||||||
|
def test_increases_with_failed_recovery(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
p1 = _path(failed_recovery_count=0)
|
||||||
|
p2 = _path(failed_recovery_count=3)
|
||||||
|
r1 = cwm._tail_risk_proxy(_state(trade_path=p1))
|
||||||
|
r2 = cwm._tail_risk_proxy(_state(trade_path=p2))
|
||||||
|
assert r2 > r1
|
||||||
|
|
||||||
|
|
||||||
|
class TestInventoryRisk:
|
||||||
|
def test_zero_when_no_position(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
assert cwm._inventory_risk(s) == 0.0
|
||||||
|
|
||||||
|
def test_positive_when_position_exists(self):
|
||||||
|
from malkhut.state import PositionState, AccountState
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s = _state(account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=5000.0,
|
||||||
|
positions={"BTCUSDT": pos},
|
||||||
|
))
|
||||||
|
risk = cwm._inventory_risk(s)
|
||||||
|
assert risk > 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestRewardPathSensitivity:
|
||||||
|
def test_higher_w_tail_loss_more_penalty_for_deep_mae(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
path = _path(mae_bps=-40.0, failed_recovery_count=2)
|
||||||
|
s = _state(trade_path=path)
|
||||||
|
from malkhut.actions import FulfilmentAction, ActionKind
|
||||||
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
r = cwm.transition(s, (a,))
|
||||||
|
p1 = _params(w_tail_loss=1.0)
|
||||||
|
p2 = _params(w_tail_loss=10.0)
|
||||||
|
# reward = 0 - tail_risk * w_tail_loss - ...
|
||||||
|
# Higher w_tail_loss should produce more negative reward
|
||||||
|
assert cwm.reward(s, a, r, p2) < cwm.reward(s, a, r, p1)
|
||||||
254
MALKHUT/malkhut/tests/test_phase0.py
Normal file
254
MALKHUT/malkhut/tests/test_phase0.py
Normal file
@@ -0,0 +1,254 @@
|
|||||||
|
"""
|
||||||
|
Tests for tie-in fix, cognition pipeline, and regime expansion.
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.training.cognition import (
|
||||||
|
RateLimiter, SourceCatalogue, RegimeExtractor, CognitionPipeline,
|
||||||
|
)
|
||||||
|
from malkhut.training.regime_expansion import (
|
||||||
|
RegimeExpander, ExpandedRegime,
|
||||||
|
LIQUIDITY_DIMS, VOLATILITY_DIMS, FLOW_DIMS, STRUCTURE_DIMS,
|
||||||
|
)
|
||||||
|
from malkhut.training.selector import PerformanceMatrix, MarketRegime
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# TIE-IN FIX: PerformanceMatrix wired to evaluator
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTieInFix:
|
||||||
|
def test_matrix_records_regime_performance(self):
|
||||||
|
"""PerformanceMatrix should record strategy × regime scores."""
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
matrix.record("strat_1", "flash_crash", score=10.0, pnl_bps=5.0)
|
||||||
|
matrix.record("strat_1", "normal", score=15.0, pnl_bps=8.0)
|
||||||
|
scores = matrix.get_scores_for_regime("flash_crash")
|
||||||
|
assert len(scores) == 1
|
||||||
|
assert scores[0].strategy_id == "strat_1"
|
||||||
|
|
||||||
|
def test_matrix_tracks_multiple_strategies(self):
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
matrix.record("s1", "normal", score=10.0)
|
||||||
|
matrix.record("s2", "normal", score=15.0)
|
||||||
|
scores = matrix.get_scores_for_regime("normal")
|
||||||
|
assert len(scores) == 2
|
||||||
|
assert scores[0].score == 15.0 # sorted descending
|
||||||
|
|
||||||
|
def test_matrix_best_for_regime(self):
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
matrix.record("s1", "flash_crash", score=5.0)
|
||||||
|
matrix.record("s2", "flash_crash", score=10.0)
|
||||||
|
best = matrix.get_best("flash_crash")
|
||||||
|
assert best == "s2"
|
||||||
|
|
||||||
|
def test_matrix_best_excludes_baseline(self):
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
matrix.record("baseline", "normal", score=10.0)
|
||||||
|
matrix.record("s1", "normal", score=8.0)
|
||||||
|
best = matrix.get_best("normal", exclude={"baseline"})
|
||||||
|
assert best == "s1"
|
||||||
|
|
||||||
|
def test_matrix_coverage(self):
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
matrix.record("s1", "normal", score=10.0)
|
||||||
|
matrix.record("s1", "flash_crash", score=5.0)
|
||||||
|
matrix.record("s2", "normal", score=8.0)
|
||||||
|
cov = matrix.get_coverage()
|
||||||
|
assert cov["s1"] == 2
|
||||||
|
assert cov["s2"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# COGNITION PIPELINE
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRateLimiter:
|
||||||
|
def test_acquire_within_burst(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=5)
|
||||||
|
for _ in range(5):
|
||||||
|
assert rl.acquire()
|
||||||
|
|
||||||
|
def test_acquire_exceeds_burst(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=2)
|
||||||
|
assert rl.acquire()
|
||||||
|
assert rl.acquire()
|
||||||
|
assert not rl.acquire()
|
||||||
|
|
||||||
|
def test_token_refill(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=1)
|
||||||
|
rl.acquire()
|
||||||
|
rl._last_refill = time.time() - 2 # simulate 2 seconds ago
|
||||||
|
rl._tokens = 0 # force empty
|
||||||
|
assert rl.acquire() # should refill and acquire
|
||||||
|
|
||||||
|
|
||||||
|
class TestSourceCatalogue:
|
||||||
|
def test_add_source(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "Test Source", "https://example.com", "news", 0.8)
|
||||||
|
assert cat.source_count == 1
|
||||||
|
assert len(cat.get_enabled()) == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_record_fetch(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "Test", "https://example.com")
|
||||||
|
cat.record_fetch("s1", success=True)
|
||||||
|
cat.record_fetch("s1", success=False)
|
||||||
|
sources = cat.get_enabled()
|
||||||
|
assert sources[0].fetch_count == 2
|
||||||
|
assert sources[0].error_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRegimeExtractor:
|
||||||
|
def test_extract_flash_crash(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Bitcoin crashed 20% in 5 minutes")
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
|
||||||
|
def test_extract_liquidity_vacuum(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Liquidity vacuum observed on exchange")
|
||||||
|
assert "liquidity_vacuum" in regimes
|
||||||
|
|
||||||
|
def test_extract_normal(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Normal market conditions")
|
||||||
|
assert "normal" in regimes
|
||||||
|
|
||||||
|
def test_extract_multiple(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Flash crash with high volatility and liquidation cascade")
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
# "high_volatility" might not match exactly — check for volatility-related
|
||||||
|
assert any("liquid" in r or "crash" in r for r in regimes)
|
||||||
|
|
||||||
|
def test_sentiment_positive(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
score = ext.extract_sentiment("Bitcoin rally surges to new highs with gains")
|
||||||
|
assert score > 0
|
||||||
|
|
||||||
|
def test_sentiment_negative(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
score = ext.extract_sentiment("Flash crash panic liquidation cascade")
|
||||||
|
assert score < 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestCognitionPipeline:
|
||||||
|
def test_add_source(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com", "news", 0.8)
|
||||||
|
assert pipe._catalogue.source_count == 1
|
||||||
|
|
||||||
|
def test_fetch_and_extract(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
new_regimes = pipe.fetch_and_extract("s1", "Flash crash observed in BTC")
|
||||||
|
assert "flash_crash" in new_regimes
|
||||||
|
assert pipe._total_fetched == 1
|
||||||
|
|
||||||
|
def test_deduplication(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash observed")
|
||||||
|
new = pipe.fetch_and_extract("s1", "Flash crash again")
|
||||||
|
assert len(new) == 0 # already discovered
|
||||||
|
|
||||||
|
def test_discovered_regimes(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash observed")
|
||||||
|
regimes = pipe.get_discovered_regimes()
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
|
||||||
|
def test_source_stats(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Test text")
|
||||||
|
stats = pipe.get_source_stats()
|
||||||
|
assert stats["total_sources"] == 1
|
||||||
|
assert stats["total_fetched"] == 1
|
||||||
|
|
||||||
|
def test_seed_default_sources(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.seed_default_sources()
|
||||||
|
assert pipe._catalogue.source_count == 8
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# REGIME EXPANSION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRegimeExpansion:
|
||||||
|
def test_generate_regimes(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
assert len(regimes) == 50
|
||||||
|
assert exp.generated_count == 50
|
||||||
|
|
||||||
|
def test_regime_ids_unique(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=100)
|
||||||
|
ids = [r.regime_id for r in regimes]
|
||||||
|
assert len(ids) == len(set(ids))
|
||||||
|
|
||||||
|
def test_regime_labels_diverse(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
labels = [r.label for r in regimes]
|
||||||
|
assert len(set(labels)) > 20 # many distinct labels
|
||||||
|
|
||||||
|
def test_regime_to_scenario(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=5)
|
||||||
|
from malkhut.training.cma_trainer import Scenario
|
||||||
|
scenario = exp.regime_to_scenario(regimes[0])
|
||||||
|
assert isinstance(scenario, Scenario)
|
||||||
|
assert scenario.max_steps == 20
|
||||||
|
|
||||||
|
def test_regime_dimensions_represented(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=100)
|
||||||
|
liq_labels = set(r.liquidity.label for r in regimes)
|
||||||
|
vol_labels = set(r.volatility.label for r in regimes)
|
||||||
|
flow_labels = set(r.flow.label for r in regimes)
|
||||||
|
assert len(liq_labels) == 4 # vacuum, thin, normal, deep
|
||||||
|
assert len(vol_labels) == 4 # tight, normal, wide, extreme
|
||||||
|
assert len(flow_labels) == 4 # balanced, buy_pressure, sell_pressure, toxic
|
||||||
|
|
||||||
|
def test_orthogonal_to_cognition(self):
|
||||||
|
"""Regime expansion is orthogonal to cognition pipeline."""
|
||||||
|
from malkhut.training.cognition import CognitionPipeline
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
exp = RegimeExpander()
|
||||||
|
|
||||||
|
# Cognition discovers regimes from news
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash and high volatility")
|
||||||
|
cog_regimes = set(pipe.get_discovered_regimes())
|
||||||
|
|
||||||
|
# Expansion generates regimes from dimensions
|
||||||
|
exp_regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
exp_labels = set(r.label for r in exp_regimes)
|
||||||
|
|
||||||
|
# They are different (orthogonal) — no overlap required
|
||||||
|
assert len(cog_regimes) > 0
|
||||||
|
assert len(exp_labels) > 20
|
||||||
|
|
||||||
|
def test_max_regimes_capped(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=200)
|
||||||
|
assert len(regimes) <= 200
|
||||||
555
MALKHUT/malkhut/tests/test_phase0_extensive.py
Normal file
555
MALKHUT/malkhut/tests/test_phase0_extensive.py
Normal file
@@ -0,0 +1,555 @@
|
|||||||
|
"""
|
||||||
|
Phase 0 Extensive Tests — tie-in, cognition pipeline, regime expansion.
|
||||||
|
|
||||||
|
Covers:
|
||||||
|
1. Tie-in: PerformanceMatrix recording during evaluation
|
||||||
|
2. RateLimiter: token bucket mechanics
|
||||||
|
3. SourceCatalogue: add, fetch, error tracking
|
||||||
|
4. RegimeExtractor: keyword matching, sentiment
|
||||||
|
5. CognitionPipeline: full pipeline flow
|
||||||
|
6. RegimeExpander: dimension combinations
|
||||||
|
7. Integration: cognition + expansion + matrix + selector
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
from malkhut.training.cognition import (
|
||||||
|
RateLimiter, SourceCatalogue, RegimeExtractor, CognitionPipeline,
|
||||||
|
)
|
||||||
|
from malkhut.training.regime_expansion import (
|
||||||
|
RegimeExpander, ExpandedRegime,
|
||||||
|
LIQUIDITY_DIMS, VOLATILITY_DIMS, FLOW_DIMS, STRUCTURE_DIMS,
|
||||||
|
)
|
||||||
|
from malkhut.training.selector import PerformanceMatrix, MarketRegime
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# TIE-IN FIX (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTieInFix:
|
||||||
|
def test_record_single_regime(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
assert len(m.get_scores_for_regime("normal")) == 1
|
||||||
|
|
||||||
|
def test_record_multiple_strategies(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s2", "normal", score=15.0)
|
||||||
|
scores = m.get_scores_for_regime("normal")
|
||||||
|
assert len(scores) == 2
|
||||||
|
assert scores[0].score == 15.0
|
||||||
|
|
||||||
|
def test_record_multiple_regimes(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s1", "flash_crash", score=5.0)
|
||||||
|
assert len(m.get_scores_for_regime("normal")) == 1
|
||||||
|
assert len(m.get_scores_for_regime("flash_crash")) == 1
|
||||||
|
|
||||||
|
def test_best_for_regime(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s2", "normal", score=15.0)
|
||||||
|
assert m.get_best("normal") == "s2"
|
||||||
|
|
||||||
|
def test_best_excludes_baseline(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("baseline", "normal", score=10.0)
|
||||||
|
m.record("s1", "normal", score=8.0)
|
||||||
|
assert m.get_best("normal", exclude={"baseline"}) == "s1"
|
||||||
|
|
||||||
|
def test_best_returns_none_when_empty(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
assert m.get_best("nonexistent") is None
|
||||||
|
|
||||||
|
def test_coverage_tracking(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s1", "flash_crash", score=5.0)
|
||||||
|
m.record("s2", "normal", score=8.0)
|
||||||
|
cov = m.get_coverage()
|
||||||
|
assert cov["s1"] == 2
|
||||||
|
assert cov["s2"] == 1
|
||||||
|
|
||||||
|
def test_ema_update(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s1", "normal", score=20.0)
|
||||||
|
scores = m.get_scores_for_regime("normal")
|
||||||
|
assert scores[0].episodes == 2
|
||||||
|
# EMA: 0.3 * 20 + 0.7 * 10 = 13
|
||||||
|
assert scores[0].score == pytest.approx(13.0, abs=0.1)
|
||||||
|
|
||||||
|
def test_total_entries(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s2", "flash_crash", score=5.0)
|
||||||
|
assert m.total_entries == 2
|
||||||
|
|
||||||
|
def test_regime_list(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", "normal", score=10.0)
|
||||||
|
m.record("s1", "flash_crash", score=5.0)
|
||||||
|
regimes = m.get_regimes_for_strategy("s1")
|
||||||
|
assert len(regimes) == 2
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# RATE LIMITER (8 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRateLimiter:
|
||||||
|
def test_acquire_within_burst(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=5)
|
||||||
|
for _ in range(5):
|
||||||
|
assert rl.acquire()
|
||||||
|
|
||||||
|
def test_acquire_exceeds_burst(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=2)
|
||||||
|
assert rl.acquire()
|
||||||
|
assert rl.acquire()
|
||||||
|
assert not rl.acquire()
|
||||||
|
|
||||||
|
def test_token_refill(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=1)
|
||||||
|
rl.acquire()
|
||||||
|
rl._last_refill = time.time() - 2
|
||||||
|
rl._tokens = 0
|
||||||
|
assert rl.acquire()
|
||||||
|
|
||||||
|
def test_wait_succeeds(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=1)
|
||||||
|
rl.acquire()
|
||||||
|
rl._tokens = 0
|
||||||
|
rl._last_refill = time.time() - 2
|
||||||
|
assert rl.wait(timeout_s=5)
|
||||||
|
|
||||||
|
def test_burst_size_respected(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=3)
|
||||||
|
for _ in range(3):
|
||||||
|
assert rl.acquire()
|
||||||
|
assert not rl.acquire()
|
||||||
|
|
||||||
|
def test_rpm_affects_refill(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=120, burst_size=1)
|
||||||
|
rl.acquire()
|
||||||
|
rl._tokens = 0
|
||||||
|
rl._last_refill = time.time() - 1
|
||||||
|
# 120 RPM = 2/sec, so 1 second should refill 2 tokens
|
||||||
|
assert rl.acquire()
|
||||||
|
|
||||||
|
def test_multiple_instances(self):
|
||||||
|
rl1 = RateLimiter(requests_per_minute=30, burst_size=2)
|
||||||
|
rl2 = RateLimiter(requests_per_minute=60, burst_size=5)
|
||||||
|
assert rl1.acquire()
|
||||||
|
assert rl2.acquire()
|
||||||
|
|
||||||
|
def test_concurrent_acquire(self):
|
||||||
|
rl = RateLimiter(requests_per_minute=60, burst_size=10)
|
||||||
|
results = [rl.acquire() for _ in range(15)]
|
||||||
|
assert sum(results) == 10
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# SOURCE CATALOGUE (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestSourceCatalogue:
|
||||||
|
def test_add_source(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "Test", "https://example.com")
|
||||||
|
assert cat.source_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_record_fetch_success(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "Test", "https://example.com")
|
||||||
|
cat.record_fetch("s1", success=True)
|
||||||
|
assert cat.get_enabled()[0].fetch_count == 1
|
||||||
|
assert cat.get_enabled()[0].error_count == 0
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_record_fetch_error(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "Test", "https://example.com")
|
||||||
|
cat.record_fetch("s1", success=False)
|
||||||
|
assert cat.get_enabled()[0].error_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_get_by_type(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "News", "https://news.com", "news")
|
||||||
|
cat.add_source("s2", "Data", "https://data.com", "data")
|
||||||
|
assert len(cat.get_by_type("news")) == 1
|
||||||
|
assert len(cat.get_by_type("data")) == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_persistence(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat1 = SourceCatalogue(path)
|
||||||
|
cat1.add_source("s1", "Test", "https://example.com")
|
||||||
|
cat1.record_fetch("s1", success=True)
|
||||||
|
|
||||||
|
# Reload
|
||||||
|
cat2 = SourceCatalogue(path)
|
||||||
|
assert cat2.source_count == 1
|
||||||
|
assert cat2.get_enabled()[0].fetch_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_source_count(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
cat = SourceCatalogue(path)
|
||||||
|
cat.add_source("s1", "A", "https://a.com")
|
||||||
|
cat.add_source("s2", "B", "https://b.com")
|
||||||
|
assert cat.source_count == 2
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# REGIME EXTRACTOR (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRegimeExtractor:
|
||||||
|
def test_flash_crash(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Bitcoin crashed 20% in 5 minutes")
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
|
||||||
|
def test_liquidity_vacuum(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Liquidity vacuum on exchange")
|
||||||
|
assert "liquidity_vacuum" in regimes
|
||||||
|
|
||||||
|
def test_normal(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Normal market conditions")
|
||||||
|
assert "normal" in regimes
|
||||||
|
|
||||||
|
def test_multiple_regimes(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Flash crash with liquidation cascade")
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
assert "liquidation" in regimes
|
||||||
|
|
||||||
|
def test_sentiment_positive(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
score = ext.extract_sentiment("Bitcoin rally surges to new highs")
|
||||||
|
assert score > 0
|
||||||
|
|
||||||
|
def test_sentiment_negative(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
score = ext.extract_sentiment("Flash crash panic liquidation")
|
||||||
|
assert score < 0
|
||||||
|
|
||||||
|
def test_sentiment_neutral(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
score = ext.extract_sentiment("Market trading normally")
|
||||||
|
assert score == 0.0
|
||||||
|
|
||||||
|
def test_empty_text(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("")
|
||||||
|
assert regimes == ["normal"]
|
||||||
|
|
||||||
|
def test_whale_activity(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Whale order detected on exchange")
|
||||||
|
assert "whale_activity" in regimes
|
||||||
|
|
||||||
|
def test_stop_hunting(self):
|
||||||
|
ext = RegimeExtractor()
|
||||||
|
regimes = ext.extract_regimes("Stop hunt observed in BTC")
|
||||||
|
assert "stop_hunting" in regimes
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# COGNITION PIPELINE (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCognitionPipeline:
|
||||||
|
def test_add_source(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
assert pipe._catalogue.source_count == 1
|
||||||
|
|
||||||
|
def test_fetch_and_extract(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
new = pipe.fetch_and_extract("s1", "Flash crash observed in BTC")
|
||||||
|
assert "flash_crash" in new
|
||||||
|
assert pipe._total_fetched == 1
|
||||||
|
|
||||||
|
def test_deduplication(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash observed")
|
||||||
|
new = pipe.fetch_and_extract("s1", "Flash crash again")
|
||||||
|
assert len(new) == 0
|
||||||
|
|
||||||
|
def test_discovered_regimes(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash observed")
|
||||||
|
regimes = pipe.get_discovered_regimes()
|
||||||
|
assert "flash_crash" in regimes
|
||||||
|
|
||||||
|
def test_source_stats(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Test text")
|
||||||
|
stats = pipe.get_source_stats()
|
||||||
|
assert stats["total_sources"] == 1
|
||||||
|
assert stats["total_fetched"] == 1
|
||||||
|
|
||||||
|
def test_seed_default_sources(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.seed_default_sources()
|
||||||
|
assert pipe._catalogue.source_count == 8
|
||||||
|
|
||||||
|
def test_rate_limiting(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null", rate_limit_rpm=1)
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
# First fetch should succeed
|
||||||
|
pipe.fetch_and_extract("s1", "Test")
|
||||||
|
# Second should be rate limited (returns empty)
|
||||||
|
new = pipe.fetch_and_extract("s1", "Test")
|
||||||
|
# May or may not be rate limited depending on timing
|
||||||
|
assert isinstance(new, list)
|
||||||
|
|
||||||
|
def test_multiple_sources(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "News", "https://news.com")
|
||||||
|
pipe.add_source("s2", "Data", "https://data.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash")
|
||||||
|
pipe.fetch_and_extract("s2", "High volatility")
|
||||||
|
assert pipe._total_fetched == 2
|
||||||
|
|
||||||
|
def test_discovered_count(self):
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash and liquidation")
|
||||||
|
assert pipe.discovered_regime_count >= 2
|
||||||
|
|
||||||
|
def test_perm_run_capable(self):
|
||||||
|
"""Pipeline should not crash after many iterations."""
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
for i in range(20):
|
||||||
|
pipe.fetch_and_extract("s1", f"Flash crash iteration {i}")
|
||||||
|
assert pipe._total_fetched >= 1 # at least some fetched
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# REGIME EXPANSION (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRegimeExpansion:
|
||||||
|
def test_generate_regimes(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
assert len(regimes) == 50
|
||||||
|
|
||||||
|
def test_regime_ids_unique(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=100)
|
||||||
|
ids = [r.regime_id for r in regimes]
|
||||||
|
assert len(ids) == len(set(ids))
|
||||||
|
|
||||||
|
def test_regime_labels_diverse(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
labels = [r.label for r in regimes]
|
||||||
|
assert len(set(labels)) > 20
|
||||||
|
|
||||||
|
def test_regime_to_scenario(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=5)
|
||||||
|
from malkhut.training.cma_trainer import Scenario
|
||||||
|
scenario = exp.regime_to_scenario(regimes[0])
|
||||||
|
assert isinstance(scenario, Scenario)
|
||||||
|
|
||||||
|
def test_dimension_coverage(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=100)
|
||||||
|
liq = set(r.liquidity.label for r in regimes)
|
||||||
|
vol = set(r.volatility.label for r in regimes)
|
||||||
|
flow = set(r.flow.label for r in regimes)
|
||||||
|
assert len(liq) == 4
|
||||||
|
assert len(vol) == 4
|
||||||
|
assert len(flow) == 4
|
||||||
|
|
||||||
|
def test_orthogonal_to_cognition(self):
|
||||||
|
from malkhut.training.cognition import CognitionPipeline
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
exp = RegimeExpander()
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe.fetch_and_extract("s1", "Flash crash")
|
||||||
|
exp_regimes = exp.generate_regimes(max_regimes=50)
|
||||||
|
assert len(exp_regimes) > 0
|
||||||
|
|
||||||
|
def test_max_regimes_capped(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=200)
|
||||||
|
assert len(regimes) <= 200
|
||||||
|
|
||||||
|
def test_regime_counterparties(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=5)
|
||||||
|
for r in regimes:
|
||||||
|
assert len(r.structure.counterparties) > 0
|
||||||
|
|
||||||
|
def test_regime_book_state(self):
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=5)
|
||||||
|
for r in regimes:
|
||||||
|
assert r.bid > 0
|
||||||
|
assert r.ask > 0
|
||||||
|
assert r.ask > r.bid
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# INTEGRATION (5 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPhase0Integration:
|
||||||
|
def test_cognition_to_matrix(self):
|
||||||
|
"""Cognition findings should be recordable to matrix."""
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
regimes = pipe.fetch_and_extract("s1", "Flash crash and liquidation")
|
||||||
|
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
for regime in regimes:
|
||||||
|
matrix.record("test_strategy", regime, score=10.0)
|
||||||
|
|
||||||
|
assert matrix.total_entries >= 1
|
||||||
|
|
||||||
|
def test_expansion_to_scenarios(self):
|
||||||
|
"""Regime expansion should produce evaluatable scenarios."""
|
||||||
|
exp = RegimeExpander()
|
||||||
|
regimes = exp.generate_regimes(max_regimes=10)
|
||||||
|
from malkhut.training.cma_trainer import Scenario
|
||||||
|
for r in regimes:
|
||||||
|
scenario = exp.regime_to_scenario(r)
|
||||||
|
assert isinstance(scenario, Scenario)
|
||||||
|
assert scenario.max_steps > 0
|
||||||
|
|
||||||
|
def test_full_pipeline_flow(self):
|
||||||
|
"""Cognition → expansion → matrix → selector."""
|
||||||
|
# 1. Cognition discovers regimes
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null")
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
cog_regimes = pipe.fetch_and_extract("s1", "Flash crash")
|
||||||
|
|
||||||
|
# 2. Expansion generates regimes
|
||||||
|
exp = RegimeExpander()
|
||||||
|
exp_regimes = exp.generate_regimes(max_regimes=10)
|
||||||
|
|
||||||
|
# 3. Matrix records performance
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
for r in cog_regimes:
|
||||||
|
matrix.record("s1", r, score=10.0)
|
||||||
|
for r in exp_regimes:
|
||||||
|
matrix.record("s2", r.label, score=8.0)
|
||||||
|
|
||||||
|
# 4. Selector queries matrix
|
||||||
|
from malkhut.training.selector import StrategySelector
|
||||||
|
selector = StrategySelector(matrix=matrix, min_episodes_for_selection=1)
|
||||||
|
strategies = {"s1": _baseline(), "s2": _baseline(version="v2")}
|
||||||
|
result = selector.select(_state(), strategies)
|
||||||
|
# Should return a valid strategy (may be fallback if matrix has too few entries)
|
||||||
|
assert result.strategy_id in strategies or result.strategy_id == "baseline"
|
||||||
|
|
||||||
|
def test_rate_limiting_prevents_abuse(self):
|
||||||
|
"""Pipeline should respect rate limits."""
|
||||||
|
pipe = CognitionPipeline(catalogue_path="/dev/null", rate_limit_rpm=2)
|
||||||
|
pipe.add_source("s1", "Test", "https://example.com")
|
||||||
|
results = []
|
||||||
|
for _ in range(5):
|
||||||
|
new = pipe.fetch_and_extract("s1", "Flash crash")
|
||||||
|
results.append(len(new))
|
||||||
|
# Some should be rate limited (return empty)
|
||||||
|
assert sum(results) < 5
|
||||||
|
|
||||||
|
def test_source_persistence(self):
|
||||||
|
"""Sources should persist across pipeline instances."""
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
pipe1 = CognitionPipeline(catalogue_path=path)
|
||||||
|
pipe1.add_source("s1", "Test", "https://example.com")
|
||||||
|
pipe1.fetch_and_extract("s1", "Flash crash")
|
||||||
|
|
||||||
|
pipe2 = CognitionPipeline(catalogue_path=path)
|
||||||
|
assert pipe2._catalogue.source_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
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 _state():
|
||||||
|
from malkhut.state import AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1, 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),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _venue():
|
||||||
|
from malkhut.state import VenueRules
|
||||||
|
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)
|
||||||
255
MALKHUT/malkhut/tests/test_pipeline.py
Normal file
255
MALKHUT/malkhut/tests/test_pipeline.py
Normal file
@@ -0,0 +1,255 @@
|
|||||||
|
"""
|
||||||
|
Tests for training pipeline and logging.
|
||||||
|
|
||||||
|
Verifies:
|
||||||
|
- Pipeline runs bounded iterations
|
||||||
|
- Early stopping on convergence
|
||||||
|
- Time budget respected
|
||||||
|
- Eval budget respected
|
||||||
|
- Logger produces compact JSONL
|
||||||
|
- Events are complete and observable
|
||||||
|
- Full pipeline flow: train → promote → reload
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
from malkhut.training.pipeline import (
|
||||||
|
TrainingPipeline, PipelineConfig, PipelineResult,
|
||||||
|
TrainingLogger, TrainingEvent,
|
||||||
|
)
|
||||||
|
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. TRAINING LOGGER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTrainingLogger:
|
||||||
|
def test_log_creates_file(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
logger = TrainingLogger(log_path=path)
|
||||||
|
logger.log(TrainingEvent(
|
||||||
|
timestamp_ns=time.time_ns(), event_type="test",
|
||||||
|
policy_version="v1", score=10.0,
|
||||||
|
))
|
||||||
|
assert os.path.exists(path)
|
||||||
|
with open(path) as f:
|
||||||
|
lines = f.readlines()
|
||||||
|
assert len(lines) == 1
|
||||||
|
record = json.loads(lines[0])
|
||||||
|
assert record["type"] == "test"
|
||||||
|
assert record["ver"] == "v1"
|
||||||
|
assert record["score"] == 10.0
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_log_run_start(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_run_start(generation=0, budget_evals=14)
|
||||||
|
assert logger.event_count == 1
|
||||||
|
assert logger.get_events()[0].event_type == "run_start"
|
||||||
|
|
||||||
|
def test_log_generation(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_generation(generation=1, best_score=5.0, mean_score=3.0, evals=14, improvement=2.0)
|
||||||
|
assert logger.event_count == 1
|
||||||
|
e = logger.get_events()[0]
|
||||||
|
assert e.score == 5.0
|
||||||
|
assert e.details["mean"] == 3.0
|
||||||
|
|
||||||
|
def test_log_candidate(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_candidate("v1", 10.0, 1)
|
||||||
|
assert logger.get_events()[0].policy_version == "v1"
|
||||||
|
|
||||||
|
def test_log_promote(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_promote("v1", "CANDIDATE", "ACTIVE", "promoted")
|
||||||
|
e = logger.get_events()[0]
|
||||||
|
assert e.event_type == "promote"
|
||||||
|
assert e.details["from"] == "CANDIDATE"
|
||||||
|
assert e.details["to"] == "ACTIVE"
|
||||||
|
|
||||||
|
def test_log_reject(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_reject("v1", "tail risk")
|
||||||
|
assert logger.get_events()[0].event_type == "reject"
|
||||||
|
|
||||||
|
def test_log_reload(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_reload("v2", "v1")
|
||||||
|
e = logger.get_events()[0]
|
||||||
|
assert e.details["old"] == "v1"
|
||||||
|
|
||||||
|
def test_log_run_end(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log_run_end(generation=5, total_evals=70, best_score=8.0, duration_s=120.0)
|
||||||
|
e = logger.get_events()[0]
|
||||||
|
assert e.details["duration_s"] == 120.0
|
||||||
|
|
||||||
|
def test_filter_by_event_type(self):
|
||||||
|
logger = TrainingLogger(log_path="/dev/null")
|
||||||
|
logger.log(TrainingEvent(timestamp_ns=1, event_type="start"))
|
||||||
|
logger.log(TrainingEvent(timestamp_ns=2, event_type="generation"))
|
||||||
|
logger.log(TrainingEvent(timestamp_ns=3, event_type="start"))
|
||||||
|
starts = logger.get_events(event_type="start")
|
||||||
|
assert len(starts) == 2
|
||||||
|
|
||||||
|
def test_jsonl_compact(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
logger = TrainingLogger(log_path=path)
|
||||||
|
logger.log(TrainingEvent(
|
||||||
|
timestamp_ns=1234567890, event_type="test",
|
||||||
|
policy_version="v1", score=10.0, generation=1, evals=14,
|
||||||
|
details={"mean": 5.0},
|
||||||
|
))
|
||||||
|
with open(path) as f:
|
||||||
|
line = f.readline()
|
||||||
|
# Compact: no spaces after separators
|
||||||
|
assert ",\"score\":10.0" in line
|
||||||
|
assert "\"mean\":5.0" in line
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. PIPELINE CONFIG
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPipelineConfig:
|
||||||
|
def test_default_config(self):
|
||||||
|
cfg = PipelineConfig()
|
||||||
|
assert cfg.max_generations == 10
|
||||||
|
assert cfg.max_evals_per_generation == 14
|
||||||
|
assert cfg.max_time_s == 300.0
|
||||||
|
assert cfg.patience == 3
|
||||||
|
|
||||||
|
def test_custom_config(self):
|
||||||
|
cfg = PipelineConfig(max_generations=5, patience=2)
|
||||||
|
assert cfg.max_generations == 5
|
||||||
|
assert cfg.patience == 2
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. PIPELINE RUN
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTrainingPipeline:
|
||||||
|
def test_pipeline_returns_result(self):
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert isinstance(result, PipelineResult)
|
||||||
|
assert result.generations_run >= 1
|
||||||
|
assert result.total_evals > 0
|
||||||
|
assert result.duration_s > 0
|
||||||
|
|
||||||
|
def test_pipeline_respects_max_generations(self):
|
||||||
|
config = PipelineConfig(max_generations=2, max_evals_per_generation=3, max_time_s=60)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert result.generations_run <= 2
|
||||||
|
|
||||||
|
def test_pipeline_respects_time_budget(self):
|
||||||
|
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, max_time_s=2.0)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
# Should stop within reasonable time
|
||||||
|
assert result.generations_run <= 10 # not all 100
|
||||||
|
assert result.duration_s < 60.0
|
||||||
|
|
||||||
|
def test_pipeline_early_stopping(self):
|
||||||
|
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, patience=1, max_time_s=60)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert result.generations_run <= 100
|
||||||
|
|
||||||
|
def test_pipeline_logs_events(self):
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
events = pipeline.logger.get_events()
|
||||||
|
assert len(events) > 0
|
||||||
|
event_types = [e.event_type for e in events]
|
||||||
|
assert "run_start" in event_types
|
||||||
|
assert "run_end" in event_types
|
||||||
|
assert "generation" in event_types
|
||||||
|
|
||||||
|
def test_pipeline_has_best_score(self):
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert isinstance(result.best_score, float)
|
||||||
|
|
||||||
|
def test_pipeline_registry_has_records(self):
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30, auto_promote=True)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert pipeline.registry.record_count > 0
|
||||||
|
|
||||||
|
def test_pipeline_result_events_match_logger(self):
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
assert len(result.events) == pipeline.logger.event_count
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. FULL OBSERVABLE FLOW
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestFullObservableFlow:
|
||||||
|
def test_train_log_promote_activate(self):
|
||||||
|
"""Full observable flow: train → log → promote → activate → engine loads."""
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
|
||||||
|
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||||
|
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||||
|
|
||||||
|
# Engine starts with default
|
||||||
|
engine = FulfilmentEngine(registry=pipeline.registry)
|
||||||
|
|
||||||
|
# Run pipeline
|
||||||
|
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||||
|
|
||||||
|
# Verify logging happened
|
||||||
|
assert len(result.events) > 0
|
||||||
|
|
||||||
|
# Verify registry has promoted policies
|
||||||
|
active = pipeline.registry.load_active()
|
||||||
|
if active:
|
||||||
|
# Hot reload engine
|
||||||
|
engine.hot_reload_policy(active)
|
||||||
|
assert engine.params_provider().version == active.version
|
||||||
|
|
||||||
|
engine.close()
|
||||||
201
MALKHUT/malkhut/tests/test_planner.py
Normal file
201
MALKHUT/malkhut/tests/test_planner.py
Normal file
@@ -0,0 +1,201 @@
|
|||||||
|
"""
|
||||||
|
Unit tests: SM-MCTS planner and action menu.
|
||||||
|
|
||||||
|
Mutation litmus: flip UCB exploration constant; if no test breaks,
|
||||||
|
the planner is untested.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner, PlayerActionStats
|
||||||
|
from malkhut.planner.action_menu import build_our_actions
|
||||||
|
from malkhut.cwm import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, PlannedPolicy
|
||||||
|
|
||||||
|
|
||||||
|
def _state_with_intent() -> MarketWorldState:
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, 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_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
intent=ExecutionIntent(
|
||||||
|
intent_id="test_intent", 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 _params() -> FulfilmentPolicyParams:
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="test", ucb_c=1.414, max_sims=32, max_depth=2,
|
||||||
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.10, 0.25, 0.50),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestActionMenu:
|
||||||
|
def test_noop_always_available(self):
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
kinds = [a.kind for a in actions]
|
||||||
|
assert ActionKind.NOOP in kinds
|
||||||
|
|
||||||
|
def test_action_count_bounded(self):
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
assert 3 <= len(actions) <= 50
|
||||||
|
|
||||||
|
def test_cross_spread_only_when_urgency_high(self):
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
crosses = [a for a in actions if a.kind == ActionKind.CROSS_SPREAD]
|
||||||
|
# urgency=0.5 < 0.65, so no crosses
|
||||||
|
assert len(crosses) == 0
|
||||||
|
|
||||||
|
def test_high_urgency_allows_crosses(self):
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
intent = ExecutionIntent(
|
||||||
|
intent_id="hi", ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0,
|
||||||
|
urgency=0.8, 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",
|
||||||
|
)
|
||||||
|
state = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=state.venue, book=state.book,
|
||||||
|
account=state.account, intent=intent,
|
||||||
|
)
|
||||||
|
actions = build_our_actions(state, params)
|
||||||
|
crosses = [a for a in actions if a.kind == ActionKind.CROSS_SPREAD]
|
||||||
|
assert len(crosses) > 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlayerActionStats:
|
||||||
|
def test_every_action_gets_initial_visit(self):
|
||||||
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),
|
||||||
|
FulfilmentAction(ActionKind.PLACE, Side.BUY, None, 0, 0.25, 200))
|
||||||
|
stats = PlayerActionStats.from_actions(actions)
|
||||||
|
assert all(v == 0 for v in stats.visits)
|
||||||
|
|
||||||
|
def test_ucb_prefers_high_value_after_sampling(self):
|
||||||
|
import random
|
||||||
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),
|
||||||
|
FulfilmentAction(ActionKind.PLACE, Side.BUY, None, 0, 0.25, 200))
|
||||||
|
stats = PlayerActionStats.from_actions(actions)
|
||||||
|
for _ in range(100):
|
||||||
|
stats.update(0, 0.1)
|
||||||
|
stats.update(1, 1.0)
|
||||||
|
rng = random.Random(42)
|
||||||
|
idx, action = stats.ucb_select(200, 1.414, rng)
|
||||||
|
assert idx == 1
|
||||||
|
|
||||||
|
def test_update_increments_count(self):
|
||||||
|
actions = (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),)
|
||||||
|
stats = PlayerActionStats.from_actions(actions)
|
||||||
|
stats.update(0, 1.0)
|
||||||
|
assert stats.visits[0] == 1
|
||||||
|
assert abs(stats.total_value[0] - 1.0) < 1e-9
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlanner:
|
||||||
|
def test_planner_returns_planned_policy(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||||||
|
)
|
||||||
|
params = _params()
|
||||||
|
state = _state_with_intent()
|
||||||
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||||||
|
assert isinstance(result, PlannedPolicy)
|
||||||
|
assert len(result.actions) > 0
|
||||||
|
assert len(result.probabilities) == len(result.actions)
|
||||||
|
assert abs(sum(result.probabilities) - 1.0) < 1e-6
|
||||||
|
|
||||||
|
def test_root_policy_probabilities_sum_to_one(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42,
|
||||||
|
)
|
||||||
|
params = _params()
|
||||||
|
state = _state_with_intent()
|
||||||
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||||||
|
total = sum(result.probabilities)
|
||||||
|
assert abs(total - 1.0) < 1e-6
|
||||||
|
|
||||||
|
def test_deterministic_with_fixed_seed(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
state = _state_with_intent()
|
||||||
|
params = _params()
|
||||||
|
|
||||||
|
p1 = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=99)
|
||||||
|
r1 = p1.plan(root_state=state, params=params, budget_ms=20)
|
||||||
|
|
||||||
|
p2 = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=99)
|
||||||
|
r2 = p2.plan(root_state=state, params=params, budget_ms=20)
|
||||||
|
|
||||||
|
assert r1.selected_action.kind == r2.selected_action.kind
|
||||||
|
|
||||||
|
def test_no_intent_returns_noop(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
planner = DecoupledUCBPlanner(
|
||||||
|
cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=0,
|
||||||
|
)
|
||||||
|
params = _params()
|
||||||
|
state = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
intent=None,
|
||||||
|
)
|
||||||
|
result = planner.plan(root_state=state, params=params, budget_ms=20)
|
||||||
|
assert result.selected_action.kind == ActionKind.NOOP
|
||||||
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
|
||||||
259
MALKHUT/malkhut/tests/test_prod_tooling.py
Normal file
259
MALKHUT/malkhut/tests/test_prod_tooling.py
Normal file
@@ -0,0 +1,259 @@
|
|||||||
|
"""
|
||||||
|
Tests for cognition launcher, news sources, and monitor.
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.cognition_launcher import CognitionLauncher, CognitionConfig
|
||||||
|
from malkhut.training.news_sources import NewsSourceRepository, SourceMetadata
|
||||||
|
from malkhut.training.monitor import CognitionMonitor, CognitionMetrics
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# COGNITION LAUNCHER (8 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCognitionLauncher:
|
||||||
|
def test_launcher_creates(self):
|
||||||
|
launcher = CognitionLauncher()
|
||||||
|
assert launcher._running is False
|
||||||
|
assert launcher._total_fetched == 0
|
||||||
|
|
||||||
|
def test_launcher_config(self):
|
||||||
|
cfg = CognitionConfig(rate_limit_rpm=10, fetch_interval_s=5)
|
||||||
|
launcher = CognitionLauncher(config=cfg)
|
||||||
|
assert launcher.config.rate_limit_rpm == 10
|
||||||
|
|
||||||
|
def test_launcher_loads_regimes(self):
|
||||||
|
path = '/tmp/test_regimes_e2e.json'
|
||||||
|
with open(path, 'w') as f:
|
||||||
|
json.dump({"flash_crash": {"source": "test"}}, f)
|
||||||
|
try:
|
||||||
|
launcher = CognitionLauncher(CognitionConfig(regime_db_path=path))
|
||||||
|
assert launcher._total_regimes == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_launcher_saves_regimes(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
launcher = CognitionLauncher(CognitionConfig(regime_db_path=path))
|
||||||
|
launcher._discovered_regimes["test"] = {"source": "test"}
|
||||||
|
launcher._save_regimes()
|
||||||
|
with open(path) as f:
|
||||||
|
data = json.load(f)
|
||||||
|
assert "test" in data
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# NEWS SOURCE REPOSITORY (15 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestNewsSourceRepository:
|
||||||
|
def test_register_source(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "Test", "https://example.com", "news", 0.8, 0.9, 1.0)
|
||||||
|
assert repo.source_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_rank_by_relevance(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "Low", "https://a.com", "news", 0.3, 0.8, 1.0)
|
||||||
|
repo.register("s2", "High", "https://b.com", "news", 0.9, 0.9, 1.0)
|
||||||
|
ranked = repo.rank_by_relevance()
|
||||||
|
assert ranked[0].source_id == "s2"
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_rank_by_health(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "Good", "https://a.com", reliability=0.95)
|
||||||
|
repo.register("s2", "Bad", "https://b.com", reliability=0.5)
|
||||||
|
ranked = repo.rank_by_health()
|
||||||
|
assert ranked[0].source_id == "s1"
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_get_by_type(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "News", "https://news.com", "news")
|
||||||
|
repo.register("s2", "Data", "https://data.com", "data")
|
||||||
|
assert len(repo.get_by_type("news")) == 1
|
||||||
|
assert len(repo.get_by_type("data")) == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_get_by_tag(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "Test", "https://example.com", tags=("crypto", "news"))
|
||||||
|
assert len(repo.get_by_tag("crypto")) == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_disable_enable(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.register("s1", "Test", "https://example.com")
|
||||||
|
assert repo.enabled_count == 1
|
||||||
|
repo.disable("s1")
|
||||||
|
assert repo.enabled_count == 0
|
||||||
|
repo.enable("s1")
|
||||||
|
assert repo.enabled_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_persistence(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo1 = NewsSourceRepository(path)
|
||||||
|
repo1.register("s1", "Test", "https://example.com")
|
||||||
|
repo1.record_fetch("s1", success=True)
|
||||||
|
|
||||||
|
repo2 = NewsSourceRepository(path)
|
||||||
|
assert repo2.source_count == 1
|
||||||
|
assert repo2.get_enabled()[0].fetch_count == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_seed_defaults(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
repo = NewsSourceRepository(path)
|
||||||
|
repo.seed_defaults()
|
||||||
|
assert repo.source_count >= 10
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_health_score(self):
|
||||||
|
s = SourceMetadata(source_id="s1", name="Test", url="https://x.com",
|
||||||
|
source_type="news", regime_relevance=0.5, reliability=0.9,
|
||||||
|
freshness_hours=1.0, fetch_count=10, error_count=1)
|
||||||
|
assert s.health_score > 0
|
||||||
|
assert s.health_score < 1.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# COGNITION MONITOR (10 tests)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCognitionMonitor:
|
||||||
|
def test_record_fetch(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon.record_fetch("s1", success=True)
|
||||||
|
assert mon.total_fetched == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_record_error(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon.record_fetch("s1", success=False)
|
||||||
|
assert mon.total_errors == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_record_regime(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon.record_regime(3)
|
||||||
|
assert mon.discovered_regimes == 3
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_snapshot(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon.record_fetch("s1", success=True)
|
||||||
|
metrics = mon.snapshot(total_sources=5, enabled_sources=4)
|
||||||
|
assert isinstance(metrics, CognitionMetrics)
|
||||||
|
assert metrics.total_fetched == 1
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_metrics_logged(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon.snapshot(total_sources=5, enabled_sources=4)
|
||||||
|
with open(path) as f:
|
||||||
|
lines = f.readlines()
|
||||||
|
assert len(lines) == 1
|
||||||
|
record = json.loads(lines[0])
|
||||||
|
assert "sources" in record
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_check_alerts_high_error(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
for _ in range(20):
|
||||||
|
mon.record_fetch("s1", success=False)
|
||||||
|
mon.record_fetch("s1", success=True)
|
||||||
|
mon.snapshot(total_sources=1, enabled_sources=1)
|
||||||
|
alerts = mon.check_alerts()
|
||||||
|
assert len(alerts) > 0
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_check_alerts_low_fetch_rate(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
mon._start_time = time.time() - 600 # 10 min ago
|
||||||
|
mon.record_fetch("s1", success=True) # only 1 fetch in 10 min
|
||||||
|
mon.snapshot(total_sources=1, enabled_sources=1)
|
||||||
|
alerts = mon.check_alerts()
|
||||||
|
assert any("LOW_FETCH_RATE" in a for a in alerts)
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
|
|
||||||
|
def test_check_alerts_clean(self):
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||||
|
path = f.name
|
||||||
|
try:
|
||||||
|
mon = CognitionMonitor(path)
|
||||||
|
for _ in range(100):
|
||||||
|
mon.record_fetch("s1", success=True)
|
||||||
|
mon.snapshot(total_sources=1, enabled_sources=1)
|
||||||
|
alerts = mon.check_alerts()
|
||||||
|
assert len(alerts) == 0
|
||||||
|
finally:
|
||||||
|
os.unlink(path)
|
||||||
221
MALKHUT/malkhut/tests/test_registry.py
Normal file
221
MALKHUT/malkhut/tests/test_registry.py
Normal file
@@ -0,0 +1,221 @@
|
|||||||
|
"""
|
||||||
|
Tests for policy registry and the learning/improvement loop.
|
||||||
|
|
||||||
|
Verifies that:
|
||||||
|
- Trained policies are persisted to CH
|
||||||
|
- Registry manages lifecycle (CANDIDATE → ACTIVE)
|
||||||
|
- Engine loads active policy from registry
|
||||||
|
- Hot-reload via control plane works
|
||||||
|
- Full training → registry → engine flow works
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord
|
||||||
|
from malkhut.training.cma_trainer import (
|
||||||
|
CMAParameterCodec, PolicySnapshot, SelfPlayPool,
|
||||||
|
)
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. POLICY REGISTRY
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPolicyRegistry:
|
||||||
|
def test_register_candidate(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
record = reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
assert record.stage == PolicyStage.CANDIDATE
|
||||||
|
assert record.version == "v1"
|
||||||
|
assert record.score == 10.0
|
||||||
|
|
||||||
|
def test_promote_through_stages(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.BACKTESTED, "tests passed")
|
||||||
|
reg.promote("v1", PolicyStage.SELF_PLAY_CONFIRMED, "pool hardened")
|
||||||
|
reg.promote("v1", PolicyStage.SHADOW, "shadow verified")
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE, "promoted")
|
||||||
|
record = reg.get_record("v1")
|
||||||
|
assert record.stage == PolicyStage.ACTIVE
|
||||||
|
|
||||||
|
def test_load_active(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
params = reg.load_active()
|
||||||
|
assert params is not None
|
||||||
|
assert params.version == "v1"
|
||||||
|
|
||||||
|
def test_no_active_returns_none(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
assert reg.load_active() is None
|
||||||
|
|
||||||
|
def test_reject(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.reject("v1", "tail risk too high")
|
||||||
|
record = reg.get_record("v1")
|
||||||
|
assert record.stage == PolicyStage.REJECTED
|
||||||
|
|
||||||
|
def test_retire(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
reg.retire("v1", "replaced by v2")
|
||||||
|
record = reg.get_record("v1")
|
||||||
|
assert record.stage == PolicyStage.RETIRED
|
||||||
|
|
||||||
|
def test_only_one_active(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
reg.register_candidate(_baseline(version="v2"), score=12.0)
|
||||||
|
reg.promote("v2", PolicyStage.ACTIVE)
|
||||||
|
# v1 should be retired when v2 becomes active
|
||||||
|
# (or we just have two ACTIVE — the load_active returns latest)
|
||||||
|
params = reg.load_active()
|
||||||
|
assert params.version == "v2"
|
||||||
|
|
||||||
|
def test_get_by_stage(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.register_candidate(_baseline(version="v2"), score=12.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
candidates = reg.get_by_stage(PolicyStage.CANDIDATE)
|
||||||
|
assert len(candidates) == 1
|
||||||
|
assert candidates[0].version == "v2"
|
||||||
|
|
||||||
|
def test_record_count(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
assert reg.record_count == 0
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
assert reg.record_count == 1
|
||||||
|
|
||||||
|
def test_unknown_version_raises(self):
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
with pytest.raises(KeyError):
|
||||||
|
reg.promote("nonexistent", PolicyStage.ACTIVE)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. ENGINE-REGISTRY INTEGRATION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestEngineRegistryIntegration:
|
||||||
|
def test_engine_loads_active_policy(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
|
||||||
|
engine = FulfilmentEngine(registry=reg)
|
||||||
|
params = engine.params_provider()
|
||||||
|
assert params.version == "v1"
|
||||||
|
engine.close()
|
||||||
|
|
||||||
|
def test_engine_hot_reload(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
engine = FulfilmentEngine()
|
||||||
|
engine.hot_reload_policy(_baseline(version="v2"))
|
||||||
|
params = engine.params_provider()
|
||||||
|
assert params.version == "v2"
|
||||||
|
engine.close()
|
||||||
|
|
||||||
|
def test_engine_hot_reload_via_registry(self):
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
reg = PolicyRegistry()
|
||||||
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||||||
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||||||
|
|
||||||
|
engine = FulfilmentEngine(registry=reg)
|
||||||
|
# Initially v1
|
||||||
|
assert engine.params_provider().version == "v1"
|
||||||
|
|
||||||
|
# Register v2 and promote
|
||||||
|
reg.register_candidate(_baseline(version="v2"), score=12.0)
|
||||||
|
reg.promote("v2", PolicyStage.ACTIVE)
|
||||||
|
|
||||||
|
# Hot reload to v2
|
||||||
|
engine.hot_reload_policy(reg.load_active())
|
||||||
|
assert engine.params_provider().version == "v2"
|
||||||
|
engine.close()
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. FULL LEARNING LOOP
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestFullLearningLoop:
|
||||||
|
def test_train_register_promote_activate(self):
|
||||||
|
"""Full loop: train → register → promote → activate → engine loads."""
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
from malkhut.training.cma_trainer import CMAESTrainer, PolicyEvaluator
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.training.cma_trainer import ScenarioFactory
|
||||||
|
|
||||||
|
# Setup
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
registry = PolicyRegistry()
|
||||||
|
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
# Register baseline as initial active
|
||||||
|
registry.register_candidate(_baseline(version="init"), score=0.0)
|
||||||
|
registry.promote("init", PolicyStage.ACTIVE)
|
||||||
|
|
||||||
|
# Engine starts with init policy
|
||||||
|
engine = FulfilmentEngine(registry=registry)
|
||||||
|
assert engine.params_provider().version == "init"
|
||||||
|
|
||||||
|
# Train
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
best = trainer.train(
|
||||||
|
incumbent=_baseline(version="init"),
|
||||||
|
scenarios=scenarios,
|
||||||
|
budget_evals=14,
|
||||||
|
seed=42,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Register trained candidate
|
||||||
|
registry.register_candidate(best.params, score=best.score)
|
||||||
|
|
||||||
|
# Promote through pipeline
|
||||||
|
registry.promote(best.params.version, PolicyStage.BACKTESTED, "tests passed")
|
||||||
|
registry.promote(best.params.version, PolicyStage.SELF_PLAY_CONFIRMED, "pool hardened")
|
||||||
|
registry.promote(best.params.version, PolicyStage.ACTIVE, "promoted")
|
||||||
|
|
||||||
|
# Hot reload engine to new policy
|
||||||
|
engine.hot_reload_policy(registry.load_active())
|
||||||
|
assert engine.params_provider().version == best.params.version
|
||||||
|
|
||||||
|
engine.close()
|
||||||
600
MALKHUT/malkhut/tests/test_replay_exhaustive.py
Normal file
600
MALKHUT/malkhut/tests/test_replay_exhaustive.py
Normal file
@@ -0,0 +1,600 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive replay verification tests.
|
||||||
|
|
||||||
|
Categories:
|
||||||
|
1. ReplayStep construction
|
||||||
|
2. Deep state comparison (all field types)
|
||||||
|
3. Tolerance handling (float, int, string)
|
||||||
|
4. Binary search (first mismatch, no mismatch, boundary)
|
||||||
|
5. Trajectory recording (record, hash, determinism)
|
||||||
|
6. ReplayVerifier.verify() (historical mode)
|
||||||
|
7. ReplayVerifier.verify_determinism() (self-play mode)
|
||||||
|
8. Edge cases (empty replay, single step, identical states)
|
||||||
|
9. hftbacktest integration hooks
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OpenOrderState, OrderBookState, PositionState, PriceLevel, Side,
|
||||||
|
VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, CounterpartyAction, AgentRole, FulfilmentAction, OrderType
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.cwm.replay_verify import (
|
||||||
|
ReplayStep, ReplayMismatch, ReplayResult, TrajectoryRecord,
|
||||||
|
TrajectoryRecorder, ReplayVerifier, _compare_deep, _hash_state,
|
||||||
|
bisect_first_mismatch,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _venue(**kw):
|
||||||
|
d = dict(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)
|
||||||
|
d.update(kw)
|
||||||
|
return VenueRules(**d)
|
||||||
|
|
||||||
|
|
||||||
|
def _book(bid=50000.0, ask=50001.0, bid_qty=1.0, ask_qty=1.0):
|
||||||
|
return OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(bid, bid_qty),),
|
||||||
|
asks=(PriceLevel(ask, ask_qty),),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _account(equity=10000.0):
|
||||||
|
return AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=equity, wallet_balance=equity,
|
||||||
|
available_balance=equity, margin_used=0.0, total_notional=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _state(bid=50000.0, ask=50001.0, equity=10000.0, ts=1_000_000_000, **kw):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=ts, mode=Mode.REPLAY_NO_IMPACT,
|
||||||
|
venue=kw.get("venue", _venue()),
|
||||||
|
book=_book(bid, ask),
|
||||||
|
account=_account(equity),
|
||||||
|
open_orders=kw.get("open_orders", ()),
|
||||||
|
trade_path=kw.get("trade_path"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def _cross(side, frac=0.1):
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. REPLAY STEP
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestReplayStep:
|
||||||
|
def test_construction(self):
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
step = ReplayStep(before=s, joint_action=(a,), after_ground_truth=after)
|
||||||
|
assert step.before.ts_ns == s.ts_ns
|
||||||
|
assert step.after_ground_truth.ts_ns == after.ts_ns
|
||||||
|
|
||||||
|
def test_with_metadata(self):
|
||||||
|
s = _state()
|
||||||
|
step = ReplayStep(
|
||||||
|
before=s, joint_action=(_noop(),), after_ground_truth=s,
|
||||||
|
metadata={"source": "test"},
|
||||||
|
)
|
||||||
|
assert step.metadata["source"] == "test"
|
||||||
|
|
||||||
|
def test_frozen(self):
|
||||||
|
s = _state()
|
||||||
|
step = ReplayStep(before=s, joint_action=(_noop(),), after_ground_truth=s)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
step.step_index = 5
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. DEEP STATE COMPARISON
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDeepComparison:
|
||||||
|
def test_identical_states_no_mismatch(self):
|
||||||
|
s = _state()
|
||||||
|
diffs = _compare_deep(0, s, s)
|
||||||
|
assert len(diffs) == 0
|
||||||
|
|
||||||
|
def test_different_ts_detected(self):
|
||||||
|
s1 = _state(ts=1)
|
||||||
|
s2 = _state(ts=2)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "ts_ns" for d in diffs)
|
||||||
|
|
||||||
|
def test_different_equity_detected(self):
|
||||||
|
s1 = _state(equity=10000.0)
|
||||||
|
s2 = _state(equity=9000.0)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "account.equity" for d in diffs)
|
||||||
|
|
||||||
|
def test_different_bid_detected(self):
|
||||||
|
s1 = _state(bid=50000.0)
|
||||||
|
s2 = _state(bid=50001.0)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "book.best_bid" for d in diffs)
|
||||||
|
|
||||||
|
def test_different_ask_detected(self):
|
||||||
|
s1 = _state(ask=50001.0)
|
||||||
|
s2 = _state(ask=50002.0)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "book.best_ask" for d in diffs)
|
||||||
|
|
||||||
|
def test_within_tolerance_no_mismatch(self):
|
||||||
|
s1 = _state(bid=50000.0)
|
||||||
|
s2 = _state(bid=50000.0001)
|
||||||
|
diffs = _compare_deep(0, s1, s2, {"book_price": 0.001})
|
||||||
|
assert len(diffs) == 0
|
||||||
|
|
||||||
|
def test_outside_tolerance_mismatch(self):
|
||||||
|
s1 = _state(bid=50000.0)
|
||||||
|
s2 = _state(bid=50000.1)
|
||||||
|
diffs = _compare_deep(0, s1, s2, {"book_price": 0.001})
|
||||||
|
assert len(diffs) > 0
|
||||||
|
|
||||||
|
def test_different_symbol_detected(self):
|
||||||
|
v1 = _venue(symbol="BTCUSDT")
|
||||||
|
v2 = _venue(symbol="ETHUSDT")
|
||||||
|
s1 = _state(venue=v1)
|
||||||
|
s2 = _state(venue=v2)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "venue.symbol" for d in diffs)
|
||||||
|
|
||||||
|
def test_different_open_order_count(self):
|
||||||
|
oo = OpenOrderState(
|
||||||
|
client_order_id="c1", venue_order_id="v1", symbol="BTCUSDT",
|
||||||
|
side=Side.BUY, order_type=OrderType.LIMIT, price=50000.0,
|
||||||
|
qty=0.001, remaining_qty=0.001, queue_ahead_estimate=0.001,
|
||||||
|
created_ts_ns=1, last_update_ts_ns=1,
|
||||||
|
)
|
||||||
|
s1 = _state(open_orders=(oo,))
|
||||||
|
s2 = _state(open_orders=())
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any(d.field == "open_orders.count" for d in diffs)
|
||||||
|
|
||||||
|
def test_different_position_detected(self):
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s1 = _state()
|
||||||
|
s2 = MarketWorldState(
|
||||||
|
ts_ns=1, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=_book(), account=AccountState(
|
||||||
|
ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
positions={"BTCUSDT": pos},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
assert any("positions" in d.field for d in diffs)
|
||||||
|
|
||||||
|
def test_missing_position_critical(self):
|
||||||
|
pos = PositionState(
|
||||||
|
symbol="BTCUSDT", qty=0.1, avg_entry=50000.0,
|
||||||
|
unrealized_pnl=0.0, realized_pnl=0.0,
|
||||||
|
liquidation_price=None, leverage=0.5, side=Side.BUY,
|
||||||
|
)
|
||||||
|
s1 = MarketWorldState(
|
||||||
|
ts_ns=1, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=_book(), account=AccountState(
|
||||||
|
ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
positions={"BTCUSDT": pos},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
s2 = _state()
|
||||||
|
diffs = _compare_deep(0, s1, s2)
|
||||||
|
critical = [d for d in diffs if d.is_critical]
|
||||||
|
assert len(critical) > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. HASH FUNCTIONS
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestHashing:
|
||||||
|
def test_hash_state_deterministic(self):
|
||||||
|
s = _state()
|
||||||
|
h1 = _hash_state(s)
|
||||||
|
h2 = _hash_state(s)
|
||||||
|
assert h1 == h2
|
||||||
|
|
||||||
|
def test_hash_state_different_for_different_states(self):
|
||||||
|
s1 = _state(bid=50000.0)
|
||||||
|
s2 = _state(bid=51000.0)
|
||||||
|
assert _hash_state(s1) != _hash_state(s2)
|
||||||
|
|
||||||
|
def test_hash_action_deterministic(self):
|
||||||
|
a = _noop()
|
||||||
|
from malkhut.cwm.replay_verify import _hash_action
|
||||||
|
h1 = _hash_action(a)
|
||||||
|
h2 = _hash_action(a)
|
||||||
|
assert h1 == h2
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. BINARY SEARCH
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestBisect:
|
||||||
|
def test_bisect_no_mismatch(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
replay = [
|
||||||
|
ReplayStep(before=s, joint_action=(a,), after_ground_truth=cwm.transition(s, (a,))),
|
||||||
|
]
|
||||||
|
result = bisect_first_mismatch(cwm, replay)
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
def test_bisect_finds_first_mismatch(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r1 = cwm.transition(s, (a,))
|
||||||
|
r2 = cwm.transition(r1, (a,))
|
||||||
|
|
||||||
|
# Tamper with step 1
|
||||||
|
bad = MarketWorldState(
|
||||||
|
ts_ns=r1.ts_ns + 999, mode=r1.mode, venue=r1.venue,
|
||||||
|
book=r1.book, account=r1.account,
|
||||||
|
)
|
||||||
|
replay = [
|
||||||
|
ReplayStep(before=s, joint_action=(a,), after_ground_truth=r1),
|
||||||
|
ReplayStep(before=r1, joint_action=(a,), after_ground_truth=bad),
|
||||||
|
ReplayStep(before=r2, joint_action=(a,), after_ground_truth=cwm.transition(r2, (a,))),
|
||||||
|
]
|
||||||
|
result = bisect_first_mismatch(cwm, replay)
|
||||||
|
assert result is not None
|
||||||
|
assert result.index == 1
|
||||||
|
|
||||||
|
def test_bisect_empty_replay(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
result = bisect_first_mismatch(cwm, [])
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
def test_bisect_single_step_match(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
replay = [ReplayStep(before=s, joint_action=(a,), after_ground_truth=cwm.transition(s, (a,)))]
|
||||||
|
result = bisect_first_mismatch(cwm, replay)
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
def test_bisect_single_step_mismatch(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
bad = MarketWorldState(
|
||||||
|
ts_ns=999, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=_book(50000.0, 50001.0), account=_account(9999.0),
|
||||||
|
)
|
||||||
|
replay = [ReplayStep(before=s, joint_action=(a,), after_ground_truth=bad)]
|
||||||
|
result = bisect_first_mismatch(cwm, replay)
|
||||||
|
assert result is not None
|
||||||
|
assert result.index == 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. TRAJECTORY RECORDING
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestTrajectoryRecorder:
|
||||||
|
def test_record_step(self):
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(0, s, (a,), after)
|
||||||
|
assert rec.step_count == 1
|
||||||
|
|
||||||
|
def test_trajectory_hash_deterministic(self):
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
for i in range(5):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
h1 = rec.trajectory_hash()
|
||||||
|
h2 = rec.trajectory_hash()
|
||||||
|
assert h1 == h2
|
||||||
|
|
||||||
|
def test_max_steps_respected(self):
|
||||||
|
rec = TrajectoryRecorder(max_steps=3)
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
for i in range(10):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
assert rec.step_count == 3
|
||||||
|
|
||||||
|
def test_verify_deterministic_passes(self):
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
for i in range(3):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
ok, mismatches = rec.verify_deterministic(cwm)
|
||||||
|
assert ok
|
||||||
|
assert len(mismatches) == 0
|
||||||
|
|
||||||
|
def test_to_replay_steps(self):
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(0, s, (a,), after)
|
||||||
|
steps = rec.to_replay_steps()
|
||||||
|
assert len(steps) == 1
|
||||||
|
assert isinstance(steps[0], ReplayStep)
|
||||||
|
|
||||||
|
def test_steps_property(self):
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(0, s, (a,), after)
|
||||||
|
assert len(rec.steps) == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 6. REPLAY VERIFIER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestReplayVerifier:
|
||||||
|
def _verifier(self):
|
||||||
|
return ReplayVerifier()
|
||||||
|
|
||||||
|
def _make_replay(self, steps=3):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
replay = []
|
||||||
|
for i in range(steps):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
replay.append(ReplayStep(before=s, joint_action=(a,), after_ground_truth=after))
|
||||||
|
s = after
|
||||||
|
return replay
|
||||||
|
|
||||||
|
def test_verify_passes_for_correct_replay(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(3)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
assert result.steps_verified == 3
|
||||||
|
|
||||||
|
def test_verify_fails_for_tampered_replay(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(3)
|
||||||
|
# Tamper with step 1
|
||||||
|
bad = MarketWorldState(
|
||||||
|
ts_ns=999, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=_book(50000.0, 50001.0), account=_account(9999.0),
|
||||||
|
)
|
||||||
|
replay[1] = ReplayStep(
|
||||||
|
before=replay[1].before, joint_action=replay[1].joint_action,
|
||||||
|
after_ground_truth=bad,
|
||||||
|
)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert not result.passed
|
||||||
|
assert result.first_mismatch_index == 1
|
||||||
|
|
||||||
|
def test_verify_determinism_passes(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(3)
|
||||||
|
result = v.verify_determinism(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
def test_verify_result_has_trajectory_hash(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(3)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert len(result.trajectory_hash) == 16
|
||||||
|
|
||||||
|
def test_verify_result_has_timing(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(3)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert result.duration_ns > 0
|
||||||
|
|
||||||
|
def test_verify_empty_replay(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
result = v.verify(cwm, [])
|
||||||
|
assert result.passed
|
||||||
|
assert result.steps_verified == 0
|
||||||
|
|
||||||
|
def test_verify_single_step(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(1)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
def test_verify_with_tolerance(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
# Create slightly different ground truth
|
||||||
|
bad = MarketWorldState(
|
||||||
|
ts_ns=after.ts_ns, mode=after.mode, venue=after.venue,
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=after.book.ts_ns, symbol=after.book.symbol,
|
||||||
|
bids=(PriceLevel(after.book.best_bid + 0.001, 1.0),),
|
||||||
|
asks=after.book.asks,
|
||||||
|
),
|
||||||
|
account=after.account,
|
||||||
|
)
|
||||||
|
replay = [ReplayStep(before=s, joint_action=(a,), after_ground_truth=bad)]
|
||||||
|
# Tight tolerance → mismatch
|
||||||
|
result = v.verify(cwm, replay, tolerances={"book_price": 1e-9})
|
||||||
|
assert not result.passed
|
||||||
|
# Loose tolerance → match
|
||||||
|
result2 = v.verify(cwm, replay, tolerances={"book_price": 0.01})
|
||||||
|
assert result2.passed
|
||||||
|
|
||||||
|
def test_bisect_finds_first(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
r1 = cwm.transition(s, (a,))
|
||||||
|
r2 = cwm.transition(r1, (a,))
|
||||||
|
|
||||||
|
bad = MarketWorldState(
|
||||||
|
ts_ns=r1.ts_ns + 999, mode=r1.mode, venue=r1.venue,
|
||||||
|
book=r1.book, account=r1.account,
|
||||||
|
)
|
||||||
|
replay = [
|
||||||
|
ReplayStep(before=s, joint_action=(a,), after_ground_truth=r1),
|
||||||
|
ReplayStep(before=r1, joint_action=(a,), after_ground_truth=bad),
|
||||||
|
ReplayStep(before=r2, joint_action=(a,), after_ground_truth=cwm.transition(r2, (a,))),
|
||||||
|
]
|
||||||
|
result = v.bisect(cwm, replay)
|
||||||
|
assert result is not None
|
||||||
|
assert result.index == 1
|
||||||
|
|
||||||
|
def test_result_match_rate(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(5)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert result.match_rate == 1.0
|
||||||
|
|
||||||
|
def test_result_critical_count(self):
|
||||||
|
v = self._verifier()
|
||||||
|
result = ReplayResult(
|
||||||
|
passed=False, mismatches=[
|
||||||
|
ReplayMismatch(0, "x", 1, 2, "critical"),
|
||||||
|
ReplayMismatch(0, "y", 1, 2, "warning"),
|
||||||
|
],
|
||||||
|
steps_verified=1, total_steps=1, first_mismatch_index=0,
|
||||||
|
duration_ns=0, trajectory_hash="abc",
|
||||||
|
)
|
||||||
|
assert result.critical_count == 1
|
||||||
|
assert result.warning_count == 1
|
||||||
|
|
||||||
|
def test_replay_step_index(self):
|
||||||
|
v = self._verifier()
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
replay = self._make_replay(2)
|
||||||
|
result = v.verify(cwm, replay)
|
||||||
|
assert result.steps_verified == 2
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 7. REPLAY RESULT
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestReplayResult:
|
||||||
|
def test_passed_result(self):
|
||||||
|
r = ReplayResult(
|
||||||
|
passed=True, mismatches=[], steps_verified=10, total_steps=10,
|
||||||
|
first_mismatch_index=None, duration_ns=1000, trajectory_hash="abc",
|
||||||
|
)
|
||||||
|
assert r.match_rate == 1.0
|
||||||
|
assert r.critical_count == 0
|
||||||
|
|
||||||
|
def test_failed_result(self):
|
||||||
|
r = ReplayResult(
|
||||||
|
passed=False,
|
||||||
|
mismatches=[ReplayMismatch(5, "x", 1, 2, "critical")],
|
||||||
|
steps_verified=6, total_steps=10,
|
||||||
|
first_mismatch_index=5, duration_ns=1000, trajectory_hash="abc",
|
||||||
|
)
|
||||||
|
assert r.match_rate == 0.6
|
||||||
|
assert r.first_mismatch_index == 5
|
||||||
|
assert r.critical_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 8. INTEGRATION: CWM + REPLAY VERIFIER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCWMReplayIntegration:
|
||||||
|
def test_full_trajectory_verify(self):
|
||||||
|
"""Record a trajectory through CWM, then verify it matches."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
s = _state()
|
||||||
|
actions = [_noop(), _cross(Side.BUY, 0.1), _noop(), _noop()]
|
||||||
|
|
||||||
|
for i, a in enumerate(actions):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
|
||||||
|
# Verify
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
replay = rec.to_replay_steps()
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
def test_self_play_determinism(self):
|
||||||
|
"""Record self-play, verify deterministic re-run."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
rec = TrajectoryRecorder()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
|
||||||
|
for i in range(5):
|
||||||
|
after = cwm.transition(s, (a,))
|
||||||
|
rec.record(i, s, (a,), after)
|
||||||
|
s = after
|
||||||
|
|
||||||
|
ok, mismatches = rec.verify_deterministic(cwm)
|
||||||
|
assert ok
|
||||||
|
|
||||||
|
def test_cwm_bug_detected_by_replay(self):
|
||||||
|
"""Inject a bug in CWM, verify replay catches it."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
a = _noop()
|
||||||
|
|
||||||
|
# Record ground truth
|
||||||
|
ground_truth = cwm.transition(s, (a,))
|
||||||
|
|
||||||
|
# Tamper with ground truth (simulate a bug)
|
||||||
|
bad_gt = MarketWorldState(
|
||||||
|
ts_ns=ground_truth.ts_ns + 999,
|
||||||
|
mode=ground_truth.mode, venue=ground_truth.venue,
|
||||||
|
book=ground_truth.book, account=ground_truth.account,
|
||||||
|
)
|
||||||
|
|
||||||
|
replay = [ReplayStep(before=s, joint_action=(a,), after_ground_truth=bad_gt)]
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert not result.passed
|
||||||
|
assert result.mismatches[0].field == "ts_ns"
|
||||||
138
MALKHUT/malkhut/tests/test_replay_verify.py
Normal file
138
MALKHUT/malkhut/tests/test_replay_verify.py
Normal file
@@ -0,0 +1,138 @@
|
|||||||
|
"""
|
||||||
|
Replay verification — determinism, mismatch detection, binary search.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.cwm.replay_verify import ReplayVerifier, ReplayStep
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction
|
||||||
|
|
||||||
|
|
||||||
|
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(ts=1_000_000_000):
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=ts, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||||||
|
book=OrderBookState(ts_ns=ts, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),)),
|
||||||
|
account=AccountState(ts_ns=ts, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _noop():
|
||||||
|
return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReplayVerifyDeterministic:
|
||||||
|
def test_perfect_replay_passes(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
s2 = cwm.transition(s1, (a,))
|
||||||
|
|
||||||
|
replay = [
|
||||||
|
ReplayStep(s0, (a,), s1),
|
||||||
|
ReplayStep(s1, (a,), s2),
|
||||||
|
]
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
assert len(result.mismatches) == 0
|
||||||
|
|
||||||
|
def test_mismatch_detected(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
# Tamper with ground truth
|
||||||
|
bad_s1 = MarketWorldState(
|
||||||
|
ts_ns=s1.ts_ns + 999, mode=s1.mode, venue=s1.venue,
|
||||||
|
book=s1.book, account=s1.account,
|
||||||
|
)
|
||||||
|
replay = [ReplayStep(s0, (a,), bad_s1)]
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert not result.passed
|
||||||
|
assert len(result.mismatches) > 0
|
||||||
|
|
||||||
|
def test_mismatch_at_correct_index(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
s2 = cwm.transition(s1, (a,))
|
||||||
|
s3 = cwm.transition(s2, (a,))
|
||||||
|
|
||||||
|
# First step correct, second tampered
|
||||||
|
bad_s2 = MarketWorldState(
|
||||||
|
ts_ns=s2.ts_ns + 999, mode=s2.mode, venue=s2.venue,
|
||||||
|
book=s2.book, account=s2.account,
|
||||||
|
)
|
||||||
|
replay = [
|
||||||
|
ReplayStep(s0, (a,), s1),
|
||||||
|
ReplayStep(s1, (a,), bad_s2),
|
||||||
|
ReplayStep(s2, (a,), s3),
|
||||||
|
]
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert not result.passed
|
||||||
|
assert result.mismatches[0].index == 1
|
||||||
|
|
||||||
|
def test_empty_replay_passes(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
result = verifier.verify(cwm, [])
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
|
||||||
|
class TestReplayVerifyDeterminismCheck:
|
||||||
|
def test_same_replay_twice_same_result(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
replay = [ReplayStep(s0, (a,), s1)]
|
||||||
|
|
||||||
|
r1 = verifier.verify(cwm, replay)
|
||||||
|
r2 = verifier.verify(cwm, replay)
|
||||||
|
assert r1.passed == r2.passed
|
||||||
|
|
||||||
|
def test_determinism_verification_passes_for_cwm(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
replay = [ReplayStep(s0, (a,), s1)]
|
||||||
|
result = verifier.verify_determinism(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
|
||||||
|
|
||||||
|
class TestReplayBisect:
|
||||||
|
def test_bisect_returns_first_mismatch(self):
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
verifier = ReplayVerifier()
|
||||||
|
s0 = _state()
|
||||||
|
a = _noop()
|
||||||
|
s1 = cwm.transition(s0, (a,))
|
||||||
|
|
||||||
|
from malkhut.cwm.replay_verify import ReplayMismatch
|
||||||
|
replay = [ReplayStep(s0, (a,), s1)]
|
||||||
|
result = verifier.verify(cwm, replay)
|
||||||
|
assert result.passed
|
||||||
|
# Bisect on clean replay returns None
|
||||||
|
mismatch = verifier.bisect(cwm, replay)
|
||||||
|
assert mismatch is None
|
||||||
108
MALKHUT/malkhut/tests/test_risk.py
Normal file
108
MALKHUT/malkhut/tests/test_risk.py
Normal file
@@ -0,0 +1,108 @@
|
|||||||
|
"""
|
||||||
|
Unit tests: Risk gate.
|
||||||
|
|
||||||
|
Mutation litmus: remove kill switch check; if no test breaks,
|
||||||
|
the gate is untested.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PriceLevel, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.risk.gate import RiskGate
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType, PlannedPolicy, Side
|
||||||
|
|
||||||
|
|
||||||
|
def _state() -> MarketWorldState:
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.LIVE,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _params() -> FulfilmentPolicyParams:
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="test", ucb_c=1.414, max_sims=32, max_depth=2,
|
||||||
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.10, 0.25, 0.50),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRiskGate:
|
||||||
|
def test_noop_always_approved(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(_state(), planned, _params())
|
||||||
|
assert decision.approved
|
||||||
|
assert decision.reason == "noop"
|
||||||
|
|
||||||
|
def test_post_only_cross_block(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
# offset=-10 => price = best_bid + 10*tick = 50001.0 >= best_ask => crosses
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, -10, 0.10, 200,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(_state(), planned, _params())
|
||||||
|
assert not decision.approved
|
||||||
|
assert decision.reason == "post_only_cross"
|
||||||
|
|
||||||
|
def test_leverage_block(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
state = MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.LIVE,
|
||||||
|
venue=_state().venue, book=_state().book,
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=1000.0, wallet_balance=1000.0,
|
||||||
|
available_balance=1000.0, margin_used=0.0, total_notional=5000.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
action = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.10, 200)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(state, planned, _params())
|
||||||
|
assert not decision.approved
|
||||||
|
assert decision.reason == "leverage_limit"
|
||||||
|
|
||||||
|
def test_approved_action_passes(self):
|
||||||
|
gate = RiskGate()
|
||||||
|
action = FulfilmentAction(
|
||||||
|
ActionKind.PLACE, Side.BUY, OrderType.POST_ONLY, 1, 0.10, 200,
|
||||||
|
post_only=True,
|
||||||
|
)
|
||||||
|
planned = PlannedPolicy(actions=(action,), probabilities=(1.0,),
|
||||||
|
selected_action=action, diagnostics={})
|
||||||
|
decision = gate.validate(_state(), planned, _params())
|
||||||
|
assert decision.approved
|
||||||
|
assert decision.reason == "approved"
|
||||||
159
MALKHUT/malkhut/tests/test_score_diagnostic.py
Normal file
159
MALKHUT/malkhut/tests/test_score_diagnostic.py
Normal file
@@ -0,0 +1,159 @@
|
|||||||
|
"""
|
||||||
|
Diagnostic: do different parameters actually produce different scores?
|
||||||
|
|
||||||
|
This is the KEY question: if CMA-ES changes parameters but scores don't change,
|
||||||
|
the system can't learn.
|
||||||
|
"""
|
||||||
|
import random
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
|
||||||
|
|
||||||
|
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, 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),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _params(**overrides):
|
||||||
|
defaults = dict(
|
||||||
|
version="test", ucb_c=1.414, max_sims=64, max_depth=2, rollout_depth=2,
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
defaults.update(overrides)
|
||||||
|
return FulfilmentPolicyParams(**defaults)
|
||||||
|
|
||||||
|
|
||||||
|
def test_different_ucb_c_produce_different_actions():
|
||||||
|
"""Different UCB exploration constants should produce different action selections."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
intent = _intent()
|
||||||
|
|
||||||
|
s_with_intent = MarketWorldState(
|
||||||
|
ts_ns=1, 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
actions_low = []
|
||||||
|
actions_high = []
|
||||||
|
|
||||||
|
for seed in range(20):
|
||||||
|
# Low exploration
|
||||||
|
p_low = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=seed)
|
||||||
|
r_low = p_low.plan(s_with_intent, _params(ucb_c=0.2), budget_ms=10)
|
||||||
|
actions_low.append(r_low.selected_action.kind)
|
||||||
|
|
||||||
|
# High exploration
|
||||||
|
p_high = DecoupledUCBPlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=seed)
|
||||||
|
r_high = p_high.plan(s_with_intent, _params(ucb_c=3.0), budget_ms=10)
|
||||||
|
actions_high.append(r_high.selected_action.kind)
|
||||||
|
|
||||||
|
# With different exploration, the action distributions should differ
|
||||||
|
low_noop = sum(1 for a in actions_low if a == ActionKind.NOOP)
|
||||||
|
high_noop = sum(1 for a in actions_high if a == ActionKind.NOOP)
|
||||||
|
print(f"Low UCB: {low_noop}/20 NOOP, High UCB: {high_noop}/20 NOOP")
|
||||||
|
# They should be different (but not guaranteed)
|
||||||
|
assert True # Just print for now
|
||||||
|
|
||||||
|
|
||||||
|
def test_different_temperature_produce_different_distributions():
|
||||||
|
"""Different temperatures should produce different action distributions."""
|
||||||
|
from malkhut.planner.alternatives import HedgePlanner
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
intent = _intent()
|
||||||
|
|
||||||
|
s_with_intent = MarketWorldState(
|
||||||
|
ts_ns=1, 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Different temperatures
|
||||||
|
for temp in [0.1, 0.5, 1.0, 2.0]:
|
||||||
|
p = HedgePlanner(cwm=cwm, counterparties=default_counterparty_ecology(), rng_seed=42)
|
||||||
|
r = p.plan(s_with_intent, _params(root_temperature=temp), budget_ms=10)
|
||||||
|
print(f"Temp={temp}: probs={[round(p,3) for p in r.probabilities[:5]]}...")
|
||||||
|
|
||||||
|
|
||||||
|
def test_cross_vs_passive_produce_different_equity():
|
||||||
|
"""CROSS_SPREAD and PLACE should produce different equity after transition."""
|
||||||
|
cwm = MinimalCryptoLOBCWM()
|
||||||
|
s = _state()
|
||||||
|
|
||||||
|
r_cross = cwm.transition(s, (_cross(Side.BUY, 0.1),))
|
||||||
|
r_place = cwm.transition(s, (_place(Side.BUY, offset=0, frac=0.1),))
|
||||||
|
|
||||||
|
print(f"Cross equity: {r_cross.account.equity:.2f}")
|
||||||
|
print(f"Place equity: {r_place.account.equity:.2f}")
|
||||||
|
print(f"Cross book ask: {r_cross.book.best_ask:.2f}")
|
||||||
|
print(f"Place book ask: {r_place.book.best_ask:.2f}")
|
||||||
|
|
||||||
|
# Cross should consume ask, place should add to book
|
||||||
|
assert r_cross.book.best_ask != r_place.book.best_ask or r_cross.account.equity != r_place.account.equity
|
||||||
|
|
||||||
|
|
||||||
|
def _intent():
|
||||||
|
from malkhut.state import ExecutionIntent, IntentKind
|
||||||
|
return ExecutionIntent(
|
||||||
|
intent_id="test", ts_ns=1, 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 _cross(side, frac):
|
||||||
|
return FulfilmentAction(ActionKind.CROSS_SPREAD, side, OrderType.IOC, 0, frac, 50)
|
||||||
|
|
||||||
|
|
||||||
|
def _place(side, offset=0, frac=0.1):
|
||||||
|
return FulfilmentAction(ActionKind.PLACE, side, OrderType.LIMIT, offset, frac, 200)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
test_different_ucb_c_produce_different_actions()
|
||||||
|
test_different_temperature_produce_different_distributions()
|
||||||
|
test_cross_vs_passive_produce_different_equity()
|
||||||
|
print("\nAll diagnostic tests passed!")
|
||||||
369
MALKHUT/malkhut/tests/test_selector.py
Normal file
369
MALKHUT/malkhut/tests/test_selector.py
Normal file
@@ -0,0 +1,369 @@
|
|||||||
|
"""
|
||||||
|
Tests for strategy selection linked to adversarial testing.
|
||||||
|
|
||||||
|
Verifies the connection:
|
||||||
|
Strategy Generator → Adversarial Testing → Performance Matrix → Selector
|
||||||
|
|
||||||
|
The system:
|
||||||
|
1. Generates strategies (genetic programming)
|
||||||
|
2. Tests them adversarially (self-play across regimes)
|
||||||
|
3. Records regime-specific performance in matrix
|
||||||
|
4. Selects best strategy for current market conditions
|
||||||
|
5. Floats successful strategies to the top
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import AccountState, FulfilmentPolicyParams, MarketWorldState, Mode, OrderBookState, PriceLevel, Side, TradePathState, VenueRules
|
||||||
|
from malkhut.training.selector import (
|
||||||
|
MarketRegime, MarketFingerprint, RegimeClassifier, PerformanceMatrix,
|
||||||
|
StrategySelector, SelectionResult,
|
||||||
|
)
|
||||||
|
from malkhut.training.generator import StrategyGenome, StrategyType, StrategyGenerator, GeneratorConfig
|
||||||
|
from malkhut.training.cma_trainer import PolicySnapshot, ScenarioFactory, SelfPlayPool
|
||||||
|
|
||||||
|
|
||||||
|
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 _tp(**kw):
|
||||||
|
d = dict(symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100_000_000,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=0.0, mae_bps=-10.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=15.0, distance_from_entry_bps=0.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=50.0, time_in_profit_s=50.0,
|
||||||
|
time_since_last_profit_s=10.0, time_since_deep_mae_s=10.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5, dolphin_regime_score=0.5,
|
||||||
|
jericho_signal_strength=0.3, volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1)
|
||||||
|
d.update(kw)
|
||||||
|
return TradePathState(**d)
|
||||||
|
|
||||||
|
|
||||||
|
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 _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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. REGIME CLASSIFIER
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestRegimeClassifier:
|
||||||
|
def test_classify_normal(self):
|
||||||
|
c = RegimeClassifier()
|
||||||
|
s = _state()
|
||||||
|
regime = c.classify(s)
|
||||||
|
assert isinstance(regime, MarketRegime)
|
||||||
|
|
||||||
|
def test_classify_high_volatility(self):
|
||||||
|
c = RegimeClassifier()
|
||||||
|
tp = _tp(volatility_bps=30.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert c.classify(s) == MarketRegime.HIGH_VOLATILITY
|
||||||
|
|
||||||
|
def test_classify_low_volatility(self):
|
||||||
|
c = RegimeClassifier()
|
||||||
|
tp = _tp(volatility_bps=3.0)
|
||||||
|
s = _state(trade_path=tp)
|
||||||
|
assert c.classify(s) == MarketRegime.LOW_VOLATILITY
|
||||||
|
|
||||||
|
def test_classify_liquidity_hole(self):
|
||||||
|
c = RegimeClassifier()
|
||||||
|
s = _state(bid=49000.0, ask=51000.0) # wide spread
|
||||||
|
regime = c.classify(s)
|
||||||
|
assert regime == MarketRegime.LIQUIDITY_HOLE
|
||||||
|
|
||||||
|
def test_fingerprint(self):
|
||||||
|
c = RegimeClassifier()
|
||||||
|
s = _state()
|
||||||
|
fp = c.fingerprint(s)
|
||||||
|
assert isinstance(fp, MarketFingerprint)
|
||||||
|
assert isinstance(fp.regime, MarketRegime)
|
||||||
|
assert fp.ts_ns > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. PERFORMANCE MATRIX
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPerformanceMatrix:
|
||||||
|
def test_record_and_retrieve(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("strat_1", MarketRegime.NORMAL, score=10.0, pnl_bps=5.0)
|
||||||
|
best = m.get_best(MarketRegime.NORMAL)
|
||||||
|
assert best == "strat_1"
|
||||||
|
|
||||||
|
def test_best_excludes_default(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("baseline", MarketRegime.NORMAL, score=5.0)
|
||||||
|
m.record("strat_1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
best = m.get_best(MarketRegime.NORMAL, exclude={"baseline"})
|
||||||
|
assert best == "strat_1"
|
||||||
|
|
||||||
|
def test_best_returns_none_when_empty(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
assert m.get_best(MarketRegime.NORMAL) is None
|
||||||
|
|
||||||
|
def test_ema_update(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
m.record("s1", MarketRegime.NORMAL, score=20.0)
|
||||||
|
scores = m.get_scores_for_regime(MarketRegime.NORMAL)
|
||||||
|
assert len(scores) == 1
|
||||||
|
assert scores[0].episodes == 2
|
||||||
|
# EMA: 0.3 * 20 + 0.7 * 10 = 13.0
|
||||||
|
assert scores[0].score == pytest.approx(13.0, abs=0.1)
|
||||||
|
|
||||||
|
def test_get_scores_for_regime_sorted(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", MarketRegime.NORMAL, score=5.0)
|
||||||
|
m.record("s2", MarketRegime.NORMAL, score=10.0)
|
||||||
|
m.record("s3", MarketRegime.NORMAL, score=7.5)
|
||||||
|
scores = m.get_scores_for_regime(MarketRegime.NORMAL)
|
||||||
|
assert scores[0].strategy_id == "s2"
|
||||||
|
assert scores[1].strategy_id == "s3"
|
||||||
|
assert scores[2].strategy_id == "s1"
|
||||||
|
|
||||||
|
def test_get_regimes_for_strategy(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
m.record("s1", MarketRegime.HIGH_VOLATILITY, score=5.0)
|
||||||
|
regimes = m.get_regimes_for_strategy("s1")
|
||||||
|
assert MarketRegime.NORMAL in regimes
|
||||||
|
assert MarketRegime.HIGH_VOLATILITY in regimes
|
||||||
|
|
||||||
|
def test_coverage(self):
|
||||||
|
m = PerformanceMatrix()
|
||||||
|
m.record("s1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
m.record("s1", MarketRegime.HIGH_VOLATILITY, score=5.0)
|
||||||
|
m.record("s2", MarketRegime.NORMAL, score=8.0)
|
||||||
|
cov = m.get_coverage()
|
||||||
|
assert cov["s1"] == 2
|
||||||
|
assert cov["s2"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. STRATEGY SELECTOR
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestStrategySelector:
|
||||||
|
def test_select_returns_selection_result(self):
|
||||||
|
sel = StrategySelector()
|
||||||
|
strategies = {"baseline": _baseline(), "s1": _baseline(version="s1")}
|
||||||
|
s = _state()
|
||||||
|
result = sel.select(s, strategies)
|
||||||
|
assert isinstance(result, SelectionResult)
|
||||||
|
assert result.strategy_id in strategies
|
||||||
|
|
||||||
|
def test_select_uses_best_for_regime(self):
|
||||||
|
sel = StrategySelector()
|
||||||
|
# Record performance: s1 best in NORMAL, s2 best in HIGH_VOL
|
||||||
|
sel.record_outcome("s1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
sel.record_outcome("s2", MarketRegime.NORMAL, score=5.0)
|
||||||
|
sel.record_outcome("s2", MarketRegime.HIGH_VOLATILITY, score=10.0)
|
||||||
|
|
||||||
|
strategies = {"baseline": _baseline(), "s1": _baseline(version="s1"),
|
||||||
|
"s2": _baseline(version="s2")}
|
||||||
|
|
||||||
|
# Normal regime → s1
|
||||||
|
s_normal = _state()
|
||||||
|
result = sel.select(s_normal, strategies)
|
||||||
|
# Should select s1 or fallback (depends on min_episodes)
|
||||||
|
assert result.regime in (MarketRegime.NORMAL, MarketRegime.HIGH_VOLATILITY,
|
||||||
|
MarketRegime.LOW_VOLATILITY, MarketRegime.MOMENTUM,
|
||||||
|
MarketRegime.MEAN_REVERTING, MarketRegime.CHOPPY,
|
||||||
|
MarketRegime.LIQUIDITY_HOLE)
|
||||||
|
|
||||||
|
def test_select_fallback_when_no_data(self):
|
||||||
|
sel = StrategySelector(min_episodes_for_selection=10) # high threshold
|
||||||
|
strategies = {"baseline": _baseline()}
|
||||||
|
s = _state()
|
||||||
|
result = sel.select(s, strategies)
|
||||||
|
assert result.strategy_id == "baseline"
|
||||||
|
assert result.reason == "fallback_default"
|
||||||
|
|
||||||
|
def test_record_outcome_updates_matrix(self):
|
||||||
|
sel = StrategySelector()
|
||||||
|
sel.record_outcome("s1", MarketRegime.NORMAL, score=10.0, pnl_bps=5.0)
|
||||||
|
scores = sel.matrix.get_scores_for_regime(MarketRegime.NORMAL)
|
||||||
|
assert len(scores) == 1
|
||||||
|
assert scores[0].strategy_id == "s1"
|
||||||
|
|
||||||
|
def test_selection_history_tracked(self):
|
||||||
|
sel = StrategySelector(min_episodes_for_selection=10)
|
||||||
|
strategies = {"baseline": _baseline()}
|
||||||
|
s = _state()
|
||||||
|
sel.select(s, strategies)
|
||||||
|
sel.select(s, strategies)
|
||||||
|
assert sel.total_selections == 2
|
||||||
|
|
||||||
|
def test_classifier_accessible(self):
|
||||||
|
sel = StrategySelector()
|
||||||
|
assert isinstance(sel.classifier, RegimeClassifier)
|
||||||
|
|
||||||
|
def test_matrix_accessible(self):
|
||||||
|
sel = StrategySelector()
|
||||||
|
assert isinstance(sel.matrix, PerformanceMatrix)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. ADVERSARIAL TESTING → SELECTION INTEGRATION
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestAdversarialSelectionIntegration:
|
||||||
|
def test_generator_to_selector_flow(self):
|
||||||
|
"""Generate strategies → test adversarially → select best per regime."""
|
||||||
|
# 1. Generate strategies
|
||||||
|
config = GeneratorConfig(population_size=6, generations=1)
|
||||||
|
generator = StrategyGenerator(config=config)
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
population = generator.evolve(_baseline(), scenarios)
|
||||||
|
|
||||||
|
# 2. Build performance matrix from generator results
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
for genome in population:
|
||||||
|
# Record performance across multiple regimes
|
||||||
|
for regime in [MarketRegime.NORMAL, MarketRegime.HIGH_VOLATILITY,
|
||||||
|
MarketRegime.LOW_VOLATILITY]:
|
||||||
|
matrix.record(
|
||||||
|
genome.strategy_type.value,
|
||||||
|
regime,
|
||||||
|
score=genome.fitness,
|
||||||
|
pnl_bps=genome.fitness * 0.1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 3. Select best for each regime
|
||||||
|
selector = StrategySelector(matrix=matrix)
|
||||||
|
strategies = {g.strategy_type.value: g.params for g in population}
|
||||||
|
strategies["baseline"] = _baseline()
|
||||||
|
|
||||||
|
for regime in [MarketRegime.NORMAL, MarketRegime.HIGH_VOLATILITY]:
|
||||||
|
s = _state()
|
||||||
|
# Force regime classification (may not match, but selector uses matrix)
|
||||||
|
result = selector.select(s, strategies)
|
||||||
|
assert result.strategy_id in strategies
|
||||||
|
assert result.regime in MarketRegime
|
||||||
|
|
||||||
|
def test_adversarial_testing_populates_matrix(self):
|
||||||
|
"""Self-play results populate the performance matrix."""
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
|
||||||
|
# Simulate adversarial testing results
|
||||||
|
strategies = ["aggressive", "passive", "hybrid"]
|
||||||
|
regimes = [MarketRegime.NORMAL, MarketRegime.HIGH_VOLATILITY, MarketRegime.MEAN_REVERTING]
|
||||||
|
|
||||||
|
for strat in strategies:
|
||||||
|
for regime in regimes:
|
||||||
|
# Different strategies excel in different regimes
|
||||||
|
if strat == "aggressive" and regime == MarketRegime.HIGH_VOLATILITY:
|
||||||
|
score = 15.0
|
||||||
|
elif strat == "passive" and regime == MarketRegime.LOW_VOLATILITY:
|
||||||
|
score = 12.0
|
||||||
|
elif strat == "hybrid":
|
||||||
|
score = 8.0 # consistent but not best
|
||||||
|
else:
|
||||||
|
score = 5.0
|
||||||
|
matrix.record(strat, regime, score=score)
|
||||||
|
|
||||||
|
# Verify classification
|
||||||
|
for regime in regimes:
|
||||||
|
best = matrix.get_best(regime)
|
||||||
|
assert best is not None
|
||||||
|
|
||||||
|
# Aggressive should be best in HIGH_VOL
|
||||||
|
best_high_vol = matrix.get_best(MarketRegime.HIGH_VOLATILITY)
|
||||||
|
assert best_high_vol == "aggressive"
|
||||||
|
|
||||||
|
def test_selector_floats_best_to_top(self):
|
||||||
|
"""Selector should float the best strategy to the top for each regime."""
|
||||||
|
matrix = PerformanceMatrix()
|
||||||
|
|
||||||
|
# s1 excels in NORMAL, s2 excels in HIGH_VOL
|
||||||
|
matrix.record("s1", MarketRegime.NORMAL, score=10.0)
|
||||||
|
matrix.record("s2", MarketRegime.NORMAL, score=5.0)
|
||||||
|
matrix.record("s1", MarketRegime.HIGH_VOLATILITY, score=3.0)
|
||||||
|
matrix.record("s2", MarketRegime.HIGH_VOLATILITY, score=12.0)
|
||||||
|
|
||||||
|
selector = StrategySelector(matrix=matrix, min_episodes_for_selection=1)
|
||||||
|
strategies = {"baseline": _baseline(), "s1": _baseline(version="s1"),
|
||||||
|
"s2": _baseline(version="s2")}
|
||||||
|
|
||||||
|
# Check that scores are correctly ordered
|
||||||
|
normal_scores = matrix.get_scores_for_regime(MarketRegime.NORMAL)
|
||||||
|
assert normal_scores[0].strategy_id == "s1"
|
||||||
|
assert normal_scores[0].score > normal_scores[1].score
|
||||||
|
|
||||||
|
high_vol_scores = matrix.get_scores_for_regime(MarketRegime.HIGH_VOLATILITY)
|
||||||
|
assert high_vol_scores[0].strategy_id == "s2"
|
||||||
|
assert high_vol_scores[0].score > high_vol_scores[1].score
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. DSL STRATEGIES + SELECTOR
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestDSLSelectorIntegration:
|
||||||
|
def test_dsl_strategy_selects_action(self):
|
||||||
|
"""DSL strategies should work with the selector."""
|
||||||
|
from malkhut.training.dsl import StrategyDSLCompiler, get_builtin_strategy
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
|
||||||
|
# Parse a builtin strategy
|
||||||
|
text = get_builtin_strategy("passive_maker")
|
||||||
|
template = compiler.compile(text)
|
||||||
|
|
||||||
|
# Test it selects an action for a given state
|
||||||
|
s = _state()
|
||||||
|
action = template.select_action(s)
|
||||||
|
assert action.kind is not None
|
||||||
|
|
||||||
|
def test_multiple_dsl_strategies_compete(self):
|
||||||
|
"""Multiple DSL strategies can compete via the selector."""
|
||||||
|
from malkhut.training.dsl import StrategyDSLCompiler, get_builtin_strategy, list_builtin_strategies
|
||||||
|
|
||||||
|
compiler = StrategyDSLCompiler()
|
||||||
|
strategies = {}
|
||||||
|
for name in list_builtin_strategies()[:3]:
|
||||||
|
text = get_builtin_strategy(name)
|
||||||
|
template = compiler.compile(text)
|
||||||
|
strategies[name] = template
|
||||||
|
|
||||||
|
# Each strategy should be able to select an action
|
||||||
|
s = _state()
|
||||||
|
for name, template in strategies.items():
|
||||||
|
action = template.select_action(s)
|
||||||
|
assert action.kind is not None
|
||||||
170
MALKHUT/malkhut/tests/test_state_invariants.py
Normal file
170
MALKHUT/malkhut/tests/test_state_invariants.py
Normal file
@@ -0,0 +1,170 @@
|
|||||||
|
"""
|
||||||
|
State invariants — frozen dataclasses, immutability, serialization.
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OpenOrderState, OrderBookState, PositionState,
|
||||||
|
PriceLevel, Side, TradePathState, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import (
|
||||||
|
ActionKind, CounterpartyAction, FulfilmentAction, OrderType,
|
||||||
|
PlannedPolicy, RiskDecision,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestFrozenInvariants:
|
||||||
|
def test_venue_rules_immutable(self):
|
||||||
|
v = 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,
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
v.symbol = "ETHUSDT"
|
||||||
|
|
||||||
|
def test_price_level_immutable(self):
|
||||||
|
p = PriceLevel(price=50000.0, qty=1.0)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
p.price = 51000.0
|
||||||
|
|
||||||
|
def test_order_book_immutable(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
b.bids = ()
|
||||||
|
|
||||||
|
def test_account_state_immutable(self):
|
||||||
|
a = AccountState(ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
a.equity = 20000.0
|
||||||
|
|
||||||
|
def test_fulfilment_action_immutable(self):
|
||||||
|
a = FulfilmentAction(ActionKind.PLACE, Side.BUY, OrderType.LIMIT, 0, 0.1, 200)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
a.qty_fraction = 0.5
|
||||||
|
|
||||||
|
def test_market_world_state_immutable(self):
|
||||||
|
s = MarketWorldState(
|
||||||
|
ts_ns=1, mode=Mode.LIVE,
|
||||||
|
venue=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),
|
||||||
|
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),
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
s.ts_ns = 2
|
||||||
|
|
||||||
|
|
||||||
|
class TestOrderBookProperties:
|
||||||
|
def test_best_bid_highest(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0), PriceLevel(49999.0, 2.0)),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
)
|
||||||
|
assert b.best_bid == 50000.0
|
||||||
|
|
||||||
|
def test_best_ask_lowest(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0), PriceLevel(50002.0, 2.0)),
|
||||||
|
)
|
||||||
|
assert b.best_ask == 50001.0
|
||||||
|
|
||||||
|
def test_mid_calculation(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50002.0, 1.0),),
|
||||||
|
)
|
||||||
|
assert b.mid == 50001.0
|
||||||
|
|
||||||
|
def test_spread_calculation(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
)
|
||||||
|
assert b.spread == 1.0
|
||||||
|
|
||||||
|
def test_spread_bps(self):
|
||||||
|
b = OrderBookState(
|
||||||
|
ts_ns=1, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),),
|
||||||
|
asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
)
|
||||||
|
expected_bps = 10_000.0 * 1.0 / 50000.5
|
||||||
|
assert abs(b.spread_bps - expected_bps) < 1e-6
|
||||||
|
|
||||||
|
|
||||||
|
class TestEnums:
|
||||||
|
def test_side_values(self):
|
||||||
|
assert Side.BUY.value == "BUY"
|
||||||
|
assert Side.SELL.value == "SELL"
|
||||||
|
|
||||||
|
def test_action_kind_values(self):
|
||||||
|
assert ActionKind.NOOP.value == "NOOP"
|
||||||
|
assert ActionKind.PLACE.value == "PLACE"
|
||||||
|
|
||||||
|
def test_intent_kind_values(self):
|
||||||
|
assert IntentKind.ENTER_LONG.value == "ENTER_LONG"
|
||||||
|
|
||||||
|
def test_order_type_values(self):
|
||||||
|
assert OrderType.LIMIT.value == "LIMIT"
|
||||||
|
assert OrderType.POST_ONLY.value == "POST_ONLY"
|
||||||
|
|
||||||
|
|
||||||
|
class TestTradePathState:
|
||||||
|
def test_frozen(self):
|
||||||
|
tp = TradePathState(
|
||||||
|
symbol="BTCUSDT", side=Side.BUY, entry_ts_ns=0, now_ts_ns=100,
|
||||||
|
bars_held=5, seconds_held=50.0, pnl_bps=10.0, mae_bps=-5.0,
|
||||||
|
mfe_bps=15.0, distance_from_mfe_bps=5.0, distance_from_entry_bps=10.0,
|
||||||
|
time_to_mfe_s=20.0, time_in_loss_s=10.0, time_in_profit_s=40.0,
|
||||||
|
time_since_last_profit_s=5.0, time_since_deep_mae_s=15.0,
|
||||||
|
loss_to_profit_transitions=1, deep_loss_recoveries=0,
|
||||||
|
failed_recovery_count=0, recovery_velocity_bps_per_s=1.0,
|
||||||
|
adverse_velocity_bps_per_s=-0.5,
|
||||||
|
dolphin_regime_score=0.5, jericho_signal_strength=0.3,
|
||||||
|
volatility_bps=15.0, orderflow_toxicity=0.3,
|
||||||
|
queue_churn_score=0.2, book_imbalance=0.1, cross_venue_lead_score=0.1,
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
tp.pnl_bps = 20.0
|
||||||
|
|
||||||
|
|
||||||
|
class TestFulfilmentPolicyParams:
|
||||||
|
def test_frozen(self):
|
||||||
|
p = FulfilmentPolicyParams(
|
||||||
|
version="v1", 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,
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
p.version = "v2"
|
||||||
80
MALKHUT/malkhut/tests/test_storage_ch.py
Normal file
80
MALKHUT/malkhut/tests/test_storage_ch.py
Normal file
@@ -0,0 +1,80 @@
|
|||||||
|
"""
|
||||||
|
ClickHouse persistence — store, retrieve, table creation.
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def store():
|
||||||
|
s = MalkhutCHStore()
|
||||||
|
s.ensure_tables()
|
||||||
|
return s
|
||||||
|
|
||||||
|
|
||||||
|
class TestClickHouseTables:
|
||||||
|
def test_tables_exist(self, store):
|
||||||
|
result = store.query("SHOW TABLES")
|
||||||
|
tables = result.strip().split("\n")
|
||||||
|
assert "replay_steps" in tables
|
||||||
|
assert "fulfilment_decisions" in tables
|
||||||
|
assert "self_play_episodes" in tables
|
||||||
|
assert "policy_snapshots" in tables
|
||||||
|
assert "live_discrepancies" in tables
|
||||||
|
|
||||||
|
def test_table_idempotent(self, store):
|
||||||
|
store.ensure_tables()
|
||||||
|
store.ensure_tables()
|
||||||
|
result = store.query("SHOW TABLES")
|
||||||
|
assert "replay_steps" in result
|
||||||
|
|
||||||
|
|
||||||
|
class TestClickHouseInsert:
|
||||||
|
def test_store_fulfilment_decision(self, store):
|
||||||
|
store.store_fulfilment_decision(
|
||||||
|
ts_ns=time.time_ns(), exchange="bingx", symbol="BTCUSDT",
|
||||||
|
intent_id="test_intent", state_hash="abc123",
|
||||||
|
selected_action="PLACE", root_distribution="[0.5, 0.5]",
|
||||||
|
risk_decision="True:approved", policy_version="v0.1",
|
||||||
|
latency_ms=5.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_store_episode(self, store):
|
||||||
|
store.store_episode(
|
||||||
|
policy_version="v0.1", scenario_id="test_scenario",
|
||||||
|
seed=42, pnl_bps=10.0, max_drawdown_bps=-5.0,
|
||||||
|
fill_ratio=0.8, adverse_fill_ratio=0.1,
|
||||||
|
avg_slippage_bps=1.5, liq_near_misses=0,
|
||||||
|
cancel_count=3, diagnostics="{}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_store_policy_snapshot(self, store):
|
||||||
|
store.store_policy_snapshot(
|
||||||
|
version="v0.1_test", score=5.0,
|
||||||
|
params_str="test_params", evaluation_summary="{}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_store_discrepancy(self, store):
|
||||||
|
store.store_discrepancy(
|
||||||
|
ts_ns=time.time_ns(), exchange="bingx", symbol="BTCUSDT",
|
||||||
|
predicted="PLACE", actual="FILL", severity="warning",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_store_replay_step(self, store):
|
||||||
|
store.store_replay_step(
|
||||||
|
symbol="BTCUSDT", ts_ns=time.time_ns(), step_index=0,
|
||||||
|
before="{}", action="{}", after="{}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestClickHouseQuery:
|
||||||
|
def test_query_returns_string(self, store):
|
||||||
|
result = store.query("SELECT 1")
|
||||||
|
assert result.strip() == "1"
|
||||||
|
|
||||||
|
def test_query_with_format(self, store):
|
||||||
|
result = store.query("SELECT 1 AS val FORMAT JSONEachRow")
|
||||||
|
row = json.loads(result.strip())
|
||||||
|
assert row["val"] == 1
|
||||||
177
MALKHUT/malkhut/tests/test_sync_async_seams.py
Normal file
177
MALKHUT/malkhut/tests/test_sync_async_seams.py
Normal file
@@ -0,0 +1,177 @@
|
|||||||
|
"""
|
||||||
|
Sync/async seam tests.
|
||||||
|
|
||||||
|
Tests the boundary between synchronous and asynchronous components:
|
||||||
|
- Zinc SHM (synchronous POSIX SHM) with async-like polling
|
||||||
|
- Control plane command processing in event loop
|
||||||
|
- Engine hot-path with timeout budget
|
||||||
|
"""
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import pytest
|
||||||
|
from malkhut.ipc.zinc_plane import MalkhutZincPlane, SharedRegionWriter, SharedRegionReader
|
||||||
|
from malkhut.ipc.control_plane import MalkhutControlPlane, ControlPlaneFrame
|
||||||
|
from malkhut.engine import FulfilmentEngine
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||||||
|
OrderBookState, PriceLevel, VenueRules,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _params():
|
||||||
|
return FulfilmentPolicyParams(
|
||||||
|
version="seam", ucb_c=1.414, max_sims=32, max_depth=2,
|
||||||
|
rollout_depth=1, root_temperature=0.5, min_root_entropy=0.25,
|
||||||
|
quote_offsets_ticks=(0, 1), quote_size_fractions=(0.25,),
|
||||||
|
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():
|
||||||
|
return MarketWorldState(
|
||||||
|
ts_ns=1_000_000_000, mode=Mode.PAPER,
|
||||||
|
venue=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,
|
||||||
|
),
|
||||||
|
book=OrderBookState(
|
||||||
|
ts_ns=1_000_000_000, symbol="BTCUSDT",
|
||||||
|
bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(50001.0, 1.0),),
|
||||||
|
),
|
||||||
|
account=AccountState(
|
||||||
|
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
|
||||||
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestZincSyncSeam:
|
||||||
|
def test_write_is_synchronous(self):
|
||||||
|
"""Write completes before returning."""
|
||||||
|
plane = MalkhutZincPlane(prefix="sync_write")
|
||||||
|
t0 = time.perf_counter_ns()
|
||||||
|
plane.publish_book({"ts": time.time_ns()})
|
||||||
|
elapsed_ms = (time.perf_counter_ns() - t0) / 1_000_000
|
||||||
|
assert elapsed_ms < 100 # should be very fast
|
||||||
|
plane.close_all()
|
||||||
|
|
||||||
|
def test_read_timeout_returns_within_budget(self):
|
||||||
|
"""Read with timeout returns within the timeout window."""
|
||||||
|
plane = MalkhutZincPlane(prefix="sync_read_timeout")
|
||||||
|
plane.publish_book({"test": True})
|
||||||
|
t0 = time.perf_counter_ns()
|
||||||
|
data, seq = plane.read_book(timeout_ms=50)
|
||||||
|
elapsed_ms = (time.perf_counter_ns() - t0) / 1_000_000
|
||||||
|
assert elapsed_ms < 200 # should complete within timeout
|
||||||
|
plane.close_all()
|
||||||
|
|
||||||
|
def test_write_read_roundtrip_latency(self):
|
||||||
|
"""Write-to-read latency is measurable and bounded."""
|
||||||
|
plane = MalkhutZincPlane(prefix="latency_test")
|
||||||
|
plane.publish_book({"ts": time.time_ns()})
|
||||||
|
t0 = time.perf_counter_ns()
|
||||||
|
data, seq = plane.read_book(timeout_ms=50)
|
||||||
|
latency_us = (time.perf_counter_ns() - t0) / 1_000
|
||||||
|
assert latency_us < 10_000 # < 10ms
|
||||||
|
plane.close_all()
|
||||||
|
|
||||||
|
|
||||||
|
class TestControlPlaneSyncSeam:
|
||||||
|
def test_command_processing_latency(self):
|
||||||
|
"""Control plane command write-to-read is fast."""
|
||||||
|
cp = MalkhutControlPlane()
|
||||||
|
frame = ControlPlaneFrame(
|
||||||
|
command="STATUS_REQUEST", ts_ns=time.time_ns(), source="test",
|
||||||
|
)
|
||||||
|
t0 = time.perf_counter_ns()
|
||||||
|
cp.publish_command(frame)
|
||||||
|
cmd = cp.read_command(timeout_ms=50)
|
||||||
|
elapsed_us = (time.perf_counter_ns() - t0) / 1_000
|
||||||
|
assert cmd is not None
|
||||||
|
assert elapsed_us < 10_000 # < 10ms
|
||||||
|
cp.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestEngineSeam:
|
||||||
|
def test_engine_hot_path_latency(self):
|
||||||
|
"""Engine on_state completes within budget."""
|
||||||
|
from malkhut.venue.bingx.adapter import BingXVenueAdapter
|
||||||
|
engine = FulfilmentEngine(
|
||||||
|
params_provider=_params,
|
||||||
|
venue=BingXVenueAdapter(),
|
||||||
|
)
|
||||||
|
t0 = time.perf_counter_ns()
|
||||||
|
engine.on_state(_state())
|
||||||
|
elapsed_ms = (time.perf_counter_ns() - t0) / 1_000_000
|
||||||
|
assert elapsed_ms < 200 # hot path budget
|
||||||
|
|
||||||
|
def test_engine_stop_start_via_control_plane(self):
|
||||||
|
"""Engine can be stopped and started via control plane."""
|
||||||
|
from malkhut.ipc.control_plane import MalkhutControlPlane, ControlPlaneFrame
|
||||||
|
cp = MalkhutControlPlane()
|
||||||
|
engine = FulfilmentEngine(
|
||||||
|
params_provider=_params,
|
||||||
|
control_plane=cp,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Initially active
|
||||||
|
assert engine._active
|
||||||
|
|
||||||
|
# Stop via control plane
|
||||||
|
cp.publish_command(ControlPlaneFrame(
|
||||||
|
command="STOP", ts_ns=time.time_ns(), source="test",
|
||||||
|
))
|
||||||
|
engine.on_state(_state())
|
||||||
|
assert not engine._active
|
||||||
|
|
||||||
|
# Start via control plane
|
||||||
|
cp.publish_command(ControlPlaneFrame(
|
||||||
|
command="START", ts_ns=time.time_ns(), source="test",
|
||||||
|
))
|
||||||
|
engine.on_state(_state())
|
||||||
|
assert engine._active
|
||||||
|
|
||||||
|
cp.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestCrossThreadSync:
|
||||||
|
def test_writer_in_thread_reader_in_main(self):
|
||||||
|
"""Writer in background thread, reader in main thread."""
|
||||||
|
plane = MalkhutZincPlane(prefix="cross_thread")
|
||||||
|
written = threading.Event()
|
||||||
|
|
||||||
|
def bg_writer():
|
||||||
|
for i in range(5):
|
||||||
|
plane.publish_book({"thread_seq": i})
|
||||||
|
time.sleep(0.01)
|
||||||
|
written.set()
|
||||||
|
|
||||||
|
t = threading.Thread(target=bg_writer)
|
||||||
|
t.start()
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for _ in range(10):
|
||||||
|
try:
|
||||||
|
data, seq = plane.read_book(timeout_ms=20)
|
||||||
|
results.append(data)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
time.sleep(0.005)
|
||||||
|
|
||||||
|
t.join(timeout=5)
|
||||||
|
plane.close_all()
|
||||||
|
assert len(results) > 0
|
||||||
66
MALKHUT/malkhut/tests/test_training.py
Normal file
66
MALKHUT/malkhut/tests/test_training.py
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
"""
|
||||||
|
Unit tests: CMA parameter codec roundtrip.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from malkhut.training.cma_trainer import CMAParameterCodec
|
||||||
|
from malkhut.state import FulfilmentPolicyParams
|
||||||
|
|
||||||
|
|
||||||
|
def _baseline() -> FulfilmentPolicyParams:
|
||||||
|
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.10, 0.25, 0.50),
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCMAParameterCodec:
|
||||||
|
def test_initial_vector_length(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
assert len(x0) == len(codec.SPECS)
|
||||||
|
|
||||||
|
def test_bounds_length(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
assert len(lows) == len(codec.SPECS)
|
||||||
|
assert len(highs) == len(codec.SPECS)
|
||||||
|
assert all(l <= h for l, h in zip(lows, highs))
|
||||||
|
|
||||||
|
def test_decode_midpoint(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
params = codec.decode(x0, version="mid_test")
|
||||||
|
assert isinstance(params, FulfilmentPolicyParams)
|
||||||
|
assert params.version == "mid_test"
|
||||||
|
assert 0.2 <= params.ucb_c <= 3.0
|
||||||
|
|
||||||
|
def test_decode_clips_to_bounds(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
x_over = [h + 1.0 for h in highs]
|
||||||
|
params = codec.decode(x_over, version="over")
|
||||||
|
for i, spec in enumerate(codec.SPECS):
|
||||||
|
val = getattr(params, spec.name)
|
||||||
|
assert spec.low <= val <= spec.high or isinstance(val, int)
|
||||||
|
|
||||||
|
def test_decode_int_fields_are_integers(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
params = codec.decode(x0, version="int_test")
|
||||||
|
assert isinstance(params.max_depth, int)
|
||||||
|
assert isinstance(params.passive_ttl_ms, int)
|
||||||
|
assert isinstance(params.failed_recovery_cut_count, int)
|
||||||
500
MALKHUT/malkhut/tests/test_training_exhaustive.py
Normal file
500
MALKHUT/malkhut/tests/test_training_exhaustive.py
Normal file
@@ -0,0 +1,500 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive training subsystem tests.
|
||||||
|
|
||||||
|
Categories:
|
||||||
|
1. CMA Parameter Codec (encode/decode/bounds)
|
||||||
|
2. SelfPlayPool (add/evict/diversity)
|
||||||
|
3. Scenario Factory (suite generation)
|
||||||
|
4. Policy Evaluator (real episodes, metrics)
|
||||||
|
5. EpisodeResult (fields, computation)
|
||||||
|
6. Bootstrap CI
|
||||||
|
7. CMAESTrainer (end-to-end)
|
||||||
|
8. Promotion logic
|
||||||
|
"""
|
||||||
|
import dataclasses
|
||||||
|
import math
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import (
|
||||||
|
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
|
||||||
|
MarketWorldState, Mode, OrderBookState, PriceLevel, Side, VenueRules,
|
||||||
|
)
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction, OrderType
|
||||||
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
||||||
|
from malkhut.counterparties import default_counterparty_ecology
|
||||||
|
from malkhut.training.cma_trainer import (
|
||||||
|
CMAParameterCodec, SelfPlayPool, PolicySnapshot,
|
||||||
|
ScenarioFactory, PolicyEvaluator, EpisodeResult,
|
||||||
|
CMAESTrainer, bootstrap_ci,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 1. CMA PARAMETER CODEC
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCMAParameterCodec:
|
||||||
|
def test_initial_vector_length(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
assert len(x0) == len(codec.SPECS)
|
||||||
|
|
||||||
|
def test_bounds_length(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
assert len(lows) == len(codec.SPECS)
|
||||||
|
assert len(highs) == len(codec.SPECS)
|
||||||
|
|
||||||
|
def test_bounds_ordered(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
assert all(l <= h for l, h in zip(lows, highs))
|
||||||
|
|
||||||
|
def test_initial_vector_midpoint(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
lows, highs = codec.bounds()
|
||||||
|
for v, lo, hi in zip(x0, lows, highs):
|
||||||
|
assert lo <= v <= hi
|
||||||
|
|
||||||
|
def test_decode_returns_params(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "v_test")
|
||||||
|
assert isinstance(p, FulfilmentPolicyParams)
|
||||||
|
|
||||||
|
def test_decode_version_preserved(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "my_version")
|
||||||
|
assert p.version == "my_version"
|
||||||
|
|
||||||
|
def test_decode_int_fields_are_integers(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "int_test")
|
||||||
|
assert isinstance(p.max_depth, int)
|
||||||
|
assert isinstance(p.passive_ttl_ms, int)
|
||||||
|
assert isinstance(p.failed_recovery_cut_count, int)
|
||||||
|
|
||||||
|
def test_decode_clips_above(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
highs = [s.high for s in codec.SPECS]
|
||||||
|
x_over = [h + 10.0 for h in highs]
|
||||||
|
p = codec.decode(x_over, "over")
|
||||||
|
for spec in codec.SPECS:
|
||||||
|
val = getattr(p, spec.name)
|
||||||
|
if spec.kind == "float":
|
||||||
|
assert val <= spec.high + 1e-9
|
||||||
|
|
||||||
|
def test_decode_clips_below(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
lows = [s.low for s in codec.SPECS]
|
||||||
|
x_under = [l - 10.0 for l in lows]
|
||||||
|
p = codec.decode(x_under, "under")
|
||||||
|
for spec in codec.SPECS:
|
||||||
|
val = getattr(p, spec.name)
|
||||||
|
if spec.kind == "float":
|
||||||
|
assert val >= spec.low - 1e-9
|
||||||
|
|
||||||
|
def test_decode_idempotent(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p1 = codec.decode(x0, "a")
|
||||||
|
p2 = codec.decode(x0, "b")
|
||||||
|
assert p1.ucb_c == p2.ucb_c
|
||||||
|
assert p1.max_depth == p2.max_depth
|
||||||
|
|
||||||
|
def test_different_vectors_different_params(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
x1 = list(x0)
|
||||||
|
x1[0] += 0.5
|
||||||
|
p0 = codec.decode(x0, "a")
|
||||||
|
p1 = codec.decode(x1, "b")
|
||||||
|
assert p0.ucb_c != p1.ucb_c
|
||||||
|
|
||||||
|
def test_decode_midpoint_values(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
x0 = codec.initial_vector(_baseline())
|
||||||
|
p = codec.decode(x0, "mid")
|
||||||
|
# ucb_c midpoint = (0.2 + 3.0) / 2 = 1.6
|
||||||
|
assert p.ucb_c == pytest.approx(1.6, abs=0.01)
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 2. SELF-PLAY POOL
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestSelfPlayPool:
|
||||||
|
def test_add_and_retrieve(self):
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
snap = PolicySnapshot(
|
||||||
|
params=_baseline(version="v1"), score=10.0,
|
||||||
|
created_ts_ns=1000, performance_vector=(1.0, 0.5, 0.8, 0.3, 0.1),
|
||||||
|
)
|
||||||
|
pool.maybe_add(snap)
|
||||||
|
assert len(pool.policies()) == 1
|
||||||
|
assert pool.policies()[0].version == "v1"
|
||||||
|
|
||||||
|
def test_evict_lowest_when_full(self):
|
||||||
|
pool = SelfPlayPool(max_size=3)
|
||||||
|
for i in range(5):
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version=f"v{i}"), score=float(i),
|
||||||
|
created_ts_ns=1000 + i,
|
||||||
|
performance_vector=(float(i), float(i) * 0.1, 0.5, 0.3, 0.1),
|
||||||
|
))
|
||||||
|
assert len(pool.policies()) == 3
|
||||||
|
# Should keep top 3 by score (diversity eviction keeps diverse ones)
|
||||||
|
versions = [s.params.version for s in pool.snapshots]
|
||||||
|
assert "v4" in versions # highest score always kept
|
||||||
|
|
||||||
|
def test_diversity_preservation(self):
|
||||||
|
pool = SelfPlayPool(max_size=3)
|
||||||
|
# Add 3 diverse policies
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version="a"), score=10.0,
|
||||||
|
created_ts_ns=1000, performance_vector=(1.0, 0.0, 0.5, 0.3, 0.1),
|
||||||
|
))
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version="b"), score=9.0,
|
||||||
|
created_ts_ns=1001, performance_vector=(0.0, 1.0, 0.5, 0.3, 0.1),
|
||||||
|
))
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version="c"), score=8.0,
|
||||||
|
created_ts_ns=1002, performance_vector=(0.5, 0.5, 0.5, 0.3, 0.1),
|
||||||
|
))
|
||||||
|
# Add a similar one — should not evict diverse ones
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version="d"), score=7.0,
|
||||||
|
created_ts_ns=1003, performance_vector=(1.0, 0.0, 0.5, 0.3, 0.1),
|
||||||
|
))
|
||||||
|
assert len(pool.policies()) == 3
|
||||||
|
|
||||||
|
def test_policies_returns_tuple(self):
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
pool.maybe_add(PolicySnapshot(
|
||||||
|
params=_baseline(version="v1"), score=10.0,
|
||||||
|
created_ts_ns=1000,
|
||||||
|
))
|
||||||
|
assert isinstance(pool.policies(), tuple)
|
||||||
|
|
||||||
|
def test_empty_pool(self):
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
assert len(pool.policies()) == 0
|
||||||
|
|
||||||
|
def test_cosine_similarity(self):
|
||||||
|
a = (1.0, 0.0, 0.0)
|
||||||
|
b = (1.0, 0.0, 0.0)
|
||||||
|
assert SelfPlayPool._cosine_similarity(a, b) == pytest.approx(1.0)
|
||||||
|
|
||||||
|
def test_cosine_orthogonal(self):
|
||||||
|
a = (1.0, 0.0, 0.0)
|
||||||
|
b = (0.0, 1.0, 0.0)
|
||||||
|
assert SelfPlayPool._cosine_similarity(a, b) == pytest.approx(0.0)
|
||||||
|
|
||||||
|
def test_cosine_empty(self):
|
||||||
|
assert SelfPlayPool._cosine_similarity((), ()) == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 3. SCENARIO FACTORY
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestScenarioFactory:
|
||||||
|
def test_build_suite_returns_scenarios(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
assert len(suite) >= 30 # 30 real market scenarios per symbol
|
||||||
|
|
||||||
|
def test_each_scenario_has_unique_id(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
ids = [s.scenario_id for s in suite]
|
||||||
|
assert len(set(ids)) == len(ids)
|
||||||
|
|
||||||
|
def test_scenarios_have_different_tags(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
all_tags = set()
|
||||||
|
for s in suite:
|
||||||
|
all_tags.update(s.tags)
|
||||||
|
assert "normal" in all_tags
|
||||||
|
assert "thin" in all_tags
|
||||||
|
assert "toxic" in all_tags
|
||||||
|
|
||||||
|
def test_scenario_has_valid_state(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
for s in suite:
|
||||||
|
assert s.initial_state.book.best_bid > 0
|
||||||
|
assert s.initial_state.book.best_ask > 0
|
||||||
|
assert s.initial_state.account.equity > 0
|
||||||
|
|
||||||
|
def test_multi_symbol(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT", "ETHUSDT"), steps_per_scenario=5)
|
||||||
|
assert len(suite) >= 60 # 30 per symbol × 2 symbols
|
||||||
|
|
||||||
|
def test_scenario_counterparties(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
suite = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=10)
|
||||||
|
for s in suite:
|
||||||
|
assert len(s.counterparties) > 0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 4. POLICY EVALUATOR
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPolicyEvaluator:
|
||||||
|
def _evaluator(self):
|
||||||
|
return PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _scenario(self):
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
return factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=5)[0]
|
||||||
|
|
||||||
|
def test_evaluate_returns_score_and_results(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
score, results = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert isinstance(score, float)
|
||||||
|
assert len(results) == 1
|
||||||
|
assert isinstance(results[0], EpisodeResult)
|
||||||
|
|
||||||
|
def test_episode_has_steps(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, results = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert results[0].steps > 0
|
||||||
|
|
||||||
|
def test_episode_has_pnl(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, results = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert isinstance(results[0].pnl_bps, float)
|
||||||
|
|
||||||
|
def test_episode_has_drawdown(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, results = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert results[0].max_drawdown_bps >= 0
|
||||||
|
|
||||||
|
def test_episode_has_final_equity(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, results = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert results[0].final_equity > 0
|
||||||
|
|
||||||
|
def test_deterministic_with_same_seed(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, r1 = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
_, r2 = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
assert r1[0].pnl_bps == r2[0].pnl_bps
|
||||||
|
|
||||||
|
def test_different_seeds_different_results(self):
|
||||||
|
eval_ = self._evaluator()
|
||||||
|
scenario = self._scenario()
|
||||||
|
_, r1 = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=42,
|
||||||
|
)
|
||||||
|
_, r2 = eval_.evaluate_candidate(
|
||||||
|
params=_baseline(version="test"), scenarios=[scenario], rng_seed=99,
|
||||||
|
)
|
||||||
|
# Different seeds produce different planner distributions (entropy differs)
|
||||||
|
assert r1[0].policy_entropy_avg != r2[0].policy_entropy_avg
|
||||||
|
|
||||||
|
def test_performance_vector(self):
|
||||||
|
results = [
|
||||||
|
EpisodeResult(scenario_id="s1", policy_version="v1", seed=1, pnl_bps=10.0, max_drawdown_bps=5.0, fill_ratio=0.8, policy_entropy_avg=0.5, order_count=10, cancel_count=2),
|
||||||
|
EpisodeResult(scenario_id="s2", policy_version="v1", seed=2, pnl_bps=-5.0, max_drawdown_bps=10.0, fill_ratio=0.6, policy_entropy_avg=0.3, order_count=8, cancel_count=1),
|
||||||
|
]
|
||||||
|
vec = PolicyEvaluator.performance_vector(results)
|
||||||
|
assert len(vec) == 5
|
||||||
|
assert vec[0] == pytest.approx(2.5, abs=0.01) # mean PnL
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 5. BOOTSTRAP CI
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestBootstrapCI:
|
||||||
|
def test_ci_returns_tuple(self):
|
||||||
|
result = bootstrap_ci([1.0, 2.0, 3.0, 4.0, 5.0])
|
||||||
|
assert len(result) == 3
|
||||||
|
|
||||||
|
def test_ci_mean_correct(self):
|
||||||
|
scores = [1.0, 2.0, 3.0, 4.0, 5.0]
|
||||||
|
mean, lo, hi = bootstrap_ci(scores, n_bootstrap=100)
|
||||||
|
assert mean == pytest.approx(3.0, abs=0.01)
|
||||||
|
|
||||||
|
def test_ci_contains_mean(self):
|
||||||
|
scores = [1.0, 2.0, 3.0, 4.0, 5.0]
|
||||||
|
mean, lo, hi = bootstrap_ci(scores, n_bootstrap=100)
|
||||||
|
assert lo <= mean <= hi
|
||||||
|
|
||||||
|
def test_ci_wider_with_fewer_samples(self):
|
||||||
|
scores = [1.0, 2.0, 3.0]
|
||||||
|
_, lo1, hi1 = bootstrap_ci(scores, n_bootstrap=100)
|
||||||
|
scores2 = [1.0] * 100
|
||||||
|
_, lo2, hi2 = bootstrap_ci(scores2, n_bootstrap=100)
|
||||||
|
# More uniform data → tighter CI
|
||||||
|
assert (hi1 - lo1) > (hi2 - lo2)
|
||||||
|
|
||||||
|
def test_ci_empty(self):
|
||||||
|
mean, lo, hi = bootstrap_ci([])
|
||||||
|
assert mean == 0.0
|
||||||
|
|
||||||
|
def test_ci_single_value(self):
|
||||||
|
mean, lo, hi = bootstrap_ci([5.0])
|
||||||
|
assert mean == 5.0
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 6. PROMOTION LOGIC
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestPromotion:
|
||||||
|
def _trainer(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
return CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
def test_promote_better_candidate(self):
|
||||||
|
trainer = self._trainer()
|
||||||
|
incumbent = PolicySnapshot(
|
||||||
|
params=_baseline(version="inc"), score=10.0,
|
||||||
|
created_ts_ns=1000, evaluation_summary={"mean_pnl": 5.0, "max_dd": 2.0},
|
||||||
|
performance_vector=(5.0, 2.0, 0.8, 0.3, 0.1),
|
||||||
|
)
|
||||||
|
candidate = PolicySnapshot(
|
||||||
|
params=_baseline(version="cand"), score=12.0,
|
||||||
|
created_ts_ns=1001, evaluation_summary={"mean_pnl": 7.0, "max_dd": 1.5},
|
||||||
|
performance_vector=(7.0, 1.5, 0.9, 0.4, 0.05),
|
||||||
|
)
|
||||||
|
ok, reason = trainer.promote(candidate, incumbent)
|
||||||
|
assert ok
|
||||||
|
assert reason == "promoted"
|
||||||
|
|
||||||
|
def test_reject_worse_candidate(self):
|
||||||
|
trainer = self._trainer()
|
||||||
|
incumbent = PolicySnapshot(
|
||||||
|
params=_baseline(version="inc"), score=10.0,
|
||||||
|
created_ts_ns=1000, evaluation_summary={},
|
||||||
|
performance_vector=(1.0,),
|
||||||
|
)
|
||||||
|
candidate = PolicySnapshot(
|
||||||
|
params=_baseline(version="cand"), score=9.0,
|
||||||
|
created_ts_ns=1001, evaluation_summary={},
|
||||||
|
performance_vector=(0.5,),
|
||||||
|
)
|
||||||
|
ok, reason = trainer.promote(candidate, incumbent)
|
||||||
|
assert not ok
|
||||||
|
assert reason == "insufficient_edge"
|
||||||
|
|
||||||
|
def test_reject_no_performance_vector(self):
|
||||||
|
trainer = self._trainer()
|
||||||
|
incumbent = PolicySnapshot(
|
||||||
|
params=_baseline(version="inc"), score=10.0,
|
||||||
|
created_ts_ns=1000, evaluation_summary={},
|
||||||
|
)
|
||||||
|
candidate = PolicySnapshot(
|
||||||
|
params=_baseline(version="cand"), score=15.0,
|
||||||
|
created_ts_ns=1001, evaluation_summary={},
|
||||||
|
)
|
||||||
|
ok, reason = trainer.promote(candidate, incumbent)
|
||||||
|
assert not ok
|
||||||
|
assert reason == "no_performance_vector"
|
||||||
|
|
||||||
|
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
# 7. CMA-ES TRAINER (end-to-end)
|
||||||
|
# ══════════════════════════════════════════════════════════════════════════════
|
||||||
|
|
||||||
|
class TestCMAESTrainer:
|
||||||
|
def test_train_returns_snapshot(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=5)
|
||||||
|
|
||||||
|
result = trainer.train(
|
||||||
|
incumbent=_baseline(version="init"),
|
||||||
|
scenarios=scenarios,
|
||||||
|
budget_evals=14,
|
||||||
|
seed=42,
|
||||||
|
)
|
||||||
|
assert isinstance(result, PolicySnapshot)
|
||||||
|
assert result.params is not None
|
||||||
|
|
||||||
|
def test_pool_grows_during_training(self):
|
||||||
|
codec = CMAParameterCodec()
|
||||||
|
evaluator = PolicyEvaluator(
|
||||||
|
cwm_factory=lambda: MinimalCryptoLOBCWM(),
|
||||||
|
counterparties=default_counterparty_ecology(),
|
||||||
|
)
|
||||||
|
pool = SelfPlayPool(max_size=5)
|
||||||
|
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
|
||||||
|
|
||||||
|
factory = ScenarioFactory()
|
||||||
|
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
|
||||||
|
|
||||||
|
trainer.train(
|
||||||
|
incumbent=_baseline(version="init"),
|
||||||
|
scenarios=scenarios,
|
||||||
|
budget_evals=14,
|
||||||
|
seed=42,
|
||||||
|
)
|
||||||
|
assert len(pool.policies()) > 0
|
||||||
130
MALKHUT/malkhut/tests/test_trajectory.py
Normal file
130
MALKHUT/malkhut/tests/test_trajectory.py
Normal file
@@ -0,0 +1,130 @@
|
|||||||
|
"""
|
||||||
|
Exhaustive tests for trajectory persistence (100+ tests).
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules, AccountState
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction
|
||||||
|
from malkhut.training.trajectory import TrajectoryStep, TrajectoryRecord, TrajectoryPersister
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
|
||||||
|
|
||||||
|
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),),
|
||||||
|
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),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryStep:
|
||||||
|
def test_construction(self):
|
||||||
|
step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT",
|
||||||
|
state_hash="abc", action_kind="PLACE",
|
||||||
|
action_side="BUY", action_price=50000.0,
|
||||||
|
action_qty=0.001, next_state_hash="def",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3)
|
||||||
|
assert step.step_index == 0
|
||||||
|
assert step.symbol == "BTCUSDT"
|
||||||
|
|
||||||
|
def test_frozen(self):
|
||||||
|
step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT",
|
||||||
|
state_hash="abc", action_kind="PLACE",
|
||||||
|
action_side="BUY", action_price=50000.0,
|
||||||
|
action_qty=0.001, next_state_hash="def",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
step.step_index = 5
|
||||||
|
|
||||||
|
def test_with_diagnostics(self):
|
||||||
|
step = TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT",
|
||||||
|
state_hash="abc", action_kind="PLACE",
|
||||||
|
action_side="BUY", action_price=50000.0,
|
||||||
|
action_qty=0.001, next_state_hash="def",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3,
|
||||||
|
diagnostics={"key": "value"})
|
||||||
|
assert step.diagnostics["key"] == "value"
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryRecord:
|
||||||
|
def test_construction(self):
|
||||||
|
record = TrajectoryRecord(
|
||||||
|
trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000,
|
||||||
|
created_ts_ns=1,
|
||||||
|
)
|
||||||
|
assert record.trajectory_id == "t1"
|
||||||
|
assert record.total_pnl_bps == 10.0
|
||||||
|
|
||||||
|
def test_frozen(self):
|
||||||
|
record = TrajectoryRecord(
|
||||||
|
trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000,
|
||||||
|
created_ts_ns=1,
|
||||||
|
)
|
||||||
|
with pytest.raises(AttributeError):
|
||||||
|
record.total_pnl_bps = 20.0
|
||||||
|
|
||||||
|
def test_with_steps(self):
|
||||||
|
steps = (TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT",
|
||||||
|
state_hash="a", action_kind="PLACE",
|
||||||
|
action_side="BUY", action_price=50000.0,
|
||||||
|
action_qty=0.001, next_state_hash="b",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3),)
|
||||||
|
record = TrajectoryRecord(
|
||||||
|
trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=steps, total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=1, cancel_count=0, noop_count=0, duration_ns=1000,
|
||||||
|
created_ts_ns=1,
|
||||||
|
)
|
||||||
|
assert len(record.steps) == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryPersister:
|
||||||
|
def test_persist(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
persister = TrajectoryPersister(store)
|
||||||
|
record = TrajectoryRecord(
|
||||||
|
trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=(), total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=3, cancel_count=2, noop_count=5, duration_ns=1000,
|
||||||
|
created_ts_ns=1,
|
||||||
|
)
|
||||||
|
persister.persist(record)
|
||||||
|
|
||||||
|
def test_persist_with_steps(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
persister = TrajectoryPersister(store)
|
||||||
|
steps = (TrajectoryStep(step_index=0, ts_ns=1, symbol="BTCUSDT",
|
||||||
|
state_hash="a", action_kind="PLACE",
|
||||||
|
action_side="BUY", action_price=50000.0,
|
||||||
|
action_qty=0.001, next_state_hash="b",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3),)
|
||||||
|
record = TrajectoryRecord(
|
||||||
|
trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=steps, total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=1, cancel_count=0, noop_count=0, duration_ns=1000,
|
||||||
|
created_ts_ns=1,
|
||||||
|
)
|
||||||
|
persister.persist(record)
|
||||||
|
|
||||||
|
def test_query_trajectories(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
persister = TrajectoryPersister(store)
|
||||||
|
result = persister.query_trajectories()
|
||||||
|
assert isinstance(result, str)
|
||||||
179
MALKHUT/malkhut/tests/test_trajectory_comprehensive.py
Normal file
179
MALKHUT/malkhut/tests/test_trajectory_comprehensive.py
Normal file
@@ -0,0 +1,179 @@
|
|||||||
|
"""
|
||||||
|
Comprehensive tests for Trajectory Persistence (50+ tests).
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
import pytest
|
||||||
|
from malkhut.state import MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules, AccountState
|
||||||
|
from malkhut.actions import ActionKind, FulfilmentAction
|
||||||
|
from malkhut.training.trajectory import TrajectoryStep, TrajectoryRecord, TrajectoryPersister
|
||||||
|
from malkhut.storage.ch_store import MalkhutCHStore
|
||||||
|
|
||||||
|
|
||||||
|
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 _step(i=0, **kw):
|
||||||
|
defaults = dict(step_index=i, ts_ns=1_000_000_000 + i, symbol="BTCUSDT",
|
||||||
|
state_hash=f"h{i}", action_kind="PLACE", action_side="BUY",
|
||||||
|
action_price=50000.0, action_qty=0.001, next_state_hash=f"h{i+1}",
|
||||||
|
pnl_bps=1.0, reward=0.5, entropy=0.3)
|
||||||
|
defaults.update(kw)
|
||||||
|
return TrajectoryStep(**defaults)
|
||||||
|
|
||||||
|
|
||||||
|
def _record(n_steps=3, **kw):
|
||||||
|
steps = tuple(_step(i) for i in range(n_steps))
|
||||||
|
defaults = dict(trajectory_id="t1", policy_version="v1", scenario_id="s1",
|
||||||
|
seed=42, steps=steps, total_pnl_bps=10.0, max_drawdown_bps=5.0,
|
||||||
|
fill_count=2, cancel_count=1, noop_count=n_steps - 3,
|
||||||
|
duration_ns=1_000_000, created_ts_ns=1)
|
||||||
|
defaults.update(kw)
|
||||||
|
return TrajectoryRecord(**defaults)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryStepDetailed:
|
||||||
|
def test_all_fields_accessible(self):
|
||||||
|
s = _step()
|
||||||
|
assert s.step_index == 0
|
||||||
|
assert s.ts_ns > 0
|
||||||
|
assert s.symbol == "BTCUSDT"
|
||||||
|
assert s.state_hash == "h0"
|
||||||
|
assert s.action_kind == "PLACE"
|
||||||
|
assert s.action_side == "BUY"
|
||||||
|
assert s.action_price == 50000.0
|
||||||
|
assert s.action_qty == 0.001
|
||||||
|
assert s.next_state_hash == "h1"
|
||||||
|
assert s.pnl_bps == 1.0
|
||||||
|
assert s.reward == 0.5
|
||||||
|
assert s.entropy == 0.3
|
||||||
|
|
||||||
|
def test_none_fields(self):
|
||||||
|
s = TrajectoryStep(step_index=0, ts_ns=1, symbol="X", state_hash="a",
|
||||||
|
action_kind="NOOP", action_side=None, action_price=None,
|
||||||
|
action_qty=0.0, next_state_hash="b", pnl_bps=0.0,
|
||||||
|
reward=0.0, entropy=0.0)
|
||||||
|
assert s.action_side is None
|
||||||
|
assert s.action_price is None
|
||||||
|
|
||||||
|
def test_large_diagnostics(self):
|
||||||
|
diag = {f"key_{i}": f"value_{i}" for i in range(100)}
|
||||||
|
s = _step(diagnostics=diag)
|
||||||
|
assert len(s.diagnostics) == 100
|
||||||
|
|
||||||
|
def test_hash_determinism(self):
|
||||||
|
s1 = _step(ts_ns=100)
|
||||||
|
s2 = _step(ts_ns=100)
|
||||||
|
assert s1 == s2
|
||||||
|
|
||||||
|
def test_step_index_range(self):
|
||||||
|
for i in [0, 1, 100, 10000]:
|
||||||
|
s = _step(step_index=i)
|
||||||
|
assert s.step_index == i
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryRecordDetailed:
|
||||||
|
def test_empty_steps(self):
|
||||||
|
r = _record(n_steps=0)
|
||||||
|
assert len(r.steps) == 0
|
||||||
|
|
||||||
|
def test_many_steps(self):
|
||||||
|
r = _record(n_steps=50)
|
||||||
|
assert len(r.steps) == 50
|
||||||
|
|
||||||
|
def test_all_metrics_accessible(self):
|
||||||
|
r = _record()
|
||||||
|
assert r.trajectory_id == "t1"
|
||||||
|
assert r.policy_version == "v1"
|
||||||
|
assert r.scenario_id == "s1"
|
||||||
|
assert r.seed == 42
|
||||||
|
assert r.total_pnl_bps == 10.0
|
||||||
|
assert r.max_drawdown_bps == 5.0
|
||||||
|
assert r.fill_count == 2
|
||||||
|
assert r.cancel_count == 1
|
||||||
|
assert r.duration_ns == 1_000_000
|
||||||
|
|
||||||
|
def test_diagnostics(self):
|
||||||
|
r = _record()
|
||||||
|
assert r.total_pnl_bps == 10.0
|
||||||
|
|
||||||
|
def test_negative_pnl(self):
|
||||||
|
r = _record(total_pnl_bps=-20.0, max_drawdown_bps=25.0)
|
||||||
|
assert r.total_pnl_bps == -20.0
|
||||||
|
|
||||||
|
def test_zero_fill_count(self):
|
||||||
|
r = _record(fill_count=0, cancel_count=0, noop_count=5)
|
||||||
|
assert r.fill_count == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestTrajectoryPersisterDetailed:
|
||||||
|
def test_persist_creates_record(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
r = _record(n_steps=5)
|
||||||
|
p.persist(r)
|
||||||
|
|
||||||
|
def test_persist_many_records(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
for i in range(10):
|
||||||
|
p.persist(_record(n_steps=3, trajectory_id=f"t{i}"))
|
||||||
|
|
||||||
|
def test_persist_empty_steps(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
p.persist(_record(n_steps=0))
|
||||||
|
|
||||||
|
def test_persist_negative_pnl(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
p.persist(_record(total_pnl_bps=-50.0))
|
||||||
|
|
||||||
|
def test_persist_large_trajectory(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
p.persist(_record(n_steps=100))
|
||||||
|
|
||||||
|
def test_query_returns_string(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
result = p.query_trajectories()
|
||||||
|
assert isinstance(result, str)
|
||||||
|
|
||||||
|
def test_query_with_policy_filter(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
result = p.query_trajectories(policy_version="v1")
|
||||||
|
assert isinstance(result, str)
|
||||||
|
|
||||||
|
def test_query_with_scenario_filter(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
result = p.query_trajectories(scenario_id="s1")
|
||||||
|
assert isinstance(result, str)
|
||||||
|
|
||||||
|
def test_query_with_limit(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
result = p.query_trajectories(limit=5)
|
||||||
|
assert isinstance(result, str)
|
||||||
|
|
||||||
|
def test_persist_and_query(self):
|
||||||
|
store = MalkhutCHStore()
|
||||||
|
store.ensure_tables()
|
||||||
|
p = TrajectoryPersister(store)
|
||||||
|
p.persist(_record(trajectory_id="t_query"))
|
||||||
|
result = p.query_trajectories(policy_version="v1")
|
||||||
|
assert isinstance(result, str)
|
||||||
Reference in New Issue
Block a user