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
370 lines
17 KiB
Python
370 lines
17 KiB
Python
"""
|
|
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
|