Initial working version, positions show.

This commit is contained in:
2026-04-01 19:17:09 -05:00
commit 28afa137cb
7 changed files with 398 additions and 0 deletions

233
backend/ib_client.py Normal file
View File

@@ -0,0 +1,233 @@
"""Lightweight IB client wrapper using ibapi.
This wraps a minimal subset of EClient/EWrapper to subscribe to market
data and track positions in-memory. It's designed to be used by a
small Flask backend that serves a browser dashboard.
Notes:
- Keep behavior simple: subscribe to market data for provided symbols.
- Store latest quotes in `self.quotes` and positions in `self.positions`.
"""
import logging
import threading
import time
from typing import Dict, List
from ibapi.client import EClient
from ibapi.contract import Contract
from ibapi.wrapper import EWrapper
class IBClient(EWrapper, EClient):
def __init__(self, host: str = "127.0.0.1", port: int = 7496, client_id: int = 1):
EClient.__init__(self, self)
# normalize/validate connection params
if not host:
host = "127.0.0.1"
if not isinstance(port, int):
try:
port = int(port)
except Exception:
port = 7496
self.host = host
self.port = port
self.client_id = client_id
self._thread = None
self._next_req_id = 1
# connection sync
self._connected_event = threading.Event()
# subscriptions requested before connection
self._pending_subscriptions = set()
# in-memory stores
self.quotes: Dict[str, Dict] = {}
self.positions: List[Dict] = []
# tracking for subscriptions: reqId -> symbol
self._req_map: Dict[int, str] = {}
self._requested_symbols = set()
# basic logging
logging.basicConfig(level=logging.INFO)
def connect_and_start(self) -> None:
if not self.isConnected():
logging.info("Connecting to IB gateway %s:%s", self.host, self.port)
self.connect(self.host, self.port, self.client_id)
self._thread = threading.Thread(target=self.run, daemon=True)
self._thread.start()
# don't block here; nextValidId will mark connection ready
def disconnect_and_stop(self) -> None:
try:
if self.isConnected():
self.disconnect()
except Exception:
pass
# ---- EWrapper overrides (minimal) ----
def nextValidId(self, orderId: int):
super().nextValidId(orderId)
# ensure we start with a reasonably large req id to avoid collisions
if orderId and orderId > self._next_req_id:
self._next_req_id = orderId
# mark connection as ready
try:
self._connected_event.set()
except Exception:
pass
# process any pending subscriptions
if self._pending_subscriptions:
pending = list(self._pending_subscriptions)
self._pending_subscriptions.clear()
for s in pending:
try:
self.subscribe_market_data(s)
except Exception:
logging.exception("Failed to process pending subscription %s", s)
def error(self, reqId, errorCode, errorString):
logging.error("IB error (req=%s code=%s): %s", reqId, errorCode, errorString)
def tickPrice(self, reqId, tickType, price, attrib):
# store last price when available
if price is None:
return
symbol = self._req_map.get(reqId)
if not symbol:
return
data = self.quotes.setdefault(symbol, {})
data["last"] = price
def tickSize(self, reqId, tickType, size):
symbol = self._req_map.get(reqId)
if not symbol:
return
data = self.quotes.setdefault(symbol, {})
data["size"] = size
def updatePortfolio(self, contract, position, marketPrice, marketValue, averageCost, unrealizedPNL, realizedPNL, accountName):
# simple representation of the position
if not contract or not contract.symbol:
return
pos = {
"symbol": contract.symbol,
"position": position,
"marketPrice": marketPrice,
"marketValue": marketValue,
}
# replace existing entry for the symbol
for i, p in enumerate(self.positions):
if p.get("symbol") == pos["symbol"]:
self.positions[i] = pos
break
else:
self.positions.append(pos)
def positionEnd(self):
# called when a batch of position updates is finished
logging.info("Position update batch complete. %d positions stored", len(self.positions))
def position(self, account, contract, position, avgCost):
"""Handle positions returned by `reqPositions()`.
Signature from EWrapper: position(self, account, contract, position, avgCost)
"""
try:
if not contract or not getattr(contract, 'symbol', None):
return
pos = {
"symbol": contract.symbol,
"position": position,
"averageCost": avgCost,
"account": account,
}
# replace existing entry for the symbol
for i, p in enumerate(self.positions):
if p.get("symbol") == pos["symbol"] and p.get("account") == account:
self.positions[i] = pos
break
else:
self.positions.append(pos)
except Exception:
logging.exception("Error handling position callback")
# ---- helpers ----
def _new_req_id(self) -> int:
rid = self._next_req_id
self._next_req_id += 1
return rid
def _make_stock_contract(self, symbol: str, exchange: str = "SMART", currency: str = "USD") -> Contract:
c = Contract()
c.symbol = symbol
c.secType = "STK"
c.exchange = exchange
c.currency = currency
return c
def subscribe_market_data(self, symbol: str) -> None:
"""Subscribe to basic market data for `symbol` (in USD on SMART exchange).
Repeated calls for the same symbol are ignored.
"""
if symbol in self._requested_symbols:
return
# if not connected yet, queue the subscription
if not self._connected_event.is_set():
logging.info("Not connected yet; queuing subscription for %s", symbol)
self._pending_subscriptions.add(symbol)
# ensure we start connection
try:
self.connect_and_start()
except Exception:
logging.exception("Failed to start connection while queuing subscription")
return
req_id = self._new_req_id()
self._req_map[req_id] = symbol
contract = self._make_stock_contract(symbol)
# request real-time market data (empty generic tick list)
try:
self.reqMarketDataType(1)
self.reqMktData(req_id, contract, "", False, False, [])
self._requested_symbols.add(symbol)
logging.info("Subscribed to market data for %s (req=%s)", symbol, req_id)
except Exception as e:
logging.exception("Failed to subscribe to %s: %s", symbol, e)
def request_positions(self) -> None:
"""Ask the IB API to send current positions (triggers updatePortfolio callbacks)."""
if not self._connected_event.is_set():
logging.info("Not connected yet; will request positions only if connection parameters are valid")
# validate host/port before trying to connect
if not isinstance(self.host, str) or not isinstance(self.port, int):
logging.error("Invalid connection parameters: host=%r port=%r", self.host, self.port)
return
try:
self.connect_and_start()
except Exception:
logging.exception("Failed to start connection for positions request")
# schedule a delayed reqPositions after connection is established
def _delayed():
if self._connected_event.wait(5):
try:
self.reqPositions()
except Exception:
logging.exception("Failed to request positions after connect")
threading.Thread(target=_delayed, daemon=True).start()
return
try:
self.reqPositions()
except Exception:
logging.exception("Failed to request positions")
def get_quotes(self) -> Dict[str, Dict]:
return self.quotes
def get_positions(self) -> List[Dict]:
return self.positions