from __future__ import annotations import asyncio import json from collections.abc import Callable, Coroutine from typing import Any from yuxi.utils.logging_config import logger try: import websockets except ImportError: websockets = None class YuanbaoMonitor: def __init__( self, ws_url: str, token_provider: Callable[[], Coroutine[Any, Any, str]], on_event: Callable[[dict[str, Any]], Coroutine[Any, Any, None]], max_reconnect: int = 10, ping_interval: float = 30, ping_timeout: float = 10, close_timeout: float = 5, auth_timeout: float = 10, ): self._ws_url = ws_url self._token_provider = token_provider self._on_event = on_event self._max_reconnect = max_reconnect self._ping_interval = ping_interval self._ping_timeout = ping_timeout self._close_timeout = close_timeout self._auth_timeout = auth_timeout self._ws = None self._running = False self._reconnect_delay = 1.0 self._max_reconnect_delay = 60.0 self._task: asyncio.Task | None = None self._reconnect_count = 0 self._auth_failed = False self._last_event_timestamp: float | None = None async def start(self) -> None: if websockets is None: logger.warning("[Yuanbao] websockets library not installed, skipping WebSocket") return self._running = True self._reconnect_delay = 1.0 self._reconnect_count = 0 self._auth_failed = False self._task = asyncio.create_task(self._listen_loop()) logger.info("[Yuanbao] WebSocket monitor started") async def stop(self) -> None: self._running = False if self._task and not self._task.done(): self._task.cancel() try: await self._task except asyncio.CancelledError: pass self._task = None if self._ws: await self._ws.close() self._ws = None logger.info("[Yuanbao] WebSocket monitor stopped") @property def is_connected(self) -> bool: return self._ws is not None and self._running @property def reconnect_count(self) -> int: return self._reconnect_count @property def healthy(self) -> bool: if not self.is_connected: return False if self._last_event_timestamp is None: return True import time idle_seconds = time.monotonic() - self._last_event_timestamp return idle_seconds < self._ping_interval * 3 @property def last_event_timestamp(self) -> float | None: return self._last_event_timestamp async def _listen_loop(self) -> None: while self._running: try: token = await self._token_provider() headers = {"Authorization": f"Bearer {token}"} async for ws in websockets.connect( self._ws_url, additional_headers=headers, ping_interval=self._ping_interval, ping_timeout=self._ping_timeout, close_timeout=self._close_timeout, ): self._ws = ws self._reconnect_delay = 1.0 self._reconnect_count = 0 logger.info(f"[Yuanbao] WebSocket connected: {self._ws_url}") auth_frame = json.dumps( { "type": "auth", "access_token": token, } ) await ws.send(auth_frame) try: auth_raw = await asyncio.wait_for(ws.recv(), timeout=self._auth_timeout) auth_data = json.loads(auth_raw) if auth_data.get("type") != "auth_ok": logger.warning(f"[Yuanbao] WS auth failed: {auth_data}") self._auth_failed = True break except TimeoutError: logger.warning("[Yuanbao] WS auth timeout") break await self._handle_messages(ws) except Exception as e: if not self._running: break if self._auth_failed: logger.error(f"[Yuanbao] WebSocket auth permanently failed, stopping") break self._reconnect_count += 1 if self._reconnect_count > self._max_reconnect: logger.error(f"[Yuanbao] Max reconnect attempts ({self._max_reconnect}) exceeded") break logger.warning( f"[Yuanbao] WebSocket disconnected, reconnecting in {self._reconnect_delay}s " f"(attempt {self._reconnect_count}/{self._max_reconnect}): {e}" ) await asyncio.sleep(self._reconnect_delay) self._reconnect_delay = min( self._reconnect_delay * 2, self._max_reconnect_delay, ) async def _handle_messages(self, ws) -> None: try: async for raw in ws: if not self._running: break try: event = json.loads(raw) event_type = event.get("type", "") if event.get("timestamp"): try: self._last_event_timestamp = float(event["timestamp"]) except (ValueError, TypeError): pass if event_type == "pong": continue await self._on_event(event) except Exception as e: logger.error(f"[Yuanbao] Failed to handle WS event: {e}", exc_info=True) except asyncio.CancelledError: raise except Exception: raise # connection closed or other transport error, triggers reconnect in _listen_loop