hataa 26800d1245
fix(channels): stream-cap and validate WeChat/WeCom inbound media downloads, fixes #5223 (#5225)
* fix(channels): stream-cap and validate WeChat/WeCom inbound media downloads, fixes #5223

* fix(channels): address WeCom APPID, decompression, and log-sanitization review findings (#5223)

Round-4 review follow-ups on the inbound-media download cap:

- The COS bucket numeric suffix is the owner's Tencent Cloud APPID and bucket
  names are user-chosen, so any Tencent Cloud account could register a
  matching ww-aibot-img-* bucket and pass the shape gate. The built-in rule
  now admits only the APPID observed in Tencent's published aibot callback
  examples (1258476243), across regions; any other account (including a
  future WeCom rotation) goes through channels.wecom.allowed_media_hosts.
- aiter_bytes() transparently decodes Content-Encoding, and the decoder
  allocates the full decompressed body before the byte cap sees a chunk (an
  ~8 KB gzip wire chunk decoding to 8 MiB reproduces it). Both URL readers
  now send Accept-Encoding: identity, refuse a response with a residual
  Content-Encoding before reading, and iterate aiter_raw().
- httpx.HTTPStatusError formats the signed URL (path + query credentials)
  into its message, so _ingest_inbound_files' reader-failure branch logs a
  sanitized summary (class + status) instead of logger.exception, and the
  WeChat extract paths catch httpx.HTTPError so the polling loop's
  per-message logger.exception can never render a media URL.

Every change ships with a red/green regression: the reviewer's 403
mock-transport repro asserted against caplog.text (fully formatted logs),
the reviewer's different-APPID bucket host, and gzip bombs driven through
real httpx mock transports in both readers. Docs (channels AGENTS.md,
README, config.example.yaml) updated for the APPID pinning and encoding
gate.

* docs(channels): document why the inbound-media cap is 50 MB, not WeCom's 100 MB ceiling

* fix(logging): redact URLs in httpx request logs down to scheme + host, fixes #5223

httpx emits 'HTTP Request: GET <full URL>' at the Gateway's INFO level
before any response handling runs, so even successful signed-media
downloads leaked their credentials. HttpxUrlQueryRedactionFilter
(installed by configure_logging) rewrites those records in place — path
and query become /<redacted>, method/status/duration observability is
preserved — which also keeps Telegram's token-bearing Bot API paths out
of the logs. Reader-level regression tests run at production INFO level
with a real MockTransport, success paths included.

* fix(logging): blank userinfo credentials in httpx request-log redaction

* fix(logging): redact authority-only URLs and cover urllib3 redirect logs

Two follow-ups from the review plus one extrapolation of the same class:

- rest is now optional in _URL_REDACT_RE, so an authority-only URL
  (scheme://user:pass@host, no path) is rewritten too — userinfo had
  nowhere else to hide and previously passed through verbatim. A bare
  credential-free origin still passes through unchanged.
- Renamed to UrlRedactionFilter / install_url_log_redaction and attached
  to the urllib3 logger as well: urllib3 logs 'Redirecting <url> -> <url>'
  at INFO with full URLs on both sides, the same leak class on a different
  library logger. No gateway path today both uses requests and redirects
  a signed URL, but the class stays closed instead of dormant.
- Unit tests now build records with the real httpx 0.28.1 format string
  ('HTTP Request: %s %s "%s %d %s"', 5 args) and httpx.URL args, per
  the nit, instead of a synthetic shape httpx never emits.

* fix(logging): install URL redaction at handler level so propagated records are covered

A logging.Filter on a logger only runs for records emitted through that
exact logger — child loggers neither inherit it nor trigger it on
propagation — so the previous attachment to the bare urllib3 logger was
dead code: urllib3 emits Redirecting via urllib3.poolmanager at INFO and
urllib3.connectionpool at DEBUG. The filter is now attached to every root
handler (mirroring _install_trace_filter, which already iterates root
handlers; handler-level filters see propagated records) in addition to
the httpx logger (httpx emits via the bare name, and emission-point
coverage survives handlers added later). The wiring is pinned by tests
that emit through the real urllib3 child loggers — a mutation removing
the handler-level install turns them red. Comments, docstrings, and
AGENTS.md now state the actual emitter names and levels.

* fix(logging): redact urllib3 DEBUG request lines, whose split shape evaded the URL regex

urllib3's per-request line (connectionpool.py:545 on 2.7.0) renders as
`scheme://host:port "METHOD /path?query HTTP/x.x" status len` — the
authority ends at a space so _URL_REDACT_RE's bare-origin early return
applies, and the quoted origin-form target has no scheme, so neither half
was rewritten. UrlRedactionFilter now runs a dedicated request-line shape
first (collapsing the target to /<redacted>, keeping scheme+host+method+
version), then the absolute-URL pass. Regressions pin the exact format
string both at unit level and through the real urllib3.connectionpool
DEBUG emit path; AGENTS.md wording now names both covered DEBUG shapes.

* fix(logging): redact urllib3 retry lines and linearize scheme scanning

Closes the two open review threads on the inbound-media log hardening:

Retry/redirect targets: urllib3 logs the request target with no scheme in
five shapes the generic absolute-URL pass cannot see - `Retry: <target>`
(connectionpool.py:954 DEBUG), `Incremented Retry for (url='<target>')`
(util/retry.py:545 DEBUG, absolute on the redirect path), `Retrying (...)
after connection broken by '<err>': <target>` (connectionpool.py:869
WARNING, above the INFO root), and origin-form halves of both Redirecting
emitters (poolmanager.py:500 INFO / connectionpool.py:922 DEBUG). Each
gets a rewrite anchored to the exact urllib3 format, collapsing the
target to /<redacted>; the generic pass's rest now stops at quote
characters so a quoted URL keeps its closing punctuation (previously the
absolute-form increment line was mangled), and the request-line method
class accepts any case. The emitter enumeration in channels AGENTS.md is
closed against the installed urllib3 2.7.0 source.

Quadratic scanning: both scheme-bearing patterns start with a character
class, so re.sub retried every suffix of a long token - 64K paths cost
~1.8s and URL-free 64K error bodies ~3.1s per record, synchronously in
every root handler. The two passes are now driven from literal "://"
occurrences: _scheme_starts walks back over the scheme charset to each
run's first letter and the pattern is attempted only there, reproducing
re.sub's leftmost-non-overlapping result in linear time (256K path:
5.6ms; worst adversarial shapes <= 28ms). Long-input regressions pin the
URL-bearing and URL-free cases with mutation-verified bounds, plus
nested-scheme and digit-headed-run equivalence cases.

Validation: tests/test_logging_config.py 12/12; scheme-pass equivalence
against the old re.sub pipeline verified by two independent 30k+ case
fuzz runs; full-suite A/B against HEAD shows zero tests that pass on HEAD
and fail with this diff.

* fix(logging): boundary-aware quote stops and whole-message Redirecting anchor

Two follow-ups on the urllib3 redaction shapes:

Embedded quotes: `rest` treated ANY quote as a closing mark, so a URL
with an apostrophe in the path kept everything after it verbatim
(`https://h/path'quoted'?token=Q` rendered the credential suffix in
full) while the class docstring claimed path/query/fragment are
replaced. A quote now closes `rest` only at a boundary - followed by
whitespace, a closing parenthesis, or end of string - so urllib3's
Incremented Retry (url='...') scaffolding keeps its ') closer while an
embedded quote stays consumed. The increment line's url capture gets
the same rule narrowed to its fixed ')' closer.

Redirecting anchoring: the origin-half pass matched `(?
<=-> )/path` as a substring, and an `-> /path` arrow is not
urllib3-owned shape - the sandbox provider's actionable mount error
(`sandbox.mounts entry <host> -> /mnt/knowledge ignored: ...`) had its
container path rewritten to /<redacted>, failing
test_setup_path_mappings_logs_actionable_error_for_missing_host_path on
CI (backend-unit-tests shard 3). The pass is now anchored to the whole
`Redirecting <t> -> <t>` message, which is exactly urllib3's record;
origin slots collapse, absolute slots stay for the generic pass.
Regression tests pin the embedded-quote shapes and the sandbox error's
byte-for-byte passthrough; both mutations verified red.

Validation: tests/test_logging_config.py 14/14; the CI-failing sandbox
test green locally; every test file asserting redaction/arrow log
content passes (attachments, support bundle, run metadata, skill
secrets, ragflow, skillscan, sandbox provider); full offline backend
suite 14084 passed / 164 failed with the failure set matching this
machine's documented Windows-environment baseline (NTFS chmod/symlink,
docker/lark/langfuse absences) - no failure involves redaction output.

