diff --git a/plugins/platforms/email/adapter.py b/plugins/platforms/email/adapter.py index 228cad281f..c67c5d7b6e 100644 --- a/plugins/platforms/email/adapter.py +++ b/plugins/platforms/email/adapter.py @@ -21,16 +21,13 @@ Environment variables: import asyncio import email as email_lib +from contextlib import contextmanager import imaplib import logging import os import re import smtplib import socket - -# Profile-scoped secret reader for multiplexing support (PR #50094) -from agent.secret_scope import UnscopedSecretError as _UnscopedSecretError -from agent.secret_scope import get_secret as _scoped_get_secret import ssl import uuid from email.header import decode_header @@ -52,42 +49,18 @@ from gateway.platforms.base import ( ) from gateway.config import Platform, PlatformConfig from utils import is_truthy_value +from gateway.platforms._shared import get_scoped_secret as _get_esecret, coerce_port logger = logging.getLogger(__name__) -def _get_esecret(name: str, default: str = "") -> str: - """Scope-aware ``EMAIL_*`` read with the default-profile startup fallback. - - Secondary profiles run under ``_profile_runtime_scope`` — the scope is - authoritative and a scoped miss returns ``default`` (no cross-profile - borrow). The DEFAULT profile's adapter constructs and sends *unscoped* - under multiplexing, where a bare ``get_secret`` would raise - ``UnscopedSecretError`` and crash its email 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 the WhatsApp - ``_get_wsecret`` fix (5438e9c629). - """ - try: - val = _scoped_get_secret(name, default) - except _UnscopedSecretError: - val = os.getenv(name) - return val if val is not None else default - - -# Backwards-compatible alias for the name used by the original #59076 hunks. +# Backwards-compatible alias. _get_secret = _get_esecret def _esecret_int(name: str, default: int) -> int: """Scope-aware integer read (``env_int`` variant of ``_get_esecret``).""" - raw = str(_get_esecret(name, "")).strip() - if not raw: - return default - try: - return int(raw) - except (ValueError, TypeError): - return default + return coerce_port(str(_get_esecret(name, "")).strip() or default, default) def _esecret_bool(name: str, default: bool = False) -> bool: @@ -106,8 +79,8 @@ _SECURITY_ALIASES = { def _normalize_security(value: Any, default: str = "tls") -> str: """Map an IMAP/SMTP security setting to ``tls`` | ``starttls`` | ``plain``. - Unknown values log a warning and fall back to *default* rather than - failing the connection, so a typo never silently downgrades to plaintext. + Unknown values warn and fall back to *default* so a typo never silently + downgrades to plaintext. """ raw = str(value or "").strip().lower().replace("-", "").replace("_", "") if not raw: @@ -148,19 +121,15 @@ MAX_MESSAGE_LENGTH = 50_000 SMTP_CONNECT_TIMEOUT = 30 +_TRUTHY = {"true", "1", "yes"} + def _close_imap(imap: "imaplib.IMAP4") -> None: - """Best-effort teardown that guarantees the underlying socket is closed. + """Best-effort teardown that guarantees the socket is closed. - ``IMAP4.logout()`` only guards against ``OSError`` internally: a broken - connection makes ``_simple_command('LOGOUT')`` raise ``IMAP4.abort`` - (which is *not* an ``OSError``), so ``logout()`` propagates before its - own ``shutdown()`` call and the TCP socket stays open. On macOS, where - the default soft fd limit is 256 and pollers may run through a local - proxy, these abandoned sockets accumulate one per failed poll until the - gateway hits ``[Errno 24] Too many open files`` (#79889). Always chase a - failed ``logout()`` with ``shutdown()``, which closes the socket - unconditionally. + ``IMAP4.logout()`` only guards ``OSError``; ``IMAP4.abort`` on a broken + connection escapes before its ``shutdown()``, leaking one fd per failed poll + (fatal on macOS's 256 soft limit). Chase it with an unconditional ``shutdown()``. """ try: imap.logout() @@ -171,18 +140,9 @@ def _close_imap(imap: "imaplib.IMAP4") -> None: pass -def _create_ipv4_connection( - host: str, - port: int, - timeout: float, - source_address: Any = None, -) -> socket.socket: - """Create a TCP connection using only IPv4 addresses. - - This mirrors ``socket.create_connection`` but constrains DNS resolution to - ``AF_INET``. It avoids mutating process-global socket functions, which - matters because email sends run in executor threads. - """ +def _create_ipv4_connection(host: str, port: int, timeout: float, source_address: Any = None) -> socket.socket: + """``socket.create_connection`` constrained to ``AF_INET`` (no process-global + socket mutation — email sends run in executor threads).""" last_error: OSError | None = None for family, socktype, proto, _canonname, sockaddr in socket.getaddrinfo( host, port, socket.AF_INET, socket.SOCK_STREAM @@ -204,38 +164,37 @@ def _create_ipv4_connection( class _IPv4SMTP(smtplib.SMTP): def _get_socket(self, host, port, timeout): # type: ignore[override] - return _create_ipv4_connection( - host, - port, - timeout, - source_address=self.source_address, - ) + return _create_ipv4_connection(host, port, timeout, source_address=self.source_address) class _IPv4SMTP_SSL(smtplib.SMTP_SSL): def _get_socket(self, host, port, timeout): # type: ignore[override] - raw_sock = _create_ipv4_connection( - host, - port, - timeout, - source_address=self.source_address, - ) - return self.context.wrap_socket( - raw_sock, - server_hostname=getattr(self, "_host", host), - ) + raw_sock = _create_ipv4_connection(host, port, timeout, source_address=self.source_address) + return self.context.wrap_socket(raw_sock, server_hostname=getattr(self, "_host", host)) + + +def _open_smtp(host: str, port: int, security: str, ctx: ssl.SSLContext, + smtp_cls: type, smtp_ssl_cls: type, **kwargs: Any) -> smtplib.SMTP: + """Open one SMTP connection with TLS established per *security*; *kwargs* go to the constructor.""" + if security == "tls": + return smtp_ssl_cls(host, port, context=ctx, **kwargs) + smtp = smtp_cls(host, port, **kwargs) + if security == "starttls": + try: + smtp.starttls(context=ctx) + except Exception: + smtp.close() + raise + return smtp + # Supported image extensions for inline detection _IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".gif", ".webp"} -def _send_imap_id(imap: "imaplib.IMAP4") -> None: - """Send RFC 2971 IMAP ID command identifying this client. - Required by 163/NetEase mailbox after LOGIN: without it, every UID - SEARCH/FETCH returns ``BYE Unsafe Login`` and disconnects. Other - IMAP servers either honor it silently or reject the unknown command; - we swallow failures so non-supporting servers keep working. - """ +def _send_imap_id(imap: "imaplib.IMAP4") -> None: + """Send RFC 2971 IMAP ID. 163/NetEase require it after LOGIN (else every UID + command returns ``BYE Unsafe Login``); other servers may reject it, so failures are swallowed.""" try: try: from hermes_cli import __version__ as _hermes_version @@ -261,24 +220,21 @@ def _is_automated_sender(address: str, headers: dict) -> bool: if value and check(value): return True return False - -def check_email_requirements() -> bool: - """Check if email platform settings are available and non-blank. - Treats blank/whitespace-only values as missing so an abandoned setup that - left empty ``EMAIL_*`` keys in ``.env`` does not enable the platform (#40715). - """ - addr = _get_secret("EMAIL_ADDRESS", "").strip() - pwd = _get_secret("EMAIL_PASSWORD", "").strip() - imap = _get_secret("EMAIL_IMAP_HOST", "").strip() - smtp = _get_secret("EMAIL_SMTP_HOST", "").strip() - return all([addr, pwd, imap, smtp]) + +def check_email_requirements() -> bool: + """True when all email settings are present and non-blank (blank ``EMAIL_*`` + keys left by an abandoned setup must not enable the platform).""" + return all( + _get_secret(name, "").strip() + for name in ("EMAIL_ADDRESS", "EMAIL_PASSWORD", "EMAIL_IMAP_HOST", "EMAIL_SMTP_HOST") + ) _CHARSET_ALIASES = { # Aliases seen in the wild that Python's codec registry doesn't know. # "unknown-8bit" / "x-unknown" are RFC 1428 placeholders some MTAs (QQ - # Mail among them) emit when the original charset was lost (#35901). + # Mail among them) emit when the original charset was lost. "unknown-8bit": "utf-8", "unknown": "utf-8", "x-unknown": "utf-8", @@ -292,15 +248,8 @@ _CHARSET_ALIASES = { def _safe_decode(payload: bytes, charset: "Optional[str]") -> str: - """Decode *payload* without ever raising. - - Unknown or malformed charset labels (``unknown-8bit``, misspelled names, - attacker-controlled garbage) previously raised ``LookupError`` from - ``bytes.decode`` — ``errors="replace"`` only guards decode errors, not a - missing codec — which aborted the whole IMAP fetch and dropped every - message in the batch (#35901, #55381, #55383). Fall back through a small - alias table, then UTF-8, then latin-1 (which never fails). - """ + """Decode *payload* without ever raising: ``errors="replace"`` does not guard a + missing codec (``LookupError``), so fall back via alias table → UTF-8 → latin-1.""" label = (charset or "utf-8").strip().strip("\"'").lower() or "utf-8" label = _CHARSET_ALIASES.get(label, label) for candidate in (label, "utf-8"): @@ -312,79 +261,64 @@ def _safe_decode(payload: bytes, charset: "Optional[str]") -> str: def _decode_header_value(raw: str) -> str: - """Decode an RFC 2047 encoded email header into a plain string. - - Never raises: malformed encoded-words or unknown charsets degrade to - replacement characters instead of crashing the fetch loop (#55381). - """ + """Decode an RFC 2047 header into a plain string; never raises.""" try: parts = decode_header(raw) except Exception: # malformed RFC 2047 structure return raw - decoded = [] - for part, charset in parts: - if isinstance(part, bytes): - decoded.append(_safe_decode(part, charset)) - else: - decoded.append(part) - return " ".join(decoded) + return " ".join( + _safe_decode(part, charset) if isinstance(part, bytes) else part + for part, charset in parts + ) + + +def _first_body_part(msg: email_lib.message.Message, content_type: str) -> str: + """Decoded text of the first non-attachment part of *content_type*, or ''.""" + for part in msg.walk(): + if "attachment" in str(part.get("Content-Disposition", "")): + continue + if part.get_content_type() == content_type: + payload = part.get_payload(decode=True) + if payload: + return _safe_decode(payload, part.get_content_charset()) + return "" def _extract_text_body(msg: email_lib.message.Message) -> str: """Extract the plain-text body from a potentially multipart email.""" if msg.is_multipart(): - for part in msg.walk(): - content_type = part.get_content_type() - disposition = str(part.get("Content-Disposition", "")) - # Skip attachments - if "attachment" in disposition: - continue - if content_type == "text/plain": - payload = part.get_payload(decode=True) - if payload: - return _safe_decode(payload, part.get_content_charset()) - # Fallback: try text/html and strip tags - for part in msg.walk(): - content_type = part.get_content_type() - disposition = str(part.get("Content-Disposition", "")) - if "attachment" in disposition: - continue - if content_type == "text/html": - payload = part.get_payload(decode=True) - if payload: - html = _safe_decode(payload, part.get_content_charset()) - return _strip_html(html) - return "" - else: - payload = msg.get_payload(decode=True) - if payload: - text = _safe_decode(payload, msg.get_content_charset()) - if msg.get_content_type() == "text/html": - return _strip_html(text) + text = _first_body_part(msg, "text/plain") + if text: return text + html = _first_body_part(msg, "text/html") + return _strip_html(html) if html else "" + payload = msg.get_payload(decode=True) + if not payload: return "" + text = _safe_decode(payload, msg.get_content_charset()) + return _strip_html(text) if msg.get_content_type() == "text/html" else text + + +# Ordered (pattern, replacement) substitutions for _strip_html. +_HTML_SUBS = ( + (re.compile(r"", re.IGNORECASE), "\n"), (re.compile(r"]*>", re.IGNORECASE), "\n"), + (re.compile(r"

