Yufeng He 3aa1ac477d
fix(channels): cap WeCom outbound content at the 20480-byte protocol limit (#5148)
* fix(channels): cap WeCom outbound content at the 20480-byte protocol limit

Both _send_ws paths sent unbounded text. Stream replies now clip on a
character boundary with a truncation marker (one stream carries the whole
reply and cannot split mid-way), and proactive pushes split into
sequential markdown messages at newline boundaries. Measured in UTF-8
bytes, matching the documented protocol cap.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* test(channels): pin emoji boundary behavior in the WeCom content limit

Review on #5148 raised 4-byte emoji cut points. Probes show the split
path already carries a byte-split character into the next chunk and
terminates on all-emoji input; these tests pin that behavior so a later
refactor cannot regress it.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* fix(channels): preserve the delimiter when splitting WeCom pushes

The boundary newline was stripped by lstrip, so the sequential markdown
messages lost one delimiter per split and could not rebuild the original
response. Keep it on the chunk's tail and assert the exact round trip in
the tests, including leading blank lines.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* fix(channels): serialize WeCom proactive chunk batches per chat

Each chunk send awaits, so two manager workers pushing long texts to the
same chat could interleave markdown messages (A1, B1, A2, B2) and break
the sequential-message contract. Hold a per-chat lock across the whole
split batch; different chats still send concurrently.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* fix(channels): keep WeCom split advancing and cap the chunk batch

Two edge cases in _split_for_byte_limit left after the delimiter fix:

- A limit narrower than one whole character made the decode window
  empty, so the hard cut became 0 and the loop appended empty chunks
  forever. Take the character anyway when the window decodes empty, so
  the loop always advances.
- A single oversized push produced one message per 20480 bytes with no
  ceiling, flooding the chat and holding the per-chat lock for the whole
  drain. Cap one push at 10 messages: keep the first nine verbatim and
  collapse the rest into one clipped tail carrying the truncation
  marker, with a warning log when the cap trips.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* fix(channels): reclaim completed WeCom send locks, cap the split before it does the work

Two leftovers from the last review round:

- _ws_send_locks kept one lock per chat forever. A guard-locked refcount
  now reclaims an entry only when no sender is queued on it, so a waiter
  can never land on a fresh lock mid-batch for the same chat. Pinned by
  three reclamation tests (single push, capped batch, 20 concurrent chats).
- _split_for_byte_limit built every chunk and then joined the discarded
  tail to clip it — quadratic work on pathological pushes. The batch cap
  now applies inside the loop: once the kept chunks are full, the
  remainder is clipped whole. A spy test pins that the clipper receives
  the unsplit remainder.

tests/test_wecom_content_limit.py 25 passed, plus tests/test_wecom_ws_text.py.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* docs(channels): describe the WeCom push cap, pin the staggered-lock race

The channels guide still described proactive pushes as an uncapped split. Also add the staggered-start regression the reviewer asked for: a waiter queuing while the holder's cleanup runs must share one lock, keep batches contiguous, and leave both lock registries empty.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

---------

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>
2026-09-12 09:59:58 +08:00

620 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.

from __future__ import annotations
import asyncio
import base64
import hashlib
import inspect
import logging
from collections.abc import Awaitable, Callable
from typing import Any, cast
from app.channels.base import Channel
from app.channels.commands import is_known_channel_command
from app.channels.connection_identity import attach_connection_identity
from app.channels.message_bus import (
InboundMessage,
InboundMessageType,
MessageBus,
OutboundMessage,
ResolvedAttachment,
)
logger = logging.getLogger(__name__)
def _file_md5(path: str) -> str:
md5_hasher = hashlib.md5()
with open(path, "rb") as file_obj:
for chunk in iter(lambda: file_obj.read(1024 * 1024), b""):
md5_hasher.update(chunk)
return md5_hasher.hexdigest()
def _open_binary(path: str):
return open(path, "rb")
# The WeCom bot protocol caps message content at 20480 UTF-8 bytes, for both
# passive stream replies and active markdown pushes.
_WECOM_MAX_CONTENT_BYTES = 20480
_TRUNCATION_MARKER = "\n\n... (truncated)"
# One push must not flood the chat with an unbounded run of messages: keep the
# first few chunks and collapse the rest into one truncated tail.
_WECOM_MAX_CHUNK_BATCH = 10
def _clip_to_byte_limit(text: str, limit: int) -> str:
"""Clip text to a UTF-8 byte budget, never splitting a character."""
if len(text.encode("utf-8")) <= limit:
return text
budget = limit - len(_TRUNCATION_MARKER.encode("utf-8"))
clipped = text.encode("utf-8")[:budget].decode("utf-8", errors="ignore")
return clipped + _TRUNCATION_MARKER
def _split_for_byte_limit(text: str, limit: int) -> list[str]:
"""Split text into chunks within the UTF-8 byte limit.
Prefers newline boundaries so markdown structure survives the split.
The batch cap applies inside the loop, so a pathological text is never
fully split just to be discarded.
"""
if len(text.encode("utf-8")) <= limit:
return [text]
chunks: list[str] = []
remaining = text
while len(remaining.encode("utf-8")) > limit:
if len(chunks) >= _WECOM_MAX_CHUNK_BATCH - 1:
logger.warning(
"WeCom push of %d bytes exceeds %d messages, capping the batch",
len(text.encode("utf-8")),
_WECOM_MAX_CHUNK_BATCH,
)
chunks.append(_clip_to_byte_limit(remaining, limit))
return chunks
window = remaining.encode("utf-8")[:limit].decode("utf-8", errors="ignore")
cut = window.rfind("\n")
if cut <= 0:
cut = len(window)
else:
# Keep the delimiter on this chunk's tail: the sequential messages
# must round-trip to the original text exactly.
cut += 1
if cut == 0:
# limit is narrower than one whole character; take it anyway so
# the loop always advances.
cut = 1
chunks.append(remaining[:cut])
remaining = remaining[cut:]
if remaining:
chunks.append(remaining)
return chunks
class WeComChannel(Channel):
def __init__(self, bus: MessageBus, config: dict[str, Any]) -> None:
super().__init__(name="wecom", bus=bus, config=config)
self._bot_id: str | None = None
self._bot_secret: str | None = None
self._ws_client = None
self._ws_task: asyncio.Task | None = None
self._ws_shutdown_task: asyncio.Future[Any] | None = None
self._lifecycle_lock = asyncio.Lock()
self._ws_frames: dict[str, dict[str, Any]] = {}
self._ws_stream_ids: dict[str, str] = {}
self._ws_send_locks: dict[str, asyncio.Lock] = {}
self._ws_send_lock_users: dict[str, int] = {}
self._ws_send_locks_guard = asyncio.Lock()
self._working_message = "Working on it..."
@property
def supports_streaming(self) -> bool:
return True
def _clear_ws_context(self, thread_ts: str | None) -> None:
if not thread_ts:
return
self._ws_frames.pop(thread_ts, None)
self._ws_stream_ids.pop(thread_ts, None)
async def _send_ws_upload_command(self, req_id: str, body: dict[str, Any], cmd: str) -> dict[str, Any]:
if not self._ws_client:
raise RuntimeError("WeCom WebSocket client is not available")
ws_manager = getattr(self._ws_client, "_ws_manager", None)
send_reply = getattr(ws_manager, "send_reply", None)
if not callable(send_reply):
raise RuntimeError("Installed wecom-aibot-python-sdk does not expose the WebSocket media upload API expected by DeerFlow. Use wecom-aibot-python-sdk==0.1.6 or update the adapter.")
send_reply_async = cast(Callable[[str, dict[str, Any], str], Awaitable[dict[str, Any]]], send_reply)
return await send_reply_async(req_id, body, cmd)
async def start(self) -> None:
async with self._lifecycle_lock:
if self._running:
return
bot_id = self.config.get("bot_id")
bot_secret = self.config.get("bot_secret")
working_message = self.config.get("working_message")
self._bot_id = bot_id if isinstance(bot_id, str) and bot_id else None
self._bot_secret = bot_secret if isinstance(bot_secret, str) and bot_secret else None
self._working_message = working_message if isinstance(working_message, str) and working_message else "Working on it..."
if not self._bot_id or not self._bot_secret:
logger.error("WeCom channel requires bot_id and bot_secret")
return
try:
from aibot import WSClient, WSClientOptions
except ImportError:
logger.error("wecom-aibot-python-sdk is not installed. Install it with: uv add wecom-aibot-python-sdk")
return
else:
self._ws_client = WSClient(WSClientOptions(bot_id=self._bot_id, secret=self._bot_secret, logger=logger))
self._ws_client.on("message.text", self._on_ws_text)
self._ws_client.on("message.mixed", self._on_ws_mixed)
self._ws_client.on("message.image", self._on_ws_image)
self._ws_client.on("message.file", self._on_ws_file)
self._ws_client.on("error", self._on_ws_error)
self._ws_client.on("disconnected", self._on_ws_disconnected)
self._ws_task = asyncio.create_task(self._ws_client.connect())
self._ws_task.add_done_callback(self._on_ws_task_done)
self._running = True
self.bus.subscribe_outbound(self._on_outbound)
logger.info("WeCom channel started")
def _on_ws_task_done(self, task: asyncio.Task) -> None:
if task.cancelled():
return
exc = task.exception()
if exc is None:
return
logger.error(
"WeCom WebSocket connection task failed: %s. Check that the network/proxy allows wss://openws.work.weixin.qq.com and that bot_id/bot_secret are valid.",
exc,
)
def _on_ws_error(self, error: Any) -> None:
logger.error("WeCom WebSocket error: %s", error)
def _on_ws_disconnected(self, *args: Any) -> None:
detail = f" ({args[0]})" if args else ""
logger.warning("WeCom WebSocket disconnected%s; SDK will attempt to reconnect", detail)
def _begin_ws_shutdown(self, ws_client: Any) -> asyncio.Future[Any] | None:
ws_manager = getattr(ws_client, "_ws_manager", None)
async_disconnect = getattr(ws_manager, "_async_disconnect", None)
stop_heartbeat = getattr(ws_manager, "_stop_heartbeat", None)
clear_pending_messages = getattr(ws_manager, "_clear_pending_messages", None)
if inspect.iscoroutinefunction(async_disconnect) and callable(stop_heartbeat) and callable(clear_pending_messages):
# wecom-aibot-python-sdk 1.0.2 makes disconnect() synchronous and
# discards the task created for _async_disconnect(). Perform its
# synchronous bookkeeping here so DeerFlow can own and await the
# actual SDK shutdown operation without scheduling a duplicate.
try:
if hasattr(ws_client, "_started"):
ws_client._started = False
ws_manager._is_manual_close = True
stop_heartbeat()
clear_pending_messages("Connection manually closed")
except Exception:
logger.exception("Failed to prepare WeCom WebSocket shutdown")
return asyncio.create_task(async_disconnect())
try:
result = ws_client.disconnect()
except Exception:
logger.exception("Failed to request WeCom WebSocket disconnect")
return None
if inspect.isawaitable(result):
return asyncio.ensure_future(result)
return None
async def stop(self) -> None:
async with self._lifecycle_lock:
self._running = False
self.bus.unsubscribe_outbound(self._on_outbound)
ws_client = self._ws_client
ws_task = self._ws_task
if ws_task and not ws_task.done():
ws_task.cancel()
shutdown_task = None
try:
shutdown_task = self._begin_ws_shutdown(ws_client) if ws_client else None
self._ws_shutdown_task = shutdown_task
tasks = [task for task in (ws_task, shutdown_task) if task is not None]
drain_future = asyncio.gather(*tasks, return_exceptions=True) if tasks else None
if drain_future is not None:
try:
await asyncio.shield(drain_future)
except asyncio.CancelledError:
# Caller cancellation still propagates, but only after
# the channel-owned tasks have completed their cleanup.
await drain_future
raise
finally:
if self._ws_task is ws_task:
self._ws_task = None
if self._ws_client is ws_client:
self._ws_client = None
if self._ws_shutdown_task is shutdown_task:
self._ws_shutdown_task = None
self._ws_frames.clear()
self._ws_stream_ids.clear()
logger.info("WeCom channel stopped")
async def send(self, msg: OutboundMessage, *, _max_retries: int = 3) -> None:
if self._ws_client:
await self._send_ws(msg, _max_retries=_max_retries)
return
logger.warning("[WeCom] send called but WebSocket client is not available")
async def _on_outbound(self, msg: OutboundMessage) -> None:
if msg.channel_name != self.name:
return
try:
await self.send(msg)
except Exception:
logger.exception("Failed to send outbound message on channel %s", self.name)
if msg.is_final:
self._clear_ws_context(msg.thread_ts)
return
for attachment in msg.attachments:
try:
success = await self.send_file(msg, attachment)
if not success:
logger.warning("[%s] file upload skipped for %s", self.name, attachment.filename)
except Exception:
logger.exception("[%s] failed to upload file %s", self.name, attachment.filename)
if msg.is_final:
self._clear_ws_context(msg.thread_ts)
async def send_file(self, msg: OutboundMessage, attachment: ResolvedAttachment) -> bool:
if not msg.is_final:
return True
if not self._ws_client:
return False
if not msg.thread_ts:
return False
frame = self._ws_frames.get(msg.thread_ts)
if not frame:
return False
media_type = "image" if attachment.is_image else "file"
size_limit = 2 * 1024 * 1024 if attachment.is_image else 20 * 1024 * 1024
if attachment.size > size_limit:
logger.warning(
"[WeCom] %s too large (%d bytes), skipping: %s",
media_type,
attachment.size,
attachment.filename,
)
return False
try:
media_id = await self._upload_media_ws(
media_type=media_type,
filename=attachment.filename,
path=str(attachment.actual_path),
size=attachment.size,
)
if not media_id:
return False
body = {media_type: {"media_id": media_id}, "msgtype": media_type}
await self._ws_client.reply(frame, body)
logger.debug("[WeCom] %s sent via ws: %s", media_type, attachment.filename)
return True
except Exception:
logger.exception("[WeCom] failed to upload/send file via ws: %s", attachment.filename)
return False
async def _on_ws_text(self, frame: dict[str, Any]) -> None:
body = frame.get("body", {}) or {}
text = ((body.get("text") or {}).get("content") or "").strip()
quote = (((body.get("quote") or {}).get("text") or {}).get("content") or "").strip()
if not text and not quote:
return
await self._publish_ws_inbound(frame, text + (f"\nQuote message: {quote}" if quote else ""))
async def _on_ws_mixed(self, frame: dict[str, Any]) -> None:
body = frame.get("body", {}) or {}
mixed = body.get("mixed") or {}
items = mixed.get("msg_item") or []
parts: list[str] = []
files: list[dict[str, Any]] = []
for item in items:
item_type = (item or {}).get("msgtype")
if item_type == "text":
content = (((item or {}).get("text") or {}).get("content") or "").strip()
if content:
parts.append(content)
elif item_type in ("image", "file"):
payload = (item or {}).get(item_type) or {}
url = payload.get("url")
aeskey = payload.get("aeskey")
if isinstance(url, str) and url:
files.append(
{
"type": item_type,
"url": url,
"aeskey": (aeskey if isinstance(aeskey, str) and aeskey else None),
}
)
text = "\n\n".join(parts).strip()
if not text and not files:
return
if not text:
text = "receive image/file"
await self._publish_ws_inbound(frame, text, files=files)
async def _on_ws_image(self, frame: dict[str, Any]) -> None:
body = frame.get("body", {}) or {}
image = body.get("image") or {}
url = image.get("url")
aeskey = image.get("aeskey")
if not isinstance(url, str) or not url:
return
await self._publish_ws_inbound(
frame,
"receive image ",
files=[
{
"type": "image",
"url": url,
"aeskey": aeskey if isinstance(aeskey, str) and aeskey else None,
}
],
)
async def _on_ws_file(self, frame: dict[str, Any]) -> None:
body = frame.get("body", {}) or {}
file_obj = body.get("file") or {}
url = file_obj.get("url")
aeskey = file_obj.get("aeskey")
if not isinstance(url, str) or not url:
return
await self._publish_ws_inbound(
frame,
"receive file",
files=[
{
"type": "file",
"url": url,
"aeskey": aeskey if isinstance(aeskey, str) and aeskey else None,
}
],
)
async def _publish_ws_inbound(
self,
frame: dict[str, Any],
text: str,
*,
files: list[dict[str, Any]] | None = None,
) -> None:
if not self._ws_client:
return
try:
from aibot import generate_req_id
except Exception:
return
body = frame.get("body", {}) or {}
msg_id = body.get("msgid")
if not msg_id:
return
user_id = (body.get("from") or {}).get("userid")
connect_code = self._pending_connect_code(text)
if connect_code:
handled = await self._bind_connection_from_connect_code(
frame=frame,
user_id=str(user_id or ""),
code=connect_code,
)
if handled:
return
inbound_type = InboundMessageType.COMMAND if is_known_channel_command(text) else InboundMessageType.CHAT
inbound = self._make_inbound(
chat_id=user_id, # keep user's conversation in memory
user_id=user_id,
text=text,
msg_type=inbound_type,
thread_ts=msg_id,
files=files or [],
metadata={
"aibotid": body.get("aibotid"),
"chattype": body.get("chattype"),
"message_id": msg_id,
},
)
inbound.topic_id = user_id # keep the same thread
stream_id = generate_req_id("stream")
self._ws_frames[msg_id] = frame
self._ws_stream_ids[msg_id] = stream_id
reservation = self._reserve_inbound(inbound)
if reservation is None:
self._ws_frames.pop(msg_id, None)
self._ws_stream_ids.pop(msg_id, None)
return
try:
try:
await self._ws_client.reply_stream(frame, stream_id, self._working_message, False)
except Exception:
pass
inbound = await self._attach_connection_identity(inbound)
self._commit_reserved_inbound(reservation, inbound)
finally:
reservation.release()
async def _attach_connection_identity(self, inbound: InboundMessage) -> InboundMessage:
return await attach_connection_identity(
inbound,
repo=self._connection_repo,
provider="wecom",
workspace_id=str(inbound.metadata.get("aibotid") or "") or None,
fallback_without_workspace=True,
)
async def _bind_connection_from_connect_code(self, *, frame: dict[str, Any], user_id: str, code: str) -> bool:
if self._connection_repo is None or not code:
return False
state = await self._connection_repo.consume_oauth_state(provider="wecom", state=code)
if state is None:
await self._send_connection_reply(frame, "WeCom connection code is invalid or expired.")
return True
if not user_id:
await self._send_connection_reply(frame, "WeCom connection could not be completed from this message.")
return True
body = frame.get("body", {}) or {}
workspace_id = str(body.get("aibotid") or "") or None
await self._connection_repo.upsert_connection(
owner_user_id=state["owner_user_id"],
provider="wecom",
external_account_id=user_id,
workspace_id=workspace_id,
metadata={
"aibotid": workspace_id,
"chattype": body.get("chattype"),
},
status="connected",
)
await self._send_connection_reply(frame, "WeCom connected to DeerFlow.")
return True
async def _send_connection_reply(self, frame: dict[str, Any], text: str) -> None:
if not self._ws_client:
return
await self._ws_client.reply(frame, {"msgtype": "text", "text": {"content": text}})
async def _send_ws(self, msg: OutboundMessage, *, _max_retries: int = 3) -> None:
if not self._ws_client:
return
try:
from aibot import generate_req_id
except Exception:
generate_req_id = None
if msg.thread_ts and msg.thread_ts in self._ws_frames:
frame = self._ws_frames[msg.thread_ts]
stream_id = self._ws_stream_ids.get(msg.thread_ts)
if not stream_id and generate_req_id:
stream_id = generate_req_id("stream")
self._ws_stream_ids[msg.thread_ts] = stream_id
if not stream_id:
return
await self._send_with_retry(
lambda: self._ws_client.reply_stream(frame, stream_id, _clip_to_byte_limit(msg.text, _WECOM_MAX_CONTENT_BYTES), bool(msg.is_final)),
max_retries=_max_retries,
log_prefix="[WeCom]",
operation_name="stream send",
)
return
# No replyable frame (e.g. a scheduled-task push): a stream reply is one
# stream per reply and cannot split mid-way, but this path can, so the
# full text goes out as sequential markdown messages. Each send awaits,
# so hold a per-chat lock across the whole batch: manager workers run
# concurrently, and two long pushes to the same chat would otherwise
# interleave chunks (A1, B1, A2, B2) and break the sequential contract.
async with self._ws_send_locks_guard:
lock = self._ws_send_locks.setdefault(msg.chat_id, asyncio.Lock())
self._ws_send_lock_users[msg.chat_id] = self._ws_send_lock_users.get(msg.chat_id, 0) + 1
try:
async with lock:
for chunk in _split_for_byte_limit(msg.text, _WECOM_MAX_CONTENT_BYTES):
body = {"msgtype": "markdown", "markdown": {"content": chunk}}
await self._send_with_retry(
lambda body=body: self._ws_client.send_message(msg.chat_id, body),
max_retries=_max_retries,
log_prefix="[WeCom]",
)
finally:
async with self._ws_send_locks_guard:
self._ws_send_lock_users[msg.chat_id] -= 1
# Reclaim only while nobody else is queued on this chat's lock;
# the guard serializes the check so a waiter can never end up
# holding a fresh lock for a chat whose batch is mid-flight.
if self._ws_send_lock_users[msg.chat_id] == 0:
self._ws_send_lock_users.pop(msg.chat_id, None)
self._ws_send_locks.pop(msg.chat_id, None)
async def _upload_media_ws(
self,
*,
media_type: str,
filename: str,
path: str,
size: int,
) -> str | None:
if not self._ws_client:
return None
try:
from aibot import generate_req_id
except Exception:
return None
chunk_size = 512 * 1024
total_chunks = (size + chunk_size - 1) // chunk_size
if total_chunks < 1 or total_chunks > 100:
logger.warning("[WeCom] invalid total_chunks=%d for %s", total_chunks, filename)
return None
md5 = await asyncio.to_thread(_file_md5, path)
init_req_id = generate_req_id("aibot_upload_media_init")
init_body = {
"type": media_type,
"filename": filename,
"total_size": int(size),
"total_chunks": int(total_chunks),
"md5": md5,
}
init_ack = await self._send_ws_upload_command(init_req_id, init_body, "aibot_upload_media_init")
upload_id = (init_ack.get("body") or {}).get("upload_id")
if not upload_id:
logger.warning("[WeCom] upload init returned no upload_id: %s", init_ack)
return None
file_obj = await asyncio.to_thread(_open_binary, path)
try:
for idx in range(total_chunks):
data = await asyncio.to_thread(file_obj.read, chunk_size)
if not data:
break
chunk_req_id = generate_req_id("aibot_upload_media_chunk")
chunk_body = {
"upload_id": upload_id,
"chunk_index": int(idx),
"base64_data": base64.b64encode(data).decode("utf-8"),
}
await self._send_ws_upload_command(chunk_req_id, chunk_body, "aibot_upload_media_chunk")
finally:
await asyncio.to_thread(file_obj.close)
finish_req_id = generate_req_id("aibot_upload_media_finish")
finish_ack = await self._send_ws_upload_command(finish_req_id, {"upload_id": upload_id}, "aibot_upload_media_finish")
media_id = (finish_ack.get("body") or {}).get("media_id")
if not media_id:
logger.warning("[WeCom] upload finish returned no media_id: %s", finish_ack)
return None
return media_id