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

346 lines
15 KiB
Python
Raw Normal View History

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