refactor(adapters/cn_group): 12368->8549; wecom split into streaming/media/send_queue mixins, dingtalk inbound parsers -> inbound.py, google_chat cards.py + setup_files.py, teams summary_writer.py; dedupe token/URL/credential helpers
This commit is contained in:
+273
-1110
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,242 @@
|
||||
"""Pure parsers for inbound DingTalk ``ChatbotMessage`` payloads (no I/O, no adapter state)."""
|
||||
|
||||
import json
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
from gateway.platforms.base import MessageType
|
||||
|
||||
# DingTalk rich-text item type → runtime content type
|
||||
DINGTALK_TYPE_MAPPING = {"picture": "image", "voice": "audio"}
|
||||
|
||||
# File extension → MIME type for DingTalk file/image messages. image/* MIMEs
|
||||
# make ``extract_media`` classify msgtype='image'/'file' payloads as PHOTO.
|
||||
EXT_MAP = {
|
||||
"pdf": "application/pdf", "doc": "application/msword", "xls": "application/vnd.ms-excel",
|
||||
"docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
"xlsx": "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||||
"png": "image/png", "jpg": "image/jpeg", "jpeg": "image/jpeg", "gif": "image/gif", "webp": "image/webp",
|
||||
"md": "text/markdown", "txt": "text/plain", "csv": "text/csv", "zip": "application/zip", "mp4": "video/mp4",
|
||||
}
|
||||
|
||||
# rich-text runtime type → (media_types entry, MessageType promotion when still TEXT)
|
||||
_RICH_MEDIA = {
|
||||
"image": ("image", MessageType.PHOTO),
|
||||
"video": ("video", MessageType.VIDEO),
|
||||
"file": ("application/octet-stream", MessageType.DOCUMENT),
|
||||
}
|
||||
|
||||
|
||||
def _extensions(message: Any) -> Any:
|
||||
return getattr(message, "extensions", {}) or {}
|
||||
|
||||
|
||||
def _ext_content(message: Any) -> Optional[dict]:
|
||||
"""``extensions['content']`` when it is a dict, else None."""
|
||||
content = _extensions(message).get("content", {})
|
||||
return content if isinstance(content, dict) else None
|
||||
|
||||
|
||||
def _rich_list(message: Any) -> Optional[list]:
|
||||
"""Rich-text item list from either SDK shape (``rich_text_content.rich_text_list`` or legacy ``rich_text``)."""
|
||||
rich_text = getattr(message, "rich_text_content", None) or getattr(message, "rich_text", None)
|
||||
if not rich_text:
|
||||
return None
|
||||
rich_list = getattr(rich_text, "rich_text_list", None) or rich_text
|
||||
return rich_list if isinstance(rich_list, list) else None
|
||||
|
||||
|
||||
def _card_text(message: Any) -> str:
|
||||
"""msgtype='card' (钉钉文档分享卡片 / link card): title + doc URL from ``extensions['card']``."""
|
||||
extensions = _extensions(message)
|
||||
content = ""
|
||||
card = extensions.get("card", {})
|
||||
if isinstance(card, dict):
|
||||
title = card.get("title", "")
|
||||
raw_content = card.get("content", "")
|
||||
doc_url = ""
|
||||
if isinstance(raw_content, dict):
|
||||
doc_url = raw_content.get("url", "") or raw_content.get("docUrl", "")
|
||||
elif isinstance(raw_content, str) and raw_content.strip():
|
||||
try:
|
||||
parsed = json.loads(raw_content.strip())
|
||||
if isinstance(parsed, dict):
|
||||
doc_url = parsed.get("url", "") or parsed.get("docUrl", "")
|
||||
except (ValueError, TypeError):
|
||||
doc_url = raw_content
|
||||
parts = ([f"[文档] {title}"] if title else []) + ([doc_url] if doc_url else [])
|
||||
if parts:
|
||||
content = " ".join(parts)
|
||||
if not content:
|
||||
# Last-resort: raw text field from extensions (if present)
|
||||
ext_text = extensions.get("text", {})
|
||||
if isinstance(ext_text, dict):
|
||||
content = (ext_text.get("content", "") or "").strip()
|
||||
return content
|
||||
|
||||
|
||||
def _interactive_card_text(message: Any) -> str:
|
||||
"""msgtype='interactiveCard': ``extensions['content']`` carries title + biz_custom_action_url."""
|
||||
ext_content = _ext_content(message)
|
||||
if not ext_content:
|
||||
return ""
|
||||
doc_url = ext_content.get("biz_custom_action_url", "")
|
||||
title = ext_content.get("title", "")
|
||||
if not (doc_url or title):
|
||||
return ""
|
||||
parts = [f"[文档卡片] {title}" if title else "[文档卡片]"] + ([doc_url] if doc_url else [])
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def _ext_field(message: Any, field: str) -> Any:
|
||||
ext_content = _ext_content(message)
|
||||
return ext_content.get(field, "") if ext_content else ""
|
||||
|
||||
|
||||
def _audio_text(message: Any) -> str:
|
||||
"""msgtype='audio': DingTalk-provided speech recognition text."""
|
||||
recognition = _ext_field(message, "recognition")
|
||||
return recognition.strip() if recognition else ""
|
||||
|
||||
|
||||
def _file_text(message: Any) -> str:
|
||||
"""msgtype='file': use fileName as text."""
|
||||
fname = _ext_field(message, "fileName")
|
||||
return f"[文件] {fname}" if fname else ""
|
||||
|
||||
|
||||
# Fallbacks by msgtype when no plain/rich text was found (types are exclusive).
|
||||
_EMPTY_TEXT_FALLBACKS = (
|
||||
("audio", _audio_text),
|
||||
("file", _file_text),
|
||||
("card", _card_text),
|
||||
("interactiveCard", _interactive_card_text),
|
||||
)
|
||||
|
||||
|
||||
def extract_text(message: Any) -> str:
|
||||
"""Extract plain text from a DingTalk chatbot message.
|
||||
|
||||
Handles both SDK payload shapes: legacy ``message.text`` dict ``{"content": ...}`` and
|
||||
>= 0.20 ``TextContent`` (whose ``__str__`` is ``"TextContent(content=...)"`` — always read
|
||||
``.content`` first); rich text via ``rich_text_content.rich_text_list`` or legacy ``rich_text``.
|
||||
"""
|
||||
text = getattr(message, "text", None) or ""
|
||||
if hasattr(text, "content"):
|
||||
content = (text.content or "").strip()
|
||||
elif isinstance(text, dict):
|
||||
content = text.get("content", "").strip()
|
||||
else:
|
||||
content = str(text).strip()
|
||||
|
||||
if not content:
|
||||
rich_list = _rich_list(message)
|
||||
if rich_list is not None:
|
||||
parts = []
|
||||
for item in rich_list:
|
||||
if isinstance(item, dict):
|
||||
t = item.get("text") or item.get("content") or ""
|
||||
if t:
|
||||
parts.append(t)
|
||||
elif hasattr(item, "text") and item.text:
|
||||
parts.append(item.text)
|
||||
content = " ".join(parts).strip()
|
||||
|
||||
if not content:
|
||||
msg_type = getattr(message, "message_type", "")
|
||||
for kind, fallback in _EMPTY_TEXT_FALLBACKS:
|
||||
if msg_type == kind:
|
||||
content = fallback(message)
|
||||
break
|
||||
|
||||
# Do NOT strip "@bot": the mention is routed structurally (callback ``isInAtList``), and
|
||||
# regex-stripping @handles would damage e-mails, SSH URLs and literal "@openai" references.
|
||||
return content
|
||||
|
||||
|
||||
def extract_media(message: Any) -> Tuple[MessageType, List[str], List[str]]:
|
||||
"""Return ``(MessageType, [download codes/urls], [mime types])`` for a message."""
|
||||
msg_type = MessageType.TEXT
|
||||
media_urls: List[str] = []
|
||||
media_types: List[str] = []
|
||||
|
||||
image_content = getattr(message, "image_content", None)
|
||||
if image_content:
|
||||
download_code = getattr(image_content, "download_code", None)
|
||||
if download_code:
|
||||
media_urls.append(download_code)
|
||||
media_types.append("image")
|
||||
msg_type = MessageType.PHOTO
|
||||
|
||||
for item in _rich_list(message) or ():
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
dl_code = item.get("downloadCode") or item.get("download_code") or ""
|
||||
item_type = item.get("type", "")
|
||||
if not dl_code:
|
||||
continue
|
||||
mapped = DINGTALK_TYPE_MAPPING.get(item_type, "file")
|
||||
media_urls.append(dl_code)
|
||||
if mapped == "audio":
|
||||
media_types.append("audio")
|
||||
if msg_type == MessageType.TEXT:
|
||||
# "voice" items are native voice notes → STT (VOICE); "audio" file uploads stay AUDIO.
|
||||
msg_type = MessageType.VOICE if item_type == "voice" else MessageType.AUDIO
|
||||
else:
|
||||
mime, promoted = _RICH_MEDIA[mapped]
|
||||
media_types.append(mime)
|
||||
if msg_type == MessageType.TEXT:
|
||||
msg_type = promoted
|
||||
|
||||
msg_type_str = getattr(message, "message_type", "") or ""
|
||||
if msg_type_str == "picture" and not media_urls:
|
||||
msg_type = MessageType.PHOTO
|
||||
elif msg_type_str == "richText":
|
||||
# Only re-derive when the scan above left TEXT — resetting a VOICE/AUDIO/VIDEO/DOCUMENT
|
||||
# promotion here dropped native voice notes back to TEXT and skipped STT.
|
||||
if msg_type == MessageType.TEXT and any("image" in t for t in media_types):
|
||||
msg_type = MessageType.PHOTO
|
||||
elif msg_type_str == "audio":
|
||||
# Voice message: recognition text is already in the text. Do NOT add media_urls, or
|
||||
# run.py's transcription enrichment overwrites it with a failed STT attempt.
|
||||
if msg_type == MessageType.TEXT:
|
||||
msg_type = MessageType.VOICE
|
||||
elif msg_type_str in ("file", "image"):
|
||||
ext_content = _ext_content(message)
|
||||
if ext_content:
|
||||
dl_code = ext_content.get("downloadCode") or ""
|
||||
fname = ext_content.get("fileName", "")
|
||||
if dl_code:
|
||||
media_urls.append(dl_code)
|
||||
mime = "application/octet-stream"
|
||||
if fname:
|
||||
ext = fname.rsplit(".", 1)[-1].lower() if "." in fname else ""
|
||||
mime = EXT_MAP.get(ext, mime)
|
||||
media_types.append(mime)
|
||||
if msg_type == MessageType.TEXT:
|
||||
# Image messages, and files with image MIME (a .png sent as attachment), → PHOTO.
|
||||
if msg_type_str == "image" or mime.startswith("image/"):
|
||||
msg_type = MessageType.PHOTO
|
||||
else:
|
||||
msg_type = MessageType.DOCUMENT
|
||||
|
||||
return msg_type, media_urls, media_types
|
||||
|
||||
|
||||
def collect_download_codes(message: Any) -> List[Tuple[Any, str]]:
|
||||
"""Return ``(container, key)`` pairs whose download code should be resolved to a URL."""
|
||||
codes: List[Tuple[Any, str]] = []
|
||||
img_content = getattr(message, "image_content", None)
|
||||
if img_content and getattr(img_content, "download_code", None):
|
||||
codes.append((img_content, "download_code"))
|
||||
rich_text = getattr(message, "rich_text_content", None)
|
||||
if rich_text:
|
||||
for item in getattr(rich_text, "rich_text_list", []) or []:
|
||||
if isinstance(item, dict):
|
||||
for key in ("downloadCode", "pictureDownloadCode", "download_code"):
|
||||
if item.get(key):
|
||||
codes.append((item, key))
|
||||
if (getattr(message, "message_type", "") or "") in ("file", "image"):
|
||||
ext_content = _ext_content(message)
|
||||
if ext_content and ext_content.get("downloadCode"):
|
||||
codes.append((ext_content, "downloadCode"))
|
||||
return codes
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,201 @@
|
||||
"""Google Chat outbound text formatting and Cards v2 rendering.
|
||||
|
||||
Extracted from ``adapter.py``; the card dict shapes and key order here are the
|
||||
wire format and must stay byte-identical.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Any, Callable, Dict, List
|
||||
|
||||
# Invisible Unicode codepoints that render as tofu (□) in Google Chat's
|
||||
# restricted font stack: ZWS/ZWNJ/ZWJ, bidi marks, word joiner, BOM and
|
||||
# Variation Selectors (Chat ignores them and often shows a blank box).
|
||||
_INVISIBLE_RE = re.compile(
|
||||
"["
|
||||
"\u200b" # Zero-Width Space
|
||||
"\u200c" # Zero-Width Non-Joiner
|
||||
"\u200d" # Zero-Width Joiner (ZWJ)
|
||||
"\u200e\u200f" # LTR / RTL marks
|
||||
"\u2060" # Word Joiner
|
||||
"\ufeff" # BOM / Zero-Width No-Break Space
|
||||
"\ufe00-\ufe0f" # Variation Selectors 1-16 (VS1–VS16)
|
||||
"\U000e0100-\U000e01ef" # Variation Selectors 17-256
|
||||
"]"
|
||||
)
|
||||
|
||||
|
||||
def format_message(content: str) -> str:
|
||||
"""Convert standard Markdown to Google Chat's dialect.
|
||||
|
||||
Chat renders only ``*bold*``, ``_italic_``, ``~strike~`` and code; ``**bold**``,
|
||||
``# headers`` and ``[text](url)`` must be converted. Fenced and inline code
|
||||
are protected via placeholders so literal asterisks/brackets inside them
|
||||
survive; invisible tofu codepoints are stripped at the end.
|
||||
"""
|
||||
if not content:
|
||||
return content
|
||||
|
||||
text = content
|
||||
placeholders: Dict[str, str] = {}
|
||||
counter = [0]
|
||||
|
||||
def _ph(value: str) -> str:
|
||||
key = f"\x00GC{counter[0]}\x00"
|
||||
counter[0] += 1
|
||||
placeholders[key] = value
|
||||
return key
|
||||
|
||||
# Protect fenced blocks first, then inline code.
|
||||
text = re.sub(r"(```(?:[^\n]*\n)?[\s\S]*?```)", lambda m: _ph(m.group(0)), text)
|
||||
text = re.sub(r"(`[^`]+`)", lambda m: _ph(m.group(0)), text)
|
||||
# Headers (## Title) → *Title* (Chat has no header support).
|
||||
text = re.sub(r"^#{1,6}\s+(.+)$", lambda m: _ph(f"*{m.group(1).strip()}*"), text, flags=re.MULTILINE)
|
||||
# ***text*** → *_text_*, then **text** → *text*.
|
||||
text = re.sub(r"\*\*\*(.+?)\*\*\*", lambda m: _ph(f"*_{m.group(1)}_*"), text)
|
||||
text = re.sub(r"\*\*(.+?)\*\*", lambda m: _ph(f"*{m.group(1)}*"), text)
|
||||
# [text](url) → <url|text> (Slack-style angle-bracket).
|
||||
text = re.sub(r"\[([^\]]+)\]\(([^)]+)\)", lambda m: _ph(f"<{m.group(2)}|{m.group(1)}>"), text)
|
||||
text = _INVISIBLE_RE.sub("", text)
|
||||
# Collapse double spaces left over from stripped chars.
|
||||
text = re.sub(r" +", " ", text)
|
||||
for key, value in placeholders.items():
|
||||
text = text.replace(key, value)
|
||||
return text
|
||||
|
||||
|
||||
def _required_str(mapping: Dict[str, Any], key: str, context: str) -> str:
|
||||
value = mapping.get(key)
|
||||
if value is None:
|
||||
raise ValueError(f"{context}.{key} is required")
|
||||
value = str(value).strip()
|
||||
if not value:
|
||||
raise ValueError(f"{context}.{key} is required")
|
||||
return value
|
||||
|
||||
|
||||
def _button_to_chat(button: Dict[str, Any]) -> Dict[str, Any]:
|
||||
text = _required_str(button, "text", "button")
|
||||
action = _required_str(button, "action", "button")
|
||||
raw_params = button.get("parameters") or {}
|
||||
if not isinstance(raw_params, dict):
|
||||
raise ValueError("button.parameters must be an object")
|
||||
parameters = [{"key": str(key), "value": str(value)} for key, value in sorted(raw_params.items())]
|
||||
return {
|
||||
"text": text,
|
||||
"onClick": {"action": {"function": action, "parameters": parameters}},
|
||||
}
|
||||
|
||||
|
||||
def _text_widget(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {"textParagraph": {"text": format_message(_required_str(widget, "text", "widget"))}}
|
||||
|
||||
|
||||
def _decorated_text_widget(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
decorated: Dict[str, Any] = {
|
||||
"text": format_message(_required_str(widget, "text", "widget")),
|
||||
"wrapText": bool(widget.get("wrap_text", True)),
|
||||
}
|
||||
if widget.get("top_label"):
|
||||
decorated["topLabel"] = str(widget["top_label"])
|
||||
if widget.get("bottom_label"):
|
||||
decorated["bottomLabel"] = str(widget["bottom_label"])
|
||||
return {"decoratedText": decorated}
|
||||
|
||||
|
||||
def _image_widget(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
image = {"imageUrl": _required_str(widget, "image_url", "widget")}
|
||||
if widget.get("alt_text"):
|
||||
image["altText"] = str(widget["alt_text"])
|
||||
return {"image": image}
|
||||
|
||||
|
||||
def _buttons_widget(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
raw_buttons = widget.get("buttons") or []
|
||||
if not isinstance(raw_buttons, list) or not raw_buttons:
|
||||
raise ValueError("button widgets require at least one button")
|
||||
return {"buttonList": {"buttons": [_button_to_chat(btn) for btn in raw_buttons]}}
|
||||
|
||||
|
||||
def _selection_widget(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
name = _required_str(widget, "name", "widget")
|
||||
raw_items = widget.get("items") or []
|
||||
if not isinstance(raw_items, list) or not raw_items:
|
||||
raise ValueError("selection widgets require at least one item")
|
||||
items: List[Dict[str, Any]] = []
|
||||
for item in raw_items:
|
||||
if not isinstance(item, dict):
|
||||
raise ValueError("selection items must be objects")
|
||||
items.append({
|
||||
"text": _required_str(item, "text", "selection item"),
|
||||
"value": _required_str(item, "value", "selection item"),
|
||||
"selected": bool(item.get("selected", False)),
|
||||
})
|
||||
return {
|
||||
"selectionInput": {
|
||||
"name": name,
|
||||
"label": str(widget.get("label") or name),
|
||||
"type": str(widget.get("selection_type") or "CHECK_BOX"),
|
||||
"items": items,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
_WIDGET_RENDERERS: Dict[str, Callable[[Dict[str, Any]], Dict[str, Any]]] = {
|
||||
"text": _text_widget,
|
||||
"text_paragraph": _text_widget,
|
||||
"decorated_text": _decorated_text_widget,
|
||||
"buttons": _buttons_widget,
|
||||
"button_list": _buttons_widget,
|
||||
"selection": _selection_widget,
|
||||
"selection_input": _selection_widget,
|
||||
"image": _image_widget,
|
||||
"divider": lambda widget: {"divider": {}},
|
||||
}
|
||||
|
||||
|
||||
def _widget_to_chat(widget: Dict[str, Any]) -> Dict[str, Any]:
|
||||
if not isinstance(widget, dict):
|
||||
raise ValueError("card widgets must be objects")
|
||||
widget_type = str(widget.get("type") or "").strip()
|
||||
renderer = _WIDGET_RENDERERS.get(widget_type)
|
||||
if renderer is None:
|
||||
raise ValueError(f"unsupported widget type: {widget_type or '<missing>'}")
|
||||
return renderer(widget)
|
||||
|
||||
|
||||
def card_spec_to_cards_v2(card_spec: Dict[str, Any]) -> Dict[str, Any]:
|
||||
if not isinstance(card_spec, dict):
|
||||
raise ValueError("card must be an object")
|
||||
raw_sections = card_spec.get("sections") or []
|
||||
if not isinstance(raw_sections, list) or not raw_sections:
|
||||
raise ValueError("card.sections must contain at least one section")
|
||||
|
||||
sections: List[Dict[str, Any]] = []
|
||||
for section in raw_sections:
|
||||
if not isinstance(section, dict):
|
||||
raise ValueError("card sections must be objects")
|
||||
widgets = section.get("widgets") or []
|
||||
if not isinstance(widgets, list) or not widgets:
|
||||
raise ValueError("card section widgets must contain at least one widget")
|
||||
rendered: Dict[str, Any] = {"widgets": [_widget_to_chat(w) for w in widgets]}
|
||||
if section.get("header"):
|
||||
rendered["header"] = str(section["header"])
|
||||
sections.append(rendered)
|
||||
|
||||
card: Dict[str, Any] = {"sections": sections}
|
||||
header = card_spec.get("header")
|
||||
if header:
|
||||
if not isinstance(header, dict):
|
||||
raise ValueError("card.header must be an object")
|
||||
rendered_header: Dict[str, Any] = {"title": _required_str(header, "title", "card.header")}
|
||||
if header.get("subtitle"):
|
||||
rendered_header["subtitle"] = str(header["subtitle"])
|
||||
if header.get("image_url"):
|
||||
rendered_header["imageUrl"] = str(header["image_url"])
|
||||
rendered_header["imageType"] = str(header.get("image_type") or "SQUARE")
|
||||
if header.get("image_alt_text"):
|
||||
rendered_header["imageAltText"] = str(header["image_alt_text"])
|
||||
card["header"] = rendered_header
|
||||
return {"cardId": str(card_spec.get("card_id") or "hermes-card"), "card": card}
|
||||
@@ -1,57 +1,26 @@
|
||||
"""User OAuth helper for the Google Chat gateway adapter.
|
||||
|
||||
Google Chat's ``media.upload`` REST endpoint hard-rejects service-account
|
||||
authentication:
|
||||
Google Chat's ``media.upload`` hard-rejects service-account auth ("This method
|
||||
doesn't support app authentication with a service account"), so for native
|
||||
file attachments each user grants the bot ``chat.messages.create`` ONCE in
|
||||
their own DM. The bot stores per-user refresh tokens and uploads *as the user*.
|
||||
See https://developers.google.com/chat/api/guides/auth/users.
|
||||
|
||||
"This method doesn't support app authentication with a service
|
||||
account. Authenticate with a user account."
|
||||
Both a library (imported by the adapter) and a CLI (driven by ``/setup-files``):
|
||||
|
||||
(See https://developers.google.com/workspace/chat/api/reference/rest/v1/media/upload
|
||||
and https://developers.google.com/chat/api/guides/auth/users.)
|
||||
|
||||
For the bot to deliver native file attachments — the same drag-and-drop
|
||||
file widget the user gets when they upload manually — each user must
|
||||
grant the bot the ``chat.messages.create`` scope ONCE in their own DM.
|
||||
The bot stores per-user refresh tokens and calls ``media.upload`` plus
|
||||
the subsequent ``messages.create`` *as the requesting user* whenever a
|
||||
file needs sending.
|
||||
|
||||
This module is BOTH a CLI tool (driven by the agent via slash commands or
|
||||
terminal commands) AND a library imported by ``google_chat.py``:
|
||||
|
||||
Library functions (called from the adapter at runtime):
|
||||
load_user_credentials(email=None) -> Credentials | None
|
||||
refresh_or_none(creds, email=None) -> Credentials | None
|
||||
build_user_chat_service(creds) -> chat_v1.Resource
|
||||
list_authorized_emails() -> List[str]
|
||||
|
||||
CLI commands (driven by the agent through the /setup-files slash
|
||||
command, modeled on skills/productivity/google-workspace/scripts/setup.py):
|
||||
--check Exit 0 if auth is valid, else 1
|
||||
--client-secret /path/to.json Persist OAuth client credentials
|
||||
--auth-url Print the OAuth URL for the user
|
||||
--auth-code CODE Exchange auth code for token
|
||||
--revoke Revoke and delete stored token
|
||||
--install-deps Install Python dependencies
|
||||
--email EMAIL Scope CLI ops to a specific user
|
||||
(defaults to legacy single-user
|
||||
mode when omitted)
|
||||
|
||||
The flow mirrors the existing google-workspace skill exactly so anyone
|
||||
familiar with that flow can read this without surprises.
|
||||
Library: load_user_credentials(email=None), refresh_or_none(creds, email=None),
|
||||
build_user_chat_service(creds), list_authorized_emails()
|
||||
CLI: --check | --client-secret PATH | --auth-url | --auth-code CODE |
|
||||
--revoke | --install-deps [--email EMAIL] (legacy single-user
|
||||
mode when --email is omitted)
|
||||
|
||||
Token storage layout
|
||||
--------------------
|
||||
- Per-user tokens (keyed by sender email):
|
||||
``${HERMES_HOME}/google_chat_user_tokens/<sanitized_email>.json``
|
||||
- Legacy single-user token (fallback, untouched for backward compat):
|
||||
``${HERMES_HOME}/google_chat_user_token.json``
|
||||
- Per-user pending OAuth state during /setup-files start → exchange:
|
||||
``${HERMES_HOME}/google_chat_user_oauth_pending/<sanitized_email>.json``
|
||||
- Legacy pending state:
|
||||
``${HERMES_HOME}/google_chat_user_oauth_pending.json``
|
||||
- OAuth client secret (profile-scoped — each profile registers its own):
|
||||
``${HERMES_HOME}/google_chat_user_client_secret.json``
|
||||
- Per-user tokens: ``${HERMES_HOME}/google_chat_user_tokens/<sanitized_email>.json``
|
||||
- Legacy single-user: ``${HERMES_HOME}/google_chat_user_token.json``
|
||||
- Per-user pending PKCE state: ``${HERMES_HOME}/google_chat_user_oauth_pending/<sanitized_email>.json``
|
||||
- Legacy pending state: ``${HERMES_HOME}/google_chat_user_oauth_pending.json``
|
||||
- OAuth client secret (profile-scoped): ``${HERMES_HOME}/google_chat_user_client_secret.json``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -63,26 +32,20 @@ import os
|
||||
import re
|
||||
import secrets
|
||||
import stat
|
||||
import subprocess
|
||||
import sys
|
||||
from importlib.metadata import version as _distribution_version
|
||||
from pathlib import Path
|
||||
from typing import Any, List, Optional, Tuple
|
||||
from typing import Any, List, NoReturn, Optional, Tuple
|
||||
|
||||
from packaging.requirements import Requirement
|
||||
|
||||
# Pin the legacy logger name so operator-side log filters keep matching
|
||||
# after the in-tree → plugin migration. See adapter.py for context.
|
||||
# Pinned legacy logger name so operator log filters keep matching (see adapter.py).
|
||||
logger = logging.getLogger("gateway.platforms.google_chat_user_oauth")
|
||||
|
||||
# Use the project's HERMES_HOME helper so the token follows the user's
|
||||
# profile (e.g. tests can override via HERMES_HOME=/tmp/...).
|
||||
try:
|
||||
from hermes_constants import display_hermes_home, get_hermes_home
|
||||
except (ModuleNotFoundError, ImportError):
|
||||
# Fallback for environments where hermes_constants isn't importable
|
||||
# (mirrors the same fallback used by the google-workspace skill's
|
||||
# _hermes_home.py shim).
|
||||
# Mirrors the google-workspace skill's _hermes_home.py shim.
|
||||
def get_hermes_home() -> Path:
|
||||
val = os.environ.get("HERMES_HOME", "").strip()
|
||||
return Path(val) if val else Path.home() / ".hermes"
|
||||
@@ -98,19 +61,12 @@ from utils import atomic_replace
|
||||
|
||||
|
||||
def _hermes_home() -> Path:
|
||||
"""Resolve HERMES_HOME at call time (NOT module import).
|
||||
|
||||
Tests and ``HERMES_HOME=...`` env overrides need this to be late-
|
||||
binding. If we cached the path at import time, switching profiles
|
||||
or tweaking env vars in tests would silently keep using the old
|
||||
path."""
|
||||
"""Resolve HERMES_HOME at call time (late-binding for tests / profile switches)."""
|
||||
return get_hermes_home()
|
||||
|
||||
|
||||
# Filesystem-safe key: lowercase, allow ``[a-z0-9._-@]``, replace anything
|
||||
# else with ``_``. ``ramon.fernandez@nttdata.com`` stays human-readable
|
||||
# (``ramon.fernandez@nttdata.com.json``) which makes admin debugging by
|
||||
# ``ls ~/.hermes/google_chat_user_tokens/`` trivial.
|
||||
# Filesystem-safe key: lowercase, keep ``[a-z0-9._-@]`` so token files stay
|
||||
# human-readable under ``ls ~/.hermes/google_chat_user_tokens/``.
|
||||
_EMAIL_FS_RE = re.compile(r"[^a-z0-9._@-]+")
|
||||
|
||||
|
||||
@@ -119,27 +75,15 @@ def _sanitize_email(email: str) -> str:
|
||||
return cleaned or "_unknown_"
|
||||
|
||||
|
||||
def _legacy_token_path() -> Path:
|
||||
return _hermes_home() / "google_chat_user_token.json"
|
||||
|
||||
|
||||
def _user_tokens_dir() -> Path:
|
||||
return _hermes_home() / "google_chat_user_tokens"
|
||||
|
||||
|
||||
def _legacy_pending_path() -> Path:
|
||||
return _hermes_home() / "google_chat_user_oauth_pending.json"
|
||||
|
||||
|
||||
def _user_pending_dir() -> Path:
|
||||
return _hermes_home() / "google_chat_user_oauth_pending"
|
||||
|
||||
|
||||
def _token_path(email: Optional[str] = None) -> Path:
|
||||
"""Return the on-disk token path for ``email`` or the legacy path."""
|
||||
"""Per-user token path for ``email``, or the legacy single-user path."""
|
||||
if email:
|
||||
return _user_tokens_dir() / f"{_sanitize_email(email)}.json"
|
||||
return _legacy_token_path()
|
||||
return _hermes_home() / "google_chat_user_token.json"
|
||||
|
||||
|
||||
def _client_secret_path() -> Path:
|
||||
@@ -148,14 +92,12 @@ def _client_secret_path() -> Path:
|
||||
|
||||
def _pending_auth_path(email: Optional[str] = None) -> Path:
|
||||
if email:
|
||||
return _user_pending_dir() / f"{_sanitize_email(email)}.json"
|
||||
return _legacy_pending_path()
|
||||
return _hermes_home() / "google_chat_user_oauth_pending" / f"{_sanitize_email(email)}.json"
|
||||
return _hermes_home() / "google_chat_user_oauth_pending.json"
|
||||
|
||||
|
||||
# Minimum scope for native Chat attachment delivery.
|
||||
# `chat.messages.create` covers BOTH `media.upload` and the subsequent
|
||||
# `messages.create` that references the attachmentDataRef. We deliberately
|
||||
# do NOT request drive.file or other scopes — least privilege.
|
||||
# Least privilege: chat.messages.create covers BOTH media.upload and the
|
||||
# subsequent messages.create; no drive.file or other scopes.
|
||||
SCOPES: List[str] = [
|
||||
"https://www.googleapis.com/auth/chat.messages.create",
|
||||
]
|
||||
@@ -171,10 +113,8 @@ _REQUIRED_PACKAGES = [
|
||||
"pyasn1==0.6.4",
|
||||
]
|
||||
|
||||
# Out-of-band redirect: Google deprecated the ``urn:ietf:wg:oauth:2.0:oob``
|
||||
# flow, so we use a localhost redirect that's expected to FAIL. The user
|
||||
# copies the auth code from the failed browser URL bar back into chat.
|
||||
# Same trick used by skills/productivity/google-workspace/scripts/setup.py.
|
||||
# Google deprecated the ``oob`` flow: use a localhost redirect that is expected
|
||||
# to FAIL; the user pastes the code from the failed browser URL back into chat.
|
||||
_REDIRECT_URI = "http://localhost:1"
|
||||
|
||||
|
||||
@@ -183,30 +123,37 @@ _REDIRECT_URI = "http://localhost:1"
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _refresh_and_persist(creds: Any, token_path: Path, request_cls: Any) -> Optional[Any]:
|
||||
"""Refresh expired creds and write them back; None when unusable or refresh fails."""
|
||||
if creds.valid:
|
||||
return creds
|
||||
if creds.expired and creds.refresh_token:
|
||||
try:
|
||||
creds.refresh(request_cls())
|
||||
except Exception as exc:
|
||||
logger.warning("[google_chat_user_oauth] token refresh failed (user should re-run /setup-files): %s", exc)
|
||||
return None
|
||||
_persist_credentials(creds, token_path)
|
||||
return creds
|
||||
# Token exists but is unusable (e.g. revoked, no refresh token).
|
||||
return None
|
||||
|
||||
|
||||
def load_user_credentials(email: Optional[str] = None) -> Optional[Any]:
|
||||
"""Load + validate persisted user OAuth credentials.
|
||||
|
||||
``email`` selects the per-user token file; ``None`` falls back to the
|
||||
legacy single-user path (left in place for installs that ran the
|
||||
pre-multi-user flow). Returns a ``google.oauth2.credentials.Credentials``
|
||||
instance ready for use, or ``None`` if no token is stored, the token
|
||||
is corrupt, or refresh fails. Adapter callers should treat ``None``
|
||||
as "user has not run /setup-files yet" and surface the setup-instructions
|
||||
fallback to the user.
|
||||
|
||||
Does NOT raise on the no-token case — that's expected.
|
||||
``None`` email → legacy single-user path. Returns ``None`` (never raises) when
|
||||
no token is stored, the token is corrupt, or refresh fails — callers treat
|
||||
that as "user has not run /setup-files yet".
|
||||
"""
|
||||
token_path = _token_path(email)
|
||||
if not token_path.exists():
|
||||
return None
|
||||
|
||||
# Same class as slack_tokens.json: hand-provisioned or legacy-written
|
||||
# token files commonly end up 0o644. Warn so the owner tightens them.
|
||||
# Hand-provisioned / legacy token files commonly end up 0o644; warn the owner.
|
||||
from utils import warn_if_credential_file_broadly_readable
|
||||
|
||||
warn_if_credential_file_broadly_readable(
|
||||
token_path, label="[google_chat_user_oauth]", log=logger
|
||||
)
|
||||
warn_if_credential_file_broadly_readable(token_path, label="[google_chat_user_oauth]", log=logger)
|
||||
|
||||
try:
|
||||
from google.oauth2.credentials import Credentials
|
||||
@@ -219,113 +166,58 @@ def load_user_credentials(email: Optional[str] = None) -> Optional[Any]:
|
||||
return None
|
||||
|
||||
try:
|
||||
# Don't pass scopes — user may have authorized only a subset, and
|
||||
# passing scopes makes refresh validate them strictly. Same logic
|
||||
# as the google-workspace skill.
|
||||
# No scopes: the user may have authorized a subset, and passing scopes
|
||||
# makes refresh validate them strictly.
|
||||
creds = Credentials.from_authorized_user_file(str(token_path))
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[google_chat_user_oauth] token at %s is corrupt: %s",
|
||||
token_path, exc,
|
||||
)
|
||||
logger.warning("[google_chat_user_oauth] token at %s is corrupt: %s", token_path, exc)
|
||||
return None
|
||||
|
||||
if creds.valid:
|
||||
return creds
|
||||
|
||||
if creds.expired and creds.refresh_token:
|
||||
try:
|
||||
creds.refresh(Request())
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[google_chat_user_oauth] token refresh failed (user "
|
||||
"should re-run /setup-files): %s", exc,
|
||||
)
|
||||
return None
|
||||
# Persist refreshed token so next start picks up the new access
|
||||
# token without an unnecessary refresh round-trip.
|
||||
_persist_credentials(creds, token_path)
|
||||
return creds
|
||||
|
||||
# Token exists but is unusable (e.g. revoked, no refresh token).
|
||||
return None
|
||||
return _refresh_and_persist(creds, token_path, Request)
|
||||
|
||||
|
||||
def refresh_or_none(creds: Any, email: Optional[str] = None) -> Optional[Any]:
|
||||
"""Refresh ``creds`` if expired. Returns the credentials or ``None``.
|
||||
|
||||
Used by the adapter just before calling media.upload to ensure the
|
||||
token is current. Returns ``None`` if refresh fails — caller falls
|
||||
back to the text-notice path. ``email`` controls where the refreshed
|
||||
token is written back; ``None`` keeps the legacy single-file path.
|
||||
"""
|
||||
"""Refresh ``creds`` if expired; ``None`` on failure (caller falls back to the
|
||||
text-notice path). ``email`` selects where the refreshed token is written."""
|
||||
if creds is None:
|
||||
return None
|
||||
|
||||
if creds.valid:
|
||||
return creds
|
||||
|
||||
try:
|
||||
from google.auth.transport.requests import Request
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
if creds.expired and creds.refresh_token:
|
||||
try:
|
||||
creds.refresh(Request())
|
||||
_persist_credentials(creds, _token_path(email))
|
||||
return creds
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[google_chat_user_oauth] refresh failed: %s", exc,
|
||||
)
|
||||
logger.warning("[google_chat_user_oauth] refresh failed: %s", exc)
|
||||
return None
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def build_user_chat_service(creds: Any) -> Any:
|
||||
"""Build a Google Chat API client authenticated as the user.
|
||||
|
||||
Used for media.upload + the subsequent messages.create that
|
||||
references the attachmentDataRef. The bot's separate SA-authed
|
||||
client (``self._chat_api`` in the adapter) is for everything else.
|
||||
"""
|
||||
"""Chat API client authenticated as the user (for media.upload + messages.create)."""
|
||||
from googleapiclient.discovery import build as build_service
|
||||
return build_service("chat", "v1", credentials=creds, cache_discovery=False)
|
||||
|
||||
|
||||
def list_authorized_emails() -> List[str]:
|
||||
"""Return the set of user emails that have stored per-user tokens.
|
||||
|
||||
Lists files in the per-user tokens dir; does NOT include the legacy
|
||||
single-user token (its owner is unknown). Sanitized filenames lose
|
||||
the ``+suffix`` part of plus-addressed emails — accept that and use
|
||||
this list only for admin display, not for trust decisions.
|
||||
"""
|
||||
"""Sanitized emails with stored per-user tokens (admin display only, not trust;
|
||||
excludes the legacy single-user token whose owner is unknown)."""
|
||||
d = _user_tokens_dir()
|
||||
if not d.exists():
|
||||
return []
|
||||
out: List[str] = []
|
||||
for f in d.iterdir():
|
||||
if f.is_file() and f.suffix == ".json":
|
||||
out.append(f.stem)
|
||||
out.sort()
|
||||
return out
|
||||
return sorted(f.stem for f in d.iterdir() if f.is_file() and f.suffix == ".json")
|
||||
|
||||
|
||||
def _persist_credentials(creds: Any, token_path: Path) -> None:
|
||||
"""Persist refreshed credentials atomically with private permissions."""
|
||||
try:
|
||||
_write_private_json(
|
||||
token_path,
|
||||
_normalize_authorized_user_payload(json.loads(creds.to_json())),
|
||||
)
|
||||
_write_private_json(token_path, _normalize_authorized_user_payload(json.loads(creds.to_json())))
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"[google_chat_user_oauth] failed to persist credentials at %s",
|
||||
token_path, exc_info=True,
|
||||
)
|
||||
logger.debug("[google_chat_user_oauth] failed to persist credentials at %s", token_path, exc_info=True)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
@@ -351,11 +243,7 @@ def _write_private_json(path: Path, data: Any) -> None:
|
||||
|
||||
tmp_path = path.with_suffix(f".tmp.{os.getpid()}.{secrets.token_hex(4)}")
|
||||
try:
|
||||
fd = os.open(
|
||||
str(tmp_path),
|
||||
os.O_WRONLY | os.O_CREAT | os.O_EXCL,
|
||||
stat.S_IRUSR | stat.S_IWUSR,
|
||||
)
|
||||
fd = os.open(str(tmp_path), os.O_WRONLY | os.O_CREAT | os.O_EXCL, stat.S_IRUSR | stat.S_IWUSR)
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as fh:
|
||||
json.dump(data, fh, indent=2, ensure_ascii=False)
|
||||
fh.flush()
|
||||
@@ -373,6 +261,13 @@ def _write_private_json(path: Path, data: Any) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _fail(*lines: str) -> NoReturn:
|
||||
"""Print CLI error lines and exit 1."""
|
||||
for line in lines:
|
||||
print(line)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def _ensure_deps() -> None:
|
||||
"""Check exact dependency versions; install if stale; exit on failure."""
|
||||
if _missing_required_packages() and not install_deps():
|
||||
@@ -409,9 +304,7 @@ def install_deps() -> bool:
|
||||
raise RuntimeError((result.stderr or "install failed").strip()[:300])
|
||||
remaining = _missing_required_packages()
|
||||
if remaining:
|
||||
raise RuntimeError(
|
||||
"dependencies remain stale after install: " + " ".join(remaining)
|
||||
)
|
||||
raise RuntimeError("dependencies remain stale after install: " + " ".join(remaining))
|
||||
print("Dependencies installed.")
|
||||
return True
|
||||
except Exception as exc:
|
||||
@@ -421,20 +314,14 @@ def install_deps() -> bool:
|
||||
|
||||
|
||||
def check_auth(email: Optional[str] = None) -> bool:
|
||||
"""Print status; return True if creds are usable.
|
||||
|
||||
Per-user when ``email`` given, legacy single-user when omitted.
|
||||
"""
|
||||
"""Print status; return True if creds are usable."""
|
||||
token_path = _token_path(email)
|
||||
if not token_path.exists():
|
||||
print(f"NOT_AUTHENTICATED: No token at {token_path}")
|
||||
return False
|
||||
|
||||
creds = load_user_credentials(email)
|
||||
if creds is None:
|
||||
if load_user_credentials(email) is None:
|
||||
print(f"TOKEN_INVALID: Re-run /setup-files (path: {token_path})")
|
||||
return False
|
||||
|
||||
print(f"AUTHENTICATED: Token valid at {token_path}")
|
||||
return True
|
||||
|
||||
@@ -443,35 +330,24 @@ def store_client_secret(path: str) -> None:
|
||||
"""Validate and copy the user's OAuth client_secret.json into HERMES_HOME."""
|
||||
src = Path(path).expanduser().resolve()
|
||||
if not src.exists():
|
||||
print(f"ERROR: File not found: {src}")
|
||||
sys.exit(1)
|
||||
|
||||
_fail(f"ERROR: File not found: {src}")
|
||||
try:
|
||||
data = json.loads(src.read_text(encoding="utf-8"))
|
||||
except json.JSONDecodeError:
|
||||
print("ERROR: File is not valid JSON.")
|
||||
sys.exit(1)
|
||||
|
||||
_fail("ERROR: File is not valid JSON.")
|
||||
if "installed" not in data and "web" not in data:
|
||||
print(
|
||||
"ERROR: Not a Google OAuth client secret file (missing "
|
||||
"'installed' or 'web' key)."
|
||||
_fail(
|
||||
"ERROR: Not a Google OAuth client secret file (missing 'installed' or 'web' key).",
|
||||
"Download from: https://console.cloud.google.com/apis/credentials",
|
||||
)
|
||||
print(
|
||||
"Download from: https://console.cloud.google.com/apis/credentials"
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
target = _client_secret_path()
|
||||
_write_private_json(target, data)
|
||||
print(f"OK: Client secret saved to {target}")
|
||||
|
||||
|
||||
def _save_pending_auth(*, state: str, code_verifier: str,
|
||||
email: Optional[str] = None) -> None:
|
||||
pending = _pending_auth_path(email)
|
||||
def _save_pending_auth(*, state: str, code_verifier: str, email: Optional[str] = None) -> None:
|
||||
_write_private_json(
|
||||
pending,
|
||||
_pending_auth_path(email),
|
||||
{
|
||||
"state": state,
|
||||
"code_verifier": code_verifier,
|
||||
@@ -484,18 +360,13 @@ def _save_pending_auth(*, state: str, code_verifier: str,
|
||||
def _load_pending_auth(email: Optional[str] = None) -> dict:
|
||||
pending = _pending_auth_path(email)
|
||||
if not pending.exists():
|
||||
print("ERROR: No pending OAuth session found. Run --auth-url first.")
|
||||
sys.exit(1)
|
||||
_fail("ERROR: No pending OAuth session found. Run --auth-url first.")
|
||||
try:
|
||||
data = json.loads(pending.read_text(encoding="utf-8"))
|
||||
except Exception as exc:
|
||||
print(f"ERROR: Could not read pending OAuth session: {exc}")
|
||||
print("Run --auth-url again to start a fresh session.")
|
||||
sys.exit(1)
|
||||
_fail(f"ERROR: Could not read pending OAuth session: {exc}", "Run --auth-url again to start a fresh session.")
|
||||
if not data.get("state") or not data.get("code_verifier"):
|
||||
print("ERROR: Pending OAuth session is missing PKCE data.")
|
||||
print("Run --auth-url again.")
|
||||
sys.exit(1)
|
||||
_fail("ERROR: Pending OAuth session is missing PKCE data.", "Run --auth-url again.")
|
||||
return data
|
||||
|
||||
|
||||
@@ -506,25 +377,21 @@ def _extract_code_and_state(code_or_url: str) -> Tuple[str, Optional[str]]:
|
||||
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
parsed = urlparse(code_or_url)
|
||||
params = parse_qs(parsed.query)
|
||||
params = parse_qs(urlparse(code_or_url).query)
|
||||
if "code" not in params:
|
||||
print("ERROR: No 'code' parameter found in URL.")
|
||||
sys.exit(1)
|
||||
state = params.get("state", [None])[0]
|
||||
return params["code"][0], state
|
||||
_fail("ERROR: No 'code' parameter found in URL.")
|
||||
return params["code"][0], params.get("state", [None])[0]
|
||||
|
||||
|
||||
def _require_client_secret() -> None:
|
||||
if not _client_secret_path().exists():
|
||||
_fail("ERROR: No client secret stored. Run --client-secret first.")
|
||||
|
||||
|
||||
def get_auth_url(email: Optional[str] = None) -> None:
|
||||
"""Print the OAuth URL for the user to visit. Persists PKCE state.
|
||||
|
||||
``email`` namespaces the pending state so two users can be mid-flow
|
||||
in parallel without trampling each other's PKCE verifier.
|
||||
"""
|
||||
if not _client_secret_path().exists():
|
||||
print("ERROR: No client secret stored. Run --client-secret first.")
|
||||
sys.exit(1)
|
||||
|
||||
"""Print the OAuth URL for the user to visit; persists PKCE state under ``email``
|
||||
so two users can be mid-flow in parallel."""
|
||||
_require_client_secret()
|
||||
_ensure_deps()
|
||||
from google_auth_oauthlib.flow import Flow
|
||||
|
||||
@@ -534,34 +401,20 @@ def get_auth_url(email: Optional[str] = None) -> None:
|
||||
redirect_uri=_REDIRECT_URI,
|
||||
autogenerate_code_verifier=True,
|
||||
)
|
||||
auth_url, state = flow.authorization_url(
|
||||
access_type="offline",
|
||||
prompt="consent",
|
||||
)
|
||||
auth_url, state = flow.authorization_url(access_type="offline", prompt="consent")
|
||||
_save_pending_auth(state=state, code_verifier=flow.code_verifier, email=email)
|
||||
print(auth_url)
|
||||
|
||||
|
||||
def exchange_auth_code(code: str, email: Optional[str] = None) -> None:
|
||||
"""Exchange an auth code (or pasted redirect URL) for a refresh token.
|
||||
|
||||
``email`` selects the destination token path. ``None`` writes to the
|
||||
legacy single-user path (kept for the existing CLI entrypoint and for
|
||||
pre-multi-user installs).
|
||||
"""
|
||||
if not _client_secret_path().exists():
|
||||
print("ERROR: No client secret stored. Run --client-secret first.")
|
||||
sys.exit(1)
|
||||
|
||||
"""Exchange an auth code (or pasted redirect URL) for a refresh token stored
|
||||
at the per-user path for ``email`` (legacy single-user path when None)."""
|
||||
_require_client_secret()
|
||||
pending_auth = _load_pending_auth(email)
|
||||
raw_callback = code
|
||||
code, returned_state = _extract_code_and_state(code)
|
||||
if returned_state and returned_state != pending_auth["state"]:
|
||||
print(
|
||||
"ERROR: OAuth state mismatch. Run --auth-url again to start a "
|
||||
"fresh session."
|
||||
)
|
||||
sys.exit(1)
|
||||
_fail("ERROR: OAuth state mismatch. Run --auth-url again to start a fresh session.")
|
||||
|
||||
_ensure_deps()
|
||||
from google_auth_oauthlib.flow import Flow
|
||||
@@ -581,24 +434,16 @@ def exchange_auth_code(code: str, email: Optional[str] = None) -> None:
|
||||
state=pending_auth["state"],
|
||||
code_verifier=pending_auth["code_verifier"],
|
||||
)
|
||||
|
||||
try:
|
||||
# Accept partial scopes — user may deselect items in the consent screen.
|
||||
os.environ["OAUTHLIB_RELAX_TOKEN_SCOPE"] = "1"
|
||||
flow.fetch_token(code=code)
|
||||
except Exception as exc:
|
||||
print(f"ERROR: Token exchange failed: {exc}")
|
||||
print("The code may have expired. Run --auth-url to get a fresh URL.")
|
||||
sys.exit(1)
|
||||
_fail(f"ERROR: Token exchange failed: {exc}", "The code may have expired. Run --auth-url to get a fresh URL.")
|
||||
|
||||
creds = flow.credentials
|
||||
token_payload = _normalize_authorized_user_payload(json.loads(creds.to_json()))
|
||||
|
||||
actually_granted = (
|
||||
list(creds.granted_scopes or [])
|
||||
if hasattr(creds, "granted_scopes") and creds.granted_scopes
|
||||
else []
|
||||
)
|
||||
actually_granted = list(creds.granted_scopes or []) if hasattr(creds, "granted_scopes") and creds.granted_scopes else []
|
||||
if actually_granted:
|
||||
token_payload["scopes"] = actually_granted
|
||||
elif granted_scopes != SCOPES:
|
||||
@@ -618,10 +463,7 @@ def exchange_auth_code(code: str, email: Optional[str] = None) -> None:
|
||||
|
||||
|
||||
def revoke(email: Optional[str] = None) -> None:
|
||||
"""Revoke the stored token with Google and delete it locally.
|
||||
|
||||
Per-user when ``email`` given, legacy single-user when omitted.
|
||||
"""
|
||||
"""Revoke the stored token with Google and delete it locally."""
|
||||
token_path = _token_path(email)
|
||||
if not token_path.exists():
|
||||
print("No token to revoke.")
|
||||
@@ -659,21 +501,16 @@ def main() -> None:
|
||||
description="Google Chat user-OAuth setup for Hermes (native attachment delivery)"
|
||||
)
|
||||
group = parser.add_mutually_exclusive_group(required=True)
|
||||
group.add_argument("--check", action="store_true",
|
||||
help="Check if auth is valid (exit 0=yes, 1=no)")
|
||||
group.add_argument("--client-secret", metavar="PATH",
|
||||
help="Store OAuth client_secret.json")
|
||||
group.add_argument("--auth-url", action="store_true",
|
||||
help="Print OAuth URL for user to visit")
|
||||
group.add_argument("--auth-code", metavar="CODE",
|
||||
help="Exchange auth code for token")
|
||||
group.add_argument("--revoke", action="store_true",
|
||||
help="Revoke and delete stored token")
|
||||
group.add_argument("--install-deps", action="store_true",
|
||||
help="Install Python dependencies")
|
||||
parser.add_argument("--email", metavar="EMAIL", default=None,
|
||||
help="Scope operation to a specific user's token "
|
||||
"(default: legacy single-user path)")
|
||||
group.add_argument("--check", action="store_true", help="Check if auth is valid (exit 0=yes, 1=no)")
|
||||
group.add_argument("--client-secret", metavar="PATH", help="Store OAuth client_secret.json")
|
||||
group.add_argument("--auth-url", action="store_true", help="Print OAuth URL for user to visit")
|
||||
group.add_argument("--auth-code", metavar="CODE", help="Exchange auth code for token")
|
||||
group.add_argument("--revoke", action="store_true", help="Revoke and delete stored token")
|
||||
group.add_argument("--install-deps", action="store_true", help="Install Python dependencies")
|
||||
parser.add_argument(
|
||||
"--email", metavar="EMAIL", default=None,
|
||||
help="Scope operation to a specific user's token (default: legacy single-user path)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
email = args.email or None
|
||||
|
||||
@@ -0,0 +1,170 @@
|
||||
"""``/setup-files`` in-chat OAuth setup flow for native attachment delivery.
|
||||
|
||||
Extracted from ``adapter.py``: ``GoogleChatAdapter._handle_setup_files_command``
|
||||
delegates here. Logs under the adapter's pinned logger name.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import io
|
||||
import logging
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
logger = logging.getLogger("gateway.platforms.google_chat")
|
||||
|
||||
_NOT_CONFIGURED_TEXT = (
|
||||
"🔧 Native attachment delivery is **not configured**.\n"
|
||||
"**Step 1 (one-time, on the host):** create OAuth client credentials at "
|
||||
"https://console.cloud.google.com/apis/credentials → *Create credentials* → "
|
||||
"*OAuth client ID* → *Desktop app*. Download the JSON. Then on the host run:\n"
|
||||
"```\npython -m plugins.platforms.google_chat.oauth --client-secret /path/to/client_secret.json\n```\n"
|
||||
"**Step 2:** come back here and send `/setup-files start`."
|
||||
)
|
||||
_START_INSTRUCTIONS = (
|
||||
"1. Open this URL in your browser and authorize:\n{auth_url}\n\n"
|
||||
"2. After clicking *Allow*, your browser will fail to load "
|
||||
"`http://localhost:1/?...&code=...`. That's expected.\n\n"
|
||||
"3. Copy the entire failed URL from the browser's URL bar and paste it back here as: "
|
||||
"`/setup-files <PASTE_URL>` (or just the `code=...` value).\n\n"
|
||||
"Tip: the URL contains your access grant — keep it private."
|
||||
)
|
||||
|
||||
|
||||
async def _run_captured(fn: Callable[..., Any], *args: Any) -> str:
|
||||
"""Run ``fn`` in a thread with stdout captured (the oauth helpers print their output)."""
|
||||
buf = io.StringIO()
|
||||
with contextlib.redirect_stdout(buf):
|
||||
await asyncio.to_thread(fn, *args)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
async def handle_setup_files_command(
|
||||
adapter: Any,
|
||||
chat_id: str,
|
||||
thread_id: Optional[str],
|
||||
raw_text: str,
|
||||
sender_email: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Run the in-chat OAuth setup flow. Returns True when the message was consumed.
|
||||
|
||||
``sender_email`` is the per-user OAuth key; ``None`` falls back to the legacy
|
||||
single-user token slot so pre-multi-user installs keep working.
|
||||
Subcommands: ``/setup-files`` (status), ``start`` (OAuth URL), ``revoke``,
|
||||
``<CODE_OR_URL>`` (exchange). Requires client_secret.json on the host.
|
||||
"""
|
||||
from . import oauth as oauth_helper
|
||||
|
||||
# Same normalization as the token-path sanitizer so cache lookups stay consistent.
|
||||
sender_key = sender_email.strip().lower() if sender_email else None
|
||||
parts = raw_text.split(maxsplit=1)
|
||||
arg = parts[1].strip() if len(parts) > 1 else ""
|
||||
|
||||
async def _reply(text: str) -> None:
|
||||
body: Dict[str, Any] = {"text": text}
|
||||
if thread_id:
|
||||
body["thread"] = {"name": thread_id}
|
||||
try:
|
||||
await adapter._create_message(chat_id, body)
|
||||
except Exception:
|
||||
logger.debug("[GoogleChat] /setup-files reply send failed", exc_info=True)
|
||||
|
||||
if not arg:
|
||||
client_secret_present = oauth_helper._client_secret_path().exists()
|
||||
token_path = oauth_helper._token_path(sender_key)
|
||||
creds = oauth_helper.load_user_credentials(sender_key) if token_path.exists() else None
|
||||
if creds is not None:
|
||||
who = sender_key or "shared (legacy)"
|
||||
await _reply(
|
||||
f"✅ Native attachment delivery is **active** for `{who}`.\n"
|
||||
f"Token: `{token_path}`\nSend `/setup-files revoke` to disable."
|
||||
)
|
||||
elif not client_secret_present:
|
||||
await _reply(_NOT_CONFIGURED_TEXT)
|
||||
else:
|
||||
await _reply(
|
||||
"🔧 Client credentials are stored but you haven't authorized yet. "
|
||||
"Send `/setup-files start` to begin."
|
||||
)
|
||||
return True
|
||||
|
||||
if arg == "start":
|
||||
if not oauth_helper._client_secret_path().exists():
|
||||
await _reply(
|
||||
"⚠️ No client credentials stored for this profile. Send "
|
||||
"`/setup-files` (no args) for setup instructions."
|
||||
)
|
||||
return True
|
||||
try:
|
||||
output = await _run_captured(oauth_helper.get_auth_url, sender_key)
|
||||
auth_url = output.strip().splitlines()[-1]
|
||||
except SystemExit:
|
||||
await _reply(
|
||||
"❌ Couldn't generate the OAuth URL. Check the gateway logs and verify "
|
||||
"the client_secret.json is valid."
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("[GoogleChat] /setup-files start failed: %s", exc)
|
||||
await _reply(f"❌ Error: {exc}")
|
||||
return True
|
||||
await _reply(_START_INSTRUCTIONS.format(auth_url=auth_url))
|
||||
return True
|
||||
|
||||
if arg == "revoke":
|
||||
try:
|
||||
output = (await _run_captured(oauth_helper.revoke, sender_key)).strip() or "Revoked."
|
||||
except SystemExit:
|
||||
output = "Revoke completed (some steps may have been skipped)."
|
||||
except Exception as exc:
|
||||
logger.warning("[GoogleChat] /setup-files revoke failed: %s", exc)
|
||||
await _reply(f"❌ Error revoking: {exc}")
|
||||
return True
|
||||
# Evict only the sender's slot: Bob revoking must not break Alice's
|
||||
# per-user token nor the shared legacy fallback.
|
||||
if sender_key:
|
||||
adapter._user_creds_by_email.pop(sender_key, None)
|
||||
adapter._user_chat_api_by_email.pop(sender_key, None)
|
||||
else:
|
||||
adapter._user_credentials = None
|
||||
adapter._user_chat_api = None
|
||||
await _reply(f"✅ Done.\n```\n{output}\n```")
|
||||
return True
|
||||
|
||||
# Anything else is the auth code or the pasted failed-redirect URL.
|
||||
try:
|
||||
output = (await _run_captured(oauth_helper.exchange_auth_code, arg, sender_key)).strip()
|
||||
except SystemExit:
|
||||
await _reply(
|
||||
"❌ Token exchange failed. The code may have expired or the URL is malformed. "
|
||||
"Send `/setup-files start` to get a fresh OAuth URL."
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("[GoogleChat] /setup-files exchange failed: %s", exc)
|
||||
await _reply(f"❌ Error: {exc}")
|
||||
return True
|
||||
|
||||
# Re-load credentials so the next file send uses them without a gateway restart.
|
||||
try:
|
||||
new_creds = await asyncio.to_thread(oauth_helper.load_user_credentials, sender_key)
|
||||
if new_creds is not None:
|
||||
new_api = await asyncio.to_thread(lambda: oauth_helper.build_user_chat_service(new_creds))
|
||||
if sender_key:
|
||||
adapter._user_creds_by_email[sender_key] = new_creds
|
||||
adapter._user_chat_api_by_email[sender_key] = new_api
|
||||
else:
|
||||
adapter._user_credentials = new_creds
|
||||
adapter._user_chat_api = new_api
|
||||
await _reply("✅ Authorized! Native attachment delivery is now active. Try asking me to send you a PDF.")
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("[GoogleChat] post-exchange creds load failed: %s", exc)
|
||||
|
||||
await _reply(
|
||||
"⚠️ Token exchanged but the gateway couldn't load the new credentials in-memory. "
|
||||
f"Restart the gateway and the token at `{oauth_helper._token_path(sender_key)}` will be picked up.\n"
|
||||
f"Helper output:\n```\n{output}\n```"
|
||||
)
|
||||
return True
|
||||
+307
-1014
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,195 @@
|
||||
"""Pipeline-facing Teams outbound delivery (meeting-summary writer).
|
||||
|
||||
Lives inside the Teams platform plugin so the meeting pipeline reuses one Teams
|
||||
integration surface. httpx is imported lazily: plugin discovery imports this
|
||||
module on every CLI start, but only ``incoming_webhook`` delivery needs it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import html
|
||||
import os
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
from gateway.config import PlatformConfig
|
||||
from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret
|
||||
|
||||
|
||||
def _parse_bool(value: Any, *, default: bool = False) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return default
|
||||
|
||||
|
||||
_LIST_SECTIONS = (("Key decisions", "key_decisions"), ("Action items", "action_items"), ("Risks", "risks"))
|
||||
|
||||
|
||||
class _StaticAccessTokenProvider:
|
||||
"""Minimal token-provider shim so outbound Graph delivery can reuse the shared client."""
|
||||
|
||||
def __init__(self, access_token: str):
|
||||
self._access_token = str(access_token or "").strip()
|
||||
|
||||
async def get_access_token(self, *, force_refresh: bool = False) -> str:
|
||||
del force_refresh
|
||||
if not self._access_token:
|
||||
raise ValueError("TEAMS_GRAPH_ACCESS_TOKEN is required for graph delivery mode.")
|
||||
return self._access_token
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class TeamsSummaryWriter:
|
||||
"""Deliver a meeting summary to Teams via incoming webhook or Graph."""
|
||||
|
||||
def __init__(
|
||||
self, platform_config: PlatformConfig | None = None, *,
|
||||
graph_client: Any | None = None, transport: httpx.AsyncBaseTransport | None = None,
|
||||
) -> None:
|
||||
self._platform_config = platform_config
|
||||
self._graph_client = graph_client
|
||||
self._transport = transport
|
||||
|
||||
async def write_summary(
|
||||
self, payload: Any, config: dict[str, Any] | None, existing_record: Optional[dict[str, Any]] = None
|
||||
) -> dict[str, Any]:
|
||||
merged = self._resolve_delivery_config(config)
|
||||
if existing_record and not _parse_bool(merged.get("force_resend"), default=False):
|
||||
return dict(existing_record)
|
||||
mode = str(merged.get("delivery_mode") or merged.get("mode") or "").strip().lower()
|
||||
if not mode:
|
||||
if merged.get("incoming_webhook_url"):
|
||||
mode = "incoming_webhook"
|
||||
elif merged.get("chat_id") or (merged.get("team_id") and merged.get("channel_id")):
|
||||
mode = "graph"
|
||||
if mode == "incoming_webhook":
|
||||
return await self._write_summary_via_incoming_webhook(payload, merged)
|
||||
if mode == "graph":
|
||||
return await self._write_summary_via_graph(payload, merged)
|
||||
raise ValueError("Teams delivery_mode must be 'incoming_webhook' or 'graph'.")
|
||||
|
||||
def _resolve_delivery_config(self, config: dict[str, Any] | None) -> dict[str, Any]:
|
||||
merged: dict[str, Any] = {}
|
||||
platform_cfg = self._platform_config
|
||||
if platform_cfg is not None:
|
||||
merged.update(dict(platform_cfg.extra or {}))
|
||||
if platform_cfg.token and "access_token" not in merged:
|
||||
merged["access_token"] = platform_cfg.token
|
||||
if platform_cfg.home_channel:
|
||||
merged.setdefault("channel_id", platform_cfg.home_channel.chat_id)
|
||||
merged.update(dict(config or {}))
|
||||
env_defaults = {
|
||||
"delivery_mode": os.getenv("TEAMS_DELIVERY_MODE", ""),
|
||||
"incoming_webhook_url": os.getenv("TEAMS_INCOMING_WEBHOOK_URL", ""),
|
||||
"access_token": _get_scoped_secret("TEAMS_GRAPH_ACCESS_TOKEN", ""),
|
||||
"team_id": os.getenv("TEAMS_TEAM_ID", ""),
|
||||
"channel_id": os.getenv("TEAMS_CHANNEL_ID", ""),
|
||||
"chat_id": os.getenv("TEAMS_CHAT_ID", ""),
|
||||
}
|
||||
for key, value in env_defaults.items():
|
||||
if value and not merged.get(key):
|
||||
merged[key] = value
|
||||
return merged
|
||||
|
||||
async def _write_summary_via_incoming_webhook(self, payload: Any, config: dict[str, Any]) -> dict[str, Any]:
|
||||
import httpx # lazy — see module docstring
|
||||
|
||||
webhook_url = str(config.get("incoming_webhook_url") or "").strip()
|
||||
if not webhook_url:
|
||||
raise ValueError("TEAMS_INCOMING_WEBHOOK_URL is required for incoming_webhook mode.")
|
||||
body = {"text": self._render_summary_markdown(payload)}
|
||||
async with httpx.AsyncClient(timeout=20.0, transport=self._transport) as client:
|
||||
response = await client.post(webhook_url, json=body)
|
||||
response.raise_for_status()
|
||||
return {
|
||||
"delivery_mode": "incoming_webhook", "webhook_url": webhook_url,
|
||||
"status_code": response.status_code, "delivered": True,
|
||||
}
|
||||
|
||||
async def _write_summary_via_graph(self, payload: Any, config: dict[str, Any]) -> dict[str, Any]:
|
||||
graph_client = self._build_graph_client(config)
|
||||
chat_id = str(config.get("chat_id") or "").strip()
|
||||
if chat_id:
|
||||
path = f"/chats/{quote(chat_id, safe='')}/messages"
|
||||
target = {"target_type": "chat", "chat_id": chat_id}
|
||||
else:
|
||||
team_id = str(config.get("team_id") or "").strip()
|
||||
channel_id = str(config.get("channel_id") or "").strip()
|
||||
if not team_id or not channel_id:
|
||||
raise ValueError("Graph delivery mode requires chat_id, or both team_id and channel_id.")
|
||||
path = f"/teams/{quote(team_id, safe='')}/channels/{quote(channel_id, safe='')}/messages"
|
||||
target = {"target_type": "channel", "team_id": team_id, "channel_id": channel_id}
|
||||
response = await graph_client.post_json(
|
||||
path,
|
||||
json_body={"body": {"contentType": "html", "content": self._render_summary_html(payload)}},
|
||||
)
|
||||
return {
|
||||
"delivery_mode": "graph", **target,
|
||||
"message_id": (response or {}).get("id"), "web_url": (response or {}).get("webUrl"),
|
||||
}
|
||||
|
||||
def _build_graph_client(self, config: dict[str, Any]) -> Any:
|
||||
if self._graph_client is not None:
|
||||
return self._graph_client
|
||||
from tools.microsoft_graph_auth import MicrosoftGraphTokenProvider
|
||||
from tools.microsoft_graph_client import MicrosoftGraphClient
|
||||
|
||||
access_token = str(config.get("access_token") or "").strip()
|
||||
if access_token:
|
||||
return MicrosoftGraphClient(_StaticAccessTokenProvider(access_token), transport=self._transport)
|
||||
return MicrosoftGraphClient(MicrosoftGraphTokenProvider.from_env(), transport=self._transport)
|
||||
|
||||
def _render_summary_markdown(self, payload: Any) -> str:
|
||||
lines = [
|
||||
f"**{self._title(payload)}**",
|
||||
"",
|
||||
f"Summary: {self._text(getattr(payload, 'summary', None), 'No summary available.')}",
|
||||
]
|
||||
for heading, attr in _LIST_SECTIONS:
|
||||
lines += ["", f"{heading}:", *self._bullet_lines(getattr(payload, attr, None))]
|
||||
return "\n".join(lines)
|
||||
|
||||
def _render_summary_html(self, payload: Any) -> str:
|
||||
sections = [
|
||||
("Summary", [self._text(getattr(payload, "summary", None), "No summary available.")]),
|
||||
*((heading, list(getattr(payload, attr, None) or [])) for heading, attr in _LIST_SECTIONS),
|
||||
]
|
||||
blocks = [f"<h2>{html.escape(self._title(payload))}</h2>"]
|
||||
for heading, items in sections:
|
||||
blocks.append(f"<h3>{html.escape(heading)}</h3>")
|
||||
if len(items) == 1 and heading == "Summary":
|
||||
blocks.append(f"<p>{html.escape(str(items[0]))}</p>")
|
||||
continue
|
||||
if items:
|
||||
rendered = "".join(f"<li>{html.escape(str(item))}</li>" for item in items if str(item).strip())
|
||||
blocks.append(rendered and f"<ul>{rendered}</ul>" or "<p>None</p>")
|
||||
else:
|
||||
blocks.append("<p>None</p>")
|
||||
return "".join(blocks)
|
||||
|
||||
@staticmethod
|
||||
def _title(payload: Any) -> str:
|
||||
title = getattr(payload, "title", None)
|
||||
if title:
|
||||
return str(title)
|
||||
meeting_ref = getattr(payload, "meeting_ref", None)
|
||||
meeting_id = getattr(meeting_ref, "meeting_id", None) if meeting_ref else None
|
||||
return f"Meeting {meeting_id or 'summary'}"
|
||||
|
||||
@staticmethod
|
||||
def _text(value: Any, default: str) -> str:
|
||||
text = str(value or "").strip()
|
||||
return text or default
|
||||
|
||||
@classmethod
|
||||
def _bullet_lines(cls, values: Any) -> list[str]:
|
||||
items = [str(item).strip() for item in (values or []) if str(item).strip()]
|
||||
return [f"- {item}" for item in items] or ["- None"]
|
||||
+306
-2644
File diff suppressed because it is too large
Load Diff
@@ -1,12 +1,8 @@
|
||||
"""WeCom callback-mode adapter for self-built enterprise applications.
|
||||
|
||||
Unlike the bot/websocket adapter in ``wecom.py``, this handles the standard
|
||||
WeCom callback flow: WeCom POSTs encrypted XML to an HTTP endpoint, the
|
||||
adapter decrypts it, queues the message for the agent, and immediately
|
||||
acknowledges. The agent's reply is delivered later via the proactive
|
||||
``message/send`` API using an access-token.
|
||||
|
||||
Supports multiple self-built apps under one gateway instance, scoped by
|
||||
WeCom POSTs encrypted XML to an HTTP endpoint; we decrypt, queue for the agent
|
||||
and ack immediately. Replies go out later via the proactive ``message/send``
|
||||
API with an access-token. Multiple apps per gateway are scoped by
|
||||
``corp_id:user_id`` to avoid cross-corp collisions.
|
||||
"""
|
||||
|
||||
@@ -17,10 +13,8 @@ import logging
|
||||
import socket as _socket
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
# Security: parse untrusted, pre-auth request bodies (WeCom callbacks) with
|
||||
# defusedxml to block billion-laughs / entity-expansion (and XXE) DoS. The
|
||||
# parsing API (fromstring) is a drop-in for the stdlib calls used below;
|
||||
# response-building XML lives in wecom_crypto.py and is not parsed here.
|
||||
|
||||
# Untrusted pre-auth bodies are parsed with defusedxml (billion-laughs / XXE).
|
||||
try:
|
||||
import defusedxml.ElementTree as ET
|
||||
|
||||
@@ -51,43 +45,26 @@ from plugins.platforms.wecom.wecom_crypto import WXBizMsgCrypt, WeComCryptoError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ``None`` → aiohttp/asyncio ``create_server`` binds one listening socket per
|
||||
# address family (IPv4 + IPv6). The old "0.0.0.0" default bound IPv4 ONLY and
|
||||
# was unreachable over IPv6-only private networks (e.g. Fly.io 6PN) — same
|
||||
# bug as the LINE adapter (NS-603) and gateway/platforms/webhook.py
|
||||
# (d542894ad). Pin a host via WECOM_CALLBACK_HOST or extra.host.
|
||||
# None → aiohttp binds one socket per address family (IPv4 + IPv6); "0.0.0.0"
|
||||
# was unreachable on IPv6-only networks. Pin via WECOM_CALLBACK_HOST / extra.host.
|
||||
DEFAULT_HOST = None
|
||||
DEFAULT_PORT = 8645
|
||||
DEFAULT_PATH = "/wecom/callback"
|
||||
# Cap pre-auth request bodies. WeCom callbacks are small encrypted XML
|
||||
# envelopes (media is delivered out-of-band via MediaId, never inline), so
|
||||
# 64 KB is ample for any legitimate message while bounding the work an
|
||||
# unauthenticated POST can force before signature verification.
|
||||
# Pre-auth body cap: callbacks are small encrypted XML envelopes (media is
|
||||
# out-of-band via MediaId), so 64 KB bounds unauthenticated work.
|
||||
_MAX_BODY = 65_536
|
||||
ACCESS_TOKEN_TTL_SECONDS = 7200
|
||||
MESSAGE_DEDUP_TTL_SECONDS = 300
|
||||
|
||||
|
||||
def check_wecom_callback_requirements() -> bool:
|
||||
"""PASSIVE probe: are aiohttp/httpx/defusedxml importable right now?
|
||||
|
||||
Registry ``check_fn`` — must never install anything. The ACTIVE
|
||||
lazy-installer is ``ensure_wecom_callback_requirements`` below.
|
||||
"""
|
||||
"""PASSIVE probe (registry ``check_fn``) — must never install anything."""
|
||||
return AIOHTTP_AVAILABLE and HTTPX_AVAILABLE and DEFUSEDXML_AVAILABLE
|
||||
|
||||
|
||||
def ensure_wecom_callback_requirements() -> bool:
|
||||
"""ACTIVE lazy-installer for the ``platform.wecom_callback`` feature.
|
||||
|
||||
Registered as ``ensure_deps_fn``: the registry's ``create_adapter()``
|
||||
runs it when the passive probe fails, right before the gateway connects
|
||||
the platform (#79812). Installs ``defusedxml`` (the only non-core dep;
|
||||
aiohttp/httpx ship with every messaging install) and rebinds the module
|
||||
globals. Before this hook existed, the passive ``check_fn`` returned
|
||||
False forever on installs without the ``wecom`` extra and the
|
||||
``platform.wecom_callback`` LAZY_DEPS entry was never exercised.
|
||||
"""
|
||||
"""ACTIVE lazy-installer (``ensure_deps_fn``): installs ``defusedxml`` — the
|
||||
only non-core dep — when the passive probe fails, and rebinds module globals."""
|
||||
if check_wecom_callback_requirements():
|
||||
return True
|
||||
|
||||
@@ -109,7 +86,6 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
def __init__(self, config: PlatformConfig):
|
||||
super().__init__(config, Platform.WECOM_CALLBACK)
|
||||
extra = config.extra or {}
|
||||
# Falsy host (None/"") collapses to the dual-stack default.
|
||||
_raw_host = extra.get("host") or DEFAULT_HOST
|
||||
self._host = str(_raw_host) if _raw_host else None
|
||||
self._port = int(extra.get("port") or DEFAULT_PORT)
|
||||
@@ -125,10 +101,6 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
self._user_app_map: Dict[str, str] = {}
|
||||
self._access_tokens: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# App normalisation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _user_app_key(corp_id: str, user_id: str) -> str:
|
||||
return f"{corp_id}:{user_id}" if corp_id else user_id
|
||||
@@ -139,29 +111,18 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
if isinstance(apps, list) and apps:
|
||||
return [dict(app) for app in apps if isinstance(app, dict)]
|
||||
if extra.get("corp_id"):
|
||||
return [
|
||||
{
|
||||
"name": extra.get("name") or "default",
|
||||
"corp_id": extra.get("corp_id", ""),
|
||||
"corp_secret": extra.get("corp_secret", ""),
|
||||
"agent_id": str(extra.get("agent_id", "")),
|
||||
"token": extra.get("token", ""),
|
||||
"encoding_aes_key": extra.get("encoding_aes_key", ""),
|
||||
}
|
||||
]
|
||||
return [{
|
||||
"name": extra.get("name") or "default",
|
||||
"corp_id": extra.get("corp_id", ""),
|
||||
"corp_secret": extra.get("corp_secret", ""),
|
||||
"agent_id": str(extra.get("agent_id", "")),
|
||||
"token": extra.get("token", ""),
|
||||
"encoding_aes_key": extra.get("encoding_aes_key", ""),
|
||||
}]
|
||||
return []
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Lifecycle
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def connect(self, *, is_reconnect: bool = False) -> bool:
|
||||
# ``is_reconnect`` is forwarded by GatewayRunner on every retry per
|
||||
# the BasePlatformAdapter.connect contract. Callback adapters have
|
||||
# no server-side queue to preserve, so the flag is accepted-and-
|
||||
# ignored — but the kwarg MUST be present or the reconnect watcher
|
||||
# dies with TypeError and the platform silently stays offline.
|
||||
del is_reconnect
|
||||
del is_reconnect # kwarg MUST exist (GatewayRunner passes it) even though unused
|
||||
if not self._apps:
|
||||
logger.warning("[WecomCallback] No callback apps configured")
|
||||
return False
|
||||
@@ -169,8 +130,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
logger.warning("[WecomCallback] aiohttp/httpx not installed")
|
||||
return False
|
||||
|
||||
# Quick port-in-use check.
|
||||
try:
|
||||
try: # quick port-in-use check
|
||||
with _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM) as sock:
|
||||
sock.settimeout(1)
|
||||
sock.connect(("127.0.0.1", self._port))
|
||||
@@ -180,11 +140,9 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
pass
|
||||
|
||||
try:
|
||||
# Tighter keepalive so idle CLOSE_WAIT drains promptly (#18451).
|
||||
from gateway.platforms._http_client_limits import platform_httpx_limits
|
||||
self._http_client = httpx.AsyncClient(timeout=20.0, limits=platform_httpx_limits())
|
||||
# client_max_size rejects oversized bodies at the aiohttp layer
|
||||
# (413) before our handler — and before any signature work — runs.
|
||||
# client_max_size → 413 before our handler / any signature work runs.
|
||||
self._app = web.Application(client_max_size=_MAX_BODY)
|
||||
self._app.router.add_get("/health", self._handle_health)
|
||||
self._app.router.add_get(self._path, self._handle_verify)
|
||||
@@ -195,10 +153,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
await self._site.start()
|
||||
self._poll_task = asyncio.create_task(self._poll_loop())
|
||||
self._mark_connected()
|
||||
logger.info(
|
||||
"[WecomCallback] HTTP server listening on %s:%s%s",
|
||||
self._host, self._port, self._path,
|
||||
)
|
||||
logger.info("[WecomCallback] HTTP server listening on %s:%s%s", self._host, self._port, self._path)
|
||||
for app in self._apps:
|
||||
try:
|
||||
await self._refresh_access_token(app)
|
||||
@@ -236,10 +191,6 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Outbound: proactive send via access-token API
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def send(
|
||||
self,
|
||||
chat_id: str,
|
||||
@@ -266,8 +217,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
data = resp.json()
|
||||
errcode = data.get("errcode")
|
||||
if errcode in {40001, 42001} and _attempt == 0:
|
||||
# WeCom rejected the token — evict the cached entry so
|
||||
# the next _get_access_token call forces a fresh fetch.
|
||||
# Token rejected — evict so the next call fetches a fresh one.
|
||||
logger.warning(
|
||||
"[WecomCallback] Token rejected for app '%s' (errcode=%s), refreshing",
|
||||
app.get("name", "default"), errcode,
|
||||
@@ -276,11 +226,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
continue
|
||||
if errcode != 0:
|
||||
return SendResult(success=False, error=str(data))
|
||||
return SendResult(
|
||||
success=True,
|
||||
message_id=str(data.get("msgid", "")),
|
||||
raw_response=data,
|
||||
)
|
||||
return SendResult(success=True, message_id=str(data.get("msgid", "")), raw_response=data)
|
||||
return SendResult(success=False, error="send failed after token refresh")
|
||||
except Exception as exc:
|
||||
return SendResult(success=False, error=str(exc))
|
||||
@@ -288,8 +234,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
def _resolve_app_for_chat(self, chat_id: str) -> Dict[str, Any]:
|
||||
"""Pick the app associated with *chat_id*, falling back sensibly."""
|
||||
app_name = self._user_app_map.get(chat_id)
|
||||
if not app_name and ":" not in chat_id:
|
||||
# Legacy bare user_id — try to find a unique match.
|
||||
if not app_name and ":" not in chat_id: # legacy bare user_id — unique match only
|
||||
matching = [k for k in self._user_app_map if k.endswith(f":{chat_id}")]
|
||||
if len(matching) == 1:
|
||||
app_name = self._user_app_map.get(matching[0])
|
||||
@@ -299,23 +244,16 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
|
||||
return {"name": chat_id, "type": "dm"}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Inbound: HTTP callback handlers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _handle_health(self, request: web.Request) -> web.Response:
|
||||
return web.json_response({"status": "ok", "platform": "wecom_callback"})
|
||||
|
||||
async def _handle_verify(self, request: web.Request) -> web.Response:
|
||||
"""GET endpoint — WeCom URL verification handshake."""
|
||||
msg_signature = request.query.get("msg_signature", "")
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
msg_signature, timestamp, nonce = self._signature_params(request)
|
||||
echostr = request.query.get("echostr", "")
|
||||
for app in self._apps:
|
||||
try:
|
||||
crypt = self._crypt_for_app(app)
|
||||
plain = crypt.verify_url(msg_signature, timestamp, nonce, echostr)
|
||||
plain = self._crypt_for_app(app).verify_url(msg_signature, timestamp, nonce, echostr)
|
||||
return web.Response(text=plain, content_type="text/plain")
|
||||
except Exception:
|
||||
continue
|
||||
@@ -323,11 +261,8 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
|
||||
async def _handle_callback(self, request: web.Request) -> web.Response:
|
||||
"""POST endpoint — receive an encrypted message callback."""
|
||||
msg_signature = request.query.get("msg_signature", "")
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
# Explicit guard in addition to client_max_size: rejects oversized
|
||||
# payloads before any XML parse / signature check (DoS, zip bombs).
|
||||
msg_signature, timestamp, nonce = self._signature_params(request)
|
||||
# Explicit guard in addition to client_max_size (DoS / zip bombs).
|
||||
body_bytes = await request.read()
|
||||
if len(body_bytes) > _MAX_BODY:
|
||||
logger.warning("[WecomCallback] Payload too large (%d bytes) — rejected", len(body_bytes))
|
||||
@@ -336,34 +271,18 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
|
||||
for app in self._apps:
|
||||
try:
|
||||
decrypted = self._decrypt_request(
|
||||
app, body, msg_signature, timestamp, nonce,
|
||||
)
|
||||
decrypted = self._decrypt_request(app, body, msg_signature, timestamp, nonce)
|
||||
event = self._build_event(app, decrypted)
|
||||
if event is not None:
|
||||
# Deduplicate: WeCom retries callbacks on timeout,
|
||||
# producing duplicate inbound messages (#10305).
|
||||
if event.message_id:
|
||||
now = time.time()
|
||||
if event.message_id in self._seen_messages:
|
||||
if now - self._seen_messages[event.message_id] < MESSAGE_DEDUP_TTL_SECONDS:
|
||||
logger.debug("[WecomCallback] Duplicate MsgId %s, skipping", event.message_id)
|
||||
return web.Response(text="success", content_type="text/plain")
|
||||
del self._seen_messages[event.message_id]
|
||||
self._seen_messages[event.message_id] = now
|
||||
# Prune expired entries when cache grows large
|
||||
if len(self._seen_messages) > 2000:
|
||||
cutoff = now - MESSAGE_DEDUP_TTL_SECONDS
|
||||
self._seen_messages = {k: v for k, v in self._seen_messages.items() if v > cutoff}
|
||||
# Record which app this user belongs to.
|
||||
# WeCom retries callbacks on timeout → duplicate inbound messages.
|
||||
if event.message_id and self._is_duplicate(event.message_id):
|
||||
logger.debug("[WecomCallback] Duplicate MsgId %s, skipping", event.message_id)
|
||||
return web.Response(text="success", content_type="text/plain")
|
||||
if event.source and event.source.user_id:
|
||||
map_key = self._user_app_key(
|
||||
str(app.get("corp_id") or ""), event.source.user_id,
|
||||
)
|
||||
map_key = self._user_app_key(str(app.get("corp_id") or ""), event.source.user_id)
|
||||
self._user_app_map[map_key] = app["name"]
|
||||
await self._message_queue.put(event)
|
||||
# Immediately acknowledge — the agent's reply will arrive
|
||||
# later via the proactive message/send API.
|
||||
# Ack immediately — the reply arrives later via proactive message/send.
|
||||
return web.Response(text="success", content_type="text/plain")
|
||||
except WeComCryptoError:
|
||||
continue
|
||||
@@ -372,6 +291,23 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
break
|
||||
return web.Response(status=400, text="invalid callback payload")
|
||||
|
||||
@staticmethod
|
||||
def _signature_params(request: web.Request):
|
||||
q = request.query
|
||||
return q.get("msg_signature", ""), q.get("timestamp", ""), q.get("nonce", "")
|
||||
|
||||
def _is_duplicate(self, message_id: str) -> bool:
|
||||
now = time.time()
|
||||
if message_id in self._seen_messages:
|
||||
if now - self._seen_messages[message_id] < MESSAGE_DEDUP_TTL_SECONDS:
|
||||
return True
|
||||
del self._seen_messages[message_id]
|
||||
self._seen_messages[message_id] = now
|
||||
if len(self._seen_messages) > 2000: # prune expired entries
|
||||
cutoff = now - MESSAGE_DEDUP_TTL_SECONDS
|
||||
self._seen_messages = {k: v for k, v in self._seen_messages.items() if v > cutoff}
|
||||
return False
|
||||
|
||||
async def _poll_loop(self) -> None:
|
||||
"""Drain the message queue and dispatch to the gateway runner."""
|
||||
while True:
|
||||
@@ -383,27 +319,16 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
except Exception:
|
||||
logger.exception("[WecomCallback] Failed to enqueue event")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# XML / crypto helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _decrypt_request(
|
||||
self, app: Dict[str, Any], body: str,
|
||||
msg_signature: str, timestamp: str, nonce: str,
|
||||
) -> str:
|
||||
root = ET.fromstring(body)
|
||||
encrypt = root.findtext("Encrypt", default="")
|
||||
crypt = self._crypt_for_app(app)
|
||||
return crypt.decrypt(msg_signature, timestamp, nonce, encrypt).decode("utf-8")
|
||||
def _decrypt_request(self, app: Dict[str, Any], body: str, msg_signature: str, timestamp: str, nonce: str) -> str:
|
||||
encrypt = ET.fromstring(body).findtext("Encrypt", default="")
|
||||
return self._crypt_for_app(app).decrypt(msg_signature, timestamp, nonce, encrypt).decode("utf-8")
|
||||
|
||||
def _build_event(self, app: Dict[str, Any], xml_text: str) -> Optional[MessageEvent]:
|
||||
root = ET.fromstring(xml_text)
|
||||
msg_type = (root.findtext("MsgType") or "").lower()
|
||||
# Silently acknowledge lifecycle events.
|
||||
if msg_type == "event":
|
||||
event_name = (root.findtext("Event") or "").lower()
|
||||
if event_name in {"enter_agent", "subscribe"}:
|
||||
return None
|
||||
# Lifecycle events are silently acknowledged.
|
||||
if msg_type == "event" and (root.findtext("Event") or "").lower() in {"enter_agent", "subscribe"}:
|
||||
return None
|
||||
if msg_type not in {"text", "event"}:
|
||||
return None
|
||||
|
||||
@@ -413,24 +338,9 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
content = root.findtext("Content", default="").strip()
|
||||
if not content and msg_type == "event":
|
||||
content = "/start"
|
||||
msg_id = (
|
||||
root.findtext("MsgId")
|
||||
or f"{user_id}:{root.findtext('CreateTime', default='0')}"
|
||||
)
|
||||
source = self.build_source(
|
||||
chat_id=scoped_chat_id,
|
||||
chat_name=user_id,
|
||||
chat_type="dm",
|
||||
user_id=user_id,
|
||||
user_name=user_id,
|
||||
)
|
||||
return MessageEvent(
|
||||
text=content,
|
||||
message_type=MessageType.TEXT,
|
||||
source=source,
|
||||
raw_message=xml_text,
|
||||
message_id=msg_id,
|
||||
)
|
||||
msg_id = root.findtext("MsgId") or f"{user_id}:{root.findtext('CreateTime', default='0')}"
|
||||
source = self.build_source(chat_id=scoped_chat_id, chat_name=user_id, chat_type="dm", user_id=user_id, user_name=user_id)
|
||||
return MessageEvent(text=content, message_type=MessageType.TEXT, source=source, raw_message=xml_text, message_id=msg_id)
|
||||
|
||||
def _crypt_for_app(self, app: Dict[str, Any]) -> WXBizMsgCrypt:
|
||||
return WXBizMsgCrypt(
|
||||
@@ -440,16 +350,7 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
)
|
||||
|
||||
def _get_app_by_name(self, name: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||
if not name:
|
||||
return None
|
||||
for app in self._apps:
|
||||
if app.get("name") == name:
|
||||
return app
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Access-token management
|
||||
# ------------------------------------------------------------------
|
||||
return next((app for app in self._apps if app.get("name") == name), None) if name else None
|
||||
|
||||
async def _get_access_token(self, app: Dict[str, Any]) -> str:
|
||||
cached = self._access_tokens.get(app["name"])
|
||||
@@ -461,24 +362,16 @@ class WecomCallbackAdapter(BasePlatformAdapter):
|
||||
async def _refresh_access_token(self, app: Dict[str, Any]) -> str:
|
||||
resp = await self._http_client.get(
|
||||
"https://qyapi.weixin.qq.com/cgi-bin/gettoken",
|
||||
params={
|
||||
"corpid": app.get("corp_id"),
|
||||
"corpsecret": app.get("corp_secret"),
|
||||
},
|
||||
params={"corpid": app.get("corp_id"), "corpsecret": app.get("corp_secret")},
|
||||
)
|
||||
data = resp.json()
|
||||
if data.get("errcode") != 0:
|
||||
raise RuntimeError(f"WeCom token refresh failed: {data}")
|
||||
token = data["access_token"]
|
||||
expires_in = int(data.get("expires_in", ACCESS_TOKEN_TTL_SECONDS))
|
||||
self._access_tokens[app["name"]] = {
|
||||
"token": token,
|
||||
"expires_at": time.time() + expires_in,
|
||||
}
|
||||
self._access_tokens[app["name"]] = {"token": token, "expires_at": time.time() + expires_in}
|
||||
logger.info(
|
||||
"[WecomCallback] Token refreshed for app '%s' (corp=%s), expires in %ss",
|
||||
app.get("name", "default"),
|
||||
app.get("corp_id", ""),
|
||||
expires_in,
|
||||
app.get("name", "default"), app.get("corp_id", ""), expires_in,
|
||||
)
|
||||
return token
|
||||
|
||||
@@ -0,0 +1,471 @@
|
||||
"""WeCom media: inbound attachment caching and outbound upload/send.
|
||||
|
||||
Mixed into :class:`WeComAdapter`. Outbound media goes through the chunked
|
||||
``aibot_upload_media_*`` flow and is then sent natively (image/video/voice/file).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
import mimetypes
|
||||
import re
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
from gateway.platforms.base import SendResult, cache_document_from_bytes, cache_image_from_bytes
|
||||
|
||||
logger = logging.getLogger("plugins.platforms.wecom.adapter")
|
||||
|
||||
APP_CMD_SEND = "aibot_send_msg"
|
||||
APP_CMD_UPLOAD_MEDIA_INIT = "aibot_upload_media_init"
|
||||
APP_CMD_UPLOAD_MEDIA_CHUNK = "aibot_upload_media_chunk"
|
||||
APP_CMD_UPLOAD_MEDIA_FINISH = "aibot_upload_media_finish"
|
||||
|
||||
IMAGE_MAX_BYTES = 10 * 1024 * 1024
|
||||
VIDEO_MAX_BYTES = 10 * 1024 * 1024
|
||||
VOICE_MAX_BYTES = 2 * 1024 * 1024
|
||||
FILE_MAX_BYTES = 20 * 1024 * 1024
|
||||
ABSOLUTE_MAX_BYTES = FILE_MAX_BYTES
|
||||
UPLOAD_CHUNK_SIZE = 512 * 1024
|
||||
MAX_UPLOAD_CHUNKS = 100
|
||||
VOICE_SUPPORTED_MIMES = {"audio/amr"}
|
||||
|
||||
|
||||
def _size_verdict(final_type: str, *, reject: Optional[str] = None, downgrade: Optional[str] = None) -> Dict[str, Any]:
|
||||
return {
|
||||
"final_type": final_type,
|
||||
"rejected": reject is not None,
|
||||
"reject_reason": reject,
|
||||
"downgraded": downgrade is not None,
|
||||
"downgrade_note": downgrade,
|
||||
}
|
||||
|
||||
|
||||
class WeComMediaMixin:
|
||||
"""Media helpers for WeComAdapter (expects ``_http_client``, ``_send_request``,
|
||||
``_send_reply_request``, ``send``, ``_reply_req_id_for_message``,
|
||||
``_last_chat_req_ids``, ``_stream_expired_chats``, ``_find_active_turn_for_chat``)."""
|
||||
|
||||
# ── Inbound ──────────────────────────────────────────────────────────
|
||||
|
||||
async def _extract_media(self, body: Dict[str, Any]) -> Tuple[List[str], List[str]]:
|
||||
"""Best-effort extraction of inbound media to local cache paths."""
|
||||
refs: List[Tuple[str, Dict[str, Any]]] = []
|
||||
msgtype = str(body.get("msgtype") or "").lower()
|
||||
|
||||
def _ref(kind: str, container: Dict[str, Any]) -> bool:
|
||||
if isinstance(container.get(kind), dict):
|
||||
refs.append((kind, container[kind]))
|
||||
return True
|
||||
return False
|
||||
|
||||
if msgtype == "mixed":
|
||||
mixed = body.get("mixed") if isinstance(body.get("mixed"), dict) else {}
|
||||
items = mixed.get("msg_item") if isinstance(mixed.get("msg_item"), list) else []
|
||||
for item in items:
|
||||
if isinstance(item, dict) and str(item.get("msgtype") or "").lower() == "image":
|
||||
_ref("image", item)
|
||||
else:
|
||||
_ref("image", body)
|
||||
if msgtype == "file":
|
||||
_ref("file", body)
|
||||
# appmsg = WeCom AI Bot attachments (PDF/Word/Excel)
|
||||
if msgtype == "appmsg" and isinstance(body.get("appmsg"), dict):
|
||||
_ref("file", body["appmsg"]) or _ref("image", body["appmsg"])
|
||||
|
||||
quote = body.get("quote") if isinstance(body.get("quote"), dict) else {}
|
||||
quote_type = str(quote.get("msgtype") or "").lower()
|
||||
if quote_type in ("image", "file"):
|
||||
_ref(quote_type, quote)
|
||||
|
||||
media_paths: List[str] = []
|
||||
media_types: List[str] = []
|
||||
for kind, ref in refs:
|
||||
cached = await self._cache_media(kind, ref)
|
||||
if cached:
|
||||
media_paths.append(cached[0])
|
||||
media_types.append(cached[1])
|
||||
return media_paths, media_types
|
||||
|
||||
async def _cache_media(self, kind: str, media: Dict[str, Any]) -> Optional[Tuple[str, str]]:
|
||||
"""Cache an inbound image/file reference (inline base64 or URL) to local storage."""
|
||||
if media.get("base64"):
|
||||
try:
|
||||
raw = self._decode_base64(media["base64"])
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] Failed to decode %s base64 media: %s", self.name, kind, exc)
|
||||
return None
|
||||
if kind == "image":
|
||||
ext = self._detect_image_ext(raw)
|
||||
return self._cache_image(raw, ext, self._mime_for_ext(ext, fallback="image/jpeg"), "")
|
||||
filename = str(media.get("filename") or media.get("name") or "wecom_file")
|
||||
return cache_document_from_bytes(raw, filename), mimetypes.guess_type(filename)[0] or "application/octet-stream"
|
||||
|
||||
url = str(media.get("url") or "").strip()
|
||||
if not url:
|
||||
return None
|
||||
try:
|
||||
raw, headers = await self._download_remote_bytes(url, max_bytes=ABSOLUTE_MAX_BYTES)
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] Failed to download %s from %s: %s", self.name, kind, url, exc)
|
||||
return None
|
||||
aes_key = str(media.get("aeskey") or "").strip()
|
||||
if aes_key:
|
||||
try:
|
||||
raw = self._decrypt_file_bytes(raw, aes_key)
|
||||
except Exception as exc:
|
||||
logger.debug("[%s] Failed to decrypt %s from %s: %s", self.name, kind, url, exc)
|
||||
return None
|
||||
content_type = str(headers.get("content-type") or "").split(";", 1)[0].strip() or "application/octet-stream"
|
||||
if kind == "image":
|
||||
ext = self._guess_extension(url, content_type, fallback=self._detect_image_ext(raw))
|
||||
return self._cache_image(raw, ext, content_type or self._mime_for_ext(ext, fallback="image/jpeg"), f" from {url}")
|
||||
filename = self._guess_filename(url, headers.get("content-disposition"), content_type)
|
||||
return cache_document_from_bytes(raw, filename), content_type
|
||||
|
||||
def _cache_image(self, raw: bytes, ext: str, mime: str, origin: str) -> Optional[Tuple[str, str]]:
|
||||
try:
|
||||
return cache_image_from_bytes(raw, ext), mime
|
||||
except ValueError as exc:
|
||||
logger.warning("[%s] Rejected non-image bytes%s: %s", self.name, origin, exc)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _decode_base64(data: str) -> bytes:
|
||||
return base64.b64decode(data.split(",", 1)[-1].strip())
|
||||
|
||||
@staticmethod
|
||||
def _detect_image_ext(data: bytes) -> str:
|
||||
if data.startswith(b"\x89PNG\r\n\x1a\n"):
|
||||
return ".png"
|
||||
if data.startswith(b"\xff\xd8\xff"):
|
||||
return ".jpg"
|
||||
if data.startswith((b"GIF87a", b"GIF89a")):
|
||||
return ".gif"
|
||||
if data.startswith(b"RIFF") and data[8:12] == b"WEBP":
|
||||
return ".webp"
|
||||
return ".jpg"
|
||||
|
||||
@staticmethod
|
||||
def _mime_for_ext(ext: str, fallback: str = "application/octet-stream") -> str:
|
||||
return mimetypes.types_map.get(ext.lower(), fallback)
|
||||
|
||||
@staticmethod
|
||||
def _guess_extension(url: str, content_type: str, fallback: str) -> str:
|
||||
ext = mimetypes.guess_extension(content_type) if content_type else None
|
||||
return ext or Path(urlparse(url).path).suffix or fallback
|
||||
|
||||
@staticmethod
|
||||
def _guess_filename(url: str, content_disposition: Optional[str], content_type: str) -> str:
|
||||
if content_disposition:
|
||||
match = re.search(r'filename="?([^";]+)"?', content_disposition)
|
||||
if match:
|
||||
return match.group(1)
|
||||
name = Path(urlparse(url).path).name or "document"
|
||||
if "." not in name:
|
||||
name = f"{name}{mimetypes.guess_extension(content_type) or '.bin'}"
|
||||
return name
|
||||
|
||||
@staticmethod
|
||||
def _decrypt_file_bytes(encrypted_data: bytes, aes_key: str) -> bytes:
|
||||
if not encrypted_data:
|
||||
raise ValueError("encrypted_data is empty")
|
||||
if not aes_key:
|
||||
raise ValueError("aes_key is required")
|
||||
# WeCom doesn't pad base64 keys; add padding if needed
|
||||
aes_key = aes_key + '=' * ((4 - len(aes_key) % 4) % 4)
|
||||
key = base64.b64decode(aes_key)
|
||||
if len(key) != 32:
|
||||
raise ValueError(f"Invalid WeCom AES key length: expected 32 bytes, got {len(key)}")
|
||||
try:
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
except ImportError as exc: # pragma: no cover - dependency is environment-specific
|
||||
raise RuntimeError("cryptography is required for WeCom media decryption") from exc
|
||||
decryptor = Cipher(algorithms.AES(key), modes.CBC(key[:16])).decryptor()
|
||||
decrypted = decryptor.update(encrypted_data) + decryptor.finalize()
|
||||
pad_len = decrypted[-1]
|
||||
if pad_len < 1 or pad_len > 32 or pad_len > len(decrypted):
|
||||
raise ValueError(f"Invalid PKCS#7 padding value: {pad_len}")
|
||||
if any(byte != pad_len for byte in decrypted[-pad_len:]):
|
||||
raise ValueError("Invalid PKCS#7 padding: padding bytes mismatch")
|
||||
return decrypted[:-pad_len]
|
||||
|
||||
async def _download_remote_bytes(self, url: str, max_bytes: int) -> Tuple[bytes, Dict[str, str]]:
|
||||
from gateway.platforms.base import _ssrf_redirect_guard
|
||||
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
|
||||
from plugins.platforms.wecom import adapter as _adapter_mod
|
||||
|
||||
if not is_safe_url(url):
|
||||
raise ValueError(f"Blocked unsafe URL (SSRF protection): {url[:80]}")
|
||||
if not _adapter_mod.HTTPX_AVAILABLE:
|
||||
raise RuntimeError("httpx is required for WeCom media download")
|
||||
|
||||
client = self._http_client or create_ssrf_safe_async_client(
|
||||
timeout=30.0, follow_redirects=True, event_hooks={"response": [_ssrf_redirect_guard]},
|
||||
)
|
||||
created_client = client is not self._http_client
|
||||
try:
|
||||
async with client.stream(
|
||||
"GET", url, headers={"User-Agent": "HermesAgent/1.0", "Accept": "*/*"},
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
headers = {key.lower(): value for key, value in response.headers.items()}
|
||||
content_length = headers.get("content-length")
|
||||
if content_length and content_length.isdigit() and int(content_length) > max_bytes:
|
||||
raise ValueError(
|
||||
f"Remote media exceeds WeCom limit: {int(content_length)} bytes > {max_bytes} bytes"
|
||||
)
|
||||
data = bytearray()
|
||||
async for chunk in response.aiter_bytes():
|
||||
data.extend(chunk)
|
||||
if len(data) > max_bytes:
|
||||
raise ValueError(
|
||||
f"Remote media exceeds WeCom limit while downloading: {len(data)} bytes > {max_bytes} bytes"
|
||||
)
|
||||
return bytes(data), headers
|
||||
finally:
|
||||
if created_client:
|
||||
await client.aclose()
|
||||
|
||||
# ── Outbound classification ──────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _guess_mime_type(filename: str) -> str:
|
||||
mime_type = mimetypes.guess_type(filename)[0]
|
||||
if mime_type:
|
||||
return mime_type
|
||||
if Path(filename).suffix.lower() == ".amr":
|
||||
return "audio/amr"
|
||||
return "application/octet-stream"
|
||||
|
||||
@staticmethod
|
||||
def _normalize_content_type(content_type: str, filename: str) -> str:
|
||||
normalized = str(content_type or "").split(";", 1)[0].strip().lower()
|
||||
if not normalized or normalized in {"application/octet-stream", "text/plain"}:
|
||||
return WeComMediaMixin._guess_mime_type(filename)
|
||||
return normalized
|
||||
|
||||
@staticmethod
|
||||
def _detect_wecom_media_type(content_type: str) -> str:
|
||||
mime_type = str(content_type or "").strip().lower()
|
||||
if mime_type.startswith("image/"):
|
||||
return "image"
|
||||
if mime_type.startswith("video/"):
|
||||
return "video"
|
||||
if mime_type.startswith("audio/") or mime_type == "application/ogg":
|
||||
return "voice"
|
||||
return "file"
|
||||
|
||||
@staticmethod
|
||||
def _apply_file_size_limits(file_size: int, detected_type: str, content_type: Optional[str] = None) -> Dict[str, Any]:
|
||||
file_size_mb = file_size / (1024 * 1024)
|
||||
normalized_type = str(detected_type or "file").lower()
|
||||
normalized_content_type = str(content_type or "").strip().lower()
|
||||
if file_size > ABSOLUTE_MAX_BYTES:
|
||||
return _size_verdict(normalized_type, reject=(
|
||||
f"文件大小 {file_size_mb:.2f}MB 超过了企业微信允许的最大限制 20MB,无法发送。"
|
||||
"请尝试压缩文件或减小文件大小。"
|
||||
))
|
||||
if normalized_type == "image" and file_size > IMAGE_MAX_BYTES:
|
||||
return _size_verdict("file", downgrade=f"图片大小 {file_size_mb:.2f}MB 超过 10MB 限制,已转为文件格式发送")
|
||||
if normalized_type == "video" and file_size > VIDEO_MAX_BYTES:
|
||||
return _size_verdict("file", downgrade=f"视频大小 {file_size_mb:.2f}MB 超过 10MB 限制,已转为文件格式发送")
|
||||
if normalized_type == "voice":
|
||||
if normalized_content_type and normalized_content_type not in VOICE_SUPPORTED_MIMES:
|
||||
return _size_verdict("file", downgrade=(
|
||||
f"语音格式 {normalized_content_type} 不支持,企微仅支持 AMR 格式,已转为文件格式发送"
|
||||
))
|
||||
if file_size > VOICE_MAX_BYTES:
|
||||
return _size_verdict("file", downgrade=f"语音大小 {file_size_mb:.2f}MB 超过 2MB 限制,已转为文件格式发送")
|
||||
return _size_verdict(normalized_type)
|
||||
|
||||
@staticmethod
|
||||
def _looks_like_url(media_source: str) -> bool:
|
||||
return urlparse(str(media_source or "")).scheme in {"http", "https"}
|
||||
|
||||
async def _load_outbound_media(self, media_source: str, file_name: Optional[str] = None) -> Tuple[bytes, str, str]:
|
||||
source = str(media_source or "").strip()
|
||||
if not source:
|
||||
raise ValueError("media source is required")
|
||||
if re.fullmatch(r"<[^>\n]+>", source):
|
||||
raise ValueError(f"Media placeholder was not replaced with a real file path: {source}")
|
||||
|
||||
parsed = urlparse(source)
|
||||
if parsed.scheme in {"http", "https"}:
|
||||
data, headers = await self._download_remote_bytes(source, max_bytes=ABSOLUTE_MAX_BYTES)
|
||||
content_disposition = headers.get("content-disposition")
|
||||
resolved_name = file_name or self._guess_filename(source, content_disposition, headers.get("content-type", ""))
|
||||
content_type = self._normalize_content_type(headers.get("content-type", ""), resolved_name)
|
||||
return data, content_type, resolved_name
|
||||
|
||||
local_path = Path(unquote(parsed.path) if parsed.scheme == "file" else source).expanduser()
|
||||
if not local_path.is_absolute():
|
||||
local_path = (Path.cwd() / local_path).resolve()
|
||||
if not local_path.exists() or not local_path.is_file():
|
||||
raise FileNotFoundError(f"Media file not found: {local_path}")
|
||||
data = local_path.read_bytes()
|
||||
resolved_name = file_name or local_path.name
|
||||
return data, self._normalize_content_type("", resolved_name), resolved_name
|
||||
|
||||
async def _prepare_outbound_media(self, media_source: str, file_name: Optional[str] = None) -> Dict[str, Any]:
|
||||
data, content_type, resolved_name = await self._load_outbound_media(media_source, file_name=file_name)
|
||||
detected_type = self._detect_wecom_media_type(content_type)
|
||||
size_check = self._apply_file_size_limits(len(data), detected_type, content_type)
|
||||
return {"data": data, "content_type": content_type, "file_name": resolved_name, "detected_type": detected_type, **size_check}
|
||||
|
||||
# ── Outbound upload + send ───────────────────────────────────────────
|
||||
|
||||
async def _upload_media_bytes(self, data: bytes, media_type: str, filename: str) -> Dict[str, Any]:
|
||||
if not data:
|
||||
raise ValueError("Cannot upload empty media")
|
||||
total_size = len(data)
|
||||
total_chunks = (total_size + UPLOAD_CHUNK_SIZE - 1) // UPLOAD_CHUNK_SIZE
|
||||
if total_chunks > MAX_UPLOAD_CHUNKS:
|
||||
raise ValueError(f"File too large: {total_chunks} chunks exceeds maximum of {MAX_UPLOAD_CHUNKS} chunks")
|
||||
|
||||
init_response = await self._send_request(APP_CMD_UPLOAD_MEDIA_INIT, {
|
||||
"type": media_type, "filename": filename, "total_size": total_size, "total_chunks": total_chunks,
|
||||
"md5": hashlib.md5(data).hexdigest(),
|
||||
})
|
||||
self._raise_for_wecom_error(init_response, "media upload init")
|
||||
init_body = init_response.get("body") if isinstance(init_response.get("body"), dict) else {}
|
||||
upload_id = str(init_body.get("upload_id") or "").strip()
|
||||
if not upload_id:
|
||||
raise RuntimeError(f"media upload init failed: missing upload_id in response {init_response}")
|
||||
|
||||
for chunk_index, start in enumerate(range(0, total_size, UPLOAD_CHUNK_SIZE)):
|
||||
chunk_response = await self._send_request(APP_CMD_UPLOAD_MEDIA_CHUNK, {
|
||||
"upload_id": upload_id,
|
||||
"chunk_index": chunk_index, # official SDK uses 0-based chunk indexes
|
||||
"base64_data": base64.b64encode(data[start : start + UPLOAD_CHUNK_SIZE]).decode("ascii"),
|
||||
})
|
||||
self._raise_for_wecom_error(chunk_response, f"media upload chunk {chunk_index}")
|
||||
|
||||
finish_response = await self._send_request(APP_CMD_UPLOAD_MEDIA_FINISH, {"upload_id": upload_id})
|
||||
self._raise_for_wecom_error(finish_response, "media upload finish")
|
||||
finish_body = finish_response.get("body") if isinstance(finish_response.get("body"), dict) else {}
|
||||
media_id = str(finish_body.get("media_id") or "").strip()
|
||||
if not media_id:
|
||||
raise RuntimeError(f"media upload finish failed: missing media_id in response {finish_response}")
|
||||
return {"type": str(finish_body.get("type") or media_type), "media_id": media_id, "created_at": finish_body.get("created_at")}
|
||||
|
||||
async def _send_media_message(self, chat_id: str, media_type: str, media_id: str) -> Dict[str, Any]:
|
||||
response = await self._send_request(
|
||||
APP_CMD_SEND, {"chatid": chat_id, "msgtype": media_type, media_type: {"media_id": media_id}},
|
||||
)
|
||||
self._raise_for_wecom_error(response, "send media message")
|
||||
return response
|
||||
|
||||
async def _send_reply_media_message(self, reply_req_id: str, media_type: str, media_id: str) -> Dict[str, Any]:
|
||||
response = await self._send_reply_request(reply_req_id, {"msgtype": media_type, media_type: {"media_id": media_id}})
|
||||
self._raise_for_wecom_error(response, "send reply media message")
|
||||
return response
|
||||
|
||||
async def _send_followup_markdown(self, chat_id: str, content: str, reply_to: Optional[str] = None) -> Optional[SendResult]:
|
||||
if not content:
|
||||
return None
|
||||
result = await self.send(chat_id=chat_id, content=content, reply_to=reply_to)
|
||||
if not result.success:
|
||||
logger.warning("[%s] Follow-up markdown send failed: %s", self.name, result.error)
|
||||
return result
|
||||
|
||||
async def _send_media_source(
|
||||
self, chat_id: str, media_source: str, caption: Optional[str] = None, file_name: Optional[str] = None,
|
||||
reply_to: Optional[str] = None,
|
||||
) -> SendResult:
|
||||
if not chat_id:
|
||||
return SendResult(success=False, error="chat_id is required")
|
||||
try:
|
||||
prepared = await self._prepare_outbound_media(media_source, file_name=file_name)
|
||||
except FileNotFoundError as exc:
|
||||
return SendResult(success=False, error=str(exc))
|
||||
except Exception as exc:
|
||||
logger.error("[%s] Failed to prepare outbound media %s: %s", self.name, media_source, exc)
|
||||
return SendResult(success=False, error=str(exc))
|
||||
|
||||
if prepared["rejected"]:
|
||||
await self._send_followup_markdown(chat_id, f"⚠️ {prepared['reject_reason']}", reply_to=reply_to)
|
||||
return SendResult(success=False, error=prepared["reject_reason"])
|
||||
|
||||
reply_req_id = self._reply_req_id_for_message(reply_to)
|
||||
if not reply_req_id and chat_id in self._last_chat_req_ids:
|
||||
reply_req_id = self._last_chat_req_ids[chat_id]
|
||||
# Media MUST use the proactive path when a stream was/is active for the
|
||||
# chat: passive replyMedia cannot overwrite a replyStream thinking bubble
|
||||
# and the stream "owns" the req_id (server ignores or never acks).
|
||||
if self._find_active_turn_for_chat(chat_id) or chat_id in self._stream_expired_chats:
|
||||
reply_req_id = None
|
||||
|
||||
try:
|
||||
upload_result = await self._upload_media_bytes(prepared["data"], prepared["final_type"], prepared["file_name"])
|
||||
logger.info("[%s] upload_media_bytes OK: media_id=%s type=%s", self.name, upload_result.get("media_id"), prepared["final_type"])
|
||||
if reply_req_id:
|
||||
media_response = await self._send_reply_media_message(reply_req_id, prepared["final_type"], upload_result["media_id"])
|
||||
logger.info("[%s] send_reply_media OK: %s", self.name, media_response)
|
||||
else:
|
||||
media_response = await self._send_media_message(chat_id, prepared["final_type"], upload_result["media_id"])
|
||||
logger.info("[%s] send_media_message OK: %s", self.name, media_response)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("[%s] TIMEOUT in _send_media_source for %s", self.name, media_source)
|
||||
return SendResult(success=False, error="Timeout sending media to WeCom")
|
||||
except Exception as exc:
|
||||
logger.error("[%s] Failed to send media %s: %s", self.name, media_source, exc)
|
||||
return SendResult(success=False, error=str(exc))
|
||||
|
||||
caption_result = downgrade_result = None
|
||||
if caption:
|
||||
caption_result = await self._send_followup_markdown(chat_id, caption, reply_to=reply_to)
|
||||
if prepared["downgraded"] and prepared["downgrade_note"]:
|
||||
downgrade_result = await self._send_followup_markdown(chat_id, f"ℹ️ {prepared['downgrade_note']}", reply_to=reply_to)
|
||||
|
||||
return SendResult(
|
||||
success=True,
|
||||
message_id=self._payload_req_id(media_response) or uuid.uuid4().hex[:12],
|
||||
raw_response={
|
||||
"upload": upload_result,
|
||||
"media": media_response,
|
||||
"caption": caption_result.raw_response if caption_result else None,
|
||||
"caption_error": caption_result.error if caption_result and not caption_result.success else None,
|
||||
"downgrade": downgrade_result.raw_response if downgrade_result else None,
|
||||
"downgrade_error": downgrade_result.error if downgrade_result and not downgrade_result.success else None,
|
||||
},
|
||||
)
|
||||
|
||||
async def send_image(
|
||||
self, chat_id: str, image_url: str, caption: Optional[str] = None, reply_to: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> SendResult:
|
||||
del metadata
|
||||
result = await self._send_media_source(chat_id=chat_id, media_source=image_url, caption=caption, reply_to=reply_to)
|
||||
if result.success or not self._looks_like_url(image_url):
|
||||
return result
|
||||
logger.warning("[%s] Falling back to text send for image URL %s: %s", self.name, image_url, result.error)
|
||||
fallback_text = f"{caption}\n{image_url}" if caption else image_url
|
||||
return await self.send(chat_id=chat_id, content=fallback_text, reply_to=reply_to)
|
||||
|
||||
async def send_image_file(self, chat_id: str, image_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs) -> SendResult:
|
||||
del kwargs
|
||||
return await self._send_media_source(chat_id=chat_id, media_source=image_path, caption=caption, reply_to=reply_to)
|
||||
|
||||
async def send_document(
|
||||
self, chat_id: str, file_path: str, caption: Optional[str] = None, file_name: Optional[str] = None,
|
||||
reply_to: Optional[str] = None, **kwargs,
|
||||
) -> SendResult:
|
||||
del kwargs
|
||||
logger.info("[%s] send_document called: chat=%s file=%s", self.name, chat_id, file_path)
|
||||
return await self._send_media_source(
|
||||
chat_id=chat_id, media_source=file_path, caption=caption, file_name=file_name, reply_to=reply_to,
|
||||
)
|
||||
|
||||
async def send_voice(self, chat_id: str, audio_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs) -> SendResult:
|
||||
del kwargs
|
||||
return await self._send_media_source(chat_id=chat_id, media_source=audio_path, caption=caption, reply_to=reply_to)
|
||||
|
||||
async def send_video(self, chat_id: str, video_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs) -> SendResult:
|
||||
del kwargs
|
||||
return await self._send_media_source(chat_id=chat_id, media_source=video_path, caption=caption, reply_to=reply_to)
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Per-chat FIFO send queues with token-bucket rate limiting for WeCom.
|
||||
|
||||
Mirrors OpenClaw's chat-queue.ts (serial per chat) plus a token bucket that
|
||||
keeps each chat under WeCom's 30 msgs/min/chat limit (errcode 846607). Two
|
||||
lanes per chat: a normal lane and a high-priority control lane (approval
|
||||
prompts, finalize frames, error notices) backed by a reserved token pool.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from typing import Dict
|
||||
|
||||
logger = logging.getLogger("plugins.platforms.wecom.adapter")
|
||||
|
||||
|
||||
class ChatSendQueueMixin:
|
||||
"""Expects ``_chat_queues/_chat_workers/_control_queues/_control_workers/_chat_token_usage`` dicts."""
|
||||
|
||||
# Token bucket: 30 tokens/min per chat, split between normal and reserved (control) quota.
|
||||
_BUCKET_MAX_TOKENS = 30
|
||||
_BUCKET_NORMAL_TOKENS = 24
|
||||
_BUCKET_RESERVED_TOKENS = 6
|
||||
|
||||
def _get_token_usage(self, chat_id: str) -> Dict[str, float]:
|
||||
"""Get or create token usage tracking for a chat."""
|
||||
key = str(chat_id or "").strip()
|
||||
if key not in self._chat_token_usage:
|
||||
self._chat_token_usage[key] = {"normal": 0.0, "reserved": 0.0, "last_reset": time.monotonic()}
|
||||
return self._chat_token_usage[key]
|
||||
|
||||
def _bucket_try_consume(self, chat_id: str, is_control: bool = False) -> float:
|
||||
"""Consume one token. Returns 0 if available, else seconds until the next minute window.
|
||||
|
||||
Normal messages only use the normal quota; control messages use normal
|
||||
quota first (don't waste reserved), then the reserved pool.
|
||||
"""
|
||||
usage = self._get_token_usage(chat_id)
|
||||
now = time.monotonic()
|
||||
if now - usage["last_reset"] > 60.0: # reset counters every minute
|
||||
usage["normal"] = 0.0
|
||||
usage["reserved"] = 0.0
|
||||
usage["last_reset"] = now
|
||||
if usage["normal"] < self._BUCKET_NORMAL_TOKENS:
|
||||
usage["normal"] += 1.0
|
||||
return 0.0
|
||||
if is_control and usage["reserved"] < self._BUCKET_RESERVED_TOKENS:
|
||||
usage["reserved"] += 1.0
|
||||
return 0.0
|
||||
return 60.0 - (now - usage["last_reset"])
|
||||
|
||||
async def _enqueue_chat_send(self, chat_id: str, coro_factory, is_control: bool = False):
|
||||
"""Enqueue a send task for a chat and await its result (FIFO per chat, parallel across chats).
|
||||
|
||||
Control-lane sends bypass the normal queue so approval prompts are never blocked.
|
||||
"""
|
||||
key = str(chat_id or "").strip()
|
||||
lane = "control" if is_control else "normal"
|
||||
queues = self._control_queues if is_control else self._chat_queues
|
||||
if key not in queues:
|
||||
logger.debug("[%s] Creating %s queue + worker for chat %s", self.name, lane, key)
|
||||
queues[key] = asyncio.Queue()
|
||||
workers = self._control_workers if is_control else self._chat_workers
|
||||
workers[key] = asyncio.create_task(self._send_worker(key, is_control))
|
||||
queue = queues[key]
|
||||
logger.debug("[%s] Enqueuing send for chat %s (lane=%s, qsize=%d)", self.name, key, lane, queue.qsize())
|
||||
future = asyncio.get_running_loop().create_future()
|
||||
await queue.put((coro_factory, future))
|
||||
return await future
|
||||
|
||||
async def _send_worker(self, chat_key: str, is_control: bool) -> None:
|
||||
"""Per-chat worker: drain one lane's queue under the token bucket."""
|
||||
if is_control:
|
||||
queue = self._control_queues[chat_key]
|
||||
else:
|
||||
queue = self._chat_queues[chat_key]
|
||||
logger.debug("[%s] Normal send worker started for chat %s", self.name, chat_key)
|
||||
try:
|
||||
while True:
|
||||
coro_factory, future = await queue.get()
|
||||
try:
|
||||
wait = self._bucket_try_consume(chat_key, is_control)
|
||||
if wait > 0:
|
||||
if not is_control:
|
||||
logger.debug(
|
||||
"[%s] Normal worker rate-limited for chat %s, waiting %.1fs",
|
||||
self.name, chat_key, wait,
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
self._bucket_try_consume(chat_key, is_control) # re-consume after wait
|
||||
result = await coro_factory()
|
||||
if not future.done():
|
||||
future.set_result(result)
|
||||
except Exception as exc:
|
||||
if not future.done():
|
||||
future.set_exception(exc)
|
||||
finally:
|
||||
queue.task_done()
|
||||
except asyncio.CancelledError:
|
||||
while not queue.empty():
|
||||
try:
|
||||
_, future = queue.get_nowait()
|
||||
if not future.done():
|
||||
future.set_exception(RuntimeError("WeCom adapter shutting down"))
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
@@ -0,0 +1,562 @@
|
||||
"""WeCom native streaming (``msgtype: stream`` via aibot_respond_msg).
|
||||
|
||||
Per-turn stream state, per-req_id ack tracking (official SDK's
|
||||
replyStreamNonBlocking semantics), the stream-level keep-alive heartbeat and
|
||||
the finalize clock fallback. Mixed into :class:`WeComAdapter`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
logger = logging.getLogger("plugins.platforms.wecom.adapter")
|
||||
|
||||
APP_CMD_RESPONSE = "aibot_respond_msg"
|
||||
|
||||
# WeCom binds a ~6-minute lifetime to each reply stream (stream_id + req_id);
|
||||
# the connection-level ping does NOT refresh it. Past that window updates come
|
||||
# back 846608 (stream update window) / 846604 (req_id reply-request window) —
|
||||
# both mean the reply flow is dead and further frames will be rejected.
|
||||
STREAM_EXPIRED_ERRCODE = 846608
|
||||
STREAM_REQUEST_EXPIRED_ERRCODE = 846604
|
||||
STREAM_NOT_SUBSCRIBED_ERRCODE = 846609 # ws connection lost the subscription
|
||||
# 6000 = finalize raced a newer frame on the same stream_id: the bubble was
|
||||
# ALREADY replaced, so for a finalize frame this is benign, not a failure.
|
||||
STREAM_VERSION_CONFLICT_ERRCODE = 6000
|
||||
MAX_STREAM_CONTENT_LENGTH = 20480 # WeCom server-enforced byte limit per frame
|
||||
# WeCom SDK has a 100-frame per-reqId queue; cap intermediates at 85 (matches
|
||||
# the openclaw plugin) so the finalize frame always has room. Past the cap
|
||||
# intermediates are silently dropped — finalize still sends unconditionally.
|
||||
MAX_INTERMEDIATE_FRAMES = 85
|
||||
|
||||
# Two independent defences against the 6-min stream window, both defaulting
|
||||
# to the safe side (see docs/wecom-stream-keepalive-*.md):
|
||||
# Layer 2 — clock fallback (always on): finalize declines the finish=true
|
||||
# frame once the stream is older than STREAM_SAFE_DURATION_SECONDS and
|
||||
# returns False so the consumer's send() fallback delivers the content.
|
||||
# Layer 1 — keep-alive heartbeat (OFF by default): every
|
||||
# STREAM_KEEPALIVE_INTERVAL_SECONDS re-send the accumulated text as a
|
||||
# finish=false frame. Never sends a placeholder. Off by default because an
|
||||
# extra intermediate frame widens the ack race the double-send
|
||||
# coordination depends on.
|
||||
STREAM_SAFE_DURATION_SECONDS = 330.0
|
||||
STREAM_KEEPALIVE_INTERVAL_SECONDS = 120.0
|
||||
STREAM_KEEPALIVE_ENABLED_DEFAULT = False
|
||||
|
||||
|
||||
class WeComStreamExpiredError(RuntimeError):
|
||||
"""Raised on errcode 846608/846604: the stream/req_id reply flow is dead.
|
||||
|
||||
Callers must fall back to a proactive ``aibot_send_msg``.
|
||||
"""
|
||||
|
||||
def __init__(self, errcode: int = STREAM_EXPIRED_ERRCODE, errmsg: str = ""):
|
||||
super().__init__(f"WeCom stream expired (errcode={errcode}): {errmsg or 'no detail'}")
|
||||
self.errcode = errcode
|
||||
self.errmsg = errmsg
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReplyFrame:
|
||||
"""A reply frame awaiting its aibot_respond_msg ack (FIFO per req_id)."""
|
||||
body: Dict[str, Any]
|
||||
future: asyncio.Future
|
||||
is_final: bool = False
|
||||
sent_at: Optional[float] = None
|
||||
|
||||
|
||||
class ReplyQueue:
|
||||
"""Per-req_id pending-ack tracker: intermediates skip while an ack is pending, finals wait."""
|
||||
def __init__(self, req_id: str):
|
||||
self.req_id = req_id
|
||||
self.pending_ack: Optional[ReplyFrame] = None
|
||||
|
||||
|
||||
class StreamTurn:
|
||||
"""Per-turn stream state so concurrent messages never share a stream."""
|
||||
def __init__(self, chat_id: str, req_id: str):
|
||||
self.chat_id = chat_id
|
||||
self.req_id = req_id
|
||||
self.stream_id = f"stream_{uuid.uuid4().hex[:12]}"
|
||||
self.accumulated_text = ""
|
||||
self.finalized = False
|
||||
self.seeded = False # seed frame sent (prevents double seed → errcode 6000)
|
||||
self.start_time = time.monotonic()
|
||||
self.expired = False
|
||||
# Last content ACTUALLY sent (not skipped) — finalize uses it to avoid a
|
||||
# duplicate-content final frame that WeCom silently drops.
|
||||
self.last_sent_content: str = ""
|
||||
self._intermediate_frames_sent: int = 0
|
||||
# Keep-alive TimerHandle; MUST be cancelled on every turn-exit path
|
||||
# (finalize / expired / error / cleanup) so it never fires on a dead turn.
|
||||
self.keepalive_handle: Optional[asyncio.TimerHandle] = None
|
||||
|
||||
|
||||
def _stream_of(body: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return body.get("stream", {}) if isinstance(body.get("stream"), dict) else {}
|
||||
|
||||
|
||||
class WeComStreamMixin:
|
||||
"""Native streaming for WeComAdapter (expects ``_ws``, ``_send_json``, ``_reply_queues``,
|
||||
``_stream_turns``, ``_stream_expired_chats``, ``_last_chat_req_ids`` and the
|
||||
``_stream_*`` config attributes set in ``__init__``)."""
|
||||
|
||||
MAX_STREAM_CONTENT_LENGTH = MAX_STREAM_CONTENT_LENGTH
|
||||
|
||||
# Ack timeout matches the official plugin's REPLY_SEND_TIMEOUT_MS = 15_000; a
|
||||
# shorter window widened the race where the final-frame ack is still in
|
||||
# flight while the gateway's normal final send fires → duplicate messages.
|
||||
_REPLY_ACK_TIMEOUT = 15.0
|
||||
|
||||
# ── Per-req_id reply queue (ack tracking) ────────────────────────────
|
||||
|
||||
async def _send_reply_queued(
|
||||
self, reply_req_id: str, body: Dict[str, Any], *, is_final: bool = False, skip_if_pending: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""Send a reply via aibot_respond_msg with per-req_id ack tracking.
|
||||
|
||||
is_final: wait for any pending ack before sending, then await our own ack.
|
||||
skip_if_pending: return ``{"skipped": True}`` if a prior frame's ack is pending.
|
||||
"""
|
||||
if not self._ws or self._ws.closed:
|
||||
raise RuntimeError("WeCom websocket is not connected")
|
||||
normalized = str(reply_req_id or "").strip()
|
||||
if not normalized:
|
||||
raise ValueError("reply_req_id is required")
|
||||
|
||||
queue = self._reply_queues.get(normalized)
|
||||
if queue is None:
|
||||
queue = ReplyQueue(normalized)
|
||||
self._reply_queues[normalized] = queue
|
||||
|
||||
if skip_if_pending and queue.pending_ack is not None:
|
||||
return {"skipped": True, "errcode": 0, "errmsg": "pending_ack"}
|
||||
|
||||
if is_final and queue.pending_ack is not None:
|
||||
pending_frame = queue.pending_ack
|
||||
_pending_stream = _stream_of(pending_frame.body)
|
||||
pending_desc = (self.name, normalized, _pending_stream.get("id", "N/A"), _pending_stream.get("finish", "N/A"))
|
||||
logger.debug(
|
||||
"[%s] _send_reply_queued: final waiting for pending ack drain — "
|
||||
"req_id=%s pending_stream_id=%s pending_finish=%s pending_sent_at=%.1fs_ago",
|
||||
*pending_desc, time.monotonic() - (pending_frame.sent_at or time.monotonic()),
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.shield(pending_frame.future), timeout=self._REPLY_ACK_TIMEOUT)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
"[%s] Reply ack timeout waiting for pending (req_id=%s) — "
|
||||
"pending_stream_id=%s pending_finish=%s elapsed=%.1fs. "
|
||||
"Possible causes: ack cmd filtered, ack req_id mismatch, or WeCom did not ack.",
|
||||
*pending_desc, time.monotonic() - (pending_frame.sent_at or time.monotonic()),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
queue.pending_ack = None # resolved or timed out either way
|
||||
|
||||
future: asyncio.Future = asyncio.get_running_loop().create_future()
|
||||
frame = ReplyFrame(body=body, future=future, is_final=is_final)
|
||||
frame.sent_at = time.monotonic()
|
||||
# Register pending BEFORE sending so an ack arriving mid-send is routed.
|
||||
# Re-attach `queue` too: while the final frame awaited the drain above, the
|
||||
# intermediate ack may have popped the whole queue out of _reply_queues,
|
||||
# leaving our local reference orphaned (its ack would then be Unrouted →
|
||||
# 15s timeout).
|
||||
self._reply_queues[normalized] = queue
|
||||
queue.pending_ack = frame
|
||||
|
||||
_stream_info = _stream_of(body)
|
||||
logger.debug(
|
||||
"[%s] _send_reply_queued: req_id=%s is_final=%s skip_if_pending=%s stream_id=%s finish=%s content_len=%d",
|
||||
self.name, normalized, is_final, skip_if_pending,
|
||||
_stream_info.get("id", "N/A"), _stream_info.get("finish", "N/A"), len(_stream_info.get("content", "") or ""),
|
||||
)
|
||||
|
||||
try:
|
||||
await self._send_json({"cmd": APP_CMD_RESPONSE, "headers": {"req_id": normalized}, "body": body})
|
||||
except Exception:
|
||||
# Nobody awaits the future on this branch — cancel it rather than
|
||||
# leave a "Future exception was never retrieved" log.
|
||||
if queue.pending_ack is frame:
|
||||
queue.pending_ack = None
|
||||
self._reply_queues.pop(normalized, None)
|
||||
if not future.done():
|
||||
future.cancel()
|
||||
raise
|
||||
|
||||
if not is_final:
|
||||
# Fire-and-forget; pending_ack stays registered so later frames can skip.
|
||||
return {"errcode": 0, "errmsg": "sent_nonblocking"}
|
||||
try:
|
||||
return await asyncio.wait_for(future, timeout=self._REPLY_ACK_TIMEOUT)
|
||||
except asyncio.TimeoutError:
|
||||
# The bytes went out (send did not raise) but the ack is late — in
|
||||
# practice WeCom has already rendered the message. Raising here made
|
||||
# the upper layer fall back to a markdown send and produced duplicates;
|
||||
# match the official plugin: warn and treat as delivered.
|
||||
logger.warning(
|
||||
"[%s] Final frame ack timeout (req_id=%s) — treating as "
|
||||
"delivered (matches official wecom-openclaw-plugin "
|
||||
"behaviour). No fallback send.",
|
||||
self.name, normalized,
|
||||
)
|
||||
return {"errcode": 0, "errmsg": "ack_timeout_assumed_delivered", "ack_pending": True}
|
||||
finally:
|
||||
self._release_pending(queue, normalized, frame)
|
||||
|
||||
def _release_pending(self, queue: ReplyQueue, req_id: str, frame: ReplyFrame) -> None:
|
||||
"""Clear ``frame`` if it is still the pending ack; drop the queue once empty."""
|
||||
if queue.pending_ack is frame:
|
||||
queue.pending_ack = None
|
||||
if queue.pending_ack is None:
|
||||
self._reply_queues.pop(req_id, None)
|
||||
|
||||
def _resolve_reply_ack(self, req_id: str, payload: Dict[str, Any]) -> bool:
|
||||
"""Resolve a pending reply ack. Returns True if handled."""
|
||||
queue = self._reply_queues.get(req_id)
|
||||
if queue is None or queue.pending_ack is None:
|
||||
return False
|
||||
frame = queue.pending_ack
|
||||
if not frame.future.done():
|
||||
_body = payload.get("body", {}) if isinstance(payload.get("body"), dict) else {}
|
||||
logger.debug(
|
||||
"[%s] _resolve_reply_ack: resolved req_id=%s is_final=%s "
|
||||
"elapsed=%.2fs errcode=%s",
|
||||
self.name, req_id, frame.is_final,
|
||||
time.monotonic() - (frame.sent_at or time.monotonic()),
|
||||
_body.get("errcode", "N/A"),
|
||||
)
|
||||
frame.future.set_result(payload)
|
||||
self._release_pending(queue, req_id, frame)
|
||||
return True
|
||||
|
||||
def _fail_reply_queues(self, error: Exception) -> None:
|
||||
"""Fail all pending reply acks (disconnect/error)."""
|
||||
for queue in list(self._reply_queues.values()):
|
||||
if queue.pending_ack and not queue.pending_ack.future.done():
|
||||
queue.pending_ack.future.set_exception(error)
|
||||
self._reply_queues.clear()
|
||||
|
||||
# ── Turn registry ────────────────────────────────────────────────────
|
||||
|
||||
def _resolve_stream_req_id(self, chat_id: str, reply_to: Optional[str]) -> Optional[str]:
|
||||
"""Explicit ``reply_to`` (cached message id) → last inbound req_id for the chat → None."""
|
||||
req_id = self._reply_req_id_for_message(reply_to)
|
||||
if req_id:
|
||||
return req_id
|
||||
return self._last_chat_req_ids.get(str(chat_id or "").strip()) or None
|
||||
|
||||
@staticmethod
|
||||
def _cancel_keepalive(turn: StreamTurn) -> None:
|
||||
if turn.keepalive_handle is not None:
|
||||
try:
|
||||
turn.keepalive_handle.cancel()
|
||||
except Exception:
|
||||
pass
|
||||
turn.keepalive_handle = None
|
||||
|
||||
def _retire_turn(self, turn: StreamTurn, turn_id: Optional[str]) -> None:
|
||||
"""Single choke point for "turn is dead": cancel the timer, then drop it from the registry."""
|
||||
self._cancel_keepalive(turn)
|
||||
self._stream_turns.pop(f"{turn.chat_id}:{turn_id or turn.req_id}", None)
|
||||
|
||||
def _expire_turn(self, turn: StreamTurn, turn_id: Optional[str]) -> None:
|
||||
turn.expired = True
|
||||
self._retire_turn(turn, turn_id)
|
||||
self._stream_expired_chats.add(turn.chat_id)
|
||||
|
||||
def _find_active_turn_for_chat(self, chat_id: str) -> Optional[StreamTurn]:
|
||||
for turn in self._stream_turns.values():
|
||||
if turn.chat_id == chat_id and not turn.finalized:
|
||||
return turn
|
||||
return None
|
||||
|
||||
# ── Stream-level keep-alive (Layer 1) ────────────────────────────────
|
||||
|
||||
def _arm_keepalive(self, turn: StreamTurn, *, turn_id: Optional[str]) -> None:
|
||||
"""Arm the keep-alive timer if enabled and not already armed (idempotent)."""
|
||||
if not self._stream_keepalive_enabled or turn.finalized or turn.expired:
|
||||
return
|
||||
if turn.keepalive_handle is not None:
|
||||
return
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return
|
||||
turn.keepalive_handle = loop.call_later(
|
||||
self._stream_keepalive_interval_seconds, self._on_keepalive_fire, turn, turn_id,
|
||||
)
|
||||
|
||||
def _on_keepalive_fire(self, turn: StreamTurn, turn_id: Optional[str]) -> None:
|
||||
turn.keepalive_handle = None
|
||||
if turn.finalized or turn.expired:
|
||||
return
|
||||
try:
|
||||
asyncio.ensure_future(self._keepalive_send(turn, turn_id))
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
async def _keepalive_send(self, turn: StreamTurn, turn_id: Optional[str]) -> None:
|
||||
"""Re-send the accumulated text as finish=false to refresh the server window, then re-arm.
|
||||
|
||||
Never sends a placeholder: with no accumulated text the tick is skipped
|
||||
(Layer 2 handles content-less turns). On 846604/846608 the turn is
|
||||
retired so finalize takes the Layer 2 fallback; no re-arm.
|
||||
"""
|
||||
if turn.finalized or turn.expired:
|
||||
return
|
||||
if turn._intermediate_frames_sent >= MAX_INTERMEDIATE_FRAMES:
|
||||
return # no room left for intermediates; let finalize / Layer 2 run
|
||||
content = turn.accumulated_text or ""
|
||||
if not content.strip():
|
||||
self._arm_keepalive(turn, turn_id=turn_id)
|
||||
return
|
||||
try:
|
||||
await self._send_stream_reply(turn.req_id, turn.stream_id, content, finish=False)
|
||||
except WeComStreamExpiredError:
|
||||
self._expire_turn(turn, turn_id)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"[%s] keep-alive send failed (chat=%s, turn=%s): %s",
|
||||
self.name, turn.chat_id, turn.stream_id, exc,
|
||||
)
|
||||
self._arm_keepalive(turn, turn_id=turn_id) # transient — retry next interval
|
||||
return
|
||||
turn.last_sent_content = content
|
||||
self._arm_keepalive(turn, turn_id=turn_id)
|
||||
|
||||
# ── Frame sending ────────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _truncate_stream_content(content: str, limit: int) -> str:
|
||||
"""Truncate to ``limit`` UTF-8 bytes (WeCom caps frames by bytes, not codepoints)."""
|
||||
encoded = content.encode("utf-8")
|
||||
if len(encoded) <= limit:
|
||||
return content
|
||||
return encoded[:limit].decode("utf-8", errors="ignore")
|
||||
|
||||
async def _send_stream_reply(
|
||||
self, reply_req_id: str, stream_id: str, content: str, finish: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""Send one ``msgtype: "stream"`` frame.
|
||||
|
||||
Intermediate frames are non-blocking with skip-if-pending (cumulative
|
||||
text means nothing is lost). The final frame drains any pending ack
|
||||
first, then awaits its own ack so 846608/6000 are detected reliably.
|
||||
Raises WeComStreamExpiredError on 846608/846604.
|
||||
"""
|
||||
truncated = self._truncate_stream_content(content or "", self.MAX_STREAM_CONTENT_LENGTH)
|
||||
if len(content or "") != len(truncated):
|
||||
logger.warning("[%s] Stream content truncated for stream_id=%s", self.name, stream_id)
|
||||
body: Dict[str, Any] = {"msgtype": "stream", "stream": {"id": stream_id, "finish": bool(finish), "content": truncated}}
|
||||
if not finish:
|
||||
return await self._send_reply_queued(reply_req_id, body, is_final=False, skip_if_pending=True)
|
||||
|
||||
response = await self._send_reply_queued(reply_req_id, body, is_final=True, skip_if_pending=False)
|
||||
errcode = response.get("errcode", 0)
|
||||
if errcode in (STREAM_EXPIRED_ERRCODE, STREAM_REQUEST_EXPIRED_ERRCODE):
|
||||
raise WeComStreamExpiredError(errcode=errcode, errmsg=str(response.get("errmsg") or ""))
|
||||
if errcode == STREAM_VERSION_CONFLICT_ERRCODE:
|
||||
# Content is already on screen; raising would pop the turn and cause
|
||||
# a duplicate standalone send(). Absorbing makes finalize retry safe.
|
||||
logger.info(
|
||||
"[%s] finalize hit errcode 6000 (version conflict) — bubble "
|
||||
"already replaced by a newer frame; treating as delivered.",
|
||||
self.name,
|
||||
)
|
||||
return response
|
||||
self._raise_for_wecom_error(response, "send stream reply")
|
||||
return response
|
||||
|
||||
async def send_stream_frame(
|
||||
self, text: str, *, finalize: bool = False, chat_id: Optional[str] = None, reply_to: Optional[str] = None, **kwargs,
|
||||
) -> bool:
|
||||
"""Entry point for the gateway streaming consumer.
|
||||
|
||||
First call for a turn resolves the req_id, creates the StreamTurn and
|
||||
seeds the typing bubble; later calls push cumulative text (not deltas);
|
||||
``finalize=True`` closes the stream and drops turn state. ``turn_id``
|
||||
(kwarg) keys the turn by (chat, turn_id) so concurrent consumers
|
||||
(/background, subagents) never share a stream.
|
||||
|
||||
Returns False when the stream is unavailable (no req_id, expired,
|
||||
transport error) — the caller should fall back to :meth:`send`.
|
||||
"""
|
||||
chat = (chat_id or "").strip()
|
||||
if not chat:
|
||||
logger.warning("[%s] send_stream_frame: chat_id required", self.name)
|
||||
return False
|
||||
turn_id = kwargs.get("turn_id")
|
||||
# Chat-level expiry only blocks NEW turn creation; a known turn_id may
|
||||
# still finalize after another turn in the chat expired.
|
||||
if not turn_id and chat in self._stream_expired_chats:
|
||||
return False
|
||||
if finalize:
|
||||
# Finalize counts toward 30/min — control lane so it is never blocked.
|
||||
return await self._enqueue_chat_send(
|
||||
chat,
|
||||
lambda: self._send_stream_frame_inner(text, chat=chat, reply_to=reply_to, finalize=True, turn_id=turn_id),
|
||||
is_control=True,
|
||||
)
|
||||
# Intermediate frames don't count toward the quota: no queue, no rate limit.
|
||||
return await self._send_stream_frame_inner(text, chat=chat, reply_to=reply_to, finalize=False, turn_id=turn_id)
|
||||
|
||||
def _locate_turn(
|
||||
self, chat: str, reply_to: Optional[str], finalize: bool, turn_id: Optional[str],
|
||||
) -> Optional[StreamTurn]:
|
||||
"""Find or create the StreamTurn for a frame; None means "stream unavailable".
|
||||
|
||||
A turn locks to its req_id at creation even if ``_last_chat_req_ids``
|
||||
changes mid-turn (e.g. the user sends /approve).
|
||||
"""
|
||||
if turn_id:
|
||||
turn = self._stream_turns.get(f"{chat}:{turn_id}")
|
||||
if turn:
|
||||
return turn
|
||||
# finalize must NOT create a turn: if it was cleaned up (e.g. 6000) the
|
||||
# caller should fall back rather than send a fresh seed + finish.
|
||||
if finalize:
|
||||
logger.debug(
|
||||
"[%s] send_stream_frame: cannot finalize non-existent turn (turn_id=%s, chat=%s)",
|
||||
self.name, turn_id, chat,
|
||||
)
|
||||
return None
|
||||
else:
|
||||
# No turn_id (direct callers): reuse the chat's active turn if any.
|
||||
existing_turn = self._find_active_turn_for_chat(chat)
|
||||
if existing_turn and not existing_turn.finalized:
|
||||
logger.debug(
|
||||
"[%s] send_stream_frame: reusing existing turn %s for chat %s",
|
||||
self.name, existing_turn.stream_id, chat,
|
||||
)
|
||||
return existing_turn
|
||||
|
||||
suffix = f" (turn_id={turn_id})" if turn_id else ""
|
||||
if chat in self._stream_expired_chats:
|
||||
logger.debug("[%s] send_stream_frame: chat %s is expired, cannot create new turn%s", self.name, chat, suffix)
|
||||
return None
|
||||
req_id = self._resolve_stream_req_id(chat, reply_to)
|
||||
if not req_id:
|
||||
logger.debug("[%s] send_stream_frame: no req_id available for chat %s%s", self.name, chat, suffix)
|
||||
return None
|
||||
key = f"{chat}:{turn_id or req_id}"
|
||||
turn = (None if turn_id else self._stream_turns.get(key)) or StreamTurn(chat, req_id)
|
||||
self._stream_turns[key] = turn
|
||||
logger.debug(
|
||||
"[%s] send_stream_frame: created new turn %s (%s) for chat %s",
|
||||
self.name, turn.stream_id, f"turn_id={turn_id}, req_id={req_id}" if turn_id else f"req_id={req_id}", chat,
|
||||
)
|
||||
return turn
|
||||
|
||||
async def _send_stream_frame_inner(
|
||||
self, text: str, *, chat: str, reply_to: Optional[str] = None, finalize: bool = False, turn_id: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Stream frame logic with per-turn state (see ``send_stream_frame``)."""
|
||||
turn: Optional[StreamTurn] = None
|
||||
try:
|
||||
turn = self._locate_turn(chat, reply_to, finalize, turn_id)
|
||||
if turn is None or turn.expired:
|
||||
return False
|
||||
|
||||
if not turn.seeded and not turn.finalized:
|
||||
# Seed with the official plugin's THINKING_MESSAGE (<think></think>)
|
||||
# so the client shows a reasoning turn; the seeded flag prevents a
|
||||
# double seed (errcode 6000) since the consumer seeds too.
|
||||
await self._send_stream_reply(turn.req_id, turn.stream_id, "<think></think>", finish=False)
|
||||
turn.seeded = True
|
||||
self._arm_keepalive(turn, turn_id=turn_id)
|
||||
if not text and not finalize:
|
||||
return True # consumer's explicit seed call — nothing more to send
|
||||
|
||||
if finalize:
|
||||
# Layer 2 clock fallback: an old stream would almost certainly hit
|
||||
# 846604/846608 on finish=true, so decline up front and let the
|
||||
# consumer's send() fallback deliver exactly once. SKIPPED when
|
||||
# Layer 1 keep-alive is on — the heartbeat has been refreshing the
|
||||
# window, so age alone does not mean dead, and declining a live
|
||||
# stream would re-deliver content already on screen. A truly dead
|
||||
# stream still raises WeComStreamExpiredError below.
|
||||
if not self._stream_keepalive_enabled:
|
||||
stream_age = time.monotonic() - turn.start_time
|
||||
if stream_age >= self._stream_safe_duration_seconds:
|
||||
logger.info(
|
||||
"[%s] Stream age %.0fs >= safe duration %.0fs for chat "
|
||||
"%s — declining finalize frame, falling back to "
|
||||
"proactive send (Layer 2 clock fallback).",
|
||||
self.name, stream_age,
|
||||
self._stream_safe_duration_seconds, chat,
|
||||
)
|
||||
self._expire_turn(turn, turn_id)
|
||||
return False
|
||||
|
||||
self._cancel_keepalive(turn)
|
||||
# WeCom silently drops (no ack) a final frame identical to the last
|
||||
# intermediate — append a zero-width space so the content differs.
|
||||
final_text = text
|
||||
if text and text == turn.last_sent_content:
|
||||
final_text = text + "\u200b"
|
||||
await self._send_stream_reply(turn.req_id, turn.stream_id, final_text, finish=True)
|
||||
turn.finalized = True
|
||||
self._stream_turns.pop(f"{chat}:{turn_id or turn.req_id}", None)
|
||||
else:
|
||||
# Fire-and-forget: the gateway decides when to push (identity dedup
|
||||
# in stream_consumer.py); no adapter-side buffering.
|
||||
turn.accumulated_text = text
|
||||
if turn._intermediate_frames_sent >= MAX_INTERMEDIATE_FRAMES:
|
||||
return True # cap reached — drop intermediates; finalize drains the rest
|
||||
if text == turn.last_sent_content:
|
||||
return True
|
||||
await self._send_stream_reply(turn.req_id, turn.stream_id, text, finish=False)
|
||||
turn._intermediate_frames_sent += 1
|
||||
turn.last_sent_content = text
|
||||
return True
|
||||
|
||||
except WeComStreamExpiredError:
|
||||
# An intermediate frame is overwritten by the next cumulative/final
|
||||
# frame anyway; flipping the turn expired here would trip the
|
||||
# consumer's send() fallback and duplicate the bubble. Only a FINAL
|
||||
# frame's expiry means content is genuinely missing.
|
||||
if not finalize:
|
||||
logger.info(
|
||||
"[%s] Intermediate stream frame expired (errcode=%d) for chat %s — dropping frame, stream stays live",
|
||||
self.name, STREAM_EXPIRED_ERRCODE, chat,
|
||||
)
|
||||
return True
|
||||
logger.info(
|
||||
"[%s] Stream expired (errcode=%d) for chat %s — switching to proactive send",
|
||||
self.name, STREAM_EXPIRED_ERRCODE, chat,
|
||||
)
|
||||
if turn is not None:
|
||||
self._expire_turn(turn, turn_id)
|
||||
else:
|
||||
self._stream_expired_chats.add(chat)
|
||||
return False
|
||||
except Exception as exc:
|
||||
if not finalize: # same intermediate/final split as above
|
||||
logger.info(
|
||||
"[%s] Intermediate stream frame failed (chat=%s): %s — dropping frame, stream stays live",
|
||||
self.name, chat, exc,
|
||||
)
|
||||
return True
|
||||
logger.warning("[%s] Stream frame failed (chat=%s): %s", self.name, chat, exc)
|
||||
if turn is not None:
|
||||
self._retire_turn(turn, turn_id)
|
||||
return False
|
||||
|
||||
def supports_native_streaming(
|
||||
self, chat_type: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> bool:
|
||||
"""Stream frames work in DMs and groups alike (groups just need a cached inbound req_id)."""
|
||||
del chat_type, metadata
|
||||
return True
|
||||
|
||||
async def send_typing(self, chat_id: str, metadata=None) -> None:
|
||||
"""No-op: the stream consumer's seed frame triggers WeCom typing; repeated
|
||||
send_typing calls would open orphan streams."""
|
||||
del chat_id, metadata
|
||||
Reference in New Issue
Block a user