""" execution/broker.py — Alpaca REST + WebSocket wrapper. Token-bucket rate limiting, SAFE MODE, reconnection with back-off. """ from __future__ import annotations import collections import datetime import json import logging import random import threading import time from typing import Callable import requests import websocket import config from contracts import BrokerError, SafeModeError from data import storage from data.downloader import TokenBucket, download_missing_bars logger = logging.getLogger("trading_system.broker") # ── Rate Limiter ───────────────────────────────────────────────────────────── _rate_limiter = TokenBucket(max_tokens=200, interval_sec=60.0) # ── Latency Tracking ──────────────────────────────────────────────────────── _latencies: collections.deque = collections.deque(maxlen=100) # ── SAFE MODE ──────────────────────────────────────────────────────────────── safe_mode_active = False _safe_mode_lock = threading.Lock() consecutive_failures = 0 def enter_safe_mode(reason: str, alert_callback: Callable[[str], None] | None = None): """Enter SAFE MODE: cancel all orders, block new submissions.""" global safe_mode_active with _safe_mode_lock: if safe_mode_active: return safe_mode_active = True logger.critical("SAFE MODE ACTIVATED: %s", reason) try: cancel_all_orders() except Exception as e: logger.error("Failed to cancel orders on safe mode: %s", e) if alert_callback: alert_callback(f"🚨 SAFE MODE ACTIVATED: {reason}") def clear_safe_mode(): """Clear SAFE MODE (requires manual confirmation).""" global safe_mode_active, consecutive_failures with _safe_mode_lock: safe_mode_active = False consecutive_failures = 0 logger.warning("SAFE MODE cleared manually") def _check_safe_mode(): if safe_mode_active: raise SafeModeError("Operation blocked: SAFE MODE is active") # ── REST API helpers ───────────────────────────────────────────────────────── def _headers() -> dict: return { "APCA-API-KEY-ID": config.ALPACA_API_KEY, "APCA-API-SECRET-KEY": config.ALPACA_SECRET_KEY, "Content-Type": "application/json", } def _request( method: str, url: str, params: dict | None = None, json_data: dict | None = None, is_data_endpoint: bool = False, ) -> dict | list | None: """Make an authenticated Alpaca API request with rate limiting and retries.""" global consecutive_failures if is_data_endpoint: _rate_limiter.acquire() for attempt in range(6): start_t = time.monotonic() try: resp = requests.request( method, url, headers=_headers(), params=params, json=json_data, timeout=30, ) latency_ms = (time.monotonic() - start_t) * 1000 _latencies.append(latency_ms) logger.debug( "API %s %s → %d (%.0fms)", method, url.split("/")[-1], resp.status_code, latency_ms, ) # Check p95 latency if len(_latencies) >= 20: sorted_lat = sorted(_latencies) p95 = sorted_lat[int(len(sorted_lat) * 0.95)] if p95 > 2000: logger.warning("API p95 latency %.0fms > 2000ms threshold", p95) if resp.status_code in (200, 201, 204): consecutive_failures = 0 if resp.status_code == 204: return None return resp.json() elif resp.status_code == 429: wait = min(2 ** attempt + random.random(), 60) logger.warning("HTTP 429, retry %d in %.1fs", attempt + 1, wait) time.sleep(wait) continue elif resp.status_code >= 500: consecutive_failures += 1 logger.error("HTTP %d: %s", resp.status_code, resp.text[:200]) if consecutive_failures >= 3: enter_safe_mode(f"3 consecutive API failures (last: HTTP {resp.status_code})") raise BrokerError(f"SAFE MODE triggered after HTTP {resp.status_code}") wait = min(2 ** attempt + random.random(), 30) time.sleep(wait) continue elif resp.status_code == 422: logger.error("HTTP 422 (Unprocessable): %s", resp.text[:500]) return {"error": resp.text, "status_code": 422} else: logger.error("HTTP %d: %s", resp.status_code, resp.text[:200]) resp.raise_for_status() except requests.exceptions.Timeout: consecutive_failures += 1 wait = min(2 ** attempt + random.random(), 30) logger.warning("Request timeout (attempt %d), retry in %.1fs", attempt + 1, wait) time.sleep(wait) except SafeModeError: raise except requests.exceptions.RequestException as e: consecutive_failures += 1 logger.error("Request error: %s", e) if consecutive_failures >= 3: enter_safe_mode(f"3 consecutive failures: {e}") raise BrokerError(str(e)) raise BrokerError(f"Failed after 6 retries: {method} {url}") _cached_account = None _cached_account_time = 0.0 def get_account() -> dict: global _cached_account, _cached_account_time now = time.monotonic() if _cached_account and (now - _cached_account_time) < 30.0: return _cached_account url = f"{config.ALPACA_BASE_URL}/v2/account" _cached_account = _request("GET", url) _cached_account_time = time.monotonic() return _cached_account def get_buying_power() -> float: acct = get_account() return float(acct.get("buying_power", 0)) def get_equity() -> float: acct = get_account() return float(acct.get("equity", 0)) def get_cash() -> float: acct = get_account() return float(acct.get("cash", 0)) # ── Orders ─────────────────────────────────────────────────────────────────── def submit_limit_order( symbol: str, side: str, qty: float, limit_price: float, time_in_force: str = "day", ) -> dict: """Submit a limit order.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "limit", "time_in_force": time_in_force, "limit_price": str(round(limit_price, 2)), } logger.info("Submitting limit order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_market_order( symbol: str, side: str, qty: float, time_in_force: str = "day", ) -> dict: """Submit a market order.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "market", "time_in_force": time_in_force, } logger.info("Submitting market order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_bracket_order( symbol: str, side: str, qty: int, limit_price: float, stop_price: float, tp_price: float, time_in_force: str = "day", is_market: bool = False, ) -> dict: """Submit a bracket order (entry + stop + take-profit). Whole qty only.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "market" if is_market else "limit", "time_in_force": time_in_force, "order_class": "bracket", "stop_loss": {"stop_price": str(round(stop_price, 2))}, "take_profit": {"limit_price": str(round(tp_price, 2))}, } if not is_market: order["limit_price"] = str(round(limit_price, 2)) logger.info("Submitting bracket order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_notional_order( symbol: str, side: str, notional: float, time_in_force: str = "day", ) -> dict: """Submit a notional (dollar amount) order for fractional shares.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "notional": str(round(notional, 2)), "side": side, "type": "market", "time_in_force": time_in_force, } logger.info("Submitting notional order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_stop_order( symbol: str, side: str, qty: float, stop_price: float, time_in_force: str = "gtc", ) -> dict: """Submit a stop order.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "stop", "time_in_force": time_in_force, "stop_price": str(round(stop_price, 2)), } logger.info("Submitting stop order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_trailing_stop_order( symbol: str, side: str, qty: float, trail_percent: float, time_in_force: str = "gtc", ) -> dict: """Submit a dynamic volatility trailing stop order.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "trailing_stop", "time_in_force": time_in_force, "trail_percent": str(round(trail_percent, 2)), } logger.info("Submitting dynamic trailing stop order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def submit_limit_tp_order( symbol: str, side: str, qty: float, limit_price: float, time_in_force: str = "gtc", ) -> dict: """Submit a take-profit limit order.""" _check_safe_mode() url = f"{config.ALPACA_BASE_URL}/v2/orders" order = { "symbol": symbol, "qty": str(qty), "side": side, "type": "limit", "time_in_force": time_in_force, "limit_price": str(round(limit_price, 2)), } logger.info("Submitting TP limit order: %s", json.dumps(order)) return _request("POST", url, json_data=order) def get_order(order_id: str) -> dict: url = f"{config.ALPACA_BASE_URL}/v2/orders/{order_id}" return _request("GET", url) def cancel_order(order_id: str) -> None: url = f"{config.ALPACA_BASE_URL}/v2/orders/{order_id}" _request("DELETE", url) logger.info("Cancelled order %s", order_id) def cancel_all_orders() -> None: url = f"{config.ALPACA_BASE_URL}/v2/orders" _request("DELETE", url) logger.warning("Cancelled ALL open orders") def list_orders(status: str = "open") -> list[dict]: url = f"{config.ALPACA_BASE_URL}/v2/orders" result = _request("GET", url, params={"status": status}) return result if isinstance(result, list) else [] # ── Positions ──────────────────────────────────────────────────────────────── def list_positions() -> list[dict]: url = f"{config.ALPACA_BASE_URL}/v2/positions" result = _request("GET", url) return result if isinstance(result, list) else [] def close_position(symbol: str) -> dict: """Close entire position for a symbol.""" url = f"{config.ALPACA_BASE_URL}/v2/positions/{symbol}" logger.warning("Closing position: %s", symbol) return _request("DELETE", url) def close_all_positions() -> None: url = f"{config.ALPACA_BASE_URL}/v2/positions" _request("DELETE", url) logger.warning("Closed ALL positions") # ── Quotes ─────────────────────────────────────────────────────────────────── def get_latest_quote(symbol: str) -> dict: url = f"{config.ALPACA_DATA_URL}/v2/stocks/{symbol}/quotes/latest" return _request("GET", url, is_data_endpoint=True) def get_latest_trade(symbol: str) -> dict: url = f"{config.ALPACA_DATA_URL}/v2/stocks/{symbol}/trades/latest" return _request("GET", url, is_data_endpoint=True) # ── Asset Info ─────────────────────────────────────────────────────────────── def get_asset(symbol: str) -> dict: url = f"{config.ALPACA_BASE_URL}/v2/assets/{symbol}" return _request("GET", url) def is_fractionable(symbol: str) -> bool: asset = get_asset(symbol) return asset.get("fractionable", False) if asset else False # ── WebSocket Stream ───────────────────────────────────────────────────────── class AlpacaWebSocket: """WebSocket stream for real-time bar updates with reconnection logic.""" def __init__( self, symbols: list[str], on_bar: Callable[[dict], None] | None = None, shutdown_event: threading.Event | None = None, alert_callback: Callable[[str], None] | None = None, ): self.symbols = symbols self.on_bar = on_bar self.shutdown_event = shutdown_event or threading.Event() self.alert_callback = alert_callback self._ws: websocket.WebSocketApp | None = None self._thread: threading.Thread | None = None self._connected = False self._reconnect_attempt = 0 self._last_message_time: float | None = None self._disconnect_time: float | None = None self._seen_timestamps: dict[str, set] = {s: set() for s in symbols} def _get_ws_url(self) -> str: # User requested to use FREE IEX data feed unconditionally return "wss://stream.data.alpaca.markets/v2/iex" def _on_open(self, ws): auth_msg = { "action": "auth", "key": config.ALPACA_API_KEY, "secret": config.ALPACA_SECRET_KEY, } ws.send(json.dumps(auth_msg)) def _on_message(self, ws, message): self._last_message_time = time.monotonic() self._reconnect_attempt = 0 data = json.loads(message) if not isinstance(data, list): data = [data] for msg in data: msg_type = msg.get("T") if msg_type == "success": if msg.get("msg") == "authenticated": self._connected = True sub_msg = { "action": "subscribe", "bars": self.symbols, } ws.send(json.dumps(sub_msg)) logger.info("WebSocket authenticated, subscribed to %s", self.symbols) # Recover missed bars on reconnect if self._disconnect_time: disconnect_duration = time.monotonic() - self._disconnect_time logger.info( "Reconnected after %.1fs disconnect", disconnect_duration ) self._recover_missed_bars() self._disconnect_time = None elif msg_type == "b": # bar symbol = msg.get("S", "") ts = msg.get("t", "") # Deduplication if ts in self._seen_timestamps.get(symbol, set()): continue if storage.bar_exists(symbol, "5Min", ts): continue if symbol in self._seen_timestamps: self._seen_timestamps[symbol].add(ts) # Limit set size if len(self._seen_timestamps[symbol]) > 10000: self._seen_timestamps[symbol] = set( list(self._seen_timestamps[symbol])[-5000:] ) bar = { "timestamp": ts, "open": msg.get("o"), "high": msg.get("h"), "low": msg.get("l"), "close": msg.get("c"), "volume": msg.get("v"), "vwap": msg.get("vw"), "adjusted": True, } storage.insert_bars(symbol, "5Min", [bar]) if self.on_bar: self.on_bar({"symbol": symbol, **bar}) def _on_error(self, ws, error): logger.error("WebSocket error: %s", error) def _on_close(self, ws, close_status, close_msg): self._connected = False self._disconnect_time = time.monotonic() logger.warning("WebSocket closed: %s %s", close_status, close_msg) def _recover_missed_bars(self): """Fetch bars missed during disconnection.""" for symbol in self.symbols: last_ts = storage.get_last_timestamp(symbol, "5Min") if last_ts: try: count = download_missing_bars(symbol, "5Min", last_ts) if count > 0: logger.info("Recovered %d bars for %s", count, symbol) except Exception as e: logger.error("Failed to recover bars for %s: %s", symbol, e) def start(self): """Start WebSocket in a daemon thread with reconnection loop.""" self._thread = threading.Thread( target=self._run_loop, daemon=True, name="WebSocketStream" ) self._thread.start() def _run_loop(self): while not self.shutdown_event.is_set(): try: self._ws = websocket.WebSocketApp( self._get_ws_url(), on_open=self._on_open, on_message=self._on_message, on_error=self._on_error, on_close=self._on_close, ) self._ws.run_forever(ping_interval=30, ping_timeout=10) except Exception as e: logger.error("WebSocket exception: %s", e) if self.shutdown_event.is_set(): break # Exponential back-off for reconnection wait = min(2 ** self._reconnect_attempt + random.random(), 120) self._reconnect_attempt += 1 logger.warning( "WebSocket reconnecting in %.1fs (attempt %d)", wait, self._reconnect_attempt, ) # Alert if disconnected > 5 minutes if self._disconnect_time: disc_duration = time.monotonic() - self._disconnect_time if disc_duration > 300 and self.alert_callback: self.alert_callback( f"⚠️ WebSocket disconnected for {disc_duration:.0f}s" ) self.shutdown_event.wait(timeout=wait) def stop(self): if self._ws: self._ws.close() self._connected = False @property def is_connected(self) -> bool: return self._connected