", re.IGNORECASE), "\n"), (re.compile(r"<[^>]+>"), ""), + (re.compile(r" "), " "), (re.compile(r"&"), "&"), (re.compile(r"<"), "<"), + (re.compile(r">"), ">"), (re.compile(r"\n{3,}"), "\n\n"), +) def _strip_html(html: str) -> str: """Naive HTML tag stripper for fallback text extraction.""" - text = re.sub(r"", "\n", html, flags=re.IGNORECASE) - text = re.sub(r"]*>", "\n", text, flags=re.IGNORECASE) - text = re.sub(r"

", "\n", text, flags=re.IGNORECASE) - text = re.sub(r"<[^>]+>", "", text) - text = re.sub(r" ", " ", text) - text = re.sub(r"&", "&", text) - text = re.sub(r"<", "<", text) - text = re.sub(r">", ">", text) - text = re.sub(r"\n{3,}", "\n\n", text) - return text.strip() + for pattern, repl in _HTML_SUBS: + html = pattern.sub(repl, html) + return html.strip() def _extract_email_address(raw: str) -> str: """Extract bare email address from 'Name ' format.""" match = re.search(r"<([^>]+)>", raw) - if match: - return match.group(1).strip().lower() - return raw.strip().lower() + return (match.group(1) if match else raw).strip().lower() def _domain_of(address: str) -> str: @@ -394,74 +328,40 @@ def _domain_of(address: str) -> str: def _domains_aligned(a: str, b: str) -> bool: - """Return True if two domains are equal or in an organizational - parent/subdomain relationship (relaxed DMARC alignment). - - DMARC relaxed alignment treats ``mail.example.com`` as aligned with - ``example.com``. We approximate organizational alignment by checking - exact equality or that one domain is a dot-suffix of the other. - """ + """Relaxed DMARC alignment: equal, or one is a dot-suffix of the other.""" a = (a or "").strip().lower().rstrip(".") b = (b or "").strip().lower().rstrip(".") - if not a or not b: - return False - if a == b: - return True - return a.endswith("." + b) or b.endswith("." + a) + return bool(a and b) and (a == b or a.endswith("." + b) or b.endswith("." + a)) -# Match a single "method=result" token in an Authentication-Results header, -# e.g. ``dmarc=pass`` or ``spf=fail``. -_AUTH_METHOD_RE = re.compile( - r"\b(dmarc|dkim|spf)\s*=\s*([a-z]+)", re.IGNORECASE -) -# Match a property value like ``header.from=example.com`` or -# ``smtp.mailfrom=user@example.com``. +# "method=result" tokens (``dmarc=pass``) and property values +# (``header.from=example.com``) in an Authentication-Results header. +_AUTH_METHOD_RE = re.compile(r"\b(dmarc|dkim|spf)\s*=\s*([a-z]+)", re.IGNORECASE) _AUTH_PROP_RE = re.compile( r"\b(header\.from|header\.d|smtp\.mailfrom|smtp\.from|envelope-from)\s*=\s*([^\s;]+)", re.IGNORECASE, ) -def _verify_sender_authentication( - msg: email_lib.message.Message, - from_addr: str, - *, - authserv_id: str = "", -) -> Tuple[bool, str]: +def _verify_sender_authentication(msg: email_lib.message.Message, from_addr: str, *, + authserv_id: str = "") -> Tuple[bool, str]: """Verify that the message's ``From:`` domain is authenticated. - The ``From:`` header is attacker-controlled and is never authenticated by - IMAP delivery, so an allowlist keyed on ``From:`` alone is trivially - spoofable (GHSA-rxqh-5572-8m77). The only trustworthy signal is the - ``Authentication-Results`` header that the *receiving* mail server (the one - we IMAP into) stamps after running SPF/DKIM/DMARC. That header is prepended - by our own server, so the topmost instance is the one we trust; any - ``Authentication-Results`` an attacker injected into the body of their - message sorts below it. + ``From:`` is attacker-controlled (GHSA-rxqh-5572-8m77); the only trustworthy + signal is the ``Authentication-Results`` header stamped by the *receiving* + server. It prepends, so the FIRST instance is trusted and any injected copy + sorts below it. Pinned to *authserv_id* when given. - Returns ``(authenticated, reason)``. ``authenticated`` is True when: - * a DMARC pass is recorded for the From domain, OR - * an SPF pass aligned with the From domain, OR - * a DKIM pass aligned (``header.d``) with the From domain. - - When no ``Authentication-Results`` header is present at all, we return - ``(False, "no Authentication-Results header")`` — fail-closed. Operators - whose mail server does not stamp this header can opt out of the check - (see ``EmailAdapter._require_authenticated_sender``). + Returns ``(authenticated, reason)``; True on a DMARC pass, an aligned SPF + pass, or an aligned DKIM (``header.d``) pass. No header → fail-closed + (opt out via ``EmailAdapter._require_authenticated_sender``). """ from_domain = _domain_of(from_addr) if not from_domain: return False, "missing From domain" - - # get_all preserves header order; the receiving server prepends its result, - # so the FIRST Authentication-Results is the trusted one. We pin to the - # configured authserv-id when provided to defend against an injected header - # that happens to sort first. headers = msg.get_all("Authentication-Results") or [] if not headers: return False, "no Authentication-Results header" - trusted = None for raw in headers: value = " ".join(str(raw).split()) @@ -477,67 +377,45 @@ def _verify_sender_authentication( methods = {m.lower(): r.lower() for m, r in _AUTH_METHOD_RE.findall(trusted)} props = {p.lower(): v.strip().strip('"') for p, v in _AUTH_PROP_RE.findall(trusted)} - - # 1) DMARC pass is the strongest signal — DMARC already enforces From - # alignment, so a pass means the From domain is authenticated. + # DMARC already enforces From alignment, so a pass is sufficient. if methods.get("dmarc") == "pass": return True, "dmarc=pass" - - # 2) SPF pass aligned with the From domain (the envelope/MAIL FROM domain - # must match the From domain). + # SPF pass: envelope/MAIL FROM domain must align with From. if methods.get("spf") == "pass": - spf_domain = _domain_of(props.get("smtp.mailfrom", "")) or props.get( - "smtp.from", "" - ) or props.get("envelope-from", "") + spf_domain = (_domain_of(props.get("smtp.mailfrom", "")) or props.get("smtp.from", "") + or props.get("envelope-from", "")) spf_domain = _domain_of(spf_domain) if "@" in spf_domain else spf_domain if _domains_aligned(spf_domain, from_domain): return True, "spf=pass aligned" - - # 3) DKIM pass aligned with the From domain (the signing domain header.d - # must align with the From domain). + # DKIM pass: signing domain header.d must align with From. if methods.get("dkim") == "pass": dkim_domain = props.get("header.d", "") or _domain_of(props.get("header.from", "")) if _domains_aligned(dkim_domain, from_domain): return True, "dkim=pass aligned" - return False, f"authentication failed ({trusted[:120]})" -def _extract_attachments( - msg: email_lib.message.Message, - skip_attachments: bool = False, -) -> List[Dict[str, Any]]: +def _extract_attachments(msg: email_lib.message.Message, skip_attachments: bool = False) -> List[Dict[str, Any]]: """Extract attachment metadata and cache files locally. - When *skip_attachments* is True, all attachment/inline parts are ignored - (useful for malware protection or bandwidth savings). + When *skip_attachments* is True, all attachment/inline parts are ignored. """ attachments = [] if not msg.is_multipart(): return attachments - for part in msg.walk(): disposition = str(part.get("Content-Disposition", "")) - if skip_attachments and ("attachment" in disposition or "inline" in disposition): - continue - if "attachment" not in disposition and "inline" not in disposition: + if skip_attachments or ("attachment" not in disposition and "inline" not in disposition): continue # Skip text/plain and text/html body parts content_type = part.get_content_type() if content_type in {"text/plain", "text/html"} and "attachment" not in disposition: continue - filename = part.get_filename() - if filename: - filename = _decode_header_value(filename) - else: - ext = part.get_content_subtype() or "bin" - filename = f"attachment.{ext}" - + filename = _decode_header_value(filename) if filename else f"attachment.{part.get_content_subtype() or 'bin'}" payload = part.get_payload(decode=True) if not payload: continue - ext = Path(filename).suffix.lower() if ext in _IMAGE_EXTS: try: @@ -545,91 +423,68 @@ def _extract_attachments( except ValueError: logger.debug("Skipping non-image attachment %s (invalid magic bytes)", filename) continue - attachments.append({ - "path": cached_path, - "filename": filename, - "type": "image", - "media_type": content_type, - }) + kind = "image" else: cached_path = cache_document_from_bytes(payload, filename) - attachments.append({ - "path": cached_path, - "filename": filename, - "type": "document", - "media_type": content_type, - }) + kind = "document" + attachments.append({"path": cached_path, "filename": filename, "type": kind, "media_type": content_type}) return attachments +def _attach_file(msg: MIMEMultipart, path: Path, filename: str) -> None: + """Attach *path* to *msg* as base64 application/octet-stream.""" + with open(path, "rb") as f: + part = MIMEBase("application", "octet-stream") + part.set_payload(f.read()) + encoders.encode_base64(part) + part.add_header("Content-Disposition", f"attachment; filename={filename}") + msg.attach(part) + + class EmailAdapter(BasePlatformAdapter): """Email gateway adapter using IMAP (receive) and SMTP (send).""" - # Per-account snapshot of seen UIDs, surviving adapter recreation. - # The gateway's reconnect watcher builds a FRESH adapter instance for - # each retry; without this, connect(is_reconnect=True) would re-mark the - # entire mailbox seen and silently skip every message that arrived - # during the outage. Keyed by account address (multiplex gateways can - # run several email accounts in one process). Same-process only by - # design — after a full restart the usual mark-all-seen baseline applies. + # Per-account seen-UID snapshot surviving adapter recreation: the reconnect + # watcher builds a FRESH adapter per retry, and without this + # connect(is_reconnect=True) would re-mark the mailbox seen and skip mail + # that arrived during the outage. Keyed by address (multiplex gateways run + # several accounts). Same-process only — a full restart re-baselines. _seen_uids_snapshot: Dict[str, set] = {} def __init__(self, config: PlatformConfig): super().__init__(config, Platform.EMAIL) - # Resolve connection settings from the env vars first, then fall back to - # PlatformConfig.extra (address/imap_host/smtp_host) — the canonical dict - # gateway.config populates and that the "connected" check, the - # send-helper, and `hermes config show` already read. Without the - # fallback a config.yaml-only setup left these empty. Host/address values - # are stripped: a stray space or newline made IMAP4_SSL raise the - # misleading ``[Errno 8] nodename nor servname`` (an unresolvable name) - # instead of an obvious "host not set" error. + # Env vars first, then PlatformConfig.extra so a config.yaml-only setup + # works. Host/address are stripped: a stray newline made IMAP4_SSL raise + # ``[Errno 8] nodename nor servname`` instead of "host not set". extra = config.extra or {} - self._address = (_get_secret("EMAIL_ADDRESS", "") or extra.get("address", "")).strip() + + def setting(env: str, key: str) -> str: + return _get_secret(env, "") or extra.get(key, "") + + def tls_verify(env: str, key: str) -> bool: + return _esecret_bool(env, is_truthy_value(extra.get(key), default=True)) + + self._address = setting("EMAIL_ADDRESS", "address").strip() self._password = _get_secret("EMAIL_PASSWORD", "") - self._imap_host = (_get_secret("EMAIL_IMAP_HOST", "") or extra.get("imap_host", "")).strip() + self._imap_host = setting("EMAIL_IMAP_HOST", "imap_host").strip() self._imap_port = _esecret_int("EMAIL_IMAP_PORT", 993) - self._imap_security = _normalize_security( - _get_secret("EMAIL_IMAP_SECURITY", "") or extra.get("imap_security", "") - ) - self._imap_tls_verify = _esecret_bool( - "EMAIL_IMAP_TLS_VERIFY", - is_truthy_value(extra.get("imap_tls_verify"), default=True), - ) - self._smtp_host = (_get_secret("EMAIL_SMTP_HOST", "") or extra.get("smtp_host", "")).strip() + self._imap_security = _normalize_security(setting("EMAIL_IMAP_SECURITY", "imap_security")) + self._imap_tls_verify = tls_verify("EMAIL_IMAP_TLS_VERIFY", "imap_tls_verify") + self._smtp_host = setting("EMAIL_SMTP_HOST", "smtp_host").strip() self._smtp_port = _esecret_int("EMAIL_SMTP_PORT", 587) - self._smtp_security = _normalize_security( - _get_secret("EMAIL_SMTP_SECURITY", "") or extra.get("smtp_security", ""), - default="tls" if self._smtp_port == 465 else "starttls", - ) - self._smtp_tls_verify = _esecret_bool( - "EMAIL_SMTP_TLS_VERIFY", - is_truthy_value(extra.get("smtp_tls_verify"), default=True), - ) + self._smtp_security = _normalize_security(setting("EMAIL_SMTP_SECURITY", "smtp_security"), + default="tls" if self._smtp_port == 465 else "starttls") + self._smtp_tls_verify = tls_verify("EMAIL_SMTP_TLS_VERIFY", "smtp_tls_verify") self._poll_interval = _esecret_int("EMAIL_POLL_INTERVAL", 15) - # Skip attachments — configured via config.yaml: - # platforms: - # email: - # skip_attachments: true + # config.yaml: platforms.email.skip_attachments: true self._skip_attachments = extra.get("skip_attachments", False) - # Require the sender's From: domain to be authenticated (SPF/DKIM/DMARC) - # before trusting it for authorization. The From: header is - # attacker-controlled and unauthenticated by IMAP, so an allowlist keyed - # on it alone is spoofable (GHSA-rxqh-5572-8m77). Default ON (fail-closed). - # - # Operators whose receiving mail server does not stamp an - # Authentication-Results header can opt out via config.yaml: - # platforms: - # email: - # require_authenticated_sender: false - # or the EMAIL_TRUST_FROM_HEADER=true env mirror (parity with the other - # EMAIL_* access-control vars). When allow-all is in effect the operator - # has already chosen to accept any sender, so the check is moot and the - # gate below is skipped. + # Require an authenticated From: domain (SPF/DKIM/DMARC) before trusting + # it for authorization (GHSA-rxqh-5572-8m77). Default ON; opt out via + # platforms.email.require_authenticated_sender: false or EMAIL_TRUST_FROM_HEADER=true. if "require_authenticated_sender" in extra: self._require_authenticated_sender = bool(extra["require_authenticated_sender"]) elif _esecret_bool("EMAIL_TRUST_FROM_HEADER", False): @@ -637,36 +492,22 @@ class EmailAdapter(BasePlatformAdapter): else: self._require_authenticated_sender = True - # Optional authserv-id to pin Authentication-Results to the operator's - # own receiving server (defends against an injected header that sorts - # first). Defaults to the From-domain of the agent's own address. - self._authserv_id = ( - extra.get("authserv_id", "") or _get_secret("EMAIL_AUTHSERV_ID", "") - ).strip().lower() + # Optional authserv-id pinning Authentication-Results to the operator's + # own receiving server (defends against an injected header sorting first). + self._authserv_id = (extra.get("authserv_id", "") or _get_secret("EMAIL_AUTHSERV_ID", "")).strip().lower() - # Track message IDs we've already processed to avoid duplicates self._seen_uids: set = set() self._seen_uids_max: int = 2000 # cap to prevent unbounded memory growth self._poll_task: Optional[asyncio.Task] = None - - # Track the last IMAP fetch attempt so the poll loop can distinguish - # "checked, nothing new" from "the check itself failed" (#80016). + # Distinguish "checked, nothing new" from "the check itself failed". self._last_fetch_failed: bool = False self._last_fetch_error: str = "" - # Map chat_id (sender email) -> last subject + message-id for threading self._thread_context: Dict[str, Dict[str, str]] = {} - logger.info("[Email] Adapter initialized for %s", self._address) def _trim_seen_uids(self) -> None: - """Keep only the most recent UIDs to prevent unbounded memory growth. - - IMAP UIDs are monotonically increasing integers. When the set grows - beyond the cap, we keep only the highest half — old UIDs are safe to - drop because new messages always have higher UIDs and IMAP's UNSEEN - flag prevents re-delivery regardless. - """ + """Keep only the highest half of UIDs once over the cap (UIDs are monotonic; UNSEEN prevents re-delivery).""" if len(self._seen_uids) <= self._seen_uids_max: return try: @@ -682,12 +523,8 @@ class EmailAdapter(BasePlatformAdapter): def _connect_imap(self) -> imaplib.IMAP4: """Create an IMAP connection using implicit TLS, STARTTLS, or plaintext.""" if self._imap_security == "tls": - return imaplib.IMAP4_SSL( - self._imap_host, - self._imap_port, - timeout=30, - ssl_context=_tls_context(self._imap_tls_verify, self._imap_host), - ) + return imaplib.IMAP4_SSL(self._imap_host, self._imap_port, timeout=30, + ssl_context=_tls_context(self._imap_tls_verify, self._imap_host)) imap = imaplib.IMAP4(self._imap_host, self._imap_port, timeout=30) if self._imap_security == "starttls": @@ -698,103 +535,53 @@ class EmailAdapter(BasePlatformAdapter): raise return imap - def _connect_smtp(self) -> smtplib.SMTP: - """Create an SMTP connection, selecting the correct protocol for the port. - - Port 465 uses implicit TLS (``SMTP_SSL``). All other ports use - ``SMTP`` + ``STARTTLS``. - - When the host resolves to an IPv6 address that is unreachable - (common on networks without IPv6 routing), the default connection can - hang until the socket timeout expires. We retry connection-level - failures through an IPv4-only socket path, without mutating global - resolver state. TLS verification errors are not retried. - - Returns a connected SMTP object with TLS established — callers - can proceed directly to ``login()``. - """ - host = self._smtp_host - port = self._smtp_port - security = self._smtp_security - ctx = _tls_context(self._smtp_tls_verify, host) - - def _connect(*, ipv4_only: bool = False) -> smtplib.SMTP: - """Attempt one SMTP connection.""" - smtp_cls = _IPv4SMTP if ipv4_only else smtplib.SMTP - smtp_ssl_cls = _IPv4SMTP_SSL if ipv4_only else smtplib.SMTP_SSL - if security == "tls": - return smtp_ssl_cls(host, port, timeout=SMTP_CONNECT_TIMEOUT, context=ctx) - smtp = smtp_cls(host, port, timeout=SMTP_CONNECT_TIMEOUT) - if security == "starttls": - try: - smtp.starttls(context=ctx) - except Exception: - smtp.close() - raise - return smtp - + @contextmanager + def _inbox(self): + """Logged-in IMAP handle on INBOX; always ``_close_imap``-ed on exit (a + login/select failure used to leak one fd per reconnect attempt).""" + imap = self._connect_imap() try: - return _connect() + imap.login(self._address, self._password) + _send_imap_id(imap) + imap.select("INBOX") + yield imap + finally: + _close_imap(imap) + + def _connect_smtp(self) -> smtplib.SMTP: + """SMTP connection with TLS established (callers go straight to ``login()``). + + An unreachable IPv6 address can hang until the socket timeout, so + connection-level failures retry through an IPv4-only socket path (no + global resolver mutation). TLS verification errors are not retried. + """ + host, port, security = self._smtp_host, self._smtp_port, self._smtp_security + ctx = _tls_context(self._smtp_tls_verify, host) + try: + return _open_smtp(host, port, security, ctx, smtplib.SMTP, smtplib.SMTP_SSL, + timeout=SMTP_CONNECT_TIMEOUT) except (socket.timeout, TimeoutError, ConnectionError, OSError) as exc: if isinstance(exc, ssl.SSLError): raise - # Connection-level failure (may be unreachable IPv6). - # Retry with IPv4 only. - return _connect(ipv4_only=True) + # Connection-level failure (may be unreachable IPv6): retry IPv4 only. + return _open_smtp(host, port, security, ctx, _IPv4SMTP, _IPv4SMTP_SSL, + timeout=SMTP_CONNECT_TIMEOUT) - async def connect(self, *, is_reconnect: bool = False) -> bool: - """Connect to the IMAP server and start polling for new messages.""" - # Validate up front so a missing host surfaces as an actionable config - # error instead of IMAP4_SSL("") raising the cryptic - # ``[Errno 8] nodename nor servname provided, or not known``. - missing = [ - name - for name, value in ( - ("EMAIL_ADDRESS", self._address), - ("EMAIL_PASSWORD", self._password), - ("EMAIL_IMAP_HOST", self._imap_host), - ("EMAIL_SMTP_HOST", self._smtp_host), - ) - if not value - ] - if missing: - message = ( - "Not configured — missing " - + ", ".join(missing) - + ". Set it via `hermes gateway setup` (env) or platforms.email " - "in config.yaml." - ) - logger.error("[Email] %s", message) - # Mark non-retryable so the gateway does NOT keep reconnecting against - # an empty host. A blank-but-present env var (e.g. ``EMAIL_IMAP_HOST=``) - # used to slip past the startup gate and drive an indefinite retry - # loop that leaked memory until the host OOM-killed (#40715). - self._set_fatal_error( - "email_missing_configuration", message, retryable=False - ) - return False + def _fail(self, log_fmt: str, err: object, code: str, detail: str, *, retryable: bool) -> bool: + """Log *err*, record a fatal error for the gateway's reconnect machinery, return False.""" + logger.error(log_fmt, err) + self._set_fatal_error(code, detail, retryable=retryable) + return False + def _probe_imap(self, is_reconnect: bool) -> bool: + """Connection test + seen-UID baseline. Sets a fatal error and returns False on failure.""" try: - # Test IMAP connection. The handle is closed in ``finally`` — - # before this, a failure in login/select/search left the TCP - # socket open with no owner, leaking one fd per connect attempt. - # Under the gateway's reconnect watcher (fresh adapter instance - # per retry) against an unreachable/proxied host this grew - # monotonically until fd exhaustion on macOS's 256 soft limit - # (#79889). - imap = None - try: - imap = self._connect_imap() - imap.login(self._address, self._password) - _send_imap_id(imap) - imap.select("INBOX") + with self._inbox() as imap: snapshot = self._seen_uids_snapshot.get(self._address) if is_reconnect and snapshot is not None: - # Reconnect within the same process: restore the previous - # adapter's seen-UID baseline instead of re-marking the whole - # mailbox. Mail that arrived during the outage stays UNSEEN - # relative to the baseline and is dispatched by the next poll - # instead of being silently skipped. + # Same-process reconnect: restore the previous adapter's + # baseline so mail that arrived during the outage stays + # eligible for the next poll instead of being skipped. self._seen_uids = set(snapshot) self._trim_seen_uids() logger.info( @@ -803,69 +590,62 @@ class EmailAdapter(BasePlatformAdapter): len(self._seen_uids), ) else: - # First connect (or no snapshot): mark all existing messages as - # seen so we only process new ones. + # First connect (or no snapshot): mark all existing messages seen. status, data = imap.uid("search", None, "ALL") if status == "OK" and data and data[0]: - for uid in data[0].split(): - self._seen_uids.add(uid) - # Keep only the most recent UIDs to prevent unbounded growth + self._seen_uids.update(data[0].split()) self._trim_seen_uids() logger.info("[Email] IMAP connection test passed. %d existing messages skipped.", len(self._seen_uids)) - finally: - if imap is not None: - _close_imap(imap) self._seen_uids_snapshot[self._address] = set(self._seen_uids) + return True except Exception as e: - logger.error("[Email] IMAP connection failed: %s", e) - # Always set an explicit fatal code (OOF-156): returning False - # with no error info made the gateway treat every IMAP failure — - # including permanently bad credentials — as transient, retrying - # forever with zero owner signal ("stuck retrying 22h"). - # Kept retryable=True deliberately: imaplib raises the same - # generic IMAP4.error for bad credentials AND transient server - # NOs (e.g. Gmail's "too many simultaneous connections"), so a - # type-based terminal classification isn't safe here. Long-lived - # loops surface via the reconnect watcher's NEEDS_ATTENTION - # escalation instead. - self._set_fatal_error( - "email_imap_connect_error", - f"IMAP connection to {self._imap_host}:{self._imap_port} failed: {e}", - retryable=True, - ) - return False + # Always set an explicit fatal code, else the gateway treats every + # failure as transient with zero owner signal. retryable=True because + # imaplib raises the same generic IMAP4.error for bad credentials AND + # transient NOs (Gmail "too many simultaneous connections"); long-lived + # loops surface via the reconnect watcher's NEEDS_ATTENTION escalation. + return self._fail("[Email] IMAP connection failed: %s", e, "email_imap_connect_error", + f"IMAP connection to {self._imap_host}:{self._imap_port} failed: {e}", retryable=True) + def _probe_smtp(self) -> bool: + """SMTP connect + login test. Sets a fatal error and returns False on failure.""" try: - # Test SMTP connection smtp = self._connect_smtp() try: smtp.login(self._address, self._password) finally: smtp.quit() logger.info("[Email] SMTP connection test passed.") + return True except smtplib.SMTPAuthenticationError as e: - logger.error("[Email] SMTP authentication failed: %s", e) - # Typed auth failure (535 & friends): bad or revoked credentials - # can never self-heal, so drop out of the reconnect queue instead - # of retrying a dead password forever (OOF-156). Type-based only — - # SMTPAuthenticationError is unambiguous, unlike IMAP4.error above. - self._set_fatal_error( - "email_auth_error", + # Typed auth failure (535 & friends) can never self-heal, so drop out + # of the reconnect queue — unambiguous, unlike IMAP4.error above. + return self._fail( + "[Email] SMTP authentication failed: %s", e, "email_auth_error", f"SMTP authentication failed for {self._address}: {e}. " "Check EMAIL_PASSWORD (for Gmail/Outlook this must be an " "app password, not the account password).", retryable=False, ) - return False except Exception as e: - logger.error("[Email] SMTP connection failed: %s", e) - self._set_fatal_error( - "email_smtp_connect_error", - f"SMTP connection to {self._smtp_host} failed: {e}", - retryable=True, - ) - return False + return self._fail("[Email] SMTP connection failed: %s", e, "email_smtp_connect_error", + f"SMTP connection to {self._smtp_host} failed: {e}", retryable=True) + async def connect(self, *, is_reconnect: bool = False) -> bool: + """Connect to the IMAP server and start polling for new messages.""" + # Validate up front so a missing host is an actionable config error, not + # IMAP4_SSL("") raising ``[Errno 8] nodename nor servname provided``. + required = (("EMAIL_ADDRESS", self._address), ("EMAIL_PASSWORD", self._password), + ("EMAIL_IMAP_HOST", self._imap_host), ("EMAIL_SMTP_HOST", self._smtp_host)) + missing = [name for name, value in required if not value] + if missing: + message = ("Not configured — missing " + ", ".join(missing) + + ". Set it via `hermes gateway setup` (env) or platforms.email in config.yaml.") + # Non-retryable: a blank-but-present env var (``EMAIL_IMAP_HOST=``) + # used to drive an indefinite retry loop that leaked until OOM. + return self._fail("[Email] %s", message, "email_missing_configuration", message, retryable=False) + if not self._probe_imap(is_reconnect) or not self._probe_smtp(): + return False self._running = True self._poll_task = asyncio.create_task(self._poll_loop()) print(f"[Email] Connected as {self._address}") @@ -898,119 +678,74 @@ class EmailAdapter(BasePlatformAdapter): async def _check_inbox(self) -> None: """Check INBOX for unseen messages and dispatch them.""" - # Run IMAP operations in a thread to avoid blocking the event loop loop = asyncio.get_running_loop() messages = await loop.run_in_executor(None, self._fetch_new_messages) - # Dispatch whatever the fetch managed to return BEFORE escalating a - # failure: on a mid-batch exception _fetch_new_messages returns the - # partial results, and dropping them here would lose those messages - # (their processing already marked them seen). + # Dispatch partial results BEFORE escalating a failure — a mid-batch + # exception returns what was fetched, and those are already marked seen. for msg_data in messages: await self._dispatch_message(msg_data) if self._last_fetch_failed: - # The IMAP check itself failed (connect/login/select/search/fetch), - # not just an empty inbox. Surface it through the fatal-error hook - # so the gateway's existing reconnect/backoff/status machinery - # re-establishes the mailbox instead of silently treating every - # failed check as "nothing new" (#80016). The handler runs in a - # detached task (gateway/run.py), so awaiting it from our own poll - # task is safe even though teardown cancels this task. + # The IMAP check itself failed (not an empty inbox): route through the + # fatal-error hook so the gateway's reconnect/backoff re-establishes the + # mailbox. The handler runs in a detached task (gateway/run.py), so + # awaiting it from our own poll task is safe despite teardown cancelling us. self._last_fetch_failed = False - self._set_fatal_error( - "email_imap_fetch_failed", - self._last_fetch_error or "IMAP fetch failed", - retryable=True, - ) + self._set_fatal_error("email_imap_fetch_failed", self._last_fetch_error or "IMAP fetch failed", retryable=True) await self._notify_fatal_error() def _fetch_new_messages(self) -> List[Dict[str, Any]]: """Fetch new (unseen) messages from IMAP. Runs in executor thread.""" results = [] - imap: Optional[imaplib.IMAP4] = None try: - imap = self._connect_imap() - try: - imap.login(self._address, self._password) - _send_imap_id(imap) - imap.select("INBOX") - + with self._inbox() as imap: status, data = imap.uid("search", None, "UNSEEN") - if status != "OK" or not data or not data[0]: - return results - - for uid in data[0].split(): + uids = data[0].split() if status == "OK" and data and data[0] else [] + for uid in uids: if uid in self._seen_uids: continue - status, msg_data = imap.uid("fetch", uid, "(RFC822)") if status != "OK": - # Transient per-UID fetch refusal: leave the UID out of - # _seen_uids so the next poll retries it. + # Transient per-UID refusal: leave UID unseen so the next poll retries. continue - - # IMAP fetch can return unexpected structures (e.g. a - # single bytes item instead of a list of tuples). Mark the - # UID seen once a response arrived (even a malformed one) - # so a garbage response is skipped once, not retried - # forever — but NOT before the fetch: a connection failure - # above must leave the remaining batch eligible for the - # next poll instead of permanently skipping it (#80032 - # review). + # Mark seen once a response arrived (even malformed) so garbage + # is skipped once, not retried forever — but NOT before the + # fetch: a connection failure must leave the rest of the batch + # eligible for the next poll. self._seen_uids.add(uid) - # Trim periodically to prevent unbounded memory growth if len(self._seen_uids) > self._seen_uids_max: self._trim_seen_uids() - try: raw_email = msg_data[0][1] except (IndexError, TypeError): - logger.warning( - "[Email] Unexpected IMAP response structure for UID %s, skipping", - uid, - ) + logger.warning("[Email] Unexpected IMAP response structure for UID %s, skipping", uid) continue if not isinstance(raw_email, (bytes, bytearray)): - logger.warning( - "[Email] Non-bytes IMAP payload for UID %s, skipping", uid - ) + logger.warning("[Email] Non-bytes IMAP payload for UID %s, skipping", uid) continue - # Per-message processing guard: one poison message - # (unparseable headers, pathological attachment, DNS - # hiccup in SPF/DKIM verification) must not abort the - # batch or escalate to a reconnect — it is already marked - # seen above, so log the UID and move on (#80032 review). + # One poison message (unparseable headers, pathological + # attachment, DNS hiccup in SPF/DKIM) must not abort the batch + # or escalate to a reconnect — it is already marked seen. try: parsed = self._parse_fetched_message(uid, raw_email) except Exception as parse_exc: - logger.error( - "[Email] Failed to process message UID %s, skipping: %s", - uid, - parse_exc, - ) + logger.error("[Email] Failed to process message UID %s, skipping: %s", uid, parse_exc) continue if parsed is not None: results.append(parsed) - finally: - # _close_imap guarantees the socket dies even when logout() - # raises IMAP4.abort on a broken connection (#79889). - _close_imap(imap) except Exception as e: logger.error("[Email] IMAP fetch error: %s", e) self._last_fetch_failed = True self._last_fetch_error = str(e) - # Keep the reconnect snapshot current with every poll so a mid-outage - # adapter recreation restores an up-to-date baseline: stale snapshots - # would re-dispatch messages this instance already processed. + # Keep the reconnect snapshot current so a mid-outage adapter recreation + # does not re-dispatch messages this instance already processed. self._seen_uids_snapshot[self._address] = set(self._seen_uids) return results def _parse_fetched_message(self, uid: bytes, raw_email: "bytes | bytearray") -> Optional[Dict[str, Any]]: """Parse one fetched RFC822 payload into a dispatchable dict. - Returns ``None`` for messages that should be silently skipped - (automated/noreply senders). Raises on pathological input — the - caller's per-message guard logs the UID and continues, so a poison - message never aborts the batch or escalates to a reconnect. + Returns ``None`` for automated/noreply senders. Raises on pathological + input — the caller's per-message guard logs the UID and continues. """ msg = email_lib.message_from_bytes(raw_email) @@ -1022,36 +757,24 @@ class EmailAdapter(BasePlatformAdapter): sender_name = sender_name.split("<")[0].strip().strip('"') subject = _decode_header_value(msg.get("Subject", "(no subject)")) - message_id = msg.get("Message-ID", "") - in_reply_to = msg.get("In-Reply-To", "") # Skip automated/noreply senders before any processing - msg_headers = dict(msg.items()) - if _is_automated_sender(sender_addr, msg_headers): + if _is_automated_sender(sender_addr, dict(msg.items())): logger.debug("[Email] Skipping automated sender: %s", sender_addr) return None - # Verify the From: domain is authenticated (SPF/DKIM/DMARC) - # while the raw message — and its trusted - # Authentication-Results header — is still in scope. The - # verdict is consumed at dispatch where authorization is - # decided. From: is attacker-controlled, so this is the only - # place a spoof can be caught (GHSA-rxqh-5572-8m77). - sender_authenticated, auth_reason = _verify_sender_authentication( - msg, sender_addr, authserv_id=self._authserv_id - ) - - body = _extract_text_body(msg) - attachments = _extract_attachments(msg, skip_attachments=self._skip_attachments) - + # Verify From: authentication while the raw message (and its trusted + # Authentication-Results header) is in scope; the verdict is consumed + # at dispatch where authorization is decided (GHSA-rxqh-5572-8m77). + sender_authenticated, auth_reason = _verify_sender_authentication(msg, sender_addr, authserv_id=self._authserv_id) return { "uid": uid, "sender_addr": sender_addr, "sender_name": sender_name, "subject": subject, - "message_id": message_id, - "in_reply_to": in_reply_to, - "body": body, - "attachments": attachments, + "message_id": msg.get("Message-ID", ""), + "in_reply_to": msg.get("In-Reply-To", ""), + "body": _extract_text_body(msg), + "attachments": _extract_attachments(msg, skip_attachments=self._skip_attachments), "date": msg.get("Date", ""), "sender_authenticated": sender_authenticated, "auth_reason": auth_reason, @@ -1059,196 +782,124 @@ class EmailAdapter(BasePlatformAdapter): @staticmethod def _allow_all_senders() -> bool: - """Return True when the operator opted into accepting any sender. - - Mirrors the gateway authz allow-all resolution: the per-platform - EMAIL_ALLOW_ALL_USERS flag or the global GATEWAY_ALLOW_ALL_USERS flag. - When either is set, sender identity is moot, so the From: authentication - gate is skipped. - """ - truthy = {"true", "1", "yes"} - return ( - _get_secret("EMAIL_ALLOW_ALL_USERS", "").strip().lower() in truthy - or os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in truthy - ) + """True when the operator opted into any sender (EMAIL_ or GATEWAY_ALLOW_ALL_USERS).""" + return (_get_secret("EMAIL_ALLOW_ALL_USERS", "").strip().lower() in _TRUTHY + or os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in _TRUTHY) @staticmethod def _allowlist_in_effect() -> bool: - """Return True when a sender allowlist gates email access. + """True when EMAIL_ALLOWED_USERS or GATEWAY_ALLOWED_USERS gates access. - Authorization keys on the From: address only when an allowlist is - configured — the per-platform EMAIL_ALLOWED_USERS or the global - GATEWAY_ALLOWED_USERS. When neither is set the gateway default-denies - every sender regardless, so the spoofable From: identity grants nothing - and the authentication gate is unnecessary. + Without one the gateway default-denies every sender, so the spoofable + From: identity grants nothing and the authentication gate is moot. """ - return bool( - _get_secret("EMAIL_ALLOWED_USERS", "").strip() - or os.getenv("GATEWAY_ALLOWED_USERS", "").strip() - ) + return bool(_get_secret("EMAIL_ALLOWED_USERS", "").strip() or os.getenv("GATEWAY_ALLOWED_USERS", "").strip()) - async def _dispatch_message(self, msg_data: Dict[str, Any]) -> None: - """Convert a fetched email into a MessageEvent and dispatch it.""" - sender_addr = msg_data["sender_addr"] - - # Skip self-messages + def _sender_accepted(self, sender_addr: str, msg_data: Dict[str, Any]) -> bool: + """Pre-dispatch sender gate: self, automated, allowlist, From: authentication.""" if sender_addr == self._address.lower(): - return - - # Never reply to automated senders + return False if _is_automated_sender(sender_addr, {}): logger.debug("[Email] Dropping automated sender at dispatch: %s", sender_addr) - return + return False - # Skip senders not in EMAIL_ALLOWED_USERS — prevents the adapter - # from creating a MessageEvent (and thus thread context) for senders - # that the gateway will never authorize. Without this early guard, - # a race between dispatch and authorization can result in the adapter - # sending a reply even though the handler returned None. + # Drop senders the gateway would never authorize before a MessageEvent + # (and thread context) exists — otherwise a race between dispatch and + # authorization can send a reply even though the handler returned None. allowed_raw = _get_secret("EMAIL_ALLOWED_USERS", "").strip() if not allowed_raw: - if _get_secret("EMAIL_ALLOW_ALL_USERS", "").strip().lower() not in {"true", "1", "yes"} and ( - os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() not in {"true", "1", "yes"} - ): - logger.debug( - "[Email] Dropping sender at dispatch — EMAIL_ALLOWED_USERS is unset " - "and open access is not opted in: %s", - sender_addr, - ) - return - else: - allowed = {addr.strip().lower() for addr in allowed_raw.split(",") if addr.strip()} - if sender_addr.lower() not in allowed: - logger.debug("[Email] Dropping non-allowlisted sender at dispatch: %s", sender_addr) - return + if not self._allow_all_senders(): + logger.debug("[Email] Dropping sender at dispatch — EMAIL_ALLOWED_USERS is unset " + "and open access is not opted in: %s", sender_addr) + return False + elif sender_addr.lower() not in {a.strip().lower() for a in allowed_raw.split(",") if a.strip()}: + logger.debug("[Email] Dropping non-allowlisted sender at dispatch: %s", sender_addr) + return False - # Reject spoofed senders. The allowlist (and the gateway's own authz) - # key on sender_addr, which comes straight from the attacker-controlled - # From: header — so an attacker can forge From: an-allowlisted@addr to - # get authorized (GHSA-rxqh-5572-8m77). This only matters when an - # allowlist is actually being used to GRANT access: if no allowlist is - # configured the gateway default-denies everyone anyway, and if allow-all - # is on the operator already accepts any sender. So enforce From: - # authentication exactly when an allowlist is in effect and allow-all is - # off. Fail-closed: an unauthenticated From: is dropped before it can be - # matched against the allowlist. - if ( - self._require_authenticated_sender - and self._allowlist_in_effect() - and not self._allow_all_senders() - and not msg_data.get("sender_authenticated", False) - ): + # Reject spoofed senders (GHSA-rxqh-5572-8m77): the allowlist keys on the + # attacker-controlled From:. Only matters when an allowlist GRANTS access + # and allow-all is off; fail-closed before matching against the allowlist. + if (self._require_authenticated_sender and self._allowlist_in_effect() + and not self._allow_all_senders() and not msg_data.get("sender_authenticated", False)): logger.warning( "[Email] Dropping sender with unauthenticated From: %s (%s). " "If your mail server does not stamp Authentication-Results, set " "platforms.email.require_authenticated_sender: false (or " "EMAIL_TRUST_FROM_HEADER=true) to accept the risk.", - sender_addr, - msg_data.get("auth_reason", "no verdict"), + sender_addr, msg_data.get("auth_reason", "no verdict"), ) + return False + return True + + async def _dispatch_message(self, msg_data: Dict[str, Any]) -> None: + """Convert a fetched email into a MessageEvent and dispatch it.""" + sender_addr = msg_data["sender_addr"] + if not self._sender_accepted(sender_addr, msg_data): return - subject = msg_data["subject"] - body = msg_data["body"].strip() - attachments = msg_data["attachments"] + subject, body, attachments = msg_data["subject"], msg_data["body"].strip(), msg_data["attachments"] + # Include subject as context unless it is a reply + text = f"[Subject: {subject}]\n\n{body}" if subject and not subject.startswith("Re:") else body - # Build message text: include subject as context - text = body - if subject and not subject.startswith("Re:"): - text = f"[Subject: {subject}]\n\n{body}" - - # Determine message type and media - media_urls = [] - media_types = [] msg_type = MessageType.TEXT - for att in attachments: - media_urls.append(att["path"]) - media_types.append(att["media_type"]) if att["type"] == "image" and msg_type == MessageType.TEXT: msg_type = MessageType.PHOTO elif att["type"] == "document": - # Document wins over PHOTO for mixed attachments: run.py's - # image handling keys off the per-path image/* mime type - # regardless of message_type, but document-context injection - # gates strictly on MessageType.DOCUMENT — so DOCUMENT is the - # only classification that surfaces both. + # Document wins over PHOTO for mixed attachments: run.py keys + # image handling off the per-path mime type regardless of + # message_type, but document-context injection gates strictly + # on MessageType.DOCUMENT — so DOCUMENT surfaces both. msg_type = MessageType.DOCUMENT # Store thread context for reply threading - self._thread_context[sender_addr] = { - "subject": subject, - "message_id": msg_data["message_id"], - } - - source = self.build_source( - chat_id=sender_addr, - chat_name=msg_data["sender_name"] or sender_addr, - chat_type="dm", - user_id=sender_addr, - user_name=msg_data["sender_name"] or sender_addr, - ) - + self._thread_context[sender_addr] = {"subject": subject, "message_id": msg_data["message_id"]} + name = msg_data["sender_name"] or sender_addr event = MessageEvent( text=text or "(empty email)", message_type=msg_type, - source=source, + source=self.build_source(chat_id=sender_addr, chat_name=name, chat_type="dm", + user_id=sender_addr, user_name=name), message_id=msg_data["message_id"], - media_urls=media_urls, - media_types=media_types, + media_urls=[att["path"] for att in attachments], + media_types=[att["media_type"] for att in attachments], reply_to_message_id=msg_data["in_reply_to"] or None, ) - logger.info("[Email] New message from %s: %s", sender_addr, subject) await self.handle_message(event) - async def send( - self, - chat_id: str, - content: str, - reply_to: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - ) -> SendResult: - """Send an email reply to the given address.""" + async def _run_send(self, fn, args: tuple, log_fmt: str, *log_args) -> SendResult: + """Run a blocking SMTP sender in the executor; wrap its Message-ID in a SendResult.""" try: - loop = asyncio.get_running_loop() - message_id = await loop.run_in_executor( - None, self._send_email, chat_id, content, reply_to - ) + message_id = await asyncio.get_running_loop().run_in_executor(None, fn, *args) return SendResult(success=True, message_id=message_id) except Exception as e: - logger.error("[Email] Send failed to %s: %s", chat_id, e) + logger.error(log_fmt, *log_args, e) return SendResult(success=False, error=str(e)) + async def send(self, chat_id: str, content: str, reply_to: Optional[str] = None, + metadata: Optional[Dict[str, Any]] = None) -> SendResult: + """Send an email reply to the given address.""" + return await self._run_send(self._send_email, (chat_id, content, reply_to), + "[Email] Send failed to %s: %s", chat_id) + def _message_id_domain(self) -> str: - """Domain part for generated Message-IDs. + """Domain for generated Message-IDs; ``localhost`` when EMAIL_ADDRESS lacks ``@``.""" + return (self._address.rsplit("@", 1)[-1] if "@" in self._address else "") or "localhost" - EMAIL_ADDRESS may lack an ``@`` (misconfiguration); fall back to - ``localhost`` instead of crashing send with an IndexError. - """ - if "@" in self._address: - return self._address.rsplit("@", 1)[-1] or "localhost" - return "localhost" - - def _send_email( - self, - to_addr: str, - body: str, - reply_to_msg_id: Optional[str] = None, - ) -> str: - """Send an email via SMTP. Runs in executor thread.""" + def _new_reply(self, to_addr: str, body: str, reply_to_msg_id: Optional[str] = None, + *, attach_empty_body: bool = False) -> Tuple[MIMEMultipart, str, str]: + """Build a threaded reply skeleton. Returns ``(msg, msg_id, subject)``.""" msg = MIMEMultipart() msg["From"] = self._address msg["To"] = to_addr - # Thread context for reply ctx = self._thread_context.get(to_addr, {}) subject = ctx.get("subject", "Hermes Agent") if not subject.startswith("Re:"): subject = f"Re: {subject}" msg["Subject"] = subject - # Threading headers original_msg_id = reply_to_msg_id or ctx.get("message_id") if original_msg_id: msg["In-Reply-To"] = original_msg_id @@ -1258,8 +909,12 @@ class EmailAdapter(BasePlatformAdapter): msg_id = f"" msg["Message-ID"] = msg_id - msg.attach(MIMEText(body, "plain", "utf-8")) + if body or attach_empty_body: + msg.attach(MIMEText(body, "plain", "utf-8")) + return msg, msg_id, subject + def _smtp_send(self, msg: MIMEMultipart) -> None: + """Login, send, and always release the SMTP connection (quit, else close).""" smtp = self._connect_smtp() try: smtp.login(self._address, self._password) @@ -1270,42 +925,26 @@ class EmailAdapter(BasePlatformAdapter): except Exception: smtp.close() + def _send_email(self, to_addr: str, body: str, reply_to_msg_id: Optional[str] = None) -> str: + """Send an email via SMTP. Runs in executor thread.""" + msg, msg_id, subject = self._new_reply(to_addr, body, reply_to_msg_id, attach_empty_body=True) + self._smtp_send(msg) logger.info("[Email] Sent reply to %s (subject: %s)", to_addr, subject) return msg_id - async def send_typing(self, chat_id: str, metadata: Optional[Dict[str, Any]] = None) -> None: - """Email has no typing indicator — no-op.""" - - async def send_image( - self, - chat_id: str, - image_url: str, - caption: Optional[str] = None, - reply_to: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, - ) -> SendResult: - """Send an image URL as part of an email body. - - ``metadata`` is accepted to honor the base-class contract; the - email body send doesn't use it. - """ + async def send_image(self, chat_id: str, image_url: str, caption: Optional[str] = None, + reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None) -> SendResult: + """Send an image URL as part of an email body (``metadata`` unused).""" text = caption or "" text += f"\n\nImage: {image_url}" return await self.send(chat_id, text.strip(), reply_to) - async def send_multiple_images( - self, - chat_id: str, - images: List[Tuple[str, str]], - metadata: Optional[Dict[str, Any]] = None, - human_delay: float = 0.0, - ) -> None: - """Send a batch of images as a single email with multiple MIME attachments. + async def send_multiple_images(self, chat_id: str, images: List[Tuple[str, str]], + metadata: Optional[Dict[str, Any]] = None, human_delay: float = 0.0) -> None: + """Send a batch of images as one email with multiple MIME attachments. - Local files are attached directly. URL images have their URL - appended to the body (email adapter does not download remote - images). No hard cap — email clients handle dozens of - attachments fine, subject to SMTP message size limits. + Local files are attached; URL images are linked in the body (no remote + download). No hard cap beyond SMTP message size limits. """ if not images: return @@ -1326,209 +965,68 @@ class EmailAdapter(BasePlatformAdapter): else: # Remote URLs just get linked in the body (parity with send_image) body_parts.append(f"Image: {image_url}") - if not local_paths and not body_parts: return - - body = "\n\n".join(body_parts) - try: - loop = asyncio.get_running_loop() - await loop.run_in_executor( - None, - self._send_email_with_attachments, - chat_id, - body, - local_paths, - ) + await asyncio.get_running_loop().run_in_executor( + None, self._send_email_with_attachments, chat_id, "\n\n".join(body_parts), local_paths) except Exception as e: logger.error("[Email] Multi-image send failed, falling back: %s", e, exc_info=True) await super().send_multiple_images(chat_id, images, metadata, human_delay) - def _send_email_with_attachments( - self, - to_addr: str, - body: str, - file_paths: List[str], - ) -> str: - """Send an email with multiple file attachments via SMTP.""" - msg = MIMEMultipart() - msg["From"] = self._address - msg["To"] = to_addr - - ctx = self._thread_context.get(to_addr, {}) - subject = ctx.get("subject", "Hermes Agent") - if not subject.startswith("Re:"): - subject = f"Re: {subject}" - msg["Subject"] = subject - - original_msg_id = ctx.get("message_id") - if original_msg_id: - msg["In-Reply-To"] = original_msg_id - msg["References"] = original_msg_id - - msg["Date"] = formatdate(localtime=True) - msg_id = f"" - msg["Message-ID"] = msg_id - - if body: - msg.attach(MIMEText(body, "plain", "utf-8")) - + def _send_email_with_attachments(self, to_addr: str, body: str, file_paths: List[str]) -> str: + """Send an email with multiple file attachments via SMTP (unattachable files are skipped).""" + msg, msg_id, _ = self._new_reply(to_addr, body) for file_path in file_paths: p = Path(file_path) try: - with open(p, "rb") as f: - part = MIMEBase("application", "octet-stream") - part.set_payload(f.read()) - encoders.encode_base64(part) - part.add_header("Content-Disposition", f"attachment; filename={p.name}") - msg.attach(part) + _attach_file(msg, p, p.name) except Exception as e: logger.warning("[Email] Failed to attach %s: %s", file_path, e) - - smtp = self._connect_smtp() - try: - smtp.login(self._address, self._password) - smtp.send_message(msg) - finally: - try: - smtp.quit() - except Exception: - smtp.close() - + self._smtp_send(msg) logger.info("[Email] Sent multi-attachment email to %s (%d files)", to_addr, len(file_paths)) return msg_id - async def send_document( - self, - chat_id: str, - file_path: str, - caption: Optional[str] = None, - file_name: Optional[str] = None, - reply_to: Optional[str] = None, - **kwargs, - ) -> SendResult: + async def send_document(self, chat_id: str, file_path: str, caption: Optional[str] = None, + file_name: Optional[str] = None, reply_to: Optional[str] = None, **kwargs) -> SendResult: """Send a file as an email attachment.""" - try: - loop = asyncio.get_running_loop() - message_id = await loop.run_in_executor( - None, - self._send_email_with_attachment, - chat_id, - caption or "", - file_path, - file_name, - ) - return SendResult(success=True, message_id=message_id) - except Exception as e: - logger.error("[Email] Send document failed: %s", e) - return SendResult(success=False, error=str(e)) + return await self._run_send(self._send_email_with_attachment, (chat_id, caption or "", file_path, file_name), + "[Email] Send document failed: %s") - def _send_email_with_attachment( - self, - to_addr: str, - body: str, - file_path: str, - file_name: Optional[str] = None, - ) -> str: + def _send_email_with_attachment(self, to_addr: str, body: str, file_path: str, + file_name: Optional[str] = None) -> str: """Send an email with a file attachment via SMTP.""" - msg = MIMEMultipart() - msg["From"] = self._address - msg["To"] = to_addr - - ctx = self._thread_context.get(to_addr, {}) - subject = ctx.get("subject", "Hermes Agent") - if not subject.startswith("Re:"): - subject = f"Re: {subject}" - msg["Subject"] = subject - - original_msg_id = ctx.get("message_id") - if original_msg_id: - msg["In-Reply-To"] = original_msg_id - msg["References"] = original_msg_id - - msg["Date"] = formatdate(localtime=True) - msg_id = f"" - msg["Message-ID"] = msg_id - - if body: - msg.attach(MIMEText(body, "plain", "utf-8")) - - # Attach file + msg, msg_id, _ = self._new_reply(to_addr, body) p = Path(file_path) - fname = file_name or p.name - with open(p, "rb") as f: - part = MIMEBase("application", "octet-stream") - part.set_payload(f.read()) - encoders.encode_base64(part) - part.add_header("Content-Disposition", f"attachment; filename={fname}") - msg.attach(part) - - smtp = self._connect_smtp() - try: - smtp.login(self._address, self._password) - smtp.send_message(msg) - finally: - try: - smtp.quit() - except Exception: - smtp.close() - + _attach_file(msg, p, file_name or p.name) + self._smtp_send(msg) return msg_id async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: """Return basic info about the email chat.""" ctx = self._thread_context.get(chat_id, {}) - return { - "name": chat_id, - "type": "dm", - "chat_id": chat_id, - "subject": ctx.get("subject", ""), - } + return {"name": chat_id, "type": "dm", "chat_id": chat_id, "subject": ctx.get("subject", "")} -# ────────────────────────────────────────────────────────────────────────── -# Plugin migration glue (#41112 / #3823) -# -# Added when the Email adapter moved from gateway/platforms/email.py into this -# bundled plugin. register() exposes the platform via the registry, replacing -# the Platform.EMAIL elif in gateway/run.py, the _PLATFORM_CONNECTED_CHECKERS -# entry in gateway/config.py, the _PLATFORMS["email"] static dict in -# hermes_cli/gateway.py, and the _send_email dispatch in -# tools/send_message_tool.py. EMAIL_* env→PlatformConfig seeding stays in core. -# ────────────────────────────────────────────────────────────────────────── +# Plugin glue: register() exposes the platform via the registry (replacing the +# former Platform.EMAIL branches in gateway/run.py, gateway/config.py, +# hermes_cli/gateway.py and tools/send_message_tool.py). EMAIL_* env → +# PlatformConfig seeding stays in core. -async def _standalone_send( - pconfig, - chat_id, - message, - *, - thread_id=None, - media_files=None, - force_document=False, -): - """Out-of-process Email delivery via SMTP (one-shot). Implements the - standalone_sender_fn contract; replaces the legacy _send_email helper.""" - import smtplib - from email.mime.text import MIMEText - from email.utils import formatdate - +async def _standalone_send(pconfig, chat_id, message, *, thread_id=None, media_files=None, force_document=False): + """Out-of-process Email delivery via SMTP (one-shot); standalone_sender_fn contract.""" extra = getattr(pconfig, "extra", {}) or {} address = extra.get("address") or _get_secret("EMAIL_ADDRESS", "") password = _get_secret("EMAIL_PASSWORD", "") smtp_host = extra.get("smtp_host") or _get_secret("EMAIL_SMTP_HOST", "") - try: - smtp_port = int(_get_secret("EMAIL_SMTP_PORT", "587") or "587") - except (ValueError, TypeError): - smtp_port = 587 + smtp_port = _esecret_int("EMAIL_SMTP_PORT", 587) smtp_security = _normalize_security( _get_secret("EMAIL_SMTP_SECURITY", "") or extra.get("smtp_security"), default="tls" if smtp_port == 465 else "starttls", ) smtp_tls_verify = _esecret_bool( - "EMAIL_SMTP_TLS_VERIFY", - is_truthy_value(extra.get("smtp_tls_verify"), default=True), + "EMAIL_SMTP_TLS_VERIFY", is_truthy_value(extra.get("smtp_tls_verify"), default=True) ) if not all([address, password, smtp_host]): @@ -1542,16 +1040,7 @@ async def _standalone_send( msg["Date"] = formatdate(localtime=True) ctx = _tls_context(smtp_tls_verify, smtp_host) - if smtp_security == "tls": - server = smtplib.SMTP_SSL(smtp_host, smtp_port, context=ctx) - else: - server = smtplib.SMTP(smtp_host, smtp_port) - if smtp_security == "starttls": - try: - server.starttls(context=ctx) - except Exception: - server.close() - raise + server = _open_smtp(smtp_host, smtp_port, smtp_security, ctx, smtplib.SMTP, smtplib.SMTP_SSL) server.login(address, password) server.send_message(msg) server.quit() @@ -1565,9 +1054,7 @@ async def _standalone_send( def _is_connected(config) -> bool: - """Email is connected when an address is configured (in PlatformConfig.extra - or via EMAIL_ADDRESS). Mirrors the legacy - _PLATFORM_CONNECTED_CHECKERS[Platform.EMAIL] = bool(extra.get('address')).""" + """Connected when an address is configured (PlatformConfig.extra or EMAIL_ADDRESS).""" extra = getattr(config, "extra", {}) or {} if extra.get("address"): return True diff --git a/plugins/platforms/homeassistant/adapter.py b/plugins/platforms/homeassistant/adapter.py index bfdd136cdf..dd185be7a9 100644 --- a/plugins/platforms/homeassistant/adapter.py +++ b/plugins/platforms/homeassistant/adapter.py @@ -1,10 +1,7 @@ -""" -Home Assistant platform adapter. +"""Home Assistant platform adapter. -Connects to the HA WebSocket API for real-time event monitoring. -State-change events are converted to MessageEvent objects and forwarded -to the agent for processing. Outbound messages are delivered as HA -persistent notifications. +Listens on the HA WebSocket API; ``state_changed`` events become MessageEvents. +Outbound messages are delivered as HA persistent notifications. Requires: - aiohttp (already in messaging extras) @@ -37,28 +34,7 @@ from gateway.platforms.base import ( SendResult, ) -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 +from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret logger = logging.getLogger(__name__) @@ -70,83 +46,71 @@ def check_ha_requirements() -> bool: def validate_ha_config(config: PlatformConfig) -> bool: - """Return True when Home Assistant has enough credential config to connect.""" - token = (getattr(config, "token", None) or _get_scoped_secret("HASS_TOKEN", "")).strip() - return bool(token) + """True when Home Assistant has enough credential config to connect.""" + return bool((getattr(config, "token", None) or _get_scoped_secret("HASS_TOKEN", "")).strip()) + + +def _domain_of(entity_id: str) -> str: + return entity_id.split(".")[0] if "." in entity_id else "" + + +def _on_off(val: str) -> str: + return "on" if val == "on" else "off" + + +def _triggered(val: str) -> str: + return "triggered" if val == "on" else "cleared" class HomeAssistantAdapter(BasePlatformAdapter): - """ - Home Assistant WebSocket adapter. - - Subscribes to ``state_changed`` events and forwards them as - MessageEvent objects. Supports domain/entity filtering and - per-entity cooldowns to avoid event floods. - """ + """HA WebSocket adapter: ``state_changed`` events → MessageEvents, with + domain/entity filtering and per-entity cooldowns against event floods.""" MAX_MESSAGE_LENGTH = 4096 - - # Reconnection backoff schedule (seconds) - _BACKOFF_STEPS = [5, 10, 30, 60] + _BACKOFF_STEPS = [5, 10, 30, 60] # reconnect backoff (seconds) def __init__(self, config: PlatformConfig): super().__init__(config, Platform.HOMEASSISTANT) - - # Connection state self._session: Optional["aiohttp.ClientSession"] = None self._ws: Optional["aiohttp.ClientWebSocketResponse"] = None self._rest_session: Optional["aiohttp.ClientSession"] = None self._listen_task: Optional[asyncio.Task] = None self._msg_id: int = 0 - # Configuration from extra extra = config.extra or {} - token = config.token or _get_scoped_secret("HASS_TOKEN", "") url = extra.get("url") or os.getenv("HASS_URL", "http://homeassistant.local:8123") self._hass_url: str = url.rstrip("/") - self._hass_token: str = token + self._hass_token: str = config.token or _get_scoped_secret("HASS_TOKEN", "") - # Event filtering self._watch_domains: Set[str] = set(extra.get("watch_domains", [])) self._watch_entities: Set[str] = set(extra.get("watch_entities", [])) self._ignore_entities: Set[str] = set(extra.get("ignore_entities", [])) self._watch_all: bool = bool(extra.get("watch_all", False)) self._cooldown_seconds: int = int(extra.get("cooldown_seconds", 30)) - - # Cooldown tracking: entity_id -> last_event_timestamp - self._last_event_time: Dict[str, float] = {} + self._last_event_time: Dict[str, float] = {} # entity_id -> last event ts def _next_id(self) -> int: - """Return the next WebSocket message ID.""" self._msg_id += 1 return self._msg_id - # ------------------------------------------------------------------ - # Connection lifecycle - # ------------------------------------------------------------------ + # -- Connection lifecycle ----------------------------------------------- async def connect(self, *, is_reconnect: bool = False) -> bool: """Connect to HA WebSocket API and subscribe to events.""" if not AIOHTTP_AVAILABLE: logger.warning("[%s] aiohttp not installed. Run: pip install aiohttp", self.name) return False - if not self._hass_token: logger.warning("[%s] No HASS_TOKEN configured", self.name) return False try: - success = await self._ws_connect() - if not success: + if not await self._ws_connect(): return False - # Dedicated REST session for send() calls self._rest_session = aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=30), - trust_env=gateway_trust_env(), + timeout=aiohttp.ClientTimeout(total=30), trust_env=gateway_trust_env(), ) - - # Warn if no event filters are configured if not self._watch_domains and not self._watch_entities and not self._watch_all: logger.warning( "[%s] No watch_domains, watch_entities, or watch_all configured. " @@ -154,69 +118,49 @@ class HomeAssistantAdapter(BasePlatformAdapter): "your HA platform config to receive events.", self.name, ) - - # Start background listener self._listen_task = asyncio.create_task(self._listen_loop()) self._running = True logger.info("[%s] Connected to %s", self.name, self._hass_url) - # Plugin-registered native handlers (ctx.register_platform_handler). self._wire_plugin_handlers(None) return True - except Exception as e: logger.error("[%s] Failed to connect: %s", self.name, e) return False async def _ws_connect(self) -> bool: - """Establish WebSocket connection and authenticate.""" + """Open the WebSocket, authenticate, and subscribe to ``state_changed``.""" ws_url = self._hass_url.replace("https://", "wss://").replace("http://", "ws://") - ws_url = f"{ws_url}/api/websocket" - self._session = aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=30), - trust_env=gateway_trust_env(), + timeout=aiohttp.ClientTimeout(total=30), trust_env=gateway_trust_env(), ) - self._ws = await self._session.ws_connect(ws_url, heartbeat=30, timeout=30) + self._ws = await self._session.ws_connect(f"{ws_url}/api/websocket", heartbeat=30, timeout=30) - # Step 1: Receive auth_required msg = await self._ws.receive_json() if msg.get("type") != "auth_required": logger.error("Expected auth_required, got: %s", msg.get("type")) await self._cleanup_ws() return False - # Step 2: Send auth - await self._ws.send_json({ - "type": "auth", - "access_token": self._hass_token, - }) - - # Step 3: Wait for auth_ok + await self._ws.send_json({"type": "auth", "access_token": self._hass_token}) msg = await self._ws.receive_json() if msg.get("type") != "auth_ok": logger.error("Auth failed: %s", msg) await self._cleanup_ws() return False - # Step 4: Subscribe to state_changed events - sub_id = self._next_id() await self._ws.send_json({ - "id": sub_id, + "id": self._next_id(), "type": "subscribe_events", "event_type": "state_changed", }) - - # Verify subscription acknowledgement msg = await self._ws.receive_json() if not msg.get("success"): logger.error("Failed to subscribe to events: %s", msg) await self._cleanup_ws() return False - return True async def _cleanup_ws(self) -> None: - """Close WebSocket and session.""" if self._ws and not self._ws.closed: await self._ws.close() self._ws = None @@ -225,7 +169,6 @@ class HomeAssistantAdapter(BasePlatformAdapter): self._session = None async def disconnect(self) -> None: - """Disconnect from Home Assistant.""" self._running = False if self._listen_task: self._listen_task.cancel() @@ -241,14 +184,11 @@ class HomeAssistantAdapter(BasePlatformAdapter): self._rest_session = None logger.info("[%s] Disconnected", self.name) - # ------------------------------------------------------------------ - # Event listener - # ------------------------------------------------------------------ + # -- Event listener ----------------------------------------------------- async def _listen_loop(self) -> None: """Main event loop with automatic reconnection.""" backoff_idx = 0 - while self._running: try: await self._read_events() @@ -259,8 +199,6 @@ class HomeAssistantAdapter(BasePlatformAdapter): if not self._running: return - - # Reconnect with backoff delay = self._BACKOFF_STEPS[min(backoff_idx, len(self._BACKOFF_STEPS) - 1)] logger.info("[%s] Reconnecting in %ds...", self.name, delay) await asyncio.sleep(delay) @@ -268,9 +206,8 @@ class HomeAssistantAdapter(BasePlatformAdapter): try: await self._cleanup_ws() - success = await self._ws_connect() - if success: - backoff_idx = 0 # Reset on successful reconnect + if await self._ws_connect(): + backoff_idx = 0 logger.info("[%s] Reconnected", self.name) except Exception as e: logger.warning("[%s] Reconnection failed: %s", self.name, e) @@ -290,46 +227,32 @@ class HomeAssistantAdapter(BasePlatformAdapter): elif ws_msg.type in {aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR}: break + def _passes_filters(self, entity_id: str) -> bool: + """Closed by default: requires watch_domains, watch_entities, or watch_all.""" + if entity_id in self._ignore_entities: + return False + if self._watch_domains or self._watch_entities: + return _domain_of(entity_id) in self._watch_domains or entity_id in self._watch_entities + return self._watch_all + async def _handle_ha_event(self, event: Dict[str, Any]) -> None: """Process a state_changed event from Home Assistant.""" event_data = event.get("data", {}) entity_id: str = event_data.get("entity_id", "") - - if not entity_id: + if not entity_id or not self._passes_filters(entity_id): return - # Apply ignore filter - if entity_id in self._ignore_entities: - return - - # Apply domain/entity watch filters (closed by default — require - # explicit watch_domains, watch_entities, or watch_all to forward) - domain = entity_id.split(".")[0] if "." in entity_id else "" - if self._watch_domains or self._watch_entities: - domain_match = domain in self._watch_domains if self._watch_domains else False - entity_match = entity_id in self._watch_entities if self._watch_entities else False - if not domain_match and not entity_match: - return - elif not self._watch_all: - # No filters configured and watch_all is off — drop the event - return - - # Apply cooldown now = time.time() - last = self._last_event_time.get(entity_id, 0) - if (now - last) < self._cooldown_seconds: + if (now - self._last_event_time.get(entity_id, 0)) < self._cooldown_seconds: return self._last_event_time[entity_id] = now - # Build human-readable message - old_state = event_data.get("old_state", {}) - new_state = event_data.get("new_state", {}) - message = self._format_state_change(entity_id, old_state, new_state) - + message = self._format_state_change( + entity_id, event_data.get("old_state", {}), event_data.get("new_state", {}), + ) if not message: return - # Build MessageEvent and forward to handler source = self.build_source( chat_id="ha_events", chat_name="Home Assistant Events", @@ -337,16 +260,13 @@ class HomeAssistantAdapter(BasePlatformAdapter): user_id="homeassistant", user_name="Home Assistant", ) - - msg_event = MessageEvent( + await self.handle_message(MessageEvent( text=message, message_type=MessageType.TEXT, source=source, message_id=f"ha_{entity_id}_{int(now)}", timestamp=datetime.now(), - ) - - await self.handle_message(msg_event) + )) @staticmethod def _format_state_change( @@ -357,62 +277,34 @@ class HomeAssistantAdapter(BasePlatformAdapter): """Convert a state_changed event into a human-readable description.""" if not new_state: return None - old_val = old_state.get("state", "unknown") if old_state else "unknown" new_val = new_state.get("state", "unknown") - - # Skip if state didn't actually change if old_val == new_val: return None - friendly_name = new_state.get("attributes", {}).get("friendly_name", entity_id) - domain = entity_id.split(".")[0] if "." in entity_id else "" + attrs = new_state.get("attributes", {}) + name = attrs.get("friendly_name", entity_id) + domain = _domain_of(entity_id) - # Domain-specific formatting if domain == "climate": - attrs = new_state.get("attributes", {}) temp = attrs.get("current_temperature", "?") target = attrs.get("temperature", "?") return ( - f"[Home Assistant] {friendly_name}: HVAC mode changed from " + f"[Home Assistant] {name}: HVAC mode changed from " f"'{old_val}' to '{new_val}' (current: {temp}, target: {target})" ) - if domain == "sensor": - unit = new_state.get("attributes", {}).get("unit_of_measurement", "") - return ( - f"[Home Assistant] {friendly_name}: changed from " - f"{old_val}{unit} to {new_val}{unit}" - ) - + unit = attrs.get("unit_of_measurement", "") + return f"[Home Assistant] {name}: changed from {old_val}{unit} to {new_val}{unit}" if domain == "binary_sensor": - return ( - f"[Home Assistant] {friendly_name}: " - f"{'triggered' if new_val == 'on' else 'cleared'} " - f"(was {'triggered' if old_val == 'on' else 'cleared'})" - ) - + return f"[Home Assistant] {name}: {_triggered(new_val)} (was {_triggered(old_val)})" if domain in {"light", "switch", "fan"}: - return ( - f"[Home Assistant] {friendly_name}: turned " - f"{'on' if new_val == 'on' else 'off'}" - ) - + return f"[Home Assistant] {name}: turned {_on_off(new_val)}" if domain == "alarm_control_panel": - return ( - f"[Home Assistant] {friendly_name}: alarm state changed from " - f"'{old_val}' to '{new_val}'" - ) + return f"[Home Assistant] {name}: alarm state changed from '{old_val}' to '{new_val}'" + return f"[Home Assistant] {name} ({entity_id}): changed from '{old_val}' to '{new_val}'" - # Generic fallback - return ( - f"[Home Assistant] {friendly_name} ({entity_id}): " - f"changed from '{old_val}' to '{new_val}'" - ) - - # ------------------------------------------------------------------ - # Outbound messaging - # ------------------------------------------------------------------ + # -- Outbound messaging ------------------------------------------------- async def send( self, @@ -423,66 +315,37 @@ class HomeAssistantAdapter(BasePlatformAdapter): ) -> SendResult: """Send a notification via HA REST API (persistent_notification.create). - Uses the REST API instead of WebSocket to avoid a race condition - with the event listener loop that reads from the same WS connection. + REST rather than the WebSocket, to avoid racing the listener loop that + reads from the same WS connection. """ url = f"{self._hass_url}/api/services/persistent_notification/create" - headers = { - "Authorization": f"Bearer {self._hass_token}", - "Content-Type": "application/json", - } - payload = { - "title": "Hermes Agent", - "message": content[:self.MAX_MESSAGE_LENGTH], - } + headers = {"Authorization": f"Bearer {self._hass_token}", "Content-Type": "application/json"} + payload = {"title": "Hermes Agent", "message": content[:self.MAX_MESSAGE_LENGTH]} + + async def _post(session) -> SendResult: + async with session.post( + url, headers=headers, json=payload, timeout=aiohttp.ClientTimeout(total=10), + ) as resp: + if resp.status < 300: + return SendResult(success=True, message_id=uuid.uuid4().hex[:12]) + body = await resp.text() + return SendResult(success=False, error=f"HTTP {resp.status}: {body}") try: if self._rest_session: - async with self._rest_session.post( - url, - headers=headers, - json=payload, - timeout=aiohttp.ClientTimeout(total=10), - ) as resp: - if resp.status < 300: - return SendResult(success=True, message_id=uuid.uuid4().hex[:12]) - else: - body = await resp.text() - return SendResult(success=False, error=f"HTTP {resp.status}: {body}") - else: - async with aiohttp.ClientSession(trust_env=gateway_trust_env()) as session: - async with session.post( - url, - headers=headers, - json=payload, - timeout=aiohttp.ClientTimeout(total=10), - ) as resp: - if resp.status < 300: - return SendResult(success=True, message_id=uuid.uuid4().hex[:12]) - else: - body = await resp.text() - return SendResult(success=False, error=f"HTTP {resp.status}: {body}") - + return await _post(self._rest_session) + async with aiohttp.ClientSession(trust_env=gateway_trust_env()) as session: + return await _post(session) except asyncio.TimeoutError: return SendResult(success=False, error="Timeout sending notification to HA") except Exception as e: return SendResult(success=False, error=str(e)) - async def send_typing(self, chat_id: str, metadata=None) -> None: - """No typing indicator for Home Assistant.""" - async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: - """Return basic info about the HA event channel.""" - return { - "name": "Home Assistant Events", - "type": "channel", - "url": self._hass_url, - } + return {"name": "Home Assistant Events", "type": "channel", "url": self._hass_url} -# --------------------------------------------------------------------------- -# Standalone (out-of-process) sender — used by cron deliver=homeassistant -# --------------------------------------------------------------------------- +# -- Standalone (out-of-process) sender — cron deliver=homeassistant --------- async def _standalone_send( @@ -494,23 +357,11 @@ async def _standalone_send( media_files: Optional[list] = None, force_document: bool = False, ) -> Dict[str, Any]: - """Send a notification via the HA ``notify.notify`` service without a - live gateway adapter. + """Send via the HA ``notify.notify`` service without a live gateway adapter. - Used by ``tools/send_message_tool._send_via_adapter`` when the gateway - runner is not in this process (typical for cron jobs running - out-of-process). The HTTP path is the same one the legacy - ``_send_homeassistant`` helper used in ``tools/send_message_tool.py`` - before this migration. - - Reads ``HASS_TOKEN`` from ``pconfig.token`` (set by the gateway config - loader from env) and falls back to the ``HASS_TOKEN`` env var. Server - URL comes from ``pconfig.extra["url"]`` (seeded by the env loader in - ``gateway/config.py``) or the ``HASS_URL`` env var. - - ``thread_id``, ``media_files`` and ``force_document`` are accepted for - signature parity with other standalone senders. HA notifications have - no native threading or attachment model — these arguments are ignored. + Token: ``pconfig.token`` then ``HASS_TOKEN``; URL: ``pconfig.extra["url"]`` + then ``HASS_URL``. ``thread_id``/``media_files``/``force_document`` are + signature parity only — HA notifications have no threads or attachments. """ if not AIOHTTP_AVAILABLE: return {"error": "aiohttp not installed. Run: pip install aiohttp"} @@ -519,69 +370,37 @@ async def _standalone_send( hass_url = (extra.get("url") or os.getenv("HASS_URL", "")).rstrip("/") token = (getattr(pconfig, "token", None) or _get_scoped_secret("HASS_TOKEN", "")).strip() if not hass_url or not token: - return { - "error": ( - "Home Assistant standalone send: HASS_URL and HASS_TOKEN " - "must both be set" - ) - } + return {"error": "Home Assistant standalone send: HASS_URL and HASS_TOKEN must both be set"} url = f"{hass_url}/api/services/notify/notify" - headers = { - "Authorization": f"Bearer {token}", - "Content-Type": "application/json", - } + headers = {"Authorization": f"Bearer {token}", "Content-Type": "application/json"} payload = {"message": message, "target": chat_id} - try: async with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=30), - trust_env=gateway_trust_env(), + timeout=aiohttp.ClientTimeout(total=30), trust_env=gateway_trust_env(), ) as session: async with session.post(url, headers=headers, json=payload) as resp: if resp.status not in {200, 201}: body = await resp.text() - return { - "error": ( - f"Home Assistant API error ({resp.status}): {body}" - ) - } - return { - "success": True, - "platform": "homeassistant", - "chat_id": chat_id, - } + return {"error": f"Home Assistant API error ({resp.status}): {body}"} + return {"success": True, "platform": "homeassistant", "chat_id": chat_id} except asyncio.TimeoutError: return {"error": "Timeout sending notification to Home Assistant"} except Exception as e: return {"error": f"Home Assistant send failed: {e}"} -# --------------------------------------------------------------------------- -# is_connected probe -# --------------------------------------------------------------------------- - - def _is_connected(config) -> bool: - """Home Assistant is considered connected when ``HASS_TOKEN`` is set. + """Connected when ``HASS_TOKEN`` is set. - Looks up via ``hermes_cli.gateway.get_env_value`` at call time (not via - the plugin's own bound import) so tests that patch - ``gateway_mod.get_env_value`` can suppress ambient ``HASS_TOKEN`` env - vars. Matches what the legacy connected-platforms check did before - this migration. + Looked up via ``hermes_cli.gateway.get_env_value`` at call time so tests + patching ``gateway_mod.get_env_value`` can suppress ambient env vars. """ import hermes_cli.gateway as gateway_mod return bool((gateway_mod.get_env_value("HASS_TOKEN") or "").strip()) -# --------------------------------------------------------------------------- -# Plugin registration entry point -# --------------------------------------------------------------------------- - - def _build_adapter(config): - """Factory wrapper that constructs HomeAssistantAdapter from a PlatformConfig.""" return HomeAssistantAdapter(config) @@ -596,15 +415,8 @@ def register(ctx) -> None: is_connected=_is_connected, required_env=["HASS_TOKEN"], install_hint="pip install aiohttp", - # Out-of-process cron delivery via the HA ``notify.notify`` service. - # Without this hook, ``deliver=homeassistant`` cron jobs would fail - # with "No live adapter" when cron runs separately from the gateway. - # Mirrors the Discord / Teams / Mattermost pattern. - standalone_sender_fn=_standalone_send, - # HA notification message cap — matches MAX_MESSAGE_LENGTH on the - # adapter class above. + standalone_sender_fn=_standalone_send, # out-of-process cron delivery via notify.notify max_message_length=HomeAssistantAdapter.MAX_MESSAGE_LENGTH, - # Display emoji="🏠", allow_update_command=True, ) diff --git a/plugins/platforms/line/adapter.py b/plugins/platforms/line/adapter.py index 1150556b4e..2129680804 100644 --- a/plugins/platforms/line/adapter.py +++ b/plugins/platforms/line/adapter.py @@ -1,61 +1,19 @@ -""" -LINE Messaging API platform adapter for Hermes Agent. +"""LINE Messaging API platform adapter for Hermes Agent. -A bundled platform plugin that runs an aiohttp webhook server, accepts LINE -webhook events (signature-verified), and relays messages to/from the agent -via the standard ``BasePlatformAdapter`` interface. +An aiohttp webhook server accepts signature-verified LINE events and relays them +through ``BasePlatformAdapter``. Design highlights: -Design highlights ------------------ - -**Reply token preferred, Push fallback.** LINE's reply token is single-use -and expires roughly 60 seconds after the inbound event. We try Reply first -(it's free) and fall back to the metered Push API when the token is absent, -expired, or rejected by the API. - -**Slow-LLM postback button (optional).** When the LLM is still running past -``slow_response_threshold`` seconds (default 45, leaving 15s margin on the -60s reply-token TTL), we burn the original reply token to send a Template -Buttons bubble — the user taps it later to receive the cached answer via a -*fresh* reply token (also free). State machine: PENDING → READY → DELIVERED, -with ERROR for cancelled runs. Set the threshold to 0 to disable the -button and always Push-fallback instead. - -**Three-allowlist gating.** Separate allowlists for users (U-prefixed), -groups (C-prefixed), and rooms (R-prefixed). ``LINE_ALLOW_ALL_USERS=true`` -is a dev-only escape hatch. - -**Media via public HTTPS.** LINE's Messaging API does *not* accept -binary uploads — images, audio, and video must be reachable HTTPS URLs. -We register registered tempfiles under ``/line/media//`` -served by the same aiohttp app, with an allowed-roots traversal guard. -``LINE_PUBLIC_URL`` (e.g. ``https://my-tunnel.example.com``) overrides -the host:port construction so URLs are reachable when the bind is a -wildcard/dual-stack listener or behind a reverse proxy. - -**5-message batching.** LINE accepts at most 5 message objects per -Reply/Push call; longer responses are smart-chunked at 4500 chars -(LINE per-bubble limit is 5000) and batched. - -Synthesis credits ------------------ - -This file is a synthesis of seven open community PRs adding LINE support -to Hermes Agent. It deliberately ports the *strongest* idea from each into -a single plugin-form module that requires zero core edits: - -* PR #18153 (leepoweii) — Template Buttons postback cache state machine, - Markdown URL preservation, system-message bypass. -* PR #8398 (yuga-hashimoto) — media URL serving with traversal guard, - send_voice / send_video, ``LINE_PUBLIC_URL`` env, macOS ``/tmp`` root. -* PR #16832 (jethac) — config wiring style, voice/image tests. -* PR #21023 (perng) — plugin-form skeleton (the only one already - modeled on ``ADDING_A_PLATFORM.md``), reply→push fallback at 50s TTL, - loading-animation indicator, source dispatcher. -* PR #14942 (soichiyo) — Cloudflare-tunnel operating model (docs only). -* PR #14988 (David-0x221Eight) — text-first scope discipline. -* PR #6676 (liyoungc) — Push-only mode (used as the ``threshold=0`` - fallback path here). +* Reply token preferred (free, single-use, ~60s TTL), metered Push as fallback. +* Optional slow-LLM postback button: past ``slow_response_threshold`` (default 45s) + the reply token is burned on a Template Buttons bubble; tapping it yields a fresh + free token carrying the cached answer (PENDING → READY → DELIVERED, ERROR on + cancel). Threshold 0 disables the button. +* Three allowlists: users (U…), groups (C…), rooms (R…); ``LINE_ALLOW_ALL_USERS`` + is a dev-only escape hatch. +* Media via public HTTPS: LINE takes no binary uploads, so local files are served + under ``/line/media//`` (allowed-roots guard); ``LINE_PUBLIC_URL`` + overrides host:port behind tunnels/proxies or wildcard binds. +* ≤5 message objects per call; text chunked at 4500 chars (bubble hard limit 5000). """ from __future__ import annotations @@ -77,58 +35,24 @@ import time import uuid from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Dict, List, Optional, Set, Tuple +from typing import Any, Callable, Dict, List, Optional, Set, Tuple from urllib.parse import quote as _urlquote -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 +from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Lazy / function-level imports for gateway internals are NOT used here — -# the plugin discovery flow imports adapter.py late enough that gateway is -# already loaded. -# --------------------------------------------------------------------------- - +# Plugin discovery imports adapter.py late enough that gateway is already loaded, +# so gateway internals are imported eagerly here. from gateway.platforms.base import ( - gateway_trust_env, - BasePlatformAdapter, - MessageEvent, - MessageType, - SendResult, - cache_audio_from_bytes, - cache_document_from_bytes, - cache_image_from_bytes, - cache_video_from_bytes, + gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, + cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes, cache_video_from_bytes, ) from gateway.config import Platform -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- +# --- Constants --- LINE_REPLY_URL = "https://api.line.me/v2/bot/message/reply" LINE_PUSH_URL = "https://api.line.me/v2/bot/message/push" @@ -137,10 +61,10 @@ LINE_CONTENT_URL_FMT = "https://api-data.line.me/v2/bot/message/{message_id}/con LINE_BOT_INFO_URL = "https://api.line.me/v2/bot/info" # LINE Messaging API hard limits -LINE_PER_BUBBLE_CHARS = 5000 # Hard limit per text message object -LINE_SAFE_BUBBLE_CHARS = 4500 # Conservative limit for chunking -LINE_MAX_MESSAGES_PER_CALL = 5 # API rejects >5 messages per Reply/Push -LINE_REPLY_TOKEN_TTL_SECONDS = 50 # Conservative cap below LINE's ~60s +LINE_PER_BUBBLE_CHARS = 5000 +LINE_SAFE_BUBBLE_CHARS = 4500 # conservative chunking limit +LINE_MAX_MESSAGES_PER_CALL = 5 +LINE_REPLY_TOKEN_TTL_SECONDS = 50 # below LINE's ~60s # Webhook hardening WEBHOOK_BODY_MAX_BYTES = 1_048_576 # 1 MiB — webhooks are tiny JSON @@ -148,34 +72,17 @@ DEFAULT_WEBHOOK_PORT = 8646 DEFAULT_WEBHOOK_PATH = "/line/webhook" DEFAULT_MEDIA_PATH_PREFIX = "/line/media" -# Default bind host. ``None`` tells aiohttp/asyncio's ``create_server`` to -# bind BOTH address families (IPv4 + IPv6) — the portable dual-stack default. -# Mirrors gateway/platforms/webhook.py DEFAULT_HOST (commit d542894ad). -# -# Why not "0.0.0.0" (the old default) or "::"? -# - "0.0.0.0" binds IPv4 ONLY. On IPv6-only private networks — notably -# Fly.io 6PN, where the hosted edge router reverse-proxies LINE ingest to -# ``.internal:8646`` over an ``fdaa:…`` IPv6 address — an IPv4-only -# listener is unreachable: dial refused → customer-visible 502 (NS-603). -# - "::" is NOT a safe fix: on hosts where the kernel sets IPV6_V6ONLY=1 -# (verified on Fly machines), binding "::" yields an IPv6-ONLY socket, -# breaking IPv4 loopback health probes. -# - ``None`` asks the event loop to create a listening socket per resolved -# family, so both 127.0.0.1 (v4) and the 6PN fdaa (v6) are served -# regardless of the bindv6only sysctl. Users can still pin a host via -# ``LINE_HOST`` or ``platforms.line.extra.host``. +# ``None`` → asyncio binds BOTH address families (mirrors gateway/platforms/webhook.py). +# "0.0.0.0" is unreachable on IPv6-only networks (Fly.io 6PN → 502s); "::" breaks IPv4 +# loopback probes on IPV6_V6ONLY=1 hosts. Pin via ``LINE_HOST`` / ``extra.host``. DEFAULT_HOST = None -# Hosts that mean "listening on every interface" — i.e. the bind address is -# not a name LINE's servers could ever fetch media from, so a public base URL -# is required for outbound media. +# Bind hosts LINE's servers could never fetch media from → public URL required. _WILDCARD_HOSTS = frozenset({"0.0.0.0", "::", ""}) # Slow-LLM postback button defaults DEFAULT_SLOW_RESPONSE_THRESHOLD = 45.0 # seconds; 0 disables -DEFAULT_PENDING_REPLY_TEXT = ( - "🤔 Still thinking. Tap below to fetch the answer when it's ready." -) +DEFAULT_PENDING_REPLY_TEXT = "🤔 Still thinking. Tap below to fetch the answer when it's ready." DEFAULT_BUTTON_LABEL = "Get answer" DEFAULT_DELIVERED_TEXT = "Already replied ✅" DEFAULT_INTERRUPTED_TEXT = "Run was interrupted before completion." @@ -185,24 +92,14 @@ MEDIA_TOKEN_TTL_SECONDS = 1800 # 30 minutes; LINE caches the URL aggressively LINE_IMAGE_MAX_BYTES = 10 * 1024 * 1024 # 10 MB per LINE docs LINE_AV_MAX_BYTES = 200 * 1024 * 1024 # 200 MB for voice/video -# Map LINE webhook message types to the normalized MessageType the gateway -# routes on. LINE has no separate "voice" type — audio messages are recorded -# voice clips, so they map to VOICE (which the gateway sends through STT), -# mirroring how Telegram/WhatsApp classify voice notes. Anything unknown -# falls back to TEXT. +# LINE message type → normalized MessageType. LINE audio is recorded voice clips → +# VOICE (STT path), like Telegram/WhatsApp. Unknown types fall back to TEXT. _LINE_MESSAGE_TYPES = { - "text": MessageType.TEXT, - "image": MessageType.PHOTO, - "video": MessageType.VIDEO, - "audio": MessageType.VOICE, - "file": MessageType.DOCUMENT, - "location": MessageType.LOCATION, - "sticker": MessageType.STICKER, + "text": MessageType.TEXT, "image": MessageType.PHOTO, "video": MessageType.VIDEO, "audio": MessageType.VOICE, + "file": MessageType.DOCUMENT, "location": MessageType.LOCATION, "sticker": MessageType.STICKER, } -# A 1×1 transparent PNG used as fallback video preview thumbnail when no -# explicit preview is supplied — LINE requires ``previewImageUrl`` for -# video messages. Sourced from the Python stdlib (no Pillow dependency). +# 1×1 transparent PNG: fallback video preview (LINE requires ``previewImageUrl``). _FALLBACK_PNG_PREVIEW = bytes.fromhex( "89504e470d0a1a0a0000000d49484452000000010000000108060000001f15c4" "890000000d49444154789c63000100000005000100377a7ff20000000049454e" @@ -210,9 +107,7 @@ _FALLBACK_PNG_PREVIEW = bytes.fromhex( ) -# --------------------------------------------------------------------------- -# Markdown stripping (URL-preserving) -# --------------------------------------------------------------------------- +# --- Markdown stripping (URL-preserving) --- _MD_LINK_RE = re.compile(r"\[([^\]]+)\]\((https?://[^\s)]+)\)") _MD_BOLD_RE = re.compile(r"\*\*(.+?)\*\*") @@ -224,55 +119,26 @@ _MD_BULLET_RE = re.compile(r"^[\s]*[-*+]\s+", re.MULTILINE) def strip_markdown_preserving_urls(text: str) -> str: - """Strip Markdown that LINE can't render, but keep URLs usable. - - LINE's text bubble has zero Markdown support — bold, italics, code - fences, headings, and bullet markers all render as literal characters. - URLs *are* auto-linked by the client, but only when they appear bare - (not inside ``[label](url)`` syntax). This converts ``[label](url)`` - to ``label (url)`` so the URL remains tappable, then strips the rest. - - Source: PR #18153 (leepoweii) — adapted to keep code-block content - visible (LINE users frequently want command snippets to land as - plain text, not be eaten by the fence). - """ + """Strip Markdown LINE can't render; ``[label](url)`` → ``label (url)`` keeps URLs + tappable (LINE auto-links bare URLs only). Code-block content is kept.""" if not text: return text - - # Code blocks first — keep the inner content, drop the fences. - def _unfence(m: re.Match) -> str: - return m.group(1).rstrip("\n") - text = _MD_CODE_BLOCK_RE.sub(_unfence, text) - - # Inline code: keep content, drop backticks. + text = _MD_CODE_BLOCK_RE.sub(lambda m: m.group(1).rstrip("\n"), text) text = _MD_CODE_INLINE_RE.sub(r"\1", text) - - # Markdown links → "label (url)" text = _MD_LINK_RE.sub(lambda m: f"{m.group(1)} ({m.group(2)})", text) - - # Bold/italic markers — strip. text = _MD_BOLD_RE.sub(r"\1", text) text = _MD_ITAL_RE.sub(r"\1", text) - - # Headings (#, ##) and bullet markers — strip the prefix only. text = _MD_HEADING_RE.sub("", text) text = _MD_BULLET_RE.sub("• ", text) - return text def split_for_line(text: str, max_chars: int = LINE_SAFE_BUBBLE_CHARS) -> List[str]: - """Split ``text`` into LINE-sized bubbles, preferring paragraph/line breaks. - - Returns at most ``LINE_MAX_MESSAGES_PER_CALL`` chunks; longer text is - truncated with an ellipsis on the final chunk to keep the response - deliverable in a single Reply/Push call. - """ + """Split into ≤5 LINE bubbles at paragraph/line/word breaks; overflow is ellipsised.""" if not text: return [] if len(text) <= max_chars: return [text] - chunks: List[str] = [] remaining = text while remaining and len(chunks) < LINE_MAX_MESSAGES_PER_CALL: @@ -280,65 +146,45 @@ def split_for_line(text: str, max_chars: int = LINE_SAFE_BUBBLE_CHARS) -> List[s chunks.append(remaining) remaining = "" break - # Try to break on the latest paragraph or newline within budget. - cut = remaining.rfind("\n\n", 0, max_chars) - if cut < int(max_chars * 0.5): - cut = remaining.rfind("\n", 0, max_chars) - if cut < int(max_chars * 0.5): - cut = remaining.rfind(" ", 0, max_chars) + for sep in ("\n\n", "\n", " "): # prefer paragraph, then line, then word breaks + cut = remaining.rfind(sep, 0, max_chars) + if cut >= int(max_chars * 0.5): + break if cut <= 0: cut = max_chars chunks.append(remaining[:cut].rstrip()) remaining = remaining[cut:].lstrip() - - if remaining: - # Truncate gracefully — caller already burned its 5-bubble budget. - if chunks: - tail = chunks[-1] - if len(tail) > max_chars - 1: - tail = tail[: max_chars - 1] - chunks[-1] = tail.rstrip() + "…" - else: - chunks.append(remaining[: max_chars - 1] + "…") + if remaining: # budget exhausted → ellipsis on the last bubble + tail = chunks[-1] + if len(tail) > max_chars - 1: + tail = tail[: max_chars - 1] + chunks[-1] = tail.rstrip() + "…" return chunks -# --------------------------------------------------------------------------- -# Webhook signature verification -# --------------------------------------------------------------------------- +# --- Webhook signature verification --- def verify_line_signature(body: bytes, signature: str, channel_secret: str) -> bool: - """Verify a LINE webhook's ``X-Line-Signature`` header. - - LINE signs the *raw* request body with HMAC-SHA256 keyed by the - channel secret, then base64-encodes the digest. Constant-time - comparison defends against timing oracles. - """ + """Verify ``X-Line-Signature``: base64(HMAC-SHA256(secret, raw body)), constant-time.""" if not signature or not channel_secret or body is None: return False try: - digest = hmac.new( - channel_secret.encode("utf-8"), - body, - hashlib.sha256, - ).digest() + digest = hmac.new(channel_secret.encode("utf-8"), body, hashlib.sha256).digest() expected = base64.b64encode(digest).decode("utf-8") except Exception: return False - # Compare as bytes: compare_digest raises TypeError on a str with - # non-ASCII characters, and the signature is a raw request header. + # Compare as bytes: compare_digest raises TypeError on non-ASCII str, and + # the signature is a raw request header. return hmac.compare_digest(expected.encode(), signature.encode()) -# --------------------------------------------------------------------------- -# Cache state machine — slow-LLM postback flow -# --------------------------------------------------------------------------- +# --- Cache state machine — slow-LLM postback flow --- class State(enum.Enum): PENDING = "pending" # button sent, LLM still running - READY = "ready" # LLM done, response cached, waiting for postback tap + READY = "ready" # response cached, waiting for postback tap DELIVERED = "delivered" - ERROR = "error" # LLM raised / interrupted; cached error text waiting + ERROR = "error" # LLM raised / interrupted; error text cached @dataclass @@ -350,22 +196,14 @@ class _CacheEntry: updated_at: float = field(default_factory=time.time) +_KEEP = object() # sentinel: leave entry.payload untouched + + class RequestCache: - """In-memory cache for slow-LLM postback retrieval. + """In-memory cache for slow-LLM postback retrieval.""" - PRs #18153 originally combined two TTLs — one for PENDING (24h) and - a shorter one for READY/DELIVERED/ERROR (1h). We keep the same model - here. - """ - - def __init__( - self, - ttl_seconds: int = 3600, - pending_ttl_seconds: int = 86400, - ) -> None: + def __init__(self) -> None: self._entries: Dict[str, _CacheEntry] = {} - self._ttl = ttl_seconds - self._pending_ttl = pending_ttl_seconds def register_pending(self, chat_id: str) -> str: rid = str(uuid.uuid4()) @@ -375,54 +213,26 @@ class RequestCache: def get(self, request_id: str) -> Optional[_CacheEntry]: return self._entries.get(request_id) - def set_ready(self, request_id: str, payload: Any) -> None: + def _transition(self, request_id: str, allowed: Set[State], state: State, payload: Any = _KEEP) -> None: entry = self._entries.get(request_id) - if entry is None or entry.state is not State.PENDING: + if entry is None or entry.state not in allowed: return - entry.state = State.READY - entry.payload = payload + entry.state = state + if payload is not _KEEP: + entry.payload = payload entry.updated_at = time.time() + def set_ready(self, request_id: str, payload: Any) -> None: + self._transition(request_id, {State.PENDING}, State.READY, payload) + def set_error(self, request_id: str, message: str) -> None: - entry = self._entries.get(request_id) - if entry is None or entry.state is not State.PENDING: - return - entry.state = State.ERROR - entry.payload = message - entry.updated_at = time.time() + self._transition(request_id, {State.PENDING}, State.ERROR, message) def mark_delivered(self, request_id: str) -> None: - entry = self._entries.get(request_id) - if entry is None or entry.state not in {State.READY, State.ERROR}: - return - entry.state = State.DELIVERED - entry.updated_at = time.time() - - def find_pending_for_chat(self, chat_id: str) -> Optional[str]: - for rid, entry in self._entries.items(): - if entry.state is State.PENDING and entry.chat_id == chat_id: - return rid - return None - - def prune(self) -> int: - now = time.time() - removed = 0 - for rid in list(self._entries.keys()): - entry = self._entries[rid] - if entry.state is State.PENDING: - if now - entry.created_at > self._pending_ttl: - del self._entries[rid] - removed += 1 - else: - if now - entry.updated_at > self._ttl: - del self._entries[rid] - removed += 1 - return removed + self._transition(request_id, {State.READY, State.ERROR}, State.DELIVERED) -# --------------------------------------------------------------------------- -# Inbound dedup -# --------------------------------------------------------------------------- +# --- Inbound dedup --- class _MessageDeduplicator: """Bounded LRU of LINE webhook event IDs to ignore at-least-once retries.""" @@ -444,28 +254,18 @@ class _MessageDeduplicator: return False -# --------------------------------------------------------------------------- -# Source / chat-id resolution -# --------------------------------------------------------------------------- +# --- Source / chat-id resolution --- + +# LINE source type → (id key, normalized chat_type) +_SOURCE_KINDS = {"group": ("groupId", "group"), "room": ("roomId", "room"), "user": ("userId", "dm")} + def _resolve_chat(source: Dict[str, Any]) -> Tuple[str, str]: - """Return ``(chat_id, chat_type)`` from a LINE event ``source`` block. - - LINE sources are one of: - * ``{"type": "user", "userId": "U..."}`` → 1:1 DM - * ``{"type": "group", "groupId": "C...", "userId": "U..."}`` → group chat - * ``{"type": "room", "roomId": "R...", "userId": "U..."}`` → multi-user room - - Source: PR #21023 (perng), unchanged. - """ - src_type = (source or {}).get("type", "") - if src_type == "group": - return source.get("groupId", ""), "group" - if src_type == "room": - return source.get("roomId", ""), "room" - if src_type == "user": - return source.get("userId", ""), "dm" - return "", "dm" + """Return ``(chat_id, chat_type)`` from a LINE event ``source`` block (user/group/room).""" + kind = _SOURCE_KINDS.get((source or {}).get("type", "")) + if kind is None: + return "", "dm" + return source.get(kind[0], ""), kind[1] def _allowed_for_source( @@ -476,92 +276,60 @@ def _allowed_for_source( group_ids: Set[str], room_ids: Set[str], ) -> bool: - """Three-list gate — credit PR #18153.""" + """Three-list gate: users, groups, rooms.""" if allow_all: return True - src_type = (source or {}).get("type", "") - if src_type == "user": - uid = source.get("userId", "") - return bool(uid) and uid in user_ids - if src_type == "group": - gid = source.get("groupId", "") - return bool(gid) and gid in group_ids - if src_type == "room": - rid = source.get("roomId", "") - return bool(rid) and rid in room_ids - return False + sid, chat_type = _resolve_chat(source) + return bool(sid) and sid in {"dm": user_ids, "group": group_ids, "room": room_ids}[chat_type] -# --------------------------------------------------------------------------- -# LINE Reply / Push HTTP client -# --------------------------------------------------------------------------- +# --- LINE Reply / Push HTTP client --- class _LineClient: - """Thin async wrapper around the LINE Messaging API. - - We use ``aiohttp`` directly to avoid a ``line-bot-sdk`` dependency - (the SDK pulls in its own httpx pin and the ergonomic gain is small - for the four endpoints we actually call). - """ + """Thin aiohttp wrapper around the LINE Messaging API (no ``line-bot-sdk`` dependency).""" def __init__(self, channel_access_token: str, *, timeout: float = 15.0) -> None: self._token = channel_access_token self._timeout = timeout - self._headers = { - "Authorization": f"Bearer {channel_access_token}", - "Content-Type": "application/json", - } + self._headers = {"Authorization": f"Bearer {channel_access_token}", "Content-Type": "application/json"} + + @staticmethod + def _session(timeout: float): + import aiohttp + return aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=timeout), trust_env=gateway_trust_env()) + + async def _post_messages(self, url: str, label: str, payload: Dict[str, Any]) -> None: + async with self._session(self._timeout) as session: + async with session.post(url, headers=self._headers, json=payload) as resp: + if resp.status >= 400: + body = await resp.text() + raise RuntimeError(f"LINE {label} {resp.status}: {body[:200]}") async def reply(self, reply_token: str, messages: List[Dict[str, Any]]) -> None: - import aiohttp - timeout = aiohttp.ClientTimeout(total=self._timeout) - async with aiohttp.ClientSession(timeout=timeout, trust_env=gateway_trust_env()) as session: - async with session.post( - LINE_REPLY_URL, - headers=self._headers, - json={"replyToken": reply_token, "messages": messages}, - ) as resp: - if resp.status >= 400: - body = await resp.text() - raise RuntimeError(f"LINE reply {resp.status}: {body[:200]}") + await self._post_messages(LINE_REPLY_URL, "reply", {"replyToken": reply_token, "messages": messages}) async def push(self, chat_id: str, messages: List[Dict[str, Any]]) -> None: - import aiohttp - timeout = aiohttp.ClientTimeout(total=self._timeout) - async with aiohttp.ClientSession(timeout=timeout, trust_env=gateway_trust_env()) as session: - async with session.post( - LINE_PUSH_URL, - headers=self._headers, - json={"to": chat_id, "messages": messages}, - ) as resp: - if resp.status >= 400: - body = await resp.text() - raise RuntimeError(f"LINE push {resp.status}: {body[:200]}") + await self._post_messages(LINE_PUSH_URL, "push", {"to": chat_id, "messages": messages}) async def loading(self, chat_id: str, seconds: int = 60) -> None: """Loading indicator (DM only). LINE rejects this for groups/rooms.""" if not chat_id or not chat_id.startswith("U"): return - import aiohttp + import aiohttp # noqa: F401 — ImportError must escape the swallow-all below # LINE caps loadingSeconds in 5-step increments, max 60. clamped = max(5, min(60, (seconds // 5) * 5 or 5)) try: - timeout = aiohttp.ClientTimeout(total=5.0) - async with aiohttp.ClientSession(timeout=timeout, trust_env=gateway_trust_env()) as session: + async with self._session(5.0) as session: await session.post( - LINE_LOADING_URL, - headers=self._headers, - json={"chatId": chat_id, "loadingSeconds": clamped}, + LINE_LOADING_URL, headers=self._headers, json={"chatId": chat_id, "loadingSeconds": clamped} ) except Exception as exc: # best-effort; never raise logger.debug("LINE loading indicator failed: %s", exc) async def fetch_content(self, message_id: str) -> bytes: """Download an inbound media message's binary content.""" - import aiohttp url = LINE_CONTENT_URL_FMT.format(message_id=message_id) - timeout = aiohttp.ClientTimeout(total=30.0) - async with aiohttp.ClientSession(timeout=timeout, trust_env=gateway_trust_env()) as session: + async with self._session(30.0) as session: async with session.get(url, headers={"Authorization": f"Bearer {self._token}"}) as resp: if resp.status >= 400: raise RuntimeError(f"LINE content {resp.status}") @@ -569,10 +337,9 @@ class _LineClient: async def get_bot_user_id(self) -> Optional[str]: """Fetch this channel's own userId so we can filter self-messages.""" - import aiohttp - timeout = aiohttp.ClientTimeout(total=10.0) + import aiohttp # noqa: F401 — ImportError must escape the swallow-all below try: - async with aiohttp.ClientSession(timeout=timeout, trust_env=gateway_trust_env()) as session: + async with self._session(10.0) as session: async with session.get(LINE_BOT_INFO_URL, headers=self._headers) as resp: if resp.status >= 400: return None @@ -582,9 +349,7 @@ class _LineClient: return None -# --------------------------------------------------------------------------- -# Message builders -# --------------------------------------------------------------------------- +# --- Message builders --- def _text_message(text: str) -> Dict[str, Any]: """Build a LINE text message object, capped to per-bubble max.""" @@ -593,73 +358,35 @@ def _text_message(text: str) -> Dict[str, Any]: return {"type": "text", "text": text} -def _image_message(original_url: str, preview_url: Optional[str] = None) -> Dict[str, Any]: - return { - "type": "image", - "originalContentUrl": original_url, - "previewImageUrl": preview_url or original_url, - } - - -def _audio_message(url: str, duration_ms: int = 1000) -> Dict[str, Any]: - return { - "type": "audio", - "originalContentUrl": url, - "duration": int(duration_ms), - } - - -def _video_message(url: str, preview_url: str) -> Dict[str, Any]: - return { - "type": "video", - "originalContentUrl": url, - "previewImageUrl": preview_url, - } +def _text_messages(content: str) -> List[Dict[str, Any]]: + """Markdown-strip, chunk and cap ``content`` into ≤5 LINE text messages.""" + chunks = split_for_line(strip_markdown_preserving_urls(content)) + return [_text_message(c) for c in chunks][:LINE_MAX_MESSAGES_PER_CALL] def build_postback_button_message( text: str, button_label: str, request_id: str ) -> Dict[str, Any]: - """Template Buttons message — the slow-LLM postback bubble. - - From PR #18153 (leepoweii). Template Buttons stay tappable from chat - history, unlike Quick Reply chips which are dismissed the moment any - new message arrives in the chat. - - LINE limits: ``text`` ≤ 160 chars, ``altText`` ≤ 400 chars. - """ + """Slow-LLM postback bubble. Template Buttons stay tappable from history (Quick + Reply chips vanish on the next message). LINE limits: text ≤160, altText ≤400.""" truncated = text if len(text) <= 160 else text[:157] + "..." alt = text if len(text) <= 400 else text[:397] + "..." + action = { + "type": "postback", + "label": button_label[:20] or "Get answer", + "data": json.dumps({"action": "show_response", "request_id": request_id}), + "displayText": button_label[:300] or "Get answer", + } return { "type": "template", "altText": alt, - "template": { - "type": "buttons", - "text": truncated, - "actions": [ - { - "type": "postback", - "label": button_label[:20] or "Get answer", - "data": json.dumps( - {"action": "show_response", "request_id": request_id} - ), - "displayText": button_label[:300] or "Get answer", - } - ], - }, + "template": {"type": "buttons", "text": truncated, "actions": [action]}, } -# Prefixes the gateway uses for system busy-acks (interrupting / queued / -# steered). When the postback cache has a PENDING entry we *bypass* the -# cache for these so they reach the user as visible bubbles instead of -# being silently swallowed. From PR #18153. -_SYSTEM_BYPASS_PREFIXES: Tuple[str, ...] = ( - "⚡ Interrupting", - "⏳ Queued", - "⏩ Steered", - "💾", # background-review summary -) +# Gateway busy-ack prefixes (interrupting / queued / steered / background review); +# these bypass a PENDING postback cache so they land as visible bubbles. +_SYSTEM_BYPASS_PREFIXES: Tuple[str, ...] = ("⚡ Interrupting", "⏳ Queued", "⏩ Steered", "💾") def _is_system_bypass(content: str) -> bool: @@ -668,9 +395,7 @@ def _is_system_bypass(content: str) -> bool: return any(content.startswith(p) for p in _SYSTEM_BYPASS_PREFIXES) -# --------------------------------------------------------------------------- -# Configuration helpers -# --------------------------------------------------------------------------- +# --- Configuration helpers --- def _csv_set(value: str) -> Set[str]: if not value: @@ -685,94 +410,76 @@ def _truthy_env(name: str, default: bool = False) -> bool: return v.strip().lower() in {"1", "true", "yes", "on"} -# --------------------------------------------------------------------------- -# Adapter -# --------------------------------------------------------------------------- +def _credentials(config) -> Tuple[str, str]: + """Return ``(channel_access_token, channel_secret)`` from scoped secrets, then ``extra``.""" + extra = getattr(config, "extra", {}) or {} + return ( + _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") or extra.get("channel_access_token", ""), + _get_scoped_secret("LINE_CHANNEL_SECRET") or extra.get("channel_secret", ""), + ) + + +def _coerce(cast: Callable[[Any], Any], value: Any, default: Any) -> Any: + try: + return cast(value) + except (TypeError, ValueError): + return default + + +# Outbound media kinds → (size cap, size error, missing-public-URL error). +_OUTBOUND_MEDIA = { + "image": ( + LINE_IMAGE_MAX_BYTES, "image exceeds 10 MB LINE limit", + "LINE_PUBLIC_URL must be set to send images (LINE only accepts publicly reachable HTTPS URLs)", + ), + "audio": (LINE_AV_MAX_BYTES, "audio exceeds 200 MB LINE limit", "LINE_PUBLIC_URL must be set to send audio"), + "video": (LINE_AV_MAX_BYTES, "video exceeds 200 MB LINE limit", "LINE_PUBLIC_URL must be set to send video"), +} + +# Inbound media kinds → cached file extension. +_INBOUND_MEDIA_EXT = {"image": ".jpg", "audio": ".m4a", "video": ".mp4", "file": ".bin"} + + +# --- Adapter --- class LineAdapter(BasePlatformAdapter): - """LINE Messaging API gateway adapter.""" - - # LINE has its own message-edit story (none) — we always send fresh - # bubbles, never edit, so REQUIRES_EDIT_FINALIZE stays False. + """LINE Messaging API gateway adapter (no message editing → REQUIRES_EDIT_FINALIZE stays False).""" def __init__(self, config, **kwargs): - platform = Platform("line") - super().__init__(config=config, platform=platform) + super().__init__(config=config, platform=Platform("line")) extra = getattr(config, "extra", {}) or {} - # Credentials - self.channel_access_token = ( - _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") - or extra.get("channel_access_token", "") - ) - self.channel_secret = ( - _get_scoped_secret("LINE_CHANNEL_SECRET") - or extra.get("channel_secret", "") - ) + def env_or(env: str, key: str, default: Any = "") -> Any: + return os.getenv(env) or extra.get(key, default) - # Webhook server. Host default is ``None`` → dual-stack bind (both - # IPv4 and IPv6); see DEFAULT_HOST above. ``LINE_HOST``/extra.host pin - # a specific address when needed; empty string collapses to None. - self.webhook_host = ( - os.getenv("LINE_HOST") or extra.get("host", DEFAULT_HOST) or DEFAULT_HOST - ) - try: - self.webhook_port = int( - os.getenv("LINE_PORT") or extra.get("port", DEFAULT_WEBHOOK_PORT) - ) - except (TypeError, ValueError): - self.webhook_port = DEFAULT_WEBHOOK_PORT + def allowlist(env: str, key: str) -> Set[str]: + return _csv_set(os.getenv(env, "")) | set(extra.get(key, [])) + + self.channel_access_token, self.channel_secret = _credentials(config) + + # Webhook server. Host default ``None`` → dual-stack bind (see DEFAULT_HOST); + # ``LINE_HOST``/extra.host pin an address, empty string collapses to None. + self.webhook_host = env_or("LINE_HOST", "host", DEFAULT_HOST) or DEFAULT_HOST + self.webhook_port = _coerce(int, env_or("LINE_PORT", "port", DEFAULT_WEBHOOK_PORT), DEFAULT_WEBHOOK_PORT) self.webhook_path = extra.get("webhook_path", DEFAULT_WEBHOOK_PATH) - # Public base URL — required for media sending when bind isn't - # publicly reachable. - self.public_base_url = ( - os.getenv("LINE_PUBLIC_URL") - or extra.get("public_url", "") - or "" - ).rstrip("/") + # Public base URL — required for media when the bind isn't publicly reachable. + self.public_base_url = (env_or("LINE_PUBLIC_URL", "public_url") or "").rstrip("/") # Three-allowlist gating - self.allow_all = _truthy_env( - "LINE_ALLOW_ALL_USERS", bool(extra.get("allow_all_users", False)) - ) - self.allowed_users = _csv_set( - os.getenv("LINE_ALLOWED_USERS", "") - ) | set(extra.get("allowed_users", [])) - self.allowed_groups = _csv_set( - os.getenv("LINE_ALLOWED_GROUPS", "") - ) | set(extra.get("allowed_groups", [])) - self.allowed_rooms = _csv_set( - os.getenv("LINE_ALLOWED_ROOMS", "") - ) | set(extra.get("allowed_rooms", [])) + self.allow_all = _truthy_env("LINE_ALLOW_ALL_USERS", bool(extra.get("allow_all_users", False))) + self.allowed_users = allowlist("LINE_ALLOWED_USERS", "allowed_users") + self.allowed_groups = allowlist("LINE_ALLOWED_GROUPS", "allowed_groups") + self.allowed_rooms = allowlist("LINE_ALLOWED_ROOMS", "allowed_rooms") - # Slow-LLM postback button threshold - try: - self.slow_response_threshold = float( - os.getenv("LINE_SLOW_RESPONSE_THRESHOLD") - or extra.get("slow_response_threshold", DEFAULT_SLOW_RESPONSE_THRESHOLD) - ) - except (TypeError, ValueError): - self.slow_response_threshold = DEFAULT_SLOW_RESPONSE_THRESHOLD - - # User-overridable copy - self.pending_text = ( - os.getenv("LINE_PENDING_TEXT") - or extra.get("pending_text", DEFAULT_PENDING_REPLY_TEXT) - ) - self.button_label = ( - os.getenv("LINE_BUTTON_LABEL") - or extra.get("button_label", DEFAULT_BUTTON_LABEL) - ) - self.delivered_text = ( - os.getenv("LINE_DELIVERED_TEXT") - or extra.get("delivered_text", DEFAULT_DELIVERED_TEXT) - ) - self.interrupted_text = ( - os.getenv("LINE_INTERRUPTED_TEXT") - or extra.get("interrupted_text", DEFAULT_INTERRUPTED_TEXT) - ) + # Slow-LLM postback button threshold + user-overridable copy + threshold = env_or("LINE_SLOW_RESPONSE_THRESHOLD", "slow_response_threshold", DEFAULT_SLOW_RESPONSE_THRESHOLD) + self.slow_response_threshold = _coerce(float, threshold, DEFAULT_SLOW_RESPONSE_THRESHOLD) + self.pending_text = env_or("LINE_PENDING_TEXT", "pending_text", DEFAULT_PENDING_REPLY_TEXT) + self.button_label = env_or("LINE_BUTTON_LABEL", "button_label", DEFAULT_BUTTON_LABEL) + self.delivered_text = env_or("LINE_DELIVERED_TEXT", "delivered_text", DEFAULT_DELIVERED_TEXT) + self.interrupted_text = env_or("LINE_INTERRUPTED_TEXT", "interrupted_text", DEFAULT_INTERRUPTED_TEXT) # Runtime state self._client: Optional[_LineClient] = None @@ -790,103 +497,68 @@ class LineAdapter(BasePlatformAdapter): self._media_temp_paths: Set[str] = set() self._media_ttl = MEDIA_TOKEN_TTL_SECONDS - # Pending-button slot per chat — ensures one outstanding postback - # button per chat at a time. Postback cache request_id keyed by chat_id. + # One outstanding postback button per chat: chat_id → cache request_id. self._pending_buttons: Dict[str, str] = {} - # ------------------------------------------------------------------ - # Connection lifecycle - # ------------------------------------------------------------------ + # --- Connection lifecycle --- + + def _fail(self, code: str, detail: str, *, retryable: bool = False) -> bool: + """Record a fatal connect error and return False.""" + self._set_fatal_error(code, detail, retryable=retryable) + return False async def connect(self, *, is_reconnect: bool = False) -> bool: if not self.channel_access_token or not self.channel_secret: - self._set_fatal_error( - "config_missing", - "LINE_CHANNEL_ACCESS_TOKEN and LINE_CHANNEL_SECRET must be set", - retryable=False, - ) - return False + return self._fail("config_missing", "LINE_CHANNEL_ACCESS_TOKEN and LINE_CHANNEL_SECRET must be set") - # Prevent two profiles from running on the same channel access token. + # One profile per channel token; lock on a hash so the secret never hits disk. try: from gateway.status import acquire_scoped_lock - # Use a hash of the token so we don't write the secret to disk. tok_hash = hashlib.sha256(self.channel_access_token.encode()).hexdigest()[:16] if not acquire_scoped_lock("line", tok_hash): - self._set_fatal_error( - "lock_conflict", - "LINE channel already in use by another profile", - retryable=False, - ) - return False + return self._fail("lock_conflict", "LINE channel already in use by another profile") self._lock_key = tok_hash except ImportError: self._lock_key = None self._client = _LineClient(self.channel_access_token) - # Best-effort: fetch our own bot userId for self-message filtering. - # If the call fails (offline tests, transient 5xx) we fall back to - # not filtering self-events; the cost is minor (LINE doesn't - # actually echo our own messages back). + # Best-effort self-userId for self-echo filtering (LINE rarely echoes anyway). try: self._bot_user_id = await self._client.get_bot_user_id() except Exception as exc: logger.debug("LINE: get_bot_user_id failed: %s", exc) self._bot_user_id = None - # Spin up the aiohttp webhook server. try: from aiohttp import web except ImportError: - self._set_fatal_error( - "missing_dep", - "aiohttp is required for the LINE adapter — install with `pip install aiohttp`", - retryable=False, - ) - return False + return self._fail("missing_dep", "aiohttp is required for the LINE adapter — install with `pip install aiohttp`") self._app = web.Application(client_max_size=WEBHOOK_BODY_MAX_BYTES) self._app.router.add_post(self.webhook_path, self._handle_webhook) # Public health probe — useful for tunnel/proxy verification. self._app.router.add_get(f"{self.webhook_path}/health", self._handle_health) - # Media serving endpoint. - self._app.router.add_get( - f"{DEFAULT_MEDIA_PATH_PREFIX}/{{token}}/{{filename}}", - self._handle_media, - ) - - # Plugin-registered native handlers (aiohttp web.Application — - # router routes). Wired before AppRunner.setup() freezes the router. + self._app.router.add_get(f"{DEFAULT_MEDIA_PATH_PREFIX}/{{token}}/{{filename}}", self._handle_media) + # Plugin-registered routes must be wired before AppRunner.setup() freezes the router. self._wire_plugin_handlers(self._app) - self._runner = web.AppRunner(self._app) try: await self._runner.setup() - # SO_REUSEADDR is platform-dependent (mirrors the generic webhook - # adapter, commits d542894ad/9420ad946): - # - macOS (BSD semantics): two wildcard/specific sockets with - # SO_REUSEADDR can silently split traffic — disable it there. - # - Linux: SO_REUSEADDR only permits rebinding past TIME_WAIT; - # disabling it would make a quick gateway restart fail to - # bind for up to ~60s — keep the default (enabled). + # SO_REUSEADDR: on macOS/BSD two sockets with it can silently split traffic → + # disable; on Linux it only allows rebinding past TIME_WAIT → keep default. self._site = web.TCPSite( - self._runner, - self.webhook_host, - self.webhook_port, + self._runner, self.webhook_host, self.webhook_port, reuse_address=False if sys.platform == "darwin" else None, ) await self._site.start() except OSError as exc: - self._set_fatal_error( + return self._fail( "bind_failed", - "Could not bind LINE webhook on " - f"{self.webhook_host or 'all IPv4+IPv6 interfaces'}:" + f"Could not bind LINE webhook on {self.webhook_host or 'all IPv4+IPv6 interfaces'}:" f"{self.webhook_port}: {exc}", retryable=True, ) - return False - self._mark_connected() logger.info( "LINE: webhook listening on %s:%s%s%s", @@ -899,30 +571,19 @@ class LineAdapter(BasePlatformAdapter): async def disconnect(self) -> None: self._mark_disconnected() - - if self._site is not None: - try: - await self._site.stop() - except Exception: - pass - self._site = None - if self._runner is not None: - try: - await self._runner.cleanup() - except Exception: - pass - self._runner = None + for attr, method in (("_site", "stop"), ("_runner", "cleanup")): + obj = getattr(self, attr) + if obj is not None: + try: + await getattr(obj, method)() + except Exception: + pass + setattr(self, attr, None) self._app = None - - # Cleanup any tracked tempfiles. for path in list(self._media_temp_paths): - try: - os.unlink(path) - except OSError: - pass + _unlink_quietly(path) self._media_temp_paths.clear() self._media_tokens.clear() - if self._lock_key: try: from gateway.status import release_scoped_lock @@ -931,9 +592,7 @@ class LineAdapter(BasePlatformAdapter): pass self._lock_key = None - # ------------------------------------------------------------------ - # Webhook handlers - # ------------------------------------------------------------------ + # --- Webhook handlers --- async def _handle_health(self, request) -> Any: from aiohttp import web @@ -942,8 +601,7 @@ class LineAdapter(BasePlatformAdapter): async def _handle_webhook(self, request) -> Any: from aiohttp import web - # Body cap defends against memory-exhaustion via crafted Content-Length - # (aiohttp's client_max_size only applies to certain body modes). + # Explicit body cap: aiohttp's client_max_size only covers some body modes. try: body = await request.read() except Exception as exc: @@ -951,23 +609,18 @@ class LineAdapter(BasePlatformAdapter): return web.Response(status=400, text="bad request") if len(body) > WEBHOOK_BODY_MAX_BYTES: return web.Response(status=413, text="payload too large") - signature = request.headers.get("X-Line-Signature", "") if not verify_line_signature(body, signature, self.channel_secret): return web.Response(status=401, text="invalid signature") - try: payload = json.loads(body.decode("utf-8")) except (UnicodeDecodeError, json.JSONDecodeError): return web.Response(status=400, text="bad json") - - events = payload.get("events", []) or [] - for event in events: + for event in payload.get("events", []) or []: try: await self._dispatch_event(event) except Exception: logger.exception("LINE: dispatch_event failed") - return web.Response(status=200, text="ok") async def _dispatch_event(self, event: Dict[str, Any]) -> None: @@ -975,27 +628,20 @@ class LineAdapter(BasePlatformAdapter): source = event.get("source") or {} webhook_event_id = event.get("webhookEventId", "") or "" - # Dedup retries (LINE webhooks may be re-delivered). + # LINE webhooks may be re-delivered. if webhook_event_id and self._dedup.is_duplicate(webhook_event_id): logger.debug("LINE: ignoring duplicate webhook event %s", webhook_event_id) return - # Filter our own messages (self-echo). - sender_user_id = source.get("userId", "") - if self._bot_user_id and sender_user_id == self._bot_user_id: + # Self-echo filter. + if self._bot_user_id and source.get("userId", "") == self._bot_user_id: return - - # Allowlist gate. if not _allowed_for_source( - source, - allow_all=self.allow_all, - user_ids=self.allowed_users, - group_ids=self.allowed_groups, - room_ids=self.allowed_rooms, + source, allow_all=self.allow_all, user_ids=self.allowed_users, + group_ids=self.allowed_groups, room_ids=self.allowed_rooms, ): logger.info("LINE: rejecting unauthorized source %s", source) return - if event_type == "message": await self._handle_message_event(event) elif event_type == "postback": @@ -1013,27 +659,16 @@ class LineAdapter(BasePlatformAdapter): source = event.get("source") or {} chat_id, chat_type = _resolve_chat(source) user_id = source.get("userId", "") or chat_id - - # Stash the reply token for outbound use. - if chat_id and reply_token: - self._reply_tokens[chat_id] = ( - reply_token, - time.time() + LINE_REPLY_TOKEN_TTL_SECONDS, - ) - - # Handle media inbound — fetch the binary, cache it, and surface a - # vision-tool-friendly local path on the MessageEvent. + if chat_id and reply_token: # stash the reply token for outbound use + self._reply_tokens[chat_id] = (reply_token, time.time() + LINE_REPLY_TOKEN_TTL_SECONDS) media_urls: List[str] = [] media_types: List[str] = [] - text = "" - if msg_type == "text": text = msg.get("text", "") or "" - elif msg_type in ("image", "audio", "video", "file"): + elif msg_type in _INBOUND_MEDIA_EXT: + # Fetch the binary, cache it, surface a vision-friendly local path. local_path, media_type = await self._download_media( - message_id, - msg_type, - filename=msg.get("fileName") or msg.get("file_name"), + message_id, msg_type, filename=msg.get("fileName") or msg.get("file_name") ) if local_path: media_urls.append(local_path) @@ -1043,101 +678,72 @@ class LineAdapter(BasePlatformAdapter): keywords = msg.get("keywords") or [] text = f"[sticker: {', '.join(keywords)}]" if keywords else "[sticker]" elif msg_type == "location": - title = msg.get("title", "") - address = msg.get("address", "") - text = f"[location: {title} {address}]".strip() + text = f"[location: {msg.get('title', '')} {msg.get('address', '')}]".strip() else: text = f"[unsupported message type: {msg_type}]" # Best-effort typing indicator (DM only). if chat_type == "dm" and self._client: asyncio.create_task(self._client.loading(chat_id)) - source_obj = self.build_source( - chat_id=chat_id, - chat_type=chat_type, - user_id=user_id, - user_name=user_id, - chat_name=chat_id, + chat_id=chat_id, chat_type=chat_type, user_id=user_id, user_name=user_id, chat_name=chat_id ) - - event_obj = MessageEvent( - text=text, - message_type=_LINE_MESSAGE_TYPES.get(msg_type, MessageType.TEXT), - source=source_obj, - raw_message=event, - message_id=message_id, - media_urls=media_urls, - media_types=media_types, - ) - - await self.handle_message(event_obj) + await self.handle_message(MessageEvent( + text=text, message_type=_LINE_MESSAGE_TYPES.get(msg_type, MessageType.TEXT), source=source_obj, + raw_message=event, message_id=message_id, media_urls=media_urls, media_types=media_types, + )) async def _handle_postback_event(self, event: Dict[str, Any]) -> None: """User tapped the slow-LLM postback button — deliver cached payload.""" postback = event.get("postback") or {} data = postback.get("data", "") or "" reply_token = event.get("replyToken", "") - source = event.get("source") or {} - chat_id, _ = _resolve_chat(source) - + chat_id, _ = _resolve_chat(event.get("source") or {}) try: parsed = json.loads(data) except (TypeError, json.JSONDecodeError): return - if parsed.get("action") != "show_response": return request_id = parsed.get("request_id", "") if not request_id: return - entry = self._cache.get(request_id) if not self._client or not reply_token or not entry: return + def _settle() -> None: + self._cache.mark_delivered(request_id) + self._pending_buttons.pop(chat_id, None) if entry.state is State.READY: - payload = entry.payload or "" - chunks = split_for_line(strip_markdown_preserving_urls(str(payload))) - messages = [_text_message(c) for c in chunks][:LINE_MAX_MESSAGES_PER_CALL] + messages = _text_messages(str(entry.payload or "")) try: await self._client.reply(reply_token, messages) - self._cache.mark_delivered(request_id) - self._pending_buttons.pop(chat_id, None) + _settle() except Exception as exc: logger.warning("LINE: postback reply failed (%s); falling back to push", exc) try: await self._client.push(chat_id, messages) - self._cache.mark_delivered(request_id) - self._pending_buttons.pop(chat_id, None) + _settle() except Exception as exc2: logger.error("LINE: postback push fallback failed: %s", exc2) elif entry.state is State.ERROR: text = str(entry.payload or self.interrupted_text) try: await self._client.reply(reply_token, [_text_message(text)]) - self._cache.mark_delivered(request_id) - self._pending_buttons.pop(chat_id, None) + _settle() except Exception as exc: logger.warning("LINE: postback ERROR reply failed: %s", exc) - elif entry.state is State.DELIVERED: + elif entry.state in (State.DELIVERED, State.PENDING): + # DELIVERED → "already replied"; PENDING → re-issue the wait notice. + text = self.delivered_text if entry.state is State.DELIVERED else self.pending_text try: - await self._client.reply(reply_token, [_text_message(self.delivered_text)]) - except Exception: - pass - elif entry.state is State.PENDING: - # Still working — re-issue the wait notice. - try: - await self._client.reply(reply_token, [_text_message(self.pending_text)]) + await self._client.reply(reply_token, [_text_message(text)]) except Exception: pass async def _download_media( - self, - message_id: str, - msg_type: str, - *, - filename: Optional[str] = None, + self, message_id: str, msg_type: str, *, filename: Optional[str] = None ) -> Tuple[Optional[str], str]: if not self._client or not message_id: return None, "" @@ -1146,95 +752,40 @@ class LineAdapter(BasePlatformAdapter): except Exception as exc: logger.warning("LINE: failed to fetch %s content for %s: %s", msg_type, message_id, exc) return None, "" - ext = { - "image": ".jpg", - "audio": ".m4a", - "video": ".mp4", - "file": ".bin", - }.get(msg_type, ".bin") + ext = _INBOUND_MEDIA_EXT.get(msg_type, ".bin") try: if msg_type == "image": return cache_image_from_bytes(data, ext=ext), "image/jpeg" - if msg_type == "audio": - media_type = mimetypes.guess_type(f"audio{ext}")[0] or "audio/mp4" - return cache_audio_from_bytes(data, ext=ext), media_type - if msg_type == "video": - media_type = mimetypes.guess_type(f"video{ext}")[0] or "video/mp4" - return cache_video_from_bytes(data, ext=ext), media_type + if msg_type in ("audio", "video"): + cache_fn = cache_audio_from_bytes if msg_type == "audio" else cache_video_from_bytes + return cache_fn(data, ext=ext), mimetypes.guess_type(f"{msg_type}{ext}")[0] or f"{msg_type}/mp4" document_name = filename or f"line_file{ext}" - return ( - cache_document_from_bytes(data, document_name), - mimetypes.guess_type(document_name)[0] or "application/octet-stream", - ) + mime = mimetypes.guess_type(document_name)[0] or "application/octet-stream" + return cache_document_from_bytes(data, document_name), mime except Exception as exc: logger.warning("LINE: failed to cache %s payload: %s", msg_type, exc) return None, "" - # ------------------------------------------------------------------ - # Outbound send (text) - # ------------------------------------------------------------------ + # --- Outbound send (text) --- async def send( - self, - chat_id: str, - content: str, - reply_to: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, + self, chat_id: str, content: str, reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None ) -> SendResult: if not self._client: return SendResult(success=False, error="LINE adapter not connected") - - # System busy-acks (interrupting / queued / steered) bypass the - # postback cache and route directly to LINE so they reach the user - # as visible bubbles. Source: PR #18153. - if _is_system_bypass(content): - return await self._send_text_chunks(chat_id, content, force_push=False) - - # If the chat has a PENDING postback button outstanding, route the - # response into the cache for the user to fetch via tap. + # With a PENDING postback button outstanding, cache the response for the tap — + # except system busy-acks, which must land as visible bubbles. pending_rid = self._pending_buttons.get(chat_id) - if pending_rid: + if pending_rid and not _is_system_bypass(content): self._cache.set_ready(pending_rid, content) return SendResult(success=True, message_id=pending_rid) - return await self._send_text_chunks(chat_id, content, force_push=False) - async def _send_text_chunks( - self, - chat_id: str, - content: str, - *, - force_push: bool, - ) -> SendResult: - if not self._client: - return SendResult(success=False, error="LINE adapter not connected") - - chunks = split_for_line(strip_markdown_preserving_urls(content)) - if not chunks: - return SendResult(success=True, message_id=None) - messages = [_text_message(c) for c in chunks][:LINE_MAX_MESSAGES_PER_CALL] - - token, used_reply = self._consume_reply_token(chat_id) - if used_reply and not force_push: - try: - await self._client.reply(token, messages) - return SendResult(success=True, message_id=token) - except Exception as exc: - logger.info("LINE: reply token rejected (%s); falling back to push", exc) - # fall through to push - - try: - await self._client.push(chat_id, messages) - return SendResult(success=True, message_id=None) - except Exception as exc: - logger.error("LINE: push send failed: %s", exc) - return SendResult(success=False, error=str(exc)) + async def _send_text_chunks(self, chat_id: str, content: str, *, force_push: bool) -> SendResult: + return await self._send_messages(chat_id, _text_messages(content), force_push=force_push, text=True) def _consume_reply_token(self, chat_id: str) -> Tuple[str, bool]: - """Consume a stashed reply token if present and unexpired. - - Returns ``(token, used_reply)``. - """ + """Consume a stashed reply token if present and unexpired → ``(token, used_reply)``.""" entry = self._reply_tokens.pop(chat_id, None) if not entry: return "", False @@ -1249,37 +800,19 @@ class LineAdapter(BasePlatformAdapter): await self._client.loading(chat_id) async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: - """Best-effort chat info derived from the chat_id prefix. - - LINE's chat-info APIs are limited and per-source-type — instead of - chasing them we infer from the well-known ID prefixes: - ``U`` = user (1:1), ``C`` = group, ``R`` = room. The agent only - needs ``name`` + ``type`` from this method. - """ - prefix = (chat_id or "")[:1] - chat_type = {"U": "dm", "C": "group", "R": "channel"}.get(prefix, "dm") + """Best-effort chat info inferred from the ID prefix (U=user, C=group, R=room).""" + chat_type = {"U": "dm", "C": "group", "R": "channel"}.get((chat_id or "")[:1], "dm") return {"name": chat_id or "", "type": chat_type} def format_message(self, content: str) -> str: """Strip Markdown that LINE can't render. URLs are preserved.""" return strip_markdown_preserving_urls(content) - # ------------------------------------------------------------------ - # Slow-LLM postback button — driven by _keep_typing - # ------------------------------------------------------------------ + # --- Slow-LLM postback button — driven by _keep_typing --- async def _keep_typing(self, chat_id: str, *args, **kwargs) -> None: - """Override the base loop to fire the postback button at threshold. - - We intentionally keep the base implementation behind us: it's - responsible for the typing-indicator heartbeat, while *this* - wrapper layers in the slow-LLM postback bubble at threshold. - """ - if ( - self.slow_response_threshold <= 0 - or not self._client - or not chat_id - ): + """Wrap the base typing heartbeat; fire the postback button at threshold.""" + if self.slow_response_threshold <= 0 or not self._client or not chat_id: await super()._keep_typing(chat_id, *args, **kwargs) return @@ -1288,11 +821,9 @@ class LineAdapter(BasePlatformAdapter): await asyncio.sleep(self.slow_response_threshold) except asyncio.CancelledError: raise - # Only fire if we still have a usable reply token. If the agent - # already responded, _consume_reply_token has cleared it. - if chat_id not in self._reply_tokens: - return - if chat_id in self._pending_buttons: + # Only fire while a usable reply token remains (the agent responding + # consumes it) and no button is already outstanding. + if chat_id not in self._reply_tokens or chat_id in self._pending_buttons: return rid = self._cache.register_pending(chat_id) self._pending_buttons[chat_id] = rid @@ -1300,16 +831,13 @@ class LineAdapter(BasePlatformAdapter): if not used: self._pending_buttons.pop(chat_id, None) return - msg = build_postback_button_message( - self.pending_text, self.button_label, rid - ) + msg = build_postback_button_message(self.pending_text, self.button_label, rid) try: await self._client.reply(token, [msg]) logger.info("LINE: sent slow-LLM postback button for chat %s (rid=%s)", chat_id, rid) except Exception as exc: logger.warning("LINE: postback button send failed: %s", exc) self._pending_buttons.pop(chat_id, None) - post_task = asyncio.create_task(_fire_postback()) try: await super()._keep_typing(chat_id, *args, **kwargs) @@ -1328,13 +856,10 @@ class LineAdapter(BasePlatformAdapter): if rid: self._cache.set_error(rid, self.interrupted_text) - # ------------------------------------------------------------------ - # Outbound media (image / voice / video) - # ------------------------------------------------------------------ + # --- Outbound media (image / voice / video) --- def _register_media(self, file_path: str, *, cleanup: bool = False) -> str: """Register a local file for HTTPS serving; return the URL token.""" - # Evict expired tokens first. now = time.time() for token in list(self._media_tokens.keys()): path, exp = self._media_tokens[token] @@ -1342,11 +867,7 @@ class LineAdapter(BasePlatformAdapter): self._media_tokens.pop(token, None) if path in self._media_temp_paths: self._media_temp_paths.discard(path) - try: - os.unlink(path) - except OSError: - pass - + _unlink_quietly(path) resolved = str(Path(file_path).resolve()) token = secrets.token_urlsafe(32) self._media_tokens[token] = (resolved, now + self._media_ttl) @@ -1355,158 +876,105 @@ class LineAdapter(BasePlatformAdapter): return token def _media_url(self, token: str, filename: str) -> str: - """Build the public HTTPS URL for a media token. PR #8398 style.""" + """Build the public HTTPS URL for a media token.""" if self.public_base_url: base = self.public_base_url else: - # A wildcard/dual-stack bind has no fetchable hostname; the - # _missing_public_url guard should have caught this earlier. - # Fall back to localhost so the URL is at least well-formed. - host = self.webhook_host - if host is None or host in _WILDCARD_HOSTS: - host = "127.0.0.1" - port = self.webhook_port - if port == 443: - base = f"https://{host}" - else: - base = f"https://{host}:{port}" + # Wildcard/dual-stack binds have no fetchable hostname (the _missing_public_url + # guard should have fired); fall back to localhost so the URL is well-formed. + host = "127.0.0.1" if self._missing_public_url() else self.webhook_host + base = f"https://{host}" if self.webhook_port == 443 else f"https://{host}:{self.webhook_port}" safe_name = _urlquote(filename, safe="") return f"{base}{DEFAULT_MEDIA_PATH_PREFIX}/{token}/{safe_name}" + def _serve_file(self, path: Path) -> str: + """Register ``path`` for serving and return its public URL.""" + return self._media_url(self._register_media(str(path.resolve())), path.name) + def _missing_public_url(self) -> bool: - """True when outbound media cannot work: no LINE_PUBLIC_URL and the - bind host is a wildcard (or the dual-stack ``None`` default), i.e. - not an address LINE's fetchers could ever reach.""" + """True when no LINE_PUBLIC_URL is set and the bind host is wildcard/dual-stack ``None``.""" if self.public_base_url: return False return self.webhook_host is None or self.webhook_host in _WILDCARD_HOSTS - async def _handle_media(self, request) -> Any: - """Serve a registered local file over HTTPS for LINE's media URLs. + def _check_media_file(self, kind: str, file_path: str) -> Tuple[Optional[Path], Optional[SendResult]]: + """Shared preflight for send_image_file/send_voice/send_video → ``(path, error)``.""" + max_bytes, size_error, url_error = _OUTBOUND_MEDIA[kind] + path = Path(file_path) + if not path.exists() or not path.is_file(): + return None, SendResult(success=False, error=f"{kind} file not found: {file_path}") + if path.stat().st_size > max_bytes: + return None, SendResult(success=False, error=size_error) + if not self._client: + return None, SendResult(success=False, error="LINE adapter not connected") + if self._missing_public_url(): + return None, SendResult(success=False, error=url_error) + return path, None - Defence-in-depth: even though ``_register_media`` is only called - from trusted internal code, we recheck the resolved path against - an allowed-roots set before serving. Sources allowed: - ``tempfile.gettempdir()``, ``/tmp`` (which resolves to - ``/private/tmp`` on macOS), and ``HERMES_HOME``. PR #8398. + async def _handle_media(self, request) -> Any: + """Serve a registered local file for LINE's media URLs. + Defence-in-depth: the resolved path is rechecked against allowed roots + (tempdir, ``/tmp`` → ``/private/tmp`` on macOS, ``HERMES_HOME``). """ from aiohttp import web - token = request.match_info["token"] entry = self._media_tokens.get(token) if not entry: return web.Response(status=404, text="not found") - file_path, expires_at = entry if time.time() > expires_at: self._media_tokens.pop(token, None) return web.Response(status=410, text="gone") - path = Path(file_path) if not path.exists() or not path.is_file(): return web.Response(status=404, text="not found") - try: from hermes_constants import get_hermes_home hermes_home = Path(get_hermes_home()).resolve() except Exception: hermes_home = Path.home().joinpath(".hermes").resolve() - - allowed_roots = { - Path(tempfile.gettempdir()).resolve(), - Path("/tmp").resolve(), # → /private/tmp on macOS - hermes_home, - } + allowed_roots = {Path(tempfile.gettempdir()).resolve(), Path("/tmp").resolve(), hermes_home} resolved = path.resolve() - if not any(_is_relative_to(resolved, r) for r in allowed_roots): + if not any(resolved.is_relative_to(r) for r in allowed_roots): logger.warning("LINE: refusing to serve outside allowed roots: %s", resolved) return web.Response(status=403, text="forbidden") - content_type, _ = mimetypes.guess_type(str(path)) - return web.FileResponse( - path, - headers={"Content-Type": content_type or "application/octet-stream"}, - ) + return web.FileResponse(path, headers={"Content-Type": content_type or "application/octet-stream"}) async def send_image_file( - self, - chat_id: str, - image_path: str, - caption: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None, + self, chat_id: str, image_path: str, caption: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None ) -> SendResult: - path = Path(image_path) - if not path.exists() or not path.is_file(): - return SendResult(success=False, error=f"image file not found: {image_path}") - if path.stat().st_size > LINE_IMAGE_MAX_BYTES: - return SendResult(success=False, error="image exceeds 10 MB LINE limit") - if not self._client: - return SendResult(success=False, error="LINE adapter not connected") - if self._missing_public_url(): - return SendResult( - success=False, - error="LINE_PUBLIC_URL must be set to send images " - "(LINE only accepts publicly reachable HTTPS URLs)", - ) - - token = self._register_media(str(path.resolve())) - url = self._media_url(token, path.name) + path, err = self._check_media_file("image", image_path) + if err: + return err + url = self._serve_file(path) if not url.lower().startswith("https://"): return SendResult(success=False, error=f"LINE image URL must be HTTPS: {url}") - msgs: List[Dict[str, Any]] = [_image_message(url)] + msgs: List[Dict[str, Any]] = [{"type": "image", "originalContentUrl": url, "previewImageUrl": url}] if caption: msgs.append(_text_message(caption)) return await self._send_messages(chat_id, msgs) async def send_voice( - self, - chat_id: str, - audio_path: str, - duration_ms: int = 1000, - metadata: Optional[Dict[str, Any]] = None, + self, chat_id: str, audio_path: str, duration_ms: int = 1000, metadata: Optional[Dict[str, Any]] = None ) -> SendResult: - path = Path(audio_path) - if not path.exists() or not path.is_file(): - return SendResult(success=False, error=f"audio file not found: {audio_path}") - if path.stat().st_size > LINE_AV_MAX_BYTES: - return SendResult(success=False, error="audio exceeds 200 MB LINE limit") - if not self._client: - return SendResult(success=False, error="LINE adapter not connected") - if self._missing_public_url(): - return SendResult( - success=False, - error="LINE_PUBLIC_URL must be set to send audio", - ) - - token = self._register_media(str(path.resolve())) - url = self._media_url(token, path.name) - return await self._send_messages(chat_id, [_audio_message(url, duration_ms)]) + path, err = self._check_media_file("audio", audio_path) + if err: + return err + url = self._serve_file(path) + return await self._send_messages(chat_id, [{"type": "audio", "originalContentUrl": url, "duration": int(duration_ms)}]) async def send_video( - self, - chat_id: str, - video_path: str, - preview_path: Optional[str] = None, + self, chat_id: str, video_path: str, preview_path: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: - path = Path(video_path) - if not path.exists() or not path.is_file(): - return SendResult(success=False, error=f"video file not found: {video_path}") - if path.stat().st_size > LINE_AV_MAX_BYTES: - return SendResult(success=False, error="video exceeds 200 MB LINE limit") - if not self._client: - return SendResult(success=False, error="LINE adapter not connected") - if self._missing_public_url(): - return SendResult( - success=False, - error="LINE_PUBLIC_URL must be set to send video", - ) + path, err = self._check_media_file("video", video_path) + if err: + return err - # LINE requires a previewImageUrl. Use one if supplied, otherwise - # write a stdlib 1×1 PNG to /tmp and serve it. PR #8398. + # LINE requires previewImageUrl: use the supplied preview, else a stdlib 1×1 PNG. if preview_path and Path(preview_path).is_file(): - preview_token = self._register_media(str(Path(preview_path).resolve())) - preview_filename = Path(preview_path).name + preview_url = self._serve_file(Path(preview_path)) else: tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False) try: @@ -1514,48 +982,43 @@ class LineAdapter(BasePlatformAdapter): tmp.flush() tmp.close() preview_token = self._register_media(tmp.name, cleanup=True) - preview_filename = "preview.png" + preview_url = self._media_url(preview_token, "preview.png") except Exception: - try: - os.unlink(tmp.name) - except OSError: - pass + _unlink_quietly(tmp.name) raise - - video_token = self._register_media(str(path.resolve())) - video_url = self._media_url(video_token, path.name) - preview_url = self._media_url(preview_token, preview_filename) - return await self._send_messages(chat_id, [_video_message(video_url, preview_url)]) + video_url = self._serve_file(path) + return await self._send_messages( + chat_id, [{"type": "video", "originalContentUrl": video_url, "previewImageUrl": preview_url}] + ) async def _send_messages( - self, - chat_id: str, - messages: List[Dict[str, Any]], + self, chat_id: str, messages: List[Dict[str, Any]], *, force_push: bool = False, text: bool = False ) -> SendResult: - """Send already-built message objects, batched at 5/call.""" + """Send built message objects, batched at 5/call: reply token first, then push. + ``text`` selects the text-send contract: reply success reports the token + as message_id and a push failure is logged at error level. + """ if not self._client: return SendResult(success=False, error="LINE adapter not connected") if not messages: return SendResult(success=True, message_id=None) - first_batch = messages[:LINE_MAX_MESSAGES_PER_CALL] rest = messages[LINE_MAX_MESSAGES_PER_CALL:] - - # First batch: try reply token, fall back to push. token, used_reply = self._consume_reply_token(chat_id) - if used_reply: + if used_reply and not force_push: try: await self._client.reply(token, first_batch) + first_batch = None + if text: + return SendResult(success=True, message_id=token) except Exception as exc: logger.info("LINE: reply token rejected (%s); falling back to push", exc) - try: - await self._client.push(chat_id, first_batch) - except Exception as exc2: - return SendResult(success=False, error=str(exc2)) - else: + if first_batch is not None: try: await self._client.push(chat_id, first_batch) except Exception as exc: + if text: + logger.error("LINE: push send failed: %s", exc) return SendResult(success=False, error=str(exc)) # Subsequent batches: always push (reply token is single-use). @@ -1567,49 +1030,32 @@ class LineAdapter(BasePlatformAdapter): except Exception as exc: logger.warning("LINE: push for follow-up batch failed: %s", exc) return SendResult(success=False, error=str(exc)) - return SendResult(success=True, message_id=None) -def _is_relative_to(child: Path, parent: Path) -> bool: - """Backport for Path.is_relative_to (Python 3.9+) — defensive against - cwd-resolution differences across CI runners.""" +def _unlink_quietly(path: str) -> None: try: - return child.resolve().is_relative_to(parent.resolve()) - except (AttributeError, ValueError): - try: - child.resolve().relative_to(parent.resolve()) - return True - except ValueError: - return False + os.unlink(path) + except OSError: + pass -# --------------------------------------------------------------------------- -# Plugin entry-point hooks -# --------------------------------------------------------------------------- +# --- Plugin entry-point hooks --- def check_requirements() -> bool: """Plugin gate: require credentials AND aiohttp at runtime.""" - if not _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN"): - return False - if not _get_scoped_secret("LINE_CHANNEL_SECRET"): + if not (_get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") and _get_scoped_secret("LINE_CHANNEL_SECRET")): return False try: import aiohttp # noqa: F401 + return True except ImportError: return False - return True def validate_config(config) -> bool: - extra = getattr(config, "extra", {}) or {} - has_token = bool( - _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") or extra.get("channel_access_token") - ) - has_secret = bool( - _get_scoped_secret("LINE_CHANNEL_SECRET") or extra.get("channel_secret") - ) - return has_token and has_secret + token, secret = _credentials(config) + return bool(token) and bool(secret) def is_connected(config) -> bool: @@ -1618,12 +1064,7 @@ def is_connected(config) -> bool: def _env_enablement() -> Optional[Dict[str, Any]]: - """Auto-seed PlatformConfig.extra from env-only setups. - - Lets ``hermes status`` reflect a LINE configuration that lives entirely - in ``.env`` without a ``platforms.line`` block in ``config.yaml``. - Mirrors the IRC plugin's pattern. - """ + """Seed PlatformConfig.extra from env-only setups so ``hermes status`` sees them.""" if not (_get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") and _get_scoped_secret("LINE_CHANNEL_SECRET")): return None seeded: Dict[str, Any] = {} @@ -1632,73 +1073,42 @@ def _env_enablement() -> Optional[Dict[str, Any]]: seeded["port"] = int(os.environ["LINE_PORT"]) except ValueError: pass - if os.getenv("LINE_HOST"): - seeded["host"] = os.environ["LINE_HOST"] - if os.getenv("LINE_PUBLIC_URL"): - seeded["public_url"] = os.environ["LINE_PUBLIC_URL"] - if os.getenv("LINE_HOME_CHANNEL"): - seeded["home_channel"] = os.environ["LINE_HOME_CHANNEL"] + for env, key in (("LINE_HOST", "host"), ("LINE_PUBLIC_URL", "public_url"), ("LINE_HOME_CHANNEL", "home_channel")): + if os.getenv(env): + seeded[key] = os.environ[env] return seeded or {} async def _standalone_send( - pconfig, - chat_id: str, - message: str, - *, - thread_id: Optional[str] = None, - media_files: Optional[List[str]] = None, - force_document: bool = False, + pconfig, chat_id: str, message: str, *, + thread_id: Optional[str] = None, media_files: Optional[List[str]] = None, force_document: bool = False, ) -> Dict[str, Any]: - """Out-of-process push delivery for cron jobs running detached from the gateway. - - Without this hook ``deliver=line`` cron jobs fail with ``no live adapter`` - when cron runs as its own process. We always Push (reply tokens require - an inbound webhook event we don't have in this path). - - ``thread_id`` is accepted for signature parity but ignored — LINE has - no native thread primitive on the channel-side API. ``media_files`` - likewise: cron-side media delivery requires a publicly-reachable URL, - which the standalone path can't construct without binding the webhook - server, so we send a text reference instead. + """Out-of-process Push delivery for cron jobs detached from the gateway. + Always Push (no inbound event → no reply token). ``thread_id`` is ignored (LINE + has no threads); ``media_files`` can't be served without the webhook server. """ extra = getattr(pconfig, "extra", {}) or {} - token = ( - _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") - or extra.get("channel_access_token", "") - ) + token = _get_scoped_secret("LINE_CHANNEL_ACCESS_TOKEN") or extra.get("channel_access_token", "") if not token or not chat_id: return {"error": "LINE standalone send: missing token or chat_id"} - - plain = strip_markdown_preserving_urls(message or "") - chunks = split_for_line(plain) or [""] - messages = [_text_message(c) for c in chunks][:LINE_MAX_MESSAGES_PER_CALL] + messages = _text_messages(message or "") or [_text_message("")] if media_files: - # Tack on a hint so the recipient knows media was generated but not delivered. + # Tell the recipient media was generated but not delivered. messages.append(_text_message(f"[{len(media_files)} attachment(s) generated; not deliverable from cron]")) messages = messages[:LINE_MAX_MESSAGES_PER_CALL] - - client = _LineClient(token) try: - await client.push(chat_id, messages) + await _LineClient(token).push(chat_id, messages) return {"success": True, "message_id": None} except Exception as exc: return {"error": str(exc)} def interactive_setup() -> None: - """Minimal stdin wizard for ``hermes setup line``. - - Mirrors the irc/teams style: prompts for the two required vars, plus - one optional public URL. Writes to ``~/.hermes/.env`` via ``hermes_cli.config``. - """ - print() - print("LINE Messaging API setup") - print("------------------------") - print("Create a Messaging API channel at https://developers.line.biz/console/") - print("then copy the values below.") - print() - + """Minimal stdin wizard for ``hermes setup line`` (writes ``~/.hermes/.env``).""" + print( + "\nLINE Messaging API setup\n------------------------\n" + "Create a Messaging API channel at https://developers.line.biz/console/\nthen copy the values below.\n" + ) try: from hermes_cli.config import get_env_value as _get_env, save_env_value as _set_env except ImportError: @@ -1719,7 +1129,6 @@ def interactive_setup() -> None: return if value: _set_env(var, value) - _prompt("LINE_CHANNEL_ACCESS_TOKEN", "Channel access token", secret=True) _prompt("LINE_CHANNEL_SECRET", "Channel secret", secret=True) _prompt("LINE_PUBLIC_URL", "Public HTTPS base URL (optional, e.g. https://my-tunnel.example.com)") diff --git a/plugins/platforms/ntfy/adapter.py b/plugins/platforms/ntfy/adapter.py index 87986416ef..ff8f5fab2c 100644 --- a/plugins/platforms/ntfy/adapter.py +++ b/plugins/platforms/ntfy/adapter.py @@ -1,15 +1,8 @@ """ntfy platform adapter (Hermes plugin). -Subscribes to a topic on ntfy.sh or any self-hosted ntfy server via -HTTP streaming (``/json`` endpoint with ``poll=false``) and publishes -replies via HTTP POST. No external SDK — only httpx, which is already -a Hermes dependency. - -This adapter ships as a Hermes platform plugin under -``plugins/platforms/ntfy/``. The Hermes plugin loader scans the -directory at startup, calls :func:`register`, and the platform becomes -available to ``gateway/run.py`` and ``tools/send_message_tool`` through -the registry — no edits to core files required. +Subscribes to a topic on ntfy.sh or a self-hosted ntfy server via HTTP +streaming (``/json`` with ``poll=false``) and publishes replies via HTTP POST. +No external SDK — only httpx. Configuration in config.yaml:: @@ -23,31 +16,27 @@ Configuration in config.yaml:: token: "..." # optional Bearer / Basic auth token markdown: true # optional — enable markdown (default: false) -Environment variables (all read at adapter construct time, env wins over -config.yaml ``extra``): +Environment variables (read at adapter construct time; env wins over ``extra``): NTFY_TOPIC Topic to subscribe to (required) NTFY_SERVER_URL Server URL (default: https://ntfy.sh) NTFY_TOKEN Bearer token or 'user:pass' for Basic auth NTFY_PUBLISH_TOPIC Reply topic (defaults to NTFY_TOPIC) NTFY_MARKDOWN "true"/"1"/"yes" enables X-Markdown header - NTFY_ALLOWED_USERS Allowlist (treated by gateway as user IDs; - on ntfy these are topic names) + NTFY_ALLOWED_USERS Allowlist (on ntfy these are topic names) NTFY_ALLOW_ALL_USERS Allow any topic — dev only NTFY_HOME_CHANNEL Default topic for cron / notification delivery NTFY_HOME_CHANNEL_NAME Human label for the home channel -Identity model: ntfy has no native authenticated user identity. The -``title`` field is publisher-controlled and is NOT used for -authorization. Each topic is treated as a single trusted channel — -``user_id`` is fixed to the topic name. Use a private topic protected -by a read token for any real trust boundary. +Identity model: ntfy has no authenticated user identity; ``title`` is +publisher-controlled and NOT used for authorization. Each topic is a single +trusted channel (``user_id`` == topic name). Protect the topic with a read +token for any real trust boundary. """ import asyncio import json import logging -import os import time import uuid from datetime import datetime, timezone @@ -68,28 +57,7 @@ from gateway.platforms.base import ( SendResult, ) -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 +from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret logger = logging.getLogger(__name__) @@ -106,19 +74,14 @@ DEDUP_MAX_SIZE = 1000 RECONNECT_BACKOFF = [2, 5, 10, 30, 60] STREAM_TIMEOUT_SECONDS = 90 # ntfy keepalive default is 55s; give margin _ECHO_TAG = "hermes-agent" # tag added to outgoing messages for echo-loop prevention +_MARKDOWN_TRUTHY = ("1", "true", "yes") def _build_auth_header(token: str) -> Dict[str, str]: - """Build an ``Authorization`` header from an ntfy token. + """``Authorization`` header from an ntfy token; ``{}`` when unset. - Shared by :class:`NtfyAdapter._auth_headers` and :func:`_standalone_send` - so both paths follow the same auth shape and whitespace-stripping rules. - - Tokens are stripped of surrounding whitespace — pasted tokens often - carry trailing newlines that would otherwise render the header - malformed (``Authorization: Bearer foo\\n``). ``user:pass`` tokens - become Basic auth; anything else is treated as a Bearer token. - Returns ``{}`` when no token is configured. + Tokens are whitespace-stripped (pasted tokens often carry newlines that + would malform the header). ``user:pass`` → Basic, anything else → Bearer. """ if not token: return {} @@ -133,11 +96,7 @@ def _build_auth_header(token: str) -> Dict[str, str]: def _truncate_body(message: str, *, context: str) -> bytes: - """Apply the ntfy 4096-char limit, logging a warning on truncation. - - ``context`` is included in the log message so adapter and standalone - truncations can be told apart in logs. - """ + """Apply the ntfy 4096-char limit, logging a warning (tagged ``context``) on truncation.""" if len(message) > MAX_MESSAGE_LENGTH: logger.warning( "%s: truncating message from %d to %d chars (ntfy limit)", @@ -146,64 +105,54 @@ def _truncate_body(message: str, *, context: str) -> bytes: return message[:MAX_MESSAGE_LENGTH].encode("utf-8") -def check_requirements() -> bool: - """Check whether the ntfy adapter is installable and minimally configured. +def _response_message_id(resp) -> str: + """ntfy's returned message id, or a random 12-hex fallback.""" + try: + return resp.json().get("id") or uuid.uuid4().hex[:12] + except Exception: + return uuid.uuid4().hex[:12] - Reads ``NTFY_TOPIC`` directly to avoid the cost of a full - ``load_gateway_config()`` (which also writes to ``os.environ``) on - every pre-flight check. - """ + +def check_requirements() -> bool: + """Installable and minimally configured (reads NTFY_TOPIC directly — no full config load).""" if not HTTPX_AVAILABLE: return False - topic = _get_scoped_secret("NTFY_TOPIC", "").strip() - return bool(topic) + return bool(_get_scoped_secret("NTFY_TOPIC", "").strip()) def validate_config(config) -> bool: - """Validate that the configured ntfy platform has a topic set.""" + """True when a topic is configured (config.yaml ``extra`` or env).""" extra = getattr(config, "extra", {}) or {} - topic = extra.get("topic") or _get_scoped_secret("NTFY_TOPIC", "") - return bool(topic) + return bool(extra.get("topic") or _get_scoped_secret("NTFY_TOPIC", "")) def is_connected(config) -> bool: """Check whether ntfy is configured (env or config.yaml).""" extra = getattr(config, "extra", {}) or {} - topic = _get_scoped_secret("NTFY_TOPIC") or extra.get("topic", "") - return bool(topic) + return bool(_get_scoped_secret("NTFY_TOPIC") or extra.get("topic", "")) class NtfyAdapter(BasePlatformAdapter): - """ntfy adapter. - - Subscribes to a topic via HTTP streaming (``/json`` endpoint) and - publishes replies via HTTP POST. No external SDK — only httpx. - """ + """ntfy adapter: HTTP-streaming subscription in, HTTP POST publish out.""" MAX_MESSAGE_LENGTH = MAX_MESSAGE_LENGTH def __init__(self, config: PlatformConfig): - platform = Platform("ntfy") - super().__init__(config=config, platform=platform) + super().__init__(config=config, platform=Platform("ntfy")) extra = config.extra or {} self._server: str = ( - extra.get("server") - or _get_scoped_secret("NTFY_SERVER_URL", DEFAULT_SERVER) + extra.get("server") or _get_scoped_secret("NTFY_SERVER_URL", DEFAULT_SERVER) ).rstrip("/") self._topic: str = extra.get("topic") or _get_scoped_secret("NTFY_TOPIC", "") self._publish_topic: str = ( - extra.get("publish_topic") - or _get_scoped_secret("NTFY_PUBLISH_TOPIC", "") - or self._topic + extra.get("publish_topic") or _get_scoped_secret("NTFY_PUBLISH_TOPIC", "") or self._topic ) self._token: str = extra.get("token") or _get_scoped_secret("NTFY_TOKEN", "") self._stream_task: Optional[asyncio.Task] = None self._http_client: Optional["httpx.AsyncClient"] = None - - # Message deduplication: msg_id -> timestamp - self._seen_messages: Dict[str, float] = {} + self._seen_messages: Dict[str, float] = {} # msg_id -> timestamp (dedup) # -- Connection lifecycle ----------------------------------------------- @@ -221,7 +170,6 @@ class NtfyAdapter(BasePlatformAdapter): self._stream_task = asyncio.create_task(self._run_stream()) self._mark_connected() logger.info("[%s] Connected — subscribing to %s/%s", self.name, self._server, self._topic) - # Plugin-registered native handlers (ctx.register_platform_handler). self._wire_plugin_handlers(None) return True except Exception as e: @@ -252,7 +200,6 @@ class NtfyAdapter(BasePlatformAdapter): if not self._running: return - # Reset backoff if stream stayed alive for at least 60s if time.monotonic() - stream_start >= 60.0: backoff_idx = 0 @@ -264,12 +211,11 @@ class NtfyAdapter(BasePlatformAdapter): async def _consume_stream(self, url: str, headers: Dict[str, str]) -> None: """Open an HTTP streaming connection and dispatch events.""" # poll=false keeps a persistent streaming connection alive with keepalive events - params = {"poll": "false"} async with self._http_client.stream( "GET", url, headers=headers, - params=params, + params={"poll": "false"}, timeout=httpx.Timeout(connect=15.0, read=STREAM_TIMEOUT_SECONDS, write=15.0, pool=15.0), ) as response: if response.status_code == 401: @@ -278,15 +224,12 @@ class NtfyAdapter(BasePlatformAdapter): self.name, ) self._set_fatal_error( - "ntfy_unauthorized", - "ntfy server rejected auth (401). Check NTFY_TOKEN.", - retryable=False, + "ntfy_unauthorized", "ntfy server rejected auth (401). Check NTFY_TOKEN.", retryable=False, ) raise _FatalStreamError("401 Unauthorized") if response.status_code == 404: logger.error( - "[%s] Topic not found (404): %s — stopping reconnect loop.", - self.name, self._topic, + "[%s] Topic not found (404): %s — stopping reconnect loop.", self.name, self._topic, ) self._set_fatal_error( "ntfy_topic_not_found", @@ -337,34 +280,19 @@ class NtfyAdapter(BasePlatformAdapter): if self._is_duplicate(msg_id): logger.debug("[%s] Duplicate message %s, skipping", self.name, msg_id) return - - # Echo-loop prevention: skip messages tagged by this adapter. - tags = event.get("tags") or [] - if _ECHO_TAG in tags: + if _ECHO_TAG in (event.get("tags") or []): logger.debug("[%s] Skipping own message (echo tag)", self.name) return - text = (event.get("message") or "").strip() if not text: logger.debug("[%s] Empty message body, skipping", self.name) return + # No native user identity on ntfy: the publisher-controlled title must + # NOT drive authorization, so user_id is fixed to the topic name. topic = event.get("topic") or self._topic - # ntfy has no native authenticated user identity. The title field is - # publisher-controlled and must NOT be used for authorization — any - # publisher who knows the topic can set title to an allowed username. - # Treat ntfy as a single trusted channel; user_id is fixed to the - # topic name. NTFY_ALLOWED_USERS is only a real trust boundary when - # the topic itself is protected by a read token. - user_id = topic - user_name = topic - source = self.build_source( - chat_id=topic, - chat_name=topic, - chat_type="dm", - user_id=user_id, - user_name=user_name, + chat_id=topic, chat_name=topic, chat_type="dm", user_id=topic, user_name=topic, ) unix_ts = event.get("time") @@ -384,19 +312,15 @@ class NtfyAdapter(BasePlatformAdapter): raw_message=event, timestamp=timestamp, ) - logger.debug("[%s] Message on topic %s: %s", self.name, topic, text[:80]) await self.handle_message(message_event) - # -- Deduplication ------------------------------------------------------ - def _is_duplicate(self, msg_id: str) -> bool: - """Return True if this message ID was already seen within the dedup window.""" + """True if this message ID was already seen within the dedup window.""" now = time.time() if len(self._seen_messages) > DEDUP_MAX_SIZE: cutoff = now - DEDUP_WINDOW_SECONDS self._seen_messages = {k: v for k, v in self._seen_messages.items() if v > cutoff} - if msg_id in self._seen_messages: return True self._seen_messages[msg_id] = now @@ -414,18 +338,16 @@ class NtfyAdapter(BasePlatformAdapter): """Publish a message to the configured publish topic.""" metadata = metadata or {} publish_topic = metadata.get("publish_topic") or self._publish_topic or chat_id - if not self._http_client: return SendResult(success=False, error="HTTP client not initialized") url = f"{self._server}/{publish_topic}" - markdown_enabled = (self.config.extra or {}).get("markdown", False) headers = { **self._auth_headers(), "Content-Type": "text/plain; charset=utf-8", "X-Tags": _ECHO_TAG, } - if markdown_enabled: + if (self.config.extra or {}).get("markdown", False): headers["X-Markdown"] = "true" if len(content) > self.MAX_MESSAGE_LENGTH: @@ -440,12 +362,7 @@ class NtfyAdapter(BasePlatformAdapter): url, content=body.encode("utf-8"), headers=headers, timeout=15.0, ) if resp.status_code < 300: - try: - data = resp.json() - returned_id = data.get("id") or uuid.uuid4().hex[:12] - except Exception: - returned_id = uuid.uuid4().hex[:12] - return SendResult(success=True, message_id=returned_id) + return SendResult(success=True, message_id=_response_message_id(resp)) body_text = resp.text logger.warning("[%s] Send failed HTTP %d: %s", self.name, resp.status_code, body_text[:200]) return SendResult(success=False, error=f"HTTP {resp.status_code}: {body_text[:200]}") @@ -455,38 +372,23 @@ class NtfyAdapter(BasePlatformAdapter): logger.error("[%s] Send error: %s", self.name, e) return SendResult(success=False, error=str(e)) - async def send_typing(self, chat_id: str, metadata=None) -> None: - """ntfy does not support typing indicators.""" - pass - async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: - """Return basic info about an ntfy topic.""" return {"name": chat_id, "type": "dm"} - # -- Helpers ------------------------------------------------------------ - def _auth_headers(self) -> Dict[str, str]: - """Build Authorization header if a token is configured.""" return _build_auth_header(self._token) -# --------------------------------------------------------------------------- -# Plugin registration -# --------------------------------------------------------------------------- +# -- Plugin registration ----------------------------------------------------- def _env_enablement() -> dict | None: """Seed ``PlatformConfig.extra`` from env vars during gateway config load. - Called by the platform registry's env-enablement hook BEFORE adapter - construction, so ``gateway status`` and ``get_connected_platforms()`` - reflect env-only configuration without instantiating the HTTP client. - Returns ``None`` when ntfy isn't minimally configured; the caller skips - auto-enabling. - - The special ``home_channel`` key in the returned dict is handled by the - core hook — it becomes a proper ``HomeChannel`` dataclass on the - ``PlatformConfig`` rather than being merged into ``extra``. + Runs BEFORE adapter construction so ``gateway status`` reflects env-only + setups without instantiating the HTTP client. ``None`` = not configured. + The ``home_channel`` key is lifted by the core hook into a ``HomeChannel`` + on the ``PlatformConfig`` instead of being merged into ``extra``. """ topic = _get_scoped_secret("NTFY_TOPIC", "").strip() if not topic: @@ -503,13 +405,10 @@ def _env_enablement() -> dict | None: seed["token"] = token markdown = _get_scoped_secret("NTFY_MARKDOWN", "").strip().lower() if markdown: - seed["markdown"] = markdown in ("1", "true", "yes") + seed["markdown"] = markdown in _MARKDOWN_TRUTHY home = _get_scoped_secret("NTFY_HOME_CHANNEL", "").strip() or topic if home: - seed["home_channel"] = { - "chat_id": home, - "name": _get_scoped_secret("NTFY_HOME_CHANNEL_NAME", home), - } + seed["home_channel"] = {"chat_id": home, "name": _get_scoped_secret("NTFY_HOME_CHANNEL_NAME", home)} return seed @@ -522,26 +421,17 @@ async def _standalone_send( media_files: Optional[List[str]] = None, force_document: bool = False, ) -> Dict[str, Any]: - """Out-of-process publish for cron / send_message_tool fallbacks. + """Out-of-process publish for cron / send_message_tool when no gateway adapter is live. - Used by ``tools/send_message_tool._send_via_adapter`` and the cron - scheduler when the gateway runner is not in this process (e.g. - ``hermes cron`` running standalone). Without this hook, - ``deliver=ntfy`` cron jobs fail with ``No live adapter for platform``. - - ``thread_id`` and ``media_files`` are accepted for signature parity - only — ntfy has no thread or attachment primitive. Markdown is - honored if ``NTFY_MARKDOWN`` is set OR ``pconfig.extra["markdown"]`` - is True. + ``thread_id``/``media_files`` are signature parity only (ntfy has no thread + or attachment primitive). Markdown is honored if ``NTFY_MARKDOWN`` is set + OR ``pconfig.extra["markdown"]`` is True. """ if not HTTPX_AVAILABLE: return {"error": "ntfy standalone send: httpx not installed"} extra = getattr(pconfig, "extra", {}) or {} - server = ( - extra.get("server") - or _get_scoped_secret("NTFY_SERVER_URL", DEFAULT_SERVER) - ).rstrip("/") + server = (extra.get("server") or _get_scoped_secret("NTFY_SERVER_URL", DEFAULT_SERVER)).rstrip("/") publish_topic = ( chat_id or extra.get("publish_topic") @@ -554,26 +444,20 @@ async def _standalone_send( token = extra.get("token") or _get_scoped_secret("NTFY_TOKEN", "") markdown_env = _get_scoped_secret("NTFY_MARKDOWN", "").strip().lower() - markdown_enabled = bool(extra.get("markdown")) or markdown_env in ("1", "true", "yes") - headers = {"Content-Type": "text/plain; charset=utf-8", "X-Tags": _ECHO_TAG, **_build_auth_header(token)} - if markdown_enabled: + if bool(extra.get("markdown")) or markdown_env in _MARKDOWN_TRUTHY: headers["X-Markdown"] = "true" body = _truncate_body(message, context="ntfy standalone") - url = f"{server}/{publish_topic}" try: async with httpx.AsyncClient(timeout=15.0) as client: resp = await client.post(url, content=body, headers=headers) if resp.status_code >= 300: return {"error": f"ntfy HTTP {resp.status_code}: {resp.text[:200]}"} - try: - data = resp.json() - msg_id = data.get("id") or uuid.uuid4().hex[:12] - except Exception: - msg_id = uuid.uuid4().hex[:12] - return {"success": True, "platform": "ntfy", "chat_id": publish_topic, "message_id": msg_id} + return { + "success": True, "platform": "ntfy", "chat_id": publish_topic, "message_id": _response_message_id(resp), + } except Exception as e: return {"error": f"ntfy standalone send failed: {e}"} @@ -589,25 +473,14 @@ def register(ctx) -> None: is_connected=is_connected, required_env=["NTFY_TOPIC"], install_hint="pip install httpx # already a Hermes dependency", - # Env-driven auto-configuration: seeds PlatformConfig.extra so - # env-only setups show up in `hermes gateway status` without - # instantiating the HTTP client. - env_enablement_fn=_env_enablement, - # Cron home-channel delivery support — `deliver=ntfy` cron jobs - # route to NTFY_HOME_CHANNEL when set. + env_enablement_fn=_env_enablement, # env-only setups show in `gateway status` cron_deliver_env_var="NTFY_HOME_CHANNEL", - # Out-of-process cron delivery. Without this hook, deliver=ntfy - # cron jobs fail with "No live adapter" when cron runs separately - # from the gateway. - standalone_sender_fn=_standalone_send, - # Auth env vars for _is_user_authorized() integration. + standalone_sender_fn=_standalone_send, # out-of-process cron delivery allowed_users_env="NTFY_ALLOWED_USERS", allow_all_env="NTFY_ALLOW_ALL_USERS", max_message_length=MAX_MESSAGE_LENGTH, emoji="🔔", - # ntfy publishers have no persistent identity — topic names are - # the only identifier, no phone numbers / emails to redact. - pii_safe=True, + pii_safe=True, # topic names only — no phone numbers / emails to redact allow_update_command=True, platform_hint=( "You are communicating via ntfy push notifications. " diff --git a/plugins/platforms/sms/adapter.py b/plugins/platforms/sms/adapter.py index 8d2592bc7b..df3724739a 100644 --- a/plugins/platforms/sms/adapter.py +++ b/plugins/platforms/sms/adapter.py @@ -1,7 +1,6 @@ """SMS (Twilio) platform adapter. -Connects to the Twilio REST API for outbound SMS and runs an aiohttp -webhook server to receive inbound messages. +Outbound SMS via the Twilio REST API; inbound via an aiohttp webhook server. Shares credentials with the optional telephony skill — same env vars: - TWILIO_ACCOUNT_SID @@ -24,6 +23,7 @@ import hashlib import hmac import logging import os +import re import urllib.parse from typing import Any, Dict, Optional @@ -37,28 +37,7 @@ from gateway.platforms.base import ( ) from gateway.platforms.helpers import redact_phone, strip_markdown -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 +from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret logger = logging.getLogger(__name__) @@ -68,6 +47,31 @@ MAX_SMS_LENGTH = 1600 # ~10 SMS segments DEFAULT_WEBHOOK_PORT = 8080 DEFAULT_WEBHOOK_HOST = "127.0.0.1" _TWILIO_WEBHOOK_MAX_BODY_BYTES = 65_536 # 64 KiB — Twilio payloads are small +_EMPTY_TWIML = '' + + +def _twiml_response(status: int = 200): + """Empty TwiML reply — replies go out via the REST API, never inline TwiML.""" + from aiohttp import web + + return web.Response(text=_EMPTY_TWIML, content_type="application/xml", status=status) + + +def _basic_auth(account_sid: str, auth_token: str) -> str: + """HTTP Basic auth header value for Twilio.""" + encoded = base64.b64encode(f"{account_sid}:{auth_token}".encode("ascii")).decode("ascii") + return f"Basic {encoded}" + + +def _twilio_form(from_number: str, to_number: str, body: str): + """Twilio Messages.json form payload (aiohttp FormData).""" + import aiohttp + + form_data = aiohttp.FormData() + form_data.add_field("From", from_number) + form_data.add_field("To", to_number) + form_data.add_field("Body", body) + return form_data def check_sms_requirements() -> bool: @@ -80,11 +84,10 @@ def check_sms_requirements() -> bool: class SmsAdapter(BasePlatformAdapter): - """ - Twilio SMS <-> Hermes gateway adapter. + """Twilio SMS <-> Hermes gateway adapter. - Each inbound phone number gets its own Hermes session (multi-tenant). - Replies are always sent from the configured TWILIO_PHONE_NUMBER. + Each inbound phone number gets its own Hermes session; replies are always + sent from the configured TWILIO_PHONE_NUMBER. """ MAX_MESSAGE_LENGTH = MAX_SMS_LENGTH @@ -94,23 +97,16 @@ class SmsAdapter(BasePlatformAdapter): self._account_sid: str = _get_scoped_secret("TWILIO_ACCOUNT_SID", "") self._auth_token: str = _get_scoped_secret("TWILIO_AUTH_TOKEN", "") self._from_number: str = os.getenv("TWILIO_PHONE_NUMBER", "") - self._webhook_port: int = int( - os.getenv("SMS_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT)) - ) + self._webhook_port: int = int(os.getenv("SMS_WEBHOOK_PORT", str(DEFAULT_WEBHOOK_PORT))) self._webhook_host: str = os.getenv("SMS_WEBHOOK_HOST", DEFAULT_WEBHOOK_HOST) self._webhook_url: str = os.getenv("SMS_WEBHOOK_URL", "").strip() self._runner = None self._http_session: Optional["aiohttp.ClientSession"] = None def _basic_auth_header(self) -> str: - """Build HTTP Basic auth header value for Twilio.""" - creds = f"{self._account_sid}:{self._auth_token}" - encoded = base64.b64encode(creds.encode("ascii")).decode("ascii") - return f"Basic {encoded}" + return _basic_auth(self._account_sid, self._auth_token) - # ------------------------------------------------------------------ - # Required abstract methods - # ------------------------------------------------------------------ + # -- Lifecycle ----------------------------------------------------------- async def connect(self, *, is_reconnect: bool = False) -> bool: import aiohttp @@ -123,7 +119,6 @@ class SmsAdapter(BasePlatformAdapter): return False insecure_no_sig = os.getenv("SMS_INSECURE_NO_SIGNATURE", "").lower() == "true" - if not self._webhook_url and not insecure_no_sig: msg = ( "[sms] Refusing to start: SMS_WEBHOOK_URL is required for Twilio " @@ -135,7 +130,6 @@ class SmsAdapter(BasePlatformAdapter): logger.error(msg) self._set_fatal_error("sms_missing_webhook_url", msg, retryable=False) return False - if insecure_no_sig and not self._webhook_url: logger.warning( "[sms] SMS_INSECURE_NO_SIGNATURE=true — Twilio signature validation " @@ -144,9 +138,8 @@ class SmsAdapter(BasePlatformAdapter): self._webhook_port, ) - # client_max_size bounds every read path — including chunked bodies - # with no Content-Length — before the handler's own 413 checks run - # (#58536/#58902/#59180 pattern). + # client_max_size bounds every read path (incl. chunked bodies with no + # Content-Length) before the handler's own 413 checks run. app = web.Application(client_max_size=_TWILIO_WEBHOOK_MAX_BODY_BYTES) app.router.add_post("/webhooks/twilio", self._handle_webhook) app.router.add_get("/health", lambda _: web.Response(text="ok")) @@ -156,18 +149,13 @@ class SmsAdapter(BasePlatformAdapter): site = web.TCPSite(self._runner, self._webhook_host, self._webhook_port) await site.start() self._http_session = aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=30), - trust_env=gateway_trust_env(), + timeout=aiohttp.ClientTimeout(total=30), trust_env=gateway_trust_env(), ) self._running = True - logger.info( "[sms] Twilio webhook server listening on %s:%d, from: %s", - self._webhook_host, - self._webhook_port, - redact_phone(self._from_number), + self._webhook_host, self._webhook_port, redact_phone(self._from_number), ) - # Plugin-registered native handlers (ctx.register_platform_handler). self._wire_plugin_handlers(None) return True @@ -181,6 +169,8 @@ class SmsAdapter(BasePlatformAdapter): self._running = False logger.info("[sms] Disconnected") + # -- Outbound ------------------------------------------------------------ + async def send( self, chat_id: str, @@ -190,26 +180,17 @@ class SmsAdapter(BasePlatformAdapter): ) -> SendResult: import aiohttp - formatted = self.format_message(content) - chunks = self.truncate_message(formatted) + chunks = self.truncate_message(self.format_message(content)) last_result = SendResult(success=True) - url = f"{TWILIO_API_BASE}/{self._account_sid}/Messages.json" - headers = { - "Authorization": self._basic_auth_header(), - } + headers = {"Authorization": self._basic_auth_header()} session = self._http_session or aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(total=30), - trust_env=gateway_trust_env(), + timeout=aiohttp.ClientTimeout(total=30), trust_env=gateway_trust_env(), ) try: for chunk in chunks: - form_data = aiohttp.FormData() - form_data.add_field("From", self._from_number) - form_data.add_field("To", chat_id) - form_data.add_field("Body", chunk) - + form_data = _twilio_form(self._from_number, chat_id, chunk) try: async with session.post(url, data=form_data, headers=headers) as resp: body = await resp.json() @@ -217,224 +198,121 @@ class SmsAdapter(BasePlatformAdapter): error_msg = body.get("message", str(body)) logger.error( "[sms] send failed to %s: %s %s", - redact_phone(chat_id), - resp.status, - error_msg, + redact_phone(chat_id), resp.status, error_msg, ) - return SendResult( - success=False, - error=f"Twilio {resp.status}: {error_msg}", - ) - msg_sid = body.get("sid", "") - last_result = SendResult(success=True, message_id=msg_sid) + return SendResult(success=False, error=f"Twilio {resp.status}: {error_msg}") + last_result = SendResult(success=True, message_id=body.get("sid", "")) except Exception as e: logger.error("[sms] send error to %s: %s", redact_phone(chat_id), e) return SendResult(success=False, error=str(e)) finally: - # Close session only if we created a fallback (no persistent session) + # Close only a fallback session we created ourselves. if not self._http_session and session: await session.close() - return last_result async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: return {"name": chat_id, "type": "dm"} - # ------------------------------------------------------------------ - # SMS-specific formatting - # ------------------------------------------------------------------ - def format_message(self, content: str) -> str: """Strip markdown — SMS renders it as literal characters.""" return strip_markdown(content) - # ------------------------------------------------------------------ - # Twilio signature validation - # ------------------------------------------------------------------ + # -- Twilio signature validation ----------------------------------------- - def _validate_twilio_signature( - self, url: str, post_params: dict, signature: str, - ) -> bool: - """Validate ``X-Twilio-Signature`` header (HMAC-SHA1, base64). + def _validate_twilio_signature(self, url: str, post_params: dict, signature: str) -> bool: + """Validate ``X-Twilio-Signature`` (HMAC-SHA1, base64). - Tries both with and without the default port for the URL scheme, - since Twilio may sign with either variant. - - Algorithm: https://www.twilio.com/docs/usage/security#validating-requests + Twilio may sign the URL with or without the scheme's default port, so + both variants are tried. https://www.twilio.com/docs/usage/security#validating-requests """ if self._check_signature(url, post_params, signature): return True - variant = self._port_variant_url(url) - if variant and self._check_signature(variant, post_params, signature): - return True + return bool(variant and self._check_signature(variant, post_params, signature)) - return False - - def _check_signature( - self, url: str, post_params: dict, signature: str, - ) -> bool: - """Compute and compare a single Twilio signature.""" - data_to_sign = url - for key in sorted(post_params.keys()): - data_to_sign += key + post_params[key] - mac = hmac.new( - self._auth_token.encode("utf-8"), - data_to_sign.encode("utf-8"), - hashlib.sha1, - ) + def _check_signature(self, url: str, post_params: dict, signature: str) -> bool: + data_to_sign = url + "".join(key + post_params[key] for key in sorted(post_params.keys())) + mac = hmac.new(self._auth_token.encode("utf-8"), data_to_sign.encode("utf-8"), hashlib.sha1) computed = base64.b64encode(mac.digest()).decode("utf-8") - # Compare as bytes: compare_digest raises TypeError on a str with - # non-ASCII characters, and the signature is a raw request header. + # Compare as bytes: compare_digest raises TypeError on non-ASCII str, + # and the signature is a raw request header. return hmac.compare_digest(computed.encode(), signature.encode()) @staticmethod def _port_variant_url(url: str) -> str | None: - """Return the URL with the default port toggled, or None. - - Only toggles default ports (443 for https, 80 for http). - Non-standard ports are never modified. - """ + """URL with the scheme's default port toggled (added/stripped); None for non-default ports.""" parsed = urllib.parse.urlparse(url) - default_ports = {"https": 443, "http": 80} - default_port = default_ports.get(parsed.scheme) + default_port = {"https": 443, "http": 80}.get(parsed.scheme) if default_port is None: return None - if parsed.port == default_port: - # Has explicit default port → strip it - return urllib.parse.urlunparse( - (parsed.scheme, parsed.hostname, parsed.path, - parsed.params, parsed.query, parsed.fragment) - ) + netloc = parsed.hostname elif parsed.port is None: - # No port → add default netloc = f"{parsed.hostname}:{default_port}" - return urllib.parse.urlunparse( - (parsed.scheme, netloc, parsed.path, - parsed.params, parsed.query, parsed.fragment) - ) + else: + return None + return urllib.parse.urlunparse( + (parsed.scheme, netloc, parsed.path, parsed.params, parsed.query, parsed.fragment) + ) - # Non-standard port — no variant - return None - - # ------------------------------------------------------------------ - # Twilio webhook handler - # ------------------------------------------------------------------ + # -- Inbound webhook ----------------------------------------------------- async def _handle_webhook(self, request) -> "aiohttp.web.Response": - from aiohttp import web - try: content_length = request.content_length if content_length is not None and content_length > _TWILIO_WEBHOOK_MAX_BODY_BYTES: - return web.Response( - text='', - content_type="application/xml", - status=413, - ) + return _twiml_response(413) raw = await request.read() if len(raw) > _TWILIO_WEBHOOK_MAX_BODY_BYTES: - return web.Response( - text='', - content_type="application/xml", - status=413, - ) + return _twiml_response(413) # Twilio sends form-encoded data, not JSON form = urllib.parse.parse_qs(raw.decode("utf-8"), keep_blank_values=True) except Exception as e: logger.error("[sms] webhook parse error: %s", e) - return web.Response( - text='', - content_type="application/xml", - status=400, - ) + return _twiml_response(400) - # Validate Twilio request signature when SMS_WEBHOOK_URL is configured if self._webhook_url: twilio_sig = request.headers.get("X-Twilio-Signature", "") if not twilio_sig: logger.warning("[sms] Rejected: missing X-Twilio-Signature header") - return web.Response( - text='', - content_type="application/xml", - status=403, - ) + return _twiml_response(403) flat_params = {k: v[0] for k, v in form.items() if v} - if not self._validate_twilio_signature( - self._webhook_url, flat_params, twilio_sig - ): + if not self._validate_twilio_signature(self._webhook_url, flat_params, twilio_sig): logger.warning("[sms] Rejected: invalid Twilio signature") - return web.Response( - text='', - content_type="application/xml", - status=403, - ) + return _twiml_response(403) - # Extract fields (parse_qs returns lists) - from_number = (form.get("From", [""]))[0].strip() - to_number = (form.get("To", [""]))[0].strip() - text = (form.get("Body", [""]))[0].strip() - message_sid = (form.get("MessageSid", [""]))[0].strip() + # parse_qs returns lists + from_number = form.get("From", [""])[0].strip() + to_number = form.get("To", [""])[0].strip() + text = form.get("Body", [""])[0].strip() + message_sid = form.get("MessageSid", [""])[0].strip() if not from_number or not text: - return web.Response( - text='', - content_type="application/xml", - ) - - # Ignore messages from our own number (echo prevention) - if from_number == self._from_number: + return _twiml_response() + if from_number == self._from_number: # echo prevention logger.debug("[sms] ignoring echo from own number %s", redact_phone(from_number)) - return web.Response( - text='', - content_type="application/xml", - ) + return _twiml_response() logger.info( - "[sms] inbound from %s -> %s: %s", - redact_phone(from_number), - redact_phone(to_number), - text[:80], + "[sms] inbound from %s -> %s: %s", redact_phone(from_number), redact_phone(to_number), text[:80], ) - source = self.build_source( - chat_id=from_number, - chat_name=from_number, - chat_type="dm", - user_id=from_number, - user_name=from_number, + chat_id=from_number, chat_name=from_number, chat_type="dm", + user_id=from_number, user_name=from_number, ) event = MessageEvent( - text=text, - message_type=MessageType.TEXT, - source=source, - raw_message=form, - message_id=message_sid, + text=text, message_type=MessageType.TEXT, source=source, raw_message=form, message_id=message_sid, ) - # Non-blocking: Twilio expects a fast response task = asyncio.create_task(self.handle_message(event)) self._background_tasks.add(task) task.add_done_callback(self._background_tasks.discard) - - # Return empty TwiML — we send replies via the REST API, not inline TwiML - return web.Response( - text='', - content_type="application/xml", - ) + return _twiml_response() -# ────────────────────────────────────────────────────────────────────────── -# Plugin migration glue (#41112 / #3823) -# -# Added when the SMS (Twilio) adapter moved from gateway/platforms/sms.py into -# this bundled plugin. register() exposes the platform via the registry, -# replacing the Platform.SMS elif in gateway/run.py, the -# _PLATFORM_CONNECTED_CHECKERS entry in gateway/config.py, the _PLATFORMS["sms"] -# static dict in hermes_cli/gateway.py, and the _send_sms dispatch in -# tools/send_message_tool.py. TWILIO_* env→PlatformConfig seeding stays in core. -# ────────────────────────────────────────────────────────────────────────── +# -- Plugin registration ----------------------------------------------------- +# TWILIO_* env→PlatformConfig seeding stays in core (gateway/config.py). def _strip_markdown_for_sms(message: str) -> str: @@ -460,14 +338,12 @@ async def _standalone_send( media_files=None, force_document=False, ): - """Out-of-process SMS delivery via the Twilio REST API. Implements the - standalone_sender_fn contract; replaces the legacy _send_sms helper.""" + """Out-of-process SMS delivery via the Twilio REST API (standalone_sender_fn contract).""" auth_token = getattr(pconfig, "api_key", None) or _get_scoped_secret("TWILIO_AUTH_TOKEN", "") try: import aiohttp except ImportError: return {"error": "aiohttp not installed. Run: pip install aiohttp"} - import base64 account_sid = _get_scoped_secret("TWILIO_ACCOUNT_SID", "") from_number = os.getenv("TWILIO_PHONE_NUMBER", "") @@ -485,17 +361,11 @@ async def _standalone_send( try: from gateway.platforms.base import resolve_proxy_url, proxy_kwargs_for_aiohttp - _proxy = resolve_proxy_url() - _sess_kw, _req_kw = proxy_kwargs_for_aiohttp(_proxy) - creds = f"{account_sid}:{auth_token}" - encoded = base64.b64encode(creds.encode("ascii")).decode("ascii") + _sess_kw, _req_kw = proxy_kwargs_for_aiohttp(resolve_proxy_url()) url = f"https://api.twilio.com/2010-04-01/Accounts/{account_sid}/Messages.json" - headers = {"Authorization": f"Basic {encoded}"} + headers = {"Authorization": _basic_auth(account_sid, auth_token)} async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30), **_sess_kw) as session: - form_data = aiohttp.FormData() - form_data.add_field("From", from_number) - form_data.add_field("To", chat_id) - form_data.add_field("Body", message) + form_data = _twilio_form(from_number, chat_id, message) async with session.post(url, data=form_data, headers=headers, **_req_kw) as resp: body = await resp.json() if resp.status >= 400: @@ -507,14 +377,12 @@ async def _standalone_send( def _is_connected(config) -> bool: - """SMS is connected when Twilio credentials are present. Mirrors the legacy - _PLATFORM_CONNECTED_CHECKERS[Platform.SMS] = bool(TWILIO_ACCOUNT_SID).""" + """SMS is connected when Twilio credentials are present (bool(TWILIO_ACCOUNT_SID)).""" import hermes_cli.gateway as gateway_mod return bool((gateway_mod.get_env_value("TWILIO_ACCOUNT_SID") or "").strip()) def _build_adapter(config): - """Factory wrapper that constructs SmsAdapter from a PlatformConfig.""" return SmsAdapter(config) diff --git a/tests/gateway/test_simplex_plugin.py b/tests/gateway/test_simplex_plugin.py index 90d3aa8ed1..465def9679 100644 --- a/tests/gateway/test_simplex_plugin.py +++ b/tests/gateway/test_simplex_plugin.py @@ -24,7 +24,6 @@ is_connected = _simplex.is_connected register = _simplex.register _env_enablement = _simplex._env_enablement _standalone_send = _simplex._standalone_send -_guess_extension = _simplex._guess_extension _is_image_ext = _simplex._is_image_ext _is_audio_ext = _simplex._is_audio_ext _CORR_PREFIX = _simplex._CORR_PREFIX @@ -104,14 +103,6 @@ def test_adapter_init_custom_url(): assert adapter._ws is None -# --------------------------------------------------------------------------- -# 5. Helper functions (magic-byte detection) -# --------------------------------------------------------------------------- - -def test_guess_extension_png(): - assert _guess_extension(b"\x89PNG\r\n\x1a\n") == ".png" - - # --------------------------------------------------------------------------- # 6. Correlation IDs # ---------------------------------------------------------------------------