raghava4u's picture
Upload folder using huggingface_hub
d53dc44 verified
Raw History Blame Contribute Delete
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
@property
def is_connected(self) -> bool:
return self._connected