""" Central treasury — position/capital limits, circuit breakers, PnL stops. Single source of truth for all risk constraints in live trading. Integrates with sim/constraints.py for the constraint logic and adds live-specific bookkeeping. """ from __future__ import annotations import time from typing import Optional from sim.constraints import ( InventoryConstraint, FundingConstraint, FeeSchedule, LiquidationRisk, CircuitBreaker, ) class Treasury: """Central risk and capital management for live trading. Tracks: - Current positions per asset - Realized and unrealized PnL - Daily trade counts - Circuit breaker state - Fee budget consumption Usage: treasury = Treasury(initial_equity=10000.0) ok = treasury.can_open("BTC", side="buy", size=0.001, mark_price=50000.0) treasury.record_fill("BTC", side="buy", size=0.001, price=50000.0, fee=10.0) """ def __init__( self, initial_equity: float = 10000.0, max_position_per_asset: float = 0.005, max_net_exposure: float = 0.01, max_daily_trades: int = 500, max_drawdown_pct: float = -10.0, max_toxic_rate: float = 0.4, cooldown_seconds: float = 300.0, maker_fee_pct: float = 0.0002, taker_fee_pct: float = 0.0005, maintenance_margin_pct: float = 0.03, ): self._initial_equity = initial_equity self._realized_pnl: float = 0.0 self._fees_paid: float = 0.0 self._daily_trades: int = 0 self._toxic_fills: int = 0 self._api_errors: int = 0 # Positions tracked as {coin: {"side": "long"|"short", "size": float, "entry_px": float}} self._positions: dict[str, dict] = {} self._mark_prices: dict[str, float] = {} self._circuit_breaker = CircuitBreaker( max_drawdown_pct=max_drawdown_pct, max_daily_trades=max_daily_trades, max_toxic_rate=max_toxic_rate, cooldown_seconds=cooldown_seconds, ) self._inventory = InventoryConstraint( max_long=max_position_per_asset, max_short=max_position_per_asset, max_net_exposure=max_net_exposure, ) self._fees = FeeSchedule(maker_fee_pct=maker_fee_pct, taker_fee_pct=taker_fee_pct) self._funding = FundingConstraint() self._liquidation = LiquidationRisk(maintenance_margin_pct=maintenance_margin_pct) self._halted: bool = False self._halt_reason: str = "" self._halted_at: float = 0.0 self._session_start: float = time.time() # ── Position management ────────────────────────────────── def can_open(self, coin: str, side: str, size: float, mark_price: float) -> dict: """Check whether a new position can be opened. Returns {allowed: bool, reason: str, fee_estimate: float} """ if self._halted: return {"allowed": False, "reason": self._halt_reason, "fee_estimate": 0.0} pos = self._positions.get(coin.upper(), {}) current_size = pos.get("size", 0.0) if pos.get("side") == side else -(pos.get("size", 0.0)) new_size = current_size + size limits = self._inventory.check( max(0.0, new_size) if side == "buy" else max(0.0, current_size), max(0.0, -new_size) if side == "sell" else max(0.0, -current_size), ) if not limits["long_ok"]: return {"allowed": False, "reason": "long limit exceeded", "fee_estimate": 0.0} if not limits["short_ok"]: return {"allowed": False, "reason": "short limit exceeded", "fee_estimate": 0.0} fee = self._fees.maker_fee(size * mark_price) return {"allowed": True, "reason": "ok", "fee_estimate": round(fee, 6)} def record_fill(self, coin: str, side: str, size: float, price: float, fee: float, pnl: float = 0.0): """Record a filled trade.""" c = coin.upper() pos = self._positions.get(c) is_close = pos and pos.get("side") != side if is_close: self._realized_pnl += pnl pos["size"] -= size if pos["size"] <= 1e-10: del self._positions[c] else: if not pos: self._positions[c] = {"side": side, "size": size, "entry_px": price} else: total = pos["size"] + size pos["entry_px"] = (pos["entry_px"] * pos["size"] + price * size) / total if total > 0 else price pos["size"] = total self._fees_paid += fee self._daily_trades += 1 self._check_breakers() def record_toxic_fill(self): self._toxic_fills += 1 def record_api_error(self): self._api_errors += 1 def update_mark_price(self, coin: str, price: float): self._mark_prices[coin.upper()] = price # ── Position queries ───────────────────────────────────── def position(self, coin: str) -> float: """Signed position (positive = long).""" pos = self._positions.get(coin.upper(), {}) raw = pos.get("size", 0.0) return raw if pos.get("side") == "buy" else -raw def position_size(self, coin: str) -> float: """Absolute position size.""" return abs(self.position(coin)) @property def all_positions(self) -> dict[str, float]: return {c: self.position(c) for c in self._positions} @property def net_exposure(self) -> float: return sum(abs(p) for p in self.all_positions.values()) # ── PnL ────────────────────────────────────────────────── def unrealized_pnl(self) -> float: pnl = 0.0 for coin, pos in self._positions.items(): mark = self._mark_prices.get(coin, pos.get("entry_px", 0)) if pos["side"] == "buy": pnl += pos["size"] * (mark - pos["entry_px"]) else: pnl += pos["size"] * (pos["entry_px"] - mark) return round(pnl, 4) def total_pnl(self) -> float: return self._realized_pnl + self.unrealized_pnl() - self._fees_paid def pnl_pct(self) -> float: return self.total_pnl() / self._initial_equity * 100 if self._initial_equity > 0 else 0 @property def equity(self) -> float: return self._initial_equity + self.total_pnl() # ── Liquidation risk ───────────────────────────────────── def liquidation_distance(self, coin: str) -> float: """Percentage distance to liquidation.""" pos = self._positions.get(coin.upper()) if not pos: return float("inf") mark = self._mark_prices.get(coin.upper(), pos["entry_px"]) liq = self._liquidation.liquidation_price( entry_price=pos["entry_px"], size=pos["size"], position_side=pos["side"], wallet_balance=self.equity, ) return self._liquidation.distance_to_liquidation_pct(mark, liq, pos["side"]) def is_liquidation_safe(self, coin: str, threshold_pct: float = 5.0) -> bool: return self.liquidation_distance(coin) >= threshold_pct # ── Circuit breaker ────────────────────────────────────── def _check_breakers(self): if self._halted: return toxic_rate = self._toxic_fills / max(self._daily_trades, 1) result = self._circuit_breaker.evaluate({ "pnl_pct": round(self.pnl_pct(), 2), "daily_trades": self._daily_trades, "toxic_rate": toxic_rate, "api_errors": self._api_errors, }) if result.get("tripped"): self._halted = True self._halt_reason = result.get("reason", "unknown") self._halted_at = time.time() def is_halted(self) -> bool: if self._halted: elapsed = time.time() - self._halted_at if elapsed > self._circuit_breaker.cooldown_seconds: self._halted = False self._halt_reason = "" self._daily_trades = 0 self._toxic_fills = 0 return self._halted @property def halt_reason(self) -> str: return self._halt_reason # ── Stats ──────────────────────────────────────────────── def summary(self) -> dict: return { "equity": round(self.equity, 2), "realized_pnl": round(self._realized_pnl, 4), "unrealized_pnl": self.unrealized_pnl(), "total_pnl": self.total_pnl(), "pnl_pct": round(self.pnl_pct(), 2), "fees_paid": round(self._fees_paid, 4), "daily_trades": self._daily_trades, "toxic_fills": self._toxic_fills, "api_errors": self._api_errors, "positions": {c: round(v, 6) for c, v in self.all_positions.items()}, "net_exposure": round(self.net_exposure, 6), "halted": self._halted, "uptime_hours": round((time.time() - self._session_start) / 3600, 1), }