Download execution/broker.py from raghava4u/Trading-Bot-M20: direct link, hf CLI and curl.
- Browser
- Download file 20.7 kB
-
https://huggingface.co/raghava4u/Trading-Bot-M20/resolve/main/execution/broker.py
- Command line
-
hf download hf://raghava4u/Trading-Bot-M20/execution/broker.py
-
curl -L -o broker.py https://huggingface.co/raghava4u/Trading-Bot-M20/resolve/main/execution/broker.py
20.7 kB
| """ | |
| 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 | |
| def is_connected(self) -> bool: | |
| return self._connected | |