diff --git a/gateway/platforms/bluebubbles.py b/gateway/platforms/bluebubbles.py index 6306c92b9e..21e96b23ac 100644 --- a/gateway/platforms/bluebubbles.py +++ b/gateway/platforms/bluebubbles.py @@ -1,11 +1,7 @@ """BlueBubbles iMessage platform adapter. Uses the local BlueBubbles macOS server for outbound REST sends and inbound -webhooks. Supports text messaging, media attachments (images, voice, video, -documents), tapback reactions, typing indicators, and read receipts. - -Architecture based on PR #5869 (benjaminsehl) with inbound attachment -downloading from PR #4588 (YuhangLin). +webhooks: text, media attachments, typing indicators, and read receipts. """ import asyncio @@ -18,19 +14,15 @@ from collections import OrderedDict from datetime import datetime from pathlib import Path from typing import Any, Dict, List, Optional -from urllib.parse import quote +from urllib.parse import parse_qs, quote import httpx from gateway.config import Platform, PlatformConfig +from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret from gateway.platforms.base import ( - BasePlatformAdapter, - MessageEvent, - MessageType, - SendResult, - cache_image_from_bytes, - cache_audio_from_bytes, - cache_document_from_bytes, + BasePlatformAdapter, MessageEvent, MessageType, SendResult, + cache_image_from_bytes, cache_audio_from_bytes, cache_document_from_bytes, ) from .media_cache import ext_for_mime from gateway.platforms.helpers import compile_mention_patterns, strip_markdown @@ -39,103 +31,51 @@ from gateway.platforms.helpers import compile_mention_patterns, strip_markdown # the shared dispatch in gateway.platforms.media_cache. Both maps are # CLOSED: unlisted mimes fall back to .jpg / .mp3 (never mimetypes). _BLUEBUBBLES_IMAGE_EXT_OVERRIDES = { - "image/jpeg": ".jpg", - "image/png": ".png", - "image/gif": ".gif", - "image/webp": ".webp", - "image/heic": ".jpg", # preserves historical bluebubbles mapping - "image/heif": ".jpg", # preserves historical bluebubbles mapping - "image/tiff": ".jpg", # preserves historical bluebubbles mapping + "image/jpeg": ".jpg", "image/png": ".png", "image/gif": ".gif", "image/webp": ".webp", + "image/heic": ".jpg", "image/heif": ".jpg", "image/tiff": ".jpg", # historical mapping } _BLUEBUBBLES_AUDIO_EXT_OVERRIDES = { - "audio/mp3": ".mp3", - "audio/mpeg": ".mp3", - "audio/ogg": ".ogg", - "audio/wav": ".wav", - "audio/x-caf": ".mp3", # preserves historical bluebubbles mapping - "audio/mp4": ".m4a", - "audio/aac": ".m4a", # preserves historical bluebubbles mapping (shared table says .aac) + "audio/mp3": ".mp3", "audio/mpeg": ".mp3", "audio/ogg": ".ogg", "audio/wav": ".wav", + "audio/x-caf": ".mp3", "audio/mp4": ".m4a", + "audio/aac": ".m4a", # historical mapping (shared table says .aac) } -from agent.secret_scope import UnscopedSecretError as _UnscopedSecretError -from agent.secret_scope import get_secret as _scoped_get_secret - - -def _get_scoped_secret(name, default=None): - """Scope-aware credential read with the default-profile startup fallback. - - Secondary profiles construct their adapters under a profile secret - scope -- the scope is authoritative and a scoped miss returns ``default`` - (no cross-profile borrow from ``os.environ``, which may hold another - profile's value). The DEFAULT profile's adapter constructs and sends - *unscoped* under multiplexing, where a bare ``get_secret`` would raise - ``UnscopedSecretError`` and crash this path; there ``os.environ`` is that - profile's own value, so fall back to it. Same pattern as the Slack - ``SLACK_APP_TOKEN`` read (#59739) and - ``gateway/platforms/whatsapp_common.py::_get_wsecret``. - """ - try: - val = _scoped_get_secret(name, default) - except _UnscopedSecretError: - val = os.getenv(name) - return val if val is not None else default - - logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - DEFAULT_WEBHOOK_HOST = "127.0.0.1" -# BlueBubbles webhook events are small JSON/form payloads; attachments come -# through the REST API, not the webhook. 1 MiB is generous headroom while -# keeping oversized/chunked bodies from being buffered unbounded. +# Webhook events are small JSON/form payloads; attachments come through the +# REST API. 1 MiB keeps oversized/chunked bodies from buffering unbounded. _WEBHOOK_MAX_BODY_BYTES = 1_048_576 DEFAULT_WEBHOOK_PORT = 8645 DEFAULT_WEBHOOK_PATH = "/bluebubbles-webhook" MAX_TEXT_LENGTH = 4000 -# BlueBubbles/iMessage does not expose a stable bot mention identity like -# Slack (<@U...>), Telegram (@botname), or Matrix (MXID). When users opt into -# group mention gating without custom aliases, use conservative Hermes wake -# words so `require_mention: true` is a one-line enablement path. +# iMessage has no stable bot mention identity (unlike <@U...>/@botname/MXID), +# so `require_mention: true` without custom aliases uses Hermes wake words. DEFAULT_MENTION_PATTERNS = [ r"(? str: """Redact phone numbers and emails from log output.""" - text = _PHONE_RE.sub("[REDACTED]", text) - text = _EMAIL_RE.sub("[REDACTED]", text) - return text + return _EMAIL_RE.sub("[REDACTED]", _PHONE_RE.sub("[REDACTED]", text)) -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - def check_bluebubbles_requirements() -> bool: try: import aiohttp # noqa: F401 @@ -154,12 +94,21 @@ def _normalize_server_url(raw: str) -> str: return value.rstrip("/") +def _closed_ext(mime: str, overrides: Dict[str, str], fallback: str) -> str: + """Historical maps were closed: unlisted mimes fall back without consulting mimetypes.""" + return ext_for_mime( + mime, overrides=overrides, use_defaults=False, use_mimetypes=False, fallback=fallback + ) or fallback +def _setting(extra: Dict[str, Any], key: str, env: str, default: str = "") -> Any: + """Config ``extra[key]`` wins over env var ``env`` (falsy values fall through).""" + return extra.get(key) or os.getenv(env, default) + + +def _temp_guid() -> str: + return f"temp-{datetime.utcnow().timestamp()}" -# --------------------------------------------------------------------------- -# Adapter -# --------------------------------------------------------------------------- class BlueBubblesAdapter(BasePlatformAdapter): platform = Platform.BLUEBUBBLES @@ -170,22 +119,13 @@ class BlueBubblesAdapter(BasePlatformAdapter): def __init__(self, config: PlatformConfig): super().__init__(config, Platform.BLUEBUBBLES) extra = config.extra or {} - self.server_url = _normalize_server_url( - extra.get("server_url") or os.getenv("BLUEBUBBLES_SERVER_URL", "") - ) + self.server_url = _normalize_server_url(_setting(extra, "server_url", "BLUEBUBBLES_SERVER_URL")) self.password = extra.get("password") or _get_scoped_secret("BLUEBUBBLES_PASSWORD", "") - self.webhook_host = ( - extra.get("webhook_host") - or os.getenv("BLUEBUBBLES_WEBHOOK_HOST", DEFAULT_WEBHOOK_HOST) - ) + self.webhook_host = _setting(extra, "webhook_host", "BLUEBUBBLES_WEBHOOK_HOST", DEFAULT_WEBHOOK_HOST) self.webhook_port = int( - extra.get("webhook_port") - or os.getenv("BLUEBUBBLES_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT)) - ) - self.webhook_path = ( - extra.get("webhook_path") - or os.getenv("BLUEBUBBLES_WEBHOOK_PATH", DEFAULT_WEBHOOK_PATH) + _setting(extra, "webhook_port", "BLUEBUBBLES_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT)) ) + self.webhook_path = _setting(extra, "webhook_path", "BLUEBUBBLES_WEBHOOK_PATH", DEFAULT_WEBHOOK_PATH) if not str(self.webhook_path).startswith("/"): self.webhook_path = f"/{self.webhook_path}" self.send_read_receipts = bool(extra.get("send_read_receipts", True)) @@ -204,9 +144,7 @@ class BlueBubblesAdapter(BasePlatformAdapter): self._helper_connected: bool = False self._guid_cache: OrderedDict[str, str] = OrderedDict() - # ------------------------------------------------------------------ - # API helpers - # ------------------------------------------------------------------ + # --- API helpers --- def _api_url(self, path: str) -> str: sep = "&" if "?" in path else "?" @@ -214,16 +152,10 @@ class BlueBubblesAdapter(BasePlatformAdapter): @staticmethod def _compile_mention_patterns(raw: Any) -> List[re.Pattern]: - """Compile group-mention wake words from config/env. - - ``raw`` is a list (from config or env JSON), a string (raw env var: - JSON list, or comma/newline-separated), or None (use Hermes defaults). - """ + """Compile group-mention wake words; ``raw`` is a list, a raw env string + (JSON list or comma/newline-separated), or None (Hermes defaults).""" return compile_mention_patterns( - raw, - log_prefix="bluebubbles", - defaults=DEFAULT_MENTION_PATTERNS, - logger_=logger, + raw, log_prefix="bluebubbles", defaults=DEFAULT_MENTION_PATTERNS, logger_=logger ) def _message_matches_mention_patterns(self, text: str) -> bool: @@ -232,11 +164,8 @@ class BlueBubblesAdapter(BasePlatformAdapter): return any(pattern.search(text) for pattern in self._mention_patterns) def _clean_mention_text(self, text: str) -> str: - """Strip a leading BlueBubbles wake word before dispatch. - - Custom mention patterns are regular expressions, so stripping only a - leading match avoids deleting ordinary words later in the prompt. - """ + """Strip a leading wake word only — patterns are regexes, so stripping + anywhere later in the prompt could delete ordinary words.""" if not text: return text for pattern in self._mention_patterns: @@ -246,31 +175,51 @@ class BlueBubblesAdapter(BasePlatformAdapter): return cleaned or text return text - async def _api_get(self, path: str) -> Dict[str, Any]: + async def _api_json(self, method: str, path: str, **kwargs) -> Dict[str, Any]: assert self.client is not None - res = await self.client.get(self._api_url(path)) + res = await getattr(self.client, method)(self._api_url(path), **kwargs) res.raise_for_status() return res.json() + async def _api_get(self, path: str) -> Dict[str, Any]: + return await self._api_json("get", path) + async def _api_post(self, path: str, payload: Dict[str, Any]) -> Dict[str, Any]: - assert self.client is not None - res = await self.client.post(self._api_url(path), json=payload) - res.raise_for_status() - return res.json() + return await self._api_json("post", path, json=payload) - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ + async def _post_message(self, path: str, payload: Dict[str, Any]) -> SendResult: + """POST a message payload and wrap the outcome as a SendResult.""" + try: + res = await self._api_post(path, payload) + data = res.get("data") or {} + msg_id = str(data.get("guid") or data.get("messageGuid") or "ok") + return SendResult(success=True, message_id=msg_id, raw_response=res) + except Exception as exc: + return SendResult(success=False, error=str(exc) or type(exc).__name__) + + async def _private_api_chat_call(self, chat_id: str, action: str, method: str) -> bool: + """Fire a private-API chat action (typing/read); True only if the call was made.""" + if not self._private_api_enabled or not self._helper_connected or not self.client: + return False + try: + guid = await self._resolve_chat_guid(chat_id) + if guid: + url = self._api_url(f"/api/v1/chat/{quote(guid, safe='')}/{action}") + await getattr(self.client, method)(url, timeout=5) + return True + except Exception: + pass + return False + + # --- Lifecycle --- async def connect(self, *, is_reconnect: bool = False) -> bool: if not self.server_url or not self.password: - logger.error( - "[bluebubbles] BLUEBUBBLES_SERVER_URL and BLUEBUBBLES_PASSWORD are required" - ) + logger.error("[bluebubbles] BLUEBUBBLES_SERVER_URL and BLUEBUBBLES_PASSWORD are required") return False from aiohttp import web - # Tighter keepalive so idle CLOSE_WAIT drains promptly (#18451). + # Tighter keepalive so idle CLOSE_WAIT drains promptly. from gateway.platforms._http_client_limits import platform_httpx_limits self.client = httpx.AsyncClient(timeout=30.0, limits=platform_httpx_limits()) try: @@ -279,55 +228,37 @@ class BlueBubblesAdapter(BasePlatformAdapter): server_data = (info or {}).get("data", {}) self._private_api_enabled = bool(server_data.get("private_api")) self._helper_connected = bool(server_data.get("helper_connected")) - logger.info( - "[bluebubbles] connected to %s (private_api=%s, helper=%s)", - self.server_url, - self._private_api_enabled, - self._helper_connected, - ) + logger.info("[bluebubbles] connected to %s (private_api=%s, helper=%s)", + self.server_url, self._private_api_enabled, self._helper_connected) except Exception as exc: - logger.error( - "[bluebubbles] cannot reach server at %s: %s", self.server_url, exc - ) + logger.error("[bluebubbles] cannot reach server at %s: %s", self.server_url, exc) if self.client: await self.client.aclose() self.client = None return False - # Explicit body cap: BlueBubbles webhook events are small JSON (or - # form-encoded) payloads. client_max_size makes aiohttp enforce the - # cap on every read path — including chunked requests that carry no - # Content-Length (same pattern as webhook.py / raft, #58536/#58902). + # client_max_size makes aiohttp enforce the cap on every read path, + # including chunked requests with no Content-Length. app = web.Application(client_max_size=_WEBHOOK_MAX_BODY_BYTES) app.router.add_get("/health", lambda _: web.Response(text="ok")) app.router.add_post(self.webhook_path, self._handle_webhook) - # The webhook auth value is carried in the query string because the - # BlueBubbles webhook API cannot send custom headers. Do not let - # aiohttp access logs write that request target to agent.log. + # The webhook auth value rides in the query string (BlueBubbles cannot + # send custom headers) — keep it out of aiohttp access logs. self._runner = web.AppRunner(app, access_log=None) await self._runner.setup() site = web.TCPSite(self._runner, self.webhook_host, self.webhook_port) await site.start() self._mark_connected() - logger.info( - "[bluebubbles] webhook listening on http://%s:%s%s", - self.webhook_host, - self.webhook_port, - self.webhook_path, - ) - - # Register webhook with BlueBubbles server - # This is required for the server to know where to send events + logger.info("[bluebubbles] webhook listening on http://%s:%s%s", + self.webhook_host, self.webhook_port, self.webhook_path) + # The server only sends events to webhooks registered via its API. await self._register_webhook() - # Plugin-registered native handlers (ctx.register_platform_handler). self._wire_plugin_handlers(None) return True async def disconnect(self) -> None: - # Unregister webhook before cleaning up await self._unregister_webhook() - if self.client: await self.client.aclose() self.client = None @@ -338,34 +269,27 @@ class BlueBubblesAdapter(BasePlatformAdapter): @property def _webhook_url(self) -> str: - """Compute the external webhook URL for BlueBubbles registration.""" - host = self.webhook_host - if host in {"0.0.0.0", "127.0.0.1", "localhost", "::"}: - host = "localhost" + """External webhook URL for BlueBubbles registration (local binds → localhost).""" + host = "localhost" if self.webhook_host in _LOCAL_HOSTS else self.webhook_host return f"http://{host}:{self.webhook_port}{self.webhook_path}" + def _webhook_register_url_with(self, password_param: str) -> str: + base = self._webhook_url + return f"{base}?password={password_param}" if self.password else base + @property def _webhook_register_url(self) -> str: - """Webhook URL registered with BlueBubbles, including the password as - a query param so inbound webhook POSTs carry credentials. + """Registered webhook URL, password embedded as a query param. - BlueBubbles posts events to the exact URL registered via - ``/api/v1/webhook``. Its webhook registration API does not support - custom headers, so embedding the password in the URL is the only - way to authenticate inbound webhooks without disabling auth. + BlueBubbles posts to the exact registered URL and its registration + API cannot set custom headers, so this is the only way to + authenticate inbound webhooks without disabling auth. """ - base = self._webhook_url - if self.password: - return f"{base}?password={quote(self.password, safe='')}" - return base + return self._webhook_register_url_with(quote(self.password, safe="")) @property def _webhook_register_url_for_log(self) -> str: - """Webhook registration URL safe for logs.""" - base = self._webhook_url - if self.password: - return f"{base}?password=***" - return base + return self._webhook_register_url_with("***") async def _find_registered_webhooks(self, url: str) -> list: """Return list of BB webhook entries matching *url*.""" @@ -379,120 +303,68 @@ class BlueBubblesAdapter(BasePlatformAdapter): return [] async def _register_webhook(self) -> bool: - """Register this webhook URL with the BlueBubbles server. - - BlueBubbles requires webhooks to be registered via API before - it will send events. Checks for an existing registration first - to avoid duplicates (e.g. after a crash without clean shutdown). - """ + """Register this webhook URL, reusing an existing registration if present + (crash resilience — avoids duplicates after an unclean shutdown).""" if not self.client: return False - webhook_url = self._webhook_register_url - - # Crash resilience — reuse an existing registration if present - existing = await self._find_registered_webhooks(webhook_url) - if existing: - logger.info( - "[bluebubbles] webhook already registered: %s", - self._webhook_register_url_for_log, - ) + log_url = self._webhook_register_url_for_log + if await self._find_registered_webhooks(webhook_url): + logger.info("[bluebubbles] webhook already registered: %s", log_url) return True - - payload = { - "url": webhook_url, - "events": ["new-message", "updated-message"], - } - + payload = {"url": webhook_url, "events": ["new-message", "updated-message"]} try: res = await self._api_post("/api/v1/webhook", payload) status = res.get("status", 0) if 200 <= status < 300: - logger.info( - "[bluebubbles] webhook registered with server: %s", - self._webhook_register_url_for_log, - ) + logger.info("[bluebubbles] webhook registered with server: %s", log_url) return True - else: - logger.warning( - "[bluebubbles] webhook registration returned status %s: %s", - status, - res.get("message"), - ) - return False + logger.warning("[bluebubbles] webhook registration returned status %s: %s", + status, res.get("message")) + return False except Exception as exc: - logger.warning( - "[bluebubbles] failed to register webhook with server: %s", - exc, - ) + logger.warning("[bluebubbles] failed to register webhook with server: %s", exc) return False async def _unregister_webhook(self) -> bool: - """Unregister this webhook URL from the BlueBubbles server. - - Removes *all* matching registrations to clean up any duplicates - left by prior crashes. - """ + """Remove *all* registrations matching our URL (cleans up crash duplicates).""" if not self.client: return False - - webhook_url = self._webhook_register_url removed = False - try: - for wh in await self._find_registered_webhooks(webhook_url): + for wh in await self._find_registered_webhooks(self._webhook_register_url): wh_id = wh.get("id") if wh_id: - res = await self.client.delete( - self._api_url(f"/api/v1/webhook/{wh_id}") - ) + res = await self.client.delete(self._api_url(f"/api/v1/webhook/{wh_id}")) res.raise_for_status() removed = True if removed: - logger.info( - "[bluebubbles] webhook unregistered: %s", - self._webhook_register_url_for_log, - ) + logger.info("[bluebubbles] webhook unregistered: %s", self._webhook_register_url_for_log) except Exception as exc: - logger.debug( - "[bluebubbles] failed to unregister webhook (non-critical): %s", - exc, - ) + logger.debug("[bluebubbles] failed to unregister webhook (non-critical): %s", exc) return removed - # ------------------------------------------------------------------ - # Chat GUID resolution - # ------------------------------------------------------------------ + # --- Chat GUID resolution --- async def _resolve_chat_guid(self, target: str) -> Optional[str]: - """Resolve an email/phone to a BlueBubbles chat GUID. + """Resolve an email/phone to a chat GUID (raw ``a;-;b`` GUIDs pass through). - If *target* already contains a semicolon (raw GUID format like - ``iMessage;-;user@example.com``), it is returned as-is. Otherwise - the adapter queries the BlueBubbles chat list and matches strictly - on ``chatIdentifier`` / ``identifier``. - - Participant membership is intentionally NOT used as a fallback: - the same contact can appear in a 1:1 DM and in any number of group - chats, so a participant match would let an outbound DM reply leak - into a group thread (see #24157). When no exact chat identity - matches, return ``None`` and let the caller create a fresh DM - explicitly via ``_create_chat_for_handle``. + Matches strictly on ``chatIdentifier`` / ``identifier``. Participant + membership is intentionally NOT a fallback: the same contact appears in + a 1:1 DM and any number of groups, so a participant match could leak a + DM reply into a group thread. Return ``None`` and let the caller create + a fresh DM via ``_create_chat_for_handle``. """ target = (target or "").strip() if not target: return None - # Already a raw GUID if ";" in target: return target if target in self._guid_cache: self._guid_cache.move_to_end(target) return self._guid_cache[target] try: - payload = await self._api_post( - "/api/v1/chat/query", - {"limit": 100, "offset": 0}, - ) + payload = await self._api_post("/api/v1/chat/query", {"limit": 100, "offset": 0}) for chat in payload.get("data", []) or []: guid = chat.get("guid") or chat.get("chatGuid") identifier = chat.get("chatIdentifier") or chat.get("identifier") @@ -506,47 +378,27 @@ class BlueBubblesAdapter(BasePlatformAdapter): pass return None - async def _create_chat_for_handle( - self, address: str, message: str - ) -> SendResult: + async def _create_chat_for_handle(self, address: str, message: str) -> SendResult: """Create a new chat by sending the first message to *address*.""" - payload = { - "addresses": [address], - "message": message, - "tempGuid": f"temp-{datetime.utcnow().timestamp()}", - } - try: - res = await self._api_post("/api/v1/chat/new", payload) - data = res.get("data") or {} - msg_id = data.get("guid") or data.get("messageGuid") or "ok" - return SendResult(success=True, message_id=str(msg_id), raw_response=res) - except Exception as exc: - return SendResult(success=False, error=str(exc) or type(exc).__name__) + payload = {"addresses": [address], "message": message, "tempGuid": _temp_guid()} + return await self._post_message("/api/v1/chat/new", payload) - # ------------------------------------------------------------------ - # Text sending - # ------------------------------------------------------------------ + # --- Text sending --- @staticmethod def truncate_message(content: str, max_length: int = MAX_TEXT_LENGTH) -> List[str]: - # Use the base splitter but skip pagination indicators — iMessage - # bubbles flow naturally without "(1/3)" suffixes. + # Base splitter minus "(1/3)" pagination suffixes — iMessage bubbles flow naturally. chunks = BasePlatformAdapter.truncate_message(content, max_length) return [re.sub(r"\s*\(\d+/\d+\)$", "", c) for c in chunks] async def send( - self, - chat_id: str, - content: str, - reply_to: Optional[str] = None, + self, chat_id: str, content: str, reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: text = self.format_message(content) if not text: return SendResult(success=False, error="BlueBubbles send requires text") - # Split on paragraph breaks first (double newlines) so each thought - # becomes its own iMessage bubble, then truncate any that are still - # too long. + # Each paragraph becomes its own iMessage bubble; truncate any still too long. paragraphs = [p.strip() for p in re.split(r'\n\s*\n', text) if p.strip()] chunks: List[str] = [] for para in (paragraphs or [text]): @@ -559,102 +411,60 @@ class BlueBubblesAdapter(BasePlatformAdapter): guid = await self._resolve_chat_guid(chat_id) if not guid: # If the target looks like an address, try creating a new chat - if self._private_api_enabled and ( - "@" in chat_id or re.match(r"^\+\d+", chat_id) - ): + if self._private_api_enabled and ("@" in chat_id or re.match(r"^\+\d+", chat_id)): return await self._create_chat_for_handle(chat_id, chunk) - return SendResult( - success=False, - error=f"BlueBubbles chat not found for target: {chat_id}", - ) - payload: Dict[str, Any] = { - "chatGuid": guid, - "tempGuid": f"temp-{datetime.utcnow().timestamp()}", - "message": chunk, - } + return SendResult(success=False, error=f"BlueBubbles chat not found for target: {chat_id}") + payload: Dict[str, Any] = {"chatGuid": guid, "tempGuid": _temp_guid(), "message": chunk} if reply_to and self._private_api_enabled and self._helper_connected: payload["method"] = "private-api" payload["selectedMessageGuid"] = reply_to payload["partIndex"] = 0 - try: - res = await self._api_post("/api/v1/message/text", payload) - data = res.get("data") or {} - msg_id = data.get("guid") or data.get("messageGuid") or "ok" - last = SendResult( - success=True, message_id=str(msg_id), raw_response=res - ) - except Exception as exc: - return SendResult(success=False, error=str(exc) or type(exc).__name__) + last = await self._post_message("/api/v1/message/text", payload) + if not last.success: + return last return last - # ------------------------------------------------------------------ - # Media sending (outbound) - # ------------------------------------------------------------------ + # --- Media sending (outbound) --- async def _send_attachment( - self, - chat_id: str, - file_path: str, - filename: Optional[str] = None, - caption: Optional[str] = None, - is_audio_message: bool = False, + self, chat_id: str, file_path: str, filename: Optional[str] = None, + caption: Optional[str] = None, is_audio_message: bool = False, ) -> SendResult: """Send a file attachment via BlueBubbles multipart upload.""" if not self.client: return SendResult(success=False, error="Not connected") if not await asyncio.to_thread(os.path.isfile, file_path): return SendResult(success=False, error=f"File not found: {file_path}") - guid = await self._resolve_chat_guid(chat_id) if not guid: return SendResult(success=False, error=f"Chat not found: {chat_id}") - fname = filename or os.path.basename(file_path) try: - # httpx's async multipart iterator reads file-like objects through - # a synchronous chunk generator. Read the file off the event-loop - # thread before handing bytes to the client. + # httpx's async multipart iterator reads file objects through a sync + # chunk generator — read the bytes off the event-loop thread first. payload = await asyncio.to_thread(Path(file_path).read_bytes) files = {"attachment": (fname, payload, "application/octet-stream")} - data: Dict[str, str] = { - "chatGuid": guid, - "name": fname, - "tempGuid": uuid.uuid4().hex, - } + data: Dict[str, str] = {"chatGuid": guid, "name": fname, "tempGuid": uuid.uuid4().hex} if is_audio_message: data["isAudioMessage"] = "true" res = await self.client.post( - self._api_url("/api/v1/message/attachment"), - files=files, - data=data, - timeout=120, + self._api_url("/api/v1/message/attachment"), files=files, data=data, timeout=120 ) res.raise_for_status() result = res.json() - if caption: await self.send(chat_id, caption) - if result.get("status") == 200: rdata = result.get("data") or {} msg_id = rdata.get("guid") if isinstance(rdata, dict) else None - return SendResult( - success=True, message_id=msg_id, raw_response=result - ) - return SendResult( - success=False, - error=result.get("message", "Attachment upload failed"), - ) + return SendResult(success=True, message_id=msg_id, raw_response=result) + return SendResult(success=False, error=result.get("message", "Attachment upload failed")) except Exception as e: return SendResult(success=False, error=str(e)) 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, + self, chat_id: str, image_url: str, caption: Optional[str] = None, + reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: try: from gateway.platforms.base import cache_image_from_url @@ -664,145 +474,51 @@ class BlueBubblesAdapter(BasePlatformAdapter): except Exception: return await super().send_image(chat_id, image_url, caption, 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: + async def send_image_file(self, chat_id, image_path, caption=None, reply_to=None, **kw) -> SendResult: return await self._send_attachment(chat_id, image_path, caption=caption) - async def send_voice( - self, - chat_id: str, - audio_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, - ) -> SendResult: - return await self._send_attachment( - chat_id, audio_path, caption=caption, is_audio_message=True - ) + async def send_voice(self, chat_id, audio_path, caption=None, reply_to=None, **kw) -> SendResult: + return await self._send_attachment(chat_id, audio_path, caption=caption, is_audio_message=True) - async def send_video( - self, - chat_id: str, - video_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, - ) -> SendResult: + async def send_video(self, chat_id, video_path, caption=None, reply_to=None, **kw) -> SendResult: return await self._send_attachment(chat_id, video_path, caption=caption) 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, + self, chat_id, file_path, caption=None, file_name=None, reply_to=None, **kw ) -> SendResult: - return await self._send_attachment( - chat_id, file_path, filename=file_name, caption=caption - ) + return await self._send_attachment(chat_id, file_path, filename=file_name, caption=caption) async def send_animation( - self, - chat_id: str, - animation_url: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, + self, chat_id, animation_url, caption=None, reply_to=None, metadata=None ) -> SendResult: - return await self.send_image( - chat_id, animation_url, caption, reply_to, metadata - ) + return await self.send_image(chat_id, animation_url, caption, reply_to, metadata) - # ------------------------------------------------------------------ - # Typing indicators - # ------------------------------------------------------------------ + # --- Typing indicators / read receipts (private API only) --- async def send_typing(self, chat_id: str, metadata=None) -> None: - if not self._private_api_enabled or not self._helper_connected or not self.client: - return - try: - guid = await self._resolve_chat_guid(chat_id) - if guid: - encoded = quote(guid, safe="") - await self.client.post( - self._api_url(f"/api/v1/chat/{encoded}/typing"), timeout=5 - ) - except Exception: - pass + await self._private_api_chat_call(chat_id, "typing", "post") async def stop_typing(self, chat_id: str) -> None: - if not self._private_api_enabled or not self._helper_connected or not self.client: - return - try: - guid = await self._resolve_chat_guid(chat_id) - if guid: - encoded = quote(guid, safe="") - await self.client.delete( - self._api_url(f"/api/v1/chat/{encoded}/typing"), timeout=5 - ) - except Exception: - pass - - # ------------------------------------------------------------------ - # Read receipts - # ------------------------------------------------------------------ + await self._private_api_chat_call(chat_id, "typing", "delete") async def mark_read(self, chat_id: str) -> bool: - if not self._private_api_enabled or not self._helper_connected or not self.client: - return False - try: - guid = await self._resolve_chat_guid(chat_id) - if guid: - encoded = quote(guid, safe="") - await self.client.post( - self._api_url(f"/api/v1/chat/{encoded}/read"), timeout=5 - ) - return True - except Exception: - pass - return False + return await self._private_api_chat_call(chat_id, "read", "post") - # ------------------------------------------------------------------ - # Tapback reactions - # ------------------------------------------------------------------ - - # ------------------------------------------------------------------ - # Chat info - # ------------------------------------------------------------------ + # --- Chat info --- async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: is_group = ";+;" in (chat_id or "") - info: Dict[str, Any] = { - "name": chat_id, - "type": "group" if is_group else "dm", - } + info: Dict[str, Any] = {"name": chat_id, "type": "group" if is_group else "dm"} try: guid = await self._resolve_chat_guid(chat_id) if guid: - encoded = quote(guid, safe="") - res = await self._api_get( - f"/api/v1/chat/{encoded}?with=participants" - ) + res = await self._api_get(f"/api/v1/chat/{quote(guid, safe='')}?with=participants") data = (res or {}).get("data", {}) - display_name = ( - data.get("displayName") - or data.get("chatIdentifier") - or chat_id - ) - participants = [] - for p in data.get("participants", []) or []: - addr = (p.get("address") or "").strip() - if addr: - participants.append(addr) - info["name"] = display_name + info["name"] = data.get("displayName") or data.get("chatIdentifier") or chat_id + participants = [ + addr for p in data.get("participants", []) or [] + if (addr := (p.get("address") or "").strip()) + ] if participants: info["participants"] = participants except Exception: @@ -812,75 +528,36 @@ class BlueBubblesAdapter(BasePlatformAdapter): def format_message(self, content: str) -> str: return strip_markdown(content) - # ------------------------------------------------------------------ - # Inbound attachment downloading (from #4588) - # ------------------------------------------------------------------ + # --- Inbound attachment downloading --- - async def _download_attachment( - self, att_guid: str, att_meta: Dict[str, Any] - ) -> Optional[str]: - """Download an attachment from BlueBubbles and cache it locally. - - Returns the local file path on success, None on failure. - """ + async def _download_attachment(self, att_guid: str, att_meta: Dict[str, Any]) -> Optional[str]: + """Download an attachment and cache it locally; local path or None on failure.""" if not self.client: return None try: - encoded = quote(att_guid, safe="") resp = await self.client.get( - self._api_url(f"/api/v1/attachment/{encoded}/download"), - timeout=60, - follow_redirects=True, + self._api_url(f"/api/v1/attachment/{quote(att_guid, safe='')}/download"), + timeout=60, follow_redirects=True, ) resp.raise_for_status() data = resp.content - mime = (att_meta.get("mimeType") or "").lower() - transfer_name = att_meta.get("transferName", "") - if mime.startswith("image/"): - ext = ext_for_mime( - mime, - overrides=_BLUEBUBBLES_IMAGE_EXT_OVERRIDES, - # Historical map was closed: any unlisted image mime - # fell back to .jpg without consulting mimetypes. - use_defaults=False, - use_mimetypes=False, - fallback=".jpg", - ) or ".jpg" + ext = _closed_ext(mime, _BLUEBUBBLES_IMAGE_EXT_OVERRIDES, ".jpg") return cache_image_from_bytes(data, ext) - if mime.startswith("audio/"): - ext = ext_for_mime( - mime, - overrides=_BLUEBUBBLES_AUDIO_EXT_OVERRIDES, - # Historical map was closed: any unlisted audio mime - # fell back to .mp3 without consulting mimetypes. - use_defaults=False, - use_mimetypes=False, - fallback=".mp3", - ) or ".mp3" + ext = _closed_ext(mime, _BLUEBUBBLES_AUDIO_EXT_OVERRIDES, ".mp3") return cache_audio_from_bytes(data, ext) - # Videos, documents, and everything else - filename = transfer_name or f"file_{uuid.uuid4().hex[:8]}" + filename = att_meta.get("transferName", "") or f"file_{uuid.uuid4().hex[:8]}" return cache_document_from_bytes(data, filename) - except Exception as exc: - logger.warning( - "[bluebubbles] failed to download attachment %s: %s", - _redact(att_guid), - exc, - ) + logger.warning("[bluebubbles] failed to download attachment %s: %s", _redact(att_guid), exc) return None - # ------------------------------------------------------------------ - # Webhook handling - # ------------------------------------------------------------------ + # --- Webhook handling --- - def _extract_payload_record( - self, payload: Dict[str, Any] - ) -> Optional[Dict[str, Any]]: + def _extract_payload_record(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]: data = payload.get("data") if isinstance(data, dict): return data @@ -899,6 +576,45 @@ class BlueBubblesAdapter(BasePlatformAdapter): return candidate.strip() return None + @staticmethod + def _parse_webhook_body(raw: bytes) -> Any: + """Decode a webhook body: JSON, else form-encoded with a JSON field.""" + body = raw.decode("utf-8", errors="replace") + try: + return json.loads(body) + except Exception: + form = parse_qs(body) + payload_str = (form.get("payload") or form.get("data") or form.get("message") or [""])[0] + return json.loads(payload_str) if payload_str else {} + + async def _collect_attachments(self, record: Dict[str, Any]): + """Download inbound attachments; returns (media_urls, media_types, msg_type).""" + media_urls: List[str] = [] + media_types: List[str] = [] + msg_type = MessageType.TEXT + for att in record.get("attachments") or []: + att_guid = att.get("guid", "") + if not att_guid: + continue + cached = await self._download_attachment(att_guid, att) + if not cached: + continue + mime = (att.get("mimeType") or "").lower() + media_urls.append(cached) + media_types.append(mime) + if mime.startswith("image/"): + msg_type = MessageType.PHOTO + elif mime.startswith("audio/") or (att.get("uti") or "").endswith("caf"): + msg_type = MessageType.VOICE + elif mime.startswith("video/"): + msg_type = MessageType.VIDEO + else: + msg_type = MessageType.DOCUMENT + # With multiple attachments, prefer PHOTO if any images present + if len(media_urls) > 1 and "image" in {(m or "").split("/")[0] for m in media_types}: + msg_type = MessageType.PHOTO + return media_urls, media_types, msg_type + async def _handle_webhook(self, request): from aiohttp import web @@ -912,21 +628,7 @@ class BlueBubblesAdapter(BasePlatformAdapter): if token != self.password: return web.json_response({"error": "unauthorized"}, status=401) try: - raw = await request.read() - body = raw.decode("utf-8", errors="replace") - try: - payload = json.loads(body) - except Exception: - from urllib.parse import parse_qs - - form = parse_qs(body) - payload_str = ( - form.get("payload") - or form.get("data") - or form.get("message") - or [""] - )[0] - payload = json.loads(payload_str) if payload_str else {} + payload = self._parse_webhook_body(await request.read()) except Exception as exc: logger.error("[bluebubbles] webhook parse error: %s", exc) return web.json_response({"error": "invalid payload"}, status=400) @@ -937,92 +639,36 @@ class BlueBubblesAdapter(BasePlatformAdapter): return web.Response(text="ok") record = self._extract_payload_record(payload) or {} - is_from_me = bool( - record.get("isFromMe") - or record.get("fromMe") - or record.get("is_from_me") - ) - if is_from_me: + if record.get("isFromMe") or record.get("fromMe") or record.get("is_from_me"): return web.Response(text="ok") - # Skip tapback reactions delivered as messages assoc_type = record.get("associatedMessageType") - if isinstance(assoc_type, int) and assoc_type in { - **_TAPBACK_ADDED, - **_TAPBACK_REMOVED, - }: + if isinstance(assoc_type, int) and assoc_type in _TAPBACK_CODES: return web.Response(text="ok") - text = ( - self._value( - record.get("text"), record.get("message"), record.get("body") - ) - or "" - ) - - # --- Inbound attachment handling --- - attachments = record.get("attachments") or [] - media_urls: List[str] = [] - media_types: List[str] = [] - msg_type = MessageType.TEXT - - for att in attachments: - att_guid = att.get("guid", "") - if not att_guid: - continue - cached = await self._download_attachment(att_guid, att) - if cached: - mime = (att.get("mimeType") or "").lower() - media_urls.append(cached) - media_types.append(mime) - if mime.startswith("image/"): - msg_type = MessageType.PHOTO - elif mime.startswith("audio/") or (att.get("uti") or "").endswith( - "caf" - ): - msg_type = MessageType.VOICE - elif mime.startswith("video/"): - msg_type = MessageType.VIDEO - else: - msg_type = MessageType.DOCUMENT - - # With multiple attachments, prefer PHOTO if any images present - if len(media_urls) > 1: - mime_prefixes = {(m or "").split("/")[0] for m in media_types} - if "image" in mime_prefixes: - msg_type = MessageType.PHOTO - + text = self._value(record.get("text"), record.get("message"), record.get("body")) or "" + media_urls, media_types, msg_type = await self._collect_attachments(record) if not text and media_urls: text = "(attachment)" - # --- End attachment handling --- chat_guid = self._value( - record.get("chatGuid"), - payload.get("chatGuid"), - record.get("chat_guid"), - payload.get("chat_guid"), - payload.get("guid"), + record.get("chatGuid"), payload.get("chatGuid"), + record.get("chat_guid"), payload.get("chat_guid"), payload.get("guid"), ) - # Fallback: BlueBubbles v1.9+ webhook payloads omit top-level chatGuid; - # the chat GUID is nested under data.chats[0].guid instead. + # BlueBubbles v1.9+ payloads omit top-level chatGuid; it's nested under data.chats[0].guid. if not chat_guid: _chats = record.get("chats") or [] if _chats and isinstance(_chats[0], dict): chat_guid = _chats[0].get("guid") or _chats[0].get("chatGuid") chat_identifier = self._value( - record.get("chatIdentifier"), - record.get("identifier"), - payload.get("chatIdentifier"), - payload.get("identifier"), + record.get("chatIdentifier"), record.get("identifier"), + payload.get("chatIdentifier"), payload.get("identifier"), ) + handle = record.get("handle") sender = ( self._value( - record.get("handle", {}).get("address") - if isinstance(record.get("handle"), dict) - else None, - record.get("sender"), - record.get("from"), - record.get("address"), + handle.get("address") if isinstance(handle, dict) else None, + record.get("sender"), record.get("from"), record.get("address"), ) or chat_identifier or chat_guid @@ -1042,36 +688,22 @@ class BlueBubblesAdapter(BasePlatformAdapter): return web.Response(text="ok") text = self._clean_mention_text(text) source = self.build_source( - chat_id=session_chat_id, - chat_name=chat_identifier or sender, - chat_type="group" if is_group else "dm", - user_id=sender, - user_name=sender, + chat_id=session_chat_id, chat_name=chat_identifier or sender, + chat_type="group" if is_group else "dm", user_id=sender, user_name=sender, chat_id_alt=chat_identifier, ) event = MessageEvent( - text=text, - message_type=msg_type, - source=source, - raw_message=payload, - message_id=self._value( - record.get("guid"), - record.get("messageGuid"), - record.get("id"), - ), + text=text, message_type=msg_type, source=source, raw_message=payload, + message_id=self._value(record.get("guid"), record.get("messageGuid"), record.get("id")), reply_to_message_id=self._value( - record.get("threadOriginatorGuid"), - record.get("associatedMessageGuid"), + record.get("threadOriginatorGuid"), record.get("associatedMessageGuid") ), - media_urls=media_urls, - media_types=media_types, + media_urls=media_urls, media_types=media_types, ) task = asyncio.create_task(self.handle_message(event)) self._background_tasks.add(task) task.add_done_callback(self._background_tasks.discard) - # Fire-and-forget read receipt if self.send_read_receipts and session_chat_id: asyncio.create_task(self.mark_read(session_chat_id)) - return web.Response(text="ok") diff --git a/gateway/platforms/msgraph_webhook.py b/gateway/platforms/msgraph_webhook.py index 6675fadf9c..eb75a5c7af 100644 --- a/gateway/platforms/msgraph_webhook.py +++ b/gateway/platforms/msgraph_webhook.py @@ -7,6 +7,7 @@ import hmac import ipaddress import json import logging +import re from collections import deque from hashlib import sha1 from typing import Any, Awaitable, Callable, Dict, Optional @@ -21,27 +22,22 @@ except ImportError: from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( - BasePlatformAdapter, - MessageEvent, - MessageType, - SendResult, - is_network_accessible, + BasePlatformAdapter, MessageEvent, MessageType, SendResult, is_network_accessible, ) 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 extra.host. The all-interfaces default still -# requires extra.allowed_source_cidrs (see _source_allowlist_required_but_missing). +# ``None`` → aiohttp binds one socket per address family (IPv4 + IPv6); the old +# "0.0.0.0" default was unreachable over IPv6-only private networks. Pin a host +# via extra.host. The all-interfaces default still requires +# extra.allowed_source_cidrs (see _source_allowlist_required_but_missing). DEFAULT_HOST = None DEFAULT_PORT = 8646 DEFAULT_WEBHOOK_PATH = "/msgraph/webhook" DEFAULT_MAX_SEEN_RECEIPTS = 5000 DEFAULT_MAX_BODY_BYTES = 1_048_576 NotificationScheduler = Callable[[Dict[str, Any], MessageEvent], Awaitable[None] | None] +_TEMPLATE_KEY_RE = re.compile(r"\{([a-zA-Z0-9_.]+)\}") def check_msgraph_webhook_requirements() -> bool: @@ -49,6 +45,63 @@ def check_msgraph_webhook_requirements() -> bool: return AIOHTTP_AVAILABLE +def _string_or_none(value: Any) -> Optional[str]: + if value is None: + return None + return str(value).strip() or None + + +def _normalize_path(path: Any) -> str: + raw = str(path or "").strip() or "/" + return raw if raw.startswith("/") else f"/{raw}" + + +def _parse_allowed_source_cidrs(raw: Any) -> list[ipaddress._BaseNetwork]: + """Parse the optional CIDR allowlist; empty/missing means "allow everything". + + When populated, requests from source IPs outside every listed CIDR are + rejected with 403 before the body is parsed (restrict to Microsoft + Graph's published webhook source ranges in production). + """ + if isinstance(raw, str): + candidates = raw.split(",") + elif isinstance(raw, (list, tuple, set)): + candidates = [str(chunk) for chunk in raw] + else: + return [] + networks: list[ipaddress._BaseNetwork] = [] + for chunk in candidates: + chunk = chunk.strip() + if not chunk: + continue + try: + networks.append(ipaddress.ip_network(chunk, strict=False)) + except ValueError: + logger.warning("[msgraph_webhook] Ignoring invalid allowed_source_cidrs entry: %r", chunk) + return networks + + +def _prefix_match(resource: str, prefix: str) -> bool: + return resource == prefix or resource.startswith(f"{prefix}/") + + +def _render_template(template: str, payload: Dict[str, Any]) -> str: + """Substitute ``{dotted.key}`` placeholders from *payload*; unknown keys stay literal.""" + + def _resolve(match: re.Match[str]) -> str: + key = match.group(1) + value: Any = payload + for part in key.split("."): + if not isinstance(value, dict): + return f"{{{key}}}" + value = value.get(part, f"{{{key}}}") + if isinstance(value, (dict, list)): + return json.dumps(value, sort_keys=True)[:2000] + return str(value) + + return _TEMPLATE_KEY_RE.sub(_resolve, template) + + class MSGraphWebhookAdapter(BasePlatformAdapter): """Receive Microsoft Graph change notifications and surface them internally.""" @@ -59,25 +112,15 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): _raw_host = extra.get("host", DEFAULT_HOST) or DEFAULT_HOST self._host: Optional[str] = str(_raw_host) if _raw_host else None self._port: int = int(extra.get("port", DEFAULT_PORT)) - self._webhook_path: str = self._normalize_path( - extra.get("webhook_path", DEFAULT_WEBHOOK_PATH) - ) - self._health_path: str = self._normalize_path(extra.get("health_path", "/health")) + self._webhook_path: str = _normalize_path(extra.get("webhook_path", DEFAULT_WEBHOOK_PATH)) + self._health_path: str = _normalize_path(extra.get("health_path", "/health")) self._accepted_resources: list[str] = [ - str(value).strip() - for value in (extra.get("accepted_resources") or []) - if str(value).strip() + str(value).strip() for value in (extra.get("accepted_resources") or []) if str(value).strip() ] - self._client_state: Optional[str] = self._string_or_none(extra.get("client_state")) - self._max_seen_receipts = max( - 1, int(extra.get("max_seen_receipts", DEFAULT_MAX_SEEN_RECEIPTS)) - ) - self._max_body_bytes = max( - 1, int(extra.get("max_body_bytes", DEFAULT_MAX_BODY_BYTES)) - ) - self._allowed_source_networks: list[ipaddress._BaseNetwork] = ( - self._parse_allowed_source_cidrs(extra.get("allowed_source_cidrs")) - ) + self._client_state: Optional[str] = _string_or_none(extra.get("client_state")) + self._max_seen_receipts = max(1, int(extra.get("max_seen_receipts", DEFAULT_MAX_SEEN_RECEIPTS))) + self._max_body_bytes = max(1, int(extra.get("max_body_bytes", DEFAULT_MAX_BODY_BYTES))) + self._allowed_source_networks = _parse_allowed_source_cidrs(extra.get("allowed_source_cidrs")) self._runner = None self._notification_scheduler: Optional[NotificationScheduler] = None self._seen_receipts: set[str] = set() @@ -85,63 +128,6 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): self._accepted_count = 0 self._duplicate_count = 0 - @staticmethod - def _string_or_none(value: Any) -> Optional[str]: - if value is None: - return None - text = str(value).strip() - return text or None - - @staticmethod - def _normalize_path(path: Any) -> str: - raw = str(path or "").strip() or "/" - return raw if raw.startswith("/") else f"/{raw}" - - @staticmethod - def _build_receipt_key(notification: Dict[str, Any]) -> Optional[str]: - explicit_id = str(notification.get("id") or "").strip() - if explicit_id: - return f"id:{explicit_id}" - return None - - @staticmethod - def _normalize_resource_value(resource: str) -> str: - return str(resource or "").strip().strip("/") - - @staticmethod - def _parse_allowed_source_cidrs( - raw: Any, - ) -> list[ipaddress._BaseNetwork]: - """Parse an optional list of CIDR ranges allowed to POST to the webhook. - - An empty or missing value means "allow everything" (same behavior as - before this field existed). When populated, requests from source IPs - outside every listed CIDR are rejected with 403 before the body is - parsed. Use this to restrict the endpoint to Microsoft Graph's - published webhook source ranges in production deployments. - """ - if raw is None: - return [] - if isinstance(raw, str): - candidates = [chunk.strip() for chunk in raw.split(",")] - elif isinstance(raw, (list, tuple, set)): - candidates = [str(chunk).strip() for chunk in raw] - else: - return [] - - networks: list[ipaddress._BaseNetwork] = [] - for chunk in candidates: - if not chunk: - continue - try: - networks.append(ipaddress.ip_network(chunk, strict=False)) - except ValueError: - logger.warning( - "[msgraph_webhook] Ignoring invalid allowed_source_cidrs entry: %r", - chunk, - ) - return networks - def set_notification_scheduler(self, scheduler: Optional[NotificationScheduler]) -> None: self._notification_scheduler = scheduler @@ -152,9 +138,7 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): async def connect(self, *, is_reconnect: bool = False) -> bool: if self._client_state is None: - logger.error( - "[msgraph_webhook] Refusing to start without extra.client_state configured" - ) + logger.error("[msgraph_webhook] Refusing to start without extra.client_state configured") return False if self._source_allowlist_required_but_missing(): logger.error( @@ -170,22 +154,14 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): app.router.add_get(self._health_path, self._handle_health) app.router.add_get(self._webhook_path, self._handle_validation) app.router.add_post(self._webhook_path, self._handle_notification) - - # Plugin-registered native handlers (aiohttp web.Application — - # router routes). Wired before AppRunner.setup() freezes the router. + # Plugin-registered native routes; wired before AppRunner.setup() freezes the router. self._wire_plugin_handlers(app) - self._runner = web.AppRunner(app) await self._runner.setup() site = web.TCPSite(self._runner, self._host, self._port) await site.start() self._mark_connected() - logger.info( - "[msgraph_webhook] Listening on %s:%d%s", - self._host, - self._port, - self._webhook_path, - ) + logger.info("[msgraph_webhook] Listening on %s:%d%s", self._host, self._port, self._webhook_path) return True async def disconnect(self) -> None: @@ -195,10 +171,7 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): self._mark_disconnected() async def send( - self, - chat_id: str, - content: str, - reply_to: Optional[str] = None, + self, chat_id: str, content: str, reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: logger.info("[msgraph_webhook] Response for %s: %s", chat_id, content[:200]) @@ -210,25 +183,17 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): async def _handle_health(self, request: "web.Request") -> "web.Response": if not self._source_ip_allowed(request): return web.Response(status=403) - return web.json_response( - { - "status": "ok", - "platform": self.platform.value, - "webhook_path": self._webhook_path, - "accepted": self._accepted_count, - "duplicates": self._duplicate_count, - } - ) + return web.json_response({ + "status": "ok", + "platform": self.platform.value, + "webhook_path": self._webhook_path, + "accepted": self._accepted_count, + "duplicates": self._duplicate_count, + }) async def _handle_validation(self, request: "web.Request") -> "web.Response": - """Handle Microsoft Graph subscription validation handshake. - - Graph validates a subscription endpoint by sending a GET with - ``validationToken`` in the query string; the service must echo the - token verbatim as ``text/plain`` within 10 seconds. Anything else - (bare GET, GET without the token) is rejected so the endpoint can't - be enumerated or mistakenly used for data exfiltration. - """ + """Graph subscription validation handshake: echo ``validationToken`` verbatim + as text/plain. Bare GETs are rejected so the endpoint can't be enumerated.""" if not self._source_ip_allowed(request): return web.Response(status=403) validation_token = request.query.get("validationToken", "") @@ -239,43 +204,16 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): async def _handle_notification(self, request: "web.Request") -> "web.Response": if not self._source_ip_allowed(request): return web.Response(status=403) - - # Graph never sends validationToken on POST, but tolerate it for - # defensive clients that replay the handshake in-band. + # Graph never sends validationToken on POST, but tolerate clients replaying it in-band. validation_token = request.query.get("validationToken", "") if validation_token: return web.Response(text=validation_token, content_type="text/plain") - try: - content_length = request.content_length - except Exception: - content_length = None - if content_length is not None and content_length > self._max_body_bytes: - return web.Response(status=413) - - try: - raw_body = await request.read() - except Exception: - return web.Response(status=400) - if len(raw_body) > self._max_body_bytes: - return web.Response(status=413) - - try: - body = json.loads(raw_body.decode("utf-8")) - except (json.JSONDecodeError, UnicodeDecodeError): - return web.Response(status=400) - if not isinstance(body, dict): - return web.Response(status=400) - - notifications = body.get("value") - if not isinstance(notifications, list): - return web.Response(status=400) - - accepted = 0 - duplicates = 0 - auth_rejected = 0 - other_rejected = 0 + status, notifications = await self._read_notifications(request) + if status: + return web.Response(status=status) + accepted = duplicates = auth_rejected = other_rejected = 0 for raw_notification in notifications: if not isinstance(raw_notification, dict): other_rejected += 1 @@ -285,54 +223,64 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): other_rejected += 1 continue if not self._verify_client_state(notification): - # Treat bad clientState as an auth failure: if the whole - # batch is forged, we want to signal 403 so the sender - # stops retrying. Legitimate Graph retries have valid - # clientState and hit the accepted/duplicate paths. + # Bad clientState is an auth failure: a fully forged batch gets 403 + # so the sender stops retrying; legitimate Graph retries carry a + # valid clientState and hit the accepted/duplicate paths. auth_rejected += 1 continue - - receipt_key = self._build_receipt_key(notification) + explicit_id = str(notification.get("id") or "").strip() + receipt_key = f"id:{explicit_id}" if explicit_id else None if receipt_key is not None: - if self._has_seen_receipt(receipt_key): + if receipt_key in self._seen_receipts: duplicates += 1 continue self._remember_receipt(receipt_key) - accepted += 1 self._accepted_count += 1 - event = self._build_message_event(notification, receipt_key) - self._schedule_notification(notification, event) + self._schedule_notification(notification, self._build_message_event(notification, receipt_key)) self._duplicate_count += duplicates - # If anything ingested OR deduped, return 202 with empty body so - # Graph acks successfully and we don't leak internal counters. If - # every item failed auth, return 403 so an attacker POSTing fake - # notifications gets a clear reject. Other failures (malformed, - # resource-not-accepted) are the sender's configuration problem, - # so 400. + # Anything ingested OR deduped → 202 with empty body (Graph acks; no + # counter leak). Every item failed auth → 403 so forged POSTs get a + # clear reject. Otherwise (malformed / resource not accepted) → 400. if accepted or duplicates: return web.Response(status=202) if auth_rejected and not other_rejected: return web.Response(status=403) return web.Response(status=400) - def _source_ip_allowed(self, request: "web.Request") -> bool: - """Return True if the request's source IP is in the configured allowlist. + async def _read_notifications(self, request: "web.Request") -> tuple[int, list]: + """Read and validate the POST body; returns (error_status, []) or (0, notifications).""" + try: + content_length = request.content_length + except Exception: + content_length = None + if content_length is not None and content_length > self._max_body_bytes: + return 413, [] + try: + raw_body = await request.read() + except Exception: + return 400, [] + if len(raw_body) > self._max_body_bytes: + return 413, [] + try: + body = json.loads(raw_body.decode("utf-8")) + except (json.JSONDecodeError, UnicodeDecodeError): + return 400, [] + notifications = body.get("value") if isinstance(body, dict) else None + if not isinstance(notifications, list): + return 400, [] + return 0, notifications - Loopback-only binds may omit ``allowed_source_cidrs`` for local reverse - proxies and dev tunnels. Network-accessible binds fail closed until an - explicit CIDR allowlist is configured. - """ + def _source_ip_allowed(self, request: "web.Request") -> bool: + """Loopback-only binds may omit ``allowed_source_cidrs`` (local proxies, + dev tunnels); network-accessible binds fail closed without one.""" if self._source_allowlist_required_but_missing(): return False if not self._allowed_source_networks: return True - peer = request.remote or "" - if not peer: - return False try: - peer_addr = ipaddress.ip_address(peer) + peer_addr = ipaddress.ip_address(request.remote or "") except ValueError: return False return any(peer_addr in network for network in self._allowed_source_networks) @@ -340,72 +288,47 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): def _resource_accepted(self, resource: str) -> bool: if not self._accepted_resources: return True - normalized_resource = self._normalize_resource_value(resource) + resource = resource.strip().strip("/") for pattern in self._accepted_resources: - normalized_pattern = self._normalize_resource_value(pattern) - if not normalized_pattern: + pattern = pattern.strip().strip("/") + if not pattern: continue - if normalized_pattern.endswith("*"): - prefix = normalized_pattern[:-1].rstrip("/") - if normalized_resource == prefix or normalized_resource.startswith(f"{prefix}/"): + if pattern.endswith("*"): + if _prefix_match(resource, pattern[:-1].rstrip("/")): return True - continue - if ( - normalized_resource == normalized_pattern - or normalized_resource.startswith(f"{normalized_pattern}/") - ): + elif _prefix_match(resource, pattern): return True return False def _verify_client_state(self, notification: Dict[str, Any]) -> bool: - """Verify the Graph-supplied clientState matches the configured secret. - - Uses ``hmac.compare_digest`` instead of ``==`` so that a mismatch - doesn't leak how many leading characters matched via string-compare - timing. The configured client_state is a shared secret (documented in - the setup guide as "generate with ``openssl rand -hex 32``"), so a - timing-safe compare is the right primitive. - """ + """Timing-safe compare of the Graph-supplied clientState against the + configured shared secret (``openssl rand -hex 32`` in the setup guide).""" expected = self._client_state if expected is None: return False - provided = self._string_or_none(notification.get("clientState")) + provided = _string_or_none(notification.get("clientState")) if provided is None: return False - # Compare as bytes: ``compare_digest`` raises TypeError on a str with - # non-ASCII characters, and clientState comes from the request body. + # Compare as bytes: compare_digest raises TypeError on non-ASCII str, + # and clientState comes from the request body. return hmac.compare_digest(provided.encode(), expected.encode()) - def _has_seen_receipt(self, receipt_key: str) -> bool: - return receipt_key in self._seen_receipts - def _remember_receipt(self, receipt_key: str) -> None: self._seen_receipts.add(receipt_key) self._seen_receipt_order.append(receipt_key) while len(self._seen_receipt_order) > self._max_seen_receipts: - oldest = self._seen_receipt_order.popleft() - self._seen_receipts.discard(oldest) + self._seen_receipts.discard(self._seen_receipt_order.popleft()) - def _build_message_event( - self, - notification: Dict[str, Any], - receipt_key: Optional[str], - ) -> MessageEvent: + def _build_message_event(self, notification: Dict[str, Any], receipt_key: Optional[str]) -> MessageEvent: message_id = receipt_key or f"sha1:{sha1(json.dumps(notification, sort_keys=True).encode('utf-8')).hexdigest()}" source = self.build_source( chat_id=f"msgraph:{notification.get('subscriptionId', 'unknown')}", - chat_name="msgraph/webhook", - chat_type="webhook", - user_id="msgraph", - user_name="Microsoft Graph", + chat_name="msgraph/webhook", chat_type="webhook", + user_id="msgraph", user_name="Microsoft Graph", ) return MessageEvent( - text=self._render_prompt(notification), - message_type=MessageType.TEXT, - source=source, - raw_message=notification, - message_id=message_id, - internal=True, + text=self._render_prompt(notification), message_type=MessageType.TEXT, source=source, + raw_message=notification, message_id=message_id, internal=True, ) def _render_prompt(self, notification: Dict[str, Any]) -> str: @@ -417,41 +340,18 @@ class MSGraphWebhookAdapter(BasePlatformAdapter): "change_type": notification.get("changeType", ""), "subscription_id": notification.get("subscriptionId", ""), } - return self._render_template(template, payload) + return _render_template(template, payload) rendered = json.dumps(notification, indent=2, sort_keys=True)[:4000] return f"Microsoft Graph change notification:\n\n```json\n{rendered}\n```" - def _render_template(self, template: str, payload: Dict[str, Any]) -> str: - import re - - def _resolve(match: "re.Match[str]") -> str: - key = match.group(1) - value: Any = payload - for part in key.split("."): - if isinstance(value, dict): - value = value.get(part, f"{{{key}}}") - else: - return f"{{{key}}}" - if isinstance(value, (dict, list)): - return json.dumps(value, sort_keys=True)[:2000] - return str(value) - - return re.sub(r"\{([a-zA-Z0-9_.]+)\}", _resolve, template) - - def _schedule_notification( - self, - notification: Dict[str, Any], - event: MessageEvent, - ) -> None: + def _schedule_notification(self, notification: Dict[str, Any], event: MessageEvent) -> None: scheduler = self._notification_scheduler - if scheduler is not None: - result = scheduler(notification, event) - if asyncio.iscoroutine(result): - task = asyncio.create_task(result) - self._background_tasks.add(task) - task.add_done_callback(self._background_tasks.discard) - return - - task = asyncio.create_task(self.handle_message(event)) + if scheduler is None: + coro = self.handle_message(event) + else: + coro = scheduler(notification, event) + if not asyncio.iscoroutine(coro): + return + task = asyncio.create_task(coro) self._background_tasks.add(task) task.add_done_callback(self._background_tasks.discard) diff --git a/gateway/platforms/qqbot/__init__.py b/gateway/platforms/qqbot/__init__.py index d755ec48df..395ba630f1 100644 --- a/gateway/platforms/qqbot/__init__.py +++ b/gateway/platforms/qqbot/__init__.py @@ -1,20 +1,13 @@ -""" -QQBot platform package. +"""QQBot platform package. -Re-exports the main adapter symbols from ``adapter.py`` (the original -``qqbot.py``) so that **all existing import paths remain unchanged**:: +Re-exports the adapter symbols from ``adapter.py`` (the original ``qqbot.py``) +so all existing import paths remain unchanged, e.g. +``from gateway.platforms.qqbot import QQAdapter, check_qq_requirements``. - from gateway.platforms.qqbot import QQAdapter # works - from gateway.platforms.qqbot import check_qq_requirements # works - -New modules: - - ``constants`` — shared constants (API URLs, timeouts, message types) - - ``utils`` — User-Agent builder, config helpers - - ``crypto`` — AES-256-GCM key generation and decryption - - ``onboard`` — QR-code scan-to-configure flow +Sub-modules: ``constants``, ``utils`` (User-Agent, config helpers), ``crypto`` +(AES-256-GCM), ``onboard`` (QR scan-to-configure), ``chunked_upload``, ``keyboards``. """ -# -- Adapter (original qqbot.py) ------------------------------------------ from .adapter import ( # noqa: F401 QQAdapter, QQCloseError, @@ -22,29 +15,16 @@ from .adapter import ( # noqa: F401 _coerce_list, _ssrf_redirect_guard, ) - -# -- Onboard (QR-code scan-to-configure) ----------------------------------- -from .onboard import ( # noqa: F401 - BindStatus, - build_connect_url, - qr_register, -) +from .onboard import BindStatus, build_connect_url, qr_register # noqa: F401 from .crypto import decrypt_secret, generate_bind_key # noqa: F401 - -# -- Utils ----------------------------------------------------------------- from .utils import build_user_agent, get_api_headers, coerce_list # noqa: F401 - -# -- Chunked upload -------------------------------------------------------- from .chunked_upload import ( # noqa: F401 ChunkedUploader, UploadDailyLimitExceededError, UploadFileTooLargeError, ) - -# -- Inline keyboards ------------------------------------------------------ from .keyboards import ( # noqa: F401 ApprovalRequest, - ApprovalSender, InlineKeyboard, InteractionEvent, build_approval_keyboard, @@ -56,36 +36,12 @@ from .keyboards import ( # noqa: F401 ) __all__ = [ - # adapter - "QQAdapter", - "QQCloseError", - "check_qq_requirements", - "_coerce_list", - "_ssrf_redirect_guard", - # onboard - "BindStatus", - "build_connect_url", - "qr_register", - # crypto - "decrypt_secret", - "generate_bind_key", - # utils - "build_user_agent", - "get_api_headers", - "coerce_list", - # chunked upload - "ChunkedUploader", - "UploadDailyLimitExceededError", - "UploadFileTooLargeError", - # keyboards - "ApprovalRequest", - "ApprovalSender", - "InlineKeyboard", - "InteractionEvent", - "build_approval_keyboard", - "build_approval_text", - "build_update_prompt_keyboard", - "parse_approval_button_data", - "parse_interaction_event", - "parse_update_prompt_button_data", + "QQAdapter", "QQCloseError", "check_qq_requirements", "_coerce_list", "_ssrf_redirect_guard", + "BindStatus", "build_connect_url", "qr_register", + "decrypt_secret", "generate_bind_key", + "build_user_agent", "get_api_headers", "coerce_list", + "ChunkedUploader", "UploadDailyLimitExceededError", "UploadFileTooLargeError", + "ApprovalRequest", "InlineKeyboard", "InteractionEvent", + "build_approval_keyboard", "build_approval_text", "build_update_prompt_keyboard", + "parse_approval_button_data", "parse_interaction_event", "parse_update_prompt_button_data", ] diff --git a/gateway/platforms/qqbot/adapter.py b/gateway/platforms/qqbot/adapter.py index b8a9470817..a78a97e96e 100644 --- a/gateway/platforms/qqbot/adapter.py +++ b/gateway/platforms/qqbot/adapter.py @@ -32,11 +32,10 @@ Reference: https://bot.q.qq.com/wiki/develop/api-v2/ from __future__ import annotations import asyncio -import base64 import json import logging -import mimetypes import os +import re import time import uuid from datetime import datetime, timezone @@ -78,10 +77,7 @@ logger = logging.getLogger(__name__) class QQCloseError(Exception): - """Raised when QQ WebSocket closes with a specific code. - - Carries the close code and reason for proper handling in the reconnect loop. - """ + """Raised when the QQ WebSocket closes; carries code + reason for the reconnect loop.""" def __init__(self, code, reason=""): self.code = int(code) if code else None @@ -89,10 +85,6 @@ class QQCloseError(Exception): super().__init__(f"WebSocket closed (code={self.code}, reason={self.reason})") -# --------------------------------------------------------------------------- -# Constants — imported from the shared constants module. -# --------------------------------------------------------------------------- - from gateway.platforms.qqbot.constants import ( API_BASE, TOKEN_URL, @@ -117,10 +109,7 @@ from gateway.platforms.qqbot.constants import ( MEDIA_TYPE_VOICE, MEDIA_TYPE_FILE, ) -from gateway.platforms.qqbot.utils import ( - coerce_list as _coerce_list_impl, - build_user_agent, -) +from gateway.platforms.qqbot.utils import coerce_list as _coerce_list, build_user_agent from gateway.platforms.qqbot.chunked_upload import ( ChunkedUploader, UploadDailyLimitExceededError, @@ -136,6 +125,7 @@ from gateway.platforms.qqbot.keyboards import ( parse_interaction_event, parse_update_prompt_button_data, ) +from gateway.platforms._shared import get_scoped_secret as _resolve_qq_secret def check_qq_requirements() -> bool: @@ -143,39 +133,13 @@ def check_qq_requirements() -> bool: return AIOHTTP_AVAILABLE and HTTPX_AVAILABLE -def _coerce_list(value: Any) -> List[str]: - """Coerce config values into a trimmed string list.""" - return _coerce_list_impl(value) - - -def _resolve_qq_secret(name: str, default: str = "") -> str: - """Resolve a per-profile ``QQ_*`` setting honoring the active secret scope. - - When a profile secret scope is installed — every secondary multiplex - profile is constructed and handled inside ``_profile_runtime_scope`` - (``gateway/run.py``), as is each per-turn inbound message — read from it so - profiles never see each other's ``os.environ`` values. This is the - cross-profile credential collision fixed for the WeChat adapter in #59662. - - The primary/active profile is constructed without a scope and legitimately - owns ``os.environ``, so fall back to it there instead of failing closed: a - bare ``get_secret`` would raise ``UnscopedSecretError`` on the active - profile's ``__init__`` and break its startup. Same pattern as the Slack - ``SLACK_APP_TOKEN`` read (#59739) and - ``gateway.platforms.whatsapp_common._get_wsecret``. - """ - from agent.secret_scope import UnscopedSecretError, get_secret - - try: - val = get_secret(name, default) - except UnscopedSecretError: - val = os.getenv(name) - return val if val is not None else default - - -# --------------------------------------------------------------------------- -# QQAdapter -# --------------------------------------------------------------------------- +_VOICE_EXTENSIONS = (".silk", ".amr", ".mp3", ".wav", ".ogg", ".m4a", ".aac", ".speex", ".flac") +_STT_PROVIDER_BASE_URLS = { + "zai": "https://open.bigmodel.cn/api/coding/paas/v4", + "openai": "https://api.openai.com/v1", + "glm": "https://open.bigmodel.cn/api/coding/paas/v4", +} +_AUDIO_URL_EXTENSIONS = {".silk", ".amr", ".mp3", ".wav", ".ogg", ".m4a", ".aac", ".flac"} class QQAdapter(BasePlatformAdapter): @@ -187,13 +151,28 @@ class QQAdapter(BasePlatformAdapter): _TYPING_INPUT_SECONDS = 60 # input_notify duration reported to QQ _TYPING_DEBOUNCE_SECONDS = 50 # refresh before it expires + # WS close codes that are unrecoverable → stop reconnecting. + _FATAL_CLOSE_CODES = { + 4001: "invalid opcode", + 4002: "invalid payload", + 4010: "invalid shard", + 4011: "sharding required", + 4012: "invalid API version", + 4013: "invalid intent", + 4014: "intent not authorized", + 4914: "offline/sandbox-only", + 4915: "banned", + } + # WS close codes that invalidate the session → clear it and re-identify on + # the next Hello. 4009 (connection timeout) is deliberately absent: it is + # resumable per the QQ protocol and must keep session state. + _SESSION_INVALID_CLOSE_CODES = {4006, 4007} | set(range(4900, 4914)) + @property def _log_tag(self) -> str: """Log prefix including app_id for multi-instance disambiguation.""" app_id = getattr(self, "_app_id", None) - if app_id: - return f"QQBot:{app_id}" - return "QQBot" + return f"QQBot:{app_id}" if app_id else "QQBot" def _fail_pending(self, reason: str) -> None: """Fail all pending response futures.""" @@ -203,20 +182,12 @@ class QQAdapter(BasePlatformAdapter): self._pending_responses.clear() def _mark_transport_disconnected(self) -> None: - """Mark QQ WS down without stopping the reconnect loop. - - BasePlatformAdapter uses _running for both process lifecycle and - connection status. QQBot needs to keep the listener task alive across - transient transport drops so it can continue reconnect attempts after a - short-lived gateway or network failure. - """ + """Mark QQ WS down without stopping the reconnect loop (base's _running + doubles as lifecycle flag; the listener must survive transient drops).""" if self.has_fatal_error: return self._write_runtime_status_safe( - "disconnected", - platform_state="disconnected", - error_code=None, - error_message=None, + "disconnected", platform_state="disconnected", error_code=None, error_message=None ) @property @@ -228,9 +199,7 @@ class QQAdapter(BasePlatformAdapter): super().__init__(config, Platform.QQBOT) extra = config.extra or {} - self._app_id = str( - extra.get("app_id") or _resolve_qq_secret("QQ_APP_ID", "") - ).strip() + self._app_id = str(extra.get("app_id") or _resolve_qq_secret("QQ_APP_ID", "")).strip() self._client_secret = str( extra.get("client_secret") or _resolve_qq_secret("QQ_CLIENT_SECRET", "") ).strip() @@ -238,9 +207,7 @@ class QQAdapter(BasePlatformAdapter): # Auth/ACL policies self._dm_policy = str(extra.get("dm_policy", "pairing")).strip().lower() - self._allow_from = _coerce_list( - extra.get("allow_from") or extra.get("allowFrom") - ) + self._allow_from = _coerce_list(extra.get("allow_from") or extra.get("allowFrom")) self._group_policy = str(extra.get("group_policy", "pairing")).strip().lower() self._group_allow_from = _coerce_list( extra.get("group_allow_from") or extra.get("groupAllowFrom") @@ -271,24 +238,12 @@ class QQAdapter(BasePlatformAdapter): self._token_expires_at: float = 0.0 self._token_lock = asyncio.Lock() - # Upload cache: content_hash -> {file_info, file_uuid, expires_at} - self._upload_cache: Dict[str, Dict[str, Any]] = {} - - # Inline-keyboard interaction routing. The callback (if set) is invoked - # for every INTERACTION_CREATE event after the adapter has already - # ACKed it. Callers (gateway wiring for approvals / update prompts) - # register via set_interaction_callback(). + # Inline-keyboard interaction routing: invoked for every INTERACTION_CREATE + # after the adapter ACKed it. Defaults to the approval/update-prompt + # dispatcher; override via set_interaction_callback() (None drops clicks). self._interaction_callback: Optional[ Callable[[InteractionEvent], Awaitable[None]] - ] = None - - # Default interaction dispatcher: routes approval-button clicks to - # tools.approval.resolve_gateway_approval() and update-prompt clicks - # to ~/.hermes/.update_response. Set here so the cross-adapter gateway - # contract (send_exec_approval / send_update_prompt) works out of the - # box; callers can override with set_interaction_callback(None) or - # register a custom handler. - self._interaction_callback = self._default_interaction_dispatch + ] = self._default_interaction_dispatch # ------------------------------------------------------------------ # Properties @@ -311,35 +266,26 @@ class QQAdapter(BasePlatformAdapter): """ Authenticate, obtain gateway URL, and open the WebSocket. - Args: - is_reconnect: False on a cold first boot; True when the - reconnect watcher is re-establishing this platform after - an outage. QQBot has no server-side update queue so this - flag is accepted for interface conformance only. + ``is_reconnect`` is accepted for interface conformance only (QQBot has no + server-side update queue). """ - if not AIOHTTP_AVAILABLE: - message = "QQ startup failed: aiohttp not installed" - self._set_fatal_error("qq_missing_dependency", message, retryable=True) - logger.warning("[%s] %s. Run: pip install aiohttp", self._log_tag, message) - return False - if not HTTPX_AVAILABLE: - message = "QQ startup failed: httpx not installed" - self._set_fatal_error("qq_missing_dependency", message, retryable=True) - logger.warning("[%s] %s. Run: pip install httpx", self._log_tag, message) - return False - if not self._app_id or not self._client_secret: - message = "QQ startup failed: QQ_APP_ID and QQ_CLIENT_SECRET are required" - self._set_fatal_error("qq_missing_credentials", message, retryable=True) - logger.warning("[%s] %s", self._log_tag, message) - return False + for ok, code, message, hint in ( + (AIOHTTP_AVAILABLE, "qq_missing_dependency", "QQ startup failed: aiohttp not installed", ". Run: pip install aiohttp"), + (HTTPX_AVAILABLE, "qq_missing_dependency", "QQ startup failed: httpx not installed", ". Run: pip install httpx"), + (self._app_id and self._client_secret, "qq_missing_credentials", + "QQ startup failed: QQ_APP_ID and QQ_CLIENT_SECRET are required", ""), + ): + if not ok: + self._set_fatal_error(code, message, retryable=True) + logger.warning("[%s] %s%s", self._log_tag, message, hint) + return False # Prevent duplicate connections with the same credentials if not self._acquire_platform_lock("qqbot-appid", self._app_id, "QQBot app ID"): return False try: - # Tighter keepalive pool so idle CLOSE_WAIT sockets drain - # faster behind proxies like Cloudflare Warp (#18451). + # Tighter keepalive pool so idle CLOSE_WAIT sockets drain faster behind proxies. from gateway.platforms._http_client_limits import platform_httpx_limits from tools.url_safety import create_ssrf_safe_async_client self._http_client = create_ssrf_safe_async_client( @@ -349,17 +295,10 @@ class QQAdapter(BasePlatformAdapter): limits=platform_httpx_limits(), ) - # 1. Get access token await self._ensure_token() - - # 2. Get WebSocket gateway URL gateway_url = await self._get_gateway_url() logger.info("[%s] Gateway URL: %s", self._log_tag, gateway_url) - - # 3. Open WebSocket await self._open_ws(gateway_url) - - # 4. Start listeners self._listen_task = asyncio.create_task(self._listen_loop()) self._heartbeat_task = asyncio.create_task(self._heartbeat_loop()) self._mark_connected() @@ -379,46 +318,39 @@ class QQAdapter(BasePlatformAdapter): """Close all connections and stop listeners.""" self._running = False self._mark_disconnected() - - if self._listen_task: - self._listen_task.cancel() - try: - await self._listen_task - except asyncio.CancelledError: - pass - self._listen_task = None - - if self._heartbeat_task: - self._heartbeat_task.cancel() - try: - await self._heartbeat_task - except asyncio.CancelledError: - pass - self._heartbeat_task = None - + self._listen_task = await self._cancel_task(self._listen_task) + self._heartbeat_task = await self._cancel_task(self._heartbeat_task) await self._cleanup() self._release_platform_lock() logger.info("[%s] Disconnected", self._log_tag) - async def _cleanup(self) -> None: - """Close WebSocket, HTTP session, and client.""" + @staticmethod + async def _cancel_task(task: Optional[asyncio.Task]) -> None: + """Cancel and await *task* (if any); always returns None for reassignment.""" + if task: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + return None + + async def _close_ws(self) -> None: + """Close the WebSocket + its aiohttp session (keeps _http_client alive).""" if self._ws and not self._ws.closed: await self._ws.close() self._ws = None - if self._session and not self._session.closed: await self._session.close() self._session = None + async def _cleanup(self) -> None: + """Close WebSocket, HTTP session, and client; fail pending futures.""" + await self._close_ws() if self._http_client: await self._http_client.aclose() self._http_client = None - - # Fail pending - for fut in self._pending_responses.values(): - if not fut.done(): - fut.set_exception(RuntimeError("Disconnected")) - self._pending_responses.clear() + self._fail_pending("Disconnected") # ------------------------------------------------------------------ # Token management @@ -447,16 +379,12 @@ class QQAdapter(BasePlatformAdapter): token = data.get("access_token") if not token: - raise RuntimeError( - f"QQ Bot token response missing access_token: {data}" - ) + raise RuntimeError(f"QQ Bot token response missing access_token: {data}") expires_in = int(data.get("expires_in", 7200)) self._access_token = token self._token_expires_at = time.time() + expires_in - logger.info( - "[%s] Access token refreshed, expires in %ds", self._log_tag, expires_in - ) + logger.info("[%s] Access token refreshed, expires in %ds", self._log_tag, expires_in) return self._access_token async def _get_gateway_url(self) -> str: @@ -465,10 +393,7 @@ class QQAdapter(BasePlatformAdapter): try: resp = await self._http_client.get( f"{API_BASE}{GATEWAY_URL_PATH}", - headers={ - "Authorization": f"QQBot {token}", - "User-Agent": build_user_agent(), - }, + headers={"Authorization": f"QQBot {token}", "User-Agent": build_user_agent()}, timeout=DEFAULT_API_TIMEOUT, ) resp.raise_for_status() @@ -487,16 +412,8 @@ class QQAdapter(BasePlatformAdapter): async def _open_ws(self, gateway_url: str) -> None: """Open a WebSocket connection to the QQ Bot gateway.""" - # Only clean up WebSocket resources — keep _http_client alive for REST API calls. - if self._ws and not self._ws.closed: - await self._ws.close() - self._ws = None - if self._session and not self._session.closed: - await self._session.close() - self._session = None - - # Honor WSL proxy env for QQ WebSocket. Hermes upgrades overwrite this - # local patch, so QQ can regress to direct-connect timeouts after update. + await self._close_ws() + # Honor proxy env vars for the WebSocket (WSL setups need this). self._session = aiohttp.ClientSession(trust_env=gateway_trust_env()) ws_proxy = ( os.getenv("WSS_PROXY") @@ -508,9 +425,7 @@ class QQAdapter(BasePlatformAdapter): ) self._ws = await self._session.ws_connect( gateway_url, - headers={ - "User-Agent": build_user_agent(), - }, + headers={"User-Agent": build_user_agent()}, timeout=CONNECT_TIMEOUT_SECONDS, proxy=ws_proxy, ) @@ -519,12 +434,8 @@ class QQAdapter(BasePlatformAdapter): async def _listen_loop(self) -> None: """Read WebSocket events and reconnect on errors. - Close code handling follows the OpenClaw qqbot reference implementation: - 4004 → invalid token, refresh and reconnect - 4006/4007/4009 → session invalid, clear session and re-identify - 4008 → rate limited, back off 60s - 4914 → bot offline/sandbox, stop reconnecting - 4915 → bot banned, stop reconnecting + Close codes: 4004 → refresh token; 4006/4007/49xx → clear session and + re-identify; 4008 → rate limited, back off; _FATAL_CLOSE_CODES → stop. """ backoff_idx = 0 connect_time = 0.0 @@ -544,10 +455,7 @@ class QQAdapter(BasePlatformAdapter): code = exc.code logger.warning( - "[%s] WebSocket closed: code=%s reason=%s", - self._log_tag, - code, - exc.reason, + "[%s] WebSocket closed: code=%s reason=%s", self._log_tag, code, exc.reason ) # Quick disconnect detection (permission issues, misconfiguration) @@ -556,9 +464,7 @@ class QQAdapter(BasePlatformAdapter): quick_disconnect_count += 1 logger.info( "[%s] Quick disconnect (%.1fs), count: %d", - self._log_tag, - duration, - quick_disconnect_count, + self._log_tag, duration, quick_disconnect_count, ) if quick_disconnect_count >= MAX_QUICK_DISCONNECT_COUNT: logger.error( @@ -578,44 +484,15 @@ class QQAdapter(BasePlatformAdapter): self._mark_transport_disconnected() self._fail_pending("Connection closed") - # Stop reconnecting for fatal codes (unrecoverable errors) - if code in { - 4001, # Invalid opcode - 4002, # Invalid payload - 4010, # Invalid shard - 4011, # Sharding required - 4012, # Invalid API version - 4013, # Invalid intent - 4014, # Intent not authorized - 4914, # Offline/sandbox-only - 4915, # Banned - }: - fatal_descriptions = { - 4001: "invalid opcode", - 4002: "invalid payload", - 4010: "invalid shard", - 4011: "sharding required", - 4012: "invalid API version", - 4013: "invalid intent", - 4014: "intent not authorized", - 4914: "offline/sandbox-only", - 4915: "banned", - } - desc = fatal_descriptions.get(code, f"fatal error (code={code})") - logger.error( - "[%s] Bot is %s. Check QQ Open Platform.", self._log_tag, desc - ) - self._set_fatal_error( - f"qq_{desc}", f"Bot is {desc}", retryable=False - ) + desc = self._FATAL_CLOSE_CODES.get(code) + if desc: + logger.error("[%s] Bot is %s. Check QQ Open Platform.", self._log_tag, desc) + self._set_fatal_error(f"qq_{desc}", f"Bot is {desc}", retryable=False) return - # Rate limited if code == 4008: logger.info( - "[%s] Rate limited (4008), waiting %ds", - self._log_tag, - RATE_LIMIT_DELAY, + "[%s] Rate limited (4008), waiting %ds", self._log_tag, RATE_LIMIT_DELAY ) if backoff_idx >= MAX_RECONNECT_ATTEMPTS: self._mark_disconnected() @@ -628,40 +505,17 @@ class QQAdapter(BasePlatformAdapter): backoff_idx += 1 continue - # Token invalid → clear cached token so _ensure_token() refreshes if code == 4004: logger.info( - "[%s] Invalid token (4004), will refresh and reconnect", - self._log_tag, + "[%s] Invalid token (4004), will refresh and reconnect", self._log_tag ) self._access_token = None self._token_expires_at = 0.0 - # Session invalid → clear session, will re-identify on next Hello - # Note: 4009 (connection timeout) is NOT included here — it is - # resumable per the QQ protocol and should preserve session state. - if code in { - 4006, - 4007, - 4900, - 4901, - 4902, - 4903, - 4904, - 4905, - 4906, - 4907, - 4908, - 4909, - 4910, - 4911, - 4912, - 4913, - }: + if code in self._SESSION_INVALID_CLOSE_CODES: logger.info( "[%s] Session error (%d), clearing session for re-identify", - self._log_tag, - code, + self._log_tag, code, ) self._session_id = None self._last_seq = None @@ -697,12 +551,7 @@ class QQAdapter(BasePlatformAdapter): async def _reconnect(self, backoff_idx: int) -> bool: """Attempt to reconnect the WebSocket. Returns True on success.""" delay = RECONNECT_BACKOFF[min(backoff_idx, len(RECONNECT_BACKOFF) - 1)] - logger.info( - "[%s] Reconnecting in %ds (attempt %d)...", - self._log_tag, - delay, - backoff_idx + 1, - ) + logger.info("[%s] Reconnecting in %ds (attempt %d)...", self._log_tag, delay, backoff_idx + 1) await asyncio.sleep(delay) self._heartbeat_interval = 30.0 # reset until Hello @@ -722,10 +571,8 @@ class QQAdapter(BasePlatformAdapter): if not self._ws: raise RuntimeError("WebSocket not connected") if self._ws.closed: - # A closed-but-non-None ws makes the while-condition false on entry, - # so this would return normally — which _listen_loop treats as a - # clean read and immediately retries with backoff reset to 0, - # producing a 100% CPU spin. Raise so the reconnect/backoff path runs. + # Returning normally here would make _listen_loop treat it as a clean + # read and retry with backoff reset → 100% CPU spin. Raise instead. raise RuntimeError("WebSocket closed") while self._running and self._ws and not self._ws.closed: @@ -734,113 +581,74 @@ class QQAdapter(BasePlatformAdapter): payload = self._parse_json(msg.data) if payload: self._dispatch_payload(payload) - elif msg.type in {aiohttp.WSMsgType.PING,}: - # aiohttp auto-replies with PONG - pass elif msg.type == aiohttp.WSMsgType.CLOSE: raise QQCloseError(msg.data, msg.extra) elif msg.type in {aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR}: raise RuntimeError("WebSocket closed") async def _heartbeat_loop(self) -> None: - """Send periodic heartbeats (QQ Gateway expects op 1 heartbeat with latest seq). - - The interval is set from the Hello (op 10) event's heartbeat_interval. - QQ's default is ~41s; we send at 80% of the interval to stay safe. - """ + """Send op 1 heartbeats with the latest seq at 80% of the Hello interval.""" try: while self._running: await asyncio.sleep(self._heartbeat_interval) if not self._ws or self._ws.closed: continue try: - # d should be the latest sequence number received, or null await self._ws.send_json({"op": 1, "d": self._last_seq}) except Exception as exc: logger.debug("[%s] Heartbeat failed: %s", self._log_tag, exc) except asyncio.CancelledError: pass + async def _send_ws_auth(self, name: str, payload: Dict[str, Any], sent_msg: str, *log_args) -> bool: + """Send an Identify/Resume payload; returns False if the send raised.""" + try: + if self._ws and not self._ws.closed: + await self._ws.send_json(payload) + logger.info("[%s] " + sent_msg, self._log_tag, *log_args) + else: + logger.warning("[%s] Cannot send %s: WebSocket not connected", self._log_tag, name) + except Exception as exc: + logger.error("[%s] Failed to send %s: %s", self._log_tag, name, exc) + return False + return True + async def _send_identify(self) -> None: - """Send op 2 Identify to authenticate the WebSocket connection. + """Send op 2 Identify (reply to Hello); server answers with READY. - After receiving op 10 Hello, the client must send op 2 Identify with - the bot token and intents. On success the server replies with a - READY dispatch event. - - Reference: https://bot.q.qq.com/wiki/develop/api-v2/dev-prepare/interface-framework/reference.html + Intents: C2C_GROUP_AT_MESSAGES | PUBLIC_GUILD_MESSAGES | DIRECT_MESSAGE | INTERACTION. """ token = await self._ensure_token() - identify_payload = { + payload = { "op": 2, "d": { "token": f"QQBot {token}", - "intents": (1 << 25) - | (1 << 30) - | (1 << 12) - | (1 << 26), # C2C_GROUP_AT_MESSAGES + PUBLIC_GUILD_MESSAGES + DIRECT_MESSAGE + INTERACTION + "intents": (1 << 25) | (1 << 30) | (1 << 12) | (1 << 26), "shard": [0, 1], - "properties": { - "$os": "macOS", - "$browser": "hermes-agent", - "$device": "hermes-agent", - }, + "properties": {"$os": "macOS", "$browser": "hermes-agent", "$device": "hermes-agent"}, }, } - try: - if self._ws and not self._ws.closed: - await self._ws.send_json(identify_payload) - logger.info("[%s] Identify sent", self._log_tag) - else: - logger.warning( - "[%s] Cannot send Identify: WebSocket not connected", self._log_tag - ) - except Exception as exc: - logger.error("[%s] Failed to send Identify: %s", self._log_tag, exc) + await self._send_ws_auth("Identify", payload, "Identify sent") async def _send_resume(self) -> None: - """Send op 6 Resume to re-authenticate after a reconnection. - - Reference: https://bot.q.qq.com/wiki/develop/api-v2/dev-prepare/interface-framework/reference.html - """ + """Send op 6 Resume after a reconnect; on failure clear session → Identify next Hello.""" token = await self._ensure_token() - resume_payload = { + payload = { "op": 6, - "d": { - "token": f"QQBot {token}", - "session_id": self._session_id, - "seq": self._last_seq, - }, + "d": {"token": f"QQBot {token}", "session_id": self._session_id, "seq": self._last_seq}, } - try: - if self._ws and not self._ws.closed: - await self._ws.send_json(resume_payload) - logger.info( - "[%s] Resume sent (session_id=%s, seq=%s)", - self._log_tag, - self._session_id, - self._last_seq, - ) - else: - logger.warning( - "[%s] Cannot send Resume: WebSocket not connected", self._log_tag - ) - except Exception as exc: - logger.error("[%s] Failed to send Resume: %s", self._log_tag, exc) - # If resume fails, clear session and fall back to identify on next Hello + if not await self._send_ws_auth( + "Resume", payload, "Resume sent (session_id=%s, seq=%s)", self._session_id, self._last_seq + ): self._session_id = None self._last_seq = None @staticmethod def _create_task(coro): - """Schedule a coroutine, silently skipping if no event loop is running. - - This avoids ``RuntimeError: no running event loop`` when tests call - ``_dispatch_payload`` synchronously outside of ``asyncio.run()``. - """ + """Schedule a coroutine; returns None (no error) when no loop is running + (tests call _dispatch_payload synchronously).""" try: - loop = asyncio.get_running_loop() - return loop.create_task(coro) + return asyncio.get_running_loop().create_task(coro) except RuntimeError: return None @@ -853,39 +661,26 @@ class QQAdapter(BasePlatformAdapter): if isinstance(s, int) and (self._last_seq is None or s > self._last_seq): self._last_seq = s - # op 10 = Hello (heartbeat interval) — must reply with Identify/Resume - if op == 10: + if op == 10: # Hello — reply with Resume (have session) or Identify d_data = d if isinstance(d, dict) else {} interval_ms = d_data.get("heartbeat_interval", 30000) - # Send heartbeats at 80% of the server interval to stay safe - self._heartbeat_interval = interval_ms / 1000.0 * 0.8 + self._heartbeat_interval = interval_ms / 1000.0 * 0.8 # 80% of server interval logger.debug( "[%s] Hello received, heartbeat_interval=%dms (sending every %.1fs)", - self._log_tag, - interval_ms, - self._heartbeat_interval, + self._log_tag, interval_ms, self._heartbeat_interval, ) - # Authenticate: send Resume if we have a session, else Identify. - # Use _create_task which is safe when no event loop is running (tests). if self._session_id and self._last_seq is not None: self._create_task(self._send_resume()) else: self._create_task(self._send_identify()) return - # op 0 = Dispatch - if op == 0 and t: + if op == 0 and t: # Dispatch if t == "READY": self._handle_ready(d) elif t == "RESUMED": logger.info("[%s] Session resumed", self._log_tag) - elif t in { - "C2C_MESSAGE_CREATE", - "GROUP_AT_MESSAGE_CREATE", - "DIRECT_MESSAGE_CREATE", - "GUILD_MESSAGE_CREATE", - "GUILD_AT_MESSAGE_CREATE", - }: + elif t in self._INBOUND_HANDLERS: asyncio.create_task(self._on_message(t, d)) elif t == "INTERACTION_CREATE": self._create_task(self._on_interaction(d)) @@ -893,32 +688,24 @@ class QQAdapter(BasePlatformAdapter): logger.debug("[%s] Unhandled dispatch: %s", self._log_tag, t) return - # op 11 = Heartbeat ACK - if op == 11: + if op == 11: # Heartbeat ACK return - # op 7 = Server Reconnect — server asks client to reconnect (e.g. - # load-balancing, maintenance). Close the WS so _read_events raises - # and the outer loop triggers a reconnect with Resume. - if op == 7: + if op == 7: # Server Reconnect — close so _read_events raises and we reconnect w/ Resume logger.info("[%s] Server requested reconnect (op 7)", self._log_tag) if self._ws and not self._ws.closed: self._create_task(self._ws.close()) return - # op 9 = Invalid Session — d=True means session is resumable, - # d=False means we must re-identify from scratch. - if op == 9: - resumable = bool(d) if d is not None else False - if not resumable: + if op == 9: # Invalid Session — d=True resumable, d=False re-identify from scratch + if d is not None and bool(d): + logger.info("[%s] Invalid session (op 9, resumable)", self._log_tag) + else: logger.info( - "[%s] Invalid session (op 9, not resumable), clearing session", - self._log_tag, + "[%s] Invalid session (op 9, not resumable), clearing session", self._log_tag ) self._session_id = None self._last_seq = None - else: - logger.info("[%s] Invalid session (op 9, resumable)", self._log_tag) if self._ws and not self._ws.closed: self._create_task(self._ws.close()) return @@ -965,28 +752,18 @@ class QQAdapter(BasePlatformAdapter): """Process an inbound QQ Bot message event.""" if not isinstance(d, dict): return - - # Extract common fields msg_id = str(d.get("id", "")) if not msg_id or self._is_duplicate(msg_id): - logger.debug( - "[%s] Duplicate or missing message id: %s", self._log_tag, msg_id - ) + logger.debug("[%s] Duplicate or missing message id: %s", self._log_tag, msg_id) return timestamp = str(d.get("timestamp", "")) content = str(d.get("content", "")).strip() author = d.get("author") if isinstance(d.get("author"), dict) else {} - # Route by event type - if event_type == "C2C_MESSAGE_CREATE": - await self._handle_c2c_message(d, msg_id, content, author, timestamp) - elif event_type in {"GROUP_AT_MESSAGE_CREATE",}: - await self._handle_group_message(d, msg_id, content, author, timestamp) - elif event_type in {"GUILD_MESSAGE_CREATE", "GUILD_AT_MESSAGE_CREATE"}: - await self._handle_guild_message(d, msg_id, content, author, timestamp) - elif event_type == "DIRECT_MESSAGE_CREATE": - await self._handle_dm_message(d, msg_id, content, author, timestamp) + handler = self._INBOUND_HANDLERS.get(event_type) + if handler: + await getattr(self, handler)(d, msg_id, content, author, timestamp) # ------------------------------------------------------------------ # Inline-keyboard interactions (INTERACTION_CREATE) @@ -996,113 +773,60 @@ class QQAdapter(BasePlatformAdapter): self, callback: Optional[Callable[[InteractionEvent], Awaitable[None]]], ) -> None: - """Register (or clear) the interaction callback. - - Invoked once per ``INTERACTION_CREATE`` event *after* the adapter has - ACKed the interaction. The callback is responsible for routing the - button click to the right subsystem (approval resolver, update-prompt - resolver, etc.) based on the ``button_data`` payload. - """ + """Register (or clear) the callback invoked per ACKed INTERACTION_CREATE.""" self._interaction_callback = callback async def _on_interaction(self, d: Any) -> None: - """Handle an ``INTERACTION_CREATE`` event. - - Responsibilities: - - 1. Parse the raw payload into an :class:`InteractionEvent`. - 2. ACK the interaction (``PUT /interactions/{id}``) so the client - stops showing a loading indicator on the button. - 3. Dispatch to the registered interaction callback, if any. - """ + """Parse INTERACTION_CREATE, ACK it promptly (else the client shows an error + icon on the button), then dispatch to the registered callback.""" if not isinstance(d, dict): return try: event = parse_interaction_event(d) except Exception as exc: - logger.warning( - "[%s] Failed to parse INTERACTION_CREATE: %s", self._log_tag, exc - ) + logger.warning("[%s] Failed to parse INTERACTION_CREATE: %s", self._log_tag, exc) return - if not event.id: - logger.warning( - "[%s] INTERACTION_CREATE missing id, skipping ACK", self._log_tag - ) + logger.warning("[%s] INTERACTION_CREATE missing id, skipping ACK", self._log_tag) return - # ACK the interaction promptly — per the QQ docs the client will show - # an error icon on the button if we don't respond quickly. try: await self._acknowledge_interaction(event.id) except Exception as exc: - logger.warning( - "[%s] Failed to ACK interaction %s: %s", - self._log_tag, event.id, exc, - ) + logger.warning("[%s] Failed to ACK interaction %s: %s", self._log_tag, event.id, exc) logger.info( "[%s] Interaction: scene=%s button_data=%r operator=%s", self._log_tag, event.scene, event.button_data, event.operator_openid, ) - callback = self._interaction_callback if callback is None: logger.debug( - "[%s] No interaction callback registered; dropping button " - "click %r", + "[%s] No interaction callback registered; dropping button click %r", self._log_tag, event.button_data, ) return try: await callback(event) except Exception as exc: - logger.error( - "[%s] Interaction callback raised: %s", - self._log_tag, exc, exc_info=True, - ) + logger.error("[%s] Interaction callback raised: %s", self._log_tag, exc, exc_info=True) - async def _acknowledge_interaction( - self, - interaction_id: str, - code: int = 0, - ) -> None: - """ACK a button interaction via ``PUT /interactions/{id}``. - - :param interaction_id: The ``id`` field from the - ``INTERACTION_CREATE`` event. - :param code: Response code (``0`` = success). - """ + async def _acknowledge_interaction(self, interaction_id: str, code: int = 0) -> None: + """ACK a button interaction via ``PUT /interactions/{id}`` (code 0 = success).""" if not self._http_client: raise RuntimeError("HTTP client not initialized — not connected?") - token = await self._ensure_token() - headers = { - "Authorization": f"QQBot {token}", - "Content-Type": "application/json", - "User-Agent": build_user_agent(), - } resp = await self._http_client.put( f"{API_BASE}/interactions/{interaction_id}", - headers=headers, + headers=await self._auth_headers(), json={"code": code}, timeout=DEFAULT_API_TIMEOUT, ) if resp.status_code >= 400: - raise RuntimeError( - f"Interaction ACK failed [{resp.status_code}]: " - f"{resp.text[:200]}" - ) + raise RuntimeError(f"Interaction ACK failed [{resp.status_code}]: {resp.text[:200]}") - # Mapping from QQ keyboard button decisions → the ``choice`` vocabulary - # accepted by ``tools.approval.resolve_gateway_approval``. QQ's 3-button - # layout (mobile-space constraint) collapses "session" and "always" into - # a single "always" button; users wanting session-only approval can fall - # back to the ``/approve session`` text command. - _APPROVAL_BUTTON_TO_CHOICE = { - "allow-once": "once", - "allow-always": "always", - "deny": "deny", - } + # Button decision → ``choice`` for tools.approval.resolve_gateway_approval. The + # 3-button layout folds "session" into "always"; ``/approve session`` still works. + _APPROVAL_BUTTON_TO_CHOICE = {"allow-once": "once", "allow-always": "always", "deny": "deny"} @staticmethod def _parse_gateway_session_key(session_key: str) -> Optional[Dict[str, str]]: @@ -1110,11 +834,7 @@ class QQAdapter(BasePlatformAdapter): parts = str(session_key or "").split(":") if len(parts) < 5 or parts[0] != "agent" or parts[1] != "main": return None - parsed = { - "platform": parts[2], - "chat_type": parts[3], - "chat_id": parts[4], - } + parsed = {"platform": parts[2], "chat_type": parts[3], "chat_id": parts[4]} if len(parts) > 5: parsed["user_id"] = parts[5] return parsed @@ -1148,21 +868,9 @@ class QQAdapter(BasePlatformAdapter): self, event: InteractionEvent, ) -> None: - """Route ``INTERACTION_CREATE`` button clicks to the right subsystem. - - - ``approve::`` → - :func:`tools.approval.resolve_gateway_approval` - (unblocks the agent thread waiting on a dangerous-command approval). - - ``update_prompt:`` → - writes the answer to ``~/.hermes/.update_response`` for the - detached ``hermes update --gateway`` process to consume. - - Anything else is logged at DEBUG and ignored. - - Installed as the adapter's default interaction callback in - ``__init__``. Callers can replace via - :meth:`set_interaction_callback` to route clicks elsewhere (or pass - ``None`` to drop them entirely). - """ + """Default interaction callback: ``approve::`` → + tools.approval.resolve_gateway_approval; ``update_prompt:`` → + ``~/.hermes/.update_response``; anything else is ignored at DEBUG.""" button_data = event.button_data if not button_data: return @@ -1185,15 +893,11 @@ class QQAdapter(BasePlatformAdapter): ) return try: - # Import lazily to keep the adapter importable in tests that - # don't exercise the approval subsystem. - from tools.approval import resolve_gateway_approval + from tools.approval import resolve_gateway_approval # lazy: keep adapter light count = resolve_gateway_approval(session_key, choice) logger.info( - "[%s] Button resolved %d approval(s) for session %s " - "(choice=%s, operator=%s)", - self._log_tag, count, session_key, choice, - event.operator_openid, + "[%s] Button resolved %d approval(s) for session %s (choice=%s, operator=%s)", + self._log_tag, count, session_key, choice, event.operator_openid, ) except Exception as exc: logger.error( @@ -1221,24 +925,15 @@ class QQAdapter(BasePlatformAdapter): @staticmethod def _write_update_response(answer: str, operator: str = "") -> None: - """Atomically write the update-prompt answer to ``.update_response``. - - Mirrors the Discord / Telegram / Feishu adapters: the detached - ``hermes update --gateway`` watcher polls this file for a ``y``/``n`` - response to its interactive prompts (stash-restore, config migration). - Writes via ``tmp + rename`` so a partial write can't fool the reader. - """ + """Atomically (tmp + rename) write the update-prompt answer to + ``.update_response``, polled by the detached ``hermes update --gateway`` watcher.""" try: from hermes_constants import get_hermes_home - home = get_hermes_home() - response_path = home / ".update_response" + response_path = get_hermes_home() / ".update_response" tmp = response_path.with_suffix(".tmp") tmp.write_text(answer, encoding="utf-8") tmp.replace(response_path) - logger.info( - "QQ update prompt answered %r by %s", - answer, operator or "(unknown)", - ) + logger.info("QQ update prompt answered %r by %s", answer, operator or "(unknown)") except Exception as exc: logger.error("Failed to write update response: %s", exc) @@ -1257,7 +952,6 @@ class QQAdapter(BasePlatformAdapter): if not self._is_dm_intake_allowed(user_openid): return - text = content attachments_raw = d.get("attachments") logger.info( "[%s] C2C message: id=%s content=%r attachments=%s", @@ -1282,60 +976,16 @@ class QQAdapter(BasePlatformAdapter): _att.get("filename", ""), ) - # Process all attachments uniformly (images, voice, files) - att_result = await self._process_attachments(attachments_raw) - image_urls = att_result["image_urls"] - image_media_types = att_result["image_media_types"] - voice_transcripts = att_result["voice_transcripts"] - attachment_info = att_result["attachment_info"] - - # Append voice transcripts to the text body - if voice_transcripts: - voice_block = "\n".join(voice_transcripts) - text = ( - (text + "\n\n" + voice_block).strip() if text.strip() else voice_block - ) - # Append non-media attachment info - if attachment_info: - text = ( - (text + "\n\n" + attachment_info).strip() - if text.strip() - else attachment_info - ) - + text, image_urls, image_media_types, n_voice = await self._absorb_attachments( + content, attachments_raw + ) logger.info( - "[%s] After processing: images=%d, voice=%d", - self._log_tag, - len(image_urls), - len(voice_transcripts), + "[%s] After processing: images=%d, voice=%d", self._log_tag, len(image_urls), n_voice ) - - # Merge any quoted-message context (message_type=103 → msg_elements[0]). - quoted = await self._process_quoted_context(d) - text = self._merge_quote_into(text, quoted["quote_block"]) - if quoted["image_urls"]: - image_urls = image_urls + quoted["image_urls"] - image_media_types = image_media_types + quoted["image_media_types"] - - if not text.strip() and not image_urls: - return - - self._chat_type_map[user_openid] = "c2c" - event = MessageEvent( - source=self.build_source( - chat_id=user_openid, - user_id=user_openid, - chat_type="dm", - ), - text=text, - message_type=self._detect_message_type(image_urls, image_media_types), - raw_message=d, - message_id=msg_id, - media_urls=image_urls, - media_types=image_media_types, - timestamp=self._parse_qq_timestamp(timestamp), + await self._emit_inbound( + d, msg_id, timestamp, text, image_urls, image_media_types, + chat_id=user_openid, qq_chat_type="c2c", user_id=user_openid, chat_type="dm", ) - await self.handle_message(event) async def _handle_group_message( self, @@ -1349,58 +999,17 @@ class QQAdapter(BasePlatformAdapter): group_openid = str(d.get("group_openid", "")) if not group_openid: return - if not self._is_group_allowed( - group_openid, str(author.get("member_openid", "")) - ): + if not self._is_group_allowed(group_openid, str(author.get("member_openid", ""))): return - # Strip the @bot mention prefix from content - text = self._strip_at_mention(content) - att_result = await self._process_attachments(d.get("attachments")) - image_urls = att_result["image_urls"] - image_media_types = att_result["image_media_types"] - voice_transcripts = att_result["voice_transcripts"] - attachment_info = att_result["attachment_info"] - - # Append voice transcripts - if voice_transcripts: - voice_block = "\n".join(voice_transcripts) - text = ( - (text + "\n\n" + voice_block).strip() if text.strip() else voice_block - ) - if attachment_info: - text = ( - (text + "\n\n" + attachment_info).strip() - if text.strip() - else attachment_info - ) - - # Merge any quoted-message context (message_type=103 → msg_elements[0]). - quoted = await self._process_quoted_context(d) - text = self._merge_quote_into(text, quoted["quote_block"]) - if quoted["image_urls"]: - image_urls = image_urls + quoted["image_urls"] - image_media_types = image_media_types + quoted["image_media_types"] - - if not text.strip() and not image_urls: - return - - self._chat_type_map[group_openid] = "group" - event = MessageEvent( - source=self.build_source( - chat_id=group_openid, - user_id=str(author.get("member_openid", "")), - chat_type="group", - ), - text=text, - message_type=self._detect_message_type(image_urls, image_media_types), - raw_message=d, - message_id=msg_id, - media_urls=image_urls, - media_types=image_media_types, - timestamp=self._parse_qq_timestamp(timestamp), + text, image_urls, image_media_types, _ = await self._absorb_attachments( + self._strip_at_mention(content), d.get("attachments") + ) + await self._emit_inbound( + d, msg_id, timestamp, text, image_urls, image_media_types, + chat_id=group_openid, qq_chat_type="group", + user_id=str(author.get("member_openid", "")), chat_type="group", ) - await self.handle_message(event) async def _handle_guild_message( self, @@ -1415,9 +1024,8 @@ class QQAdapter(BasePlatformAdapter): if not channel_id: return - # Apply group_policy ACL — guild channels are group-like contexts. - # Without this check any member of any guild the bot is in could - # bypass the configured allowlist. + # group_policy ACL — guild channels are group-like; without it any guild + # member could bypass the allowlist. guild_id = str(d.get("guild_id", "")) author_id = str(author.get("id", "")) if not self._is_group_allowed(guild_id or channel_id, author_id): @@ -1430,52 +1038,14 @@ class QQAdapter(BasePlatformAdapter): member = d.get("member") if isinstance(d.get("member"), dict) else {} nick = str(member.get("nick", "")) or str(author.get("username", "")) - text = content - att_result = await self._process_attachments(d.get("attachments")) - image_urls = att_result["image_urls"] - image_media_types = att_result["image_media_types"] - voice_transcripts = att_result["voice_transcripts"] - attachment_info = att_result["attachment_info"] - - if voice_transcripts: - voice_block = "\n".join(voice_transcripts) - text = ( - (text + "\n\n" + voice_block).strip() if text.strip() else voice_block - ) - if attachment_info: - text = ( - (text + "\n\n" + attachment_info).strip() - if text.strip() - else attachment_info - ) - - # Merge any quoted-message context (message_type=103 → msg_elements[0]). - quoted = await self._process_quoted_context(d) - text = self._merge_quote_into(text, quoted["quote_block"]) - if quoted["image_urls"]: - image_urls = image_urls + quoted["image_urls"] - image_media_types = image_media_types + quoted["image_media_types"] - - if not text.strip() and not image_urls: - return - - self._chat_type_map[channel_id] = "guild" - event = MessageEvent( - source=self.build_source( - chat_id=channel_id, - user_id=str(author.get("id", "")), - user_name=nick or None, - chat_type="group", - ), - text=text, - message_type=self._detect_message_type(image_urls, image_media_types), - raw_message=d, - message_id=msg_id, - media_urls=image_urls, - media_types=image_media_types, - timestamp=self._parse_qq_timestamp(timestamp), + text, image_urls, image_media_types, _ = await self._absorb_attachments( + content, d.get("attachments") + ) + await self._emit_inbound( + d, msg_id, timestamp, text, image_urls, image_media_types, + chat_id=channel_id, qq_chat_type="guild", + user_id=str(author.get("id", "")), user_name=nick or None, chat_type="group", ) - await self.handle_message(event) async def _handle_dm_message( self, @@ -1490,9 +1060,7 @@ class QQAdapter(BasePlatformAdapter): if not guild_id: return - # Apply dm_policy ACL — guild DMs were previously unauthenticated. - # Without this check any member of any guild the bot is in could - # bypass the configured allowlist via direct messages. + # dm_policy ACL — without it any guild member could bypass the allowlist via DM. author_id = str(author.get("id", "")) if not self._is_dm_intake_allowed(author_id): logger.debug( @@ -1501,26 +1069,58 @@ class QQAdapter(BasePlatformAdapter): ) return - text = content - att_result = await self._process_attachments(d.get("attachments")) - image_urls = att_result["image_urls"] - image_media_types = att_result["image_media_types"] - voice_transcripts = att_result["voice_transcripts"] - attachment_info = att_result["attachment_info"] + text, image_urls, image_media_types, _ = await self._absorb_attachments( + content, d.get("attachments") + ) + await self._emit_inbound( + d, msg_id, timestamp, text, image_urls, image_media_types, + chat_id=guild_id, qq_chat_type="dm", user_id=str(author.get("id", "")), chat_type="dm", + ) + _INBOUND_HANDLERS = { + "C2C_MESSAGE_CREATE": "_handle_c2c_message", + "GROUP_AT_MESSAGE_CREATE": "_handle_group_message", + "GUILD_MESSAGE_CREATE": "_handle_guild_message", + "GUILD_AT_MESSAGE_CREATE": "_handle_guild_message", + "DIRECT_MESSAGE_CREATE": "_handle_dm_message", + } + + # ------------------------------------------------------------------ + # Shared inbound pipeline (all four message kinds) + # ------------------------------------------------------------------ + + @staticmethod + def _append_block(text: str, block: str) -> str: + """Append *block* to *text* after a blank line (or return block alone if text is blank).""" + return (text + "\n\n" + block).strip() if text.strip() else block + + async def _absorb_attachments( + self, text: str, attachments: Any, + ) -> Tuple[str, List[str], List[str], int]: + """Run attachments through _process_attachments and fold transcripts/file + info into *text*. Returns (text, image_urls, image_media_types, n_voice).""" + att = await self._process_attachments(attachments) + voice_transcripts = att["voice_transcripts"] if voice_transcripts: - voice_block = "\n".join(voice_transcripts) - text = ( - (text + "\n\n" + voice_block).strip() if text.strip() else voice_block - ) - if attachment_info: - text = ( - (text + "\n\n" + attachment_info).strip() - if text.strip() - else attachment_info - ) + text = self._append_block(text, "\n".join(voice_transcripts)) + if att["attachment_info"]: + text = self._append_block(text, att["attachment_info"]) + return text, att["image_urls"], att["image_media_types"], len(voice_transcripts) - # Merge any quoted-message context (message_type=103 → msg_elements[0]). + async def _emit_inbound( + self, + d: Dict[str, Any], + msg_id: str, + timestamp: str, + text: str, + image_urls: List[str], + image_media_types: List[str], + *, + chat_id: str, + qq_chat_type: str, + **source_kwargs: Any, + ) -> None: + """Merge quoted context, drop empty events, remember the QQ chat kind and dispatch.""" quoted = await self._process_quoted_context(d) text = self._merge_quote_into(text, quoted["quote_block"]) if quoted["image_urls"]: @@ -1530,13 +1130,9 @@ class QQAdapter(BasePlatformAdapter): if not text.strip() and not image_urls: return - self._chat_type_map[guild_id] = "dm" + self._chat_type_map[chat_id] = qq_chat_type event = MessageEvent( - source=self.build_source( - chat_id=guild_id, - user_id=str(author.get("id", "")), - chat_type="dm", - ), + source=self.build_source(chat_id=chat_id, **source_kwargs), text=text, message_type=self._detect_message_type(image_urls, image_media_types), raw_message=d, @@ -1557,37 +1153,15 @@ class QQAdapter(BasePlatformAdapter): ) -> Dict[str, Any]: """Process the quoted message a user is replying to. - When a user replies while quoting another message, the platform sets - ``message_type = 103`` and pushes the referenced message's content and - attachments inside ``msg_elements[0]``. The old adapter ignored - ``msg_elements`` entirely, so: + A quote-reply has ``message_type == 103`` with the referenced message's + content + attachments in ``msg_elements`` (normally just [0]). Quoted + attachments run through the same _process_attachments pipeline, so + quoted voice gets STT and quoted images are cached identically. - - Quoted text was surfaced only when the user typed something of - their own — bare quote-replies showed nothing. - - Quoted attachments (images, voice, files) were never downloaded - or described. - - Quoted voice messages specifically produced no transcript, so the - LLM had no way to see what the user was referring to. - - This method parses ``msg_elements`` and runs the quoted attachments - through the same :meth:`_process_attachments` pipeline as the main - message body, so quoted voice messages get STT transcripts and - quoted images are cached identically. - - :param d: Raw inbound message dict (from the WS dispatch payload). - :returns: Dict with keys: - - - ``quote_block``: string to prepend to the user's text body - (empty when there's nothing quoted). - - ``image_urls``: list of cached quoted-image paths. - - ``image_media_types``: parallel list of image MIME types. + Returns ``{"quote_block", "image_urls", "image_media_types"}``; + quote_block is "" when nothing is quoted. """ - empty = { - "quote_block": "", - "image_urls": [], - "image_media_types": [], - } - # Short-circuit: only message_type 103 indicates a quote. + empty = {"quote_block": "", "image_urls": [], "image_media_types": []} try: if int(d.get("message_type", 0) or 0) != 103: return empty @@ -1598,9 +1172,6 @@ class QQAdapter(BasePlatformAdapter): if not isinstance(elements, list) or not elements: return empty - # msg_elements[0] carries the referenced message. Additional elements - # (if any) are very rare in practice; we concatenate their text and - # union their attachments for completeness. quoted_text_parts: List[str] = [] all_attachments: List[Dict[str, Any]] = [] for elem in elements: @@ -1616,33 +1187,25 @@ class QQAdapter(BasePlatformAdapter): all_attachments.append(a) att_result = await self._process_attachments(all_attachments) - quoted_voice = att_result.get("voice_transcripts") or [] - quoted_info = att_result.get("attachment_info") or "" quoted_images = att_result.get("image_urls") or [] - quoted_image_types = att_result.get("image_media_types") or [] lines: List[str] = [] if quoted_text_parts: lines.append(" ".join(quoted_text_parts)) - for t in quoted_voice: - lines.append(t) - if quoted_info: - lines.append(quoted_info) + lines.extend(att_result.get("voice_transcripts") or []) + if att_result.get("attachment_info"): + lines.append(att_result["attachment_info"]) if not lines and not quoted_images: return empty - - if lines: - quote_block = "[Quoted message]:\n" + "\n".join(lines) - else: - # Images-only quote: give the LLM at least a marker so it knows - # context was referenced. - quote_block = "[Quoted message]: (image)" - + # Images-only quote still gets a marker so the LLM knows context was referenced. + quote_block = ( + "[Quoted message]:\n" + "\n".join(lines) if lines else "[Quoted message]: (image)" + ) return { "quote_block": quote_block, "image_urls": quoted_images, - "image_media_types": quoted_image_types, + "image_media_types": att_result.get("image_media_types") or [], } @staticmethod @@ -1665,89 +1228,57 @@ class QQAdapter(BasePlatformAdapter): return MessageType.TEXT if not media_types: return MessageType.PHOTO - first_type = media_types[0].lower() if media_types else "" + first_type = media_types[0].lower() if "audio" in first_type or "voice" in first_type or "silk" in first_type: return MessageType.VOICE if "video" in first_type: return MessageType.VIDEO if "image" in first_type or "photo" in first_type: return MessageType.PHOTO - logger.debug( - "Unknown media content_type '%s', defaulting to TEXT", - first_type, - ) + logger.debug("Unknown media content_type '%s', defaulting to TEXT", first_type) return MessageType.TEXT async def _process_attachments( self, attachments: Any, ) -> Dict[str, Any]: - """Process inbound attachments (all message types). + """Process inbound attachments uniformly (images, voice, other files). - Mirrors OpenClaw's ``processAttachments`` — handles images, voice, and - other files uniformly. - - Returns a dict with: - - image_urls: list[str] — cached local image paths - - image_media_types: list[str] — MIME types of cached images - - voice_transcripts: list[str] — STT transcripts for voice messages - - attachment_info: str — text description of non-image, non-voice attachments + Returns ``{"image_urls", "image_media_types", "voice_transcripts", "attachment_info"}`` + (cached image paths + MIME types, "[Voice] ..." transcripts, and a text + description of non-image/non-voice files). """ - if not isinstance(attachments, list): - return { - "image_urls": [], - "image_media_types": [], - "voice_transcripts": [], - "attachment_info": "", - } - image_urls: List[str] = [] image_media_types: List[str] = [] voice_transcripts: List[str] = [] other_attachments: List[str] = [] - for att in attachments: + for att in attachments if isinstance(attachments, list) else (): if not isinstance(att, dict): continue ct = str(att.get("content_type", "")).strip().lower() - url_raw = str(att.get("url", "")).strip() + url = str(att.get("url", "")).strip() filename = str(att.get("filename", "")) - if url_raw.startswith("//"): - url = f"https:{url_raw}" - elif url_raw: - url = url_raw - else: - url = "" + if not url: continue + if url.startswith("//"): + url = f"https:{url}" logger.debug( "[%s] Processing attachment: content_type=%s, url=%s, filename=%s", - self._log_tag, - ct, - url[:80], - filename, + self._log_tag, ct, url[:80], filename, ) if self._is_voice_content_type(ct, filename): - # Voice: use QQ's asr_refer_text first, then voice_wav_url, then STT. - asr_refer = ( - str(att.get("asr_refer_text", "")).strip() - if isinstance(att.get("asr_refer_text"), str) - else "" - ) - voice_wav_url = ( - str(att.get("voice_wav_url", "")).strip() - if isinstance(att.get("voice_wav_url"), str) - else "" - ) - + asr_refer = att.get("asr_refer_text") + voice_wav_url = att.get("voice_wav_url") transcript = await self._stt_voice_attachment( url, ct, filename, - asr_refer_text=asr_refer or None, - voice_wav_url=voice_wav_url or None, + asr_refer_text=(asr_refer.strip() if isinstance(asr_refer, str) else "") or None, + voice_wav_url=(voice_wav_url.strip() if isinstance(voice_wav_url, str) else "") or None, ) if transcript: voice_transcripts.append(f"[Voice] {transcript}") @@ -1756,7 +1287,6 @@ class QQAdapter(BasePlatformAdapter): logger.warning("[%s] Voice STT failed for %s", self._log_tag, url[:60]) voice_transcripts.append("[Voice] [语音识别失败]") elif ct.startswith("image/"): - # Image: download and cache locally. try: cached_path = await self._download_and_cache(url, ct, filename) if cached_path and os.path.isfile(cached_path): @@ -1764,112 +1294,71 @@ class QQAdapter(BasePlatformAdapter): image_media_types.append(ct or "image/jpeg") elif cached_path: logger.warning( - "[%s] Cached image path does not exist: %s", - self._log_tag, - cached_path, + "[%s] Cached image path does not exist: %s", self._log_tag, cached_path ) except Exception as exc: logger.debug("[%s] Failed to cache image: %s", self._log_tag, exc) else: - # Other attachments (video, file, etc.): download and record with path. try: cached_path = await self._download_and_cache(url, ct, filename) if cached_path: - name = filename or ct - if ct.startswith("video/"): - other_attachments.append(f"[video: {name} ({cached_path})]") - else: - other_attachments.append(f"[file: {name} ({cached_path})]") + label = "video" if ct.startswith("video/") else "file" + other_attachments.append(f"[{label}: {filename or ct} ({cached_path})]") except Exception as exc: logger.debug("[%s] Failed to cache attachment: %s", self._log_tag, exc) - attachment_info = "\n".join(other_attachments) if other_attachments else "" return { "image_urls": image_urls, "image_media_types": image_media_types, "voice_transcripts": voice_transcripts, - "attachment_info": attachment_info, + "attachment_info": "\n".join(other_attachments), } async def _download_and_cache( self, url: str, content_type: str, original_name: str = "", ) -> Optional[str]: - """Download a URL and cache it locally. - - :param original_name: Preferred filename from attachment metadata. - Falls back to the URL path basename if empty. - """ + """Download a URL and cache it locally (``original_name`` falls back to the URL basename).""" from tools.url_safety import is_safe_url if not is_safe_url(url): raise ValueError(f"Blocked unsafe URL: {url[:80]}") - if not self._http_client: return None try: - resp = await self._http_client.get( - url, - timeout=30.0, - headers=self._qq_media_headers(), - ) + resp = await self._http_client.get(url, timeout=30.0, headers=self._qq_media_headers()) resp.raise_for_status() data = resp.content except Exception as exc: - logger.debug( - "[%s] Download failed for %s: %s", self._log_tag, url[:80], exc - ) + logger.debug("[%s] Download failed for %s: %s", self._log_tag, url[:80], exc) return None if content_type.startswith("image/"): - # preserves historical qqbot mapping: trust mimetypes' - # guess (never the shared table) and fall back to .jpg. + # Historical qqbot mapping: trust mimetypes' guess (never the shared + # table) and fall back to .jpg. ext = ext_for_mime( - content_type, - use_defaults=False, - use_mimetypes=True, - fallback=".jpg", + content_type, use_defaults=False, use_mimetypes=True, fallback=".jpg" ) or ".jpg" return cache_image_from_bytes(data, ext) - elif content_type == "voice" or content_type.startswith("audio/"): - # QQ voice messages are typically .amr or .silk format. - # Convert to .wav using ffmpeg so STT engines can process it. + if content_type == "voice" or content_type.startswith("audio/"): + # QQ voice is usually .amr/.silk — convert to .wav for STT engines. return await self._convert_audio_to_wav(data, url) - else: - filename = ( - original_name - or Path(urlparse(url).path).name - or "qq_attachment" - ) - return cache_document_from_bytes(data, filename) + filename = original_name or Path(urlparse(url).path).name or "qq_attachment" + return cache_document_from_bytes(data, filename) @staticmethod def _is_voice_content_type(content_type: str, filename: str) -> bool: """Check if an attachment is a voice/audio message.""" ct = content_type.strip().lower() - fn = filename.strip().lower() if ct == "voice" or ct.startswith("audio/"): return True - # QQ file uploads have content_type="file". Without this guard, - # any uploaded audio file (e.g. .wav, .mp3) would be misrouted into - # the STT pipeline and never be received as a normal file attachment. + # content_type="file" is an explicit upload: never route .wav/.mp3 files into STT. if ct == "file": return False - _VOICE_EXTENSIONS = ( - ".silk", ".amr", ".mp3", ".wav", ".ogg", - ".m4a", ".aac", ".speex", ".flac", - ) - if any(fn.endswith(ext) for ext in _VOICE_EXTENSIONS): - return True - return False + return filename.strip().lower().endswith(_VOICE_EXTENSIONS) def _qq_media_headers(self) -> Dict[str, str]: - """Return Authorization headers for QQ multimedia CDN downloads. - - QQ's multimedia URLs (multimedia.nt.qq.com.cn) require the bot's - access token in an Authorization header, otherwise the download - returns a non-200 status. - """ + """Authorization header for QQ multimedia CDN downloads (required, else non-200).""" if self._access_token: return {"Authorization": f"QQBot {self._access_token}"} return {} @@ -1883,23 +1372,13 @@ class QQAdapter(BasePlatformAdapter): asr_refer_text: Optional[str] = None, voice_wav_url: Optional[str] = None, ) -> Optional[str]: - """Download a voice attachment, convert to wav, and transcribe. - - Priority: - 1. QQ's built-in ``asr_refer_text`` (Tencent's own ASR — free, no API call). - 2. Self-hosted STT on ``voice_wav_url`` (pre-converted WAV from QQ, avoids SILK decoding). - 3. Self-hosted STT on the original attachment URL (requires SILK→WAV conversion). - - Returns the transcript text, or None on failure. - """ - # 1. Use QQ's built-in ASR text if available + """Transcribe a voice attachment. Priority: QQ's free ``asr_refer_text`` → + STT on ``voice_wav_url`` (pre-converted WAV, no SILK decode) → STT on the + original URL (SILK→WAV). Returns the transcript or None.""" if asr_refer_text: - logger.debug( - "[%s] STT: using QQ asr_refer_text: %r", self._log_tag, asr_refer_text[:100] - ) + logger.debug("[%s] STT: using QQ asr_refer_text: %r", self._log_tag, asr_refer_text[:100]) return asr_refer_text - # Determine which URL to download (prefer voice_wav_url — already WAV) download_url = url is_pre_wav = False if voice_wav_url: @@ -1915,74 +1394,49 @@ class QQAdapter(BasePlatformAdapter): return None try: - # 2. Download audio (QQ CDN requires Authorization header) if not self._http_client: logger.warning("[%s] STT: no HTTP client", self._log_tag) return None - download_headers = self._qq_media_headers() + download_headers = self._qq_media_headers() # QQ CDN requires Authorization logger.debug( "[%s] STT: downloading voice from %s (pre_wav=%s, headers=%s)", - self._log_tag, - download_url[:80], - is_pre_wav, - bool(download_headers), + self._log_tag, download_url[:80], is_pre_wav, bool(download_headers), ) resp = await self._http_client.get( - download_url, - timeout=30.0, - headers=download_headers, - follow_redirects=True, + download_url, timeout=30.0, headers=download_headers, follow_redirects=True ) resp.raise_for_status() audio_data = resp.content logger.debug( "[%s] STT: downloaded %d bytes, content_type=%s", - self._log_tag, - len(audio_data), - resp.headers.get("content-type", "unknown"), + self._log_tag, len(audio_data), resp.headers.get("content-type", "unknown"), ) - if len(audio_data) < 10: logger.warning( "[%s] STT: downloaded data too small (%d bytes), skipping", - self._log_tag, - len(audio_data), + self._log_tag, len(audio_data), ) return None - # 3. Convert to wav (skip if we already have a pre-converted WAV) if is_pre_wav: - import tempfile - - with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: - tmp.write(audio_data) - wav_path = tmp.name + wav_path = self._write_temp(audio_data, ".wav") logger.debug( "[%s] STT: using pre-converted WAV directly (%d bytes)", - self._log_tag, - len(audio_data), + self._log_tag, len(audio_data), ) else: - logger.debug( - "[%s] STT: converting to wav, filename=%r", self._log_tag, filename - ) + logger.debug("[%s] STT: converting to wav, filename=%r", self._log_tag, filename) wav_path = await self._convert_audio_to_wav_file(audio_data, filename) if not wav_path or not Path(wav_path).exists(): - logger.warning( - "[%s] STT: ffmpeg conversion produced no output", self._log_tag - ) + logger.warning("[%s] STT: ffmpeg conversion produced no output", self._log_tag) return None - # 4. Call STT API and always clean up the temp WAV afterward. logger.debug("[%s] STT: calling ASR on %s", self._log_tag, wav_path) try: transcript = await self._call_stt(wav_path) finally: - try: - os.unlink(wav_path) - except OSError: - pass + self._unlink_quiet(wav_path) if transcript: logger.debug("[%s] STT success: %r", self._log_tag, transcript[:100]) @@ -1992,68 +1446,56 @@ class QQAdapter(BasePlatformAdapter): except (httpx.HTTPStatusError, httpx.TransportError, IOError) as exc: logger.warning( "[%s] STT failed for voice attachment: %s: %s", - self._log_tag, - type(exc).__name__, - exc, + self._log_tag, type(exc).__name__, exc, ) return None + @staticmethod + def _write_temp(data: bytes, suffix: str) -> str: + """Write *data* to a persistent NamedTemporaryFile and return its path.""" + import tempfile + + with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp: + tmp.write(data) + return tmp.name + + @staticmethod + def _unlink_quiet(path: str) -> None: + try: + os.unlink(path) + except OSError: + pass + + @staticmethod + def _wav_ok(wav_path: str) -> bool: + """True when *wav_path* exists and holds more than a bare 44-byte header.""" + return Path(wav_path).exists() and Path(wav_path).stat().st_size > 44 + async def _convert_audio_to_wav_file( self, audio_data: bytes, filename: str ) -> Optional[str]: - """Convert audio bytes to a temp .wav file using pilk (SILK) or ffmpeg. - - QQ voice messages are typically SILK format which ffmpeg cannot decode. - Strategy: always try pilk first, fall back to ffmpeg if pilk fails. - - Returns the wav file path, or None on failure. - """ - import tempfile - - ext = ( - Path(filename).suffix.lower() - if Path(filename).suffix - else self._guess_ext_from_data(audio_data) - ) + """Convert audio bytes to a temp .wav: pilk (SILK, which ffmpeg can't decode) + → ffmpeg → raw-PCM last resort. Returns the wav path or None.""" + ext = Path(filename).suffix.lower() or self._guess_ext_from_data(audio_data) logger.info( "[%s] STT: audio_data size=%d, ext=%r, first_20_bytes=%r", - self._log_tag, - len(audio_data), - ext, - audio_data[:20], + self._log_tag, len(audio_data), ext, audio_data[:20], ) - - with tempfile.NamedTemporaryFile(suffix=ext, delete=False) as tmp_src: - tmp_src.write(audio_data) - src_path = tmp_src.name - + src_path = self._write_temp(audio_data, ext) wav_path = src_path.rsplit(".", 1)[0] + ".wav" - # Try pilk first (handles SILK and many other formats) - result = await self._convert_silk_to_wav(src_path, wav_path) - - # If pilk failed, try ffmpeg - if not result: - result = await self._convert_ffmpeg_to_wav(src_path, wav_path) - - # If ffmpeg also failed, try writing raw PCM as WAV (last resort) - if not result: - result = await self._convert_raw_to_wav(audio_data, wav_path) - - # Cleanup source file - try: - os.unlink(src_path) - except OSError: - pass - + result = ( + await self._convert_silk_to_wav(src_path, wav_path) + or await self._convert_ffmpeg_to_wav(src_path, wav_path) + or await self._convert_raw_to_wav(audio_data, wav_path) + ) + self._unlink_quiet(src_path) return result @staticmethod def _guess_ext_from_data(data: bytes) -> str: - """Guess file extension from magic bytes.""" - if data[:9] == b"#!SILK_V3" or data[:6] == b"#!SILK": - return ".silk" - if data[:2] == b"\x02!": + """Guess file extension from magic bytes (unknown → .amr, QQ's most common).""" + if data[:6] == b"#!SILK" or data[:2] == b"\x02!": return ".silk" if data[:4] == b"RIFF": return ".wav" @@ -2063,22 +1505,15 @@ class QQAdapter(BasePlatformAdapter): return ".mp3" if data[:4] == b"\x30\x26\xb2\x75" or data[:4] == b"\x4f\x67\x67\x53": return ".ogg" - if data[:4] == b"\x00\x00\x00\x20" or data[:4] == b"\x00\x00\x00\x1c": - return ".amr" - # Default to .amr for unknown (QQ's most common voice format) return ".amr" @staticmethod def _looks_like_silk(data: bytes) -> bool: """Check if bytes look like a SILK audio file.""" - return data[:6] == b"#!SILK" or data[:2] == b"\x02!" or data[:9] == b"#!SILK_V3" + return data[:6] == b"#!SILK" or data[:2] == b"\x02!" async def _convert_silk_to_wav(self, src_path: str, wav_path: str) -> Optional[str]: - """Convert audio file to WAV using the pilk library. - - Tries the file as-is first, then as .silk if the extension differs. - pilk can handle SILK files with various headers (or no header). - """ + """Convert to WAV with pilk: as-is first, then copied to .silk (pilk checks the extension).""" try: import pilk except ImportError: @@ -2088,51 +1523,38 @@ class QQAdapter(BasePlatformAdapter): ) return None - # Try converting the file as-is try: pilk.silk_to_wav(src_path, wav_path, rate=16000) - if Path(wav_path).exists() and Path(wav_path).stat().st_size > 44: + if self._wav_ok(wav_path): logger.debug( "[%s] pilk converted %s to wav (%d bytes)", - self._log_tag, - Path(src_path).name, - Path(wav_path).stat().st_size, + self._log_tag, Path(src_path).name, Path(wav_path).stat().st_size, ) return wav_path except Exception as exc: logger.debug("[%s] pilk direct conversion failed: %s", self._log_tag, exc) - # Try renaming to .silk and converting (pilk checks the extension) silk_path = src_path.rsplit(".", 1)[0] + ".silk" try: import shutil shutil.copy2(src_path, silk_path) pilk.silk_to_wav(silk_path, wav_path, rate=16000) - if Path(wav_path).exists() and Path(wav_path).stat().st_size > 44: + if self._wav_ok(wav_path): logger.debug( "[%s] pilk converted %s (as .silk) to wav (%d bytes)", - self._log_tag, - Path(src_path).name, - Path(wav_path).stat().st_size, + self._log_tag, Path(src_path).name, Path(wav_path).stat().st_size, ) return wav_path except Exception as exc: logger.debug("[%s] pilk .silk conversion failed: %s", self._log_tag, exc) finally: - try: - os.unlink(silk_path) - except OSError: - pass - + self._unlink_quiet(silk_path) return None async def _convert_raw_to_wav(self, audio_data: bytes, wav_path: str) -> Optional[str]: - """Last resort: try writing audio data as raw PCM 16-bit mono 16kHz WAV. - - This will produce garbage if the data isn't raw PCM, but at least - the ASR engine won't crash — it'll just return empty. - """ + """Last resort: wrap bytes as raw PCM 16-bit mono 16kHz WAV (garbage if not + PCM, but the ASR engine returns empty instead of crashing).""" try: import wave @@ -2150,15 +1572,7 @@ class QQAdapter(BasePlatformAdapter): """Convert audio file to WAV using ffmpeg.""" try: proc = await asyncio.create_subprocess_exec( - "ffmpeg", - "-y", - "-i", - src_path, - "-ar", - "16000", - "-ac", - "1", - wav_path, + "ffmpeg", "-y", "-i", src_path, "-ar", "16000", "-ac", "1", wav_path, stdout=asyncio.subprocess.DEVNULL, stderr=asyncio.subprocess.PIPE, ) @@ -2167,105 +1581,62 @@ class QQAdapter(BasePlatformAdapter): stderr = await proc.stderr.read() if proc.stderr else b"" logger.warning( "[%s] ffmpeg failed for %s: %s", - self._log_tag, - Path(src_path).name, - stderr[:200].decode(errors="replace"), + self._log_tag, Path(src_path).name, stderr[:200].decode(errors="replace"), ) return None except (asyncio.TimeoutError, FileNotFoundError) as exc: logger.warning("[%s] ffmpeg conversion error: %s", self._log_tag, exc) return None - if not Path(wav_path).exists() or Path(wav_path).stat().st_size <= 44: + if not self._wav_ok(wav_path): logger.warning( - "[%s] ffmpeg produced no/small output for %s", - self._log_tag, - Path(src_path).name, + "[%s] ffmpeg produced no/small output for %s", self._log_tag, Path(src_path).name ) return None logger.debug( "[%s] ffmpeg converted %s to wav (%d bytes)", - self._log_tag, - Path(src_path).name, - Path(wav_path).stat().st_size, + self._log_tag, Path(src_path).name, Path(wav_path).stat().st_size, ) return wav_path def _resolve_stt_config(self) -> Optional[Dict[str, str]]: - """Resolve STT backend configuration from config/environment. - - Priority: - 1. Plugin-specific: ``channels.qqbot.stt`` in config.yaml → ``self.config.extra["stt"]`` - 2. QQ-specific env vars: ``QQ_STT_API_KEY`` / ``QQ_STT_BASE_URL`` / ``QQ_STT_MODEL`` - 3. Return None if nothing is configured (STT will be skipped, QQ built-in ASR still works). - """ - extra = self.config.extra or {} - - # 1. Plugin-specific STT config (matches OpenClaw's channels.qqbot.stt) - stt_cfg = extra.get("stt") + """Resolve STT backend: ``extra["stt"]`` config first, then ``QQ_STT_*`` env + vars; None when unconfigured (QQ's built-in ASR still works).""" + stt_cfg = (self.config.extra or {}).get("stt") if isinstance(stt_cfg, dict) and stt_cfg.get("enabled") is not False: base_url = stt_cfg.get("baseUrl") or stt_cfg.get("base_url", "") api_key = stt_cfg.get("apiKey") or stt_cfg.get("api_key", "") model = stt_cfg.get("model", "") if base_url and api_key: - return { - "base_url": base_url.rstrip("/"), - "api_key": api_key, - "model": model or "whisper-1", - } - # Provider-only config: just model name, use default provider - if api_key: + return {"base_url": base_url.rstrip("/"), "api_key": api_key, "model": model or "whisper-1"} + if api_key: # provider-only config provider = stt_cfg.get("provider", "zai") - # Map provider to base URL - _PROVIDER_BASE_URLS = { - "zai": "https://open.bigmodel.cn/api/coding/paas/v4", - "openai": "https://api.openai.com/v1", - "glm": "https://open.bigmodel.cn/api/coding/paas/v4", - } - base_url = _PROVIDER_BASE_URLS.get(provider, "") + base_url = _STT_PROVIDER_BASE_URLS.get(provider, "") if base_url: return { "base_url": base_url, "api_key": api_key, - "model": model - or ("glm-asr" if provider in {"zai", "glm"} else "whisper-1"), + "model": model or ("glm-asr" if provider in {"zai", "glm"} else "whisper-1"), } - # 2. QQ-specific env vars (set by `hermes setup gateway` / `hermes gateway`) qq_stt_key = _resolve_qq_secret("QQ_STT_API_KEY", "") if qq_stt_key: - base_url = _resolve_qq_secret( - "QQ_STT_BASE_URL", - "https://open.bigmodel.cn/api/coding/paas/v4", - ) - model = _resolve_qq_secret("QQ_STT_MODEL", "glm-asr") + base_url = _resolve_qq_secret("QQ_STT_BASE_URL", _STT_PROVIDER_BASE_URLS["zai"]) return { "base_url": base_url.rstrip("/"), "api_key": qq_stt_key, - "model": model, + "model": _resolve_qq_secret("QQ_STT_MODEL", "glm-asr"), } - return None async def _call_stt(self, wav_path: str) -> Optional[str]: - """Call an OpenAI-compatible STT API to transcribe a wav file. - - Uses the provider configured in ``channels.qqbot.stt`` config, - falling back to QQ's built-in ``asr_refer_text`` if not configured. - Returns None if STT is not configured or the call fails. - """ + """Transcribe a wav via an OpenAI-compatible STT API; None if unconfigured/failed.""" stt_cfg = self._resolve_stt_config() if not stt_cfg: - logger.warning( - "[%s] STT not configured (no stt config or QQ_STT_API_KEY)", - self._log_tag, - ) + logger.warning("[%s] STT not configured (no stt config or QQ_STT_API_KEY)", self._log_tag) return None - base_url = stt_cfg["base_url"] - api_key = stt_cfg["api_key"] - model = stt_cfg["model"] - + base_url, api_key, model = stt_cfg["base_url"], stt_cfg["api_key"], stt_cfg["model"] try: with open(wav_path, "rb") as f: resp = await self._http_client.post( @@ -2277,80 +1648,48 @@ class QQAdapter(BasePlatformAdapter): ) resp.raise_for_status() result = resp.json() - # Zhipu/GLM format: {"choices": [{"message": {"content": "transcript text"}}]} + # Zhipu/GLM: {"choices": [{"message": {"content": ...}}]}; OpenAI/Whisper: {"text": ...} choices = result.get("choices", []) if choices: content = choices[0].get("message", {}).get("content", "") if content.strip(): return content.strip() - # OpenAI/Whisper format: {"text": "transcript text"} text = result.get("text", "") - if text.strip(): - return text.strip() - return None + return text.strip() or None except (httpx.HTTPStatusError, IOError) as exc: logger.warning( "[%s] STT API call failed (model=%s, base=%s): %s", - self._log_tag, - model, - base_url[:50], - exc, + self._log_tag, model, base_url[:50], exc, ) return None async def _convert_audio_to_wav( self, audio_data: bytes, source_url: str ) -> Optional[str]: - """Convert audio bytes to .wav using pilk (SILK) or ffmpeg, caching the result.""" - import tempfile - - # Determine source format from magic bytes or URL - ext = ( - Path(urlparse(source_url).path).suffix.lower() - if urlparse(source_url).path - else "" - ) - if not ext or ext not in { - ".silk", - ".amr", - ".mp3", - ".wav", - ".ogg", - ".m4a", - ".aac", - ".flac", - }: + """Convert audio bytes to .wav (pilk for SILK, else ffmpeg) and cache the result; + on conversion failure the original bytes are cached as ``qq_voice``.""" + ext = Path(urlparse(source_url).path).suffix.lower() + if ext not in _AUDIO_URL_EXTENSIONS: ext = self._guess_ext_from_data(audio_data) - with tempfile.NamedTemporaryFile(suffix=ext, delete=False) as tmp_src: - tmp_src.write(audio_data) - src_path = tmp_src.name - + src_path = self._write_temp(audio_data, ext) wav_path = src_path.rsplit(".", 1)[0] + ".wav" try: - is_silk = ext == ".silk" or self._looks_like_silk(audio_data) - if is_silk: + if ext == ".silk" or self._looks_like_silk(audio_data): result = await self._convert_silk_to_wav(src_path, wav_path) else: result = await self._convert_ffmpeg_to_wav(src_path, wav_path) - if not result: logger.warning( "[%s] audio conversion failed for %s (format=%s)", - self._log_tag, - source_url[:60], - ext, + self._log_tag, source_url[:60], ext, ) return cache_document_from_bytes(audio_data, f"qq_voice{ext}") except Exception: return cache_document_from_bytes(audio_data, f"qq_voice{ext}") finally: - try: - os.unlink(src_path) - except OSError: - pass + self._unlink_quiet(src_path) - # Verify output and cache try: wav_data = Path(wav_path).read_bytes() os.unlink(wav_path) @@ -2373,32 +1712,29 @@ class QQAdapter(BasePlatformAdapter): """Make an authenticated REST API request to QQ Bot API.""" if not self._http_client: raise RuntimeError("HTTP client not initialized — not connected?") - - token = await self._ensure_token() - headers = { - "Authorization": f"QQBot {token}", - "Content-Type": "application/json", - "User-Agent": build_user_agent(), - } - + headers = await self._auth_headers() try: resp = await self._http_client.request( - method, - f"{API_BASE}{path}", - headers=headers, - json=body, - timeout=timeout, + method, f"{API_BASE}{path}", headers=headers, json=body, timeout=timeout ) data = resp.json() if resp.status_code >= 400: raise RuntimeError( - f"QQ Bot API error [{resp.status_code}] {path}: " - f"{data.get('message', data)}" + f"QQ Bot API error [{resp.status_code}] {path}: {data.get('message', data)}" ) return data except httpx.TimeoutException as exc: raise RuntimeError(f"QQ Bot API timeout [{path}]: {exc}") from exc + async def _auth_headers(self) -> Dict[str, str]: + """JSON REST headers with a fresh bot token.""" + token = await self._ensure_token() + return { + "Authorization": f"QQBot {token}", + "Content-Type": "application/json", + "User-Agent": build_user_agent(), + } + async def _upload_media( self, target_type: str, @@ -2410,16 +1746,9 @@ class QQAdapter(BasePlatformAdapter): file_name: Optional[str] = None, ) -> Dict[str, Any]: """Upload media and return file_info.""" - path = ( - f"/v2/users/{target_id}/files" - if target_type == "c2c" - else f"/v2/groups/{target_id}/files" - ) - - body: Dict[str, Any] = { - "file_type": file_type, - "srv_send_msg": srv_send_msg, - } + kind = "users" if target_type == "c2c" else "groups" + path = f"/v2/{kind}/{target_id}/files" + body: Dict[str, Any] = {"file_type": file_type, "srv_send_msg": srv_send_msg} if url: body["url"] = url elif file_data: @@ -2427,18 +1756,12 @@ class QQAdapter(BasePlatformAdapter): if file_type == MEDIA_TYPE_FILE and file_name: body["file_name"] = file_name - # Retry transient upload failures - for attempt in range(3): + for attempt in range(3): # retry transient upload failures try: - return await self._api_request( - "POST", path, body, timeout=FILE_UPLOAD_TIMEOUT - ) + return await self._api_request("POST", path, body, timeout=FILE_UPLOAD_TIMEOUT) except RuntimeError as exc: err_msg = str(exc) - if any( - kw in err_msg - for kw in ("400", "401", "Invalid", "timeout", "Timeout") - ): + if any(kw in err_msg for kw in ("400", "401", "Invalid", "timeout", "Timeout")): raise if attempt < 2: await asyncio.sleep(1.5 * (attempt + 1)) @@ -2451,15 +1774,8 @@ class QQAdapter(BasePlatformAdapter): _RECONNECT_POLL_INTERVAL = 0.5 async def _wait_for_reconnection(self) -> bool: - """Wait for the WebSocket listener to reconnect. - - The listener loop (_listen_loop) auto-reconnects on disconnect, but - there is a race window where send() is called right after a disconnect - and before the reconnect completes. This method polls is_connected - for up to _RECONNECT_WAIT_SECONDS. - - Returns True if reconnected, False if still disconnected. - """ + """Poll is_connected for up to _RECONNECT_WAIT_SECONDS — covers the race where + send() lands between a disconnect and _listen_loop's reconnect.""" logger.info("[%s] Not connected — waiting for reconnection (up to %.0fs)", self._log_tag, self._RECONNECT_WAIT_SECONDS) waited = 0.0 @@ -2472,6 +1788,14 @@ class QQAdapter(BasePlatformAdapter): logger.warning("[%s] Still not connected after %.0fs", self._log_tag, self._RECONNECT_WAIT_SECONDS) return False + @property + def _NOT_CONNECTED(self) -> SendResult: + return SendResult(success=False, error="Not connected", retryable=True) + + async def _ensure_connected(self) -> bool: + """True when connected now or after waiting for the listener to reconnect.""" + return self.is_connected or await self._wait_for_reconnection() + async def send( self, chat_id: str, @@ -2479,17 +1803,10 @@ class QQAdapter(BasePlatformAdapter): reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: - """Send a text or markdown message to a QQ user or group. - - Applies format_message(), splits long messages via truncate_message(), - and retries transient failures with exponential backoff. - """ + """Send text/markdown: format, split via truncate_message(), retry transient failures.""" del metadata - - if not self.is_connected: - if not await self._wait_for_reconnection(): - return SendResult(success=False, error="Not connected", retryable=True) - + if not await self._ensure_connected(): + return self._NOT_CONNECTED if not content or not content.strip(): return SendResult(success=True) @@ -2524,27 +1841,17 @@ class QQAdapter(BasePlatformAdapter): elif chat_type == "guild": return await self._send_guild_text(chat_id, content, reply_to) else: - return SendResult( - success=False, error=f"Unknown chat type for {chat_id}" - ) + return SendResult(success=False, error=f"Unknown chat type for {chat_id}") except Exception as exc: last_exc = exc err = str(exc).lower() - # Permanent errors — don't retry - if any( - k in err - for k in ("invalid", "forbidden", "not found", "bad request") - ): - break - # Transient — back off and retry + if any(k in err for k in ("invalid", "forbidden", "not found", "bad request")): + break # permanent — don't retry if attempt < 2: delay = 1.0 * (2 ** attempt) logger.warning( "[%s] send retry %d/3 after %.1fs: %s", - self._log_tag, - attempt + 1, - delay, - exc, + self._log_tag, attempt + 1, delay, exc, ) await asyncio.sleep(delay) @@ -2555,51 +1862,41 @@ class QQAdapter(BasePlatformAdapter): ) return SendResult(success=False, error=error_msg, retryable=retryable) - async def _send_c2c_text( - self, - openid: str, - content: str, - reply_to: Optional[str] = None, - keyboard: Optional[InlineKeyboard] = None, - ) -> SendResult: - """Send text to a C2C user via REST API. + @staticmethod + def _messages_path(chat_type: str, target_id: str) -> str: + """REST path for outbound messages to a c2c user or a group.""" + kind = "users" if chat_type == "c2c" else "groups" + return f"/v2/{kind}/{target_id}/messages" - :param keyboard: Optional inline keyboard attached to the message. - """ - self._next_msg_seq(reply_to or openid) - body = self._build_text_body(content, reply_to) - if reply_to: - body["msg_id"] = reply_to - if keyboard is not None: - body["keyboard"] = keyboard.to_dict() - - data = await self._api_request("POST", f"/v2/users/{openid}/messages", body) - msg_id = str(data.get("id", uuid.uuid4().hex[:12])) - return SendResult(success=True, message_id=msg_id, raw_response=data) - - async def _send_group_text( - self, - group_openid: str, - content: str, - reply_to: Optional[str] = None, - keyboard: Optional[InlineKeyboard] = None, - ) -> SendResult: - """Send text to a group via REST API. - - :param keyboard: Optional inline keyboard attached to the message. - """ - self._next_msg_seq(reply_to or group_openid) - body = self._build_text_body(content, reply_to) - if reply_to: - body["msg_id"] = reply_to - if keyboard is not None: - body["keyboard"] = keyboard.to_dict() - - data = await self._api_request( - "POST", f"/v2/groups/{group_openid}/messages", body + async def _post_message(self, path: str, body: Dict[str, Any]) -> SendResult: + """POST a message body and wrap the response as a successful SendResult.""" + data = await self._api_request("POST", path, body) + return SendResult( + success=True, message_id=str(data.get("id", uuid.uuid4().hex[:12])), raw_response=data ) - msg_id = str(data.get("id", uuid.uuid4().hex[:12])) - return SendResult(success=True, message_id=msg_id, raw_response=data) + + async def _send_text_to( + self, + chat_type: str, + target_id: str, + content: str, + reply_to: Optional[str] = None, + keyboard: Optional[InlineKeyboard] = None, + ) -> SendResult: + """Send text (optionally with an inline keyboard) to a c2c user or group.""" + self._next_msg_seq(reply_to or target_id) + body = self._build_text_body(content, reply_to) + if reply_to: + body["msg_id"] = reply_to + if keyboard is not None: + body["keyboard"] = keyboard.to_dict() + return await self._post_message(self._messages_path(chat_type, target_id), body) + + async def _send_c2c_text(self, openid, content, reply_to=None, keyboard=None) -> SendResult: + return await self._send_text_to("c2c", openid, content, reply_to, keyboard) + + async def _send_group_text(self, group_openid, content, reply_to=None, keyboard=None) -> SendResult: + return await self._send_text_to("group", group_openid, content, reply_to, keyboard) async def _send_guild_text( self, channel_id: str, content: str, reply_to: Optional[str] = None @@ -2608,10 +1905,7 @@ class QQAdapter(BasePlatformAdapter): body: Dict[str, Any] = {"content": content[: self.MAX_MESSAGE_LENGTH]} if reply_to: body["msg_id"] = reply_to - - data = await self._api_request("POST", f"/channels/{channel_id}/messages", body) - msg_id = str(data.get("id", uuid.uuid4().hex[:12])) - return SendResult(success=True, message_id=msg_id, raw_response=data) + return await self._post_message(f"/channels/{channel_id}/messages", body) # ------------------------------------------------------------------ # Inline-keyboard outbound helpers (approval / update-prompt flows) @@ -2624,46 +1918,24 @@ class QQAdapter(BasePlatformAdapter): keyboard: InlineKeyboard, reply_to: Optional[str] = None, ) -> SendResult: - """Send a single text message with an inline keyboard attached. - - Unlike :meth:`send`, this does NOT split long content into chunks — - a keyboard message has exactly one interactive surface, and splitting - would orphan the buttons from the first chunk. Callers should keep - approval/update-prompt bodies short. - - Guild (channel) chats don't support inline keyboards; returns a - non-retryable failure for those. - """ - if not self.is_connected: - if not await self._wait_for_reconnection(): - return SendResult( - success=False, error="Not connected", retryable=True - ) - + """Send ONE text message with an inline keyboard (no chunking — splitting + would orphan the buttons; keep bodies short). Guild chats are unsupported.""" + if not await self._ensure_connected(): + return self._NOT_CONNECTED chat_type = self._guess_chat_type(chat_id) - formatted = self.format_message(content) - truncated = formatted[: self.MAX_MESSAGE_LENGTH] + truncated = self.format_message(content)[: self.MAX_MESSAGE_LENGTH] try: if chat_type == "c2c": - return await self._send_c2c_text( - chat_id, truncated, reply_to, keyboard=keyboard, - ) + return await self._send_c2c_text(chat_id, truncated, reply_to, keyboard=keyboard) if chat_type == "group": - return await self._send_group_text( - chat_id, truncated, reply_to, keyboard=keyboard, - ) + return await self._send_group_text(chat_id, truncated, reply_to, keyboard=keyboard) return SendResult( success=False, - error=( - f"Inline keyboards not supported for chat_type " - f"{chat_type!r}" - ), + error=f"Inline keyboards not supported for chat_type {chat_type!r}", retryable=False, ) except Exception as exc: - logger.error( - "[%s] send_with_keyboard failed: %s", self._log_tag, exc - ) + logger.error("[%s] send_with_keyboard failed: %s", self._log_tag, exc) return SendResult(success=False, error=str(exc) or type(exc).__name__) async def send_approval_request( @@ -2672,34 +1944,20 @@ class QQAdapter(BasePlatformAdapter): req: ApprovalRequest, reply_to: Optional[str] = None, ) -> SendResult: - """Send a 3-button approval request (``allow-once / allow-always / deny``). - - The rendered text comes from :func:`build_approval_text`; callers can - override by passing a custom :class:`ApprovalRequest`. - - Users click the button → ``INTERACTION_CREATE`` fires → the adapter's - registered :meth:`set_interaction_callback` handler decodes - ``button_data`` via :func:`parse_approval_button_data`. - """ + """Send a 3-button approval request (allow-once / allow-always / deny); + clicks come back as INTERACTION_CREATE decoded by parse_approval_button_data.""" from gateway.platforms.qqbot.keyboards import build_approval_text return await self.send_with_keyboard( chat_id, build_approval_text(req), build_approval_keyboard( - req.session_key, - allow_permanent=getattr(req, "allow_permanent", True), + req.session_key, allow_permanent=getattr(req, "allow_permanent", True) ), reply_to=reply_to, ) - # ------------------------------------------------------------------ - # Cross-adapter gateway contract — send_exec_approval + send_update_prompt - # ------------------------------------------------------------------ - # - # These mirror the signatures that gateway/run.py detects on the adapter - # class (e.g. type(adapter).send_exec_approval, type(adapter).send_update_prompt) - # for button-based approval / update-confirm UX. Discord, Telegram, Slack, - # Matrix, and Feishu already implement the same contract. + # Cross-adapter gateway contract: gateway/run.py detects send_exec_approval / + # send_update_prompt on the adapter class for button-based approval/update UX. async def send_exec_approval( self, @@ -2712,23 +1970,13 @@ class QQAdapter(BasePlatformAdapter): allow_session: bool = True, smart_denied: bool = False, ) -> SendResult: - """Send a button-based exec-approval prompt for a dangerous command. - - Called by ``gateway/run.py``'s ``_approval_notify_sync`` when the - agent is blocked waiting for approval. Button clicks resolve via - :func:`tools.approval.resolve_gateway_approval` — dispatched by the - adapter's interaction callback (:meth:`_default_interaction_dispatch`). - """ - del metadata # QQ doesn't have thread_id / DM targeting overrides. - del allow_session # QQ's 3-button keyboard has no session tier (once/always/deny). + """Button-based exec-approval prompt (called by gateway/run.py while the + agent blocks on approval); clicks resolve via _default_interaction_dispatch.""" + del metadata # QQ has no thread_id / DM targeting overrides. + del allow_session # QQ's 3-button keyboard has no session tier. if smart_denied: description += " Owner override applies to this one operation only." - # Use the reply-to message for passive-message context when we have one. - # QQ requires a msg_id on outbound messages to a user we've never - # seen; the last inbound msg_id is the natural choice. - msg_id = self._last_msg_id.get(chat_id) - req = ApprovalRequest( session_key=session_key, title="Execute this command?", @@ -2737,9 +1985,8 @@ class QQAdapter(BasePlatformAdapter): timeout_sec=self._APPROVAL_TIMEOUT_SECONDS, allow_permanent=allow_permanent and not smart_denied, ) - return await self.send_approval_request( - chat_id, req, reply_to=msg_id, - ) + # QQ requires a msg_id for passive replies; the last inbound id is the natural one. + return await self.send_approval_request(chat_id, req, reply_to=self._last_msg_id.get(chat_id)) _APPROVAL_TIMEOUT_SECONDS = 300 # matches gateway's default gateway_timeout @@ -2751,26 +1998,14 @@ class QQAdapter(BasePlatformAdapter): session_key: str = "", metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: - """Send a Yes/No update-confirmation prompt with inline buttons. - - Matches the cross-adapter contract used by - ``gateway/run.py``'s ``hermes update --gateway`` watcher. Button - clicks surface as ``INTERACTION_CREATE`` with - ``button_data = 'update_prompt:y'`` or ``'update_prompt:n'``; - the adapter's interaction callback writes the answer to - ``~/.hermes/.update_response`` so the detached update process - can read it. - """ + """Yes/No update-confirmation prompt; button clicks (``update_prompt:y|n``) + are written to ``~/.hermes/.update_response`` by the interaction callback.""" del session_key, metadata # present for contract parity only. default_hint = f" (default: {default})" if default else "" content = f"⚕ **Update Needs Your Input**\n\n{prompt}{default_hint}" - msg_id = self._last_msg_id.get(chat_id) return await self.send_with_keyboard( - chat_id, - content, - build_update_prompt_keyboard(), - reply_to=msg_id, + chat_id, content, build_update_prompt_keyboard(), reply_to=self._last_msg_id.get(chat_id) ) def _build_text_body( @@ -2778,25 +2013,12 @@ class QQAdapter(BasePlatformAdapter): ) -> Dict[str, Any]: """Build the message body for C2C/group text sending.""" msg_seq = self._next_msg_seq(reply_to or "default") - + text = content[: self.MAX_MESSAGE_LENGTH] if self._markdown_support: - body: Dict[str, Any] = { - "markdown": {"content": content[: self.MAX_MESSAGE_LENGTH]}, - "msg_type": MSG_TYPE_MARKDOWN, - "msg_seq": msg_seq, - } - else: - body = { - "content": content[: self.MAX_MESSAGE_LENGTH], - "msg_type": MSG_TYPE_TEXT, - "msg_seq": msg_seq, - } - + return {"markdown": {"content": text}, "msg_type": MSG_TYPE_MARKDOWN, "msg_seq": msg_seq} + body: Dict[str, Any] = {"content": text, "msg_type": MSG_TYPE_TEXT, "msg_seq": msg_seq} if reply_to: - # For non-markdown mode, add message_reference - if not self._markdown_support: - body["message_reference"] = {"message_id": reply_to} - + body["message_reference"] = {"message_id": reply_to} return body # ------------------------------------------------------------------ @@ -2811,85 +2033,35 @@ class QQAdapter(BasePlatformAdapter): reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: - """Send an image natively via QQ Bot API upload.""" + """Send an image natively via QQ Bot API upload; URL sources fall back to text.""" del metadata - - result = await self._send_media( - chat_id, image_url, MEDIA_TYPE_IMAGE, "image", caption, reply_to - ) + result = await self._send_media(chat_id, image_url, MEDIA_TYPE_IMAGE, "image", caption, reply_to) if result.success or not self._is_url(image_url): return result - - # Fallback to text URL logger.warning( - "[%s] Image send failed, falling back to text: %s", - self._log_tag, - result.error, + "[%s] Image send failed, falling back to text: %s", self._log_tag, result.error ) fallback = f"{caption}\n{image_url}" if caption else image_url return await self.send(chat_id=chat_id, content=fallback, 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: + async def send_image_file(self, chat_id, image_path, caption=None, reply_to=None, **kwargs) -> SendResult: """Send a local image file natively.""" - del kwargs - return await self._send_media( - chat_id, image_path, MEDIA_TYPE_IMAGE, "image", caption, reply_to - ) + return await self._send_media(chat_id, image_path, MEDIA_TYPE_IMAGE, "image", caption, reply_to) - async def send_voice( - self, - chat_id: str, - audio_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, - ) -> SendResult: + async def send_voice(self, chat_id, audio_path, caption=None, reply_to=None, **kwargs) -> SendResult: """Send a voice message natively.""" - del kwargs - return await self._send_media( - chat_id, audio_path, MEDIA_TYPE_VOICE, "voice", caption, reply_to - ) + return await self._send_media(chat_id, audio_path, MEDIA_TYPE_VOICE, "voice", caption, reply_to) - async def send_video( - self, - chat_id: str, - video_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, - ) -> SendResult: + async def send_video(self, chat_id, video_path, caption=None, reply_to=None, **kwargs) -> SendResult: """Send a video natively.""" - del kwargs - return await self._send_media( - chat_id, video_path, MEDIA_TYPE_VIDEO, "video", caption, reply_to - ) + return await self._send_media(chat_id, video_path, MEDIA_TYPE_VIDEO, "video", caption, 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, + self, chat_id, file_path, caption=None, file_name=None, reply_to=None, **kwargs ) -> SendResult: """Send a file/document natively.""" - del kwargs return await self._send_media( - chat_id, - file_path, - MEDIA_TYPE_FILE, - "file", - caption, - reply_to, - file_name=file_name, + chat_id, file_path, MEDIA_TYPE_FILE, "file", caption, reply_to, file_name=file_name ) async def _send_media( @@ -2904,35 +2076,19 @@ class QQAdapter(BasePlatformAdapter): ) -> SendResult: """Upload media and send as a native message. - Upload strategy: - - - **HTTP(S) URLs** → single ``POST /v2/{users|groups}/{id}/files`` - with ``url=...``. The QQ platform fetches the URL directly; fastest - path when the source is already hosted. - - **Local files** → three-step chunked upload (prepare / PUT parts / - complete). Handles files up to the platform's ~100 MB per-file - limit without the ~10 MB inline-base64 cap of the old adapter. + HTTP(S) URLs → single ``POST .../files`` with ``url=`` (QQ fetches it). + Local files → chunked upload (prepare / PUT parts / complete), up to the + platform's ~100 MB per-file limit. """ - if not self.is_connected: - if not await self._wait_for_reconnection(): - return SendResult(success=False, error="Not connected", retryable=True) - + if not await self._ensure_connected(): + return self._NOT_CONNECTED chat_type = self._guess_chat_type(chat_id) if chat_type == "guild": - # Guild channels don't support native media upload in the same way. - return SendResult( - success=False, - error="Guild media send not supported via this path", - ) + return SendResult(success=False, error="Guild media send not supported via this path") try: if self._is_url(media_source): - # URL upload — let the platform fetch it directly. - resolved_name = ( - file_name - or Path(urlparse(media_source).path).name - or "media" - ) + resolved_name = file_name or Path(urlparse(media_source).path).name or "media" upload = await self._upload_media( chat_type, chat_id, @@ -2942,53 +2098,26 @@ class QQAdapter(BasePlatformAdapter): file_name=resolved_name if file_type == MEDIA_TYPE_FILE else None, ) else: - # Local file — chunked upload (prepare / PUT parts / complete). resolved_name, upload = await self._upload_local_file( - chat_type, - chat_id, - media_source, - file_type, - file_name, + chat_type, chat_id, media_source, file_type, file_name ) - file_info = upload.get("file_info") or ( - upload.get("data", {}) or {} - ).get("file_info") + file_info = upload.get("file_info") or (upload.get("data", {}) or {}).get("file_info") if not file_info: - return SendResult( - success=False, - error=f"Upload returned no file_info: {upload}", - ) + return SendResult(success=False, error=f"Upload returned no file_info: {upload}") - # Send media message - msg_seq = self._next_msg_seq(chat_id) body: Dict[str, Any] = { "msg_type": MSG_TYPE_MEDIA, "media": {"file_info": file_info}, - "msg_seq": msg_seq, + "msg_seq": self._next_msg_seq(chat_id), } if caption: body["content"] = caption[: self.MAX_MESSAGE_LENGTH] if reply_to: body["msg_id"] = reply_to - - send_data = await self._api_request( - "POST", - ( - f"/v2/users/{chat_id}/messages" - if chat_type == "c2c" - else f"/v2/groups/{chat_id}/messages" - ), - body, - ) - return SendResult( - success=True, - message_id=str(send_data.get("id", uuid.uuid4().hex[:12])), - raw_response=send_data, - ) + return await self._post_message(self._messages_path(chat_type, chat_id), body) except UploadDailyLimitExceededError as exc: - # Non-retryable: daily quota hit. Give the caller actionable text - # so the model can compose a helpful reply. + # Non-retryable quota hit; give the model actionable text. logger.warning( "[%s] Daily upload limit exceeded for %s (%s)", self._log_tag, exc.file_name, exc.file_size_human, @@ -3026,16 +2155,11 @@ class QQAdapter(BasePlatformAdapter): file_type: int, file_name: Optional[str], ) -> Tuple[str, Dict[str, Any]]: - """Chunked-upload a local file and return ``(resolved_name, complete_response)``. + """Chunked-upload a local file; returns ``(resolved_name, complete_response)`` + whose ``file_info`` goes into the RichMedia body. - The returned ``complete_response`` contains the ``file_info`` token - that goes into the subsequent RichMedia message body. - - :raises UploadDailyLimitExceededError: On biz_code 40093002. - :raises UploadFileTooLargeError: When the file exceeds the platform limit. - :raises FileNotFoundError: If the path does not exist. - :raises ValueError: If the path looks like a placeholder (````). - :raises RuntimeError: If the HTTP client is not initialized. + Raises UploadDailyLimitExceededError / UploadFileTooLargeError from the + uploader, ValueError for placeholder paths like ````, FileNotFoundError. """ if not self._http_client: raise RuntimeError("HTTP client not initialized — not connected?") @@ -3046,16 +2170,12 @@ class QQAdapter(BasePlatformAdapter): if not local_path.exists() or not local_path.is_file(): if media_source.startswith("<") or len(media_source) < 3: - raise ValueError( - f"Invalid media source (looks like a placeholder): {media_source!r}" - ) + raise ValueError(f"Invalid media source (looks like a placeholder): {media_source!r}") raise FileNotFoundError(f"Media file not found: {local_path}") resolved_name = file_name or local_path.name uploader = ChunkedUploader( - api_request=self._api_request, - http_put=self._http_client.put, - log_tag=self._log_tag, + api_request=self._api_request, http_put=self._http_client.put, log_tag=self._log_tag ) complete = await uploader.upload( chat_type=chat_type, @@ -3066,82 +2186,28 @@ class QQAdapter(BasePlatformAdapter): ) return resolved_name, complete - async def _load_media( - self, source: str, file_name: Optional[str] = None - ) -> Tuple[str, str, str]: - """Load media from URL or local path. Returns (base64_or_url, content_type, filename).""" - source = str(source).strip() - if not source: - raise ValueError("Media source is required") - - parsed = urlparse(source) - if parsed.scheme in {"http", "https"}: - # For URLs, pass through directly to the upload API - content_type = mimetypes.guess_type(source)[0] or "application/octet-stream" - resolved_name = file_name or Path(parsed.path).name or "media" - return source, content_type, resolved_name - - # Local file — encode as raw base64 for QQ Bot API file_data field. - # The QQ API expects plain base64, NOT a data URI. - local_path = Path(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(): - # Guard against placeholder paths like "" that the LLM - # sometimes emits instead of real file paths. - if source.startswith("<") or len(source) < 3: - raise ValueError( - f"Invalid media source (looks like a placeholder): {source!r}" - ) - raise FileNotFoundError(f"Media file not found: {local_path}") - - raw = local_path.read_bytes() - resolved_name = file_name or local_path.name - content_type = ( - mimetypes.guess_type(str(local_path))[0] or "application/octet-stream" - ) - b64 = base64.b64encode(raw).decode("ascii") - return b64, content_type, resolved_name - # ------------------------------------------------------------------ # Typing indicator # ------------------------------------------------------------------ async def send_typing(self, chat_id: str, metadata=None) -> None: - """Send an input notify to a C2C user (only supported for C2C). - - Debounced to one request per ~50s (the API sets a 60s indicator). - The QQ API requires the originating message ID — retrieved from - ``_last_msg_id`` which is populated by ``_on_message``. - """ - if not self.is_connected: + """C2C-only input notify, debounced to ~50s (API shows a 60s indicator); + needs the last inbound msg_id from ``_last_msg_id``.""" + if not self.is_connected or self._guess_chat_type(chat_id) != "c2c": return - - chat_type = self._guess_chat_type(chat_id) - if chat_type != "c2c": - return - msg_id = self._last_msg_id.get(chat_id) if not msg_id: return - - # Debounce — skip if we sent recently now = time.time() - last_sent = self._typing_sent_at.get(chat_id, 0.0) - if now - last_sent < self._TYPING_DEBOUNCE_SECONDS: + if now - self._typing_sent_at.get(chat_id, 0.0) < self._TYPING_DEBOUNCE_SECONDS: return try: - msg_seq = self._next_msg_seq(chat_id) body = { "msg_type": MSG_TYPE_INPUT_NOTIFY, "msg_id": msg_id, - "input_notify": { - "input_type": 1, - "input_second": self._TYPING_INPUT_SECONDS, - }, - "msg_seq": msg_seq, + "input_notify": {"input_type": 1, "input_second": self._TYPING_INPUT_SECONDS}, + "msg_seq": self._next_msg_seq(chat_id), } await self._api_request("POST", f"/v2/users/{chat_id}/messages", body) self._typing_sent_at[chat_id] = now @@ -3153,14 +2219,8 @@ class QQAdapter(BasePlatformAdapter): # ------------------------------------------------------------------ def format_message(self, content: str) -> str: - """Format message for QQ. - - When markdown_support is enabled, content is sent as-is (QQ renders it). - When disabled, strip markdown via shared helper (same as BlueBubbles/SMS). - """ - if self._markdown_support: - return content - return strip_markdown(content) + """Pass markdown through when supported, else strip it (as BlueBubbles/SMS do).""" + return content if self._markdown_support else strip_markdown(content) # ------------------------------------------------------------------ # Chat info @@ -3184,18 +2244,12 @@ class QQAdapter(BasePlatformAdapter): def _guess_chat_type(self, chat_id: str) -> str: """Determine chat type from stored inbound metadata, fallback to 'c2c'.""" - if chat_id in self._chat_type_map: - return self._chat_type_map[chat_id] - return "c2c" + return self._chat_type_map.get(chat_id, "c2c") @staticmethod def _strip_at_mention(content: str) -> str: """Strip the @bot mention prefix from group message content.""" - # QQ group @-messages may have the bot's QQ/ID as prefix - import re - - stripped = re.sub(r"^@\S+\s*", "", content.strip()) - return stripped + return re.sub(r"^@\S+\s*", "", content.strip()) def _open_dm_opted_in(self) -> bool: if os.getenv("GATEWAY_ALLOW_ALL_USERS", "").lower() in {"true", "1", "yes"}: @@ -3232,25 +2286,15 @@ class QQAdapter(BasePlatformAdapter): return self._entry_matches(self._group_allow_from, group_id) if self._group_policy == "pairing": return False - if self._group_policy == "open": - return True - return False + return self._group_policy == "open" @staticmethod def _entry_matches(entries: List[str], target: str) -> bool: normalized_target = str(target).strip().lower() - for entry in entries: - normalized = str(entry).strip().lower() - if normalized == "*" or normalized == normalized_target: - return True - return False + return any(str(e).strip().lower() in ("*", normalized_target) for e in entries) def _parse_qq_timestamp(self, raw: str) -> datetime: - """Parse QQ API timestamp (ISO 8601 string or integer ms). - - The QQ API changed from integer milliseconds to ISO 8601 strings. - This handles both formats gracefully. - """ + """Parse a QQ timestamp — ISO 8601 string (current) or integer ms (legacy).""" if not raw: return datetime.now(tz=timezone.utc) try: @@ -3267,9 +2311,7 @@ class QQAdapter(BasePlatformAdapter): now = time.time() if len(self._seen_messages) > DEDUP_MAX_SIZE: cutoff = now - DEDUP_WINDOW_SECONDS - self._seen_messages = { - key: ts for key, ts in self._seen_messages.items() if ts > cutoff - } + self._seen_messages = {k: ts for k, ts in self._seen_messages.items() if ts > cutoff} if msg_id in self._seen_messages: return True self._seen_messages[msg_id] = now diff --git a/gateway/platforms/qqbot/chunked_upload.py b/gateway/platforms/qqbot/chunked_upload.py index 6979bd4cb7..32d62b517e 100644 --- a/gateway/platforms/qqbot/chunked_upload.py +++ b/gateway/platforms/qqbot/chunked_upload.py @@ -1,56 +1,40 @@ """QQ Bot chunked upload flow. The QQ v2 API caps inline base64 uploads (``file_data`` / ``url``) at ~10 MB. -For files between 10 MB and ~100 MB we have to use the three-step chunked -upload flow:: +Files between 10 MB and ~100 MB use the three-step chunked flow:: 1. POST /v2/{users|groups}/{id}/upload_prepare - → returns upload_id, block_size, and an array of pre-signed COS part URLs. - 2. For each part: - PUT the part bytes to its pre-signed COS URL, - then POST /v2/{users|groups}/{id}/upload_part_finish to acknowledge. + → upload_id, block_size, and pre-signed COS part URLs. + 2. Per part: PUT bytes to the COS URL, then POST .../upload_part_finish. 3. POST /v2/{users|groups}/{id}/files with {"upload_id": ...} - → returns the ``file_info`` token the caller uses in a RichMedia - message. + → ``file_info`` token used in a RichMedia message. -Error-code semantics (from the QQ Bot v2 API spec): +Error codes (QQ Bot v2 spec): ``40093001`` — ``upload_part_finish`` retryable +until the server's ``retry_timeout`` (or a local cap) elapses; ``40093002`` — +daily upload quota exceeded, surfaced as :class:`UploadDailyLimitExceededError`. +Other API/I/O failures raise ``RuntimeError``. -- ``40093001`` — ``upload_part_finish`` retryable. Retry until the server-provided - ``retry_timeout`` elapses (or a local cap). -- ``40093002`` — daily cumulative upload quota exceeded. Not retryable; surface - as :class:`UploadDailyLimitExceededError` so the caller can build a - user-friendly reply. - -Exceptions: - -- :class:`UploadDailyLimitExceededError` — daily quota hit (non-retryable). -- :class:`UploadFileTooLargeError` — file exceeds the platform per-file limit. -- :class:`RuntimeError` — generic upload failure (network, part PUT, complete). - -Ported from WideLee's qqbot-agent-sdk v1.2.2 (``media_loader.py::ChunkedUploader``) -so the heavy-upload path stays in-tree. Authorship preserved via Co-authored-by. +Ported from WideLee's qqbot-agent-sdk v1.2.2 (``media_loader.py::ChunkedUploader``). +Authorship preserved via Co-authored-by. """ from __future__ import annotations import asyncio -import functools import hashlib import logging from dataclasses import dataclass from pathlib import Path -from typing import Any, Awaitable, Callable, Dict, List, Optional +from typing import Any, Awaitable, Callable, Dict, List from gateway.platforms.qqbot.constants import FILE_UPLOAD_TIMEOUT logger = logging.getLogger(__name__) -# ── Error codes ────────────────────────────────────────────────────── _BIZ_CODE_DAILY_LIMIT = 40093002 # upload_prepare: daily cumulative limit _BIZ_CODE_PART_RETRYABLE = 40093001 # upload_part_finish: transient -# ── Part upload tuning ─────────────────────────────────────────────── _DEFAULT_CONCURRENT_PARTS = 1 _MAX_CONCURRENT_PARTS = 10 @@ -70,19 +54,12 @@ _MD5_10M_SIZE = 10_002_432 # ── Exceptions ─────────────────────────────────────────────────────── class UploadDailyLimitExceededError(Exception): - """Raised when ``upload_prepare`` returns biz_code 40093002. - - The daily cumulative upload quota for this bot has been reached. Callers - should surface :attr:`file_name` + :attr:`file_size_human` so the model - can compose a helpful reply. - """ + """Raised when ``upload_prepare`` returns biz_code 40093002 (daily quota hit).""" def __init__(self, file_name: str, file_size: int, message: str = "") -> None: self.file_name = file_name self.file_size = file_size - super().__init__( - message or f"Daily upload limit exceeded for {file_name!r}" - ) + super().__init__(message or f"Daily upload limit exceeded for {file_name!r}") @property def file_size_human(self) -> str: @@ -92,23 +69,13 @@ class UploadDailyLimitExceededError(Exception): class UploadFileTooLargeError(Exception): """Raised when a file exceeds the platform per-file size limit.""" - def __init__( - self, - file_name: str, - file_size: int, - limit_bytes: int = 0, - message: str = "", - ) -> None: + def __init__(self, file_name: str, file_size: int, limit_bytes: int = 0, message: str = "") -> None: self.file_name = file_name self.file_size = file_size self.limit_bytes = limit_bytes limit_str = f" ({format_size(limit_bytes)})" if limit_bytes else "" super().__init__( - message - or ( - f"File {file_name!r} ({format_size(file_size)}) " - f"exceeds platform limit{limit_str}" - ) + message or f"File {file_name!r} ({format_size(file_size)}) exceeds platform limit{limit_str}" ) @property @@ -120,16 +87,6 @@ class UploadFileTooLargeError(Exception): return format_size(self.limit_bytes) if self.limit_bytes else "unknown" -# ── Progress tracking ──────────────────────────────────────────────── - -@dataclass -class _UploadProgress: - total_parts: int = 0 - total_bytes: int = 0 - completed_parts: int = 0 - uploaded_bytes: int = 0 - - # ── Prepare-response shape ─────────────────────────────────────────── @dataclass @@ -149,35 +106,24 @@ class _PrepareResult: def _parse_prepare_response(raw: Dict[str, Any]) -> _PrepareResult: - """Parse the upload_prepare API response into a normalized shape. - - The API may return the response directly or wrapped in ``data``. - """ + """Parse upload_prepare response (either bare or wrapped in ``data``).""" src = raw.get("data") if isinstance(raw.get("data"), dict) else raw upload_id = str(src.get("upload_id", "")) if not upload_id: - raise ValueError( - f"upload_prepare response missing upload_id: {str(raw)[:200]}" - ) + raise ValueError(f"upload_prepare response missing upload_id: {str(raw)[:200]}") block_size = int(src.get("block_size", 0)) raw_parts = src.get("parts") or src.get("part_list") or [] if not isinstance(raw_parts, list) or not raw_parts: - raise ValueError( - f"upload_prepare response missing parts: {str(raw)[:200]}" - ) - parts: List[_PreparePart] = [] - for p in raw_parts: - if not isinstance(p, dict): - continue - parts.append( - _PreparePart( - index=int(p.get("part_index") or p.get("index") or 0), - presigned_url=str( - p.get("presigned_url") or p.get("url") or "" - ), - block_size=int(p.get("block_size", 0)), - ) + raise ValueError(f"upload_prepare response missing parts: {str(raw)[:200]}") + parts = [ + _PreparePart( + index=int(p.get("part_index") or p.get("index") or 0), + presigned_url=str(p.get("presigned_url") or p.get("url") or ""), + block_size=int(p.get("block_size", 0)), ) + for p in raw_parts + if isinstance(p, dict) + ] return _PrepareResult( upload_id=upload_id, block_size=block_size, @@ -187,30 +133,28 @@ def _parse_prepare_response(raw: Dict[str, Any]) -> _PrepareResult: ) +def _api_path(chat_type: str, target_id: str, endpoint: str) -> str: + base = "/v2/users" if chat_type == "c2c" else "/v2/groups" + return f"{base}/{target_id}/{endpoint}" + + # ── Chunked upload driver ──────────────────────────────────────────── -ApiRequestFn = Callable[..., Awaitable[Dict[str, Any]]] -"""Signature of the adapter's ``_api_request`` callable. - -We pass the bound method in rather than importing the adapter, to avoid -circular imports and keep this module testable in isolation. -""" - - class ChunkedUploader: """Run the prepare → PUT parts → complete sequence. :param api_request: Bound ``_api_request(method, path, body=..., timeout=...)`` - coroutine from the adapter. Must raise ``RuntimeError`` with the biz_code - embedded in the message on API errors. - :param http_put: Coroutine ``(url, data, headers, timeout) -> response`` for - COS part uploads. Typically wraps ``httpx.AsyncClient.put``. + coroutine from the adapter (passed in rather than imported to avoid + circular imports). Must raise ``RuntimeError`` with the biz_code in + the message on API errors. + :param http_put: Coroutine ``(url, data, headers) -> httpx-like response`` + for COS part uploads. :param log_tag: Log prefix. """ def __init__( self, - api_request: ApiRequestFn, + api_request: Callable[..., Awaitable[Dict[str, Any]]], http_put: Callable[..., Awaitable[Any]], log_tag: str = "QQBot", ) -> None: @@ -229,38 +173,26 @@ class ChunkedUploader: """Run the full chunked upload and return the ``complete_upload`` response. :param chat_type: ``'c2c'`` or ``'group'``. - :param target_id: User or group openid. - :param file_path: Absolute path to a local file. :param file_type: ``MEDIA_TYPE_*`` constant. - :param file_name: Original filename (for upload_prepare). - :returns: The raw response dict from ``complete_upload`` — contains - ``file_info`` that the caller uses in a RichMedia message body. + :returns: Raw ``complete_upload`` response dict (contains ``file_info``). :raises UploadDailyLimitExceededError: On biz_code 40093002. :raises UploadFileTooLargeError: When the file exceeds the platform limit. :raises RuntimeError: On other API or I/O failures. """ if chat_type not in {"c2c", "group"}: - raise ValueError( - f"ChunkedUploader: unsupported chat_type {chat_type!r}" - ) - - path = Path(file_path) - file_size = path.stat().st_size + raise ValueError(f"ChunkedUploader: unsupported chat_type {chat_type!r}") + file_size = Path(file_path).stat().st_size logger.info( "[%s] Chunked upload start: file=%s size=%s type=%d", self._log_tag, file_name, format_size(file_size), file_type, ) - # Step 1: compute hashes (blocking I/O → executor). + # Hashing is blocking I/O → executor. hashes = await asyncio.get_running_loop().run_in_executor( None, _compute_file_hashes, file_path, file_size ) - - # Step 2: upload_prepare. - prepare = await self._prepare( - chat_type, target_id, file_type, file_name, file_size, hashes - ) + prepare = await self._prepare(chat_type, target_id, file_type, file_name, file_size, hashes) max_concurrent = min(prepare.concurrency, _MAX_CONCURRENT_PARTS) retry_timeout = min( prepare.retry_timeout if prepare.retry_timeout > 0 else _PART_FINISH_DEFAULT_TIMEOUT, @@ -272,41 +204,21 @@ class ChunkedUploader: len(prepare.parts), max_concurrent, ) - progress = _UploadProgress( - total_parts=len(prepare.parts), - total_bytes=file_size, - ) + total_parts = len(prepare.parts) + completed = [0] # shared counter for progress logging + sem = asyncio.Semaphore(max(max_concurrent, 1)) - # Step 3: PUT each part + notify. - tasks: List[Callable[[], Awaitable[None]]] = [ - functools.partial( - self._upload_one_part, - chat_type=chat_type, - target_id=target_id, - file_path=file_path, - file_size=file_size, - upload_id=prepare.upload_id, - rsp_block_size=prepare.block_size, - part=part, - retry_timeout=retry_timeout, - progress=progress, - ) - for part in prepare.parts - ] - await _run_with_concurrency(tasks, max_concurrent) + async def _run(part: _PreparePart) -> None: + async with sem: + await self._upload_one_part( + chat_type, target_id, file_path, file_size, prepare.upload_id, + prepare.block_size, part, retry_timeout, total_parts, completed, + ) - logger.info( - "[%s] All %d parts uploaded, completing…", - self._log_tag, len(prepare.parts), - ) - - # Step 4: complete_upload (retry on transient errors). + await asyncio.gather(*(_run(p) for p in prepare.parts)) + logger.info("[%s] All %d parts uploaded, completing…", self._log_tag, total_parts) return await self._complete(chat_type, target_id, prepare.upload_id) - # ────────────────────────────────────────────────────────────────── - # Step 1 — upload_prepare - # ────────────────────────────────────────────────────────────────── - async def _prepare( self, chat_type: str, @@ -316,8 +228,6 @@ class ChunkedUploader: file_size: int, hashes: Dict[str, str], ) -> _PrepareResult: - base = "/v2/users" if chat_type == "c2c" else "/v2/groups" - path = f"{base}/{target_id}/upload_prepare" body = { "file_type": file_type, "file_name": file_name, @@ -328,21 +238,16 @@ class ChunkedUploader: } try: raw = await self._api_request( - "POST", path, body=body, timeout=FILE_UPLOAD_TIMEOUT + "POST", _api_path(chat_type, target_id, "upload_prepare"), + body=body, timeout=FILE_UPLOAD_TIMEOUT, ) except RuntimeError as exc: err_msg = str(exc) if f"{_BIZ_CODE_DAILY_LIMIT}" in err_msg: - raise UploadDailyLimitExceededError( - file_name, file_size, err_msg - ) from exc + raise UploadDailyLimitExceededError(file_name, file_size, err_msg) from exc raise return _parse_prepare_response(raw) - # ────────────────────────────────────────────────────────────────── - # Step 2 — PUT one part + part_finish - # ────────────────────────────────────────────────────────────────── - async def _upload_one_part( self, chat_type: str, @@ -353,7 +258,8 @@ class ChunkedUploader: rsp_block_size: int, part: _PreparePart, retry_timeout: float, - progress: _UploadProgress, + total_parts: int, + completed: List[int], ) -> None: """PUT one part to COS, then call ``upload_part_finish``.""" part_index = part.index @@ -362,82 +268,74 @@ class ChunkedUploader: offset = (part_index - 1) * rsp_block_size length = min(actual_block_size, file_size - offset) - # Read this slice of the file (blocking → executor). data = await asyncio.get_running_loop().run_in_executor( None, _read_file_chunk, file_path, offset, length ) md5_hex = hashlib.md5(data).hexdigest() - logger.debug( "[%s] Part %d/%d: uploading %s (offset=%d md5=%s)", - self._log_tag, part_index, progress.total_parts, - format_size(length), offset, md5_hex, + self._log_tag, part_index, total_parts, format_size(length), offset, md5_hex, ) - await self._put_to_presigned_url( - part.presigned_url, data, part_index, progress.total_parts - ) + await self._put_to_presigned_url(part.presigned_url, data, part_index, total_parts) await self._part_finish_with_retry( - chat_type, target_id, upload_id, - part_index, length, md5_hex, retry_timeout, + chat_type, target_id, upload_id, part_index, length, md5_hex, retry_timeout, ) - progress.completed_parts += 1 - progress.uploaded_bytes += length + completed[0] += 1 logger.debug( "[%s] Part %d/%d done (%d/%d total)", - self._log_tag, part_index, progress.total_parts, - progress.completed_parts, progress.total_parts, + self._log_tag, part_index, total_parts, completed[0], total_parts, ) - async def _put_to_presigned_url( + async def _with_retries( self, - url: str, - data: bytes, - part_index: int, - total_parts: int, - ) -> None: - """PUT part data to a pre-signed COS URL with retry.""" - last_exc: Optional[Exception] = None - for attempt in range(_PART_UPLOAD_MAX_RETRIES + 1): + attempt_fn: Callable[[], Awaitable[Any]], + *, + max_retries: int, + base_delay: float, + label: str, + failure_label: str, + ) -> Any: + """Run *attempt_fn* up to ``max_retries + 1`` times with exponential backoff.""" + last_exc: Exception | None = None + for attempt in range(max_retries + 1): try: - resp = await asyncio.wait_for( - self._http_put( - url, - data=data, - headers={"Content-Length": str(len(data))}, - ), - timeout=_PART_UPLOAD_TIMEOUT, - ) - # Caller's http_put is expected to return an httpx-like response. - status = getattr(resp, "status_code", 0) - if 200 <= status < 300: - logger.debug( - "[%s] PUT part %d/%d: %d OK", - self._log_tag, part_index, total_parts, status, - ) - return - body_preview = "" - try: - body_preview = getattr(resp, "text", "")[:200] - except Exception: # pragma: no cover — defensive - pass - raise RuntimeError( - f"COS PUT returned {status}: {body_preview}" - ) + return await attempt_fn() except Exception as exc: last_exc = exc - if attempt < _PART_UPLOAD_MAX_RETRIES: - delay = 1.0 * (2 ** attempt) + if attempt < max_retries: + delay = base_delay * (2 ** attempt) logger.warning( - "[%s] PUT part %d/%d attempt %d failed, retry in %.1fs: %s", - self._log_tag, part_index, total_parts, - attempt + 1, delay, exc, + "[%s] %s attempt %d failed, retry in %.1fs: %s", + self._log_tag, label, attempt + 1, delay, exc, ) await asyncio.sleep(delay) - raise RuntimeError( - f"Part {part_index}/{total_parts} upload failed after " - f"{_PART_UPLOAD_MAX_RETRIES + 1} attempts: {last_exc}" + raise RuntimeError(f"{failure_label} failed after {max_retries + 1} attempts: {last_exc}") + + async def _put_to_presigned_url(self, url: str, data: bytes, part_index: int, total_parts: int) -> None: + """PUT part data to a pre-signed COS URL with retry.""" + + async def _attempt() -> None: + resp = await asyncio.wait_for( + self._http_put(url, data=data, headers={"Content-Length": str(len(data))}), + timeout=_PART_UPLOAD_TIMEOUT, + ) + status = getattr(resp, "status_code", 0) + if 200 <= status < 300: + logger.debug("[%s] PUT part %d/%d: %d OK", self._log_tag, part_index, total_parts, status) + return + body_preview = "" + try: + body_preview = getattr(resp, "text", "")[:200] + except Exception: # pragma: no cover — defensive + pass + raise RuntimeError(f"COS PUT returned {status}: {body_preview}") + + await self._with_retries( + _attempt, max_retries=_PART_UPLOAD_MAX_RETRIES, base_delay=1.0, + label=f"PUT part {part_index}/{total_parts}", + failure_label=f"Part {part_index}/{total_parts} upload", ) async def _part_finish_with_retry( @@ -450,28 +348,18 @@ class ChunkedUploader: md5: str, retry_timeout: float, ) -> None: - """Call ``upload_part_finish``, retrying on biz_code 40093001.""" - base = "/v2/users" if chat_type == "c2c" else "/v2/groups" - path = f"{base}/{target_id}/upload_part_finish" - body = { - "upload_id": upload_id, - "part_index": part_index, - "block_size": block_size, - "md5": md5, - } - + """Call ``upload_part_finish``, retrying on biz_code 40093001 until *retry_timeout*.""" + path = _api_path(chat_type, target_id, "upload_part_finish") + body = {"upload_id": upload_id, "part_index": part_index, "block_size": block_size, "md5": md5} loop = asyncio.get_running_loop() start = loop.time() attempt = 0 while True: try: - await self._api_request( - "POST", path, body=body, timeout=FILE_UPLOAD_TIMEOUT - ) + await self._api_request("POST", path, body=body, timeout=FILE_UPLOAD_TIMEOUT) return except RuntimeError as exc: - err_msg = str(exc) - if f"{_BIZ_CODE_PART_RETRYABLE}" not in err_msg: + if f"{_BIZ_CODE_PART_RETRYABLE}" not in str(exc): raise elapsed = loop.time() - start if elapsed >= retry_timeout: @@ -481,50 +369,23 @@ class ChunkedUploader: ) from exc attempt += 1 logger.debug( - "[%s] part_finish retryable error, attempt %d, " - "elapsed=%.1fs: %s", + "[%s] part_finish retryable error, attempt %d, elapsed=%.1fs: %s", self._log_tag, attempt, elapsed, exc, ) await asyncio.sleep(_PART_FINISH_RETRY_INTERVAL) - # ────────────────────────────────────────────────────────────────── - # Step 3 — complete_upload - # ────────────────────────────────────────────────────────────────── - - async def _complete( - self, - chat_type: str, - target_id: str, - upload_id: str, - ) -> Dict[str, Any]: + async def _complete(self, chat_type: str, target_id: str, upload_id: str) -> Dict[str, Any]: """Call ``complete_upload`` with retry. - This reuses the ``/files`` endpoint (same as the simple URL-based upload) - but signals the chunked-completion path by sending only ``upload_id``. + Reuses the ``/files`` endpoint (same as the simple URL-based upload) but + signals the chunked-completion path by sending only ``upload_id``. """ - base = "/v2/users" if chat_type == "c2c" else "/v2/groups" - path = f"{base}/{target_id}/files" + path = _api_path(chat_type, target_id, "files") body = {"upload_id": upload_id} - - last_exc: Optional[Exception] = None - for attempt in range(_COMPLETE_UPLOAD_MAX_RETRIES + 1): - try: - return await self._api_request( - "POST", path, body=body, timeout=FILE_UPLOAD_TIMEOUT - ) - except Exception as exc: - last_exc = exc - if attempt < _COMPLETE_UPLOAD_MAX_RETRIES: - delay = _COMPLETE_UPLOAD_BASE_DELAY * (2 ** attempt) - logger.warning( - "[%s] complete_upload attempt %d failed, " - "retry in %.1fs: %s", - self._log_tag, attempt + 1, delay, exc, - ) - await asyncio.sleep(delay) - raise RuntimeError( - f"complete_upload failed after " - f"{_COMPLETE_UPLOAD_MAX_RETRIES + 1} attempts: {last_exc}" + return await self._with_retries( + lambda: self._api_request("POST", path, body=body, timeout=FILE_UPLOAD_TIMEOUT), + max_retries=_COMPLETE_UPLOAD_MAX_RETRIES, base_delay=_COMPLETE_UPLOAD_BASE_DELAY, + label="complete_upload", failure_label="complete_upload", ) @@ -541,10 +402,7 @@ def format_size(size_bytes: int) -> str: def _read_file_chunk(file_path: str, offset: int, length: int) -> bytes: - """Read *length* bytes from *file_path* starting at *offset*. - - :raises IOError: If fewer bytes were read than expected (truncated file). - """ + """Read *length* bytes at *offset*; raises IOError on a short read (truncated file).""" with open(file_path, "rb") as fh: fh.seek(offset) data = fh.read(length) @@ -561,15 +419,10 @@ def _compute_file_hashes(file_path: str, file_size: int) -> Dict[str, str]: md5 = hashlib.md5() sha1 = hashlib.sha1() md5_10m = hashlib.md5() - need_10m = file_size > _MD5_10M_SIZE bytes_read = 0 - with open(file_path, "rb") as fh: - while True: - chunk = fh.read(65536) - if not chunk: - break + while chunk := fh.read(65536): md5.update(chunk) sha1.update(chunk) if need_10m: @@ -577,7 +430,6 @@ def _compute_file_hashes(file_path: str, file_size: int) -> Dict[str, str]: if remaining > 0: md5_10m.update(chunk[:remaining]) bytes_read += len(chunk) - full_md5 = md5.hexdigest() return { "md5": full_md5, @@ -585,18 +437,3 @@ def _compute_file_hashes(file_path: str, file_size: int) -> Dict[str, str]: # For small files the "10m" hash is just the full md5. "md5_10m": md5_10m.hexdigest() if need_10m else full_md5, } - - -async def _run_with_concurrency( - tasks: List[Callable[[], Awaitable[None]]], - concurrency: int, -) -> None: - """Run a list of thunks with a bounded number in flight at once.""" - concurrency = max(concurrency, 1) - sem = asyncio.Semaphore(concurrency) - - async def _wrap(thunk: Callable[[], Awaitable[None]]) -> None: - async with sem: - await thunk() - - await asyncio.gather(*(_wrap(t) for t in tasks)) diff --git a/gateway/platforms/qqbot/constants.py b/gateway/platforms/qqbot/constants.py index ddae3c133e..9ceb09c8e9 100644 --- a/gateway/platforms/qqbot/constants.py +++ b/gateway/platforms/qqbot/constants.py @@ -4,18 +4,11 @@ from __future__ import annotations import os -# --------------------------------------------------------------------------- -# QQBot adapter version — bump on functional changes to the adapter package. -# --------------------------------------------------------------------------- - +# Bump on functional changes to the adapter package. QQBOT_VERSION = "1.1.0" -# --------------------------------------------------------------------------- -# API endpoints -# --------------------------------------------------------------------------- - -# The portal domain is configurable via QQ_API_HOST for corporate proxies -# or test environments. Default: q.qq.com (production). +# ── API endpoints ── +# Portal domain is overridable (QQ_PORTAL_HOST) for corporate proxies / test environments. PORTAL_HOST = os.getenv("QQ_PORTAL_HOST", "q.qq.com") API_BASE = "https://api.sgroup.qq.com" @@ -25,15 +18,9 @@ GATEWAY_URL_PATH = "/gateway" # QR-code onboard endpoints (on the portal host) ONBOARD_CREATE_PATH = "/lite/create_bind_task" ONBOARD_POLL_PATH = "/lite/poll_bind_result" -QR_URL_TEMPLATE = ( - "https://q.qq.com/qqbot/openclaw/connect.html" - "?task_id={task_id}&_wv=2&source=hermes" -) - -# --------------------------------------------------------------------------- -# Timeouts & retry -# --------------------------------------------------------------------------- +QR_URL_TEMPLATE = "https://q.qq.com/qqbot/openclaw/connect.html?task_id={task_id}&_wv=2&source=hermes" +# ── Timeouts & retry ── DEFAULT_API_TIMEOUT = 30.0 FILE_UPLOAD_TIMEOUT = 120.0 CONNECT_TIMEOUT_SECONDS = 20.0 @@ -47,27 +34,18 @@ MAX_QUICK_DISCONNECT_COUNT = 3 ONBOARD_POLL_INTERVAL = 2.0 # seconds between poll_bind_result calls ONBOARD_API_TIMEOUT = 10.0 -# --------------------------------------------------------------------------- -# Message limits -# --------------------------------------------------------------------------- - +# ── Message limits ── MAX_MESSAGE_LENGTH = 4000 DEDUP_WINDOW_SECONDS = 300 DEDUP_MAX_SIZE = 1000 -# --------------------------------------------------------------------------- -# QQ Bot message types -# --------------------------------------------------------------------------- - +# ── QQ Bot message types ── MSG_TYPE_TEXT = 0 MSG_TYPE_MARKDOWN = 2 MSG_TYPE_MEDIA = 7 MSG_TYPE_INPUT_NOTIFY = 6 -# --------------------------------------------------------------------------- -# QQ Bot file media types -# --------------------------------------------------------------------------- - +# ── QQ Bot file media types ── MEDIA_TYPE_IMAGE = 1 MEDIA_TYPE_VIDEO = 2 MEDIA_TYPE_VOICE = 3 diff --git a/gateway/platforms/qqbot/crypto.py b/gateway/platforms/qqbot/crypto.py index 426bd29de5..8eefc8ce1b 100644 --- a/gateway/platforms/qqbot/crypto.py +++ b/gateway/platforms/qqbot/crypto.py @@ -7,39 +7,19 @@ import os def generate_bind_key() -> str: - """Generate a 256-bit random AES key and return it as base64. + """Generate a random 256-bit AES key as base64. - The key is passed to ``create_bind_task`` so the server can encrypt - the bot's *client_secret* before returning it. Only this CLI holds - the key, ensuring the secret never travels in plaintext. + Passed to ``create_bind_task`` so the server encrypts the bot's + *client_secret*; only this CLI holds the key, so the secret never travels + in plaintext. """ return base64.b64encode(os.urandom(32)).decode() def decrypt_secret(encrypted_base64: str, key_base64: str) -> str: - """Decrypt a base64-encoded AES-256-GCM ciphertext. - - Ciphertext layout (after base64-decoding):: - - IV (12 bytes) ‖ ciphertext (N bytes) ‖ AuthTag (16 bytes) - - Args: - encrypted_base64: The ``bot_encrypt_secret`` value from - ``poll_bind_result``. - key_base64: The base64 AES key generated by - :func:`generate_bind_key`. - - Returns: - The decrypted *client_secret* as a UTF-8 string. - """ + """Decrypt ``bot_encrypt_secret`` (base64 of ``IV(12) ‖ ciphertext ‖ tag(16)``) to a UTF-8 string.""" from cryptography.hazmat.primitives.ciphers.aead import AESGCM - key = base64.b64decode(key_base64) raw = base64.b64decode(encrypted_base64) - - iv = raw[:12] - ciphertext_with_tag = raw[12:] # AESGCM expects ciphertext + tag concatenated - - aesgcm = AESGCM(key) - plaintext = aesgcm.decrypt(iv, ciphertext_with_tag, None) - return plaintext.decode("utf-8") + # AESGCM expects ciphertext + tag concatenated. + return AESGCM(base64.b64decode(key_base64)).decrypt(raw[:12], raw[12:], None).decode("utf-8") diff --git a/gateway/platforms/qqbot/keyboards.py b/gateway/platforms/qqbot/keyboards.py index 9c1afbbfa3..4e8f65f28c 100644 --- a/gateway/platforms/qqbot/keyboards.py +++ b/gateway/platforms/qqbot/keyboards.py @@ -1,22 +1,9 @@ -"""QQ Bot inline keyboards + approval / update-prompt senders. +"""QQ Bot inline keyboards + approval / update-prompt helpers. -QQ Bot v2 supports attaching inline keyboards to outbound messages. When a -user clicks a button, the platform dispatches an ``INTERACTION_CREATE`` -gateway event containing the button's ``data`` payload. The bot must ACK the -interaction promptly via ``PUT /interactions/{id}`` or the user sees an -error indicator on the button. - -This module provides: - -- :class:`InlineKeyboard` + button dataclasses — serialized into the - ``keyboard`` field of the outbound message body. -- :func:`build_approval_keyboard` — 3-button ✅ once / ⭐ always / ❌ deny - keyboard for tool-approval flows. -- :func:`build_update_prompt_keyboard` — Yes/No keyboard for update confirms. -- :func:`parse_approval_button_data` / :func:`parse_update_prompt_button_data` - — decode the ``button_data`` payload from ``INTERACTION_CREATE``. -- :class:`ApprovalRequest` + :class:`ApprovalSender` — high-level helper that - builds an approval message with keyboard and posts it to a c2c / group chat. +QQ Bot v2 attaches inline keyboards to outbound messages. A button click +dispatches an ``INTERACTION_CREATE`` gateway event carrying the button's +``data`` payload; the bot must ACK promptly via ``PUT /interactions/{id}`` or +the user sees an error indicator on the button. ``button_data`` formats:: @@ -29,250 +16,146 @@ keyboard types). Authorship preserved via Co-authored-by. from __future__ import annotations -import logging import re -from dataclasses import dataclass, field -from typing import Any, Awaitable, Callable, Dict, List, Optional - -logger = logging.getLogger(__name__) - -# ── button_data prefixes + patterns ────────────────────────────────── +from dataclasses import dataclass, field, fields, is_dataclass +from typing import Any, Dict, List, Optional APPROVAL_BUTTON_PREFIX = "approve:" UPDATE_PROMPT_PREFIX = "update_prompt:" -# Pattern: approve:: # session_key may itself contain colons (e.g. agent:main:qqbot:c2c:OPENID), # so the session_key group is greedy but trails the decision. -_APPROVAL_DATA_RE = re.compile( - r"^approve:(.+):(allow-once|allow-always|deny)$" -) - -# Pattern: update_prompt:y | update_prompt:n +_APPROVAL_DATA_RE = re.compile(r"^approve:(.+):(allow-once|allow-always|deny)$") _UPDATE_PROMPT_RE = re.compile(r"^update_prompt:(y|n)$") # ── Keyboard dataclasses ───────────────────────────────────────────── +def _to_dict(value: Any) -> Any: + """Serialize a dataclass tree in field-declaration order (the wire shape).""" + if is_dataclass(value): + return {f.name: _to_dict(getattr(value, f.name)) for f in fields(value)} + if isinstance(value, list): + return [_to_dict(v) for v in value] + return value + + +class _Serializable: + def to_dict(self) -> Dict[str, Any]: + return _to_dict(self) + + @dataclass -class KeyboardButtonPermission: +class KeyboardButtonPermission(_Serializable): """Button permission metadata. ``type=2`` means all users can click.""" type: int = 2 - def to_dict(self) -> Dict[str, Any]: - return {"type": self.type} - @dataclass -class KeyboardButtonAction: - """What happens when the button is clicked. - - :param type: ``1`` (Callback — triggers ``INTERACTION_CREATE``) or - ``2`` (Link — opens a URL). - :param data: Payload delivered in ``data.resolved.button_data`` when - ``type=1``. - :param permission: :class:`KeyboardButtonPermission`. - :param click_limit: Max clicks per user (``1`` = single-use). +class KeyboardButtonAction(_Serializable): + """Click behaviour: ``type`` 1 = Callback (INTERACTION_CREATE with ``data``), 2 = Link. + ``click_limit=1`` = single-use. """ type: int data: str - permission: KeyboardButtonPermission = field( - default_factory=KeyboardButtonPermission - ) + permission: KeyboardButtonPermission = field(default_factory=KeyboardButtonPermission) click_limit: int = 1 - def to_dict(self) -> Dict[str, Any]: - return { - "type": self.type, - "data": self.data, - "permission": self.permission.to_dict(), - "click_limit": self.click_limit, - } - @dataclass -class KeyboardButtonRenderData: - """Visual rendering of a button. - - :param label: Pre-click label. - :param visited_label: Post-click label (button stays greyed in place). - :param style: ``0`` = grey, ``1`` = blue. - """ +class KeyboardButtonRenderData(_Serializable): + """Visual rendering: pre/post-click labels; ``style`` 0 = grey, 1 = blue.""" label: str visited_label: str style: int = 1 - def to_dict(self) -> Dict[str, Any]: - return { - "label": self.label, - "visited_label": self.visited_label, - "style": self.style, - } - @dataclass -class KeyboardButton: - """One button in a keyboard. - - :param group_id: Buttons sharing a ``group_id`` are mutually exclusive — - clicking one greys the rest. - """ +class KeyboardButton(_Serializable): + """One button; buttons sharing a ``group_id`` are mutually exclusive.""" id: str render_data: KeyboardButtonRenderData action: KeyboardButtonAction group_id: str = "default" - def to_dict(self) -> Dict[str, Any]: - return { - "id": self.id, - "render_data": self.render_data.to_dict(), - "action": self.action.to_dict(), - "group_id": self.group_id, - } - @dataclass -class KeyboardRow: +class KeyboardRow(_Serializable): buttons: List[KeyboardButton] = field(default_factory=list) - def to_dict(self) -> Dict[str, Any]: - return {"buttons": [b.to_dict() for b in self.buttons]} - @dataclass -class KeyboardContent: +class KeyboardContent(_Serializable): rows: List[KeyboardRow] = field(default_factory=list) - def to_dict(self) -> Dict[str, Any]: - return {"rows": [r.to_dict() for r in self.rows]} - @dataclass -class InlineKeyboard: +class InlineKeyboard(_Serializable): """Top-level keyboard payload — goes into ``MessageToCreate.keyboard``.""" content: KeyboardContent = field(default_factory=KeyboardContent) - def to_dict(self) -> Dict[str, Any]: - return {"content": self.content.to_dict()} - # ── INTERACTION_CREATE parsing ─────────────────────────────────────── def parse_approval_button_data(button_data: str) -> Optional[tuple[str, str]]: - """Parse approval ``button_data`` into ``(session_key, decision)``. - - :param button_data: Raw ``data.resolved.button_data`` from - ``INTERACTION_CREATE``. - :returns: ``(session_key, decision)`` or ``None`` if not an approval button. - """ + """Parse approval ``button_data`` into ``(session_key, decision)`` or ``None``.""" m = _APPROVAL_DATA_RE.match(button_data or "") - if not m: - return None - return m.group(1), m.group(2) + return (m.group(1), m.group(2)) if m else None def parse_update_prompt_button_data(button_data: str) -> Optional[str]: - """Parse update-prompt ``button_data`` into ``'y'`` or ``'n'``.""" + """Parse update-prompt ``button_data`` into ``'y'`` / ``'n'`` or ``None``.""" m = _UPDATE_PROMPT_RE.match(button_data or "") - if not m: - return None - return m.group(1) + return m.group(1) if m else None # ── Keyboard builders ──────────────────────────────────────────────── def _make_callback_button( - btn_id: str, - label: str, - visited_label: str, - data: str, - style: int, - group_id: str, + btn_id: str, label: str, visited_label: str, data: str, style: int, group_id: str, ) -> KeyboardButton: return KeyboardButton( id=btn_id, - render_data=KeyboardButtonRenderData( - label=label, - visited_label=visited_label, - style=style, - ), + render_data=KeyboardButtonRenderData(label=label, visited_label=visited_label, style=style), action=KeyboardButtonAction(type=1, data=data), group_id=group_id, ) -def build_approval_keyboard(session_key: str, *, allow_permanent: bool = True) -> InlineKeyboard: - """Build the approval keyboard, hiding persistent scope when unavailable. - - Layout: ``[✅ 允许一次] [⭐ 始终允许] [❌ 拒绝]`` — all three share - ``group_id='approval'`` so clicking one greys out the rest. - - :param session_key: Embedded into ``button_data`` so the decision - routes back to the right pending approval. - """ - buttons = [ - _make_callback_button( - btn_id="allow", label="✅ 允许一次", visited_label="已允许", - data=f"{APPROVAL_BUTTON_PREFIX}{session_key}:allow-once", - style=1, group_id="approval", - ) - ] - if allow_permanent: - buttons.append(_make_callback_button( - btn_id="always", label="⭐ 始终允许", visited_label="已始终允许", - data=f"{APPROVAL_BUTTON_PREFIX}{session_key}:allow-always", - style=1, group_id="approval", - )) - buttons.append(_make_callback_button( - btn_id="deny", label="❌ 拒绝", visited_label="已拒绝", - data=f"{APPROVAL_BUTTON_PREFIX}{session_key}:deny", - style=0, group_id="approval", - )) +def _single_row_keyboard(buttons: List[KeyboardButton]) -> InlineKeyboard: return InlineKeyboard(content=KeyboardContent(rows=[KeyboardRow(buttons=buttons)])) +def build_approval_keyboard(session_key: str, *, allow_permanent: bool = True) -> InlineKeyboard: + """Build ``[✅ 允许一次] [⭐ 始终允许] [❌ 拒绝]`` (one group, so a click greys the rest). + + The ⭐ button is hidden when persistent scope is unavailable. *session_key* + is embedded in ``button_data`` so the decision routes to the right approval. + """ + prefix = f"{APPROVAL_BUTTON_PREFIX}{session_key}" + buttons = [_make_callback_button("allow", "✅ 允许一次", "已允许", f"{prefix}:allow-once", 1, "approval")] + if allow_permanent: + buttons.append(_make_callback_button("always", "⭐ 始终允许", "已始终允许", f"{prefix}:allow-always", 1, "approval")) + buttons.append(_make_callback_button("deny", "❌ 拒绝", "已拒绝", f"{prefix}:deny", 0, "approval")) + return _single_row_keyboard(buttons) + + def build_update_prompt_keyboard() -> InlineKeyboard: """Build a Yes/No keyboard for update confirmation prompts.""" - return InlineKeyboard( - content=KeyboardContent( - rows=[ - KeyboardRow(buttons=[ - _make_callback_button( - btn_id="yes", - label="✓ 确认", - visited_label="已确认", - data=f"{UPDATE_PROMPT_PREFIX}y", - style=1, - group_id="update_prompt", - ), - _make_callback_button( - btn_id="no", - label="✗ 取消", - visited_label="已取消", - data=f"{UPDATE_PROMPT_PREFIX}n", - style=0, - group_id="update_prompt", - ), - ]), - ] - ) - ) + return _single_row_keyboard([ + _make_callback_button("yes", "✓ 确认", "已确认", f"{UPDATE_PROMPT_PREFIX}y", 1, "update_prompt"), + _make_callback_button("no", "✗ 取消", "已取消", f"{UPDATE_PROMPT_PREFIX}n", 0, "update_prompt"), + ]) # ── ApprovalRequest + text builder ─────────────────────────────────── @dataclass class ApprovalRequest: - """Structured approval-request display data. + """Approval-request display data. - :param session_key: Routes the decision back to the waiting caller. - :param title: Short title at the top. - :param description: Optional longer description. - :param command_preview: Command text (exec approvals). - :param cwd: Working directory (exec approvals). - :param tool_name: Tool name (plugin approvals). - :param severity: ``'critical' | 'info' | ''``. - :param timeout_sec: Seconds until the approval expires. + ``command_preview`` / ``cwd`` are set for exec approvals, ``tool_name`` for + plugin approvals; ``severity`` is ``'critical' | 'info' | ''``. """ session_key: str title: str @@ -285,146 +168,48 @@ class ApprovalRequest: allow_permanent: bool = True +_SEVERITY_ICONS = {"critical": "🔴", "info": "🔵"} + + def build_approval_text(req: ApprovalRequest) -> str: """Render an :class:`ApprovalRequest` into the message body (markdown).""" if req.command_preview or req.cwd: - return _build_exec_text(req) - return _build_plugin_text(req) - - -def _build_exec_text(req: ApprovalRequest) -> str: - lines: List[str] = ["🔐 **命令执行审批**", ""] - if req.command_preview: - preview = req.command_preview[:300] - lines.append(f"```\n{preview}\n```") - if req.cwd: - lines.append(f"📁 目录: {req.cwd}") - if req.title and req.title != req.command_preview: - lines.append(f"📋 {req.title}") - if req.description: - lines.append(f"📝 {req.description}") - lines.append("") - lines.append(f"⏱️ 超时: {req.timeout_sec} 秒") + lines = ["🔐 **命令执行审批**", ""] + if req.command_preview: + lines.append(f"```\n{req.command_preview[:300]}\n```") + if req.cwd: + lines.append(f"📁 目录: {req.cwd}") + if req.title and req.title != req.command_preview: + lines.append(f"📋 {req.title}") + if req.description: + lines.append(f"📝 {req.description}") + else: + lines = [f"{_SEVERITY_ICONS.get(req.severity, '🟡')} **审批请求**", "", f"📋 {req.title}"] + if req.description: + lines.append(f"📝 {req.description}") + if req.tool_name: + lines.append(f"🔧 工具: {req.tool_name}") + lines += ["", f"⏱️ 超时: {req.timeout_sec} 秒"] return "\n".join(lines) -def _build_plugin_text(req: ApprovalRequest) -> str: - icon = ( - "🔴" if req.severity == "critical" - else "🔵" if req.severity == "info" - else "🟡" - ) - lines: List[str] = [f"{icon} **审批请求**", ""] - lines.append(f"📋 {req.title}") - if req.description: - lines.append(f"📝 {req.description}") - if req.tool_name: - lines.append(f"🔧 工具: {req.tool_name}") - lines.append("") - lines.append(f"⏱️ 超时: {req.timeout_sec} 秒") - return "\n".join(lines) - - -# ── ApprovalSender ─────────────────────────────────────────────────── - -PostMessageFn = Callable[..., Awaitable[Dict[str, Any]]] -"""Signature of an async POST to ``/v2/{users|groups}/{id}/messages``. - -Implementations accept a body dict and return the raw API response. -""" - - -class ApprovalSender: - """Send an approval-request message with an inline keyboard. - - Decoupled from the adapter via callables so it can be unit-tested in - isolation. Pass the adapter's ``_send_message_with_keyboard`` helper - (or any equivalent) as ``post_message``. - """ - - def __init__( - self, - post_c2c: PostMessageFn, - post_group: PostMessageFn, - log_tag: str = "QQBot", - ) -> None: - self._post_c2c = post_c2c - self._post_group = post_group - self._log_tag = log_tag - - async def send( - self, - chat_type: str, - chat_id: str, - req: ApprovalRequest, - msg_id: Optional[str] = None, - ) -> bool: - """Send an approval message to *chat_id*. - - :param chat_type: ``'c2c'`` or ``'group'``. - :param chat_id: User openid or group openid. - :param req: :class:`ApprovalRequest`. - :param msg_id: Reply-to message id (required for passive messages). - :returns: ``True`` on success, ``False`` on failure. - """ - text = build_approval_text(req) - keyboard = build_approval_keyboard(req.session_key) - - logger.info( - "[%s] Sending approval request to %s:%s (session=%.20s…)", - self._log_tag, chat_type, chat_id, req.session_key, - ) - - try: - if chat_type == "c2c": - await self._post_c2c(chat_id, text, msg_id, keyboard) - elif chat_type == "group": - await self._post_group(chat_id, text, msg_id, keyboard) - else: - logger.warning( - "[%s] Approval: unsupported chat_type %r", - self._log_tag, chat_type, - ) - return False - logger.info( - "[%s] Approval message sent to %s:%s", - self._log_tag, chat_type, chat_id, - ) - return True - except Exception as exc: - logger.error( - "[%s] Failed to send approval message to %s:%s: %s", - self._log_tag, chat_type, chat_id, exc, - ) - return False - - # ── INTERACTION_CREATE event shape ─────────────────────────────────── @dataclass class InteractionEvent: - """Parsed ``INTERACTION_CREATE`` event payload. + """Parsed ``INTERACTION_CREATE`` payload. See https://bot.q.qq.com/wiki/develop/api-v2/dev-prepare/interface-framework/event-emit.html """ - id: str = "" - """Interaction event id — required for the ``PUT /interactions/{id}`` ACK.""" - - type: int = 0 - """Event type code (``11`` = message button).""" - - chat_type: int = 0 - """``0`` = guild, ``1`` = group, ``2`` = c2c.""" - - scene: str = "" - """``'guild'`` | ``'group'`` | ``'c2c'`` — human-readable scene.""" - + id: str = "" # required for the ``PUT /interactions/{id}`` ACK + type: int = 0 # event type code (11 = message button) + chat_type: int = 0 # 0 = guild, 1 = group, 2 = c2c + scene: str = "" # 'guild' | 'group' | 'c2c' group_openid: str = "" group_member_openid: str = "" user_openid: str = "" channel_id: str = "" guild_id: str = "" - button_data: str = "" button_id: str = "" resolver_user_id: str = "" @@ -432,11 +217,10 @@ class InteractionEvent: @property def operator_openid(self) -> str: """Best available operator openid (group → member; c2c → user).""" - return ( - self.group_member_openid - or self.user_openid - or self.resolver_user_id - ) + return self.group_member_openid or self.user_openid or self.resolver_user_id + + +_SCENE_NAMES = {0: "guild", 1: "group", 2: "c2c"} def parse_interaction_event(raw: Dict[str, Any]) -> InteractionEvent: @@ -444,12 +228,11 @@ def parse_interaction_event(raw: Dict[str, Any]) -> InteractionEvent: data_raw = raw.get("data") or {} resolved = data_raw.get("resolved") or {} scene_code = int(raw.get("chat_type", 0) or 0) - scene = {0: "guild", 1: "group", 2: "c2c"}.get(scene_code, "") return InteractionEvent( id=str(raw.get("id", "")), type=int(data_raw.get("type", 0) or 0), chat_type=scene_code, - scene=scene, + scene=_SCENE_NAMES.get(scene_code, ""), group_openid=str(raw.get("group_openid", "")), group_member_openid=str(raw.get("group_member_openid", "")), user_openid=str(raw.get("user_openid", "")), diff --git a/gateway/platforms/qqbot/onboard.py b/gateway/platforms/qqbot/onboard.py index 6fd80c29fb..7fc1041079 100644 --- a/gateway/platforms/qqbot/onboard.py +++ b/gateway/platforms/qqbot/onboard.py @@ -1,14 +1,8 @@ -""" -QQBot scan-to-configure (QR code onboard) module. +"""QQBot scan-to-configure (QR code onboard) module. -Mirrors the Feishu onboarding pattern: synchronous HTTP + a single public -entry-point ``qr_register()`` that handles the full flow (create task → -display QR code → poll → decrypt credentials). - -Calls the ``q.qq.com`` ``create_bind_task`` / ``poll_bind_result`` APIs to -generate a QR-code URL and poll for scan completion. On success the caller -receives the bot's *app_id*, *client_secret* (decrypted locally), and the -scanner's *user_openid* — enough to fully configure the QQBot gateway. +Mirrors the Feishu onboarding pattern: synchronous HTTP + one public entry-point +``qr_register()`` (create task → display QR → poll → decrypt credentials) against +the ``q.qq.com`` ``create_bind_task`` / ``poll_bind_result`` APIs. Reference: https://bot.q.qq.com/wiki/develop/api-v2/ """ @@ -35,11 +29,6 @@ from .utils import get_api_headers logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Bind status -# --------------------------------------------------------------------------- - - class BindStatus(IntEnum): """Status codes returned by ``_poll_bind_result``.""" @@ -49,10 +38,6 @@ class BindStatus(IntEnum): EXPIRED = 3 -# --------------------------------------------------------------------------- -# QR rendering -# --------------------------------------------------------------------------- - try: import qrcode as _qrcode_mod except (ImportError, TypeError): @@ -64,10 +49,7 @@ def _render_qr(url: str) -> bool: if _qrcode_mod is None: return False try: - qr = _qrcode_mod.QRCode( - error_correction=_qrcode_mod.constants.ERROR_CORRECT_M, - border=2, - ) + qr = _qrcode_mod.QRCode(error_correction=_qrcode_mod.constants.ERROR_CORRECT_M, border=2) qr.add_data(url) qr.make(fit=True) qr.print_ascii(invert=True) @@ -76,62 +58,33 @@ def _render_qr(url: str) -> bool: return False -# --------------------------------------------------------------------------- -# Synchronous HTTP helpers (mirrors Feishu _post_registration pattern) -# --------------------------------------------------------------------------- +def _portal_post(path: str, payload: dict, timeout: float, fail_msg: str) -> dict: + """Synchronous POST to the portal host; raises RuntimeError on non-zero ``retcode``.""" + import httpx + + with httpx.Client(timeout=timeout, follow_redirects=True) as client: + resp = client.post(f"https://{PORTAL_HOST}{path}", json=payload, headers=get_api_headers()) + resp.raise_for_status() + data = resp.json() + if data.get("retcode") != 0: + raise RuntimeError(data.get("msg", fail_msg)) + return data def _create_bind_task(timeout: float = ONBOARD_API_TIMEOUT) -> Tuple[str, str]: - """Create a bind task and return *(task_id, aes_key_base64)*. - - Raises: - RuntimeError: If the API returns a non-zero ``retcode``. - """ - import httpx - - url = f"https://{PORTAL_HOST}{ONBOARD_CREATE_PATH}" + """Create a bind task and return *(task_id, aes_key_base64)*.""" key = generate_bind_key() - - with httpx.Client(timeout=timeout, follow_redirects=True) as client: - resp = client.post(url, json={"key": key}, headers=get_api_headers()) - resp.raise_for_status() - data = resp.json() - - if data.get("retcode") != 0: - raise RuntimeError(data.get("msg", "create_bind_task failed")) - + data = _portal_post(ONBOARD_CREATE_PATH, {"key": key}, timeout, "create_bind_task failed") task_id = (data.get("data") or {}).get("task_id") if not task_id: raise RuntimeError("create_bind_task: missing task_id in response") - logger.debug("create_bind_task ok: task_id=%s", task_id) return task_id, key -def _poll_bind_result( - task_id: str, - timeout: float = ONBOARD_API_TIMEOUT, -) -> Tuple[BindStatus, str, str, str]: - """Poll the bind result for *task_id*. - - Returns: - A 4-tuple of ``(status, bot_appid, bot_encrypt_secret, user_openid)``. - - Raises: - RuntimeError: If the API returns a non-zero ``retcode``. - """ - import httpx - - url = f"https://{PORTAL_HOST}{ONBOARD_POLL_PATH}" - - with httpx.Client(timeout=timeout, follow_redirects=True) as client: - resp = client.post(url, json={"task_id": task_id}, headers=get_api_headers()) - resp.raise_for_status() - data = resp.json() - - if data.get("retcode") != 0: - raise RuntimeError(data.get("msg", "poll_bind_result failed")) - +def _poll_bind_result(task_id: str, timeout: float = ONBOARD_API_TIMEOUT) -> Tuple[BindStatus, str, str, str]: + """Poll *task_id*; returns ``(status, bot_appid, bot_encrypt_secret, user_openid)``.""" + data = _portal_post(ONBOARD_POLL_PATH, {"task_id": task_id}, timeout, "poll_bind_result failed") d = data.get("data", {}) return ( BindStatus(d.get("status", 0)), @@ -146,27 +99,20 @@ def build_connect_url(task_id: str) -> str: return QR_URL_TEMPLATE.format(task_id=quote(task_id)) -# --------------------------------------------------------------------------- -# Public entry-point -# --------------------------------------------------------------------------- - _MAX_REFRESHES = 3 def qr_register(timeout_seconds: int = 600) -> Optional[dict]: """Run the QQBot scan-to-configure QR registration flow. - Mirrors ``feishu.qr_register()``: handles create → display → poll → - decrypt in one call. Unexpected errors propagate to the caller. + Unexpected errors propagate to the caller. - :returns: - ``{"app_id": ..., "client_secret": ..., "user_openid": ...}`` on - success, or ``None`` on failure / expiry / cancellation. + :returns: ``{"app_id", "client_secret", "user_openid"}`` on success, or + ``None`` on failure / expiry / cancellation. """ deadline = time.monotonic() + timeout_seconds for refresh_count in range(_MAX_REFRESHES + 1): - # ── Create bind task ── try: task_id, aes_key = _create_bind_task() except Exception as exc: @@ -174,8 +120,6 @@ def qr_register(timeout_seconds: int = 600) -> Optional[dict]: return None url = build_connect_url(task_id) - - # ── Display QR code + URL ── print() if _render_qr(url): print(f" Scan the QR code above, or open this URL directly:\n {url}") @@ -184,7 +128,6 @@ def qr_register(timeout_seconds: int = 600) -> Optional[dict]: print(" Tip: pip install qrcode to display a scannable QR code here") print() - # ── Poll loop ── while time.monotonic() < deadline: try: status, app_id, encrypted_secret, user_openid = _poll_bind_result(task_id) @@ -198,11 +141,7 @@ def qr_register(timeout_seconds: int = 600) -> Optional[dict]: print(f" QR scan complete! (App ID: {app_id})") if user_openid: print(f" Scanner's OpenID: {user_openid}") - return { - "app_id": app_id, - "client_secret": client_secret, - "user_openid": user_openid, - } + return {"app_id": app_id, "client_secret": client_secret, "user_openid": user_openid} if status == BindStatus.EXPIRED: if refresh_count >= _MAX_REFRESHES: @@ -213,7 +152,6 @@ def qr_register(timeout_seconds: int = 600) -> Optional[dict]: time.sleep(ONBOARD_POLL_INTERVAL) else: - # deadline reached without completing logger.warning("[QQBot onboard] Poll timed out after %ds", timeout_seconds) return None diff --git a/gateway/platforms/qqbot/utils.py b/gateway/platforms/qqbot/utils.py index 873e58d2a5..156fbc68cc 100644 --- a/gateway/platforms/qqbot/utils.py +++ b/gateway/platforms/qqbot/utils.py @@ -9,10 +9,6 @@ from typing import Any, Dict, List from .constants import QQBOT_VERSION -# --------------------------------------------------------------------------- -# User-Agent -# --------------------------------------------------------------------------- - def _get_hermes_version() -> str: """Return the hermes-agent package version, or 'dev' if unavailable.""" try: @@ -23,28 +19,19 @@ def _get_hermes_version() -> str: def build_user_agent() -> str: - """Build a descriptive User-Agent string. - - Format:: - - QQBotAdapter/ (Python/; ; Hermes/) - - Example:: - - QQBotAdapter/1.0.0 (Python/3.11.15; darwin; Hermes/0.9.0) - """ - py_version = f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}" - os_name = platform.system().lower() - hermes_version = _get_hermes_version() - return f"QQBotAdapter/{QQBOT_VERSION} (Python/{py_version}; {os_name}; Hermes/{hermes_version})" + """``QQBotAdapter/ (Python/; ; Hermes/)``.""" + v = sys.version_info + return ( + f"QQBotAdapter/{QQBOT_VERSION} (Python/{v.major}.{v.minor}.{v.micro}; " + f"{platform.system().lower()}; Hermes/{_get_hermes_version()})" + ) def get_api_headers() -> Dict[str, str]: - """Return standard HTTP headers for QQBot API requests. + """Standard QQBot API headers. - Includes ``Content-Type``, ``Accept``, and a dynamic ``User-Agent``. - ``q.qq.com`` requires ``Accept: application/json`` — without it, - the server returns a JavaScript anti-bot challenge page. + ``q.qq.com`` requires ``Accept: application/json`` — without it the server + returns a JavaScript anti-bot challenge page. """ return { "Content-Type": "application/json", @@ -53,15 +40,8 @@ def get_api_headers() -> Dict[str, str]: } -# --------------------------------------------------------------------------- -# Config helpers -# --------------------------------------------------------------------------- - def coerce_list(value: Any) -> List[str]: - """Coerce config values into a trimmed string list. - - Accepts comma-separated strings, lists, tuples, sets, or single values. - """ + """Coerce a comma-separated string / list / tuple / set / scalar into a trimmed string list.""" if value is None: return [] if isinstance(value, str): diff --git a/gateway/platforms/signal.py b/gateway/platforms/signal.py index 4e46f2b2b2..5a13feeaf0 100644 --- a/gateway/platforms/signal.py +++ b/gateway/platforms/signal.py @@ -1,10 +1,7 @@ """Signal messenger platform adapter. -Connects to a signal-cli daemon running in HTTP mode. -Inbound messages arrive via SSE (Server-Sent Events) streaming. -Outbound messages and actions use JSON-RPC 2.0 over HTTP. - -Based on PR #268 by ibhagwan, rebuilt with bug fixes. +Connects to a signal-cli daemon running in HTTP mode. Inbound messages arrive +via SSE streaming; outbound messages and actions use JSON-RPC 2.0 over HTTP. Requires: - signal-cli installed and running: signal-cli daemon --http 127.0.0.1:8080 @@ -13,6 +10,7 @@ Requires: import asyncio import base64 +import itertools import json import logging import os @@ -44,7 +42,7 @@ from gateway.platforms.base import ( utf16_len, ) from gateway.platforms.helpers import redact_phone -from gateway.platforms.media_cache import DEFAULT_EXT_TO_MIME, mime_for_ext +from gateway.platforms.media_cache import mime_for_ext from tools.audio_container import CONTAINER_TO_EXT, sniff_container from gateway.platforms.signal_format import markdown_to_signal from gateway.platforms.signal_rate_limit import ( @@ -58,24 +56,33 @@ from gateway.platforms.signal_rate_limit import ( _signal_send_timeout, get_scheduler, ) +from gateway.platforms._shared import get_scoped_secret as _sig_secret logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- SIGNAL_MAX_ATTACHMENT_SIZE = 100 * 1024 * 1024 # 100 MB MAX_MESSAGE_LENGTH = 8000 # Signal message size limit -TYPING_INTERVAL = 8.0 # seconds between typing indicator refreshes SSE_RETRY_DELAY_INITIAL = 2.0 SSE_RETRY_DELAY_MAX = 60.0 HEALTH_CHECK_INTERVAL = 30.0 # seconds between health checks HEALTH_CHECK_STALE_THRESHOLD = 120.0 # seconds without SSE activity before concern - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- +# Magic-byte prefixes checked before delegating to the shared audio/AV sniffer. +_MAGIC_EXTENSIONS = ( + (b"\x89PNG", ".png"), + (b"\xff\xd8", ".jpg"), + (b"GIF8", ".gif"), + (b"%PDF", ".pdf"), +) +_MEDIA_TYPE_BY_MIME_PREFIX = ( + ("audio/", MessageType.VOICE), + ("image/", MessageType.PHOTO), + ("video/", MessageType.VIDEO), +) +_OUTCOME_REACTION = {ProcessingOutcome.SUCCESS: "✅", ProcessingOutcome.FAILURE: "❌"} +_QUOTE_AUTHOR_KEYS = ( + "author", "authorNumber", "authorUuid", "authorAci", "authorServiceId", "authorServiceIdString", +) def _parse_comma_list(value: str) -> List[str]: @@ -86,28 +93,14 @@ def _parse_comma_list(value: str) -> List[str]: def _guess_extension(data: bytes) -> str: """Guess file extension from magic bytes. - Android Signal delivers voice notes as raw ADTS AAC frames, which share - the ``0xFF 0xFx`` sync word with MPEG-1/2 Layer 3 (MP3). The byte-1 - layout disambiguates: ADTS packs ``ID layer protection_absent`` into - bits 3-0, where ``ID`` is 0 for MPEG-2/4 AAC and ``layer`` is always - 0 for ADTS. A real MP3 frame has ``ID=1`` and ``layer`` in {1, 2, 3}. + WEBP is claimed before the shared audio/AV sniffer (it shares RIFF with WAVE); + tools/audio_container.py owns the MP3-vs-ADTS-AAC sync-word disambiguation. """ - if data[:4] == b"\x89PNG": - return ".png" - if data[:2] == b"\xff\xd8": - return ".jpg" - if data[:4] == b"GIF8": - return ".gif" + for magic, ext in _MAGIC_EXTENSIONS: + if data.startswith(magic): + return ext if len(data) >= 12 and data[:4] == b"RIFF" and data[8:12] == b"WEBP": return ".webp" - if data[:4] == b"%PDF": - return ".pdf" - # Audio/AV containers: delegate to the shared central sniffer - # (tools/audio_container.py) — ONE module owns magic-byte container - # detection. It handles the brand/form-type disambiguations this - # function used to carry locally: RIFF/WAVE vs WEBP (WEBP is claimed - # above, before delegation), ftyp audio brands ("M4A ", "M4B ") vs - # video brands (isom/mp42/avc1/qt), and MP3 vs ADTS AAC sync words. container = sniff_container(data) if container is not None: return CONTAINER_TO_EXT[container] @@ -124,38 +117,22 @@ def _is_audio_ext(ext: str) -> bool: return ext.lower() in {".mp3", ".wav", ".ogg", ".m4a", ".aac"} -# Historical Signal ext→mime table now lives in -# gateway.platforms.media_cache.DEFAULT_EXT_TO_MIME (byte-identical); -# kept as a module alias for backwards compatibility with any callers -# that referenced the private name. -_EXT_TO_MIME = DEFAULT_EXT_TO_MIME - - def _ext_to_mime(ext: str) -> str: - """Map file extension to MIME type.""" - # preserves historical signal mapping (shared table matches verbatim) + """Map file extension to MIME type (shared table matches Signal's historical map).""" return mime_for_ext(ext, fallback="application/octet-stream") def _remux_aac_to_m4a(aac_data: bytes) -> Optional[Tuple[bytes, str]]: - """Losslessly remux raw ADTS AAC bytes into an MP4 (.m4a) container. + """Losslessly remux raw ADTS AAC (Android voice notes, rejected by most STT APIs) to .m4a. - Used by the Signal attachment cache so Android voice notes land on disk - in a container that every major STT API (Groq, OpenAI, xAI, Mistral - Voxtral) will accept. ``ffmpeg -c:a copy`` is a single demux/remux — - no re-encode, no quality loss, sub-100ms for typical voice-note sizes. - - Returns ``(m4a_bytes, ".m4a")`` on success, or ``None`` if ffmpeg is - missing, input is invalid, or remux fails for any reason. Callers - must treat ``None`` as "pass through unchanged" and not raise. + Returns ``(m4a_bytes, ".m4a")`` or ``None`` when ffmpeg is missing or fails — + callers must then pass the input through unchanged. """ - ffmpeg = shutil.which("ffmpeg") - if not ffmpeg: - # Common Homebrew/local prefixes on macOS dev hosts. - for prefix in ("/opt/homebrew/bin/ffmpeg", "/usr/local/bin/ffmpeg"): - if os.path.isfile(prefix) and os.access(prefix, os.X_OK): - ffmpeg = prefix - break + # Fall back to common Homebrew/local prefixes on macOS dev hosts. + ffmpeg = shutil.which("ffmpeg") or next( + (p for p in ("/opt/homebrew/bin/ffmpeg", "/usr/local/bin/ffmpeg") + if os.path.isfile(p) and os.access(p, os.X_OK)), None, + ) if not ffmpeg: logger.debug("Signal: ffmpeg not found, skipping AAC→M4A remux") return None @@ -193,23 +170,15 @@ def _remux_aac_to_m4a(aac_data: bytes) -> Optional[Tuple[bytes, str]]: def _render_mentions(text: str, mentions: list) -> str: - """Replace Signal mention placeholders (\\uFFFC) with readable @identifiers. - - Signal encodes @mentions as the Unicode object replacement character - with out-of-band metadata containing the mentioned user's UUID/number. - """ + """Replace Signal mention placeholders (\\uFFFC + out-of-band start/length/number + metadata) with readable @identifiers. Replace from the end so indices hold.""" if not mentions or "\uFFFC" not in text: return text - # Sort mentions by start position (reverse) to replace from end to start - # so indices don't shift as we replace - sorted_mentions = sorted(mentions, key=lambda m: m.get("start", 0), reverse=True) - for mention in sorted_mentions: + for mention in sorted(mentions, key=lambda m: m.get("start", 0), reverse=True): start = mention.get("start", 0) length = mention.get("length", 1) - # Use the mention's number or UUID as the replacement identifier = mention.get("number") or mention.get("uuid") or "user" - replacement = f"@{identifier}" - text = text[:start] + replacement + text[start + length:] + text = text[:start] + f"@{identifier}" + text[start + length:] return text @@ -247,136 +216,72 @@ def validate_signal_config(config: PlatformConfig) -> bool: return bool(http_url and account) -# --------------------------------------------------------------------------- -# Signal Adapter -# --------------------------------------------------------------------------- - -def _sig_secret(name: str, default: str = "") -> str: - """Resolve a per-profile ``SIGNAL_*`` setting honoring the active secret - scope (#93522). - - Under ``gateway.multiplex_profiles`` every secondary profile is - constructed inside ``_profile_runtime_scope`` (``gateway/run.py``) and - its ``.env`` lives in that scope — raw ``os.getenv`` misses it and - leaks the default profile's values instead. The primary/active profile - is constructed without a scope and legitimately owns ``os.environ`` - (same canonical shape as QQ's ``_resolve_qq_secret``). - """ - from agent.secret_scope import UnscopedSecretError, get_secret - - try: - val = get_secret(name, default) - except UnscopedSecretError: - val = os.getenv(name) - return val if val is not None else default - - class SignalAdapter(BasePlatformAdapter): """Signal messenger adapter using signal-cli HTTP daemon.""" platform = Platform.SIGNAL MAX_MESSAGE_LENGTH = MAX_MESSAGE_LENGTH splits_long_messages = True # send() chunks after markdown → Signal formatting conversion - # Signal has no real edit API for already-sent messages. Mark it explicitly - # so streaming suppresses the visible cursor instead of leaving a stale tofu - # square behind in chat clients when edit attempts fail. + # Signal has no real edit API; declaring it lets streaming suppress the visible + # cursor instead of leaving a stale tofu square behind when edits fail. SUPPORTS_MESSAGE_EDITING = False def __init__(self, config: PlatformConfig): super().__init__(config, Platform.SIGNAL) - extra = config.extra or {} self.http_url = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/") self.account = extra.get("account", "") self.ignore_stories = extra.get("ignore_stories", True) - - # Parse allowlists — group policy is derived from presence of group allowlist - # Scoped reads (#93522): allowlists are per-profile authorization - # config; raw os.getenv misses secondary profiles' .env values and - # leaks the default profile's list into them. - group_allowed_str = _sig_secret("SIGNAL_GROUP_ALLOWED_USERS", "") - self.group_allow_from = set(_parse_comma_list(group_allowed_str)) - + # Allowlists are per-profile: scoped reads so secondary profiles don't inherit the + # default profile's list. Group policy derives from the group allowlist's presence. + self.group_allow_from = set(_parse_comma_list(_sig_secret("SIGNAL_GROUP_ALLOWED_USERS", ""))) # Mention filter — only respond in groups when the bot account is @mentioned. - # Read from config extra first, then SIGNAL_REQUIRE_MENTION env var. _rm_cfg = extra.get("require_mention") if _rm_cfg is not None: self.require_mention = bool(_rm_cfg) else: self.require_mention = os.getenv("SIGNAL_REQUIRE_MENTION", "false").lower() in ("true", "1", "yes", "on") - - # DM allowlist — mirrors SIGNAL_ALLOWED_USERS checked by run.py. - # Stored here so the reaction hooks can skip unauthorized senders - # (reactions fire before run.py's auth gate, so without this check - # every inbound DM from any contact gets a 👀 reaction). - # "*" means all users allowed (open mode); empty means no restriction - # recorded at adapter level (run.py still enforces auth separately). - dm_allowed_str = _sig_secret("SIGNAL_ALLOWED_USERS", "*") - self.dm_allow_from = set(_parse_comma_list(dm_allowed_str)) - - # HTTP client + # DM allowlist mirrors run.py's SIGNAL_ALLOWED_USERS check so the reaction hooks + # (which fire before run.py's auth gate) can skip unauthorized senders. "*" = open. + self.dm_allow_from = set(_parse_comma_list(_sig_secret("SIGNAL_ALLOWED_USERS", "*"))) self.client: Optional[httpx.AsyncClient] = None - - # Background tasks self._sse_task: Optional[asyncio.Task] = None self._health_monitor_task: Optional[asyncio.Task] = None self._typing_tasks: Dict[str, asyncio.Task] = {} - # Per-chat typing-indicator backoff. When signal-cli reports - # NETWORK_FAILURE (recipient offline / unroutable), base.py's - # _keep_typing refresh loop would otherwise hammer sendTyping every - # ~2s indefinitely, producing WARNING-level log spam and pointless - # RPC traffic. We track consecutive failures per chat and skip the - # RPC during a cooldown window instead. + # Per-chat typing-indicator backoff: when signal-cli reports NETWORK_FAILURE, + # base.py's _keep_typing loop would otherwise hammer sendTyping every ~2s. self._typing_failures: Dict[str, int] = {} self._typing_skip_until: Dict[str, float] = {} self._running = False self._last_sse_activity = 0.0 self._sse_response: Optional[httpx.Response] = None - - # Normalize account for self-message filtering self._account_normalized = self.account.strip() - - # Track recently sent message timestamps to prevent echo-back loops - # in Note to Self / self-chat mode and linked-device group sync-sents. - # OrderedDict[timestamp_ms -> insertion_monotonic_seconds] gives us - # LRU eviction (popitem(last=False) drops oldest) plus a TTL so that - # under chatty groups a still-pending echo cannot be evicted just - # because >50 outbounds happened. With a 5-minute TTL the cap only - # matters for runaway producers, not normal traffic bursts. + # Recently sent timestamps filter echo-backs (Note to Self / linked-device group + # sync-sents). LRU + TTL so a still-pending echo in a chatty group isn't evicted + # just because many outbounds happened; the cap only guards runaway producers. self._recent_sent_timestamps: "OrderedDict[int, float]" = OrderedDict() self._max_recent_timestamps = 512 self._recent_sent_ttl_seconds = 300.0 - # Keep a separate bounded cache of outbound Signal message timestamps. - # Signal quote.id is the timestamp of the quoted message, so this lets - # inbound replies identify that the user replied to a message sent by - # this bot even after the self-sync echo was filtered above. - # OrderedDict (not set) so the cap evicts the OLDEST timestamp in FIFO - # order — a plain set.pop() removes an arbitrary element, which could - # drop a still-recent timestamp and miss a genuine reply-to-own-message. + # Separate FIFO cache of outbound timestamps: Signal quote.id is the quoted + # message's timestamp, so replies to this bot are recognised even after the + # self-sync echo above was consumed. self._sent_message_timestamps: "OrderedDict[str, None]" = OrderedDict() self._max_sent_message_timestamps = 500 - # Signal increasingly exposes ACI/PNI UUIDs as stable recipient IDs. - # Keep a best-effort mapping so outbound sends can upgrade from a - # phone number to the corresponding UUID when signal-cli prefers it. + # Best-effort number↔ACI/PNI UUID mapping so outbound sends can upgrade a + # phone number to the UUID signal-cli prefers. self._recipient_uuid_by_number: Dict[str, str] = {} self._recipient_number_by_uuid: Dict[str, str] = {} self._recipient_cache_lock = asyncio.Lock() - logger.info("Signal adapter initialized: url=%s account=%s groups=%s", self.http_url, redact_phone(self.account), "enabled" if self.group_allow_from else "disabled") - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - async def connect(self, *, is_reconnect: bool = False) -> bool: """Connect to signal-cli daemon and start SSE listener.""" if not self.http_url or not self.account: logger.error("Signal: SIGNAL_HTTP_URL and SIGNAL_ACCOUNT are required") return False - - # Acquire scoped lock to prevent duplicate Signal listeners for the same phone + # Scoped lock prevents duplicate Signal listeners for the same phone. lock_acquired = False try: if not self._acquire_platform_lock('signal-phone', self.account, 'Signal account'): @@ -384,12 +289,10 @@ class SignalAdapter(BasePlatformAdapter): lock_acquired = True except Exception as e: logger.warning("Signal: Could not acquire phone lock (non-fatal): %s", e) - - # Tighter keepalive so idle CLOSE_WAIT drains promptly (#18451). + # Tighter keepalive so idle CLOSE_WAIT drains promptly. from gateway.platforms._http_client_limits import platform_httpx_limits self.client = httpx.AsyncClient(timeout=30.0, limits=platform_httpx_limits()) try: - # Health check — verify signal-cli daemon is reachable try: resp = await self.client.get(f"{self.http_url}/api/v1/check", timeout=10.0) if resp.status_code != 200: @@ -398,12 +301,10 @@ class SignalAdapter(BasePlatformAdapter): except Exception as e: logger.error("Signal: cannot reach signal-cli at %s: %s", self.http_url, e) return False - self._running = True self._last_sse_activity = time.time() self._sse_task = asyncio.create_task(self._sse_listener()) self._health_monitor_task = asyncio.create_task(self._health_monitor()) - logger.info("Signal: connected to %s", self.http_url) # Plugin-registered native handlers (ctx.register_platform_handler). self._wire_plugin_handlers(None) @@ -416,59 +317,43 @@ class SignalAdapter(BasePlatformAdapter): if lock_acquired: self._release_platform_lock() + @staticmethod + async def _cancel_task(task: Optional[asyncio.Task]) -> None: + if task: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + async def disconnect(self) -> None: """Stop SSE listener and clean up.""" self._running = False - - if self._sse_task: - self._sse_task.cancel() - try: - await self._sse_task - except asyncio.CancelledError: - pass - - if self._health_monitor_task: - self._health_monitor_task.cancel() - try: - await self._health_monitor_task - except asyncio.CancelledError: - pass - - # Cancel all typing tasks + for task in (self._sse_task, self._health_monitor_task): + await self._cancel_task(task) for task in self._typing_tasks.values(): task.cancel() self._typing_tasks.clear() - if self.client: await self.client.aclose() self.client = None - self._release_platform_lock() - logger.info("Signal: disconnected") - # ------------------------------------------------------------------ - # SSE Streaming (inbound messages) - # ------------------------------------------------------------------ - async def _sse_listener(self) -> None: """Listen for SSE events from signal-cli daemon.""" url = f"{self.http_url}/api/v1/events?account={quote(self.account, safe='')}" backoff = SSE_RETRY_DELAY_INITIAL - while self._running: try: logger.debug("Signal SSE: connecting to %s", url) async with self.client.stream( - "GET", url, - headers={"Accept": "text/event-stream"}, - timeout=None, + "GET", url, headers={"Accept": "text/event-stream"}, timeout=None, ) as response: self._sse_response = response backoff = SSE_RETRY_DELAY_INITIAL # Reset on successful connection self._last_sse_activity = time.time() logger.info("Signal SSE: connected") - buffer = "" async for chunk in response.aiter_text(): if not self._running: @@ -479,13 +364,11 @@ class SignalAdapter(BasePlatformAdapter): line = line.strip() if not line: continue - # SSE keepalive comments (":") prove the connection - # is alive — update activity so the health monitor - # doesn't report false idle warnings. + # Keepalive comments (":") prove the connection is alive — + # count them as activity so the health monitor stays quiet. if line.startswith(":"): self._last_sse_activity = time.time() continue - # Parse SSE data lines if line.startswith("data:"): data_str = line[5:].strip() if not data_str: @@ -498,7 +381,6 @@ class SignalAdapter(BasePlatformAdapter): logger.debug("Signal SSE: invalid JSON: %s", data_str[:100]) except Exception: logger.exception("Signal SSE: error handling event") - except asyncio.CancelledError: break except httpx.HTTPError as e: @@ -507,36 +389,26 @@ class SignalAdapter(BasePlatformAdapter): except Exception as e: if self._running: logger.warning("Signal SSE: error: %s (reconnecting in %.0fs)", e, backoff) - if self._running: - # Add 20% jitter to prevent thundering herd on reconnection + # 20% jitter prevents thundering herd on reconnection jitter = backoff * 0.2 * random.random() await asyncio.sleep(backoff + jitter) backoff = min(backoff * 2, SSE_RETRY_DELAY_MAX) - self._sse_response = None - # ------------------------------------------------------------------ - # Health Monitor - # ------------------------------------------------------------------ - async def _health_monitor(self) -> None: """Monitor SSE connection health and force reconnect if stale.""" while self._running: await asyncio.sleep(HEALTH_CHECK_INTERVAL) if not self._running: break - elapsed = time.time() - self._last_sse_activity if elapsed > HEALTH_CHECK_STALE_THRESHOLD: logger.warning("Signal: SSE idle for %.0fs, checking daemon health", elapsed) try: - resp = await self.client.get( - f"{self.http_url}/api/v1/check", timeout=10.0 - ) + resp = await self.client.get(f"{self.http_url}/api/v1/check", timeout=10.0) if resp.status_code == 200: - # Daemon is alive but SSE is idle — update activity to - # avoid repeated warnings (connection may just be quiet) + # Daemon alive but SSE quiet — reset activity to avoid repeated warnings self._last_sse_activity = time.time() logger.debug("Signal: daemon healthy, SSE idle") else: @@ -557,38 +429,80 @@ class SignalAdapter(BasePlatformAdapter): pass self._sse_response = None - # ------------------------------------------------------------------ - # Message Handling - # ------------------------------------------------------------------ + def _unwrap_sync_message(self, envelope_data: dict) -> Optional[dict]: + """Promote a genuine "Note to Self" / group sync-sent to a dataMessage envelope; + None for other sync events (read receipts, typing, our own outbound echoes).""" + sync_msg = envelope_data.get("syncMessage") + if not sync_msg or not isinstance(sync_msg, dict): + return None + sent_msg = sync_msg.get("sentMessage") + if not sent_msg or not isinstance(sent_msg, dict): + return None + dest = sent_msg.get("destinationNumber") or sent_msg.get("destination") + sent_msg_group_info = sent_msg.get("groupInfo") or {} + sent_msg_group_id = sent_msg_group_info.get("groupId") if sent_msg_group_info else None + if dest != self._account_normalized and not sent_msg_group_id: + return None + if self._consume_sent_timestamp(sent_msg.get("timestamp")): + return None # echo of our own outbound reply + return {**envelope_data, "dataMessage": sent_msg} + + def _apply_group_mention_rules(self, text: str, data_message: dict) -> Tuple[bool, str]: + """Gate on require_mention (False = drop) and strip the bot's own @mention. + + The self-mention is stripped from every group message so the agent doesn't + read "@+155****4567 say hello" as a directive to contact that number. + """ + account_norm = self._account_normalized + if self.require_mention: + mentioned_in_text = account_norm and (f"@{account_norm}" in (text or "")) + mentioned_in_metadata = any( + m.get("number") == account_norm or m.get("uuid") == account_norm + for m in (data_message.get("mentions") or []) + ) + if not mentioned_in_text and not mentioned_in_metadata: + logger.debug("Signal: ignoring group message (require_mention=true, bot not mentioned)") + return False, text + if text and account_norm: + text = text.replace(f"@{account_norm}", "") + bot_uuid = self._recipient_uuid_by_number.get(account_norm) + if bot_uuid: + text = text.replace(f"@{bot_uuid}", "") + # Collapse only the doubled space the removal introduced; intentional + # newlines in multi-line messages are preserved. + text = text.replace(" ", " ").strip() + return True, text + + async def _collect_attachments(self, attachments_data: list) -> Tuple[List[str], List[str]]: + """Fetch + cache inbound attachments; returns (media_urls, media_types).""" + media_urls: List[str] = [] + media_types: List[str] = [] + for att in attachments_data: + att_id = att.get("id") + att_size = att.get("size", 0) + if not att_id: + continue + if att_size > SIGNAL_MAX_ATTACHMENT_SIZE: + logger.warning("Signal: attachment too large (%d bytes), skipping", att_size) + continue + try: + cached_path, ext = await self._fetch_attachment(att_id) + if cached_path: + media_urls.append(cached_path) + media_types.append(att.get("contentType") or _ext_to_mime(ext)) + except Exception: + logger.exception("Signal: failed to fetch attachment %s", att_id) + return media_urls, media_types async def _handle_envelope(self, envelope: dict) -> None: """Process an incoming signal-cli envelope.""" - # Unwrap nested envelope if present envelope_data = envelope.get("envelope", envelope) - - # Handle syncMessage: extract "Note to Self" messages (sent to own account) - # while still filtering other sync events (read receipts, typing, etc.) is_note_to_self = False if "syncMessage" in envelope_data: - sync_msg = envelope_data.get("syncMessage") - if sync_msg and isinstance(sync_msg, dict): - sent_msg = sync_msg.get("sentMessage") - if sent_msg and isinstance(sent_msg, dict): - dest = sent_msg.get("destinationNumber") or sent_msg.get("destination") - sent_ts = sent_msg.get("timestamp") - sent_msg_group_info = sent_msg.get("groupInfo") or {} - sent_msg_group_id = sent_msg_group_info.get("groupId") if sent_msg_group_info else None - if dest == self._account_normalized or sent_msg_group_id: - # Check if this is an echo of our own outbound reply - if self._consume_sent_timestamp(sent_ts): - return - # Genuine user Note to Self — promote to dataMessage - is_note_to_self = True - envelope_data = {**envelope_data, "dataMessage": sent_msg} - if not is_note_to_self: + envelope_data = self._unwrap_sync_message(envelope_data) + if envelope_data is None: return - - # Extract sender info + is_note_to_self = True sender = ( envelope_data.get("sourceNumber") or envelope_data.get("sourceUuid") @@ -597,38 +511,26 @@ class SignalAdapter(BasePlatformAdapter): sender_name = envelope_data.get("sourceName", "") sender_uuid = envelope_data.get("sourceUuid", "") self._remember_recipient_identifiers(sender, sender_uuid) - if not sender: logger.debug("Signal: ignoring envelope with no sender") return - - # Self-message filtering — prevent reply loops (but allow Note to Self) + # Self-message filtering prevents reply loops (Note to Self is allowed) if self._account_normalized and sender == self._account_normalized and not is_note_to_self: return - - # Filter stories if self.ignore_stories and envelope_data.get("storyMessage"): return - - # Get data message — also check editMessage (edited messages contain - # their updated dataMessage inside editMessage.dataMessage) + # Edited messages carry their updated dataMessage inside editMessage data_message = ( envelope_data.get("dataMessage") or (envelope_data.get("editMessage") or {}).get("dataMessage") ) if not data_message: return - - # Check for group message group_info = data_message.get("groupInfo") group_id = group_info.get("groupId") if group_info else None is_group = bool(group_id) - - # Group message filtering — derived from SIGNAL_GROUP_ALLOWED_USERS: - # - No env var set → groups disabled (default safe behavior) - # - Env var set with group IDs → only those groups allowed - # - Env var set with "*" → all groups allowed - # DM auth is fully handled by run.py (_is_user_authorized) + # Group policy derives from SIGNAL_GROUP_ALLOWED_USERS: unset → groups disabled; + # IDs → only those groups; "*" → all. DM auth is run.py's (_is_user_authorized). if is_group: if not self.group_allow_from: logger.debug("Signal: ignoring group message (no SIGNAL_GROUP_ALLOWED_USERS)") @@ -636,103 +538,35 @@ class SignalAdapter(BasePlatformAdapter): if "*" not in self.group_allow_from and group_id not in self.group_allow_from: logger.debug("Signal: group %s not in allowlist", group_id[:8] if group_id else "?") return - - # Build chat info chat_id = sender if not is_group else f"group:{group_id}" chat_type = "group" if is_group else "dm" - - # Extract text and render mentions text = data_message.get("message", "") mentions = data_message.get("mentions", []) if text and mentions: text = _render_mentions(text, mentions) - - # Mention filter: in groups, only process messages that @mention the bot account - if is_group and self.require_mention: - account_norm = self._account_normalized - # Check rendered mention tags OR raw mention metadata - mentioned_in_text = account_norm and ( - f"@{account_norm}" in (text or "") - ) - mentioned_in_metadata = any( - m.get("number") == account_norm or m.get("uuid") == account_norm - for m in (data_message.get("mentions") or []) - ) - if not mentioned_in_text and not mentioned_in_metadata: - logger.debug( - "Signal: ignoring group message (require_mention=true, bot not mentioned)" - ) + if is_group: + mentioned, text = self._apply_group_mention_rules(text, data_message) + if not mentioned: return - - # Strip the bot's own @mention from any group message so the agent - # doesn't misinterpret "@+155****4567 say hello" as a directive to - # contact that phone number. _render_mentions replaces the Signal - #  placeholder with @, which looks like an - # addressee to the LLM rather than a self-reference. Applies to every - # group (not just require_mention groups) so the self-mention is - # cleaned wherever it appears. - if is_group and text: - account_norm = self._account_normalized - if account_norm: - text = text.replace(f"@{account_norm}", "") - # Also strip if the mention was rendered using the bot's UUID - bot_uuid = self._recipient_uuid_by_number.get(account_norm) - if bot_uuid: - text = text.replace(f"@{bot_uuid}", "") - # Tidy the spacing the removed mention left behind: collapse the - # double-space at a mid-sentence removal and trim the ends. - # Only touches the doubled space the removal introduced, so - # intentional newlines in a multi-line message are preserved. - text = text.replace(" ", " ").strip() - - # Extract quote (reply-to) context from Signal dataMessage. Signal's - # quote.id is the timestamp of the quoted message; quote.author points - # at the quoted sender when available. Preserve both so the gateway can - # tell the agent when the user replied to a specific assistant message. + # Signal's quote.id is the quoted message's timestamp; quote.author the quoted + # sender. Preserve both so the gateway can tell the agent which message the + # user replied to. quote_data = data_message.get("quote") or {} reply_to_id = str(quote_data.get("id")) if quote_data.get("id") else None - reply_to_text = quote_data.get("text") reply_to_author = self._extract_quote_author(quote_data) - reply_to_author_name = quote_data.get("authorName") or quote_data.get("authorProfileName") - reply_to_is_own = self._quote_references_own_message(reply_to_id, reply_to_author) - - # Process attachments attachments_data = data_message.get("attachments", []) - media_urls = [] - media_types = [] - + media_urls: List[str] = [] + media_types: List[str] = [] if attachments_data and not getattr(self, "ignore_attachments", False): - for att in attachments_data: - att_id = att.get("id") - att_size = att.get("size", 0) - if not att_id: - continue - if att_size > SIGNAL_MAX_ATTACHMENT_SIZE: - logger.warning("Signal: attachment too large (%d bytes), skipping", att_size) - continue - try: - cached_path, ext = await self._fetch_attachment(att_id) - if cached_path: - # Use contentType from Signal if available, else map from extension - content_type = att.get("contentType") or _ext_to_mime(ext) - media_urls.append(cached_path) - media_types.append(content_type) - except Exception: - logger.exception("Signal: failed to fetch attachment %s", att_id) - - # Skip envelopes with no meaningful content (no text, no attachments). - # Catches profile key updates, empty messages, and other metadata-only - # envelopes that still carry a dataMessage wrapper but have nothing - # worth processing. See issue: signal-cli logs "Profile key update" + - # Hermes receives msg='' triggering a full agent turn for nothing. + media_urls, media_types = await self._collect_attachments(attachments_data) + # Skip contentless envelopes (profile key updates, empty messages) that still + # carry a dataMessage wrapper — otherwise msg='' triggers a full agent turn. if (not text or not text.strip()) and not media_urls: logger.debug( "Signal: skipping contentless envelope from %s (%d attachments)", redact_phone(sender), len(media_urls) if media_urls else 0, ) return - - # Build session source source = self.build_source( chat_id=chat_id, chat_name=group_info.get("groupName") if group_info else sender_name, @@ -742,36 +576,25 @@ class SignalAdapter(BasePlatformAdapter): user_id_alt=sender_uuid if sender_uuid else None, chat_id_alt=group_id if is_group else None, ) - - # Determine message type from media + # First matching MIME prefix wins; everything else (application/*, text/*, + # unknown) is a DOCUMENT so run.py's document-context injection surfaces the + # cached path to the agent. msg_type = MessageType.TEXT if media_types: - if any(mt.startswith("audio/") for mt in media_types): - msg_type = MessageType.VOICE - elif any(mt.startswith("image/") for mt in media_types): - msg_type = MessageType.PHOTO - elif any(mt.startswith("video/") for mt in media_types): - msg_type = MessageType.VIDEO - else: - # Catch-all: application/*, text/*, and unknown MIME types are - # treated as documents so run.py's document-context injection - # surfaces the cached file path to the agent (same pattern as - # WhatsApp/Slack/BlueBubbles/Mattermost). - msg_type = MessageType.DOCUMENT - - # Parse timestamp from envelope data (milliseconds since epoch) - ts_ms = envelope_data.get("timestamp", 0) + msg_type = next( + (mt for prefix, mt in _MEDIA_TYPE_BY_MIME_PREFIX + if any(m.startswith(prefix) for m in media_types)), + MessageType.DOCUMENT, + ) + ts_ms = envelope_data.get("timestamp", 0) # milliseconds since epoch + timestamp = datetime.now(tz=timezone.utc) if ts_ms: try: timestamp = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) except (ValueError, OSError): - timestamp = datetime.now(tz=timezone.utc) - else: - timestamp = datetime.now(tz=timezone.utc) - - # Build and dispatch event. - # Store raw envelope data in raw_message so on_processing_start/complete - # can extract targetAuthor + targetTimestamp for sendReaction. + pass + # raw_message keeps sender + timestamp_ms so the processing hooks can build + # sendReaction targets. event = MessageEvent( source=source, text=text or "", @@ -785,15 +608,13 @@ class SignalAdapter(BasePlatformAdapter): "quote": quote_data if quote_data else None, }, reply_to_message_id=reply_to_id, - reply_to_text=reply_to_text, + reply_to_text=quote_data.get("text"), reply_to_author_id=reply_to_author, - reply_to_author_name=reply_to_author_name, - reply_to_is_own_message=reply_to_is_own, + reply_to_author_name=quote_data.get("authorName") or quote_data.get("authorProfileName"), + reply_to_is_own_message=self._quote_references_own_message(reply_to_id, reply_to_author), ) - logger.debug("Signal: message from %s in %s: %s", redact_phone(sender), chat_id[:20], (text or "")[:50]) - await self.handle_message(event) def _remember_recipient_identifiers(self, number: Optional[str], service_id: Optional[str]) -> None: @@ -808,24 +629,13 @@ class SignalAdapter(BasePlatformAdapter): """Return the best available Signal sender identifier from quote metadata.""" if not isinstance(quote_data, dict): return None - for key in ( - "author", - "authorNumber", - "authorUuid", - "authorAci", - "authorServiceId", - "authorServiceIdString", - ): + for key in _QUOTE_AUTHOR_KEYS: value = quote_data.get(key) if value: return str(value) return None - def _quote_references_own_message( - self, - reply_to_id: Optional[str], - reply_to_author: Optional[str], - ) -> bool: + def _quote_references_own_message(self, reply_to_id: Optional[str], reply_to_author: Optional[str]) -> bool: """True when a Signal quote points at this adapter's outbound message.""" if reply_to_id and str(reply_to_id) in self._sent_message_timestamps: return True @@ -845,11 +655,9 @@ class SignalAdapter(BasePlatformAdapter): if timestamp is None: return key = str(timestamp) - # Re-insert to mark most-recently-used so eviction drops genuinely old - # timestamps, not a recently re-seen one. + # Re-insert to mark most-recently-used so eviction drops genuinely old entries. self._sent_message_timestamps.pop(key, None) self._sent_message_timestamps[key] = None - # FIFO-evict the oldest entry once over the cap. while len(self._sent_message_timestamps) > self._max_sent_message_timestamps: self._sent_message_timestamps.popitem(last=False) @@ -857,19 +665,15 @@ class SignalAdapter(BasePlatformAdapter): """Best-effort extraction of a Signal service ID from listContacts output.""" if not isinstance(contact, dict): return None - - number = contact.get("number") - recipient = contact.get("recipient") service_id = contact.get("uuid") or contact.get("serviceId") if not service_id: profile = contact.get("profile") if isinstance(profile, dict): service_id = profile.get("serviceId") or profile.get("uuid") - - if service_id and _is_signal_service_id(service_id): - matches_number = number == phone_number or recipient == phone_number - if matches_number: - return service_id + if service_id and _is_signal_service_id(service_id) and ( + contact.get("number") == phone_number or contact.get("recipient") == phone_number + ): + return service_id return None async def _resolve_recipient(self, chat_id: str) -> str: @@ -881,79 +685,60 @@ class SignalAdapter(BasePlatformAdapter): or not _looks_like_e164_number(chat_id) ): return chat_id - cached = self._recipient_uuid_by_number.get(chat_id) if cached: return cached - async with self._recipient_cache_lock: cached = self._recipient_uuid_by_number.get(chat_id) if cached: return cached - - contacts = await self._rpc("listContacts", { - "account": self.account, - "allRecipients": True, - }) + contacts = await self._rpc("listContacts", {"account": self.account, "allRecipients": True}) if isinstance(contacts, list): for contact in contacts: number = contact.get("number") if isinstance(contact, dict) else None service_id = self._extract_contact_uuid(contact, chat_id) if number and service_id: self._remember_recipient_identifiers(number, service_id) - return self._recipient_uuid_by_number.get(chat_id, chat_id) - # ------------------------------------------------------------------ - # Attachment Handling - # ------------------------------------------------------------------ + async def _with_target(self, params: Dict[str, Any], chat_id: str, *, resolve: bool = True) -> Dict[str, Any]: + """Add the groupId / recipient routing key for *chat_id* to *params* (in place).""" + if chat_id.startswith("group:"): + params["groupId"] = chat_id[6:] + elif resolve: + params["recipient"] = [await self._resolve_recipient(chat_id)] + else: + params["recipient"] = [chat_id] + return params async def _fetch_attachment(self, attachment_id: str) -> tuple: """Fetch an attachment via JSON-RPC and cache it. Returns (path, ext).""" - result = await self._rpc("getAttachment", { - "account": self.account, - "id": attachment_id, - }) - + result = await self._rpc("getAttachment", {"account": self.account, "id": attachment_id}) if not result: return None, "" - - # Handle dict response (signal-cli returns {"data": "base64..."}) + # signal-cli returns {"data": "base64..."} if isinstance(result, dict): result = result.get("data") if not result: logger.warning("Signal: attachment response missing 'data' key") return None, "" - - # Result is base64-encoded file content raw_data = base64.b64decode(result) ext = _guess_extension(raw_data) - - # Android Signal voice notes are raw ADTS AAC streams. Most STT - # providers (Groq Whisper, OpenAI Whisper) reject raw ADTS — they - # require AAC to be muxed into an MP4 container. Remux losslessly - # with ``ffmpeg -c:a copy`` so the cached file is a normal .m4a. - # No re-encode, sub-100ms on a Pi 5. Graceful no-op if ffmpeg is - # absent: the raw ADTS file is cached as-is and STT may reject it - # (there is no downstream sniff-and-remux fallback). + # Android voice notes are raw ADTS AAC, which Whisper-style STT rejects; remux + # losslessly to .m4a. If ffmpeg is absent the raw file is cached as-is (there is + # no downstream sniff-and-remux fallback). if ext == ".aac": remuxed: Optional[Tuple[bytes, str]] = await asyncio.to_thread(_remux_aac_to_m4a, raw_data) if remuxed is not None: raw_data, ext = remuxed - if _is_image_ext(ext): path = cache_image_from_bytes(raw_data, ext) elif _is_audio_ext(ext): path = cache_audio_from_bytes(raw_data, ext) else: path = cache_document_from_bytes(raw_data, ext) - return path, ext - # ------------------------------------------------------------------ - # JSON-RPC Communication - # ------------------------------------------------------------------ - async def _rpc( self, method: str, @@ -966,99 +751,57 @@ class SignalAdapter(BasePlatformAdapter): ) -> Any: """Send a JSON-RPC 2.0 request to signal-cli daemon. - When ``log_failures=False``, error and exception paths log at DEBUG - instead of WARNING — used by the typing-indicator path to silence - repeated NETWORK_FAILURE spam for unreachable recipients while - still preserving visibility for the first occurrence and for - unrelated RPCs. - - When ``raise_on_rate_limit=True``, a Signal ``[429]`` / - ``RateLimitException`` response raises ``SignalRateLimitError`` - instead of being swallowed — lets callers (multi-attachment send) - opt into backoff-retry without changing default behaviour. + ``log_failures=False`` logs failures at DEBUG (typing path: silence repeated + NETWORK_FAILURE spam). ``raise_on_rate_limit=True`` raises ``SignalRateLimitError`` + on a 429 / RateLimitException instead of swallowing it (multi-attachment sends). """ if not self.client: logger.warning("Signal: RPC called but client not connected") return None - if rpc_id is None: rpc_id = f"{method}_{int(time.time() * 1000)}" - - payload = { - "jsonrpc": "2.0", - "method": method, - "params": params, - "id": rpc_id, - } - + payload = {"jsonrpc": "2.0", "method": method, "params": params, "id": rpc_id} + fail_level = logging.WARNING if log_failures else logging.DEBUG try: - resp = await self.client.post( - f"{self.http_url}/api/v1/rpc", - json=payload, - timeout=timeout, - ) + resp = await self.client.post(f"{self.http_url}/api/v1/rpc", json=payload, timeout=timeout) resp.raise_for_status() data = resp.json() - if "error" in data: err = data["error"] - if raise_on_rate_limit: - if _is_signal_rate_limit_error(err): - err_msg = str(err.get("message", "")) if isinstance(err, dict) else str(err) - retry_after = _extract_retry_after_seconds(err) - raise SignalRateLimitError(err_msg, retry_after=retry_after) - if log_failures: - logger.warning("Signal RPC error (%s): %s", method, err) - else: - logger.debug("Signal RPC error (%s): %s", method, err) + if raise_on_rate_limit and _is_signal_rate_limit_error(err): + err_msg = str(err.get("message", "")) if isinstance(err, dict) else str(err) + raise SignalRateLimitError(err_msg, retry_after=_extract_retry_after_seconds(err)) + logger.log(fail_level, "Signal RPC error (%s): %s", method, err) return None - result = data.get("result") if isinstance(result, dict) and raise_on_rate_limit: results = result.get("results") if isinstance(results, list): for r in results: if isinstance(r, dict) and r.get("type") == "RATE_LIMIT_FAILURE": - retry_after = r.get("retryAfterSeconds") - raise SignalRateLimitError("Rate limit exceeded for recipient", retry_after=retry_after) - + raise SignalRateLimitError( + "Rate limit exceeded for recipient", retry_after=r.get("retryAfterSeconds") + ) return result - except SignalRateLimitError: raise except Exception as e: - if log_failures: - logger.warning("Signal RPC %s failed: %s", method, e) - else: - logger.debug("Signal RPC %s failed: %s", method, e) + logger.log(fail_level, "Signal RPC %s failed: %s", method, e) return None - # ------------------------------------------------------------------ - # Formatting — markdown → Signal body ranges - # ------------------------------------------------------------------ - @staticmethod def _markdown_to_signal(text: str) -> tuple[str, list[str]]: """Backward-compatible wrapper around shared Signal formatting helper.""" return markdown_to_signal(text) def format_message(self, content: str) -> str: - """Strip markdown for plain-text fallback (used by base class). - - The actual rich formatting happens in send() via _markdown_to_signal(). - """ - # This is only called if someone uses the base-class send path. - # Our send() override bypasses this entirely. + """Plain-text fallback for the base-class send path; send() applies rich styles itself.""" return content def _validate_send_result(self, result: Any) -> tuple[bool, Optional[str]]: - """Validate signal-cli send response results. - - Returns (success, error_message). - """ + """Validate signal-cli send response results. Returns (success, error_message).""" if not result or not isinstance(result, dict): return True, None - results = result.get("results") if isinstance(results, list): for r in results: @@ -1068,32 +811,16 @@ class SignalAdapter(BasePlatformAdapter): if rtype and rtype != "SUCCESS": return False, str(rtype) if "success" in r and not r.get("success"): - fail = r.get("failure") - if fail: - return False, str(fail) - return False, "Recipient delivery failed" + return False, str(r.get("failure") or "Recipient delivery failed") return True, None - # ------------------------------------------------------------------ - # Sending - # ------------------------------------------------------------------ - @staticmethod def _utf16_offsets(text: str) -> list[int]: """Return cumulative UTF-16 offsets for every Python character boundary.""" - offsets = [0] - total = 0 - for char in text: - total += utf16_len(char) - offsets.append(total) - return offsets + return [0, *itertools.accumulate(utf16_len(char) for char in text)] @staticmethod - def _styles_for_chunk( - text_styles: list[str], - chunk_start: int, - chunk_end: int, - ) -> list[str]: + def _styles_for_chunk(text_styles: list[str], chunk_start: int, chunk_end: int) -> list[str]: """Translate full-message Signal styles into a chunk-local range list.""" adjusted: list[str] = [] for style_string in text_styles: @@ -1107,60 +834,56 @@ class SignalAdapter(BasePlatformAdapter): overlap_start = max(style_start, chunk_start) overlap_end = min(style_end, chunk_end) if overlap_start < overlap_end: - adjusted.append( - f"{overlap_start - chunk_start}:{overlap_end - overlap_start}:{style_type}" - ) + adjusted.append(f"{overlap_start - chunk_start}:{overlap_end - overlap_start}:{style_type}") return adjusted @classmethod def _split_signal_formatted_message( - cls, - plain_text: str, - text_styles: list[str], - max_length: int, + cls, plain_text: str, text_styles: list[str], max_length: int, ) -> list[tuple[str, list[str]]]: - """Split formatted Signal text without breaking markdown before conversion. + """Split converted Signal text into chunks, translating body ranges per chunk. - ``markdown_to_signal`` emits UTF-16 body ranges for the fully converted - plain text. Split that plain text afterwards and translate overlapping - ranges into each chunk so styles that cross a chunk boundary are - preserved instead of leaking literal Markdown markers. + Splitting after conversion (not before) keeps styles that cross a chunk boundary + intact instead of leaking literal Markdown markers. """ if utf16_len(plain_text) <= max_length: return [(plain_text, text_styles)] - indicator_reserve = 10 # Mirrors BasePlatformAdapter.truncate_message(). body_limit = max(1, max_length - indicator_reserve) offsets = cls._utf16_offsets(plain_text) - chunks: list[tuple[str, list[str], int, int]] = [] + chunks: list[tuple[str, list[str]]] = [] start_idx = 0 total_u16 = offsets[-1] while offsets[start_idx] < total_u16: - start_u16 = offsets[start_idx] - end_budget = min(total_u16, start_u16 + body_limit) + end_budget = min(total_u16, offsets[start_idx] + body_limit) end_idx = start_idx + 1 while end_idx < len(offsets) and offsets[end_idx] <= end_budget: end_idx += 1 end_idx -= 1 if end_idx <= start_idx: end_idx = start_idx + 1 - chunk_start = offsets[start_idx] - chunk_end = offsets[end_idx] - chunk_text = plain_text[start_idx:end_idx] - chunk_styles = cls._styles_for_chunk(text_styles, chunk_start, chunk_end) - chunks.append((chunk_text, chunk_styles, chunk_start, chunk_end)) + chunk_styles = cls._styles_for_chunk(text_styles, offsets[start_idx], offsets[end_idx]) + chunks.append((plain_text[start_idx:end_idx], chunk_styles)) start_idx = end_idx - if len(chunks) == 1: - chunk_text, chunk_styles, _, _ = chunks[0] - return [(chunk_text, chunk_styles)] - + return chunks total = len(chunks) return [ (f"{chunk_text} ({idx}/{total})", chunk_styles) - for idx, (chunk_text, chunk_styles, _, _) in enumerate(chunks, start=1) + for idx, (chunk_text, chunk_styles) in enumerate(chunks, start=1) ] + async def _rpc_send(self, params: Dict[str, Any], fail_error: str) -> Tuple[Any, Optional[SendResult]]: + """Run a ``send`` RPC, validate and track it; ``(result, None)`` or ``(None, failed SendResult)``.""" + result = await self._rpc("send", params) + if result is None: + return None, SendResult(success=False, error=fail_error) + success, err_msg = self._validate_send_result(result) + if not success: + return None, SendResult(success=False, error=err_msg, raw_response=result) + self._track_sent_timestamp(result) + return result, None + async def send( self, chat_id: str, @@ -1172,49 +895,26 @@ class SignalAdapter(BasePlatformAdapter): await self._stop_typing_indicator(chat_id) if not content or not content.strip(): return SendResult(success=True, message_id=None) - - base_params: Dict[str, Any] = {"account": self.account} - if chat_id.startswith("group:"): - base_params["groupId"] = chat_id[6:] - else: - base_params["recipient"] = [await self._resolve_recipient(chat_id)] - + base_params = await self._with_target({"account": self.account}, chat_id) plain_message, message_styles = self._markdown_to_signal(content) - chunks = self._split_signal_formatted_message( - plain_message, - message_styles, - self.MAX_MESSAGE_LENGTH, - ) + chunks = self._split_signal_formatted_message(plain_message, message_styles, self.MAX_MESSAGE_LENGTH) last_result = None - for idx, (plain_text, text_styles) in enumerate(chunks, start=1): params: Dict[str, Any] = dict(base_params, message=plain_text) - if text_styles: if len(text_styles) == 1: params["textStyle"] = text_styles[0] else: params["textStyles"] = text_styles - logger.info( "[Signal] Sending response chunk %d/%d (%d chars) to %s", - idx, - len(chunks), - len(plain_text), - chat_id, + idx, len(chunks), len(plain_text), chat_id, ) - result = await self._rpc("send", params) - if result is None: - return SendResult(success=False, error="RPC send failed") - success, err_msg = self._validate_send_result(result) - if not success: - return SendResult(success=False, error=err_msg, raw_response=result) - self._track_sent_timestamp(result) - last_result = result - - # Signal has no editable message identifier. Returning None keeps the - # stream consumer on the non-edit fallback path instead of pretending - # future edits can remove an in-progress cursor from the chat thread. + last_result, err = await self._rpc_send(params, "RPC send failed") + if err: + return err + # Signal has no editable message identifier; message_id=None keeps the stream + # consumer on the non-edit fallback path. return SendResult(success=True, message_id=None, raw_response=last_result) def _track_sent_timestamp(self, rpc_result) -> None: @@ -1226,15 +926,14 @@ class SignalAdapter(BasePlatformAdapter): # Re-insert to mark as most-recently-used. self._recent_sent_timestamps.pop(ts, None) self._recent_sent_timestamps[ts] = now - # Drop entries older than TTL first (cheap O(k) where k=expired). + # Drop entries older than TTL first, then enforce the hard cap. cutoff = now - self._recent_sent_ttl_seconds while self._recent_sent_timestamps: - oldest_ts, oldest_at = next(iter(self._recent_sent_timestamps.items())) + _, oldest_at = next(iter(self._recent_sent_timestamps.items())) if oldest_at < cutoff: self._recent_sent_timestamps.popitem(last=False) else: break - # Hard cap as a last-resort guard against runaway producers. while len(self._recent_sent_timestamps) > self._max_recent_timestamps: self._recent_sent_timestamps.popitem(last=False) @@ -1246,58 +945,44 @@ class SignalAdapter(BasePlatformAdapter): return False async def send_typing(self, chat_id: str, metadata=None) -> None: - """Send a typing indicator. + """Send a typing indicator (called every ~2s by base.py's ``_keep_typing``). - base.py's ``_keep_typing`` refresh loop calls this every ~2s while - the agent is processing. If signal-cli returns NETWORK_FAILURE for - this recipient (offline, unroutable, group membership lost, etc.) - the unmitigated behaviour is: a WARNING log every 2 seconds for as - long as the agent keeps running. Instead we: - - - silence the WARNING after the first consecutive failure (subsequent - attempts log at DEBUG) so transport issues are still visible once - but don't flood the log, - - skip the RPC entirely during an exponential cooldown window once - three consecutive failures have happened, so we stop hammering - signal-cli with requests it can't deliver. - - A successful sendTyping clears the counters. + On NETWORK_FAILURE only the first consecutive failure logs at WARNING, and after + three failures the RPC is skipped for an exponential cooldown; success resets. """ now = time.monotonic() - skip_until = self._typing_skip_until.get(chat_id, 0.0) - if now < skip_until: + if now < self._typing_skip_until.get(chat_id, 0.0): return - - params: Dict[str, Any] = { - "account": self.account, - } - - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [await self._resolve_recipient(chat_id)] - + params = await self._with_target({"account": self.account}, chat_id) fails = self._typing_failures.get(chat_id, 0) - result = await self._rpc( - "sendTyping", - params, - rpc_id="typing", - log_failures=(fails == 0), - ) - + result = await self._rpc("sendTyping", params, rpc_id="typing", log_failures=(fails == 0)) if result is None: fails += 1 self._typing_failures[chat_id] = fails - # After 3 consecutive failures, back off exponentially (16s, - # 32s, 60s cap) to stop spamming signal-cli for a recipient - # that clearly isn't reachable right now. + # After 3 consecutive failures back off exponentially (16s, 32s, 60s cap). if fails >= 3: - backoff = min(60.0, 16.0 * (2 ** (fails - 3))) - self._typing_skip_until[chat_id] = now + backoff + self._typing_skip_until[chat_id] = now + min(60.0, 16.0 * (2 ** (fails - 3))) else: self._typing_failures.pop(chat_id, None) self._typing_skip_until.pop(chat_id, None) + async def _resolve_image_path(self, image_url: str) -> Tuple[Optional[str], Optional[str], Any]: + """Resolve an http(s):// or file:// image URL to ``(path, None, None)``, or + ``(None, reason, detail)`` with reason download (exc) / missing / oversize (size).""" + if image_url.startswith("file://"): + file_path = unquote(image_url[7:]) + else: + try: + file_path = await cache_image_from_url(image_url) + except Exception as e: + return None, "download", e + if not file_path or not Path(file_path).exists(): + return None, "missing", None + file_size = Path(file_path).stat().st_size + if file_size > SIGNAL_MAX_ATTACHMENT_SIZE: + return None, "oversize", file_size + return file_path, None, None + async def send_multiple_images( self, chat_id: str, @@ -1307,179 +992,106 @@ class SignalAdapter(BasePlatformAdapter): ) -> None: """Send a batch of images via chunked Signal RPC calls. - Per-image alt texts are dropped — Signal's send RPC only carries - one shared message body. Bad images (download failure, missing - file, oversize) are skipped with a warning so one bad URL - doesn't lose the rest of the batch. ``human_delay`` is ignored: - the rate-limit scheduler handles inter-batch pacing. + Alt texts are dropped (one shared body per send). Bad images are skipped with a + warning. ``human_delay`` is ignored: the rate-limit scheduler paces batches. """ if not images: return - scheduler = get_scheduler() logger.info( - "Signal send_multiple_images: received %d image(s) for %s — " - "scheduler state: %s", + "Signal send_multiple_images: received %d image(s) for %s — scheduler state: %s", len(images), chat_id[:30], scheduler.state(), ) - await self._stop_typing_indicator(chat_id) - attachments: List[str] = [] - skipped_download = 0 - skipped_missing = 0 - skipped_oversize = 0 + skipped = {"download": 0, "missing": 0, "oversize": 0} for image_url, _alt_text in images: - if image_url.startswith("file://"): - file_path = unquote(image_url[7:]) - else: - try: - file_path = await cache_image_from_url(image_url) - except Exception as e: - logger.warning("Signal: failed to download image %s: %s", image_url, e) - skipped_download += 1 - continue - - if not file_path or not Path(file_path).exists(): + file_path, reason, detail = await self._resolve_image_path(image_url) + if reason == "download": + logger.warning("Signal: failed to download image %s: %s", image_url, detail) + elif reason == "missing": logger.warning("Signal: image file not found for %s", image_url) - skipped_missing += 1 + elif reason == "oversize": + logger.warning("Signal: image too large (%d bytes), skipping %s", detail, image_url) + if reason: + skipped[reason] += 1 continue - - file_size = Path(file_path).stat().st_size - if file_size > SIGNAL_MAX_ATTACHMENT_SIZE: - logger.warning( - "Signal: image too large (%d bytes), skipping %s", file_size, image_url - ) - skipped_oversize += 1 - continue - attachments.append(file_path) - if not attachments: logger.error( - "Signal: no valid images in batch of %d " - "(download=%d missing=%d oversize=%d)", - len(images), skipped_download, skipped_missing, skipped_oversize, + "Signal: no valid images in batch of %d (download=%d missing=%d oversize=%d)", + len(images), skipped["download"], skipped["missing"], skipped["oversize"], ) return - logger.info( "Signal send_multiple_images: %d/%d images valid, sending in chunks", len(attachments), len(images), ) - - base_params: Dict[str, Any] = { - "account": self.account, - "message": "", - } - if chat_id.startswith("group:"): - base_params["groupId"] = chat_id[6:] - else: - base_params["recipient"] = [await self._resolve_recipient(chat_id)] - + base_params = await self._with_target({"account": self.account, "message": ""}, chat_id) att_batches = [ attachments[i:i + SIGNAL_MAX_ATTACHMENTS_PER_MSG] for i in range(0, len(attachments), SIGNAL_MAX_ATTACHMENTS_PER_MSG) ] - - for idx, att_batch in enumerate(att_batches): + n_batches = len(att_batches) + for idx, att_batch in enumerate(att_batches, start=1): n = len(att_batch) estimated = scheduler.estimate_wait(n) - logger.debug( - "Signal batch %d/%d: %d attachments, estimated wait=%.1fs", - idx + 1, len(att_batches), n, estimated, - ) + logger.debug("Signal batch %d/%d: %d attachments, estimated wait=%.1fs", idx, n_batches, n, estimated) if estimated >= SIGNAL_BATCH_PACING_NOTICE_THRESHOLD: - await self._notify_batch_pacing( - chat_id, idx + 1, len(att_batches), estimated + await self._notify_batch_pacing(chat_id, idx, n_batches, estimated) + await self._send_attachment_batch( + scheduler, dict(base_params, attachments=att_batch), n, f"{idx}/{n_batches}", + ) + + async def _send_attachment_batch(self, scheduler, params: Dict[str, Any], n: int, label: str) -> None: + """Send one attachment batch with rate-limit pacing and a single transient retry. + + Tokens are deducted only on validated success (a None result means the server + never accepted the batch); 429s feed the scheduler before the retry. + """ + send_timeout = _signal_send_timeout(n) + for attempt in range(1, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS + 1): + await scheduler.acquire(n) + try: + _rpc_t0 = time.monotonic() + result = await self._rpc("send", params, raise_on_rate_limit=True, timeout=send_timeout) + _rpc_duration = time.monotonic() - _rpc_t0 + success, err_msg = self._validate_send_result(result) if result is not None else (False, None) + if success: + self._track_sent_timestamp(result) + await scheduler.report_rpc_duration(_rpc_duration, n) + logger.info( + "Signal batch %s: %d attachments sent in %.1fs (attempt %d/%d)", + label, n, _rpc_duration, attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, + ) + return + logger.error( + "Signal: RPC send failed for batch %s (%d attachments, attempt %d/%d, rpc_duration=%.1fs)%s", + label, n, attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, _rpc_duration, + f": {err_msg}" if result is not None else "", + ) + if attempt >= SIGNAL_RATE_LIMIT_MAX_ATTEMPTS: + return + backoff = 2.0 ** attempt + logger.info("Signal: retrying batch %s after %.1fs backoff", label, backoff) + await asyncio.sleep(backoff) + except SignalRateLimitError as e: + scheduler.feedback(e.retry_after, n) + retry_after = f"{e.retry_after:.0f}s" if e.retry_after else "unknown" + if attempt >= SIGNAL_RATE_LIMIT_MAX_ATTEMPTS: + logger.error( + "Signal: rate-limit retries exhausted on batch %s (%d attachments lost, server retry_after=%s)", + label, n, retry_after, + ) + return + logger.warning( + "Signal: rate-limited on batch %s (attempt %d/%d, server retry_after=%s); " + "scheduler will pace the retry", + label, attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, retry_after, ) - params = dict(base_params, attachments=att_batch) - send_timeout = _signal_send_timeout(n) - - for attempt in range(1, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS + 1): - await scheduler.acquire(n) - try: - _rpc_t0 = time.monotonic() - result = await self._rpc( - "send", params, raise_on_rate_limit=True, timeout=send_timeout, - ) - _rpc_duration = time.monotonic() - _rpc_t0 - if result is not None: - success, err_msg = self._validate_send_result(result) - if success: - self._track_sent_timestamp(result) - await scheduler.report_rpc_duration(_rpc_duration, n) - logger.info( - "Signal batch %d/%d: %d attachments sent in %.1fs " - "(attempt %d/%d)", - idx + 1, len(att_batches), n, _rpc_duration, - attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, - ) - else: - logger.error( - "Signal: RPC send failed for batch %d/%d (%d attachments, " - "attempt %d/%d, rpc_duration=%.1fs): %s", - idx + 1, len(att_batches), n, - attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, - _rpc_duration, err_msg, - ) - # Retry transient (non-rate-limit) failures once - if attempt < SIGNAL_RATE_LIMIT_MAX_ATTEMPTS: - backoff = 2.0 ** attempt - logger.info( - "Signal: retrying batch %d/%d after %.1fs backoff", - idx + 1, len(att_batches), backoff, - ) - await asyncio.sleep(backoff) - continue - else: - # Assume the server didn't accept the batch, don't deduce tokens - logger.error( - "Signal: RPC send failed for batch %d/%d (%d attachments, " - "attempt %d/%d, rpc_duration=%.1fs)", - idx + 1, len(att_batches), n, - attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, - _rpc_duration, - ) - # Retry transient (non-rate-limit) failures once - if attempt < SIGNAL_RATE_LIMIT_MAX_ATTEMPTS: - backoff = 2.0 ** attempt - logger.info( - "Signal: retrying batch %d/%d after %.1fs backoff", - idx + 1, len(att_batches), backoff, - ) - await asyncio.sleep(backoff) - continue - break - except SignalRateLimitError as e: - scheduler.feedback(e.retry_after, n) - if attempt >= SIGNAL_RATE_LIMIT_MAX_ATTEMPTS: - logger.error( - "Signal: rate-limit retries exhausted on batch %d/%d " - "(%d attachments lost, server retry_after=%s)", - idx + 1, len(att_batches), n, - f"{e.retry_after:.0f}s" if e.retry_after else "unknown", - ) - break - logger.warning( - "Signal: rate-limited on batch %d/%d " - "(attempt %d/%d, server retry_after=%s); " - "scheduler will pace the retry", - idx + 1, len(att_batches), - attempt, SIGNAL_RATE_LIMIT_MAX_ATTEMPTS, - f"{e.retry_after:.0f}s" if e.retry_after else "unknown", - ) - - async def _notify_batch_pacing( - self, - chat_id: str, - next_batch_idx: int, - total_batches: int, - wait_s: float, - ) -> None: - """Inform the user when an inter-batch pacing wait crosses the - notice threshold. Best-effort; logs and continues on failure.""" + async def _notify_batch_pacing(self, chat_id: str, next_batch_idx: int, total_batches: int, wait_s: float) -> None: + """Tell the user about an inter-batch pacing wait over the notice threshold (best-effort).""" try: await self.send( chat_id, @@ -1489,282 +1101,115 @@ class SignalAdapter(BasePlatformAdapter): except Exception as e: logger.warning("Signal: failed to send pacing notice: %s", e) - async def send_image( - self, - chat_id: str, - image_url: str, - caption: Optional[str] = None, - **kwargs, - ) -> SendResult: + async def send_image(self, chat_id: str, image_url: str, caption: Optional[str] = None, **kwargs) -> SendResult: """Send an image. Supports http(s):// and file:// URLs.""" await self._stop_typing_indicator(chat_id) - - # Resolve image to local path - if image_url.startswith("file://"): - file_path = unquote(image_url[7:]) - else: - # Download remote image to cache - try: - file_path = await cache_image_from_url(image_url) - except Exception as e: - logger.warning("Signal: failed to download image: %s", e) - return SendResult(success=False, error=str(e)) - - if not file_path or not Path(file_path).exists(): + file_path, reason, detail = await self._resolve_image_path(image_url) + if reason == "download": + logger.warning("Signal: failed to download image: %s", detail) + return SendResult(success=False, error=str(detail)) + if reason == "missing": return SendResult(success=False, error="Image file not found") + if reason == "oversize": + return SendResult(success=False, error=f"Image too large ({detail} bytes)") + return await self._send_file(chat_id, file_path, caption, "RPC send with attachment failed") - # Validate size - file_size = Path(file_path).stat().st_size - if file_size > SIGNAL_MAX_ATTACHMENT_SIZE: - return SendResult(success=False, error=f"Image too large ({file_size} bytes)") - - params: Dict[str, Any] = { - "account": self.account, - "message": caption or "", - "attachments": [file_path], - } - - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [await self._resolve_recipient(chat_id)] - - result = await self._rpc("send", params) - if result is not None: - success, err_msg = self._validate_send_result(result) - if not success: - return SendResult(success=False, error=err_msg, raw_response=result) - self._track_sent_timestamp(result) - return SendResult(success=True) - return SendResult(success=False, error="RPC send with attachment failed") + async def _send_file(self, chat_id: str, file_path: str, caption: Optional[str], fail_error: str) -> SendResult: + """Send one local file as a Signal attachment via the ``send`` RPC.""" + params = await self._with_target( + {"account": self.account, "message": caption or "", "attachments": [file_path]}, chat_id + ) + _, err = await self._rpc_send(params, fail_error) + return err or SendResult(success=True) async def _send_attachment( - self, - chat_id: str, - file_path: str, - media_label: str, - caption: Optional[str] = None, + self, chat_id: str, file_path: str, media_label: str, caption: Optional[str] = None, ) -> SendResult: - """Send any file as a Signal attachment via RPC. - - Shared implementation for send_document, send_image_file, send_voice, - and send_video — avoids duplicating the validation/routing/RPC logic. - """ + """Send any local file as a Signal attachment (shared by send_document/image_file/voice/video).""" await self._stop_typing_indicator(chat_id) - try: file_size = Path(file_path).stat().st_size except FileNotFoundError: return SendResult(success=False, error=f"{media_label} file not found: {file_path}") - if file_size > SIGNAL_MAX_ATTACHMENT_SIZE: return SendResult(success=False, error=f"{media_label} too large ({file_size} bytes)") - - params: Dict[str, Any] = { - "account": self.account, - "message": caption or "", - "attachments": [file_path], - } - - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [await self._resolve_recipient(chat_id)] - - result = await self._rpc("send", params) - if result is not None: - success, err_msg = self._validate_send_result(result) - if not success: - return SendResult(success=False, error=err_msg, raw_response=result) - self._track_sent_timestamp(result) - return SendResult(success=True) - return SendResult(success=False, error=f"RPC send {media_label.lower()} failed") + return await self._send_file(chat_id, file_path, caption, f"RPC send {media_label.lower()} failed") async def send_document( - self, - chat_id: str, - file_path: str, - caption: Optional[str] = None, - filename: Optional[str] = None, - **kwargs, + self, chat_id: str, file_path: str, caption: Optional[str] = None, filename: Optional[str] = None, **kwargs, ) -> SendResult: """Send a document/file attachment.""" return await self._send_attachment(chat_id, file_path, "File", caption) async def send_image_file( - self, - chat_id: str, - image_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, + self, chat_id: str, image_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs, ) -> SendResult: - """Send a local image file as a native Signal attachment. - - Called by the gateway media delivery flow when MEDIA: tags containing - image paths are extracted from agent responses. - """ + """Send a local image file as a native Signal attachment (gateway MEDIA: delivery path).""" return await self._send_attachment(chat_id, image_path, "Image", caption) async def send_voice( - self, - chat_id: str, - audio_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, + self, chat_id: str, audio_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs, ) -> SendResult: - """Send an audio file as a Signal attachment. - - Signal does not distinguish voice messages from file attachments at - the API level, so this routes through the same RPC send path. - """ + """Send an audio file as a Signal attachment (Signal has no distinct voice-message API).""" return await self._send_attachment(chat_id, audio_path, "Audio", caption) async def send_video( - self, - chat_id: str, - video_path: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, + self, chat_id: str, video_path: str, caption: Optional[str] = None, reply_to: Optional[str] = None, **kwargs, ) -> SendResult: """Send a video file as a Signal attachment.""" return await self._send_attachment(chat_id, video_path, "Video", caption) - # ------------------------------------------------------------------ - # Typing Indicators - # ------------------------------------------------------------------ - async def _stop_typing_indicator(self, chat_id: str) -> None: """Stop a typing indicator loop for a chat.""" - task = self._typing_tasks.pop(chat_id, None) - if task: - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - - # Send an explicit stop-typing RPC so the recipient's device drops the - # indicator immediately instead of waiting for Signal's ~5s built-in - # timeout. Failures are best-effort — the backoff state must still be - # cleared so the next agent turn starts clean. + await self._cancel_task(self._typing_tasks.pop(chat_id, None)) + # Explicit stop-typing RPC so the recipient drops the indicator now instead of + # after Signal's ~5s built-in timeout. Best-effort: any RPC or recipient- + # resolution failure must not prevent the backoff cleanup below. try: - params: Dict[str, Any] = {"account": self.account} - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [await self._resolve_recipient(chat_id)] + params = await self._with_target({"account": self.account}, chat_id) params["stop"] = True - await self._rpc( - "sendTyping", - params, - rpc_id="typing-stop", - log_failures=False, - ) + await self._rpc("sendTyping", params, rpc_id="typing-stop", log_failures=False) except Exception: - # Best-effort: any RPC failure (or recipient-resolution failure) - # must not prevent backoff cleanup. pass - self._typing_failures.pop(chat_id, None) self._typing_skip_until.pop(chat_id, None) async def stop_typing(self, chat_id: str) -> None: - """Public interface for stopping typing — called by base adapter's - _keep_typing finally block to clean up platform-level typing tasks.""" + """Public stop-typing hook called from the base adapter's _keep_typing finally block.""" await self._stop_typing_indicator(chat_id) - # ------------------------------------------------------------------ - # Reactions - # ------------------------------------------------------------------ + async def _send_reaction_rpc(self, chat_id: str, params: Dict[str, Any]) -> bool: + """Route a ``sendReaction`` RPC to *chat_id* (no UUID upgrade — author IDs come from the envelope).""" + await self._with_target(params, chat_id, resolve=False) + return await self._rpc("sendReaction", params) is not None - async def send_reaction( - self, - chat_id: str, - emoji: str, - target_author: str, - target_timestamp: int, - ) -> bool: - """Send a reaction emoji to a specific message via signal-cli RPC. - - Args: - chat_id: The chat (phone number or "group:") - emoji: Reaction emoji string (e.g. "👀", "✅") - target_author: Phone number / UUID of the message author - target_timestamp: Signal timestamp (ms) of the message to react to - """ - params: Dict[str, Any] = { - "account": self.account, - "emoji": emoji, - "targetAuthor": target_author, + async def send_reaction(self, chat_id: str, emoji: str, target_author: str, target_timestamp: int) -> bool: + """React to the message (author number/UUID, Signal ms timestamp) via signal-cli RPC.""" + ok = await self._send_reaction_rpc(chat_id, { + "account": self.account, "emoji": emoji, "targetAuthor": target_author, "targetTimestamp": target_timestamp, - } + }) + if not ok: + logger.debug("Signal: sendReaction failed (chat=%s, emoji=%s)", chat_id[:20], emoji) + return ok - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [chat_id] - - result = await self._rpc("sendReaction", params) - if result is not None: - return True - logger.debug("Signal: sendReaction failed (chat=%s, emoji=%s)", chat_id[:20], emoji) - return False - - async def remove_reaction( - self, - chat_id: str, - target_author: str, - target_timestamp: int, - ) -> bool: + async def remove_reaction(self, chat_id: str, target_author: str, target_timestamp: int) -> bool: """Remove a reaction by sending an empty-string emoji.""" - params: Dict[str, Any] = { - "account": self.account, - "emoji": "", - "targetAuthor": target_author, - "targetTimestamp": target_timestamp, - "remove": True, - } - - if chat_id.startswith("group:"): - params["groupId"] = chat_id[6:] - else: - params["recipient"] = [chat_id] - - result = await self._rpc("sendReaction", params) - return result is not None - - # ------------------------------------------------------------------ - # Processing Lifecycle Hooks (reactions as progress indicators) - # ------------------------------------------------------------------ + return await self._send_reaction_rpc(chat_id, { + "account": self.account, "emoji": "", "targetAuthor": target_author, + "targetTimestamp": target_timestamp, "remove": True, + }) def _extract_reaction_target(self, event: MessageEvent) -> Optional[tuple]: - """Extract (target_author, target_timestamp) from a MessageEvent. - - Returns None if the event doesn't carry the raw Signal envelope data - needed for sendReaction. - """ + """Extract (target_author, target_timestamp) from a MessageEvent, or None.""" raw = event.raw_message - if not isinstance(raw, dict): - return None - author = raw.get("sender") - ts = raw.get("timestamp_ms") - if not author or not ts: - return None - return (author, ts) + if isinstance(raw, dict) and raw.get("sender") and raw.get("timestamp_ms"): + return (raw["sender"], raw["timestamp_ms"]) + return None def _reactions_enabled(self, event: "MessageEvent" = None) -> bool: - """Check if message reactions are enabled for this event. - - Two gates: - 1. SIGNAL_REACTIONS env var — set to false/0/no to disable globally. - 2. DM allowlist — if SIGNAL_ALLOWED_USERS is set, only react to - messages from senders in that list. This prevents unauthorized - contacts from seeing the 👀 reaction (which fires before run.py's - auth gate and would otherwise reveal that a bot is listening). - """ + """SIGNAL_REACTIONS env gate, then the DM allowlist: reactions fire before run.py's + auth gate, so an unauthorized contact's 👀 would otherwise reveal a listening bot.""" if os.getenv("SIGNAL_REACTIONS", "true").lower() in {"false", "0", "no"}: return False if event is not None: @@ -1782,11 +1227,7 @@ class SignalAdapter(BasePlatformAdapter): await self.send_reaction(event.source.chat_id, "👀", *target) async def on_processing_complete(self, event: MessageEvent, outcome: "ProcessingOutcome") -> None: - """Swap the 👀 reaction for ✅ (success) or ❌ (failure). - - On CANCELLED we leave the 👀 in place — no terminal outcome means - the reaction should keep reflecting "in progress" (matches Telegram). - """ + """Swap 👀 for ✅/❌; on CANCELLED the 👀 stays to keep reflecting "in progress" (matches Telegram).""" if not self._reactions_enabled(event): return if outcome == ProcessingOutcome.CANCELLED: @@ -1795,38 +1236,17 @@ class SignalAdapter(BasePlatformAdapter): if not target: return chat_id = event.source.chat_id - # Remove the in-progress reaction, then add the final one await self.remove_reaction(chat_id, *target) - if outcome == ProcessingOutcome.SUCCESS: - await self.send_reaction(chat_id, "✅", *target) - elif outcome == ProcessingOutcome.FAILURE: - await self.send_reaction(chat_id, "❌", *target) - - # ------------------------------------------------------------------ - # Chat Info - # ------------------------------------------------------------------ + emoji = _OUTCOME_REACTION.get(outcome) + if emoji: + await self.send_reaction(chat_id, emoji, *target) async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: """Get information about a chat/contact.""" if chat_id.startswith("group:"): - return { - "name": chat_id, - "type": "group", - "chat_id": chat_id, - } - - # Try to resolve contact name - result = await self._rpc("getContact", { - "account": self.account, - "contactAddress": chat_id, - }) - + return {"name": chat_id, "type": "group", "chat_id": chat_id} + result = await self._rpc("getContact", {"account": self.account, "contactAddress": chat_id}) name = chat_id if result and isinstance(result, dict): name = result.get("name") or result.get("profileName") or chat_id - - return { - "name": name, - "type": "dm", - "chat_id": chat_id, - } + return {"name": name, "type": "dm", "chat_id": chat_id} diff --git a/gateway/platforms/signal_format.py b/gateway/platforms/signal_format.py index e8539549bf..bf152b454c 100644 --- a/gateway/platforms/signal_format.py +++ b/gateway/platforms/signal_format.py @@ -1,95 +1,79 @@ """Shared Signal formatting helpers. -Keep markdown → Signal native formatting conversion in one place so both the -live Signal adapter and standalone send paths emit the same bodyRanges. +Markdown → Signal native formatting lives here so both the live adapter and the +standalone send paths emit the same bodyRanges. """ from __future__ import annotations import re +_CODE_BLOCK_RE = re.compile(r"```[a-zA-Z0-9_+-]*\n?(.*?)```", re.DOTALL) +_HEADING_RE = re.compile(r"^#{1,6}\s+", re.MULTILINE) +_INLINE_PATTERNS = [ + (re.compile(r"\*\*(.+?)\*\*", re.DOTALL), "BOLD"), + (re.compile(r"__(.+?)__", re.DOTALL), "BOLD"), + (re.compile(r"~~(.+?)~~", re.DOTALL), "STRIKETHROUGH"), + (re.compile(r"`(.+?)`"), "MONOSPACE"), + (re.compile(r"(? int: + """Length of *s* in UTF-16 code units.""" + return len(s.encode("utf-16-le")) // 2 + + +def _normalize_bullet_markers(source: str) -> str: + """Replace Markdown bullet markers with plain Unicode bullets. + + Signal renders ``- item`` / ``* item`` literally. Fenced code blocks are kept + byte-for-byte; list-looking lines inside code are code, not prose bullets. + """ + parts = re.split(r"(```.*?```)", source, flags=re.DOTALL) + for idx, part in enumerate(parts): + if idx % 2 == 0: + parts[idx] = re.sub(r"(?m)^([ \t]{0,3})[-*+]\s+", r"\1• ", part) + return "".join(parts) + def markdown_to_signal(text: str) -> tuple[str, list[str]]: """Convert markdown to plain text + Signal textStyles list. - Signal doesn't render markdown. Instead it uses ``bodyRanges`` (exposed by - signal-cli as ``textStyle`` / ``textStyles`` params) with the format - ``start:length:STYLE``. - - Positions are measured in UTF-16 code units because that's what the Signal - protocol uses. - + Signal uses ``bodyRanges`` (signal-cli ``textStyle`` / ``textStyles`` params) in + the form ``start:length:STYLE``, positions in UTF-16 code units. Supported styles: BOLD, ITALIC, STRIKETHROUGH, MONOSPACE. """ - - def _utf16_len(s: str) -> int: - """Length of *s* in UTF-16 code units.""" - return len(s.encode("utf-16-le")) // 2 - - def _normalize_bullet_markers(source: str) -> str: - """Replace Markdown bullet markers with plain Unicode bullets. - - Signal does not render Markdown list syntax, so ``- item`` and - ``* item`` otherwise arrive as literal Markdown markers. Preserve - fenced code blocks byte-for-byte; list-looking lines inside code are - code, not prose bullets. - """ - parts = re.split(r"(```.*?```)", source, flags=re.DOTALL) - for idx, part in enumerate(parts): - if idx % 2 == 1: - continue - parts[idx] = re.sub(r"(?m)^([ \t]{0,3})[-*+]\s+", r"\1• ", part) - return "".join(parts) - - text = re.sub(r"\n{3,}", "\n\n", text) - text = text.strip() + text = re.sub(r"\n{3,}", "\n\n", text).strip() text = _normalize_bullet_markers(text) - styles: list[tuple[int, int, str]] = [] - - code_block = re.compile(r"```[a-zA-Z0-9_+-]*\n?(.*?)```", re.DOTALL) - while match := code_block.search(text): + while match := _CODE_BLOCK_RE.search(text): inner = match.group(1).rstrip("\n") start = match.start() text = text[: match.start()] + inner + text[match.end() :] styles.append((start, len(inner), "MONOSPACE")) - - heading = re.compile(r"^#{1,6}\s+", re.MULTILINE) new_text = "" last_end = 0 - for match in heading.finditer(text): + for match in _HEADING_RE.finditer(text): new_text += text[last_end : match.start()] - last_end = match.end() eol = text.find("\n", match.end()) if eol == -1: eol = len(text) heading_text = text[match.end() : eol] - start = len(new_text) + styles.append((len(new_text), len(heading_text), "BOLD")) new_text += heading_text - styles.append((start, len(heading_text), "BOLD")) last_end = eol - new_text += text[last_end:] - text = new_text - - patterns = [ - (re.compile(r"\*\*(.+?)\*\*", re.DOTALL), "BOLD"), - (re.compile(r"__(.+?)__", re.DOTALL), "BOLD"), - (re.compile(r"~~(.+?)~~", re.DOTALL), "STRIKETHROUGH"), - (re.compile(r"`(.+?)`"), "MONOSPACE"), - (re.compile(r"(? os for os, oe in occupied): all_matches.append((ms, me, match.start(1), match.end(1), style)) occupied.append((ms, me)) all_matches.sort() - removals: list[tuple[int, int]] = [] for ms, me, g1s, g1e, _ in all_matches: if g1s > ms: @@ -97,7 +81,6 @@ def markdown_to_signal(text: str) -> tuple[str, list[str]]: if me > g1e: removals.append((g1e, me - g1e)) removals.sort() - def _adjust(pos: int) -> int: shift = 0 for remove_pos, remove_len in removals: @@ -106,35 +89,27 @@ def markdown_to_signal(text: str) -> tuple[str, list[str]]: else: break return pos - shift - adjusted_prior: list[tuple[int, int, str]] = [] for start, length, style in styles: new_start = _adjust(start) new_end = _adjust(start + length) if new_end > new_start: adjusted_prior.append((new_start, new_end - new_start, style)) - result = "" last_end = 0 inline_styles: list[tuple[int, int, str]] = [] for ms, me, g1s, g1e, style in all_matches: result += text[last_end:ms] - pos = len(result) inner = text[g1s:g1e] + inline_styles.append((len(result), len(inner), style)) result += inner - inline_styles.append((pos, len(inner), style)) last_end = me - result += text[last_end:] - text = result - - styles = adjusted_prior + inline_styles - + text = result + text[last_end:] style_strings: list[str] = [] - for cp_start, cp_len, style_type in sorted(styles): + for cp_start, cp_len, style_type in sorted(adjusted_prior + inline_styles): if cp_start < 0 or cp_start + cp_len > len(text): continue u16_start = _utf16_len(text[:cp_start]) u16_len = _utf16_len(text[cp_start : cp_start + cp_len]) style_strings.append(f"{u16_start}:{u16_len}:{style_type}") - return text, style_strings diff --git a/gateway/platforms/signal_rate_limit.py b/gateway/platforms/signal_rate_limit.py index 7ddbce3b5b..5e81d226a3 100644 --- a/gateway/platforms/signal_rate_limit.py +++ b/gateway/platforms/signal_rate_limit.py @@ -1,16 +1,12 @@ """ Signal attachment rate-limit scheduler. -Process-wide token-bucket simulator that mirrors the per-account -attachment rate limit signal-cli/Signal-Server enforce. Producers -(``SignalAdapter.send_multiple_images`` and the ``send_message`` tool's -Signal path) call ``acquire(n)`` before an attachment send; on a 429 -they call ``feedback(retry_after, n)`` so the model recalibrates from -the server's authoritative hint. - -The scheduler serializes concurrent calls through an ``asyncio.Lock``, -giving FIFO fairness across agent sessions sharing one signal-cli -daemon. +Process-wide token-bucket simulator mirroring the per-account attachment rate +limit signal-cli/Signal-Server enforce. Producers (``SignalAdapter.send_multiple_images`` +and the ``send_message`` tool's Signal path) call ``acquire(n)`` before an +attachment send; on a 429 they call ``feedback(retry_after, n)`` so the model +recalibrates from the server's authoritative hint. Concurrent calls serialize +through an ``asyncio.Lock`` — FIFO fairness across sessions sharing one daemon. """ from __future__ import annotations @@ -25,11 +21,6 @@ from agent.retry_utils import parse_retry_after_seconds logger = logging.getLogger(__name__) - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - SIGNAL_MAX_ATTACHMENTS_PER_MSG = 32 # per-message attachment cap (source: Signal-{Android,Desktop} source code) SIGNAL_RATE_LIMIT_BUCKET_CAPACITY = 50 # server-side token-bucket capacity for attachments rate limiting SIGNAL_RATE_LIMIT_DEFAULT_RETRY_AFTER = 4 # fallback token refill interval for signal-cli < v0.14.3 @@ -38,18 +29,11 @@ SIGNAL_BATCH_PACING_NOTICE_THRESHOLD = 10.0 # if estimated waiting time > 10s, SIGNAL_RPC_ERROR_RATELIMIT = -5 # signal-cli (v0.14.3+) JSON-RPC error code for RateLimitException -# --------------------------------------------------------------------------- -# Errors -# --------------------------------------------------------------------------- - class SignalRateLimitError(Exception): - """ - Raised by ``SignalAdapter._rpc`` for rate-limit responses when the - caller has opted in via ``raise_on_rate_limit=True``. + """Raised by ``SignalAdapter._rpc`` for rate-limit responses when ``raise_on_rate_limit=True``. - Carries the server-supplied per-token Retry-After (in seconds) on - signal-cli ≥ v0.14.3 - ``retry_after`` is None when the version doesn't expose it. + ``retry_after`` is the server-supplied per-token Retry-After in seconds + (signal-cli ≥ v0.14.3), or None when the version doesn't expose it. """ def __init__(self, message: str, retry_after: Optional[float] = None) -> None: @@ -60,40 +44,27 @@ class SignalRateLimitError(Exception): class SignalSchedulerError(Exception): pass -# --------------------------------------------------------------------------- -# Detection helpers — used to fish a 429 out of signal-cli's various error -# shapes (typed code, [429] substring, libsignal-net RetryLaterException -# leaked through AttachmentInvalidException). -# --------------------------------------------------------------------------- -# "Retry after 4 seconds" / "retry after 4 second" — libsignal-net's -# RetryLaterException string form, surfaced when 429s hit during -# attachment upload (signal-cli wraps these as AttachmentInvalidException -# rather than RateLimitException, so the typed path doesn't fire). +# "Retry after 4 seconds" — libsignal-net's RetryLaterException string form, surfaced +# when 429s hit during attachment upload (signal-cli wraps these as +# AttachmentInvalidException rather than RateLimitException, so the typed path doesn't fire). _RETRY_AFTER_RE = re.compile(r"Retry after (\d+(?:\.\d+)?)\s*second", re.IGNORECASE) +def _error_message(err: Any) -> str: + return str(err.get("message", "")) if isinstance(err, dict) else str(err) + + def _extract_retry_after_seconds(err: Any) -> Optional[float]: """Pull the per-token Retry-After window from a signal-cli rate-limit error. - Tries two sources, in order: - 1. ``error.data.response.results[*].retryAfterSeconds`` — the - structured field signal-cli ≥ v0.14.3 surfaces for plain - RateLimitException. - 2. ``"Retry after N seconds"`` parsed out of the message — covers - libsignal-net's RetryLaterException that gets wrapped as - AttachmentInvalidException during attachment upload, where the - structured field stays null. - - Numeric parsing delegates to the shared - :func:`agent.retry_utils.parse_retry_after_seconds` core. - Returns None when neither source yields a value. + Sources, in order: ``error.data.response.results[*].retryAfterSeconds`` (structured + field, signal-cli ≥ v0.14.3), then ``"Retry after N seconds"`` parsed from the + message (libsignal-net RetryLaterException wrapped as AttachmentInvalidException, + where the structured field stays null). None when neither yields a value. """ - msg = "" if isinstance(err, dict): - data = err.get("data") or {} - response = data.get("response") or {} - results = response.get("results") or [] + results = ((err.get("data") or {}).get("response") or {}).get("results") or [] candidates = [ parse_retry_after_seconds(r.get("retryAfterSeconds")) for r in results if isinstance(r, dict) and r.get("retryAfterSeconds") @@ -101,33 +72,21 @@ def _extract_retry_after_seconds(err: Any) -> Optional[float]: candidates = [c for c in candidates if c is not None] if candidates: return max(candidates) - msg = str(err.get("message", "")) - else: - msg = str(err) - match = _RETRY_AFTER_RE.search(msg) + match = _RETRY_AFTER_RE.search(_error_message(err)) return parse_retry_after_seconds(match.group(1)) if match else None def _is_signal_rate_limit_error(err: Any) -> bool: """True if a signal-cli RPC error reflects a rate-limit failure. - Matches three layers: - - typed ``RATELIMIT_ERROR`` code (signal-cli ≥ v0.14.3, plain - RateLimitException) - - legacy ``[429] / RateLimitException`` substrings - - libsignal-net's ``RetryLaterException`` / ``Retry after N seconds`` - surfaced inside ``AttachmentInvalidException`` when the rate - limit is hit during attachment upload — signal-cli never re-tags - these as RateLimitException, so substring is the only signal. + Matches the typed ``RATELIMIT_ERROR`` code (signal-cli ≥ v0.14.3), the legacy + ``[429]`` / ``RateLimitException`` substrings, and libsignal-net's + ``RetryLaterException`` / ``Retry after N seconds`` leaked through + AttachmentInvalidException during upload (substring is the only signal there). """ if isinstance(err, dict) and err.get("code") == SIGNAL_RPC_ERROR_RATELIMIT: return True - - message = ( - str(err.get("message", "")) - if isinstance(err, dict) - else str(err) - ) + message = _error_message(err) msg_lower = message.lower() return ( "[429]" in message @@ -137,10 +96,6 @@ def _is_signal_rate_limit_error(err: Any) -> bool: ) -# --------------------------------------------------------------------------- -# Misc helpers -# --------------------------------------------------------------------------- - def _format_wait(seconds: float) -> str: """Human-friendly wait label for user-facing pacing notices.""" s = max(0.0, seconds) @@ -152,34 +107,22 @@ def _format_wait(seconds: float) -> str: def _signal_send_timeout(num_attachments: int) -> float: """HTTP timeout for a Signal ``send`` RPC. - signal-cli uploads attachments serially during the call, so the - server-side time scales with batch size. Default 30s is fine for - text-only sends but truncates large attachment batches mid-upload — - we then log a phantom failure even though signal-cli completes the - send a few seconds later. Scale at 5s/attachment with a 60s floor. + signal-cli uploads attachments serially inside the call, so the default 30s + truncates large batches mid-upload and we log a phantom failure even though the + send completes seconds later. Scale at 5s/attachment with a 60s floor. """ if num_attachments <= 0: return 30.0 return max(60.0, 5.0 * num_attachments) -# --------------------------------------------------------------------------- -# Scheduler -# --------------------------------------------------------------------------- - class SignalAttachmentScheduler: """Process-wide token-bucket simulator for Signal attachment sends. - The bucket holds up to ``capacity`` tokens (default 50, matching - Signal's server-side rate-limit bucket size). Each attachment consumes one - token. Tokens refill at ``refill_rate`` tokens/second, calibrated - from the per-token Retry-After hint we get from the server when a - 429 fires. Until we've observed one, we use the documented default - (1 token / 4 seconds). - - Concurrent ``acquire(n)`` calls serialize through an - ``asyncio.Lock`` — natural FIFO across agent sessions hitting the - same daemon. + Holds up to ``capacity`` tokens (default 50 = Signal's server bucket); each + attachment consumes one. Tokens refill at ``refill_rate``/s, calibrated from the + server's per-token Retry-After once a 429 has been observed (default 1 token / 4s). + ``acquire(n)`` calls serialize through an ``asyncio.Lock`` — FIFO across sessions. """ def __init__( @@ -193,9 +136,13 @@ class SignalAttachmentScheduler: self.last_refill = time.monotonic() self._lock = asyncio.Lock() - # ------------------------------------------------------------------ - # Internals - # ------------------------------------------------------------------ + def _projected_tokens(self) -> float: + """Tokens the bucket would hold now, without mutating state.""" + elapsed = time.monotonic() - self.last_refill + projected = self.tokens + if elapsed > 0 and projected < self.capacity: + projected = min(self.capacity, projected + elapsed * self.refill_rate) + return projected def _refill(self) -> None: now = time.monotonic() @@ -204,43 +151,25 @@ class SignalAttachmentScheduler: self.tokens = min(self.capacity, self.tokens + elapsed * self.refill_rate) self.last_refill = now - # ------------------------------------------------------------------ - # Public API - # ------------------------------------------------------------------ - def estimate_wait(self, n: int) -> float: - """Best-effort estimate of the seconds until ``n`` tokens would - be available. Used to decide whether to emit a user-facing - pacing notice *before* committing to an ``acquire`` that may - block silently. Lock-free; small races vs. concurrent acquires - are benign for an informational notice. + """Seconds until ``n`` tokens would be available (lock-free, informational). + + Used to decide whether to emit a user-facing pacing notice *before* an + ``acquire`` that may block silently; races vs. concurrent acquires are benign. """ - now = time.monotonic() - elapsed = now - self.last_refill - projected = self.tokens - if elapsed > 0 and projected < self.capacity: - projected = min(self.capacity, projected + elapsed * self.refill_rate) - deficit = n - projected + deficit = n - self._projected_tokens() if deficit <= 0: return 0.0 return deficit / self.refill_rate async def acquire(self, n: int) -> float: - """Block until at least ``n`` tokens are available, return the - seconds slept. + """Block until at least ``n`` tokens are available; return the seconds slept. - Does **not** deduct tokens — the bucket is a read-only model of - server-side capacity. Call ``report_rpc_duration()`` after the - RPC to synchronise the model with the server timeline. - - Not perfect in case lots of coroutines try to acquire for big - uploads (``report_rpc_duration`` will take a long time to get hit) - but this is just a simulation. Signal server is ground truth and - will raise rate-limit exceptions triggering requeues. - - The lock is released during ``asyncio.sleep`` so other callers - can interleave. A retry loop re-checks after each sleep in - case the deadline was pessimistic. + Does **not** deduct tokens — the bucket is a read-only model of server-side + capacity; call ``report_rpc_duration()`` after the RPC to sync. The lock is + released during ``asyncio.sleep`` so other callers interleave, and the loop + re-checks after each sleep in case the deadline was pessimistic. Signal's + server is ground truth and will 429 (→ requeue) if the model drifts. """ if n <= 0: return 0.0 @@ -249,7 +178,6 @@ class SignalAttachmentScheduler: f"Signal scheduler was called requesting {n} tokens " f"(max is {self.capacity})", ) - total_slept = 0.0 first_pass = True while True: @@ -277,20 +205,14 @@ class SignalAttachmentScheduler: total_slept += wait async def report_rpc_duration(self, rpc_duration: float, n_attachments: int) -> None: - """Record an attachment-send RPC that just completed. + """Deduct ``n_attachments`` tokens for a completed send RPC. - Deducts ``n_attachments`` tokens without crediting refill during - the upload window. Signal's server checks the bucket at RPC start - and does *not* refill during request processing — refill resumes - after the response. Crediting upload-time refill causes cumulative - drift that eventually triggers 429s. - - Advances ``last_refill`` so the next ``acquire`` / ``_refill`` - starts counting from this point. + No refill is credited for the upload window: Signal's server checks the bucket + at RPC start and resumes refill only after the response, so crediting it causes + cumulative drift that eventually triggers 429s. Advances ``last_refill``. """ if n_attachments <= 0: return - async with self._lock: now = time.monotonic() token_before = self.tokens @@ -306,15 +228,8 @@ class SignalAttachmentScheduler: ) def feedback(self, retry_after: Optional[float], n_attempted: int) -> None: - """Apply server feedback after a 429. - - ``retry_after`` is the per-*token* refill window the server - reports (None when signal-cli is older than v0.14.3 and didn't - surface it). - - When present we calibrate ``refill_rate`` from it: - the server is authoritative. - """ + """Apply server feedback after a 429: empty the bucket and, when ``retry_after`` + (per-token refill window) is present, calibrate ``refill_rate`` from it.""" if retry_after and retry_after > 0: new_rate = 1.0 / float(retry_after) if new_rate != self.refill_rate: @@ -328,28 +243,15 @@ class SignalAttachmentScheduler: self.last_refill = time.monotonic() def state(self) -> dict: - """Return current scheduler state for diagnostic logging (read-only). - - Does not advance ``last_refill`` — safe to call from logging paths - without perturbing the bucket. - """ - now = time.monotonic() - elapsed = now - self.last_refill - projected = self.tokens - if elapsed > 0 and projected < self.capacity: - projected = min(self.capacity, projected + elapsed * self.refill_rate) + """Current scheduler state for diagnostic logging (read-only, doesn't advance ``last_refill``).""" return { - "tokens": round(projected, 1), + "tokens": round(self._projected_tokens(), 1), "capacity": int(self.capacity), "refill_rate": round(self.refill_rate, 4), "refill_seconds_per_token": round(1.0 / self.refill_rate, 1) if self.refill_rate > 0 else float("inf"), } -# --------------------------------------------------------------------------- -# Process-wide singleton -# --------------------------------------------------------------------------- - _scheduler: Optional[SignalAttachmentScheduler] = None @@ -368,7 +270,6 @@ def get_scheduler() -> SignalAttachmentScheduler: def _reset_scheduler() -> None: - """Drop the cached scheduler so the next ``get_scheduler`` call - builds a fresh one. Test-only — never call from production paths.""" + """Drop the cached scheduler so the next ``get_scheduler`` builds a fresh one. Test-only.""" global _scheduler _scheduler = None