refactor(adapters/qq_signal): 8720->6042; qqbot keyboards/chunked_upload/onboard dedupe, signal send/receive helpers, bluebubbles + msgraph_webhook helper unification

This commit is contained in:
Teknium
2026-09-02 14:06:37 -07:00
parent 791417590d
commit 40532d1b7f
13 changed files with 1679 additions and 4357 deletions
+151 -251
View File
@@ -7,6 +7,7 @@ import hmac
import ipaddress
import json
import logging
import re
from collections import deque
from hashlib import sha1
from typing import Any, Awaitable, Callable, Dict, Optional
@@ -21,27 +22,22 @@ except ImportError:
from gateway.config import Platform, PlatformConfig
from gateway.platforms.base import (
BasePlatformAdapter,
MessageEvent,
MessageType,
SendResult,
is_network_accessible,
BasePlatformAdapter, MessageEvent, MessageType, SendResult, is_network_accessible,
)
logger = logging.getLogger(__name__)
# ``None`` → aiohttp/asyncio ``create_server`` binds one listening socket per
# address family (IPv4 + IPv6). The old "0.0.0.0" default bound IPv4 ONLY and
# was unreachable over IPv6-only private networks (e.g. Fly.io 6PN) — same
# bug as the LINE adapter (NS-603) and gateway/platforms/webhook.py
# (d542894ad). Pin a host via extra.host. The all-interfaces default still
# requires extra.allowed_source_cidrs (see _source_allowlist_required_but_missing).
# ``None`` → aiohttp binds one socket per address family (IPv4 + IPv6); the old
# "0.0.0.0" default was unreachable over IPv6-only private networks. Pin a host
# via extra.host. The all-interfaces default still requires
# extra.allowed_source_cidrs (see _source_allowlist_required_but_missing).
DEFAULT_HOST = None
DEFAULT_PORT = 8646
DEFAULT_WEBHOOK_PATH = "/msgraph/webhook"
DEFAULT_MAX_SEEN_RECEIPTS = 5000
DEFAULT_MAX_BODY_BYTES = 1_048_576
NotificationScheduler = Callable[[Dict[str, Any], MessageEvent], Awaitable[None] | None]
_TEMPLATE_KEY_RE = re.compile(r"\{([a-zA-Z0-9_.]+)\}")
def check_msgraph_webhook_requirements() -> bool:
@@ -49,6 +45,63 @@ def check_msgraph_webhook_requirements() -> bool:
return AIOHTTP_AVAILABLE
def _string_or_none(value: Any) -> Optional[str]:
if value is None:
return None
return str(value).strip() or None
def _normalize_path(path: Any) -> str:
raw = str(path or "").strip() or "/"
return raw if raw.startswith("/") else f"/{raw}"
def _parse_allowed_source_cidrs(raw: Any) -> list[ipaddress._BaseNetwork]:
"""Parse the optional CIDR allowlist; empty/missing means "allow everything".
When populated, requests from source IPs outside every listed CIDR are
rejected with 403 before the body is parsed (restrict to Microsoft
Graph's published webhook source ranges in production).
"""
if isinstance(raw, str):
candidates = raw.split(",")
elif isinstance(raw, (list, tuple, set)):
candidates = [str(chunk) for chunk in raw]
else:
return []
networks: list[ipaddress._BaseNetwork] = []
for chunk in candidates:
chunk = chunk.strip()
if not chunk:
continue
try:
networks.append(ipaddress.ip_network(chunk, strict=False))
except ValueError:
logger.warning("[msgraph_webhook] Ignoring invalid allowed_source_cidrs entry: %r", chunk)
return networks
def _prefix_match(resource: str, prefix: str) -> bool:
return resource == prefix or resource.startswith(f"{prefix}/")
def _render_template(template: str, payload: Dict[str, Any]) -> str:
"""Substitute ``{dotted.key}`` placeholders from *payload*; unknown keys stay literal."""
def _resolve(match: re.Match[str]) -> str:
key = match.group(1)
value: Any = payload
for part in key.split("."):
if not isinstance(value, dict):
return f"{{{key}}}"
value = value.get(part, f"{{{key}}}")
if isinstance(value, (dict, list)):
return json.dumps(value, sort_keys=True)[:2000]
return str(value)
return _TEMPLATE_KEY_RE.sub(_resolve, template)
class MSGraphWebhookAdapter(BasePlatformAdapter):
"""Receive Microsoft Graph change notifications and surface them internally."""
@@ -59,25 +112,15 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
_raw_host = extra.get("host", DEFAULT_HOST) or DEFAULT_HOST
self._host: Optional[str] = str(_raw_host) if _raw_host else None
self._port: int = int(extra.get("port", DEFAULT_PORT))
self._webhook_path: str = self._normalize_path(
extra.get("webhook_path", DEFAULT_WEBHOOK_PATH)
)
self._health_path: str = self._normalize_path(extra.get("health_path", "/health"))
self._webhook_path: str = _normalize_path(extra.get("webhook_path", DEFAULT_WEBHOOK_PATH))
self._health_path: str = _normalize_path(extra.get("health_path", "/health"))
self._accepted_resources: list[str] = [
str(value).strip()
for value in (extra.get("accepted_resources") or [])
if str(value).strip()
str(value).strip() for value in (extra.get("accepted_resources") or []) if str(value).strip()
]
self._client_state: Optional[str] = self._string_or_none(extra.get("client_state"))
self._max_seen_receipts = max(
1, int(extra.get("max_seen_receipts", DEFAULT_MAX_SEEN_RECEIPTS))
)
self._max_body_bytes = max(
1, int(extra.get("max_body_bytes", DEFAULT_MAX_BODY_BYTES))
)
self._allowed_source_networks: list[ipaddress._BaseNetwork] = (
self._parse_allowed_source_cidrs(extra.get("allowed_source_cidrs"))
)
self._client_state: Optional[str] = _string_or_none(extra.get("client_state"))
self._max_seen_receipts = max(1, int(extra.get("max_seen_receipts", DEFAULT_MAX_SEEN_RECEIPTS)))
self._max_body_bytes = max(1, int(extra.get("max_body_bytes", DEFAULT_MAX_BODY_BYTES)))
self._allowed_source_networks = _parse_allowed_source_cidrs(extra.get("allowed_source_cidrs"))
self._runner = None
self._notification_scheduler: Optional[NotificationScheduler] = None
self._seen_receipts: set[str] = set()
@@ -85,63 +128,6 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
self._accepted_count = 0
self._duplicate_count = 0
@staticmethod
def _string_or_none(value: Any) -> Optional[str]:
if value is None:
return None
text = str(value).strip()
return text or None
@staticmethod
def _normalize_path(path: Any) -> str:
raw = str(path or "").strip() or "/"
return raw if raw.startswith("/") else f"/{raw}"
@staticmethod
def _build_receipt_key(notification: Dict[str, Any]) -> Optional[str]:
explicit_id = str(notification.get("id") or "").strip()
if explicit_id:
return f"id:{explicit_id}"
return None
@staticmethod
def _normalize_resource_value(resource: str) -> str:
return str(resource or "").strip().strip("/")
@staticmethod
def _parse_allowed_source_cidrs(
raw: Any,
) -> list[ipaddress._BaseNetwork]:
"""Parse an optional list of CIDR ranges allowed to POST to the webhook.
An empty or missing value means "allow everything" (same behavior as
before this field existed). When populated, requests from source IPs
outside every listed CIDR are rejected with 403 before the body is
parsed. Use this to restrict the endpoint to Microsoft Graph's
published webhook source ranges in production deployments.
"""
if raw is None:
return []
if isinstance(raw, str):
candidates = [chunk.strip() for chunk in raw.split(",")]
elif isinstance(raw, (list, tuple, set)):
candidates = [str(chunk).strip() for chunk in raw]
else:
return []
networks: list[ipaddress._BaseNetwork] = []
for chunk in candidates:
if not chunk:
continue
try:
networks.append(ipaddress.ip_network(chunk, strict=False))
except ValueError:
logger.warning(
"[msgraph_webhook] Ignoring invalid allowed_source_cidrs entry: %r",
chunk,
)
return networks
def set_notification_scheduler(self, scheduler: Optional[NotificationScheduler]) -> None:
self._notification_scheduler = scheduler
@@ -152,9 +138,7 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
async def connect(self, *, is_reconnect: bool = False) -> bool:
if self._client_state is None:
logger.error(
"[msgraph_webhook] Refusing to start without extra.client_state configured"
)
logger.error("[msgraph_webhook] Refusing to start without extra.client_state configured")
return False
if self._source_allowlist_required_but_missing():
logger.error(
@@ -170,22 +154,14 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
app.router.add_get(self._health_path, self._handle_health)
app.router.add_get(self._webhook_path, self._handle_validation)
app.router.add_post(self._webhook_path, self._handle_notification)
# Plugin-registered native handlers (aiohttp web.Application —
# router routes). Wired before AppRunner.setup() freezes the router.
# Plugin-registered native routes; wired before AppRunner.setup() freezes the router.
self._wire_plugin_handlers(app)
self._runner = web.AppRunner(app)
await self._runner.setup()
site = web.TCPSite(self._runner, self._host, self._port)
await site.start()
self._mark_connected()
logger.info(
"[msgraph_webhook] Listening on %s:%d%s",
self._host,
self._port,
self._webhook_path,
)
logger.info("[msgraph_webhook] Listening on %s:%d%s", self._host, self._port, self._webhook_path)
return True
async def disconnect(self) -> None:
@@ -195,10 +171,7 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
self._mark_disconnected()
async def send(
self,
chat_id: str,
content: str,
reply_to: Optional[str] = None,
self, chat_id: str, content: str, reply_to: Optional[str] = None,
metadata: Optional[Dict[str, Any]] = None,
) -> SendResult:
logger.info("[msgraph_webhook] Response for %s: %s", chat_id, content[:200])
@@ -210,25 +183,17 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
async def _handle_health(self, request: "web.Request") -> "web.Response":
if not self._source_ip_allowed(request):
return web.Response(status=403)
return web.json_response(
{
"status": "ok",
"platform": self.platform.value,
"webhook_path": self._webhook_path,
"accepted": self._accepted_count,
"duplicates": self._duplicate_count,
}
)
return web.json_response({
"status": "ok",
"platform": self.platform.value,
"webhook_path": self._webhook_path,
"accepted": self._accepted_count,
"duplicates": self._duplicate_count,
})
async def _handle_validation(self, request: "web.Request") -> "web.Response":
"""Handle Microsoft Graph subscription validation handshake.
Graph validates a subscription endpoint by sending a GET with
``validationToken`` in the query string; the service must echo the
token verbatim as ``text/plain`` within 10 seconds. Anything else
(bare GET, GET without the token) is rejected so the endpoint can't
be enumerated or mistakenly used for data exfiltration.
"""
"""Graph subscription validation handshake: echo ``validationToken`` verbatim
as text/plain. Bare GETs are rejected so the endpoint can't be enumerated."""
if not self._source_ip_allowed(request):
return web.Response(status=403)
validation_token = request.query.get("validationToken", "")
@@ -239,43 +204,16 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
async def _handle_notification(self, request: "web.Request") -> "web.Response":
if not self._source_ip_allowed(request):
return web.Response(status=403)
# Graph never sends validationToken on POST, but tolerate it for
# defensive clients that replay the handshake in-band.
# Graph never sends validationToken on POST, but tolerate clients replaying it in-band.
validation_token = request.query.get("validationToken", "")
if validation_token:
return web.Response(text=validation_token, content_type="text/plain")
try:
content_length = request.content_length
except Exception:
content_length = None
if content_length is not None and content_length > self._max_body_bytes:
return web.Response(status=413)
try:
raw_body = await request.read()
except Exception:
return web.Response(status=400)
if len(raw_body) > self._max_body_bytes:
return web.Response(status=413)
try:
body = json.loads(raw_body.decode("utf-8"))
except (json.JSONDecodeError, UnicodeDecodeError):
return web.Response(status=400)
if not isinstance(body, dict):
return web.Response(status=400)
notifications = body.get("value")
if not isinstance(notifications, list):
return web.Response(status=400)
accepted = 0
duplicates = 0
auth_rejected = 0
other_rejected = 0
status, notifications = await self._read_notifications(request)
if status:
return web.Response(status=status)
accepted = duplicates = auth_rejected = other_rejected = 0
for raw_notification in notifications:
if not isinstance(raw_notification, dict):
other_rejected += 1
@@ -285,54 +223,64 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
other_rejected += 1
continue
if not self._verify_client_state(notification):
# Treat bad clientState as an auth failure: if the whole
# batch is forged, we want to signal 403 so the sender
# stops retrying. Legitimate Graph retries have valid
# clientState and hit the accepted/duplicate paths.
# Bad clientState is an auth failure: a fully forged batch gets 403
# so the sender stops retrying; legitimate Graph retries carry a
# valid clientState and hit the accepted/duplicate paths.
auth_rejected += 1
continue
receipt_key = self._build_receipt_key(notification)
explicit_id = str(notification.get("id") or "").strip()
receipt_key = f"id:{explicit_id}" if explicit_id else None
if receipt_key is not None:
if self._has_seen_receipt(receipt_key):
if receipt_key in self._seen_receipts:
duplicates += 1
continue
self._remember_receipt(receipt_key)
accepted += 1
self._accepted_count += 1
event = self._build_message_event(notification, receipt_key)
self._schedule_notification(notification, event)
self._schedule_notification(notification, self._build_message_event(notification, receipt_key))
self._duplicate_count += duplicates
# If anything ingested OR deduped, return 202 with empty body so
# Graph acks successfully and we don't leak internal counters. If
# every item failed auth, return 403 so an attacker POSTing fake
# notifications gets a clear reject. Other failures (malformed,
# resource-not-accepted) are the sender's configuration problem,
# so 400.
# Anything ingested OR deduped → 202 with empty body (Graph acks; no
# counter leak). Every item failed auth → 403 so forged POSTs get a
# clear reject. Otherwise (malformed / resource not accepted) → 400.
if accepted or duplicates:
return web.Response(status=202)
if auth_rejected and not other_rejected:
return web.Response(status=403)
return web.Response(status=400)
def _source_ip_allowed(self, request: "web.Request") -> bool:
"""Return True if the request's source IP is in the configured allowlist.
async def _read_notifications(self, request: "web.Request") -> tuple[int, list]:
"""Read and validate the POST body; returns (error_status, []) or (0, notifications)."""
try:
content_length = request.content_length
except Exception:
content_length = None
if content_length is not None and content_length > self._max_body_bytes:
return 413, []
try:
raw_body = await request.read()
except Exception:
return 400, []
if len(raw_body) > self._max_body_bytes:
return 413, []
try:
body = json.loads(raw_body.decode("utf-8"))
except (json.JSONDecodeError, UnicodeDecodeError):
return 400, []
notifications = body.get("value") if isinstance(body, dict) else None
if not isinstance(notifications, list):
return 400, []
return 0, notifications
Loopback-only binds may omit ``allowed_source_cidrs`` for local reverse
proxies and dev tunnels. Network-accessible binds fail closed until an
explicit CIDR allowlist is configured.
"""
def _source_ip_allowed(self, request: "web.Request") -> bool:
"""Loopback-only binds may omit ``allowed_source_cidrs`` (local proxies,
dev tunnels); network-accessible binds fail closed without one."""
if self._source_allowlist_required_but_missing():
return False
if not self._allowed_source_networks:
return True
peer = request.remote or ""
if not peer:
return False
try:
peer_addr = ipaddress.ip_address(peer)
peer_addr = ipaddress.ip_address(request.remote or "")
except ValueError:
return False
return any(peer_addr in network for network in self._allowed_source_networks)
@@ -340,72 +288,47 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
def _resource_accepted(self, resource: str) -> bool:
if not self._accepted_resources:
return True
normalized_resource = self._normalize_resource_value(resource)
resource = resource.strip().strip("/")
for pattern in self._accepted_resources:
normalized_pattern = self._normalize_resource_value(pattern)
if not normalized_pattern:
pattern = pattern.strip().strip("/")
if not pattern:
continue
if normalized_pattern.endswith("*"):
prefix = normalized_pattern[:-1].rstrip("/")
if normalized_resource == prefix or normalized_resource.startswith(f"{prefix}/"):
if pattern.endswith("*"):
if _prefix_match(resource, pattern[:-1].rstrip("/")):
return True
continue
if (
normalized_resource == normalized_pattern
or normalized_resource.startswith(f"{normalized_pattern}/")
):
elif _prefix_match(resource, pattern):
return True
return False
def _verify_client_state(self, notification: Dict[str, Any]) -> bool:
"""Verify the Graph-supplied clientState matches the configured secret.
Uses ``hmac.compare_digest`` instead of ``==`` so that a mismatch
doesn't leak how many leading characters matched via string-compare
timing. The configured client_state is a shared secret (documented in
the setup guide as "generate with ``openssl rand -hex 32``"), so a
timing-safe compare is the right primitive.
"""
"""Timing-safe compare of the Graph-supplied clientState against the
configured shared secret (``openssl rand -hex 32`` in the setup guide)."""
expected = self._client_state
if expected is None:
return False
provided = self._string_or_none(notification.get("clientState"))
provided = _string_or_none(notification.get("clientState"))
if provided is None:
return False
# Compare as bytes: ``compare_digest`` raises TypeError on a str with
# non-ASCII characters, and clientState comes from the request body.
# Compare as bytes: compare_digest raises TypeError on non-ASCII str,
# and clientState comes from the request body.
return hmac.compare_digest(provided.encode(), expected.encode())
def _has_seen_receipt(self, receipt_key: str) -> bool:
return receipt_key in self._seen_receipts
def _remember_receipt(self, receipt_key: str) -> None:
self._seen_receipts.add(receipt_key)
self._seen_receipt_order.append(receipt_key)
while len(self._seen_receipt_order) > self._max_seen_receipts:
oldest = self._seen_receipt_order.popleft()
self._seen_receipts.discard(oldest)
self._seen_receipts.discard(self._seen_receipt_order.popleft())
def _build_message_event(
self,
notification: Dict[str, Any],
receipt_key: Optional[str],
) -> MessageEvent:
def _build_message_event(self, notification: Dict[str, Any], receipt_key: Optional[str]) -> MessageEvent:
message_id = receipt_key or f"sha1:{sha1(json.dumps(notification, sort_keys=True).encode('utf-8')).hexdigest()}"
source = self.build_source(
chat_id=f"msgraph:{notification.get('subscriptionId', 'unknown')}",
chat_name="msgraph/webhook",
chat_type="webhook",
user_id="msgraph",
user_name="Microsoft Graph",
chat_name="msgraph/webhook", chat_type="webhook",
user_id="msgraph", user_name="Microsoft Graph",
)
return MessageEvent(
text=self._render_prompt(notification),
message_type=MessageType.TEXT,
source=source,
raw_message=notification,
message_id=message_id,
internal=True,
text=self._render_prompt(notification), message_type=MessageType.TEXT, source=source,
raw_message=notification, message_id=message_id, internal=True,
)
def _render_prompt(self, notification: Dict[str, Any]) -> str:
@@ -417,41 +340,18 @@ class MSGraphWebhookAdapter(BasePlatformAdapter):
"change_type": notification.get("changeType", ""),
"subscription_id": notification.get("subscriptionId", ""),
}
return self._render_template(template, payload)
return _render_template(template, payload)
rendered = json.dumps(notification, indent=2, sort_keys=True)[:4000]
return f"Microsoft Graph change notification:\n\n```json\n{rendered}\n```"
def _render_template(self, template: str, payload: Dict[str, Any]) -> str:
import re
def _resolve(match: "re.Match[str]") -> str:
key = match.group(1)
value: Any = payload
for part in key.split("."):
if isinstance(value, dict):
value = value.get(part, f"{{{key}}}")
else:
return f"{{{key}}}"
if isinstance(value, (dict, list)):
return json.dumps(value, sort_keys=True)[:2000]
return str(value)
return re.sub(r"\{([a-zA-Z0-9_.]+)\}", _resolve, template)
def _schedule_notification(
self,
notification: Dict[str, Any],
event: MessageEvent,
) -> None:
def _schedule_notification(self, notification: Dict[str, Any], event: MessageEvent) -> None:
scheduler = self._notification_scheduler
if scheduler is not None:
result = scheduler(notification, event)
if asyncio.iscoroutine(result):
task = asyncio.create_task(result)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
return
task = asyncio.create_task(self.handle_message(event))
if scheduler is None:
coro = self.handle_message(event)
else:
coro = scheduler(notification, event)
if not asyncio.iscoroutine(coro):
return
task = asyncio.create_task(coro)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)