314 lines
13 KiB
Python
314 lines
13 KiB
Python
|
|
"""
|
||
|
|
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
|