Files
dy/backend/rpa_engine/douyin_im/ws_client.py
T
2026-09-01 15:31:05 +08:00

616 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import asyncio
import gzip
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]]
def _safe_frame_metadata(payload: bytes) -> str:
"""Return non-content protobuf metadata for early connection diagnostics."""
try:
from .static import Live_pb2, Response_pb2
frame = Live_pb2.PushFrame()
frame.ParseFromString(payload)
body = bytes(frame.payload)
if str(frame.payloadEncoding or "").lower() == "gzip":
body = gzip.decompress(body)
response = Response_pb2.Response()
response.ParseFromString(body)
fields = [field.name for field, _ in response.body.ListFields()]
message = str(response.message or response.error_desc or "")[:80]
return (
f"service={frame.service} method={frame.method} "
f"encoding={frame.payloadEncoding or 'none'} "
f"type={frame.payloadType or 'none'} payload_bytes={len(body)} "
f"cmd={response.cmd} body={','.join(fields) or 'none'} "
f"status={message or 'ok'}"
)
except Exception as exc:
return f"metadata_unavailable={type(exc).__name__}"
# 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()
)
# frontier 按 device_id 寻址推送:两个托管账号共用同一个设备号时,两条长连接会
# 订阅到同一个地址并互相收到对方的私信。真正的拦截在 service 的会话归属校验里,
# 这里只负责把「为什么会串号」明确告诉用户。持弱引用,账号停管后自动失效。
_FRONTIER_DEVICE_OWNERS: "dict[str, weakref.ref[DouyinImWsClient]]" = {}
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
self._received_frame_count = 0
self._heartbeat_ack_logged = False
self._frontier_device_id = ""
self._blocked_device_owner_id: Optional[int] = 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
self._release_frontier_device()
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)
state.system_log_last_at.pop((account_key, "device_taken"), 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")
if self._claim_frontier_device(connect_url):
logger.info("Connecting IM WebSocket: %s...", connect_url[:100])
await self._run_connection(connect_url)
else:
# 设备号已被另一个在跑的账号占用:绝不并连同一个推送地址,
# 本账号本轮退回 HTTP 轮询兜底(connected 保持 False
# service 会自动切到更快的会话对账节奏),并在退避后重试,
# 等占用方停管时自动接管。
self._report_frontier_device_taken(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
self._release_frontier_device()
def _frontier_device_owner(self, device_id: str) -> "Optional[DouyinImWsClient]":
"""当前仍活着的设备号占用方(run 循环任务还在跑才算数)。"""
reference = _FRONTIER_DEVICE_OWNERS.get(device_id)
owner = reference() if reference is not None else None
if owner is None or owner is self:
return None
task = owner._task
if not owner._running or task is None or task.done():
return None
return owner
def _claim_frontier_device(self, url: str) -> bool:
"""独占本账号的 frontier 设备地址;已被别的账号占用时返回 False。
frontier 按 device_id 寻址推送。两个账号共用同一个设备号时,同时建连
会让两条连接互相收到对方的私信(串号的根因),且抖音也可能只保留最后
一条连接、把先连上的那个账号踢成「连着但收不到」。所以同一个设备地址
永远只允许一个账号建连,另一个账号走 HTTP 轮询兜底。
"""
from .frontier import ws_device_id
device_id = ws_device_id(url)
if not device_id:
# 判不出设备号(自建地址/异常格式)时不阻断连接,交给会话归属校验兜底。
return True
owner = self._frontier_device_owner(device_id)
if owner is not None and int(owner.account_id or 0) != int(self.account_id or 0):
self._blocked_device_owner_id = owner.account_id
return False
_FRONTIER_DEVICE_OWNERS[device_id] = weakref.ref(self)
self._frontier_device_id = device_id
self._blocked_device_owner_id = None
return True
def _report_frontier_device_taken(self, url: str) -> None:
from .frontier import ws_device_id
device_id = ws_device_id(url)
owner_id = self._blocked_device_owner_id
logger.error(
"Account %s cannot open frontier device_id %s: already held by "
"account %s; falling back to HTTP polling this round",
self.account_id,
device_id,
owner_id,
)
self._record_connection_system_event(
"device_taken",
"实时接收已让出:与另一个账号共用长连接设备号",
detail=(
f"本账号与账号 {owner_id} 的 frontier 设备号相同(device_id={device_id})。"
"同一个设备地址只允许一个账号建立长连接,否则两个账号会互相收到对方的"
"私信。本账号本轮不建连,改由 HTTP 会话轮询接收(有几十秒级延迟),"
"并在对方停止托管后自动接管。"
"根治办法:为每个账号在独立的浏览器配置/设备上重新采集凭证。"
),
level="error",
)
def _release_frontier_device(self) -> None:
device_id = self._frontier_device_id
self._frontier_device_id = ""
if not device_id:
return
reference = _FRONTIER_DEVICE_OWNERS.get(device_id)
if reference is not None and reference() is self:
_FRONTIER_DEVICE_OWNERS.pop(device_id, None)
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
@staticmethod
def _uses_browser_frontier(url: str) -> bool:
return "zijieapi.com" in url and "access_key=" in url
async def _run_browser_heartbeat(self, websocket) -> None:
"""Mirror Frontier's browser SDK application-level ``hi`` heartbeat."""
while self._running:
await websocket.send("hi")
await asyncio.sleep(30)
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
heartbeat_task: Optional[asyncio.Task] = None
browser_frontier = self._uses_browser_frontier(url)
source_ip = str(getattr(self.session, "egress_source_ip", "") or "").strip()
connect_kwargs = {"local_addr": (source_ip, 0)} if source_ip else {}
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,
# The current Douyin browser Frontier SDK uses a text ``hi``
# heartbeat instead of RFC WebSocket ping frames.
ping_interval=None if browser_frontier else 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,
**connect_kwargs,
) as websocket:
connection = websocket
self._connection = websocket
connected_at = loop.time()
self.connected = True
logger.info(
"IM WebSocket connected: subprotocol=%s",
getattr(websocket, "subprotocol", None) or "none",
)
self._record_connection_system_event(
"connected",
"实时接收通道已连接",
detail="frontier WebSocket 已建立,可实时接收私信。",
level="success",
)
if browser_frontier:
heartbeat_task = asyncio.create_task(
self._run_browser_heartbeat(websocket),
name=f"im-ws-heartbeat-{self.account_id or 'na'}",
)
try:
async for raw in websocket:
if not self._running:
break
await self._dispatch(raw)
finally:
if heartbeat_task and not heartbeat_task.done():
heartbeat_task.cancel()
try:
await heartbeat_task
except asyncio.CancelledError:
pass
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):
if raw == "hi":
if not self._heartbeat_ack_logged:
logger.info("IM WebSocket application heartbeat acknowledged")
self._heartbeat_ack_logged = True
return
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)
self._received_frame_count += 1
if self._received_frame_count <= 3:
metadata = _safe_frame_metadata(payload) if not items else "parsed-message"
logger.info(
"IM WebSocket frame received: seq=%d kind=%s bytes=%d parsed=%d %s",
self._received_frame_count,
"text" if isinstance(raw, str) else "binary",
len(payload),
len(items),
metadata,
)
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()