60 lines
2 KiB
Python
60 lines
2 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_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
|
|
) -> None:
|
|
if source is None or not source.supports_history():
|
|
return
|
|
bars = await source.history(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:
|
|
self.status = "replay" if self.source.name == "replay" else "connected"
|
|
async for bar in self.source.stream(self.symbol):
|
|
await self._emit(bar)
|
|
if self._stop.is_set():
|
|
break
|
|
if self.source.name == "replay":
|
|
return
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception:
|
|
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"
|