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