Files
sentiment-engine/MALKHUT/malkhut/tests/test_dsl_new_features.py

314 lines
13 KiB
Python
Raw Normal View History

"""
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