#!/usr/bin/env python3 """ MALKHUT MCTS Planner E2E — uses actual MCTS planner, not random actions. This is what makes strategies ADAPTIVE — the planner explores the action space. """ from __future__ import annotations import json, math, os, random, resource, sys, time from collections import defaultdict from typing import Any, Dict, List, Optional _HERE = os.path.dirname(os.path.abspath(__file__)) if _HERE not in sys.path: sys.path.insert(0, _HERE) from malkhut.state import ( AccountState, ActionKind, FulfilmentPolicyParams, MarketWorldState, OrderType, PositionState, Side, ExecutionIntent, IntentKind, ) from malkhut.actions import FulfilmentAction, PlannedPolicy from malkhut.cwm.hft_cwm import HftBacktestCWM from malkhut.risk.gate import RiskGate from malkhut.planner.alternatives import create_planner from malkhut.training.cma_trainer import ScenarioFactory, SelfPlayPool from malkhut.counterparties import ToxicTakerPolicy, PassiveMakerPolicy, LatencyArbPolicy, NoiseTraderPolicy from malkhut.counterparties_extended import * DURATION_S = 300 # 5 minutes (smoke test) STEPS = 20; SEED = 42 def _params(): return FulfilmentPolicyParams( version='mcts_test', ucb_c=1.414, max_sims=64, max_depth=2, rollout_depth=2, root_temperature=0.5, min_root_entropy=0.25, quote_offsets_ticks=(0,1,2), quote_size_fractions=(0.10,0.25,0.50), passive_ttl_ms=200, aggressive_ttl_ms=50, maker_edge_min_bps=0.5, cross_spread_edge_min_bps=5.0, adverse_toxicity_cancel_threshold=0.5, queue_churn_cancel_threshold=0.5, mae_tail_cut_bps=50.0, mfe_giveback_cut_fraction=0.5, max_time_in_loss_s=300.0, failed_recovery_cut_count=3, recovery_velocity_min_bps_per_s=0.0, max_symbol_notional_fraction=0.20, max_single_order_notional_fraction=0.05, reduce_when_global_up_fraction=0.30, session_profit_lock_fraction=0.02, w_expected_pnl=1.0, w_fill_probability=0.5, w_adverse_selection=2.0, w_queue_priority=0.5, w_inventory_risk=1.5, w_tail_loss=5.0, w_fee_quality=0.5, w_time_decay=0.3, w_policy_entropy=0.5, robust_tail_weight=2.0, toxic_counterparty_weight=3.0, low_liquidity_weight=2.0, latency_stress_weight=1.0, wait_to_retry_ms=100, chase_enabled=True, chase_offset_ticks=2, chase_max_retries=3, urgency_taker_threshold=0.65, urgency_taker_penalty_bps=2.0, execution_friction_threshold_bps=3.0, ) def _intent(symbol: str = "BTCUSDT", urgency: float = 0.5) -> ExecutionIntent: return ExecutionIntent( intent_id="e2e", ts_ns=1_000_000_000, symbol=symbol, kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0, urgency=urgency, alpha_horizon_s=60.0, alpha_bps=2.0, max_slippage_bps=5.0, prefer_maker=False, reduce_only=False, ttl_s=300.0, reason="e2e_test", ) def run_episode_mcts(cwm, scenario, params, steps, seed, risk_gate): """Run episode with MCTS planner instead of random actions.""" import random as _random from malkhut.planner.alternatives import create_planner state = scenario.initial_state rng = _random.Random(seed) planner = create_planner( "sm_mcts", cwm=cwm, counterparties=scenario.counterparties, rng_seed=seed, ) fq_s = fq_e = fq_sr = fq_fv = fq_fc = fq_t = 0 action_counts = defaultdict(int) for step in range(steps): # Set intent on state intent = _intent(scenario.symbol, urgency=0.8) state_with_intent = MarketWorldState( ts_ns=state.ts_ns, mode=state.mode, venue=state.venue, book=state.book, account=state.account, open_orders=state.open_orders, trade_path=state.trade_path, intent=intent, funding_bps=state.funding_bps, volatility_state=state.volatility_state, market_regime=state.market_regime, ) # Plan with MCTS try: planned = planner.plan(root_state=state_with_intent, params=params, budget_ms=25) action = planned.selected_action except Exception: action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) action_counts[action.kind.value] += 1 # Risk gate if action.kind != ActionKind.NOOP: p = PlannedPolicy(actions=(action,), probabilities=(1.0,), selected_action=action, diagnostics={}) if not risk_gate.validate(state, p, params).approved: action = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) # Transition cp = tuple(c.rollout_action(state, rng) for c in scenario.counterparties) fq_t += 1 state = cwm.transition(state, (action, *cp)) if state.fill_quality: fq = state.fill_quality fq_s += fq.slippage_bps; fq_e += fq.expected_slippage_bps fq_sr += fq.slippage_bps - fq.expected_slippage_bps fq_fv += fq.fill_value_score if fq.filled: fq_fc += 1 n = max(fq_t, 1) pnl = (state.account.equity - 10000.0) / 10000.0 * 10_000 return pnl, fq_s/n, fq_e/n, fq_sr/n, fq_fv/n, fq_fc/max(fq_t,1), dict(action_counts) def main(): t_start = time.time() DURATION_S = 300 print("=" * 80) print("MALKHUT MCTS PLANNER E2E — 5-MINUTE SMOKE TEST") print(" Uses actual DecoupledUCBPlanner (not random actions)") print(f" CWM: HftBacktestCWM (dynamic book)") print(f" Swarm: 11 diverse opponents") print(f" Duration: {DURATION_S}s") print("=" * 80, flush=True) risk_gate = RiskGate() params = _params() rng = random.Random(SEED) all_pnls = []; all_fv = []; all_sr = []; all_fr = [] all_actions = defaultdict(int) all_slippage = [] print(f"\nRunning MCTS episodes...", flush=True) n_done = 0 while time.time() < t_start + DURATION_S - 30: for sym in ["BTCUSDT", "ETHUSDT", "DOGEUSDT"]: cwm = HftBacktestCWM(use_queue_model=True, use_dynamic_book=True) factory = ScenarioFactory(exchange_id="bingx") scenarios = factory.build_suite(symbols=[sym], steps_per_scenario=STEPS, seed=SEED) sc = scenarios[n_done % len(scenarios)] try: pnl, sl, e, sr, fv, fr, ac = run_episode_mcts( cwm, sc, params, STEPS, SEED + n_done, risk_gate, ) all_pnls.append(pnl) all_fv.append(fv) all_sr.append(sr) all_fr.append(fr) all_slippage.append(sl) for k, v in ac.items(): all_actions[k] += v n_done += 1 elapsed = time.time() - t_start h = elapsed / 3600; m = (elapsed % 3600) / 60 rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024 avg_fv = sum(all_fv[-20:]) / min(len(all_fv), 20) print(f" [{h:.0f}m{m:.0f}s] ep={n_done} | fv={avg_fv:.3f} | " f"pnl={sum(all_pnls[-20:])/min(len(all_pnls),20):+.0f} | " f"sl={sum(all_slippage[-20:])/min(len(all_slippage),20):.3f} | " f"RSS={rss:.0f}MB", flush=True) except Exception as e: print(f" Error: {e}", flush=True) if time.time() > t_start + DURATION_S - 30: break # Final report n = len(all_pnls) elapsed = time.time() - t_start print(f"\n{'='*80}") print(f" MCTS PLANNER E2E — FINAL ({elapsed:.0f}s, {n} episodes)") print(f"{'='*80}") if n > 0: print(f"\n PERFORMANCE") print(f" Avg PnL: {sum(all_pnls)/n:+.1f} bps") print(f" Win rate: {sum(1 for p in all_pnls if p>0)/n*100:.1f}%") print(f" Fill value: {sum(all_fv)/n:.4f}") print(f" Surprise: {sum(all_sr)/n:+.3f} bps") print(f" Fill rate: {sum(all_fr)/n*100:.1f}%") print(f"\n ACTION DISTRIBUTION (MCTS planner)") total = sum(all_actions.values()) for k, v in sorted(all_actions.items(), key=lambda x: -x[1]): if v > 0: print(f" {k:20s}: {v:5d} ({v/max(total,1)*100:.1f}%)") print(f"\n COMPARISON: MCTS vs RANDOM") print(f" MCTS fill_value: {sum(all_fv)/n:.4f}") print(f" RANDOM fill_value: 0.392 (from 3h run)") print(f" STATIC fill_value: 0.890 (from smoke test)") print(f"{'='*80}", flush=True) if __name__ == "__main__": try: main() except Exception as e: import traceback; print(f"\nFATAL: {e}", flush=True); traceback.print_exc(); sys.exit(1)