From 6e562efd37a0643b17b5de1e9827a1b7195eb69b Mon Sep 17 00:00:00 2001 From: chengyongru Date: Thu, 25 Jun 2026 11:04:38 +0800 Subject: [PATCH] refactor: register webhook providers internally maintainer edit: route provider-specific webhook auth, context, and default prompt details through a small internal registry so new built-in platforms can be added without changing the router flow. --- nanobot/config/schema.py | 6 +- nanobot/webhooks.py | 180 +++++++++++++++++++++++++++------------ tests/test_webhooks.py | 15 ++++ 3 files changed, 144 insertions(+), 57 deletions(-) diff --git a/nanobot/config/schema.py b/nanobot/config/schema.py index ab3be0bbd..125eec9c2 100644 --- a/nanobot/config/schema.py +++ b/nanobot/config/schema.py @@ -307,7 +307,7 @@ class WebhookRouteConfig(Base): enabled: bool = True path: str = "" # Defaults to /webhooks/. - provider: Literal["generic", "github"] = "generic" + provider: str = "generic" auth: Literal["secret", "none"] = "secret" secret: str = Field(default="", repr=False) to: str = "" # Required when enabled, e.g. "websocket:github" or "telegram:12345". @@ -319,6 +319,10 @@ class WebhookRouteConfig(Base): @model_validator(mode="after") def _validate_route(self) -> "WebhookRouteConfig": + if re.fullmatch(r"[A-Za-z0-9_.-]{1,64}", self.provider) is None: + raise ValueError( + "webhook provider names may contain only letters, numbers, '_', '.', and '-'" + ) if self.path: if not self.path.startswith("/"): raise ValueError("webhook route path must start with '/'") diff --git a/nanobot/webhooks.py b/nanobot/webhooks.py index d9c6229ea..b50aca1e5 100644 --- a/nanobot/webhooks.py +++ b/nanobot/webhooks.py @@ -6,7 +6,7 @@ import hashlib import hmac import json import time -from collections.abc import Mapping +from collections.abc import Callable, Mapping from dataclasses import dataclass from typing import Any @@ -46,6 +46,14 @@ class WebhookError(Exception): self.message = message +@dataclass(frozen=True) +class WebhookProvider: + verify_secret: Callable[[str, Mapping[str, str], bytes], None] + context: Callable[[Mapping[str, str], Mapping[str, Any]], dict[str, Any]] + default_prompt_lines: Callable[[dict[str, Any]], list[str]] + require_json: bool = False + + class WebhookRouter: """Validate webhook requests and enqueue accepted events on the message bus.""" @@ -67,6 +75,10 @@ class WebhookRouter: for name, route in config.routes.items(): if not route.enabled: continue + try: + _webhook_provider(route.provider) + except WebhookError as exc: + raise ValueError(exc.message) from exc self._routes[_route_path(name, route)] = (name, route) @property @@ -132,7 +144,7 @@ class WebhookRouter: raise WebhookError(405, "webhook routes require POST") if len(body) > route.max_body_bytes: raise WebhookError(413, "webhook body is too large") - self._verify_auth(route, headers, body) + _verify_auth(route, headers, body) payload, body_text = _decode_body(route, body) context = _template_context( name=name, @@ -184,34 +196,6 @@ class WebhookRouter: "delivery_id": delivery_id or None, } - def _verify_auth( - self, - route: WebhookRouteConfig, - headers: Mapping[str, str], - body: bytes, - ) -> None: - if route.auth == "none": - return - secret = route.secret.strip() - if not secret: - raise WebhookError(500, "webhook route secret is not configured") - if route.provider == "github": - signature = headers.get("x-hub-signature-256", "") - if _hmac_matches(signature, secret, body): - return - raise WebhookError(401, "invalid GitHub webhook signature") - - signature = headers.get("x-nanobot-signature-256", "") - if signature and _hmac_matches(signature, secret, body): - return - bearer = _bearer_token(headers.get("authorization", "")) - header_token = headers.get("x-nanobot-auth", "") - if (bearer and hmac.compare_digest(bearer, secret)) or ( - header_token and hmac.compare_digest(header_token, secret) - ): - return - raise WebhookError(401, "invalid webhook secret") - def _is_duplicate( self, route_name: str, @@ -266,6 +250,39 @@ def _hmac_matches(signature: str, secret: str, body: bytes) -> bool: return hmac.compare_digest(supplied, expected) +def _verify_auth( + route: WebhookRouteConfig, + headers: Mapping[str, str], + body: bytes, +) -> None: + if route.auth == "none": + return + secret = route.secret.strip() + if not secret: + raise WebhookError(500, "webhook route secret is not configured") + _webhook_provider(route.provider).verify_secret(secret, headers, body) + + +def _verify_generic_secret(secret: str, headers: Mapping[str, str], body: bytes) -> None: + signature = headers.get("x-nanobot-signature-256", "") + if signature and _hmac_matches(signature, secret, body): + return + bearer = _bearer_token(headers.get("authorization", "")) + header_token = headers.get("x-nanobot-auth", "") + if (bearer and hmac.compare_digest(bearer, secret)) or ( + header_token and hmac.compare_digest(header_token, secret) + ): + return + raise WebhookError(401, "invalid webhook secret") + + +def _verify_github_secret(secret: str, headers: Mapping[str, str], body: bytes) -> None: + signature = headers.get("x-hub-signature-256", "") + if _hmac_matches(signature, secret, body): + return + raise WebhookError(401, "invalid GitHub webhook signature") + + def _decode_body(route: WebhookRouteConfig, body: bytes) -> tuple[Any, str]: try: text = body.decode("utf-8") @@ -276,8 +293,8 @@ def _decode_body(route: WebhookRouteConfig, body: bytes) -> tuple[Any, str]: try: return json.loads(text), text except json.JSONDecodeError as exc: - if route.provider == "github": - raise WebhookError(400, "GitHub webhook body must be JSON") from exc + if _webhook_provider(route.provider).require_json: + raise WebhookError(400, f"{route.provider} webhook body must be JSON") from exc return None, text @@ -291,15 +308,9 @@ def _template_context( remote: str | None, ) -> dict[str, Any]: event = payload if isinstance(payload, dict) else {} - github = _github_context(headers, event) if route.provider == "github" else {} - event_name = github.get("event") or headers.get("x-nanobot-event") or "" - delivery_id = ( - github.get("delivery_id") - or headers.get("x-nanobot-delivery") - or headers.get("x-webhook-id") - or headers.get("x-request-id") - or "" - ) + provider_context = _webhook_provider(route.provider).context(headers, event) + event_name = provider_context.pop("event_name", "") + delivery_id = provider_context.pop("delivery_id", "") return { "route": { "name": name, @@ -315,12 +326,59 @@ def _template_context( "body": body_text, "headers": _safe_headers(headers), "remote": remote or "", - "github": github, + "github": provider_context.get("github", {}), "event_name": event_name, "delivery_id": delivery_id, + **provider_context, } +def _generic_context(headers: Mapping[str, str], _payload: Mapping[str, Any]) -> dict[str, Any]: + return { + "event_name": headers.get("x-nanobot-event", ""), + "delivery_id": ( + headers.get("x-nanobot-delivery") + or headers.get("x-webhook-id") + or headers.get("x-request-id") + or "" + ), + } + + +def _generic_prompt_lines(_context: dict[str, Any]) -> list[str]: + return [] + + +def _github_provider_context( + headers: Mapping[str, str], + payload: Mapping[str, Any], +) -> dict[str, Any]: + github = _github_context(headers, payload) + return { + "github": github, + "event_name": github.get("event", ""), + "delivery_id": github.get("delivery_id", ""), + } + + +def _github_prompt_lines(context: dict[str, Any]) -> list[str]: + lines: list[str] = [] + github = context.get("github") or {} + if github.get("repository_full_name"): + lines.append(f"Repository: {github['repository_full_name']}") + if github.get("action"): + lines.append(f"Action: {github['action']}") + if github.get("sender_login"): + lines.append(f"Sender: {github['sender_login']}") + if github.get("ref"): + lines.append(f"Ref: {github['ref']}") + if github.get("pull_request_title"): + lines.append(f"Pull request: {github['pull_request_title']}") + elif github.get("issue_title"): + lines.append(f"Issue: {github['issue_title']}") + return lines + + def _github_context(headers: Mapping[str, str], payload: Mapping[str, Any]) -> dict[str, Any]: repo = payload.get("repository") sender = payload.get("sender") @@ -342,6 +400,29 @@ def _github_context(headers: Mapping[str, str], payload: Mapping[str, Any]) -> d } +_WEBHOOK_PROVIDERS: dict[str, WebhookProvider] = { + # ponytail: internal registry, add entry-point loading if third-party providers appear. + "generic": WebhookProvider( + verify_secret=_verify_generic_secret, + context=_generic_context, + default_prompt_lines=_generic_prompt_lines, + ), + "github": WebhookProvider( + verify_secret=_verify_github_secret, + context=_github_provider_context, + default_prompt_lines=_github_prompt_lines, + require_json=True, + ), +} + + +def _webhook_provider(name: str) -> WebhookProvider: + try: + return _WEBHOOK_PROVIDERS[name] + except KeyError as exc: + raise WebhookError(500, f"webhook provider {name!r} is not registered") from exc + + def _safe_headers(headers: Mapping[str, str]) -> dict[str, str]: safe: dict[str, str] = {} for key, value in headers.items(): @@ -399,20 +480,7 @@ def _default_prompt(context: dict[str, Any]) -> str: lines.append(f"Event: {event_name}") if delivery_id: lines.append(f"Delivery ID: {delivery_id}") - if provider == "github": - github = context.get("github") or {} - if github.get("repository_full_name"): - lines.append(f"Repository: {github['repository_full_name']}") - if github.get("action"): - lines.append(f"Action: {github['action']}") - if github.get("sender_login"): - lines.append(f"Sender: {github['sender_login']}") - if github.get("ref"): - lines.append(f"Ref: {github['ref']}") - if github.get("pull_request_title"): - lines.append(f"Pull request: {github['pull_request_title']}") - elif github.get("issue_title"): - lines.append(f"Issue: {github['issue_title']}") + lines.extend(_webhook_provider(provider).default_prompt_lines(context)) lines.extend(["", "Payload:", _format_payload(context.get("payload"), context.get("body", ""))]) return "\n".join(lines) diff --git a/tests/test_webhooks.py b/tests/test_webhooks.py index ef2d9f0ef..d2444aa66 100644 --- a/tests/test_webhooks.py +++ b/tests/test_webhooks.py @@ -359,6 +359,21 @@ def test_webhook_config_requires_target_for_enabled_routes() -> None: WebhooksConfig(routes={"bad": WebhookRouteConfig(auth="none", to="telegram")}) +def test_webhook_router_rejects_unregistered_provider() -> None: + config = WebhooksConfig( + routes={ + "stripe": WebhookRouteConfig( + provider="stripe", + auth="none", + to="telegram:1", + ) + } + ) + + with pytest.raises(ValueError, match="not registered"): + WebhookRouter(config, MessageBus()) + + def test_webhook_config_allows_incomplete_routes_when_webhooks_disabled() -> None: config = WebhooksConfig( enabled=False,