WechatOnCloud/bridge/woc_bridge/messaging/streamer.py
Kris e490c51510 refactor(bridge): 优化SSE订阅与日志,修复FastAPI注解问题
1. 修复with_db_retry装饰器在from __future__ annotations下的FastAPI参数识别问题
2. 重构SSE日志逻辑,统一广播日志避免重复打印
3. 新增订阅者上限限制,超限时剔除最早订阅并发送kicked事件
4. 替换剪贴板粘贴方案为xdotool type适配新版微信UI
5. 补充详细的UI操作日志便于调试
2026-07-09 19:59:50 +08:00

437 lines
18 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.

"""消息流推送器:内部轮询 message DB 并通过 SSE 广播给订阅者。
微信本地 DB 是 SQLite没有 pub/sub 机制bridge 只能轮询。本模块把轮询
从客户端移到服务端:后台 watcher 周期性检查 message_0.db 的 mtime变化时
读取增量消息并广播给所有 SSE 订阅者。
优化点:
- 无订阅者时降频到 5s仍维持全局游标避免新连接被历史消息淹没
- mtime 感知DB/WAL mtime 未变化时跳过 SQL 查询
- 批量读取(默认 200 条),消息洪峰时减少广播次数
- 全局游标对齐当前最新 create_time新连接只收"从现在起"的新消息
- 背压:订阅队列满时丢弃旧事件并记日志,避免单客户端拖垮全局
事件类型SSE
- sync连接建立时发送data.cursor 为当前全局游标
- 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
import asyncio
import json
import logging
import time
from dataclasses import dataclass
from typing import Awaitable, Callable, Optional
logger = logging.getLogger("woc-bridge")
# 有订阅者时的轮询间隔(秒)
POLL_INTERVAL_ACTIVE = 1.0
# 无订阅者时的轮询间隔(秒,仅维持全局游标)
POLL_INTERVAL_IDLE = 5.0
# 单次读取上限
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]]]]
@dataclass
class StreamEvent:
"""SSE 事件。
Attributes:
event: 事件类型sync / messages / status / heartbeat
data: 事件数据(序列化为 JSON
id: 可选事件 ID客户端断线重连时可用 Last-Event-ID 补全
"""
event: str
data: dict
id: Optional[str] = None
def format_sse(event: StreamEvent) -> str:
"""把 StreamEvent 格式化为 SSE 文本帧。
格式遵循 SSE 规范:每行 `field: value`,事件间用空行分隔。
"""
lines: list[str] = []
if event.id:
lines.append(f"id: {event.id}")
if event.event:
lines.append(f"event: {event.event}")
lines.append(f"data: {json.dumps(event.data, ensure_ascii=False)}")
return "\n".join(lines) + "\n\n"
class MessageStreamer:
"""消息流推送器。
维护一个全局游标 global_cursor后台 watcher 轮询 DB把新消息以
StreamEvent 广播给所有订阅者。每个订阅者对应一个 asyncio.Queue
SSE 路由从队列取事件推给客户端。
使用:
streamer = MessageStreamer(db_reader, lambda: _resolve_db_state_tuple())
await streamer.start() # lifespan 启动
q = streamer.subscribe() # SSE 路由订阅
streamer.unsubscribe(q) # SSE 路由断开
await streamer.stop() # lifespan 停止
"""
def __init__(
self,
db_reader: "DbReader",
resolve_db_state: DbStateResolver,
extract_key: Optional[Callable[[], Awaitable[Optional[dict[str, str]]]]] = None,
) -> None:
"""初始化推送器。
Args:
db_reader: DbReader 实例(用于读消息与查 mtime
resolve_db_state: 异步回调,返回 (db_accessible, db_error_code)
供 watcher 判断 DB 是否可读并广播状态变化
extract_key: 异步回调,强制重新提取 DB 密钥。
get_messages_since 遇到 DB_ENCRYPTED 时调用,成功后重试一次。
None 时不做重试(仅记日志)
"""
self._db_reader = db_reader
self._resolve_db_state = resolve_db_state
self._extract_key = extract_key
# value 为订阅时间戳monotonic用于在超限时剔除最早订阅者
self._subscribers: dict[asyncio.Queue[StreamEvent], float] = {}
self._global_cursor: int = 0
self._watcher: asyncio.Task | None = None
# DB mtime 缓存:变化时才查 SQL
self._last_db_mtime: float = 0.0
self._last_wal_mtime: float = 0.0
# 上次广播的 DB 状态码,变化时发 status 事件
self._last_db_error_code: Optional[str] = "__init__"
self._cursor_inited: bool = False
@property
def global_cursor(self) -> int:
"""当前全局游标(最近一次推送的 create_time"""
return self._global_cursor
def subscriber_count(self) -> int:
"""当前订阅者数量(供 /api/status 暴露给客户端)。"""
return len(self._subscribers)
def subscribe(self) -> asyncio.Queue[StreamEvent]:
"""订阅事件流。
订阅后立即向队列放入一个 sync 事件,携带当前全局游标。客户端据此
判断是否需要用 /api/messages/since 补全断线期间消息。
当订阅者数量达到 MAX_SUBSCRIBERS 上限时,剔除最早订阅者(向其队列
投递 kicked 事件),再添加新订阅者。防止僵尸连接累积。
Returns:
asyncio.Queue事件队列容量 QUEUE_MAXSIZE满时丢弃旧事件
"""
q: asyncio.Queue[StreamEvent] = asyncio.Queue(maxsize=QUEUE_MAXSIZE)
# 超上限:剔除最早订阅者(向其发 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(
event="sync",
data={"cursor": self._global_cursor},
))
except asyncio.QueueFull:
pass
logger.info("SSE 订阅,当前订阅者 %dcursor=%d", self.subscriber_count(), self._global_cursor)
return q
def unsubscribe(self, q: asyncio.Queue[StreamEvent]) -> None:
"""取消订阅。"""
self._subscribers.pop(q, None)
logger.info("SSE 取消订阅,当前订阅者 %d", self.subscriber_count())
async def start(self) -> None:
"""启动 watcher 后台任务。"""
if self._watcher is None or self._watcher.done():
self._watcher = asyncio.create_task(self._watch_loop())
logger.info("MessageStreamer watcher 已启动")
async def stop(self) -> None:
"""停止 watcher 后台任务。"""
if self._watcher is not None and not self._watcher.done():
self._watcher.cancel()
try:
await self._watcher
except asyncio.CancelledError:
pass
self._watcher = None
logger.info("MessageStreamer watcher 已停止")
# ------------------------------------------------------------------
# 内部:广播与轮询
# ------------------------------------------------------------------
def _broadcast(self, event: StreamEvent) -> None:
"""广播事件给所有订阅者。
队列满时丢弃最旧事件再放入,避免单客户端消费慢拖垮全局。
日志在此处统一打印一次,避免在每订阅者生成器循环内重复打印。
"""
# sync 事件只发给单个订阅者(在 subscribe 内直接 put不走广播
if event.event != "sync":
self._log_broadcast_once(event)
for q in list(self._subscribers):
try:
q.put_nowait(event)
except asyncio.QueueFull:
# 背压:丢掉最旧事件,腾出位置
try:
q.get_nowait()
except asyncio.QueueEmpty:
pass
try:
q.put_nowait(event)
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。
避免新连接收到全量历史消息。DB 不可读时保持 0等可读后再对齐。
Returns:
True 表示成功对齐游标或游标已对齐False 表示 DB 不可读,
需要在后续轮询中重试
"""
try:
max_ct = await asyncio.to_thread(self._db_reader.get_max_create_time)
if max_ct and max_ct > 0:
self._global_cursor = max_ct
self._cursor_inited = True
logger.info("MessageStreamer 初始游标对齐到 %d", self._global_cursor)
return True
# DB 可读但 message 表为空或返回 None视为已对齐cursor=0 合理)
self._cursor_inited = True
return True
except Exception as e:
logger.warning("MessageStreamer 初始化游标失败: %s", e)
return False
async def _watch_loop(self) -> None:
"""watcher 主循环:轮询 DB 并广播新消息。
- 首次运行时对齐全局游标
- 按 POLL_INTERVAL_ACTIVE/IDLE 间隔轮询
- DB 状态变化时发 status 事件
- DB mtime 变化时读增量消息并广播
- 每 HEARTBEAT_INTERVAL 秒发一次心跳
"""
# 初始化游标DB 可能尚不可读,失败则后续轮询中重试)
await self._init_cursor()
last_heartbeat = time.monotonic()
while True:
try:
interval = (
POLL_INTERVAL_ACTIVE if self._subscribers else POLL_INTERVAL_IDLE
)
await asyncio.sleep(interval)
# 心跳(仅有订阅者时发,无订阅者发心跳无意义)
now = time.monotonic()
if self._subscribers and now - last_heartbeat >= HEARTBEAT_INTERVAL:
self._broadcast(StreamEvent(event="heartbeat", data={}))
last_heartbeat = now
await self._poll_once()
except asyncio.CancelledError:
raise
except Exception as e:
logger.exception("MessageStreamer watch loop 异常: %s", e)
await asyncio.sleep(2.0)
async def _poll_once(self) -> None:
"""单次轮询:检查 DB 状态 + mtime + 读增量消息。"""
# 1. 解析 DB 状态
try:
db_accessible, db_error_code = await self._resolve_db_state()
except Exception as e:
logger.warning("MessageStreamer: resolve_db_state 失败: %s", e)
return
# 2. 状态变化通知
if db_error_code != self._last_db_error_code:
self._last_db_error_code = db_error_code
self._broadcast(StreamEvent(
event="status",
data={
"db_accessible": db_accessible,
"db_error_code": db_error_code,
},
))
if not db_accessible:
return
# 3. 游标未初始化DB 刚恢复可读)时补对齐
if not self._cursor_inited:
await self._init_cursor()
return
# 4. mtime 感知DB/WAL 未变化则跳过 SQL
mtimes = await asyncio.to_thread(
self._db_reader.get_db_mtime, "message/message_0.db"
)
if mtimes is None:
return
db_mtime, wal_mtime = mtimes
if db_mtime == self._last_db_mtime and wal_mtime == self._last_wal_mtime:
return
self._last_db_mtime = db_mtime
self._last_wal_mtime = wal_mtime
# 5. 读取增量消息has_more=True 时循环读完剩余批次,不受 mtime 缓存影响)
while True:
try:
result = await asyncio.to_thread(
self._db_reader.get_messages_since, self._global_cursor, BATCH_LIMIT
)
except Exception as e:
# DB_ENCRYPTED 时尝试重新提取 key 并重试一次
if getattr(e, "code", None) == "DB_ENCRYPTED" and self._extract_key is not None:
logger.warning(
"MessageStreamer: get_messages_since DB_ENCRYPTED触发 key 重新提取后重试"
)
extracted = await self._extract_key()
if not extracted:
logger.warning("MessageStreamer: key 重新提取失败,跳过本轮轮询")
return
try:
result = await asyncio.to_thread(
self._db_reader.get_messages_since, self._global_cursor, BATCH_LIMIT
)
except Exception as e2:
logger.warning("MessageStreamer: 重试后仍失败: %s", e2)
return
else:
logger.warning("MessageStreamer: get_messages_since 失败: %s", e)
return
messages = result.get("messages", [])
if not messages:
return
next_cursor = result.get("next_cursor", self._global_cursor)
has_more = result.get("has_more", False)
self._global_cursor = next_cursor
# 6. 广播
self._broadcast(StreamEvent(
event="messages",
id=str(next_cursor),
data={
"messages": messages,
"next_cursor": next_cursor,
"has_more": has_more,
},
))
# 逐条打印通过 SSE 推送出去的消息详情
for msg in messages:
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,
)
logger.info(
"MessageStreamer 推送 %d 条消息, cursor=%d, 订阅者=%d",
len(messages), next_cursor, self.subscriber_count(),
)
if not has_more:
return