refactor(bridge): 优化SSE订阅与日志,修复FastAPI注解问题

1. 修复with_db_retry装饰器在from __future__ annotations下的FastAPI参数识别问题
2. 重构SSE日志逻辑,统一广播日志避免重复打印
3. 新增订阅者上限限制,超限时剔除最早订阅并发送kicked事件
4. 替换剪贴板粘贴方案为xdotool type适配新版微信UI
5. 补充详细的UI操作日志便于调试
This commit is contained in:
Kris 2026-07-09 19:59:50 +08:00
parent 102b98adea
commit e490c51510
4 changed files with 141 additions and 85 deletions

View File

@ -8,11 +8,12 @@ from __future__ import annotations
import asyncio
import functools
import inspect
import logging
import os
import time
from dataclasses import dataclass
from typing import Optional
from typing import Optional, get_type_hints
from woc_bridge.config import _state, _require_db_reader, _require_xdotool
from woc_bridge.db.decryptor import Decryptor, _resolve_page1_key_material
@ -55,6 +56,27 @@ def with_db_retry(func):
raise
logger.info("%s key 重新提取成功,重试请求", func.__name__)
return await func(*args, **kwargs)
# functools.wraps 复制的 __annotations__ 是字符串化注解(因 from __future__
# import annotationsFastAPI 解析 wrapper 签名时拿到字符串注解,无法识别
# Pydantic body 参数,会把它当 query 参数处理导致 422。用 get_type_hints
# 解析为真实类型,并据此重建 __signature__让 FastAPI 能正确识别参数类型。
try:
hints = get_type_hints(func, include_extras=True)
wrapper.__annotations__ = hints
# 基于原始函数签名重建 wrapper 的 __signature__替换 (*args, **kwargs)
orig_sig = inspect.signature(func)
new_params = []
for name, param in orig_sig.parameters.items():
annotation = hints.get(name, param.annotation)
new_params.append(param.replace(annotation=annotation))
wrapper.__signature__ = orig_sig.replace(
parameters=new_params,
return_annotation=hints.get("return", orig_sig.return_annotation),
)
except Exception:
# 解析失败时保持 functools.wraps 的默认行为
pass
return wrapper

View File

