231 lines
8.2 KiB
Python
231 lines
8.2 KiB
Python
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_ws_thread(self, url: str):
|
||
"""在独立守护线程中跑 run_forever,直到连接断开/关闭。
|
||
|
||
不能用共享默认线程池(run_in_executor(None)/asyncio.to_thread):
|
||
WS 长连接会永久占用一个池线程,账号数超过池大小(默认 64)后,
|
||
所有账号的签名/轮询任务被饿死,表现为“启动几十个账号后全部卡死超时”。
|
||
"""
|
||
loop = asyncio.get_running_loop()
|
||
done = asyncio.Event()
|
||
error: list[BaseException] = []
|
||
|
||
def _runner():
|
||
try:
|
||
self._connect_sync(url)
|
||
except BaseException as e:
|
||
error.append(e)
|
||
finally:
|
||
try:
|
||
loop.call_soon_threadsafe(done.set)
|
||
except RuntimeError:
|
||
pass # 事件循环已关闭
|
||
|
||
thread = threading.Thread(
|
||
target=_runner,
|
||
name=f"im-ws-{self.account_id or 'na'}",
|
||
daemon=True,
|
||
)
|
||
thread.start()
|
||
try:
|
||
await done.wait()
|
||
except asyncio.CancelledError:
|
||
# stop() 会 close ws_app 使 run_forever 退出,线程随之结束
|
||
raise
|
||
if error:
|
||
raise error[0]
|
||
|
||
async def _run_loop(self, url: str):
|
||
retry = 0
|
||
while self._running:
|
||
from .frontier import ensure_frontier_ws
|
||
from .traffic_control import get_traffic_controller
|
||
|
||
# ensure_frontier_ws 可能触发签名/HTTP(阻塞),放线程池避免卡事件循环
|
||
controller = get_traffic_controller()
|
||
async with controller.background_slot(self.account_id or 0, "websocket prepare"):
|
||
await asyncio.to_thread(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._run_ws_thread(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
|
||
# Stable per-account jitter prevents every hosted account from
|
||
# reconnecting in the same second after a shared network outage.
|
||
jitter = ((int(self.account_id or 0) * 2654435761) % 5000) / 1000.0
|
||
wait = min(30.0, 2.0 * retry) + jitter
|
||
logger.info(f"IM WebSocket reconnect in {wait:.1f}s...")
|
||
system_logger.record(
|
||
f"实时接收断开,{wait:.1f}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,
|
||
)
|