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:
parent
102b98adea
commit
e490c51510
@ -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 annotations),FastAPI 解析 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
|
||||
|
||||
|
||||
|
||||
@ -16,6 +16,8 @@
|
||||
- messages:一批新消息,data 含 messages/next_cursor/has_more
|
||||
- status:DB 状态变化,data 含 db_accessible/db_error_code
|
||||
- heartbeat:30 秒心跳,保活 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。
|
||||
|
||||
|
||||
@ -78,48 +78,6 @@ async def get_messages_since(cursor: int = 0, limit: int = 50) -> MessagesRespon
|
||||
# ---------------------------------------------------------------------------
|
||||
# 路由:GET /api/messages/stream(SSE 实时消息推送)
|
||||
# ---------------------------------------------------------------------------
|
||||
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:
|
||||
- status:DB 状态变化时发送,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)
|
||||
|
||||
|
||||
@ -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 stdin(create_subprocess_exec 不经 shell,
|
||||
无转义问题,因此无需 base64 编码)。并行等待 stdin 写入与进程退出,
|
||||
避免大文本时管道缓冲阻塞导致的死锁。
|
||||
历史背景:原实现用 xclip 写剪贴板 + Ctrl+V 粘贴,但微信 4.x Linux
|
||||
自绘 UI(RadiumWMPF 框架,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
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
Loading…
Reference in New Issue
Block a user