import asyncio import logging import threading from typing import Awaitable, Callable, Optional from websocket import WebSocketApp from utils import system_logger from .protocol import parse_ws_payload from .session import DouyinImSession logger = logging.getLogger("douyin_im.ws") MessageHandler = Callable[[dict], Awaitable[None]] class DouyinImWsClient: """直连 frontier-im WebSocket(websocket-client,与 DouYin_Spider 一致)""" def __init__( self, session: DouyinImSession, on_message: MessageHandler, account_id: int | None = None, ): self.session = session self.on_message = on_message self.account_id = account_id self._running = False self._task: Optional[asyncio.Task] = None self._loop: Optional[asyncio.AbstractEventLoop] = None self._ws_app: Optional[WebSocketApp] = None self._ws_lock = threading.Lock() async def start(self): url = self.session.frontier_ws_url() if not url: logger.warning("No frontier WebSocket URL captured; WS listener disabled") system_logger.record( "实时接收未启用:未获取到 frontier WebSocket 地址", detail="缺少有效的 device_id 或 sessionid,无法建立实时私信通道,将仅依赖 HTTP 轮询。", level="warning", category="ws", account_id=self.account_id, ) return self._running = True self._loop = asyncio.get_running_loop() self._task = asyncio.create_task(self._run_loop(url)) async def stop(self): self._running = False with self._ws_lock: if self._ws_app: try: self._ws_app.close() except Exception: pass self._ws_app = None if self._task: self._task.cancel() try: await self._task except asyncio.CancelledError: pass self._task = None async def _run_loop(self, url: str): retry = 0 while self._running: from .frontier import ensure_frontier_ws ensure_frontier_ws(self.session) connect_url = self.session.frontier_ws_url() or url try: logger.info(f"Connecting IM WebSocket: {connect_url[:100]}...") await self._loop.run_in_executor(None, self._connect_sync, connect_url) retry = 0 except asyncio.CancelledError: break except Exception as e: logger.warning(f"IM WebSocket error: {e}") system_logger.record( "实时接收连接异常", detail=f"建立 frontier WebSocket 失败:{e}", level="error", category="ws", account_id=self.account_id, ) if not self._running: break retry += 1 wait = min(30, 2 * retry) logger.info(f"IM WebSocket reconnect in {wait}s...") system_logger.record( f"实时接收断开,{wait}s 后重连", detail="frontier WebSocket 连接已断开,正在自动重连。", level="warning", category="ws", account_id=self.account_id, ) await asyncio.sleep(wait) def _connect_sync(self, url: str): if not self._loop: return def on_open(_ws): logger.info("IM WebSocket connected") system_logger.record( "实时接收通道已连接", detail="frontier WebSocket 已建立,可实时接收私信。", level="success", category="ws", account_id=self.account_id, ) def on_message(_ws, message): asyncio.run_coroutine_threadsafe(self._dispatch(message), self._loop) def on_error(_ws, error): if self._running: logger.warning(f"IM WebSocket error: {error}") system_logger.record( "实时接收通道报错", detail=f"{error}", level="error", category="ws", account_id=self.account_id, ) def on_close(_ws, code, msg): logger.info(f"IM WebSocket closed: code={code}, msg={msg}") if self._running: system_logger.record( "实时接收通道关闭", detail=f"code={code}, msg={msg}", level="warning", category="ws", account_id=self.account_id, ) headers = { "User-Agent": self.session.user_agent, "Pragma": "no-cache", "Cache-Control": "no-cache", "Accept-Language": "zh-CN,zh;q=0.9,en;q=0.8", "Sec-WebSocket-Protocol": "binary, base64, pbbp2", "Sec-WebSocket-Extensions": "permessage-deflate; client_max_window_bits", } ws_app = WebSocketApp( url, header=headers, cookie=self.session.cookie_header(), on_open=on_open, on_message=on_message, on_error=on_error, on_close=on_close, ) with self._ws_lock: self._ws_app = ws_app try: ws_app.run_forever(origin="https://www.douyin.com", ping_interval=20, ping_timeout=10) finally: with self._ws_lock: if self._ws_app is ws_app: self._ws_app = None async def _dispatch(self, raw): if isinstance(raw, str): payload = raw.encode("utf-8", errors="ignore") else: payload = raw items = parse_ws_payload(payload) for item in items: try: await self.on_message(item) except Exception as e: logger.debug(f"WS message handler error: {e}") system_logger.record( "实时消息处理失败", detail=f"处理收到的私信时出错:{e}", level="error", category="recv", account_id=self.account_id, )