import asyncio import logging from collections.abc import Awaitable, Callable from app.bars.models import Bar, Timeframe from app.market.base import MarketDataSource logger = logging.getLogger(__name__) BarHandler = Callable[[Bar], Awaitable[None]] class StreamService: def __init__(self, source: MarketDataSource, symbol: str): self.source = source self.symbol = symbol self.status = "disconnected" self.last_error: str | None = None self.last_bar_t: int | None = None self._handlers: list[BarHandler] = [] self._stop = asyncio.Event() def add_handler(self, handler: BarHandler) -> None: self._handlers.append(handler) async def seed( self, source: MarketDataSource | None, tf: Timeframe, range_: str, symbol: str | None = None, ) -> None: # The seed source names the instrument differently from the live one: # Yahoo says ES=F where Schwab says /ES. if source is None or not source.supports_history(): return bars = await source.history(symbol or self.symbol, tf, None, None, range_=range_) for bar in bars: await self._emit(bar) async def _emit(self, bar: Bar) -> None: self.last_bar_t = max(self.last_bar_t or bar.t, bar.t) for handler in self._handlers: await handler(bar) async def run(self) -> None: while not self._stop.is_set(): try: async for bar in self.source.stream(self.symbol): self.status = "replay" if self.source.name == "replay" else "connected" self.last_error = None await self._emit(bar) if self._stop.is_set(): break if self.source.name == "replay": return except asyncio.CancelledError: raise except Exception as exc: self.last_error = str(exc) logger.exception("Market stream failed; reconnecting") self.status = "disconnected" try: await asyncio.wait_for(self._stop.wait(), timeout=5) except TimeoutError: pass def stop(self) -> None: self._stop.set() self.status = "disconnected"