WechatOnCloud/bridge/woc_bridge/messaging/streamer.py

437 lines
18 KiB
Python
Raw Normal View History

"""消息流推送器:内部轮询 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