@ -16,6 +16,8 @@
- messages一批新消息data messages/next_cursor/has_more
- statusDB 状态变化data db_accessible/db_error_code
- heartbeat30 秒心跳保活 NAT/代理连接
- kicked订阅者被服务端剔除超过 MAX_SUBSCRIBERS 上限时剔除最早订阅者
data.reason 给出剔除原因客户端收到后应关闭连接并按需重连
"""
from __future__ import annotations
@ -39,6 +41,8 @@ BATCH_LIMIT = 200
QUEUE_MAXSIZE = 100
# 心跳间隔(秒)
HEARTBEAT_INTERVAL = 30.0
# 最大并发订阅者数量:超过则剔除最早订阅者(防止僵尸连接累积)
MAX_SUBSCRIBERS = 3
# DB 状态解析器返回类型:(db_accessible, db_error_code)
DbStateResolver = Callable[[], Awaitable[tuple[bool, Optional[str]]]]
@ -107,7 +111,8 @@ class MessageStreamer:
self._db_reader = db_reader
self._resolve_db_state = resolve_db_state
self._extract_key = extract_key
self._subscribers: set[asyncio.Queue[StreamEvent]] = set()
# value 为订阅时间戳monotonic用于在超限时剔除最早订阅者
self._subscribers: dict[asyncio.Queue[StreamEvent], float] = {}
self._global_cursor: int = 0
self._watcher: asyncio.Task | None = None
# DB mtime 缓存:变化时才查 SQL
@ -132,11 +137,39 @@ class MessageStreamer:
订阅后立即向队列放入一个 sync 事件携带当前全局游标客户端据此
判断是否需要用 /api/messages/since 补全断线期间消息
当订阅者数量达到 MAX_SUBSCRIBERS 上限时剔除最早订阅者向其队列
投递 kicked 事件再添加新订阅者防止僵尸连接累积
Returns:
asyncio.Queue事件队列容量 QUEUE_MAXSIZE满时丢弃旧事件
"""
q: asyncio.Queue[StreamEvent] = asyncio.Queue(maxsize=QUEUE_MAXSIZE)
self._subscribers.add(q)
# 超上限:剔除最早订阅者(向其发 kicked 事件,由生成器循环 break 退出)
if len(self._subscribers) >= MAX_SUBSCRIBERS:
oldest_q = min(self._subscribers, key=lambda k: self._subscribers[k])
del self._subscribers[oldest_q]
try:
oldest_q.put_nowait(StreamEvent(
event="kicked",
data={"reason": "max_subscribers_reached"},
))
except asyncio.QueueFull:
# 队列满也要剔除:直接丢弃最旧事件后塞入 kicked
try:
oldest_q.get_nowait()
except asyncio.QueueEmpty:
pass
try:
oldest_q.put_nowait(StreamEvent(
event="kicked",
data={"reason": "max_subscribers_reached"},
))
except asyncio.QueueFull:
pass
logger.warning(
"SSE 订阅者已达上限 %d,剔除最早订阅者以腾出位置", MAX_SUBSCRIBERS,
)
self._subscribers[q] = time.monotonic()
# 立即发送 sync 事件(只给这个订阅者)
try:
q.put_nowait(StreamEvent(
@ -150,7 +183,7 @@ class MessageStreamer:
def unsubscribe(self, q: asyncio.Queue[StreamEvent]) -> None:
"""取消订阅。"""
self._subscribers.discard(q)
self._subscribers.pop(q, None)
logger.info("SSE 取消订阅,当前订阅者 %d", self.subscriber_count())
async def start(self) -> None:
@ -177,7 +210,11 @@ class MessageStreamer:
"""广播事件给所有订阅者。
队列满时丢弃最旧事件再放入避免单客户端消费慢拖垮全局
日志在此处统一打印一次避免在每订阅者生成器循环内重复打印
"""
# sync 事件只发给单个订阅者(在 subscribe 内直接 put不走广播
if event.event != "sync":
self._log_broadcast_once(event)
for q in list(self._subscribers):
try:
q.put_nowait(event)
@ -192,6 +229,43 @@ class MessageStreamer:
except asyncio.QueueFull:
logger.warning("SSE 订阅队列持续满,丢弃事件 event=%s", event.event)
def _log_broadcast_once(self, event: StreamEvent) -> None:
"""广播时统一打印一次事件摘要,替代每订阅者循环内的重复日志。"""
if event.event == "messages":
msgs = event.data.get("messages", []) or []
logger.info(
"SSE 广播 messages: 条数=%d next_cursor=%s has_more=%s 订阅者=%d",
len(msgs), event.data.get("next_cursor"), event.data.get("has_more"),
self.subscriber_count(),
)
for msg in msgs:
content_preview = (msg.get("content") or "").replace("\n", "\\n")
if len(content_preview) > 100:
content_preview = content_preview[:100] + "..."
logger.info(
"SSE 广播 消息: msg_id=%s talker=%s sender=%s is_sender=%s "
"type=%s render=%s session=%s create_time=%s content=%s",
msg.get("msg_id"),
msg.get("talker"),
msg.get("sender") or "-",
msg.get("is_sender"),
msg.get("type"),
msg.get("render_type"),
msg.get("session_type"),
msg.get("create_time"),
content_preview,
)
elif event.event == "status":
logger.info(
"SSE 广播 status: db_accessible=%s db_error_code=%s 订阅者=%d",
event.data.get("db_accessible"), event.data.get("db_error_code"),
self.subscriber_count(),
)
elif event.event == "heartbeat":
logger.info("SSE 广播 heartbeat 订阅者=%d", self.subscriber_count())
else:
logger.info("SSE 广播 %s: %s 订阅者=%d", event.event, event.data, self.subscriber_count())
async def _init_cursor(self) -> bool:
"""把全局游标对齐到当前 DB 最大 create_time。

View File

@ -78,48 +78,6 @@ async def get_messages_since(cursor: int = 0, limit: int = 50) -> MessagesRespon
# ---------------------------------------------------------------------------
# 路由GET /api/messages/streamSSE 实时消息推送)
# ---------------------------------------------------------------------------
def _log_sse_event_pushed(event: StreamEvent) -> None:
"""详细打印通过 SSE 推送给客户端的事件。
按事件类型分别打印关键字段
- messages消息条数 + 逐条 msg_id/talker/sender/content 预览
- sync / status完整 data
- heartbeat仅标记
"""
if event.event == "messages":
msgs = event.data.get("messages", []) or []
logger.info(
"SSE→客户端 messages: 条数=%d next_cursor=%s has_more=%s",
len(msgs), event.data.get("next_cursor"), event.data.get("has_more"),
)
for msg in msgs:
content_preview = (msg.get("content") or "").replace("\n", "\\n")
if len(content_preview) > 100:
content_preview = content_preview[:100] + "..."
logger.info(
"SSE→客户端 消息: msg_id=%s talker=%s sender=%s is_sender=%s "
"type=%s render=%s session=%s create_time=%s content=%s",
msg.get("msg_id"),
msg.get("talker"),
msg.get("sender") or "-",
msg.get("is_sender"),
msg.get("type"),
msg.get("render_type"),
msg.get("session_type"),
msg.get("create_time"),
content_preview,
)
elif event.event == "sync":
logger.info("SSE→客户端 sync: cursor=%s", event.data.get("cursor"))
elif event.event == "status":
logger.info(
"SSE→客户端 status: db_accessible=%s db_error_code=%s",
event.data.get("db_accessible"), event.data.get("db_error_code"),
)
elif event.event == "heartbeat":
logger.info("SSE→客户端 heartbeat")
else:
logger.info("SSE→客户端 %s: %s", event.event, event.data)
@router.get("/api/messages/stream")
@ -138,6 +96,8 @@ async def stream_messages(request: Request) -> StreamingResponse:
- statusDB 状态变化时发送data db_accessible / db_error_code
DB 加密/退出/恢复可读时触发客户端据此决定是否降级到轮询
- heartbeat 30 秒心跳保活 NAT/代理连接data 为空对象
- kicked订阅者被服务端剔除超过 MAX_SUBSCRIBERS 上限data.reason
给出原因客户端收到后应关闭连接并按需重连
断线补偿
SSE 不保证 100% 投递客户端断线重连后应
@ -154,6 +114,8 @@ async def stream_messages(request: Request) -> StreamingResponse:
- 响应头设置 X-Accel-Buffering: no 关闭 nginx 缓冲避免事件被攒批
- 客户端断开时自动取消订阅释放队列
- 不抛业务错误即使 DB 不可读也建立连接通过 status 事件告知
- 订阅者上限由 MessageStreamer.MAX_SUBSCRIBERS 控制超限时最早
订阅者会被剔除并收到 kicked 事件
"""
streamer = _require_message_streamer()
@ -169,9 +131,11 @@ async def stream_messages(request: Request) -> StreamingResponse:
except asyncio.TimeoutError:
# 队列 30s 无事件,发心跳保活
event = StreamEvent(event="heartbeat", data={})
# kicked 事件:被服务端剔除(订阅者超上限),发送后立即退出
if event.event == "kicked":
yield format_sse(event)
break
yield format_sse(event)
# 详细打印通过 SSE 推送给客户端的事件
_log_sse_event_pushed(event)
finally:
streamer.unsubscribe(queue)

View File

@ -70,10 +70,13 @@ class XdotoolDriver:
async def _key(self, key: str) -> None:
"""执行 xdotool key <key>。"""
await self._run(["xdotool", "key", key])
logger.info("[ui] key: %s", key)
rc, stdout, stderr = await self._run(["xdotool", "key", key])
logger.info("[ui] key %s -> rc=%s", key, rc)
async def _sleep(self, seconds: float) -> None:
"""asyncio.sleep 封装,便于测试与统一调速。"""
logger.info("[ui] sleep %.2fs", seconds)
await asyncio.sleep(seconds)
# ------------------------------------------------------------------
@ -293,53 +296,34 @@ class XdotoolDriver:
return None
# ------------------------------------------------------------------
# 剪贴板粘贴
# 文本输入xdotool type
# ------------------------------------------------------------------
async def _paste_via_xclip(self, text: str) -> None:
"""通过 xclip 写入剪贴板并触发 Ctrl+V 粘贴
"""通过 xdotool type 直接向聚焦控件输入文本
直接把原始字节写入 xclip stdincreate_subprocess_exec 不经 shell
无转义问题因此无需 base64 编码并行等待 stdin 写入与进程退出
避免大文本时管道缓冲阻塞导致的死锁
历史背景原实现用 xclip 写剪贴板 + Ctrl+V 粘贴但微信 4.x Linux
自绘 UIRadiumWMPF 框架0 子窗口不响应 X11 core keyboard
修饰键组合Ctrl+V只响应单字符键盘事件xev 验证 KeyPress 事件
确实到达窗口但微信不处理 xdotool type逐字符 XSendEvent
正常输入中英文因此改用 xdotool type 直接打字
方法名保留 _paste_via_xclip 以维持调用点稳定内部实现已切换
Args:
text: 待粘贴文本
text: 输入文本支持中英文路径
Raises:
BridgeError(SEND_FAILED): xclip 退出码非 0
BridgeError(SEND_FAILED): xdotool type 执行失败
"""
xclip_proc = await asyncio.create_subprocess_exec(
"xclip", "-selection", "clipboard",
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.DEVNULL,
stderr=asyncio.subprocess.PIPE,
env=self._env(),
)
assert xclip_proc.stdin is not None
async def _feed() -> None:
try:
xclip_proc.stdin.write(text.encode("utf-8"))
await xclip_proc.stdin.drain()
except (BrokenPipeError, ConnectionResetError):
# xclip 已退出,忽略
pass
finally:
try:
xclip_proc.stdin.close()
except Exception:
pass
# 并行:喂 stdin + 等进程退出,避免管道缓冲满死锁
_, (_, stderr) = await asyncio.gather(_feed(), xclip_proc.communicate())
if xclip_proc.returncode != 0:
logger.info("[ui] type text: len=%d", len(text))
rc, _, stderr = await self._run(["xdotool", "type", text])
if rc != 0:
err = stderr.decode(errors="ignore").strip() if stderr else "unknown"
raise BridgeError(
code="SEND_FAILED",
message=f"xclip 写入剪贴板失败 (code={xclip_proc.returncode}): {err}",
message=f"xdotool type 失败 (code={rc}): {err}",
)
# 触发粘贴
await self._key("ctrl+v")
logger.info("[ui] type done rc=%s", rc)
# ------------------------------------------------------------------
# 会话定位Ctrl+F 搜索)
@ -379,6 +363,7 @@ class XdotoolDriver:
async def _step(coro: Awaitable[None], desc: str) -> None:
"""执行单个 UI 步骤并加超时保护。"""
remaining = deadline - time.monotonic()
logger.info("[ui] step start: %s", desc)
if remaining <= 0:
raise BridgeError(
code="SEND_FAILED",
@ -387,13 +372,16 @@ class XdotoolDriver:
try:
await asyncio.wait_for(coro, timeout=max(1.0, remaining))
except asyncio.TimeoutError as exc:
logger.error("[ui] step timeout: %s", desc)
raise BridgeError(
code="SEND_FAILED",
message=f"定位会话步骤超时: {desc}",
) from exc
logger.info("[ui] step done: %s", desc)
# 1. 激活窗口(非阻塞,避免 --sync 死等)
window_id = await self.find_wechat_window()
logger.info("[ui] found window_id=%s", window_id)
if window_id is None:
raise BridgeError(
code="WINDOW_NOT_FOUND",
@ -404,6 +392,8 @@ class XdotoolDriver:
# 2. 关闭可能存在的搜索框/弹窗
await _step(self._key("Escape"), "关闭搜索框")
await _step(self._key("Escape"), "再按一次 Esc")
await _step(self._key("Escape"), "第三次 Esc 确保退出")
await _step(self._sleep(0.2), "等待 Esc 生效")
# 3. 打开搜索
@ -450,14 +440,20 @@ class XdotoolDriver:
Returns:
本地生成的 channel_msg_id格式 local_<unix秒>_<随机>
"""
name = display_name if display_name else to_wxid
logger.info("[send_text] to=%s name=%s content_len=%s", to_wxid, name, len(content))
# 1. 定位会话
await self._open_session_by_name(display_name if display_name else to_wxid)
await self._open_session_by_name(name)
logger.info("[send_text] session opened")
# 2. 粘贴内容并发送
await self._paste_via_xclip(content)
logger.info("[send_text] content pasted")
await self._sleep(0.2)
await self._key("Return")
logger.info("[send_text] return pressed")
# 3. 生成 local_send_id本地 ID非微信原生 msg_id
local_send_id = f"local_{int(time.time())}_{random.randint(0, 0xFFFFFF):06x}"
logger.info("[send_text] done local_send_id=%s", local_send_id)
return local_send_id
# ------------------------------------------------------------------