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