* fix(logging): redact redirects with spaced locations

* fix(logging): grammar-complete Redirecting anchor; neutral WeChat guard labels

Round-13 P3 (Redirecting anchor strictness): the whole-message anchor kept
the ^Redirecting prefix (the urllib3-owned literal that stops the sandbox
false positive) but required BOTH slots whitespace-free, so a Location
header with an interior space voided the pass and leaked the origin-form
request target in the first slot - redirect_location is the raw header
string and interior spaces are legal field syntax. The tail is now loose
(\S.*$) and the first slot gets the same grammar treatment (\S.*?): the
recursive urlopen frame passes the previous raw Location as its url, so
t1 can carry interior spaces too, lazy-split at the first arrow the way
the line is constructed. A space-carrying slot collapses whole when it
starts with /; the sandbox mount error keeps passing through untouched.

Round-14 nit (None conflation): _download_cdn_bytes returns None for two
reasons (in-flight cap abort, Content-Encoding refusal) but both image and
file callers labeled it "exceeds size limit (N bytes)" - contradicting the
accurate encoding line right above it, and reporting the plaintext limit
for a ciphertext-cap decision. Callers now log a neutral
"skipped by download guard" line (the manager reader callers' shape);
the accurate reason stays inside the download function. The same sweep
also logs _stage_downloaded_file's silent None (no state dir configured),
which made an attachment vanish with no log line at all.

Also anchors the emitter-enumeration closure to its urllib3 version: the
closure reopens if an upgrade changes these format strings, so the comment
now says so explicitly.

Validation: logging 15/15 and attachments 60/62 (the two pre-existing
Windows symlink-privilege failures documented in the PR body); three
mutations verified red (old wording, strict t1, silent staging None);
ruff clean. Full offline suite run before push (per round-11 lesson).

---------

Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
2026-09-16 16:33:47 +08:00

649 lines
26 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..."
raw_hosts = config.get("allowed_media_hosts")
if isinstance(raw_hosts, str):
host_values: list[Any] = [raw_hosts]
elif isinstance(raw_hosts, (list, tuple, set, frozenset)):
host_values = list(raw_hosts)
else:
host_values = []
# Host suffixes; a leading ``*.`` is stripped so ``*.example.com`` and
# ``example.com`` behave identically (mirrors WechatChannel._coerce_host_suffixes).
normalized_hosts: set[str] = set()
for host in host_values:
text = str(host).strip().lower().lstrip(".")
if text.startswith("*."):
text = text[2:]
if text:
normalized_hosts.add(text)
self._allowed_media_host_suffixes = frozenset(normalized_hosts)
@property
def allowed_media_host_suffixes(self) -> frozenset[str]:
"""Operator host suffixes the manager-side inbound-media gate merges in.
``channels.wecom.allowed_media_hosts``: extra host suffixes inbound
media URLs may be downloaded from, on top of the built-in ``qq.com``
family and the pinned WeCom COS bucket shape. Gives deployments that
proxy or mirror WeCom media an escape hatch that does not require
widening the hard-coded pattern for everyone.
"""
return self._allowed_media_host_suffixes
@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