73 lines
2.5 KiB
Python
73 lines
2.5 KiB
Python
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()
|
|
self.on_drop = None
|
|
|
|
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:
|
|
was_up = self.status == "connected"
|
|
self.last_error = str(exc)
|
|
logger.exception("Market stream failed; reconnecting")
|
|
if was_up and self.on_drop:
|
|
self.on_drop(str(exc))
|
|
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"
|