452 lines
17 KiB
Python
452 lines
17 KiB
Python
import asyncio
|
|
import logging
|
|
import os
|
|
import weakref
|
|
from typing import Awaitable, Callable, Optional
|
|
|
|
from websockets.legacy.client import WebSocketClientProtocol, connect as websocket_connect
|
|
|
|
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]]
|
|
|
|
# Both stages are finite. The transport queue gives the receive coroutine a
|
|
# small amount of breathing room, while the application queue decouples Pong /
|
|
# frame reads from potentially slow database and reply work. Once both fill,
|
|
# backpressure intentionally reaches TCP instead of allocating more tasks.
|
|
_TRANSPORT_MAX_QUEUE = 4
|
|
_APPLICATION_QUEUE_SIZE = 8
|
|
_INCOMING_MAX_SIZE = 2**20
|
|
_STABLE_CONNECTION_SECONDS = 60.0
|
|
_MAX_RECONNECT_BASE_SECONDS = 60.0
|
|
_CLOSE_GRACE_SECONDS = 2.0
|
|
_PING_TIMEOUT_SECONDS = 120.0
|
|
_HANDLER_CONCURRENCY_ENV = "KEFU_WS_HANDLER_CONCURRENCY"
|
|
_SYSTEM_LOG_THROTTLE_ENV = "KEFU_WS_SYSTEM_LOG_THROTTLE_SECONDS"
|
|
|
|
|
|
def _env_int_clamped(name: str, default: int, minimum: int, maximum: int) -> int:
|
|
try:
|
|
value = int(os.getenv(name, str(default)) or default)
|
|
except (TypeError, ValueError):
|
|
value = default
|
|
return max(minimum, min(maximum, value))
|
|
|
|
|
|
def _env_float_clamped(
|
|
name: str,
|
|
default: float,
|
|
minimum: float,
|
|
maximum: float,
|
|
) -> float:
|
|
try:
|
|
value = float(os.getenv(name, str(default)) or default)
|
|
except (TypeError, ValueError):
|
|
value = default
|
|
return max(minimum, min(maximum, value))
|
|
|
|
|
|
def _handler_concurrency_limit() -> int:
|
|
# SQLite serializes writes. Eight allows unrelated parsing / reads to
|
|
# progress without letting an accidental value such as 500 recreate the
|
|
# original event-loop and database stampede.
|
|
return _env_int_clamped(_HANDLER_CONCURRENCY_ENV, 8, 1, 32)
|
|
|
|
|
|
def _system_log_throttle_seconds() -> float:
|
|
return _env_float_clamped(_SYSTEM_LOG_THROTTLE_ENV, 300.0, 10.0, 3600.0)
|
|
|
|
|
|
class _LoopWsState:
|
|
"""Shared limits for all WS clients owned by one asyncio event loop."""
|
|
|
|
def __init__(self) -> None:
|
|
self.handler_slots = asyncio.Semaphore(_handler_concurrency_limit())
|
|
self.system_log_last_at: dict[tuple[int, str], float] = {}
|
|
|
|
|
|
# asyncio synchronization primitives belong to their creating event loop.
|
|
# Keeping one weakly-keyed state per loop gives production a process-wide
|
|
# limit while keeping isolated test loops and uncommon threaded loops safe.
|
|
_LOOP_STATES: "weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, _LoopWsState]" = (
|
|
weakref.WeakKeyDictionary()
|
|
)
|
|
|
|
|
|
def _get_loop_state() -> _LoopWsState:
|
|
loop = asyncio.get_running_loop()
|
|
state = _LOOP_STATES.get(loop)
|
|
if state is None:
|
|
state = _LoopWsState()
|
|
_LOOP_STATES[loop] = state
|
|
return state
|
|
|
|
|
|
def _reconnect_delay(account_id: int | None, retry: int) -> float:
|
|
"""Return exponential backoff with stable, account-specific full jitter.
|
|
|
|
A connection that is accepted and immediately closed is still a failed
|
|
attempt. The old client reset its retry counter whenever ``run_forever``
|
|
returned normally, which kept those accounts reconnecting every 2-7s.
|
|
This delay reaches a 60-90s range after repeated short-lived connections.
|
|
"""
|
|
|
|
attempt = max(1, int(retry or 1))
|
|
base = min(_MAX_RECONNECT_BASE_SECONDS, float(2 ** min(attempt, 6)))
|
|
# Spread later retries across half of the base interval. Keep at least a
|
|
# five-second spread on early retries so a shared outage doesn't reconnect
|
|
# every account in the same instant.
|
|
spread = max(5.0, base / 2.0)
|
|
# Keep one account's fraction stable across attempts. This preserves the
|
|
# exponential ordering while different accounts remain spread apart.
|
|
seed = (int(account_id or 0) * 2654435761) & 0xFFFFFFFF
|
|
fraction = (seed % 10000) / 10000.0
|
|
return base + (spread * fraction)
|
|
|
|
|
|
class DouyinImWsClient:
|
|
"""Async frontier-im WebSocket client with bounded message backpressure."""
|
|
|
|
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.connected = False
|
|
self._task: Optional[asyncio.Task] = None
|
|
self._connection: Optional[WebSocketClientProtocol] = None
|
|
self._last_connection_lifetime = 0.0
|
|
self._message_queue: Optional[asyncio.Queue[dict]] = None
|
|
self._dispatcher_task: Optional[asyncio.Task] = None
|
|
|
|
async def start(self):
|
|
if self._task and not self._task.done():
|
|
return
|
|
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._ensure_dispatcher()
|
|
self._task = asyncio.create_task(
|
|
self._run_loop(url),
|
|
name=f"im-ws-{self.account_id or 'na'}",
|
|
)
|
|
|
|
async def stop(self):
|
|
self._running = False
|
|
self.connected = False
|
|
|
|
connection = self._connection
|
|
if connection is not None:
|
|
try:
|
|
await asyncio.wait_for(
|
|
connection.close(code=1000, reason="client stopping"),
|
|
timeout=_CLOSE_GRACE_SECONDS,
|
|
)
|
|
except asyncio.TimeoutError:
|
|
# Shutdown iterates over every hosted account. One broken
|
|
# peer must not consume close_timeout repeatedly and turn a
|
|
# 500-account shutdown into a many-minute operation.
|
|
logger.debug("Timed out closing IM WebSocket; aborting transport")
|
|
try:
|
|
connection.fail_connection()
|
|
except Exception:
|
|
pass
|
|
except Exception:
|
|
logger.debug("Failed to close IM WebSocket cleanly", exc_info=True)
|
|
|
|
task = self._task
|
|
if task and task is not asyncio.current_task() and not task.done():
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
if self._task is task:
|
|
self._task = None
|
|
self._connection = None
|
|
await self._stop_dispatcher()
|
|
|
|
def _record_connection_system_event(
|
|
self,
|
|
event_key: str,
|
|
message: str,
|
|
*,
|
|
detail: str,
|
|
level: str,
|
|
) -> bool:
|
|
"""Persist at most one repeated lifecycle event per account/window."""
|
|
|
|
loop = asyncio.get_running_loop()
|
|
state = _get_loop_state()
|
|
key = (int(self.account_id or 0), event_key)
|
|
now = loop.time()
|
|
last_at = state.system_log_last_at.get(key)
|
|
if last_at is not None and now - last_at < _system_log_throttle_seconds():
|
|
return False
|
|
state.system_log_last_at[key] = now
|
|
system_logger.record(
|
|
message,
|
|
detail=detail,
|
|
level=level,
|
|
category="ws",
|
|
account_id=self.account_id,
|
|
)
|
|
return True
|
|
|
|
def _reset_connection_system_log_throttle(self) -> None:
|
|
loop = asyncio.get_running_loop()
|
|
state = _LOOP_STATES.get(loop)
|
|
if state is None:
|
|
return
|
|
account_key = int(self.account_id or 0)
|
|
state.system_log_last_at.pop((account_key, "connected"), None)
|
|
state.system_log_last_at.pop((account_key, "retry"), None)
|
|
|
|
def _ensure_dispatcher(self) -> None:
|
|
if self._dispatcher_task and not self._dispatcher_task.done():
|
|
return
|
|
if self._message_queue is None:
|
|
self._message_queue = asyncio.Queue(maxsize=_APPLICATION_QUEUE_SIZE)
|
|
self._dispatcher_task = asyncio.create_task(
|
|
self._dispatch_loop(),
|
|
name=f"im-ws-dispatch-{self.account_id or 'na'}",
|
|
)
|
|
|
|
async def _stop_dispatcher(self) -> None:
|
|
task = self._dispatcher_task
|
|
self._dispatcher_task = None
|
|
if task and task is not asyncio.current_task() and not task.done():
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
queue = self._message_queue
|
|
self._message_queue = None
|
|
if queue is not None:
|
|
# Dropped messages must decrement the unfinished counter so tests,
|
|
# diagnostics, and a later restart can never hang on queue.join().
|
|
while True:
|
|
try:
|
|
queue.get_nowait()
|
|
except asyncio.QueueEmpty:
|
|
break
|
|
else:
|
|
queue.task_done()
|
|
|
|
async def _prepare_url(self, fallback_url: str) -> str:
|
|
from .frontier import ensure_frontier_ws
|
|
from .traffic_control import get_traffic_controller
|
|
|
|
# Frontier discovery can perform synchronous signing / HTTP work. It
|
|
# remains in the shared background lane and off the FastAPI event loop.
|
|
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)
|
|
return self.session.frontier_ws_url() or fallback_url
|
|
|
|
async def _run_loop(self, initial_url: str):
|
|
retry = 0
|
|
first_attempt = True
|
|
while self._running:
|
|
self._last_connection_lifetime = 0.0
|
|
try:
|
|
# Startup validation already prepared the captured URL. Avoid
|
|
# repeating signing / frontier discovery for all 500 accounts
|
|
# on their first connect; refresh only after a disconnect.
|
|
if first_attempt and initial_url:
|
|
connect_url = initial_url
|
|
else:
|
|
connect_url = await self._prepare_url(initial_url)
|
|
first_attempt = False
|
|
if not connect_url:
|
|
raise RuntimeError("frontier WebSocket URL is unavailable")
|
|
logger.info("Connecting IM WebSocket: %s...", connect_url[:100])
|
|
await self._run_connection(connect_url)
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception as exc:
|
|
if self._running:
|
|
logger.warning("IM WebSocket error: %s", exc)
|
|
self._record_connection_system_event(
|
|
"retry",
|
|
"实时接收连接异常",
|
|
detail=f"建立 frontier WebSocket 失败:{exc}",
|
|
level="error",
|
|
)
|
|
|
|
# Only a genuinely stable connection earns a retry reset. A
|
|
# successful handshake followed by an immediate normal close must
|
|
# continue exponential backoff rather than reconnect forever at
|
|
# the first delay.
|
|
if self._last_connection_lifetime >= _STABLE_CONNECTION_SECONDS:
|
|
retry = 0
|
|
if not self._running:
|
|
break
|
|
|
|
retry += 1
|
|
wait = _reconnect_delay(self.account_id, retry)
|
|
logger.info("IM WebSocket reconnect in %.1fs...", wait)
|
|
self._record_connection_system_event(
|
|
"retry",
|
|
f"实时接收断开,{wait:.1f}s 后重连",
|
|
detail="frontier WebSocket 连接已断开,正在自动重连。",
|
|
level="warning",
|
|
)
|
|
try:
|
|
await asyncio.sleep(wait)
|
|
except asyncio.CancelledError:
|
|
break
|
|
|
|
def _connection_headers(self) -> list[tuple[str, str]]:
|
|
headers = [
|
|
("Pragma", "no-cache"),
|
|
("Cache-Control", "no-cache"),
|
|
("Accept-Language", "zh-CN,zh;q=0.9,en;q=0.8"),
|
|
]
|
|
cookie = self.session.cookie_header()
|
|
if cookie:
|
|
headers.append(("Cookie", cookie))
|
|
return headers
|
|
|
|
async def _run_connection(self, url: str) -> None:
|
|
"""Open one connection and dispatch messages sequentially.
|
|
|
|
``max_queue`` bounds the library's receive buffer and ``_dispatch``
|
|
feeds one lifecycle-owned, bounded application queue. This receive
|
|
loop therefore remains responsive to control frames during ordinary
|
|
database stalls without creating one task per incoming frame.
|
|
"""
|
|
|
|
loop = asyncio.get_running_loop()
|
|
connected_at: float | None = None
|
|
connection: Optional[WebSocketClientProtocol] = None
|
|
try:
|
|
async with websocket_connect(
|
|
url,
|
|
origin="https://www.douyin.com",
|
|
subprotocols=["binary", "base64", "pbbp2"],
|
|
extra_headers=self._connection_headers(),
|
|
user_agent_header=self.session.user_agent,
|
|
compression="deflate",
|
|
open_timeout=10,
|
|
ping_interval=20,
|
|
# A handler may legitimately wait up to SQLite's 30s busy
|
|
# timeout. Leave enough headroom for queued work so a healthy
|
|
# socket isn't mistaken for a dead peer during that stall.
|
|
ping_timeout=_PING_TIMEOUT_SECONDS,
|
|
close_timeout=3,
|
|
# Frontier frames contain metadata and media URLs rather than
|
|
# media bytes. A finite frame limit plus a finite queue makes
|
|
# receive memory genuinely bounded across hundreds of peers.
|
|
max_size=_INCOMING_MAX_SIZE,
|
|
max_queue=_TRANSPORT_MAX_QUEUE,
|
|
) as websocket:
|
|
connection = websocket
|
|
self._connection = websocket
|
|
connected_at = loop.time()
|
|
self.connected = True
|
|
logger.info("IM WebSocket connected")
|
|
self._record_connection_system_event(
|
|
"connected",
|
|
"实时接收通道已连接",
|
|
detail="frontier WebSocket 已建立,可实时接收私信。",
|
|
level="success",
|
|
)
|
|
|
|
async for raw in websocket:
|
|
if not self._running:
|
|
break
|
|
await self._dispatch(raw)
|
|
finally:
|
|
if connected_at is not None:
|
|
self._last_connection_lifetime = max(0.0, loop.time() - connected_at)
|
|
if self._last_connection_lifetime >= _STABLE_CONNECTION_SECONDS:
|
|
# A genuinely healthy session starts a new lifecycle. Its
|
|
# next outage should be visible immediately rather than
|
|
# hidden by an old retry window.
|
|
self._reset_connection_system_log_throttle()
|
|
self.connected = False
|
|
if self._connection is connection:
|
|
self._connection = None
|
|
|
|
if connection is not None:
|
|
code = connection.close_code
|
|
reason = connection.close_reason
|
|
logger.info("IM WebSocket closed: code=%s, msg=%s", code, reason)
|
|
if self._running:
|
|
self._record_connection_system_event(
|
|
"retry",
|
|
"实时接收通道关闭",
|
|
detail=f"code={code}, msg={reason}",
|
|
level="warning",
|
|
)
|
|
|
|
async def _dispatch(self, raw):
|
|
self._ensure_dispatcher()
|
|
queue = self._message_queue
|
|
if queue is None:
|
|
return
|
|
if isinstance(raw, str):
|
|
payload = raw.encode("utf-8", errors="ignore")
|
|
else:
|
|
payload = raw
|
|
items = parse_ws_payload(payload)
|
|
for item in items:
|
|
if not self._running:
|
|
return
|
|
await queue.put(item)
|
|
|
|
async def _dispatch_loop(self) -> None:
|
|
queue = self._message_queue
|
|
if queue is None:
|
|
return
|
|
handler_slots = _get_loop_state().handler_slots
|
|
while True:
|
|
item = await queue.get()
|
|
try:
|
|
if self._running:
|
|
# Every account owns one dispatcher, preserving its FIFO.
|
|
# The shared semaphore prevents 500 dispatchers from
|
|
# entering SQLite / reply work at the same instant. A
|
|
# dispatcher waiting here is directly cancellable by
|
|
# stop(); no detached per-message task is created.
|
|
async with handler_slots:
|
|
if self._running:
|
|
await self.on_message(item)
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception as exc:
|
|
logger.debug("WS message handler error: %s", exc)
|
|
system_logger.record(
|
|
"实时消息处理失败",
|
|
detail=f"处理收到的私信时出错:{exc}",
|
|
level="error",
|
|
category="recv",
|
|
account_id=self.account_id,
|
|
)
|
|
finally:
|
|
queue.task_done()
|