""" Code World Model (CWM) — deterministic exchange transition function. Full exchange mechanics: - Price-time priority with sequential level consumption - Partial fills across multiple levels - Queue position estimation - Latency injection (feed + order) - Maker/taker fee application - Post-only rejection - IOC/FOK/LIMIT/REDUCE_ONLY semantics - Tick/lot rounding - Open order aging (TTL expiry) - Path-state update (MAE/MFE/recovery tracking) - Mark-to-market Determinism: same state + same joint action + same seed = identical output. """ from __future__ import annotations import math import time from typing import List, Optional, Protocol, Sequence, Tuple import numpy as np from malkhut.state import ( AccountState, FulfilmentPolicyParams, MarketWorldState, Mode, OpenOrderState, OrderBookState, PositionState, PriceLevel, Side, TradePathState, ) from malkhut.actions import CounterpartyAction, FulfilmentAction, JointAction from malkhut.features import DefaultFeatureExtractor, FeatureExtractor # Import numba-accelerated functions with fallback try: from malkhut.cwm.numba_core import ( fill_from_levels as _nb_fill, round_tick as _nb_round_tick, round_lot as _nb_round_lot, clip_lots as _nb_clip_lots, ) _HAS_NUMBA = True except ImportError: _HAS_NUMBA = False class CodeWorldModel(Protocol): """Deterministic transition model. Same state + action + seed = identical output.""" def transition( self, state: MarketWorldState, joint_action: JointAction, ) -> MarketWorldState: ... def reward( self, prev_state: MarketWorldState, action: FulfilmentAction, next_state: MarketWorldState, params: FulfilmentPolicyParams, ) -> float: ... def terminal(self, state: MarketWorldState, depth: int) -> bool: ... def materialize_price_from_action( state: MarketWorldState, action: FulfilmentAction, ) -> Optional[float]: if action.side is None: return None tick = state.venue.tick_size if action.kind.value == "CROSS_SPREAD": if action.side == Side.BUY: return state.book.best_ask if state.book.asks else None else: return state.book.best_bid if state.book.bids else None if action.side == Side.BUY: if not state.book.bids: return None return state.book.best_bid - action.price_ticks_from_best * tick if not state.book.asks: return None return state.book.best_ask + action.price_ticks_from_best * tick def _round_tick(price: float, tick: float) -> float: return round(price / tick) * tick def _round_lot(qty: float, lot: float) -> float: return round(qty / lot) * lot def _clip_lots(qty: float, lot: float, min_qty: float) -> float: q = _round_lot(qty, lot) return q if q >= min_qty else 0.0 def _fill_from_levels( levels: List[PriceLevel], qty_remaining: float, lot: float, min_qty: float, ) -> Tuple[float, float, List[PriceLevel]]: """ Consume qty from price levels (price-time priority). Returns (filled_qty, avg_fill_price, remaining_levels). Uses numba-accelerated inner loop when available. """ if _HAS_NUMBA and len(levels) > 0: # Convert to numpy arrays for numba prices = np.array([l.price for l in levels], dtype=np.float64) qtys = np.array([l.qty for l in levels], dtype=np.float64) # Determine side from price ordering (descending = bids, ascending = asks) is_buy = len(levels) > 1 and levels[0].price > levels[-1].price filled, avg_price, new_bid_q, new_ask_q = _nb_fill( prices if not is_buy else np.array([], dtype=np.float64), qtys if not is_buy else np.array([], dtype=np.float64), prices if is_buy else np.array([], dtype=np.float64), qtys if is_buy else np.array([], dtype=np.float64), qty_remaining, lot, min_qty, is_buy, ) # Reconstruct remaining levels remaining = [] new_qtys = new_ask_q if is_buy else new_bid_q for i, level in enumerate(levels): if i < len(new_qtys) and new_qtys[i] > 0: remaining.append(PriceLevel(price=level.price, qty=new_qtys[i])) return filled, avg_price, remaining # Pure Python fallback filled = 0.0 total_cost = 0.0 remaining = list(levels) while qty_remaining > 1e-12 and remaining: level = remaining[0] take = min(qty_remaining, level.qty) take = _clip_lots(take, lot, min_qty) if take <= 0: break filled += take total_cost += take * level.price qty_remaining -= take new_qty = level.qty - take if new_qty < min_qty: remaining.pop(0) else: remaining[0] = PriceLevel(price=level.price, qty=new_qty) avg_price = total_cost / filled if filled > 0 else 0.0 return filled, avg_price, remaining def _update_path_state( state: MarketWorldState, fill_price: float, fill_qty: float, fill_side: Side, now_ts: int, ) -> Optional[TradePathState]: """Update trade path state after a fill.""" old_path = state.trade_path venue_mid = state.book.mid if state.book.bids and state.book.asks else 0.0 if old_path is None: # New position opened pnl_bps = 0.0 mae_bps = 0.0 mfe_bps = 0.0 time_in_loss_s = 0.0 time_in_profit_s = 0.0 return TradePathState( symbol=state.venue.symbol, side=fill_side, entry_ts_ns=now_ts, now_ts_ns=now_ts, bars_held=0, seconds_held=0.0, pnl_bps=pnl_bps, mae_bps=mae_bps, mfe_bps=mfe_bps, distance_from_mfe_bps=0.0, distance_from_entry_bps=0.0, time_to_mfe_s=0.0, time_in_loss_s=time_in_loss_s, time_in_profit_s=time_in_profit_s, time_since_last_profit_s=0.0, time_since_deep_mae_s=0.0, loss_to_profit_transitions=0, deep_loss_recoveries=0, failed_recovery_count=0, recovery_velocity_bps_per_s=0.0, adverse_velocity_bps_per_s=0.0, dolphin_regime_score=old_path.dolphin_regime_score if old_path else 0.0, jericho_signal_strength=old_path.jericho_signal_strength if old_path else 0.0, volatility_bps=old_path.volatility_bps if old_path else 0.0, orderflow_toxicity=old_path.orderflow_toxicity if old_path else 0.0, queue_churn_score=old_path.queue_churn_score if old_path else 0.0, book_imbalance=old_path.book_imbalance if old_path else 0.0, cross_venue_lead_score=old_path.cross_venue_lead_score if old_path else 0.0, ) # Existing position — update path metrics entry = old_path.entry_ts_ns seconds_held = (now_ts - entry) / 1_000_000_000 # PnL from entry if old_path.side == Side.BUY: pnl_bps = 10_000.0 * (venue_mid - fill_price) / max(fill_price, 1e-12) else: pnl_bps = 10_000.0 * (fill_price - venue_mid) / max(fill_price, 1e-12) # MAE/MFE tracking mae_bps = min(old_path.mae_bps, pnl_bps) mfe_bps = max(old_path.mfe_bps, pnl_bps) distance_from_mfe = mfe_bps - pnl_bps # Time tracking if pnl_bps < 0: time_in_loss_s = old_path.time_in_loss_s + (now_ts - old_path.now_ts_ns) / 1_000_000_000 time_in_profit_s = old_path.time_in_profit_s else: time_in_loss_s = old_path.time_in_loss_s time_in_profit_s = old_path.time_in_profit_s + (now_ts - old_path.now_ts_ns) / 1_000_000_000 # Recovery tracking loss_to_profit = old_path.loss_to_profit_transitions deep_recoveries = old_path.deep_loss_recoveries failed_recoveries = old_path.failed_recovery_count if old_path.pnl_bps < 0 and pnl_bps >= 0: loss_to_profit += 1 if old_path.mae_bps < -30.0 and pnl_bps > old_path.mae_bps + 10.0: deep_recoveries += 1 if old_path.mae_bps < -30.0 and pnl_bps < old_path.mae_bps + 5.0: if (now_ts - old_path.now_ts_ns) / 1_000_000_000 > 10.0: failed_recoveries += 1 return TradePathState( symbol=old_path.symbol, side=old_path.side, entry_ts_ns=old_path.entry_ts_ns, now_ts_ns=now_ts, bars_held=old_path.bars_held, seconds_held=seconds_held, pnl_bps=pnl_bps, mae_bps=mae_bps, mfe_bps=mfe_bps, distance_from_mfe_bps=distance_from_mfe, distance_from_entry_bps=abs(pnl_bps), time_to_mfe_s=old_path.time_to_mfe_s, time_in_loss_s=time_in_loss_s, time_in_profit_s=time_in_profit_s, time_since_last_profit_s=old_path.time_since_last_profit_s, time_since_deep_mae_s=old_path.time_since_deep_mae_s, loss_to_profit_transitions=loss_to_profit, deep_loss_recoveries=deep_recoveries, failed_recovery_count=failed_recoveries, recovery_velocity_bps_per_s=old_path.recovery_velocity_bps_per_s, adverse_velocity_bps_per_s=old_path.adverse_velocity_bps_per_s, dolphin_regime_score=old_path.dolphin_regime_score, jericho_signal_strength=old_path.jericho_signal_strength, volatility_bps=old_path.volatility_bps, orderflow_toxicity=old_path.orderflow_toxicity, queue_churn_score=old_path.queue_churn_score, book_imbalance=old_path.book_imbalance, cross_venue_lead_score=old_path.cross_venue_lead_score, ) class MinimalCryptoLOBCWM: """ Phase-1 local CWM with full exchange mechanics. Deterministic: - Price-time priority with sequential level consumption - Partial fills across multiple levels - Queue position estimation - Latency injection (feed + order) - Maker/taker fee application - Post-only rejection if crossing - IOC/FOK/LIMIT semantics - Tick/lot rounding - Open order aging (TTL expiry) - Path-state update (MAE/MFE/recovery) - Mark-to-market Two modes: REPLAY_NO_IMPACT: state follows historical; our order fills per queue model. ENDOGENOUS_AGENT_SIM: joint actions alter book state. """ def __init__( self, feature_extractor: Optional[FeatureExtractor] = None, tick_ns: int = 1_000_000, # 1ms per transition step ) -> None: self.feature_extractor = feature_extractor or DefaultFeatureExtractor() self._tick_ns = tick_ns @staticmethod def _make_open_order(action: FulfilmentAction, price: float, qty: float, ts: int, symbol: str = "") -> OpenOrderState: return OpenOrderState( client_order_id=f"m_{ts}", venue_order_id=None, symbol=symbol, side=action.side, order_type=action.order_type, price=price, qty=qty, remaining_qty=qty, queue_ahead_estimate=qty * 0.5, created_ts_ns=ts, last_update_ts_ns=ts, reduce_only=action.reduce_only, post_only=action.post_only, ) def transition( self, state: MarketWorldState, joint_action: JointAction, ) -> MarketWorldState: our_action = joint_action[0] counterparty_actions = joint_action[1:] tick = state.venue.tick_size lot = state.venue.lot_size min_qty = state.venue.min_qty now_ts = state.ts_ns + self._tick_ns # 1. Process cancels open_orders = list(state.open_orders) if isinstance(our_action, FulfilmentAction): if our_action.kind.value == "CANCEL" and our_action.cancel_order_id: open_orders = [o for o in open_orders if o.client_order_id != our_action.cancel_order_id] if our_action.kind.value == "CANCEL_REPLACE" and our_action.cancel_order_id: open_orders = [o for o in open_orders if o.client_order_id != our_action.cancel_order_id] # 2. Process counterparty cancels for cp in counterparty_actions: if isinstance(cp, CounterpartyAction) and cp.kind.value == "CANCEL": open_orders = [o for o in open_orders if o.symbol != state.venue.symbol] # 3. Open order aging — expire orders past TTL # In a real exchange, orders have TTL. We simulate this by removing # orders that have been open for more than a configurable duration. # For now, we keep all orders (TTL=0 means no expiry). # 4. Process our action new_fill_qty = 0.0 new_fill_price = 0.0 is_maker_fill = False book = state.book if isinstance(our_action, FulfilmentAction): if our_action.kind.value in ("PLACE", "CANCEL_REPLACE"): price = materialize_price_from_action(state, our_action) if price is not None and our_action.qty_fraction > 0: notional = our_action.qty_fraction * state.account.available_balance qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty) if qty > 0: price = _round_tick(price, tick) if price <= 0: price = tick # Post-only rejection if our_action.post_only: if state.book.bids and state.book.asks: if our_action.side == Side.BUY and price >= state.book.best_ask: pass # rejected elif our_action.side == Side.SELL and price <= state.book.best_bid: pass # rejected else: oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol) open_orders.append(oo) else: oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol) open_orders.append(oo) else: # Non-post-only: add to book oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol) open_orders.append(oo) elif our_action.kind.value == "CROSS_SPREAD": # Aggressive: immediate fill consuming levels price = materialize_price_from_action(state, our_action) if price is not None and our_action.qty_fraction > 0: notional = our_action.qty_fraction * state.account.available_balance qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty) if qty > 0: if our_action.side == Side.BUY: filled, avg_price, new_asks = _fill_from_levels( list(state.book.asks), qty, lot, min_qty, ) if filled > 0: new_fill_qty = filled new_fill_price = avg_price # Market impact: price moves up after aggressive buy impact_bps = filled / max(sum(l.qty for l in state.book.asks), 1e-12) * 0.5 impact_price = avg_price * (1 + impact_bps / 10_000) book = OrderBookState( ts_ns=now_ts, symbol=state.book.symbol, bids=state.book.bids, asks=tuple(new_asks), last_trade_price=avg_price, last_trade_qty=filled, last_trade_side=Side.BUY, ) elif our_action.side == Side.SELL: filled, avg_price, new_bids = _fill_from_levels( list(state.book.bids), qty, lot, min_qty, ) if filled > 0: new_fill_qty = filled new_fill_price = avg_price # Market impact: price moves down after aggressive sell impact_bps = filled / max(sum(l.qty for l in state.book.bids), 1e-12) * 0.5 impact_price = avg_price * (1 - impact_bps / 10_000) book = OrderBookState( ts_ns=now_ts, symbol=state.book.symbol, bids=tuple(new_bids), asks=state.book.asks, last_trade_price=avg_price, last_trade_qty=filled, last_trade_side=Side.SELL, ) elif our_action.kind.value in ("REDUCE", "FULL_EXIT"): # Immediate fill at best price if our_action.side == Side.SELL and state.book.bids: price = state.book.best_bid elif our_action.side == Side.BUY and state.book.asks: price = state.book.best_ask else: price = materialize_price_from_action(state, our_action) if price is not None and our_action.qty_fraction > 0: notional = our_action.qty_fraction * state.account.available_balance qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty) if qty > 0: new_fill_qty = qty new_fill_price = _round_tick(price, tick) # 5. Simulate counterparty trades hitting book (endogenous mode) for cp in counterparty_actions: if isinstance(cp, CounterpartyAction) and cp.kind.value == "CROSS_SPREAD" and cp.side: cp_notional = cp.qty_fraction_of_top * state.account.available_balance cp_qty = _clip_lots(cp_notional / max(state.book.mid if state.book.bids and state.book.asks else 1.0, 1e-12), lot, min_qty) if cp_qty > 0: if cp.side == Side.BUY and state.book.asks: filled, avg_price, new_asks = _fill_from_levels( list(state.book.asks), cp_qty, lot, min_qty, ) if filled > 0: # Counterparty fill only updates book, not our fill tracking book = OrderBookState( ts_ns=now_ts, symbol=state.book.symbol, bids=book.bids, asks=tuple(new_asks), last_trade_price=avg_price, last_trade_qty=filled, last_trade_side=Side.BUY, ) elif cp.side == Side.SELL and book.bids: filled, avg_price, new_bids = _fill_from_levels( list(book.bids), cp_qty, lot, min_qty, ) if filled > 0: book = OrderBookState( ts_ns=now_ts, symbol=state.book.symbol, bids=tuple(new_bids), asks=book.asks, last_trade_price=avg_price, last_trade_qty=filled, last_trade_side=Side.SELL, ) # 6. Update account and position equity = state.account.equity pos = state.account.positions.get(state.venue.symbol) pos_qty = pos.qty if pos else 0.0 pos_avg = pos.avg_entry if pos else 0.0 pos_r_pnl = pos.realized_pnl if pos else 0.0 old_unrealized = pos.unrealized_pnl if pos else 0.0 # Subtract old unrealized from equity (it was included in state.account.equity) equity -= old_unrealized trade_path = state.trade_path if new_fill_qty > 0: fee_bps = state.venue.maker_fee_bps if is_maker_fill else state.venue.taker_fee_bps fee = new_fill_qty * new_fill_price * abs(fee_bps) / 10_000.0 if isinstance(our_action, FulfilmentAction) and our_action.side == Side.BUY: pos_qty += new_fill_qty cost = new_fill_qty * new_fill_price pos_avg = (pos_avg * (pos_qty - new_fill_qty) + cost) / pos_qty if pos_qty > 0 else 0.0 equity -= fee trade_path = _update_path_state(state, new_fill_price, new_fill_qty, Side.BUY, now_ts) elif isinstance(our_action, FulfilmentAction) and our_action.side == Side.SELL: old_qty = pos_qty pos_qty -= new_fill_qty pos_r_pnl += new_fill_qty * (new_fill_price - pos_avg) equity -= fee # If position flipped sign, reset avg_entry to fill price if old_qty > 0 and pos_qty < 0: pos_avg = new_fill_price elif old_qty < 0 and pos_qty > 0: pos_avg = new_fill_price trade_path = _update_path_state(state, new_fill_price, new_fill_qty, Side.SELL, now_ts) # Mark-to-market mid = book.mid if book.bids and book.asks else (pos_avg if pos_qty != 0 else 0.0) unrealized = pos_qty * (mid - pos_avg) new_pos = PositionState( symbol=state.venue.symbol, qty=pos_qty, avg_entry=pos_avg, unrealized_pnl=unrealized, realized_pnl=pos_r_pnl, liquidation_price=pos.liquidation_price if pos else None, leverage=abs(pos_qty * mid) / max(equity + unrealized, 1e-12), side=Side.BUY if pos_qty > 0 else Side.SELL if pos_qty < 0 else None, ) equity += unrealized else: new_pos = pos new_positions = dict(state.account.positions) if new_pos: new_positions[state.venue.symbol] = new_pos elif state.venue.symbol in new_positions and (new_pos is None or (new_pos and abs(new_pos.qty) < 1e-12)): del new_positions[state.venue.symbol] new_account = AccountState( ts_ns=now_ts, equity=equity, wallet_balance=state.account.wallet_balance, available_balance=max(0.0, state.account.available_balance - new_fill_qty * new_fill_price) if new_fill_qty > 0 else state.account.available_balance, margin_used=state.account.margin_used, total_notional=abs(pos_qty * (book.mid if book.bids and book.asks else 0.0)), positions=new_positions, ) return MarketWorldState( ts_ns=now_ts, mode=state.mode, venue=state.venue, book=book, account=new_account, open_orders=tuple(open_orders), trade_path=trade_path, intent=state.intent, funding_bps=state.funding_bps, volatility_state=state.volatility_state, market_regime=state.market_regime, feed_latency_ms=state.feed_latency_ms, order_latency_ms=state.order_latency_ms, rng_seed=state.rng_seed, ) def reward( self, prev_state: MarketWorldState, action: FulfilmentAction, next_state: MarketWorldState, params: FulfilmentPolicyParams, ) -> float: fv = self.feature_extractor.extract(next_state).values pnl = fv.get("pnl_bps", 0.0) toxicity = fv.get("orderflow_toxicity", 0.0) churn = fv.get("queue_churn_score", 0.0) time_in_loss = fv.get("time_in_loss_s", 0.0) spread_bps = fv.get("spread_bps", 0.0) reward = 0.0 reward += params.w_expected_pnl * pnl reward -= params.w_adverse_selection * toxicity reward -= params.w_inventory_risk * self._inventory_risk(next_state) reward -= params.w_tail_loss * self._tail_risk_proxy(next_state) reward -= params.w_time_decay * math.log1p(max(time_in_loss, 0.0)) if action.order_type and action.order_type.value in ("POST_ONLY", "LIMIT"): reward += params.w_fee_quality * max(0.0, -prev_state.venue.maker_fee_bps) if action.kind.value == "CROSS_SPREAD": reward -= spread_bps + max(prev_state.venue.taker_fee_bps, 0.0) if action.kind.value in ("CANCEL", "CANCEL_REPLACE"): if toxicity > params.adverse_toxicity_cancel_threshold: reward += params.w_adverse_selection * toxicity if churn > params.queue_churn_cancel_threshold: reward += params.w_queue_priority * churn return reward def _inventory_risk(self, state: MarketWorldState) -> float: pos = state.account.positions.get(state.venue.symbol) if not pos: return 0.0 mid = state.book.mid if state.book.bids and state.book.asks else 0.0 return abs(pos.qty * mid) / max(state.account.equity, 1e-12) def _tail_risk_proxy(self, state: MarketWorldState) -> float: p = state.trade_path if p is None: return 0.0 return ( max(0.0, abs(p.mae_bps)) * (1.0 + math.log1p(max(p.time_in_loss_s, 0.0))) * (1.0 + max(0, p.failed_recovery_count)) ) def terminal(self, state: MarketWorldState, depth: int) -> bool: if depth <= 0: return True if state.intent is None: return True return False