refactor(platforms): a2a/buzz/dingtalk/email/google_chat/feishu-aux/discord-aux 11440->8812; dead code, dispatch tables, unified helpers

This commit is contained in:
Teknium
2026-09-02 23:33:41 -07:00
parent 113f04616b
commit 192058fda4
19 changed files with 2365 additions and 5032 deletions
+24 -42
View File
@@ -1,10 +1,5 @@
"""
A2A (Agent-to-Agent) plugin for Hermes Agent.
Registers the ``a2a`` platform adapter (inbound: exposes Hermes as an A2A v1.0
agent) and five client tools in the ``a2a`` toolset (outbound: call other
agents). Zero core edits — everything goes through the public PluginContext.
"""
"""A2A (Agent-to-Agent) plugin: registers the inbound ``a2a`` platform adapter and the
five outbound client tools of the ``a2a`` toolset through the public PluginContext."""
from __future__ import annotations
@@ -46,48 +41,44 @@ def is_connected(config) -> bool:
def interactive_setup() -> None:
"""`hermes gateway setup` flow for A2A."""
from hermes_cli.setup import (
prompt, prompt_yes_no, save_env_value, get_env_value, print_header, print_info, print_warning,
)
from hermes_cli.setup import prompt, prompt_yes_no, save_env_value, get_env_value, print_header, print_info, print_warning
print_header("A2A (Agent-to-Agent)")
print_info("Expose Hermes as an A2A-discoverable agent and call other A2A agents.")
print_info("Uses Python stdlib — no extra packages needed.")
print()
def ask(label: str, env: str) -> str:
"""Prompt with the current env value as default; save the stripped answer when non-blank."""
value = prompt(label, default=get_env_value(env) or "")
if value:
save_env_value(env, value.strip())
return value
port = prompt("Inbound A2A port (default 9900)", default=get_env_value("A2A_PORT") or "")
if port:
try:
save_env_value("A2A_PORT", str(int(port)))
except ValueError:
print_warning("Invalid port — using default 9900")
name = prompt("Agent name to advertise (blank = hostname-derived)", default=get_env_value("A2A_AGENT_NAME") or "")
if name:
save_env_value("A2A_AGENT_NAME", name.strip())
ask("Agent name to advertise (blank = hostname-derived)", "A2A_AGENT_NAME")
print()
print_info("Security: with NO token configured the server binds to 127.0.0.1 only.")
print_info("Prefer per-peer tokens (A2A_PEER_TOKENS=\"alice:tok1,bob:tok2\") so each")
print_info("remote agent has its own authenticated identity.")
for line in ("Security: with NO token configured the server binds to 127.0.0.1 only.",
"Prefer per-peer tokens (A2A_PEER_TOKENS=\"alice:tok1,bob:tok2\") so each",
"remote agent has its own authenticated identity."):
print_info(line)
if prompt_yes_no("Configure tokens to allow REMOTE A2A peers?", False):
peer_tokens = prompt(
"Per-peer tokens (name:token, comma-separated; blank to skip)",
default=get_env_value("A2A_PEER_TOKENS") or "",
)
if peer_tokens:
save_env_value("A2A_PEER_TOKENS", peer_tokens.strip())
peer_tokens = ask("Per-peer tokens (name:token, comma-separated; blank to skip)", "A2A_PEER_TOKENS")
token = prompt("Shared bearer token (blank to skip)", password=True)
if token:
save_env_value("A2A_BEARER_TOKEN", token)
if peer_tokens or token:
host = prompt("Bind host for remote access (e.g. 0.0.0.0)", default=get_env_value("A2A_HOST") or "")
if host:
save_env_value("A2A_HOST", host.strip())
ask("Bind host for remote access (e.g. 0.0.0.0)", "A2A_HOST")
else:
print_warning("No tokens entered — staying localhost-only.")
def register(ctx) -> None:
"""Plugin entry point — called by the Hermes plugin system."""
# Client tools register even when the inbound platform is disabled so the
# agent can call peers without exposing itself.
"""Plugin entry point. Client tools register even when the inbound platform is disabled
so the agent can call peers without exposing itself."""
try:
from .tools import register_tools
register_tools(ctx)
@@ -96,21 +87,12 @@ def register(ctx) -> None:
try:
from .adapter import A2AAdapter
ctx.register_platform(
name="a2a",
label="A2A",
adapter_factory=lambda cfg: A2AAdapter(cfg),
check_fn=check_requirements,
validate_config=validate_config,
is_connected=is_connected,
required_env=[],
install_hint="No extra packages needed (stdlib only)",
setup_fn=interactive_setup,
name="a2a", label="A2A", adapter_factory=lambda cfg: A2AAdapter(cfg),
check_fn=check_requirements, validate_config=validate_config, is_connected=is_connected,
required_env=[], install_hint="No extra packages needed (stdlib only)", setup_fn=interactive_setup,
emoji="\U0001f9e9", # puzzle piece
allowed_users_env="A2A_ALLOWED_USERS",
allow_all_env="A2A_ALLOW_ALL_USERS",
cron_deliver_env_var="A2A_HOME_CHANNEL",
allow_update_command=False,
platform_hint=_PLATFORM_HINT,
allowed_users_env="A2A_ALLOWED_USERS", allow_all_env="A2A_ALLOW_ALL_USERS",
cron_deliver_env_var="A2A_HOME_CHANNEL", allow_update_command=False, platform_hint=_PLATFORM_HINT,
)
except Exception:
logger.warning("A2A: failed to register platform adapter", exc_info=True)
+204 -335
View File
@@ -1,25 +1,12 @@
"""
A2A inbound platform adapter — exposes Hermes as an A2A-discoverable agent.
Stdlib http.server in a daemon thread (no a2a-sdk, no asyncio dependency at
register() time). Serves the v1.0 Agent Card at GET /.well-known/agent-card.json
(legacy agent.json too), /metrics, and JSON-RPC at POST /: message/send,
message/stream (SSE), tasks/{get,list,cancel,subscribe},
tasks/pushNotificationConfig/{create,get,list,delete}. Push payloads are v1.0
StreamResponse objects, HMAC-signed; configs may arrive inline in message/send.
Each inbound task is filtered + framed (security.wrap_inbound) and routed into
the agent's LIVE gateway session via the normal MessageEvent path (full memory,
not a clone). The reply returns through ``adapter.send()``, which fulfils the
per-task Future the HTTP handler blocks on; ``on_processing_complete`` resolves
failures promptly. Every exchange is persisted and audit-logged.
Bind safety: with no token configured, the server binds 127.0.0.1 only.
"""
"""A2A inbound adapter: stdlib http.server (daemon thread) serving the Agent Card, /metrics and
JSON-RPC (message/send, message/stream SSE, tasks/*, push-config CRUD). Inbound tasks are framed
(security.wrap_inbound) and routed into the LIVE gateway session; ``send()`` fulfils the per-task
Future the HTTP handler blocks on. No token configured => binds 127.0.0.1 only."""
from __future__ import annotations
import asyncio
import contextlib
import json
import logging
import os
@@ -38,15 +25,14 @@ from typing import Any, Dict, Optional
from gateway.platforms.base import BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult
from gateway.config import Platform
from gateway.platforms._shared import profile_scoped as _profile_scoped
from gateway.platforms._shared import coerce_port as _to_int, profile_scoped as _profile_scoped
from . import protocol, security
logger = logging.getLogger(__name__)
_DEFAULT_PORT = 9900
_ORPHAN_TIMEOUT = 300 # seconds before a pending task is considered orphaned
_WATCHDOG_INTERVAL = 60 # seconds between orphaned task watchdog runs
_ORPHAN_TIMEOUT, _WATCHDOG_INTERVAL = 300, 60 # seconds: pending task considered orphaned / watchdog period
_MAX_BODY = 1_048_576 # 1MB max request body — prevents DoS via memory exhaustion
_SSE_KEEPALIVE = 5 # seconds between SSE keepalive comments
_DEFAULT_DESCRIPTION = "Hermes Agent — a general-purpose agent reachable over A2A."
@@ -82,16 +68,8 @@ def _reply_timeout() -> float:
return 300.0
def _to_int(value: Any, default: Any) -> Any:
try:
return int(value) if value is not None else default
except (TypeError, ValueError):
return default
def _default_agent_name() -> str:
# Scope-aware: in a secondary multiplex profile os.environ holds the DEFAULT profile's
# A2A_AGENT_NAME — use the hostname default rather than another profile's identity.
# Scope-aware: a secondary multiplex profile must not borrow the default profile's A2A_AGENT_NAME.
name = "" if _profile_scoped() else os.getenv("A2A_AGENT_NAME", "").strip()
if name:
return name
@@ -103,15 +81,14 @@ def _default_agent_name() -> str:
def _clean_slug(value: str) -> str:
"""Return a URL-safe-ish single-segment slug for a served agent."""
"""URL-safe-ish single-segment slug for a served agent ("" for the root agent)."""
slug = str(value or "").strip().strip("/")
return "" if slug in ("", "default", "root") else slug.split("/")[0]
def _join_url(base: str, prefix: str) -> str:
base = (base or "").strip() or "/"
if not base.endswith("/"):
base += "/"
base = base if base.endswith("/") else base + "/"
prefix = (prefix or "").strip("/")
return urllib.parse.urljoin(base, prefix + "/") if prefix else base
@@ -125,17 +102,21 @@ def _active_profile_name() -> str:
def _profile_home(profile: str) -> Optional[str]:
try:
with contextlib.suppress(Exception):
from hermes_cli.profiles import get_profile_dir
return str(get_profile_dir(profile))
except Exception:
if not profile or profile == "default":
try:
from hermes_cli.config import get_hermes_home
return str(get_hermes_home())
except Exception:
return None
if profile and profile != "default":
return os.path.expanduser(f"~/.hermes/profiles/{profile}")
with contextlib.suppress(Exception):
from hermes_cli.config import get_hermes_home
return str(get_hermes_home())
return None
def _daemon_thread(target, name: str) -> threading.Thread:
t = threading.Thread(target=target, name=name, daemon=True)
t.start()
return t
def _safe_context_slug(value: str, max_len: int = 96) -> str:
@@ -147,16 +128,15 @@ def _safe_context_slug(value: str, max_len: int = 96) -> str:
def _state_db(profile: str, sql: str, params: tuple, log_msg: str, *, commit: bool = False) -> str:
"""Run one statement against a profile's state.db; first column of the first row or ""."""
home = _profile_home(profile)
db = os.path.join(home, "state.db") if home else None
db = os.path.join(home, "state.db") if home else ""
if not db or not os.path.exists(db):
return ""
try:
con = sqlite3.connect(db, timeout=5)
cur = con.execute(sql, params)
row = None if commit else cur.fetchone()
if commit:
con.commit()
con.close()
with contextlib.closing(sqlite3.connect(db, timeout=5)) as con:
cur = con.execute(sql, params)
row = None if commit else cur.fetchone()
if commit:
con.commit()
return str(row[0]) if row else ""
except Exception:
logger.debug(log_msg, exc_info=True)
@@ -164,8 +144,7 @@ def _state_db(profile: str, sql: str, params: tuple, log_msg: str, *, commit: bo
class A2ARequestHandler(BaseHTTPRequestHandler):
"""HTTP handler for the A2A JSON-RPC surface. Module-level so routing is
unit-testable; all state lives on ``self.server.adapter`` (set in connect())."""
"""HTTP handler for the A2A JSON-RPC surface; all state lives on ``self.server.adapter``."""
@property
def adapter(self) -> "A2AAdapter":
@@ -177,8 +156,8 @@ class A2ARequestHandler(BaseHTTPRequestHandler):
def _json(self, code: int, payload: dict):
body = json.dumps(payload).encode("utf-8")
self.send_response(code)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
for k, v in (("Content-Type", "application/json"), ("Content-Length", str(len(body)))):
self.send_header(k, v)
self.end_headers()
self.wfile.write(body)
@@ -189,42 +168,36 @@ class A2ARequestHandler(BaseHTTPRequestHandler):
return self.client_address[0] if self.client_address else ""
def _request_public_url(self) -> str:
"""Routable URL for this request: A2A_PUBLIC_URL > X-Forwarded-Host / Host
(scheme from X-Forwarded-Proto) > "" (caller falls back to bind host)."""
"""A2A_PUBLIC_URL > X-Forwarded-Host / Host (scheme from X-Forwarded-Proto) > "" (bind host)."""
explicit = os.getenv("A2A_PUBLIC_URL", "").strip()
if explicit:
return explicit
host = self.headers.get("X-Forwarded-Host", "") or self.headers.get("Host", "")
if not host:
return ""
host = host.split(",")[0].strip()
host = (self.headers.get("X-Forwarded-Host", "") or self.headers.get("Host", "")).split(",")[0].strip()
scheme = (self.headers.get("X-Forwarded-Proto", "") or "http").split(",")[0].strip()
return f"{scheme}://{host}/"
return f"{scheme}://{host}/" if host else ""
def do_GET(self): # noqa: N802
adapter = self.adapter
route = adapter._route_for_path(self.path)
agent = route["agent"]
subpath = route["subpath"].rstrip("/") or "/"
public_url = self._request_public_url() or None
if subpath in ("/.well-known/agent.json", "/.well-known/agent-card.json"):
self._json(200, adapter._build_card(self._request_public_url() or None, agent=agent))
elif subpath in ("/", "/health"):
payload = {"status": "ok", "agent": agent.get("name") or adapter.agent_name}
# Agent Cards are intentionally public; profile/tenant topology is not
# leaked on remote unauthenticated GETs.
sec = adapter._security_context
if sec.localhost_only() or sec.authenticate(self.headers.get("Authorization"), self._client_ip()) is not None:
payload["served_agents"] = adapter._served_agent_summary(public_url=self._request_public_url() or None)
self._json(200, payload)
elif subpath == "/metrics":
self._json(200, protocol.metrics.snapshot())
else:
self._json(404, {"error": "not found"})
return self._json(200, adapter._build_card(public_url, agent=agent))
if subpath == "/metrics":
return self._json(200, protocol.metrics.snapshot())
if subpath not in ("/", "/health"):
return self._json(404, {"error": "not found"})
payload = {"status": "ok", "agent": agent.get("name") or adapter.agent_name}
# Agent Cards are public; profile/tenant topology is not leaked on remote unauthenticated GETs.
sec = adapter._security_context
if sec.localhost_only() or sec.authenticate(self.headers.get("Authorization"), self._client_ip()) is not None:
payload["served_agents"] = adapter._served_agent_summary(public_url=public_url)
self._json(200, payload)
def do_POST(self): # noqa: N802
adapter = self.adapter
# Identity comes from the presented credential (or the socket in
# localhost-only mode) — never from the request body.
# Identity comes from the credential (or the socket in localhost-only mode) — never the body.
identity = adapter._security_context.authenticate(self.headers.get("Authorization"), self._client_ip())
if identity is None:
return self._error(401, None, protocol.ERR_UNAUTHORIZED, "unauthorized")
@@ -232,32 +205,32 @@ class A2ARequestHandler(BaseHTTPRequestHandler):
length = int(self.headers.get("Content-Length", 0))
if length > _MAX_BODY:
return self._error(413, None, protocol.ERR_PARSE, "payload too large")
raw = self.rfile.read(length) if length else b"{}"
req = json.loads(raw.decode("utf-8"))
req = json.loads((self.rfile.read(length) if length else b"{}").decode("utf-8"))
except Exception:
return self._error(400, None, protocol.ERR_PARSE, "parse error")
if not isinstance(req, dict):
return self._error(400, None, protocol.ERR_INVALID_PARAMS, "JSON-RPC request must be an object")
req_id = req.get("id")
method = str(req.get("method", ""))
req_id, method = req.get("id"), str(req.get("method", ""))
params = req["params"] if req.get("params") is not None else {}
if not isinstance(params, dict):
return self._error(200, req_id, protocol.ERR_INVALID_PARAMS, "params must be an object")
version = (self.headers.get("A2A-Version") or "").strip()
if version and version not in {"1.0", "1.0.0"}:
return self._error(200, req_id, protocol.ERR_INVALID_PARAMS, f"unsupported A2A-Version: {version}")
route = adapter._route_for_request(self.path, params) if isinstance(params, dict) else {}
handler_name, is_v1 = _METHODS.get(method, ("", False))
route = adapter._route_for_request(self.path, params)
if route.get("error"):
return self._error(400, req_id, protocol.ERR_INVALID_PARAMS, route["error"])
# Ordered, lazily-evaluated checks -> (http status, error code, message); first failure wins
# (the rate limiter must not be consulted for requests rejected before it).
checks = (
(lambda: not isinstance(params, dict), 200, protocol.ERR_INVALID_PARAMS, "params must be an object"),
(lambda: version and version not in {"1.0", "1.0.0"}, 200, protocol.ERR_INVALID_PARAMS, f"unsupported A2A-Version: {version}"),
(lambda: route.get("error"), 400, protocol.ERR_INVALID_PARAMS, route.get("error")),
(lambda: not adapter._rate_limiter.allow(identity), 429, protocol.ERR_RATE_LIMITED, "rate limit exceeded"),
(lambda: not adapter._security_context.is_trusted_peer(identity), 403, protocol.ERR_UNTRUSTED_PEER, f"peer '{identity}' not trusted"),
(lambda: not handler_name, 200, protocol.ERR_METHOD_NOT_FOUND, f"method not found: {method}"),
)
for failed, http, code, msg in checks:
if failed():
if code == protocol.ERR_RATE_LIMITED:
protocol.metrics.rate_limit_triggers += 1
return self._error(http, req_id, code, msg)
agent = route["agent"]
if not adapter._rate_limiter.allow(identity):
protocol.metrics.rate_limit_triggers += 1
return self._error(429, req_id, protocol.ERR_RATE_LIMITED, "rate limit exceeded")
if not adapter._security_context.is_trusted_peer(identity):
return self._error(403, req_id, protocol.ERR_UNTRUSTED_PEER, f"peer '{identity}' not trusted")
if not handler_name:
return self._error(200, req_id, protocol.ERR_METHOD_NOT_FOUND, f"method not found: {method}")
if handler_name == "_rpc_message_send":
self._json(200, adapter._rpc_message_send(req_id, params, identity, agent=agent, v1_response=is_v1))
elif handler_name == "_rpc_message_stream":
@@ -274,9 +247,8 @@ class A2AAdapter(BasePlatformAdapter):
def __init__(self, config, **kwargs):
super().__init__(config=config, platform=Platform("a2a"))
extra = getattr(config, "extra", {}) or {}
# Scope-aware: a secondary multiplex profile must not borrow the default profile's
# bridged A2A_PORT (falls closed to the module default). advertised_toolsets is
# deliberately left unscoped while its None-vs-empty-list semantics are in flux.
# Scope-aware: a secondary multiplex profile must not borrow the default profile's bridged
# A2A_PORT (falls closed to the module default). advertised_toolsets is deliberately unscoped.
self._security_context = security.A2ASecurityContext.capture()
_port_env = None if _profile_scoped() else os.getenv("A2A_PORT")
self.port = int(_port_env or extra.get("port", _DEFAULT_PORT))
@@ -287,21 +259,15 @@ class A2AAdapter(BasePlatformAdapter):
self._active_profile = _active_profile_name()
self._agents = self._load_served_agents(extra)
self._httpd: Optional[ThreadingHTTPServer] = None
self._server_thread: Optional[threading.Thread] = None
self._server_thread = self._watchdog_thread = None # type: Optional[threading.Thread]
self._loop: Optional[asyncio.AbstractEventLoop] = None
self._watchdog_stop = threading.Event()
self._watchdog_thread: Optional[threading.Thread] = None
# Per-adapter protocol state (not module-global).
self.tasks = protocol.TaskStore()
self._turns = protocol.TurnTracker()
self._rate_limiter = protocol.RateLimiter()
self.tasks, self._turns, self._rate_limiter = protocol.TaskStore(), protocol.TurnTracker(), protocol.RateLimiter()
# Forwarded profile sessions: (profile, agent_slug, context_id) -> session_id.
self._profile_sessions: Dict[tuple[str, str, str], str] = {}
self._profile_session_locks: Dict[tuple[str, str, str], threading.Lock] = {}
self._profile_session_locks_guard = threading.Lock()
# Pending reply futures: task_id -> (context_id, Future). _pending_order keeps per-context
# FIFO so adapter.send() — which only knows the context — resolves the oldest task.
self._pending: Dict[str, tuple[str, Future]] = {}
@@ -314,20 +280,13 @@ class A2AAdapter(BasePlatformAdapter):
@property
def authorization_is_upstream(self) -> bool:
"""Every request is authenticated in ``do_POST`` before dispatch; without this
the gateway's ``A2A_ALLOWED_USERS`` allow-list would reject peers, whose identity
is a token-derived name or IP. Not fail-open: a wrong credential is still 401'd."""
"""Requests are authenticated in ``do_POST``; the gateway's A2A_ALLOWED_USERS list would
otherwise reject peers (identity is a token-derived name or IP). Wrong credentials still 401."""
return True
# ── Lifecycle ─────────────────────────────────────────────────────────
async def connect(self, **_kwargs) -> bool:
# **_kwargs: base lifecycle passes ``is_reconnect`` etc. Capture the gateway loop so
# the HTTP thread can marshal events onto it via run_coroutine_threadsafe.
try:
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None
# Capture the gateway loop so the HTTP thread can marshal events via run_coroutine_threadsafe.
self._loop = asyncio.get_running_loop()
try:
self._httpd = ThreadingHTTPServer((self.host, self.port), A2ARequestHandler)
except OSError as e:
@@ -336,34 +295,27 @@ class A2AAdapter(BasePlatformAdapter):
return False
self._httpd.daemon_threads = True
self._httpd.adapter = self # type: ignore[attr-defined]
self._server_thread = threading.Thread(target=self._httpd.serve_forever, name="a2a-http", daemon=True)
self._server_thread.start()
self._server_thread = _daemon_thread(self._httpd.serve_forever, "a2a-http")
self._watchdog_stop.clear() # disconnect sets it; reset for reconnection
self._watchdog_thread = threading.Thread(target=self._watchdog_loop, name="a2a-watchdog", daemon=True)
self._watchdog_thread.start()
self._watchdog_thread = _daemon_thread(self._watchdog_loop, "a2a-watchdog")
self._mark_connected()
exposure = "localhost-only" if self._security_context.localhost_only() else "REMOTE (bearer auth)"
logger.info("A2A: serving Agent Card + JSON-RPC on http://%s:%s (%s) as %r; %d routed agent(s)",
self.host, self.port, exposure, self.agent_name, len(self._agents))
# Plugin-registered native handlers (ctx.register_platform_handler).
self._wire_plugin_handlers(None)
logger.info("A2A: serving Agent Card + JSON-RPC on http://%s:%s (%s) as %r; %d routed agent(s)", self.host, self.port,
"localhost-only" if self._security_context.localhost_only() else "REMOTE (bearer auth)", self.agent_name, len(self._agents))
self._wire_plugin_handlers(None) # plugin-registered native handlers
return True
async def disconnect(self) -> None:
self._mark_disconnected()
self._watchdog_stop.set()
if self._httpd is not None:
try:
with contextlib.suppress(Exception):
self._httpd.shutdown()
self._httpd.server_close()
except Exception:
pass
self._httpd = None
# Fail any in-flight replies so blocked HTTP threads don't hang.
with self._pending_lock:
for _ctx, fut in self._pending.values():
if not fut.done():
fut.set_result((protocol.STATE_FAILED, "[agent shutting down]"))
for tid in list(self._pending):
self._resolve_locked(tid, protocol.STATE_FAILED, "[agent shutting down]")
self._pending.clear()
self._pending_order.clear()
@@ -377,8 +329,6 @@ class A2AAdapter(BasePlatformAdapter):
except Exception:
logger.debug("A2A: watchdog error", exc_info=True)
# ── Agent routing + Agent Cards ───────────────────────────────────────
def _load_served_agents(self, extra: dict) -> dict[str, dict]:
"""Served-agent routing from ``platforms.a2a.extra.agents`` (top-level ``a2a_served_agents``
fallback for scripts/tests). Root/default always maps to the live gateway session."""
@@ -387,11 +337,10 @@ class A2AAdapter(BasePlatformAdapter):
try:
from hermes_cli.config import load_config
cfg = load_config() or {}
cfg = cfg if isinstance(cfg, dict) else {}
except Exception:
cfg = {}
cfg = cfg if isinstance(cfg, dict) else {}
raw = cfg.get("a2a_served_agents") or (cfg.get("a2a") or {}).get("served_agents")
# Scope-aware like port: a secondary profile must not inherit A2A_AGENT_DESCRIPTION.
default_desc = _DEFAULT_DESCRIPTION if _profile_scoped() else os.getenv("A2A_AGENT_DESCRIPTION", _DEFAULT_DESCRIPTION)
agents: dict[str, dict] = {"": {
@@ -415,7 +364,6 @@ class A2AAdapter(BasePlatformAdapter):
toolsets = val.get("advertised_toolsets") or val.get("toolsets") or val.get("capabilities") or []
if isinstance(toolsets, str):
toolsets = [t.strip() for t in toolsets.split(",") if t.strip()]
local = bool(val.get("local")) or profile in ("", "default", self._active_profile)
tenant = str(val.get("tenant") or slug).strip()
if tenant:
if tenant in tenants:
@@ -424,7 +372,8 @@ class A2AAdapter(BasePlatformAdapter):
continue
tenants[tenant] = slug
agents[slug] = {
"slug": slug, "path": "/" + path_segment, "tenant": tenant, "profile": profile or slug, "local": local,
"slug": slug, "path": "/" + path_segment, "tenant": tenant, "profile": profile or slug,
"local": bool(val.get("local")) or profile in ("", "default", self._active_profile),
"name": str(val.get("name") or f"Hermes {slug}"),
"description": str(val.get("description") or f"Hermes profile '{profile or slug}' exposed over A2A."),
"advertised_toolsets": list(toolsets or []),
@@ -437,18 +386,15 @@ class A2AAdapter(BasePlatformAdapter):
def _served_agent_summary(self, public_url: Optional[str] = None) -> list[dict]:
base = self._base_url(public_url)
return [
{"slug": a["slug"] or "default", "name": a.get("name"), "url": _join_url(base, a.get("path", "")),
"tenant": a.get("tenant") or None, "profile": a.get("profile"), "local": bool(a.get("local"))}
for a in self._agents.values()
]
return [{"slug": a["slug"] or "default", "name": a.get("name"), "url": _join_url(base, a.get("path", "")),
"tenant": a.get("tenant") or None, "profile": a.get("profile"), "local": bool(a.get("local"))}
for a in self._agents.values()]
def _route_for_path(self, raw_path: str) -> dict:
path = urllib.parse.urlsplit(raw_path or "/").path or "/"
# Longest prefix wins. Default/root agent is the fallback.
for agent in sorted(self._agents.values(), key=lambda a: len(a.get("path", "")), reverse=True):
prefix = agent.get("path", "") or ""
if prefix and (path == prefix or path.startswith(prefix + "/")):
if (prefix := agent.get("path", "") or "") and (path == prefix or path.startswith(prefix + "/")):
return {"agent": agent, "subpath": path[len(prefix):] or "/"}
return {"agent": self._agents[""], "subpath": path}
@@ -457,51 +403,39 @@ class A2AAdapter(BasePlatformAdapter):
agent = route["agent"]
tenant = str((params or {}).get("tenant") or "")
# If no URL prefix chose a non-default agent, allow v1.0 tenant routing.
if agent.get("slug") == "" and tenant:
matches = [a for a in self._agents.values() if a.get("tenant") == tenant]
if matches:
route = {"agent": matches[0], "subpath": route["subpath"]}
agent = matches[0]
if agent.get("slug") == "" and tenant and (matches := [a for a in self._agents.values() if a.get("tenant") == tenant]):
agent = matches[0]
route = {"agent": agent, "subpath": route["subpath"]}
expected = str(agent.get("tenant") or "")
if tenant and expected and tenant != expected:
return {"error": f"tenant {tenant!r} does not match routed agent {agent.get('slug') or 'default'}"}
return route
def _build_card(self, public_url: Optional[str] = None, agent: Optional[dict] = None) -> dict:
# Per-request public URL (X-Forwarded-Host / Host / A2A_PUBLIC_URL) beats
# the bind host so peers behind a reverse proxy can call back.
# Per-request public URL beats the bind host so peers behind a reverse proxy can call back.
agent = agent or self._agents[""]
return protocol.build_agent_card(
name=agent.get("name") or self.agent_name,
url=_join_url(self._base_url(public_url), agent.get("path", "")),
description=agent.get("description") or _DEFAULT_DESCRIPTION,
skills=self._advertised_skills(agent),
streaming=bool(agent.get("local", True)),
push_notifications=True,
auth_required=not self._security_context.localhost_only(),
tenant=str(agent.get("tenant") or ""),
name=agent.get("name") or self.agent_name, url=_join_url(self._base_url(public_url), agent.get("path", "")),
description=agent.get("description") or _DEFAULT_DESCRIPTION, skills=self._advertised_skills(agent),
streaming=bool(agent.get("local", True)), push_notifications=True,
auth_required=not self._security_context.localhost_only(), tenant=str(agent.get("tenant") or ""),
)
def _advertised_skills(self, agent: Optional[dict] = None) -> list[dict]:
"""Dynamic Agent Card skills from the live tool registry, restricted by
``advertised_toolsets`` / A2A_ADVERTISED_TOOLSETS; static fallback without a registry."""
"""Agent Card skills from the live tool registry, restricted by ``advertised_toolsets``;
static fallback without a registry."""
configured = (agent or {}).get("advertised_toolsets") if agent else self._advertised_toolsets
try:
from tools.registry import registry as tool_registry
allowed = set(configured or []) or None
mapping = {
n: tool_registry.get_tool_names_for_toolset(n)
for n in tool_registry.get_registered_toolset_names()
if allowed is None or n in allowed
}
mapping = {n: tool_registry.get_tool_names_for_toolset(n)
for n in tool_registry.get_registered_toolset_names() if allowed is None or n in allowed}
if mapping:
return protocol.skills_from_toolsets(mapping)
except Exception:
logger.debug("A2A: tool registry unavailable for Agent Card", exc_info=True)
return protocol.skills_from_toolsets(configured or [])
# ── Pending reply plumbing ────────────────────────────────────────────
def _add_pending(self, task_id: str, context_id: str) -> Future:
fut: Future = Future()
with self._pending_lock:
@@ -513,18 +447,17 @@ class A2AAdapter(BasePlatformAdapter):
with self._pending_lock:
entry = self._pending.pop(task_id, None)
order = self._pending_order.get(entry[0]) if entry else None
if order:
if task_id in order:
order.remove(task_id)
if not order:
self._pending_order.pop(entry[0], None)
if order and task_id in order:
order.remove(task_id)
if order is not None and not order:
self._pending_order.pop(entry[0], None)
def _resolve_locked(self, task_id: str, state: str, text: str) -> bool:
entry = self._pending.get(task_id)
if entry and not entry[1].done():
entry[1].set_result((state, text))
return True
return False
if not entry or entry[1].done():
return False
entry[1].set_result((state, text))
return True
def _resolve_task(self, task_id: str, state: str, text: str) -> bool:
with self._pending_lock:
@@ -535,20 +468,16 @@ class A2AAdapter(BasePlatformAdapter):
return any(self._resolve_locked(tid, state, text) for tid in self._pending_order.get(context_id, ()))
def _scope_for_agent(self, agent: Optional[dict]) -> tuple[str, str]:
agent = agent or self._agents[""]
return str(agent.get("slug") or ""), str(agent.get("tenant") or "")
return tuple(str((agent or self._agents[""]).get(k) or "") for k in ("slug", "tenant"))
def _forward_lock(self, key: tuple[str, str, str]) -> threading.Lock:
with self._profile_session_locks_guard:
return self._profile_session_locks.setdefault(key, threading.Lock())
# ── Inbound task handling ─────────────────────────────────────────────
def _end_task(self, rec: dict, state: str, text: str, stored_reply: str = "") -> tuple[dict, None]:
"""Complete a task immediately (rejected / not ready) and build its terminal Task."""
self.tasks.complete(rec["task_id"], state, stored_reply)
if state == protocol.STATE_FAILED:
protocol.metrics.tasks_failed += 1
protocol.metrics.tasks_failed += state == protocol.STATE_FAILED
return protocol.build_task(rec["task_id"], rec["context_id"], state, text, created_at=rec["created_iso"]), None
def _prepare_task(self, params: dict, peer: str, agent: Optional[dict] = None) -> tuple[Optional[dict], Optional[dict]]:
@@ -564,11 +493,8 @@ class A2AAdapter(BasePlatformAdapter):
if turn > max_turns:
protocol.metrics.anti_loop_triggers += 1
logger.warning("A2A: anti-loop triggered for context %s (turn %d > %d)", context_id, turn, max_turns)
return self._end_task(
rec, protocol.STATE_REJECTED,
f"Anti-loop protection: context {context_id} exceeded {max_turns} turns. "
f"Start a new context or increase A2A_MAX_PINGPONG_TURNS.",
)
return self._end_task(rec, protocol.STATE_REJECTED, f"Anti-loop protection: context {context_id} exceeded "
f"{max_turns} turns. Start a new context or increase A2A_MAX_PINGPONG_TURNS.")
if not text:
return self._end_task(rec, protocol.STATE_REJECTED, "Empty task — nothing to do.")
framed = security.wrap_inbound(peer, text)
@@ -583,12 +509,8 @@ class A2AAdapter(BasePlatformAdapter):
if self._loop is None or self._message_handler is None:
return self._end_task(rec, protocol.STATE_FAILED, "Agent gateway not ready to accept A2A tasks.")
fut = self._add_pending(task_id, context_id)
event = MessageEvent(
text=framed,
message_type=MessageType.TEXT,
source=self.build_source(chat_id=context_id, chat_name=f"a2a:{peer}", chat_type="dm", user_id=peer, user_name=peer),
message_id=task_id,
)
event = MessageEvent(text=framed, message_type=MessageType.TEXT, message_id=task_id,
source=self.build_source(chat_id=context_id, chat_name=f"a2a:{peer}", chat_type="dm", user_id=peer, user_name=peer))
try:
asyncio.run_coroutine_threadsafe(self.handle_message(event), self._loop)
except Exception as e:
@@ -596,13 +518,11 @@ class A2AAdapter(BasePlatformAdapter):
msg = security.redact_outbound(f"Dispatch failed: {e}")
return self._end_task(rec, protocol.STATE_FAILED, msg, stored_reply=msg)
self.tasks.set_state(task_id, protocol.STATE_WORKING)
return None, {"task_id": task_id, "context_id": context_id, "peer": peer, "future": fut,
"created_iso": rec["created_iso"], "started": time.time()}
return None, {"task_id": task_id, "context_id": context_id, "peer": peer, "future": fut, "created_iso": rec["created_iso"], "started": time.time()}
def _forward_to_profile(self, agent: dict, peer: str, context_id: str, framed_text: str) -> tuple[str, str]:
"""Forward a routed A2A task to another local Hermes profile via ``hermes chat``.
First contact creates a ``source=a2a`` session, records its id and titles it
deterministically; later turns ``--resume`` that concrete id (stable multi-turn continuity)."""
"""Forward a routed task to another local profile via ``hermes chat``. First contact creates a
``source=a2a`` session and titles it deterministically; later turns ``--resume`` that id."""
profile = str(agent.get("profile") or agent.get("slug") or "").strip()
slug = str(agent.get("slug") or profile or "agent")
safe_ctx = _safe_context_slug(context_id)
@@ -612,16 +532,11 @@ class A2AAdapter(BasePlatformAdapter):
with self._forward_lock(key):
session_id = self._profile_sessions.get(key) or _state_db(
profile, "SELECT id FROM sessions WHERE title = ? ORDER BY started_at DESC LIMIT 1",
(session_title,), "A2A: could not lookup forwarded session",
)
cmd = ["hermes", "chat", "-q", framed_text, "-Q", "--source", "a2a"]
if session_id:
cmd.extend(["--resume", session_id])
env = os.environ.copy()
home = _profile_home(profile)
if home:
(session_title,), "A2A: could not lookup forwarded session")
cmd = ["hermes", "chat", "-q", framed_text, "-Q", "--source", "a2a"] + (["--resume", session_id] if session_id else [])
env = {**os.environ, "HERMES_A2A_PEER": peer}
if home := _profile_home(profile):
env["HERMES_HOME"] = home
env["HERMES_A2A_PEER"] = peer
start = time.time()
try:
proc = subprocess.run(cmd, capture_output=True, text=True, encoding="utf-8", errors="replace",
@@ -633,15 +548,12 @@ class A2AAdapter(BasePlatformAdapter):
if proc.returncode != 0:
msg = (proc.stderr or proc.stdout or f"profile exited {proc.returncode}").strip()
return security.redact_outbound(msg[-2000:]), protocol.STATE_FAILED
if not session_id:
session_id = _state_db(
if not session_id and (session_id := _state_db(
profile, "SELECT id FROM sessions WHERE source = 'a2a' AND started_at >= ? ORDER BY started_at DESC LIMIT 1",
(start - 2.0,), "A2A: could not find latest forwarded session",
)
if session_id:
self._profile_sessions[key] = session_id
_state_db(profile, "UPDATE sessions SET title = ? WHERE id = ?", (session_title, session_id),
"A2A: could not title forwarded session", commit=True)
(start - 2.0,), "A2A: could not find latest forwarded session")):
self._profile_sessions[key] = session_id
_state_db(profile, "UPDATE sessions SET title = ? WHERE id = ?", (session_title, session_id),
"A2A: could not title forwarded session", commit=True)
return security.redact_outbound((proc.stdout or "").strip()), protocol.STATE_COMPLETED
def _record_outcome(self, task_id: str, context_id: str, peer: str, state: str, reply: str,
@@ -649,49 +561,49 @@ class A2AAdapter(BasePlatformAdapter):
"""Persist + audit + count a finished task, mark it terminal, and fire its push callback."""
protocol.persist_message(context_id, "agent", reply, task_id)
security.audit("outbound", peer, task_id, reply)
m = protocol.metrics
if state in (protocol.STATE_COMPLETED, protocol.STATE_INPUT_REQUIRED):
protocol.metrics.outbound_total += 1
protocol.metrics.tasks_completed += 1
m.outbound_total, m.tasks_completed = m.outbound_total + 1, m.tasks_completed + 1
if started is not None:
protocol.metrics.record_latency(time.time() - started)
m.record_latency(time.time() - started)
else:
protocol.metrics.tasks_failed += 1
m.tasks_failed += 1
self.tasks.complete(task_id, state, reply)
self._send_push_notification(task_id, context_id, reply, state)
def _finalize_task(self, pending: dict, state: str, reply: str) -> tuple[str, str]:
"""Record the outcome of a dispatched task; returns (state, reply) after
redaction and input-required detection."""
"""Record a dispatched task's outcome; returns (state, reply) after redaction and
input-required detection (a leading marker flags a clarification request)."""
task_id, context_id, peer = pending["task_id"], pending["context_id"], pending["peer"]
self._pop_pending(task_id)
reply = security.redact_outbound(reply or "")
# A leading marker flags a clarification request -> A2A input-required.
if state == protocol.STATE_COMPLETED:
stripped = reply.lstrip()
if stripped.upper().startswith(protocol.INPUT_REQUIRED_MARKER):
state = protocol.STATE_INPUT_REQUIRED
reply = stripped[len(protocol.INPUT_REQUIRED_MARKER):].strip()
stripped = reply.lstrip()
if state == protocol.STATE_COMPLETED and stripped.upper().startswith(protocol.INPUT_REQUIRED_MARKER):
state, reply = protocol.STATE_INPUT_REQUIRED, stripped[len(protocol.INPUT_REQUIRED_MARKER):].strip()
self._record_outcome(task_id, context_id, peer, state, reply, started=pending["started"])
return state, reply
def _await_reply(self, pending: dict, keepalive=None) -> tuple[str, str]:
"""Block until the task's future resolves (or times out). ``keepalive`` runs every
_SSE_KEEPALIVE seconds while waiting; if it raises, the client is gone and we stop."""
fut: Future = pending["future"]
deadline = pending["started"] + _reply_timeout()
@staticmethod
def _await_future(fut: Future, deadline: float, keepalive, on_timeout: tuple[str, str]) -> tuple[str, str]:
"""Block until ``fut`` resolves or ``deadline`` passes (-> ``on_timeout``). ``keepalive`` runs
every _SSE_KEEPALIVE seconds while waiting; if it raises, the client is gone and we stop."""
while True:
try:
return fut.result(timeout=_SSE_KEEPALIVE if keepalive else max(0.0, deadline - time.time()))
except FuturesTimeout:
if time.time() >= deadline:
return (protocol.STATE_FAILED, "[agent did not reply in time]")
return on_timeout
if keepalive:
try:
keepalive()
except Exception:
return (protocol.STATE_FAILED, "[client disconnected]")
except Exception:
return (protocol.STATE_FAILED, "[agent did not reply in time]")
return on_timeout
def _await_reply(self, pending: dict, keepalive=None) -> tuple[str, str]:
return self._await_future(pending["future"], pending["started"] + _reply_timeout(), keepalive,
(protocol.STATE_FAILED, "[agent did not reply in time]"))
def _rpc_message_send(self, req_id: Any, params: dict, peer: str, agent: Optional[dict] = None, v1_response: bool = False) -> dict:
task, pending = self._prepare_task(params, peer, agent=agent)
@@ -700,31 +612,31 @@ class A2AAdapter(BasePlatformAdapter):
task = protocol.build_task(pending["task_id"], pending["context_id"], state, reply, created_at=pending["created_iso"])
return _ok(req_id, protocol.send_message_response(task) if v1_response else task)
# ── Streaming (SSE) ───────────────────────────────────────────────────
@staticmethod
def _sse_headers(handler) -> None:
handler.send_response(200)
handler.send_header("Content-Type", "text/event-stream")
handler.send_header("Cache-Control", "no-cache")
for k, v in (("Content-Type", "text/event-stream"), ("Cache-Control", "no-cache")):
handler.send_header(k, v)
handler.end_headers()
# v1.0: stream closure signals the terminal state, so the socket must
# actually close once we emit the done event.
handler.close_connection = True
handler.close_connection = True # v1.0: stream closure signals the terminal state
@staticmethod
def _sse_write(handler, chunk: str) -> None:
handler.wfile.write(chunk.encode("utf-8"))
handler.wfile.flush()
@classmethod
def _keepalive(cls, handler):
return lambda: cls._sse_write(handler, ": keepalive\n\n")
def _emit_terminal(self, handler, task_id: str, context_id: str, state: str, reply: str, req_id: Any = None) -> None:
"""Emit the final artifact/status events and close the stream (v1.0: closure
signals terminal state). ``req_id`` threads into the JSON-RPC SSE envelope (§9.4)."""
if reply and state == protocol.STATE_COMPLETED:
self._sse_write(handler, protocol.sse_data(protocol.artifact_update(task_id, context_id, reply), req_id))
self._sse_write(handler, protocol.sse_data(protocol.status_update(task_id, context_id, state), req_id))
else:
self._sse_write(handler, protocol.sse_data(protocol.status_update(task_id, context_id, state, reply), req_id))
"""Emit the final artifact/status events and the closure marker. ``req_id`` threads into the
JSON-RPC SSE envelope (§9.4)."""
completed = bool(reply) and state == protocol.STATE_COMPLETED
events = ([protocol.artifact_update(task_id, context_id, reply)] if completed else []) + [
protocol.status_update(task_id, context_id, state, "" if completed else reply)]
for ev in events:
self._sse_write(handler, protocol.sse_data(ev, req_id))
self._sse_write(handler, protocol.sse_done())
def _rpc_message_stream(self, handler, req_id: Any, params: dict, peer: str, agent: Optional[dict] = None) -> None:
@@ -734,14 +646,13 @@ class A2AAdapter(BasePlatformAdapter):
try:
terminal, pending = self._prepare_task(params, peer, agent=agent)
if terminal is not None:
text = protocol.extract_text(terminal.get("status", {}).get("message", {}) or {})
return self._emit_terminal(handler, terminal["id"], terminal["contextId"], terminal["status"]["state"], text, req_id=req_id)
return self._emit_terminal(handler, terminal["id"], terminal["contextId"], terminal["status"]["state"],
protocol.extract_text(terminal.get("status", {}).get("message", {}) or {}), req_id=req_id)
task_id, context_id = pending["task_id"], pending["context_id"]
submitted = protocol.build_task(task_id, context_id, protocol.STATE_SUBMITTED, created_at=pending["created_iso"])
self._sse_write(handler, protocol.sse_data(protocol.stream_task(submitted), req_id))
self._sse_write(handler, protocol.sse_data(protocol.status_update(task_id, context_id, protocol.STATE_WORKING), req_id))
state, reply = self._await_reply(pending, keepalive=lambda: self._sse_write(handler, ": keepalive\n\n"))
state, reply = self._finalize_task(pending, state, reply)
state, reply = self._finalize_task(pending, *self._await_reply(pending, keepalive=self._keepalive(handler)))
self._emit_terminal(handler, task_id, context_id, state, reply, req_id=req_id)
except (BrokenPipeError, ConnectionResetError):
logger.debug("A2A: stream client disconnected")
@@ -753,39 +664,23 @@ class A2AAdapter(BasePlatformAdapter):
return handler._json(200, error)
self._sse_headers(handler)
try:
fut = self.tasks.watch(task_id, *self._scope_for_agent(agent))
if fut is None:
if (fut := self.tasks.watch(task_id, *self._scope_for_agent(agent))) is None:
return self._sse_write(handler, protocol.sse_done())
deadline = time.time() + _reply_timeout()
while True:
try:
state, reply = fut.result(timeout=_SSE_KEEPALIVE)
break
except FuturesTimeout:
if time.time() >= deadline:
state, reply = rec["state"], rec.get("reply", "")
break
self._sse_write(handler, ": keepalive\n\n")
state, reply = self._await_future(fut, time.time() + _reply_timeout(), self._keepalive(handler),
(rec["state"], rec.get("reply", "")))
self._emit_terminal(handler, task_id, rec["context_id"], state, reply, req_id=req_id)
except (BrokenPipeError, ConnectionResetError):
logger.debug("A2A: subscribe client disconnected")
# ── Task queries ──────────────────────────────────────────────────────
def _find_task(self, req_id: Any, params: dict, agent: Optional[dict]) -> tuple[str, Optional[dict], Optional[dict]]:
"""(task_id, record, None) for a visible task, else (task_id, None, jsonrpc_error)."""
task_id = str(params.get("taskId") or params.get("id") or "")
rec = self.tasks.get(task_id, *self._scope_for_agent(agent))
if not rec:
return task_id, None, _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"task not found: {task_id}")
return task_id, rec, None
return task_id, rec, None if rec else _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"task not found: {task_id}")
def _rpc_tasks_get(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
_task_id, rec, error = self._find_task(req_id, params, agent)
if error:
return error
history_len = _to_int(params.get("historyLength"), None)
return _ok(req_id, protocol.TaskStore.to_task(rec, history_length=history_len))
return error or _ok(req_id, protocol.TaskStore.to_task(rec))
def _rpc_tasks_list(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
offset = _to_int(params.get("pageToken") or 0, 0)
@@ -793,16 +688,11 @@ class A2AAdapter(BasePlatformAdapter):
agent_slug, tenant = self._scope_for_agent(agent)
recs, next_offset, total = self.tasks.list(
context_id=str(params.get("contextId") or ""), state=str(params.get("status") or params.get("state") or ""),
page_size=page_size, offset=max(0, offset), agent_slug=agent_slug, tenant=tenant, with_total=True,
)
page_size=page_size, offset=max(0, offset), agent_slug=agent_slug, tenant=tenant, with_total=True)
include_artifacts = bool(params.get("includeArtifacts", False))
history_len = _to_int(params.get("historyLength"), None)
return _ok(req_id, {
"tasks": [protocol.TaskStore.to_task(r, history_length=history_len, include_artifacts=include_artifacts) for r in recs],
"nextPageToken": str(next_offset) if next_offset else "",
"pageSize": max(1, min(page_size, 100)),
"totalSize": total,
})
return _ok(req_id, {"tasks": [protocol.TaskStore.to_task(r, include_artifacts=include_artifacts) for r in recs],
"nextPageToken": str(next_offset) if next_offset else "",
"pageSize": max(1, min(page_size, 100)), "totalSize": total})
def _rpc_tasks_cancel(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
task_id, rec, error = self._find_task(req_id, params, agent)
@@ -816,94 +706,75 @@ class A2AAdapter(BasePlatformAdapter):
rec = self.tasks.get(task_id, *self._scope_for_agent(agent)) or rec
return _ok(req_id, protocol.TaskStore.to_task(rec))
# ── Push notifications ────────────────────────────────────────────────
def _register_inline_push(self, task_id: str, params: dict, agent: Optional[dict] = None) -> None:
"""v1.0: message/send can carry configuration.taskPushNotificationConfig."""
cfg = (params.get("configuration") or {}).get("taskPushNotificationConfig") or {}
if not isinstance(cfg, dict):
return
url = cfg.get("url") or (cfg.get("pushNotificationConfig") or {}).get("url") or ""
url = (cfg.get("url") or (cfg.get("pushNotificationConfig") or {}).get("url") or "") if isinstance(cfg, dict) else ""
if url:
self.tasks.set_push_config(task_id, str(url), *self._scope_for_agent(agent))
def _rpc_push_config_create(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
task_id = str(params.get("taskId") or "")
cfg = params.get("pushNotificationConfig") or params.get("config") or {}
url = str((cfg or {}).get("url") or "")
url = str((params.get("pushNotificationConfig") or params.get("config") or {}).get("url") or "")
if not task_id or not url:
return _err(req_id, protocol.ERR_INVALID_PARAMS, "taskId and pushNotificationConfig.url required")
stored = self.tasks.set_push_config(task_id, url, *self._scope_for_agent(agent))
if stored is None:
return _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"task not found: {task_id}")
return _ok(req_id, stored)
return _ok(req_id, stored) if stored is not None else _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"task not found: {task_id}")
def _push_config_op(self, req_id: Any, params: dict, agent: Optional[dict], op, render) -> dict:
"""Shared get/list/delete: ``op(task_id, config_id, slug, tenant)`` falsy => not found."""
task_id = str(params.get("taskId") or "")
if not task_id:
return _err(req_id, protocol.ERR_INVALID_PARAMS, "taskId required")
config_id = str(params.get("id") or params.get("configId") or "")
found = op(task_id, config_id, *self._scope_for_agent(agent))
if not found:
return _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"push config not found for task: {task_id}")
return _ok(req_id, render(found))
found = op(task_id, str(params.get("id") or params.get("configId") or ""), *self._scope_for_agent(agent))
return _ok(req_id, render(found)) if found else _err(req_id, protocol.ERR_TASK_NOT_FOUND, f"push config not found for task: {task_id}")
def _rpc_push_config_get(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
return self._push_config_op(req_id, params, agent, self.tasks.get_push_config, lambda cfg: cfg)
def _rpc_push_config_list(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
task_id = str(params.get("taskId") or "")
if not task_id:
return _err(req_id, protocol.ERR_INVALID_PARAMS, "taskId required")
configs = self.tasks.list_push_configs(task_id, *self._scope_for_agent(agent))
return _ok(req_id, {"configs": configs, "nextPageToken": ""})
# An empty list is a valid (non-error) result, hence the ``or [[]]`` sentinel through the shared op.
return self._push_config_op(req_id, params, agent, lambda tid, _cid, *scope: self.tasks.list_push_configs(tid, *scope) or [[]],
lambda found: {"configs": [c for c in found if c], "nextPageToken": ""})
def _rpc_push_config_delete(self, req_id: Any, params: dict, agent: Optional[dict] = None) -> dict:
return self._push_config_op(req_id, params, agent, self.tasks.delete_push_config, lambda _: {"deleted": True})
def _send_push_notification(self, task_id: str, context_id: str, reply: str, state: str) -> None:
"""POST a v1.0 StreamResponse payload to the task's registered callback.
The URL is SSRF-checked (internal/private/loopback blocked unless localhost-only mode)."""
"""POST a v1.0 StreamResponse payload to the task's registered callback (SSRF-checked URL)."""
def fail(msg: str, *args) -> None:
protocol.metrics.push_failed += 1
logger.warning("A2A: push notification for task %s " + msg, task_id, *args)
callback_url = self.tasks.pop_push_url(task_id)
if not callback_url:
return
if not security.is_safe_callback_url(callback_url, localhost_mode=self._security_context.localhost_only()):
logger.warning("A2A: push notification for task %s blocked — unsafe callback URL: %s", task_id, callback_url)
protocol.metrics.push_failed += 1
return
return fail("blocked — unsafe callback URL: %s", callback_url)
payload = protocol.status_update(task_id, context_id, state, (reply or "")[:2000])
signature = self._security_context.sign_push_payload(payload)
headers = {"Content-Type": "application/json"}
if signature:
if signature := self._security_context.sign_push_payload(payload):
headers["X-A2A-Signature"] = signature
try:
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(callback_url, data=data, headers=headers, method="POST")
req = urllib.request.Request(callback_url, data=json.dumps(payload).encode("utf-8"), headers=headers, method="POST")
with urllib.request.urlopen(req, timeout=10) as resp: # noqa: S310
if 200 <= resp.status < 300:
protocol.metrics.push_sent += 1
logger.debug("A2A: push notification sent for task %s", task_id)
else:
protocol.metrics.push_failed += 1
logger.warning("A2A: push notification for task %s got HTTP %d", task_id, resp.status)
status = resp.status
except Exception as e:
protocol.metrics.push_failed += 1
logger.warning("A2A: push notification for task %s failed: %s", task_id, e)
# ── Sending (the agent's reply path) ──────────────────────────────────
return fail("failed: %s", e)
if not 200 <= status < 300:
return fail("got HTTP %d", status)
protocol.metrics.push_sent += 1
logger.debug("A2A: push notification sent for task %s", task_id)
async def send(self, chat_id: str, content: str, reply_to: Optional[str] = None, metadata: Optional[Dict[str, Any]] = None):
"""Fulfil the oldest pending reply Future for this context (``chat_id`` = A2A context id).
Only sends carrying ``metadata['notify']`` (the base adapter's final-reply marker,
``_mark_notify_metadata``) satisfy the caller; progress/status/preview sends must not."""
message_id = str(int(time.time() * 1000))
Only sends carrying ``metadata['notify']`` (the base adapter's final-reply marker) satisfy
the caller; progress/status/preview sends must not."""
if not (metadata or {}).get("notify"):
logger.debug("A2A: ignoring non-final send for context %s", chat_id)
return SendResult(success=True, message_id=message_id)
if not self._resolve_oldest_for_context(chat_id, protocol.STATE_COMPLETED, content or ""):
elif not self._resolve_oldest_for_context(chat_id, protocol.STATE_COMPLETED, content or ""):
logger.debug("A2A: send() for context %s had no pending waiter", chat_id) # late chunk / out-of-band
return SendResult(success=True, message_id=message_id)
return SendResult(success=True, message_id=str(int(time.time() * 1000)))
async def send_typing(self, chat_id: str, metadata=None) -> None:
return None
@@ -912,13 +783,11 @@ class A2AAdapter(BasePlatformAdapter):
return {"name": f"a2a:{chat_id}", "type": "dm"}
async def on_processing_complete(self, event: MessageEvent, outcome: ProcessingOutcome) -> None:
"""Resolve the task future when processing ends without a reply send
(failures, cancellations, empty runs) so the HTTP thread returns promptly."""
"""Resolve the task future when processing ends without a reply send (failures,
cancellations, empty runs) so the HTTP thread returns promptly."""
task_id = str(getattr(event, "message_id", "") or "")
if not task_id:
return
state, text = {
ProcessingOutcome.FAILURE: (protocol.STATE_FAILED, "[agent processing failed]"),
ProcessingOutcome.CANCELLED: (protocol.STATE_CANCELED, ""),
}.get(outcome, (protocol.STATE_COMPLETED, ""))
self._resolve_task(task_id, state, text)
if task_id:
self._resolve_task(task_id, *{
ProcessingOutcome.FAILURE: (protocol.STATE_FAILED, "[agent processing failed]"),
ProcessingOutcome.CANCELLED: (protocol.STATE_CANCELED, ""),
}.get(outcome, (protocol.STATE_COMPLETED, "")))
+102 -228
View File
@@ -1,17 +1,11 @@
"""
A2A protocol helpers — Agent Card, JSON-RPC framing, task store, conversation persistence.
Wire shape is A2A v1.0 (JSON-RPC 2.0 over HTTP): SCREAMING_SNAKE_CASE states/roles;
Parts and StreamResponse events (``statusUpdate`` / ``artifactUpdate``) are
discriminated by member presence (no ``kind`` / ``final`` fields); SSE stream
closure signals the terminal state; push configs carry ``configId`` + ``createdAt``.
Stdlib only (no a2a-sdk). ``extract_text`` stays tolerant of v0.3 peers.
"""
"""A2A protocol helpers — Agent Card, JSON-RPC framing, task store, conversation persistence.
Wire shape is A2A v1.0: SCREAMING_SNAKE_CASE states/roles; Parts and StreamResponse events are
discriminated by member presence (no ``kind``/``final``); SSE closure signals the terminal state.
Stdlib only. ``extract_text`` stays tolerant of v0.3 peers."""
from __future__ import annotations
import json
import copy
import os
import threading
import time
@@ -22,52 +16,33 @@ from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Optional
from gateway.platforms._shared import coerce_port as _coerce_int
PROTOCOL_VERSION = "1.0"
# A2A v1.0 task lifecycle states.
STATE_SUBMITTED = "TASK_STATE_SUBMITTED"
STATE_WORKING = "TASK_STATE_WORKING"
STATE_INPUT_REQUIRED = "TASK_STATE_INPUT_REQUIRED"
STATE_AUTH_REQUIRED = "TASK_STATE_AUTH_REQUIRED"
STATE_COMPLETED = "TASK_STATE_COMPLETED"
STATE_FAILED = "TASK_STATE_FAILED"
STATE_CANCELED = "TASK_STATE_CANCELED"
STATE_REJECTED = "TASK_STATE_REJECTED"
# A2A v1.0 task lifecycle states + message roles.
STATE_SUBMITTED, STATE_WORKING, STATE_INPUT_REQUIRED = "TASK_STATE_SUBMITTED", "TASK_STATE_WORKING", "TASK_STATE_INPUT_REQUIRED"
STATE_COMPLETED, STATE_FAILED = "TASK_STATE_COMPLETED", "TASK_STATE_FAILED"
STATE_CANCELED, STATE_REJECTED = "TASK_STATE_CANCELED", "TASK_STATE_REJECTED"
TERMINAL_STATES = frozenset({STATE_COMPLETED, STATE_FAILED, STATE_CANCELED, STATE_REJECTED})
ROLE_USER, ROLE_AGENT = "ROLE_USER", "ROLE_AGENT"
# A2A v1.0 message roles.
ROLE_USER = "ROLE_USER"
ROLE_AGENT = "ROLE_AGENT"
# The agent starts its reply with this marker when it needs clarification; the
# adapter maps such replies to TASK_STATE_INPUT_REQUIRED (marker stripped).
# A reply starting with this marker is a clarification request -> TASK_STATE_INPUT_REQUIRED (marker stripped).
INPUT_REQUIRED_MARKER = "[INPUT_REQUIRED]"
# JSON-RPC / A2A error codes. -32001..-32003 are A2A spec-defined; custom errors
# live at -32050..-32059 (implementation-defined space, clear of the A2A block).
ERR_PARSE = -32700
ERR_INVALID_PARAMS = -32602
ERR_METHOD_NOT_FOUND = -32601
ERR_TASK_NOT_FOUND = -32001 # A2A spec: TaskNotFoundError
ERR_TASK_NOT_CANCELABLE = -32002 # A2A spec: TaskNotCancelableError
ERR_UNAUTHORIZED = -32050
ERR_RATE_LIMITED = -32051
ERR_UNTRUSTED_PEER = -32052
ERR_PARSE, ERR_INVALID_PARAMS, ERR_METHOD_NOT_FOUND = -32700, -32602, -32601
ERR_TASK_NOT_FOUND, ERR_TASK_NOT_CANCELABLE = -32001, -32002 # A2A spec: TaskNotFoundError / TaskNotCancelableError
ERR_UNAUTHORIZED, ERR_RATE_LIMITED, ERR_UNTRUSTED_PEER = -32050, -32051, -32052
# Anti-loop: max inbound turns per context. A2A_MAX_PINGPONG_TURNS env, capped at 20.
_DEFAULT_MAX_PINGPONG = 5
_HARD_MAX_PINGPONG = 20
_RATE_LIMIT_DEFAULT = 60 # requests per minute
_RATE_WINDOW = 60.0 # seconds
_DEFAULT_MAX_PINGPONG, _HARD_MAX_PINGPONG = 5, 20
_RATE_LIMIT_DEFAULT, _RATE_WINDOW = 60, 60.0 # requests per minute, window seconds
def _env_int(name: str, default: int) -> int:
try:
return int(os.getenv(name, str(default)))
except (ValueError, TypeError):
return default
return _coerce_int(os.getenv(name, default), default)
def max_pingpong_turns() -> int:
@@ -75,10 +50,6 @@ def max_pingpong_turns() -> int:
return max(1, min(v, _HARD_MAX_PINGPONG))
def _rate_limit_per_minute() -> int:
return max(1, _env_int("A2A_RATE_LIMIT", _RATE_LIMIT_DEFAULT))
def now_iso() -> str:
"""ISO 8601 UTC timestamp with millisecond precision (A2A v1.0)."""
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
@@ -92,19 +63,12 @@ def _hermes_home() -> Path:
return Path(os.path.expanduser("~/.hermes"))
# ── Agent Card (v1.0) ─────────────────────────────────────────────────────────
def build_agent_card(*, name: str, url: str, description: str, skills: Optional[list[dict]] = None,
streaming: bool = False, push_notifications: bool = False, auth_required: bool = False,
tenant: str = "") -> dict:
"""Construct an A2A v1.0 Agent Card.
``tenant`` is the optional v1.0 multi-tenancy routing key on AgentInterface;
when present, clients MUST echo it in request params.
"""
iface: dict[str, Any] = {"url": url, "protocolBinding": "JSONRPC", "protocolVersion": PROTOCOL_VERSION}
if tenant:
iface["tenant"] = tenant
"""A2A v1.0 Agent Card. ``tenant`` is the optional multi-tenancy routing key on
AgentInterface; when present, clients MUST echo it in request params."""
iface: dict[str, Any] = {"url": url, "protocolBinding": "JSONRPC", "protocolVersion": PROTOCOL_VERSION, **({"tenant": tenant} if tenant else {})}
card: dict[str, Any] = {
"name": name,
"description": description,
@@ -114,9 +78,7 @@ def build_agent_card(*, name: str, url: str, description: str, skills: Optional[
"supportedInterfaces": [iface],
"capabilities": {"streaming": streaming, "pushNotifications": push_notifications,
"stateTransitionHistory": False, "extendedAgentCard": False},
"defaultInputModes": ["text/plain"],
"defaultOutputModes": ["text/plain"],
"skills": skills or [],
"defaultInputModes": ["text/plain"], "defaultOutputModes": ["text/plain"], "skills": skills or [],
}
if auth_required:
card["securitySchemes"] = {"bearer": {"type": "http", "scheme": "bearer"}}
@@ -125,20 +87,15 @@ def build_agent_card(*, name: str, url: str, description: str, skills: Optional[
def skills_from_toolsets(toolsets: "list[str] | dict[str, list[str]] | None") -> list[dict]:
"""Derive A2A skill descriptors from toolset names, or a toolset → tool-names
mapping (tool names become tags, max 10, so peers can match tasks to us)."""
"""A2A skill descriptors from toolset names or a toolset -> tool-names mapping (tool names
become tags, max 10)."""
if not isinstance(toolsets, dict):
toolsets = {ts: [] for ts in set(toolsets or [])}
skills = [
{"id": f"toolset.{name}", "name": name, "description": f"Hermes '{name}' capabilities",
"tags": [name] + [str(t) for t in (toolsets[name] or [])][:10]}
for name in sorted(toolsets)
]
skills = [{"id": f"toolset.{name}", "name": name, "description": f"Hermes '{name}' capabilities",
"tags": [name] + [str(t) for t in (toolsets[name] or [])][:10]} for name in sorted(toolsets)]
return skills or [{"id": "general", "name": "general", "description": "General-purpose conversational agent", "tags": ["general"]}]
# ── JSON-RPC framing + message / part builders ────────────────────────────────
def jsonrpc_result(req_id: Any, result: Any) -> dict:
return {"jsonrpc": "2.0", "id": req_id, "result": result}
@@ -148,15 +105,14 @@ def jsonrpc_error(req_id: Any, code: int, message: str) -> dict:
def send_message_response(payload: dict) -> dict:
"""A2A v1.0 SendMessageResponse oneof wrapper: exactly one of ``task`` /
``message``. Legacy methods still return bare payloads."""
"""v1.0 SendMessageResponse oneof: exactly one of ``task`` / ``message``."""
if isinstance(payload, dict) and payload.get("status") and payload.get("id"):
return {"task": payload}
return {"message": payload}
def unwrap_send_message_response(result: Any) -> Any:
"""Return the Task/Message inside a v1.0 response, or pass legacy through."""
"""Task/Message inside a v1.0 response; legacy bare payloads pass through."""
if isinstance(result, dict):
if isinstance(result.get("task"), dict):
return result["task"]
@@ -183,37 +139,14 @@ def text_part(text: str) -> dict:
return {"text": text, "mediaType": "text/plain"}
def file_part(url: str = "", raw: str = "", filename: str = "",
media_type: str = "application/octet-stream") -> dict:
"""v1.0 file Part: ``url`` (reference) or ``raw`` (base64 bytes)."""
part: dict[str, Any] = {"mediaType": media_type}
if filename:
part["filename"] = filename
if url:
part["url"] = url
elif raw:
part["raw"] = raw
return part
def data_part(data: Any, media_type: str = "application/json") -> dict:
"""v1.0 data Part (structured data, no ``kind`` field)."""
return {"data": data, "mediaType": media_type}
def message_with_parts(role: str, parts: list[dict], context_id: str = "") -> dict:
"""A2A v1.0 Message with arbitrary Parts (text, file, data)."""
msg: dict[str, Any] = {"role": role, "parts": parts, "messageId": uuid.uuid4().hex}
def text_message(role: str, text: str, context_id: str = "") -> dict:
"""A2A v1.0 Message with a single text Part."""
msg: dict[str, Any] = {"role": role, "parts": [text_part(text)], "messageId": uuid.uuid4().hex}
if context_id:
msg["contextId"] = context_id
return msg
def text_message(role: str, text: str, context_id: str = "") -> dict:
"""A2A v1.0 Message with a single text Part."""
return message_with_parts(role, [text_part(text)], context_id)
def _file_note(fname: str, body: str, mtype: str) -> str:
label = f"[file: {fname}]" if fname else "[file]"
return f"{label} {body}" + (f" ({mtype})" if mtype else "")
@@ -227,15 +160,12 @@ def _json_or_str(data: Any) -> str:
def extract_text(message_or_params: dict) -> str:
"""Concatenated text from an A2A Message / Task-result / params payload.
v1.0, v0.3 (``kind``) and pre-0.3 (``type``) Parts all carry a string ``text``
member. File Parts render as URL/filename (raw base64 noted, not decoded);
data Parts render their JSON so the agent sees them."""
"""Concatenated text from an A2A Message / Task-result / params payload. v1.0, v0.3
(``kind``) and pre-0.3 (``type``) Parts all carry ``text``; file Parts render as
URL/filename (raw base64 noted, not decoded); data Parts render their JSON."""
msg = message_or_params.get("message", message_or_params)
parts = msg.get("parts", []) if isinstance(msg, dict) else []
chunks = []
for part in parts:
for part in msg.get("parts", []) if isinstance(msg, dict) else []:
if not isinstance(part, dict):
continue
if isinstance(txt := part.get("text"), str):
@@ -256,13 +186,12 @@ def extract_text(message_or_params: dict) -> str:
def extract_context_id(params: dict) -> str:
"""v1.0 puts contextId inside the Message; tolerate legacy top-level."""
msg = params.get("message") or {}
ctx = str(msg.get("contextId") or "") if isinstance(msg, dict) else ""
return ctx or str(params.get("contextId") or "")
return (str(msg.get("contextId") or "") if isinstance(msg, dict) else "") or str(params.get("contextId") or "")
def build_task(task_id: str, context_id: str, state: str, agent_text: str = "", *, created_at: str = "") -> dict:
"""A2A v1.0 Task object. ``created_at`` is accepted but NOT serialized: the v1.0
Task proto has no createdAt and strict ProtoJSON parsers (a2a-sdk) reject unknown fields."""
"""A2A v1.0 Task. ``created_at`` is accepted but NOT serialized: the v1.0 Task proto has no
createdAt and strict ProtoJSON parsers (a2a-sdk) reject unknown fields."""
task: dict[str, Any] = {"id": task_id, "contextId": context_id, "status": {"state": state, "timestamp": now_iso()}}
if agent_text:
task["status"]["message"] = text_message(ROLE_AGENT, agent_text, context_id)
@@ -271,8 +200,6 @@ def build_task(task_id: str, context_id: str, state: str, agent_text: str = "",
return task
# ── Streaming (v1.0 StreamResponse events) ────────────────────────────────────
def status_update(task_id: str, context_id: str, state: str, text: str = "") -> dict:
"""v1.0 StreamResponse with a statusUpdate member."""
status: dict[str, Any] = {"state": state, "timestamp": now_iso()}
@@ -288,46 +215,38 @@ def artifact_update(task_id: str, context_id: str, text: str) -> dict:
def sse_data(payload: dict, req_id: Any = None) -> str:
"""One StreamResponse as an SSE data frame. §9.4 requires a full JSON-RPC envelope
(a2a-sdk breaks on bare StreamResponses); ``req_id=None`` is the legacy no-envelope fallback."""
"""One StreamResponse as an SSE data frame. §9.4 requires a full JSON-RPC envelope (a2a-sdk
breaks on bare StreamResponses); ``req_id=None`` is the legacy no-envelope fallback."""
envelope = jsonrpc_result(req_id, payload) if req_id is not None else payload
return f"data: {json.dumps(envelope, ensure_ascii=False)}\n\n"
def sse_done() -> str:
"""Stream-closure marker as an SSE *comment* — ``data: {}`` would make
JSON-RPC clients try to parse an empty response."""
"""Stream-closure marker as an SSE *comment* — ``data: {}`` would make JSON-RPC clients parse."""
return ": done\n\n"
# ── Anti-loop ping-pong protection (per-adapter instance) ─────────────────────
class TurnTracker:
"""Counts inbound turns per context_id; beyond max_pingpong_turns() the
adapter rejects further messages for that context."""
"""Counts inbound turns per context_id; beyond max_pingpong_turns() the adapter rejects."""
_TTL = 3600 # prune contexts idle longer than 1 hour
def __init__(self) -> None:
self._counts: dict[str, int] = defaultdict(int)
self._timestamps: dict[str, float] = {}
self._turns: dict[str, tuple[int, float]] = {} # context_id -> (count, last_seen)
self._lock = threading.Lock()
def track(self, context_id: str) -> int:
"""Increment and return the turn count; prunes stale contexts."""
with self._lock:
now = time.time()
for cid in [cid for cid, ts in self._timestamps.items() if now - ts > self._TTL]:
self._counts.pop(cid, None)
self._timestamps.pop(cid, None)
self._counts[context_id] += 1
self._timestamps[context_id] = now
return self._counts[context_id]
self._turns = {cid: v for cid, v in self._turns.items() if now - v[1] <= self._TTL}
count = self._turns.get(context_id, (0, now))[0] + 1
self._turns[context_id] = (count, now)
return count
def reset(self, context_id: str) -> None:
with self._lock:
self._counts.pop(context_id, None)
self._timestamps.pop(context_id, None)
self._turns.pop(context_id, None)
class RateLimiter:
@@ -339,7 +258,7 @@ class RateLimiter:
def allow(self, identity: str) -> bool:
with self._lock:
limit = _rate_limit_per_minute()
limit = max(1, _env_int("A2A_RATE_LIMIT", _RATE_LIMIT_DEFAULT))
now = time.time()
bucket = self._buckets[identity]
while bucket and now - bucket[0] > _RATE_WINDOW:
@@ -350,18 +269,16 @@ class RateLimiter:
return True
# ── Metrics collection ────────────────────────────────────────────────────────
class Metrics:
"""Simple counters for A2A operations (module singleton ``metrics`` is shared
by the inbound adapter and outbound tools; not persisted)."""
"""Counters for A2A operations (module singleton ``metrics`` shared by the inbound adapter
and outbound tools; not persisted)."""
_COUNTERS = ("inbound_total", "outbound_total", "streams_started", "push_sent", "push_failed",
"tasks_completed", "tasks_failed", "anti_loop_triggers", "rate_limit_triggers")
def __init__(self) -> None:
self.inbound_total = self.outbound_total = self.streams_started = self.push_sent = self.push_failed = 0
self.tasks_completed = self.tasks_failed = self.anti_loop_triggers = self.rate_limit_triggers = 0
for name in self._COUNTERS:
setattr(self, name, 0)
self._start_time = time.time()
self._latencies: deque[float] = deque(maxlen=100) # last 100 completed inbound tasks
@@ -372,19 +289,16 @@ class Metrics:
return sum(self._latencies) / len(self._latencies) if self._latencies else 0.0
def snapshot(self) -> dict[str, Any]:
return {"uptime_seconds": round(time.time() - self._start_time, 1),
**{name: getattr(self, name) for name in self._COUNTERS},
return {"uptime_seconds": round(time.time() - self._start_time, 1), **{n: getattr(self, n) for n in self._COUNTERS},
"avg_latency_ms": round(self.avg_latency() * 1000, 1)}
metrics = Metrics()
# ── Task store — pending AND completed tasks (queryable via tasks/get, tasks/list) ───
class TaskStore:
"""In-memory store of A2A tasks, kept after completion for tasks/get. Records carry
agent slug + tenant; readers pass a scope and get not-found outside it (spec authz rule)."""
"""In-memory A2A tasks, kept after completion for tasks/get. Records carry agent slug +
tenant; readers pass a scope and get not-found outside it (spec authz rule)."""
_MAX_TERMINAL = 500
@@ -395,9 +309,7 @@ class TaskStore:
@staticmethod
def _in_scope(rec: dict, agent_slug: str = "", tenant: str = "") -> bool:
if agent_slug and rec.get("agent_slug", "") != agent_slug:
return False
return not (tenant and rec.get("tenant", "") != tenant)
return not ((agent_slug and rec.get("agent_slug", "") != agent_slug) or (tenant and rec.get("tenant", "") != tenant))
def _scoped(self, task_id: str, agent_slug: str = "", tenant: str = "") -> Optional[dict]:
"""Live record if visible in scope. Caller holds the lock."""
@@ -407,11 +319,9 @@ class TaskStore:
def _push_rec(self, task_id: str, config_id: str = "", agent_slug: str = "", tenant: str = "") -> Optional[dict]:
"""Scoped record that has a push config (matching ``config_id`` if given). Caller holds the lock."""
rec = self._scoped(task_id, agent_slug, tenant)
if not rec or not rec.get("push_url"):
return None
if config_id and rec.get("push_config_id") != config_id:
return None
return rec
if rec and rec.get("push_url") and (not config_id or rec.get("push_config_id") == config_id):
return rec
return None
@staticmethod
def _push_config_view(rec: dict) -> dict:
@@ -419,62 +329,50 @@ class TaskStore:
"createdAt": rec.get("created_iso", ""), "pushNotificationConfig": {"url": rec.get("push_url") or ""}}
def create(self, task_id: str, context_id: str, peer: str, agent_slug: str = "", tenant: str = "") -> dict:
rec = {
"task_id": task_id, "context_id": context_id, "peer": peer,
"agent_slug": agent_slug or "", "tenant": tenant or "", "state": STATE_SUBMITTED, "reply": "",
"created_at": time.time(), "created_iso": now_iso(), "push_url": "", "push_config_id": "",
}
rec = {"task_id": task_id, "context_id": context_id, "peer": peer, "agent_slug": agent_slug or "", "tenant": tenant or "",
"state": STATE_SUBMITTED, "reply": "", "created_at": time.time(), "created_iso": now_iso(), "push_url": "", "push_config_id": ""}
with self._lock:
self._tasks[task_id] = rec
return dict(rec)
def set_state(self, task_id: str, state: str) -> None:
with self._lock:
rec = self._tasks.get(task_id)
if rec and rec["state"] not in TERMINAL_STATES:
if (rec := self._tasks.get(task_id)) and rec["state"] not in TERMINAL_STATES:
rec["state"] = state
def set_push_config(self, task_id: str, url: str, agent_slug: str = "", tenant: str = "") -> Optional[dict]:
"""Attach a push notification config; returns the stored config or None."""
with self._lock:
rec = self._scoped(task_id, agent_slug, tenant)
if not rec:
if not (rec := self._scoped(task_id, agent_slug, tenant)):
return None
rec["push_url"] = url
rec["push_config_id"] = "cfg-" + uuid.uuid4().hex[:12]
rec["push_url"], rec["push_config_id"] = url, "cfg-" + uuid.uuid4().hex[:12]
return self._push_config_view(rec)
def get_push_config(self, task_id: str, config_id: str = "", agent_slug: str = "", tenant: str = "") -> Optional[dict]:
with self._lock:
rec = self._push_rec(task_id, config_id, agent_slug, tenant)
return self._push_config_view(rec) if rec else None
return self._push_config_view(rec) if (rec := self._push_rec(task_id, config_id, agent_slug, tenant)) else None
def list_push_configs(self, task_id: str, agent_slug: str = "", tenant: str = "") -> list[dict]:
with self._lock:
rec = self._push_rec(task_id, "", agent_slug, tenant)
return [self._push_config_view(rec)] if rec else []
cfg = self.get_push_config(task_id, "", agent_slug, tenant)
return [cfg] if cfg else []
def delete_push_config(self, task_id: str, config_id: str = "", agent_slug: str = "", tenant: str = "") -> bool:
with self._lock:
rec = self._push_rec(task_id, config_id, agent_slug, tenant)
if not rec:
return False
rec["push_url"] = ""
rec["push_config_id"] = ""
return True
if rec:
rec["push_url"] = rec["push_config_id"] = ""
return rec is not None
def pop_push_url(self, task_id: str) -> str:
with self._lock:
rec = self._tasks.get(task_id)
if not rec:
return ""
url, rec["push_url"] = rec["push_url"], ""
return url
if rec:
url, rec["push_url"] = rec["push_url"], ""
return url if rec else ""
def get(self, task_id: str, agent_slug: str = "", tenant: str = "") -> Optional[dict]:
with self._lock:
rec = self._scoped(task_id, agent_slug, tenant)
return dict(rec) if rec else None
return dict(rec) if (rec := self._scoped(task_id, agent_slug, tenant)) else None
def complete(self, task_id: str, state: str, reply: str = "") -> Optional[dict]:
"""Transition a task to a terminal state. Idempotent."""
@@ -482,9 +380,7 @@ class TaskStore:
rec = self._tasks.get(task_id)
if not rec or rec["state"] in TERMINAL_STATES:
return None
rec["state"] = state
rec["reply"] = reply
rec["completed_at"] = time.time()
rec.update(state=state, reply=reply, completed_at=time.time())
watchers = self._watchers.pop(task_id, [])
self._trim_locked()
out = dict(rec)
@@ -495,8 +391,7 @@ class TaskStore:
def watch(self, task_id: str, agent_slug: str = "", tenant: str = "") -> Optional[Future]:
with self._lock:
rec = self._scoped(task_id, agent_slug, tenant)
if not rec:
if not (rec := self._scoped(task_id, agent_slug, tenant)):
return None
fut: Future = Future()
if rec["state"] in TERMINAL_STATES:
@@ -511,25 +406,18 @@ class TaskStore:
``(records, next_offset, total)`` with ``with_total`` (v1.0 ListTasks totalSize)."""
page_size = max(1, min(int(page_size or 50), 100))
with self._lock:
recs = [dict(r) for r in reversed(self._tasks.values())]
if agent_slug or tenant:
recs = [r for r in recs if self._in_scope(r, agent_slug, tenant)]
if context_id:
recs = [r for r in recs if r["context_id"] == context_id]
if state:
recs = [r for r in recs if r["state"] == state]
recs = [dict(r) for r in reversed(self._tasks.values())
if self._in_scope(r, agent_slug, tenant)
and (not context_id or r["context_id"] == context_id) and (not state or r["state"] == state)]
total = len(recs)
page = recs[offset:offset + page_size]
next_offset = offset + page_size if offset + page_size < total else 0
if with_total:
return page, next_offset, total
return page, next_offset
return (page, next_offset, total) if with_total else (page, next_offset)
def fail_orphans(self, timeout_seconds: int = 300) -> list[str]:
with self._lock:
now = time.time()
stale = [tid for tid, rec in self._tasks.items()
if rec["state"] not in TERMINAL_STATES and now - rec["created_at"] > timeout_seconds]
if rec["state"] not in TERMINAL_STATES and time.time() - rec["created_at"] > timeout_seconds]
return [tid for tid in stale if self.complete(tid, STATE_FAILED, "[task orphaned — no reply produced]")]
def _trim_locked(self) -> None:
@@ -538,61 +426,47 @@ class TaskStore:
self._tasks.pop(tid, None)
@staticmethod
def to_task(rec: dict, history_length: Optional[int] = None, include_artifacts: bool = True) -> dict:
"""Render a stored record as an A2A v1.0 Task object."""
def to_task(rec: dict, include_artifacts: bool = True) -> dict:
"""Render a stored record as an A2A v1.0 Task."""
task = build_task(rec["task_id"], rec["context_id"], rec["state"], rec.get("reply", ""),
created_at=rec.get("created_iso", ""))
if not include_artifacts:
task.pop("artifacts", None)
if history_length == 0:
task.pop("history", None)
return copy.deepcopy(task)
# ── Conversation persistence (outside the context-compaction pipeline) ────────
def _conv_dir() -> Path:
return _hermes_home() / "a2a_conversations"
return task
def _conv_path(context_id: str) -> Path:
safe = "".join(c for c in (context_id or "default") if c.isalnum() or c in "-_") or "default"
return _conv_dir() / f"{safe}.jsonl"
return _hermes_home() / "a2a_conversations" / f"{safe}.jsonl"
def persist_message(context_id: str, role: str, text: str, task_id: str = "") -> None:
"""Append one message to the context's on-disk conversation log."""
"""Append one message to the context's on-disk conversation log. Never raises."""
try:
_conv_dir().mkdir(parents=True, exist_ok=True)
rec = {"ts": time.time(), "role": role, "text": text, "task_id": task_id}
with _conv_path(context_id).open("a", encoding="utf-8") as fh:
fh.write(json.dumps(rec, ensure_ascii=False) + "\n")
path = _conv_path(context_id)
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as fh:
fh.write(json.dumps({"ts": time.time(), "role": role, "text": text, "task_id": task_id}, ensure_ascii=False) + "\n")
except Exception:
pass
def load_conversation(context_id: str, limit: int = 50) -> list[dict]:
"""Load the last *limit* messages for a context (empty list if none)."""
path = _conv_path(context_id)
if not path.exists():
return []
out: list[dict] = []
"""Last *limit* messages for a context (empty list if none / unreadable)."""
try:
with path.open("r", encoding="utf-8") as fh:
for line in fh:
if line.strip():
try:
out.append(json.loads(line))
except json.JSONDecodeError:
pass
lines = _conv_path(context_id).read_text(encoding="utf-8").splitlines()
except Exception:
return []
out: list[dict] = []
for line in lines:
if line.strip():
try:
out.append(json.loads(line))
except json.JSONDecodeError:
pass
return out[-limit:]
def list_conversations() -> list[str]:
"""Return known context-ids that have persisted conversations."""
d = _conv_dir()
if not d.exists():
return []
return sorted(p.stem for p in d.glob("*.jsonl"))
"""Context-ids that have persisted conversations."""
return sorted(p.stem for p in (_hermes_home() / "a2a_conversations").glob("*.jsonl"))
+21 -71
View File
@@ -1,13 +1,7 @@
"""
A2A security primitives — shared by the inbound adapter and the client tools.
A2A is a *network* surface (adversarial peers in, private context out). Layers,
opt-out only by explicit config: bind safety (no token => 127.0.0.1 only); peer
identity (A2A_PEER_TOKENS token->name, shared A2A_BEARER_TOKEN => ip:<addr>; rate
limiting and trust key on this, never on the body); inbound injection filtering;
outbound credential redaction; JSONL audit log; optional trusted-peer allow-list;
HMAC-SHA256 push signing + SSRF-safe callback URLs.
"""
"""A2A security primitives (adapter + client tools). A2A is a *network* surface: bind safety (no
token => 127.0.0.1 only); peer identity from credentials, never the body (A2A_PEER_TOKENS
token->name, shared A2A_BEARER_TOKEN => ip:<addr>); inbound injection filtering; outbound
credential redaction; JSONL audit; trusted-peer allow-list; HMAC push signing; SSRF-safe URLs."""
from __future__ import annotations
@@ -28,11 +22,8 @@ logger = logging.getLogger(__name__)
def _startup_env(name: str) -> str:
"""Read one A2A setting from the active profile's scope, else the env.
Inside a secondary profile's scope the scope is authoritative: a miss yields
"" and never falls through to ``os.environ`` (the default profile's tokens).
"""
"""One A2A setting from the active profile's scope, else the env. Inside a secondary
profile's scope a miss yields "" and never falls through to the default profile's env."""
if _profile_scoped():
from agent.secret_scope import get_secret
return (get_secret(name) or "").strip()
@@ -41,14 +32,8 @@ def _startup_env(name: str) -> str:
def _parse_peer_tokens(raw: str) -> dict[str, str]:
""""alice:tok1,bob:tok2" -> {token: peer_name}."""
out: dict[str, str] = {}
for pair in raw.split(","):
if ":" not in pair:
continue
name, token = (s.strip() for s in pair.split(":", 1))
if name and token:
out[token] = name
return out
pairs = [tuple(s.strip() for s in pair.split(":", 1)) for pair in raw.split(",") if ":" in pair]
return {token: name for name, token in pairs if name and token}
def _configured_trusted_peers() -> frozenset[str]:
@@ -80,13 +65,10 @@ class A2ASecurityContext:
@classmethod
def capture(cls) -> "A2ASecurityContext":
bearer_token = _startup_env("A2A_BEARER_TOKEN")
return cls(
bearer_token=bearer_token,
peer_tokens=tuple(_parse_peer_tokens(_startup_env("A2A_PEER_TOKENS")).items()),
trusted_peers=_configured_trusted_peers(),
allow_all_users=_startup_env("A2A_ALLOW_ALL_USERS").lower() in {"1", "true", "yes"},
requested_host=_startup_env("A2A_HOST") or "127.0.0.1",
push_secret=_startup_env("A2A_PUSH_SECRET") or bearer_token)
return cls(bearer_token=bearer_token, peer_tokens=tuple(_parse_peer_tokens(_startup_env("A2A_PEER_TOKENS")).items()),
trusted_peers=_configured_trusted_peers(),
allow_all_users=_startup_env("A2A_ALLOW_ALL_USERS").lower() in {"1", "true", "yes"},
requested_host=_startup_env("A2A_HOST") or "127.0.0.1", push_secret=_startup_env("A2A_PUSH_SECRET") or bearer_token)
def localhost_only(self) -> bool:
return not (self.bearer_token or self.peer_tokens)
@@ -131,29 +113,11 @@ class A2ASecurityContext:
return hmac.new(self.push_secret.encode("utf-8"), body, hashlib.sha256).hexdigest()
# Module-level conveniences: capture a fresh context from the current scope.
def authenticate(auth_header: Optional[str], client_ip: str = "") -> Optional[str]:
return A2ASecurityContext.capture().authenticate(auth_header, client_ip)
def localhost_only() -> bool:
"""Fresh-context convenience for callers outside the adapter."""
return A2ASecurityContext.capture().localhost_only()
def resolve_bind_host() -> str:
return A2ASecurityContext.capture().resolve_bind_host()
def is_trusted_peer(identity: str) -> bool:
return A2ASecurityContext.capture().is_trusted_peer(identity)
def sign_push_payload(payload: dict) -> str:
return A2ASecurityContext.capture().sign_push_payload(payload)
# ── Inbound injection filtering ───────────────────────────────────────────────
# Neutralise (don't reject) so a task that merely *mentions* these still gets through.
_INJECTION_PATTERNS: tuple[re.Pattern[str], ...] = (
re.compile(r"<\|im_(start|end)\|>", re.IGNORECASE),
@@ -190,9 +154,7 @@ _REDACTION_PATTERNS: tuple[tuple[re.Pattern[str], str], ...] = (
def filter_inbound(text: str) -> str:
"""Defang prompt-injection markers in inbound task text."""
if not text:
return text
for pat in _INJECTION_PATTERNS:
for pat in _INJECTION_PATTERNS if text else ():
text = pat.sub("[filtered]", text)
return text
@@ -205,18 +167,13 @@ def wrap_inbound(peer: str, text: str) -> str:
def redact_outbound(text: str) -> str:
"""Scrub credential-shaped substrings before sending text to a peer."""
if not text:
return text
for pat, repl in _REDACTION_PATTERNS:
for pat, repl in _REDACTION_PATTERNS if text else ():
text = pat.sub(repl, text)
return text
# ── SSRF protection for push notification callback URLs ───────────────────────
# Blocked even in localhost-only mode — a remote peer must not make us probe internal
# services: link-local/AWS metadata, loopback, RFC1918, unspecified, IPv6 loopback/link-local/ULA.
# Loopback is allowed only in localhost mode (local testing).
# Blocked even in localhost-only mode — a remote peer must not make us probe internal services
# (link-local/AWS metadata, RFC1918, unspecified, IPv6 link-local/ULA). Loopback only in localhost mode.
_BLOCKED_PREFIXES = ("169.254.", "127.", "10.", *(f"172.{i}." for i in range(16, 32)), "192.168.",
"0.0.0.0", "::1", "fe80:", "fc00:", "fd00:")
@@ -225,15 +182,11 @@ def is_safe_callback_url(url: str, *, localhost_mode: Optional[bool] = None) ->
"""True when a push callback URL is http(s) and not internal/private/loopback."""
if localhost_mode is None:
localhost_mode = localhost_only()
if not url or not isinstance(url, str):
return False
try:
parsed = urllib.parse.urlparse(url)
parsed = urllib.parse.urlparse(url) if url and isinstance(url, str) else None
except Exception:
return False
if parsed.scheme not in ("http", "https"):
return False
hostname = parsed.hostname or ""
hostname = (parsed.hostname or "") if parsed and parsed.scheme in ("http", "https") else ""
if not hostname:
return False
hostname_lower = hostname.lower()
@@ -251,16 +204,13 @@ def is_safe_callback_url(url: str, *, localhost_mode: Optional[bool] = None) ->
return True
# ── Audit log ─────────────────────────────────────────────────────────────────
def audit(direction: str, peer: str, task_id: str, summary: str) -> None:
"""Append an audit record (direction: inbound | outbound | push). Never raises."""
try:
from .protocol import _hermes_home
rec = {"ts": time.time(), "direction": direction, "peer": peer, "task_id": task_id, "summary": (summary or "")[:500]}
path = _hermes_home() / "a2a_audit.jsonl"
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as fh:
_hermes_home().mkdir(parents=True, exist_ok=True)
with (_hermes_home() / "a2a_audit.jsonl").open("a", encoding="utf-8") as fh:
fh.write(json.dumps(rec, ensure_ascii=False) + "\n")
except Exception:
logger.debug("A2A: audit write failed", exc_info=True)
+75 -157
View File
@@ -1,22 +1,10 @@
"""
A2A client tools (``a2a`` toolset) — let the Hermes agent talk to *other* agents:
a2a_discover, a2a_call, a2a_list, a2a_history, a2a_orchestrate.
Peers are resolved from config.yaml under ``a2a_agents``::
a2a_agents:
researcher:
url: "http://localhost:9999"
auth: { type: bearer, token: "sk-..." }
timeout: 120
capabilities: [web_search, research]
Transport is stdlib urllib (no a2a-sdk). Wire format is A2A v1.0 JSON-RPC
``SendMessage``; replies from v0.3 peers still parse.
"""
"""A2A client tools (``a2a`` toolset): a2a_discover/call/list/history/orchestrate talk to *other*
agents. Peers come from config.yaml ``a2a_agents: {name: {url, auth: {type: bearer, token}, timeout,
capabilities}}``. Stdlib urllib; wire format is A2A v1.0 ``SendMessage`` (v0.3 replies still parse)."""
from __future__ import annotations
import contextlib
import json
import logging
import os
@@ -25,6 +13,8 @@ import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Any, Optional
from gateway.platforms._shared import coerce_port as _coerce_int
from . import protocol, security
logger = logging.getLogger(__name__)
@@ -33,8 +23,6 @@ _DEFAULT_TIMEOUT = 120
_ORCHESTRATE_MAX_WORKERS = 6 # max parallel peers for fan-out
# ── Peer resolution ───────────────────────────────────────────────────────────
def _load_config() -> dict:
try:
from hermes_cli.config import load_config
@@ -53,23 +41,17 @@ def _peer_from_entry(entry: dict, **extra: Any) -> dict:
def _resolve_peer(agent: str) -> Optional[dict]:
"""Resolve a peer name to {url, auth, timeout, capabilities, tenant}, or treat ``agent`` as a URL."""
if agent.startswith("http://") or agent.startswith("https://"):
"""Peer name -> {url, auth, timeout, capabilities, tenant}, or treat ``agent`` as a URL."""
if agent.startswith(("http://", "https://")):
return {"url": agent, "auth": {}, "timeout": _DEFAULT_TIMEOUT, "capabilities": []}
entry = _configured_peers().get(agent)
if not entry:
return None
return _peer_from_entry(entry, capabilities=entry.get("capabilities", []) or [], tenant=entry.get("tenant", ""))
return _peer_from_entry(entry, capabilities=entry.get("capabilities", []) or [], tenant=entry.get("tenant", "")) if entry else None
def _auth_header(auth: dict) -> dict:
if auth and auth.get("type") == "bearer" and auth.get("token"):
return {"Authorization": f"Bearer {auth['token']}"}
return {}
return {"Authorization": f"Bearer {auth['token']}"} if auth and auth.get("type") == "bearer" and auth.get("token") else {}
# ── HTTP + Agent Card discovery ───────────────────────────────────────────────
def _http_json(url: str, headers: dict, timeout: int, method: str, data: Optional[bytes] = None) -> dict:
req = urllib.request.Request(url, data=data, headers=headers, method=method)
with urllib.request.urlopen(req, timeout=timeout) as resp: # noqa: S310 (configured peers)
@@ -106,36 +88,28 @@ def _select_jsonrpc_interface(card: Optional[dict]) -> Optional[dict]:
def _rpc_url(base_url: str, card: Optional[dict]) -> str:
"""Card's JSONRPC interface (v1.0 supportedInterfaces) > card's legacy top-level url > base."""
iface = _select_jsonrpc_interface(card)
if iface:
if iface := _select_jsonrpc_interface(card):
return str(iface["url"])
if isinstance(card, dict) and isinstance(card.get("url"), str) and card["url"]:
return card["url"]
return base_url.rstrip("/")
# ── Shared send path (used by a2a_call and a2a_orchestrate) ───────────────────
def _send_task(agent_label: str, peer: dict, message: str, context_id: str) -> tuple[str, str, str]:
"""Send one SendMessage to a peer -> (reply_text, context_id, state). Raises urllib errors /
"""One SendMessage to a peer -> (reply_text, context_id, state). Raises urllib errors /
ValueError for the caller to format; handles redaction, audit, persistence, metrics."""
base_url = peer.get("url", "")
headers = _auth_header(peer.get("auth", {}) or {})
timeout = int(peer.get("timeout", _DEFAULT_TIMEOUT))
card = None # best-effort, to learn the rpc URL
try:
card = _fetch_card(base_url, headers, min(timeout, 30))
card = _fetch_card(base_url, headers, min(timeout, 30)) # best-effort, to learn the rpc URL
except Exception:
pass
card = None
ctx = context_id or protocol.new_context_id()
safe_message = security.redact_outbound(message)
# v1.0: contextId lives inside the Message, not at the params top level.
rpc_body = {
"jsonrpc": "2.0",
"id": protocol.new_task_id(),
"method": "SendMessage",
"params": {"message": protocol.text_message(protocol.ROLE_USER, safe_message, context_id=ctx)},
}
rpc_body = {"jsonrpc": "2.0", "id": protocol.new_task_id(), "method": "SendMessage",
"params": {"message": protocol.text_message(protocol.ROLE_USER, safe_message, context_id=ctx)}}
iface = _select_jsonrpc_interface(card)
tenant = str(iface["tenant"]) if iface and iface.get("tenant") else str(peer.get("tenant") or "")
if tenant:
@@ -145,8 +119,7 @@ def _send_task(agent_label: str, peer: dict, message: str, context_id: str) -> t
protocol.metrics.outbound_total += 1
resp = _http_post_json(_rpc_url(base_url, card), rpc_body, headers, timeout)
if "error" in resp:
err = resp["error"]
raise ValueError(f"Peer '{agent_label}' returned an error: {err.get('message', err)}")
raise ValueError(f"Peer '{agent_label}' returned an error: {resp['error'].get('message', resp['error'])}")
payload = protocol.unwrap_send_message_response(resp.get("result", {}))
reply = _reply_text_from_result(payload)
reply_ctx, state = ctx, ""
@@ -162,18 +135,16 @@ def _reply_text_from_result(result: Any) -> str:
result = protocol.unwrap_send_message_response(result)
if not isinstance(result, dict):
return str(result)
# Artifacts first (final output), then status message (interim/clarify).
# Artifacts first (final output), then status message (interim/clarify), else bare Message.
for artifact in result.get("artifacts", []) or []:
txt = protocol.extract_text(artifact)
if txt:
return txt
msg = (result.get("status", {}) or {}).get("message")
if msg:
return protocol.extract_text(msg)
return protocol.extract_text(result) # bare Message result
return protocol.extract_text((result.get("status", {}) or {}).get("message") or result)
# ── Tool handlers ─────────────────────────────────────────────────────────────
_AUTH_ERR = "Error: peer '{agent}' rejected auth (HTTP {code}). Check the configured token."
_HTTP_CALL_ERRORS = {401: _AUTH_ERR, 403: _AUTH_ERR, 429: "Error: peer '{agent}' rate limited us (HTTP 429). Retry later."}
def a2a_discover(args: dict, **_: Any) -> str:
"""Fetch and summarize the Agent Card at ``url``."""
@@ -193,14 +164,10 @@ def a2a_discover(args: dict, **_: Any) -> str:
f"{i.get('protocolBinding', '?')} v{i.get('protocolVersion', '?')}"
for i in (card.get("supportedInterfaces", []) or []) if isinstance(i, dict)
) or f"v{card.get('protocolVersion', '?')} (pre-1.0 card)"
lines = [
f"Agent: {card.get('name', '?')}",
f"Description: {card.get('description', '')}",
f"URL: {_rpc_url(url, card)}",
f"Protocol: {proto}",
f"Streaming: {bool(caps.get('streaming'))} Push: {bool(caps.get('pushNotifications'))} Auth required: {auth}",
f"Skills ({len(skills)}):",
]
lines = [f"Agent: {card.get('name', '?')}", f"Description: {card.get('description', '')}", f"URL: {_rpc_url(url, card)}",
f"Protocol: {proto}",
f"Streaming: {bool(caps.get('streaming'))} Push: {bool(caps.get('pushNotifications'))} Auth required: {auth}",
f"Skills ({len(skills)}):"]
lines.extend(f" - {s.get('name', s.get('id', '?'))}: {s.get('description', '')}" for s in skills[:20])
return "\n".join(lines)
@@ -219,11 +186,7 @@ def a2a_call(args: dict, **_: Any) -> str:
try:
reply, reply_ctx, state = _send_task(agent, peer, message, context_id)
except urllib.error.HTTPError as e:
if e.code in (401, 403):
return f"Error: peer '{agent}' rejected auth (HTTP {e.code}). Check the configured token."
if e.code == 429:
return f"Error: peer '{agent}' rate limited us (HTTP 429). Retry later."
return f"Error: call to '{agent}' failed — HTTP {e.code}."
return _HTTP_CALL_ERRORS.get(e.code, "Error: call to '{agent}' failed — HTTP {code}.").format(agent=agent, code=e.code)
except ValueError as e:
return str(e)
except Exception as e:
@@ -243,14 +206,12 @@ def a2a_list(args: dict | None = None, **_: Any) -> str:
if peers:
lines.append(f"Configured peers ({len(peers)}):")
for name, entry in peers.items():
auth = (entry.get("auth") or {}).get("type", "none")
caps = entry.get("capabilities", [])
cap_str = f" caps: {', '.join(caps)}" if caps else ""
lines.append(f" - {name}: {entry.get('url', '?')} (auth: {auth}){cap_str}")
lines.append(f" - {name}: {entry.get('url', '?')} (auth: {(entry.get('auth') or {}).get('type', 'none')})"
+ (f" caps: {', '.join(caps)}" if caps else ""))
else:
lines.append("No peers configured. Add them under 'a2a_agents' in config.yaml.")
convos = protocol.list_conversations()
if convos:
if convos := protocol.list_conversations():
lines.append("")
lines.append(f"Persisted conversations ({len(convos)}) — recall with a2a_history:")
lines.extend(f" - {c}" for c in convos[:25])
@@ -263,39 +224,29 @@ def a2a_list(args: dict | None = None, **_: Any) -> str:
def a2a_history(args: dict, **_: Any) -> str:
"""Recall a persisted A2A conversation (~/.hermes/a2a_conversations/<context>.jsonl)
— how prior exchanges survive compaction/restarts."""
"""Recall a persisted A2A conversation (survives compaction/restarts)."""
context_id = str(args.get("context_id") or args.get("contextId") or "").strip()
if not context_id:
return "Error: 'context_id' is required (see a2a_list for known conversations)."
try:
limit = max(1, min(int(args.get("limit") or 50), 200))
except (ValueError, TypeError):
limit = 50
limit = max(1, min(_coerce_int(args.get("limit") or 50, 50), 200))
messages = protocol.load_conversation(context_id, limit=limit)
if not messages:
return f"No persisted conversation for context '{context_id}'."
lines = [f"Conversation {context_id} (last {len(messages)} messages):"]
for m in messages:
text = (m.get("text") or "").strip()
if len(text) > 1000:
text = text[:1000] + " …[truncated]"
lines.append(f"[{m.get('role', '?')}] {text}")
lines.append(f"[{m.get('role', '?')}] {text[:1000] + ' …[truncated]' if len(text) > 1000 else text}")
return "\n".join(lines)
# ── a2a_orchestrate: capability-based routing with fan-out ────────────────────
def _match_peers_by_capability(capability: str) -> list[tuple[str, dict]]:
"""Configured peers that advertise the capability ('*' matches all)."""
return [
(name, entry) for name, entry in _configured_peers().items()
if capability in (entry.get("capabilities", []) or []) or capability == "*"
]
return [(name, entry) for name, entry in _configured_peers().items()
if capability in (entry.get("capabilities", []) or []) or capability == "*"]
def _call_peer_sync(agent_name: str, peer_entry: dict, message: str, context_id: str = "") -> tuple[str, str]:
"""Call a single peer synchronously. Returns (agent_name, reply_text)."""
"""Call a single peer synchronously -> (agent_name, reply_text)."""
try:
reply, _ctx, _state = _send_task(agent_name, _peer_from_entry(peer_entry), message, context_id)
return (agent_name, reply or "(no reply)")
@@ -304,21 +255,19 @@ def _call_peer_sync(agent_name: str, peer_entry: dict, message: str, context_id:
def a2a_orchestrate(args: dict, **_: Any) -> str:
"""Fan-out a task to peers matching a capability (``a2a_agents.*.capabilities``). Modes: ``all``,
``first`` (first successful), ``best`` (longest successful — coarse; use ``all`` to judge yourself)."""
"""Fan-out a task to peers matching a capability. Modes: ``all``, ``first`` (first successful),
``best`` (longest successful — coarse; use ``all`` to judge yourself)."""
capability = str(args.get("capability") or "").strip()
message = str(args.get("message") or args.get("task") or "").strip()
mode = str(args.get("mode") or "all").strip().lower()
mode = mode if mode in ("all", "first", "best") else "all"
context_id = str(args.get("context_id") or "").strip()
if not message:
return "Error: 'message' is required."
if not capability:
return "Error: 'capability' is required (or use '*' for all peers)."
matches = _match_peers_by_capability(capability)
if not matches:
if not (matches := _match_peers_by_capability(capability)):
return f"Error: no configured peers advertise capability '{capability}'."
if mode not in ("all", "first", "best"):
mode = "all"
results: list[tuple[str, str]] = []
with ThreadPoolExecutor(max_workers=min(len(matches), _ORCHESTRATE_MAX_WORKERS)) as pool:
futures = {pool.submit(_call_peer_sync, name, entry, message, context_id): name for name, entry in matches}
@@ -339,87 +288,62 @@ def a2a_orchestrate(args: dict, **_: Any) -> str:
return "\n".join(["All peers failed:"] + [f" {name}: {reply}" for name, reply in results])
name, reply = max(successes, key=lambda r: len(r[1])) if mode == "best" else successes[0]
return f"[{mode}: {name}]\n{reply}"
lines = [f"Orchestrated '{capability}' to {len(matches)} peer(s):"]
for name, reply in results:
lines.append(f"\n--- {name} ---")
lines.append(reply)
return "\n".join(lines)
return "\n".join([f"Orchestrated '{capability}' to {len(matches)} peer(s):"]
+ [line for name, reply in results for line in (f"\n--- {name} ---", reply)])
# ── Tool schemas + registration ───────────────────────────────────────────────
def _str(description: str) -> dict:
return {"type": "string", "description": description}
# name -> (handler, description, properties, required)
_TOOLS: dict[str, tuple[Any, str, dict, list[str]]] = {
"a2a_discover": (
a2a_discover,
"Fetch and summarize another agent's A2A Agent Card from a URL (its name, description, "
"capabilities, and skills). Use this to find out what a remote agent can do before calling it.",
{"url": _str("Base URL of the remote A2A agent, e.g. http://localhost:9999")},
["url"],
),
"a2a_call": (
a2a_call,
"Send a natural-language task to a remote A2A agent and return its reply. The agent is a peer "
"(any A2A-compliant framework), not a sub-agent you control. Pass 'context_id' from a previous "
"reply to continue a multi-turn exchange.",
{
"agent": _str("Configured peer name (from a2a_agents) or a full http(s):// URL."),
"message": _str("The task / message to send the peer, in natural language."),
"context_id": _str("Optional: context id from a prior reply, to continue the conversation."),
},
["agent", "message"],
),
"a2a_discover": (a2a_discover,
"Fetch and summarize another agent's A2A Agent Card from a URL (its name, description, "
"capabilities, and skills). Use this to find out what a remote agent can do before calling it.",
{"url": _str("Base URL of the remote A2A agent, e.g. http://localhost:9999")}, ["url"]),
"a2a_call": (a2a_call,
"Send a natural-language task to a remote A2A agent and return its reply. The agent is a peer "
"(any A2A-compliant framework), not a sub-agent you control. Pass 'context_id' from a previous "
"reply to continue a multi-turn exchange.",
{"agent": _str("Configured peer name (from a2a_agents) or a full http(s):// URL."),
"message": _str("The task / message to send the peer, in natural language."),
"context_id": _str("Optional: context id from a prior reply, to continue the conversation.")},
["agent", "message"]),
"a2a_list": (a2a_list, "List configured A2A peer agents, persisted A2A conversations, and metrics.", {}, []),
"a2a_history": (
a2a_history,
"Recall a persisted A2A conversation transcript by context_id (survives restarts and "
"context compaction). Use a2a_list to see known context ids.",
{
"context_id": _str("Context id of the conversation to recall."),
"limit": {"type": "integer", "description": "Max messages to return (default 50, max 200)."},
},
["context_id"],
),
"a2a_orchestrate": (
a2a_orchestrate,
"Fan-out a task to multiple peer agents by capability. Peers are matched from config.yaml "
"a2a_agents.*.capabilities. Modes: 'all' (return all replies), 'first' (first successful), "
"'best' (longest successful reply).",
{
"capability": _str("Capability to match (e.g. 'research', 'code') or '*' for all peers."),
"message": _str("The task to send to all matching peers."),
"mode": {"type": "string", "enum": ["all", "first", "best"], "description": "How to aggregate results. Default: 'all'."},
"context_id": _str("Optional: shared context id for all peers."),
},
["capability", "message"],
),
"a2a_history": (a2a_history,
"Recall a persisted A2A conversation transcript by context_id (survives restarts and "
"context compaction). Use a2a_list to see known context ids.",
{"context_id": _str("Context id of the conversation to recall."),
"limit": {"type": "integer", "description": "Max messages to return (default 50, max 200)."}},
["context_id"]),
"a2a_orchestrate": (a2a_orchestrate,
"Fan-out a task to multiple peer agents by capability. Peers are matched from config.yaml "
"a2a_agents.*.capabilities. Modes: 'all' (return all replies), 'first' (first successful), "
"'best' (longest successful reply).",
{"capability": _str("Capability to match (e.g. 'research', 'code') or '*' for all peers."),
"message": _str("The task to send to all matching peers."),
"mode": {"type": "string", "enum": ["all", "first", "best"], "description": "How to aggregate results. Default: 'all'."},
"context_id": _str("Optional: shared context id for all peers.")},
["capability", "message"]),
}
def _a2a_tools_available() -> bool:
"""check_fn: serve the client tools ONLY when the operator opted into A2A (peers under
``a2a_agents``, inbound platform enabled, or A2A_PORT set). Unconditional registration cost
every session ~561 tok/call for tools that can only say 'no peers configured'. Fail closed."""
``a2a_agents``, inbound platform enabled, or A2A_PORT set). Fail closed."""
cfg = {}
try:
with contextlib.suppress(Exception):
cfg = _load_config()
if cfg.get("a2a_agents"):
return True
except Exception: # noqa: BLE001
pass
try:
if os.getenv("A2A_PORT"):
return True
a2a_cfg = (cfg.get("platforms") or {}).get("a2a") or {}
if isinstance(a2a_cfg, dict) and a2a_cfg.get("enabled"):
return True
return bool(isinstance(a2a_cfg, dict) and a2a_cfg.get("enabled"))
except Exception: # noqa: BLE001
pass
return False
return False
def register_tools(ctx) -> None:
@@ -428,12 +352,6 @@ def register_tools(ctx) -> None:
parameters: dict[str, Any] = {"type": "object", "properties": properties}
if required:
parameters["required"] = required
ctx.register_tool(
name=name,
toolset="a2a",
schema={"name": name, "description": description, "parameters": parameters},
handler=handler,
description=description,
emoji="\U0001f9e9", # puzzle piece
check_fn=_a2a_tools_available,
)
ctx.register_tool(name=name, toolset="a2a", handler=handler, description=description,
schema={"name": name, "description": description, "parameters": parameters},
emoji="\U0001f9e9", check_fn=_a2a_tools_available) # puzzle piece
File diff suppressed because it is too large Load Diff
+26 -35
View File
@@ -1,4 +1,4 @@
"""Dependency-free Nostr signing for Buzz WebSocket authentication."""
"""Dependency-free Nostr signing (secp256k1 / BIP-340) for Buzz WebSocket authentication."""
from __future__ import annotations
@@ -16,17 +16,17 @@ GENERATOR = (
0x483ADA7726A3C4655DA4FBFC0E1108A8FD17B448A68554199C47D08FFB10D4B8,
)
BECH32_CHARSET = "qpzry9x8gf2tvdw0s3jn54khce6mua7l"
_BECH32_GENERATORS = (0x3B6A57B2, 0x26508E6D, 0x1EA119FA, 0x3D4233DD, 0x2A1462B3)
Point = Optional[tuple[int, int]]
def _bech32_polymod(values: list[int]) -> int:
generators = (0x3B6A57B2, 0x26508E6D, 0x1EA119FA, 0x3D4233DD, 0x2A1462B3)
checksum = 1
for value in values:
top = checksum >> 25
checksum = ((checksum & 0x1FFFFFF) << 5) ^ value
for index, generator in enumerate(generators):
for index, generator in enumerate(_BECH32_GENERATORS):
if (top >> index) & 1:
checksum ^= generator
return checksum
@@ -52,8 +52,7 @@ def _decode_nsec(value: str) -> bytes:
raise ValueError("invalid character in nsec") from exc
if _bech32_polymod(_bech32_hrp_expand(hrp) + data) != 1:
raise ValueError("invalid nsec checksum")
accumulator = 0
bits = 0
accumulator = bits = 0
decoded = bytearray()
for value5 in data[:-6]:
accumulator = (accumulator << 5) | value5
@@ -86,10 +85,8 @@ def decode_private_key(value: str) -> int:
def _point_add(left: Point, right: Point) -> Point:
if left is None:
return right
if right is None:
return left
if left is None or right is None:
return right if left is None else left
x1, y1 = left
x2, y2 = right
if x1 == x2:
@@ -100,8 +97,7 @@ def _point_add(left: Point, right: Point) -> Point:
slope = (y2 - y1) * pow(x2 - x1, FIELD_ORDER - 2, FIELD_ORDER)
slope %= FIELD_ORDER
x3 = (slope * slope - x1 - x2) % FIELD_ORDER
y3 = (slope * (x1 - x3) - y1) % FIELD_ORDER
return x3, y3
return x3, (slope * (x1 - x3) - y1) % FIELD_ORDER
def _point_multiply(scalar: int, point: Point = GENERATOR) -> Point:
@@ -127,9 +123,7 @@ def public_key_hex(private_key: str) -> str:
return point[0].to_bytes(32, "big").hex()
def schnorr_sign(
message: bytes, private_key: str, *, auxiliary_randomness: Optional[bytes] = None,
) -> bytes:
def schnorr_sign(message: bytes, private_key: str, *, auxiliary_randomness: Optional[bytes] = None) -> bytes:
if len(message) != 32:
raise ValueError("BIP-340 signs a 32-byte message")
secret = decode_private_key(private_key)
@@ -138,7 +132,7 @@ def schnorr_sign(
raise ValueError("invalid private key")
public_x = public_point[0].to_bytes(32, "big")
adjusted_secret = secret if public_point[1] % 2 == 0 else CURVE_ORDER - secret
aux = (auxiliary_randomness if auxiliary_randomness is not None else secrets.token_bytes(32))
aux = auxiliary_randomness if auxiliary_randomness is not None else secrets.token_bytes(32)
if len(aux) != 32:
raise ValueError("auxiliary randomness must be 32 bytes")
masked = (adjusted_secret ^ int.from_bytes(_tagged_hash("BIP0340/aux", aux), "big")).to_bytes(32, "big")
@@ -148,12 +142,22 @@ def schnorr_sign(
nonce_point = _point_multiply(nonce)
if nonce_point is None: # pragma: no cover
raise RuntimeError("BIP-340 produced an invalid nonce point")
adjusted_nonce = nonce if nonce_point[1] % 2 == 0 else CURVE_ORDER - nonce
nonce_x = nonce_point[0].to_bytes(32, "big")
challenge_hash = _tagged_hash("BIP0340/challenge", nonce_x + public_x + message)
challenge = int.from_bytes(challenge_hash, "big") % CURVE_ORDER
signature_scalar = (adjusted_nonce + challenge * adjusted_secret) % CURVE_ORDER
return nonce_x + signature_scalar.to_bytes(32, "big")
adjusted_nonce = nonce if nonce_point[1] % 2 == 0 else CURVE_ORDER - nonce
challenge = int.from_bytes(_tagged_hash("BIP0340/challenge", nonce_x + public_x + message), "big") % CURVE_ORDER
return nonce_x + ((adjusted_nonce + challenge * adjusted_secret) % CURVE_ORDER).to_bytes(32, "big")
def parse_auth_tag(raw: Any, label: str) -> list[str]:
"""Validate a NIP-OA owner-attestation tag (JSON text or list) -> ``["auth", a, b, c]``."""
if isinstance(raw, str):
try:
raw = json.loads(raw)
except json.JSONDecodeError as exc:
raise ValueError(f"{label} is not valid JSON") from exc
if not isinstance(raw, list) or len(raw) != 4 or raw[0] != "auth" or not all(isinstance(p, str) for p in raw):
raise ValueError(f"{label} must be a four-string auth tag")
return raw
def build_auth_event(
@@ -162,25 +166,12 @@ def build_auth_event(
) -> dict[str, Any]:
tags: list[list[str]] = [["relay", relay_url], ["challenge", challenge]]
if auth_tag_json.strip():
try:
auth_tag = json.loads(auth_tag_json)
except json.JSONDecodeError as exc:
raise ValueError("BUZZ_AUTH_TAG is not valid JSON") from exc
if not isinstance(auth_tag, list) or len(auth_tag) != 4 or auth_tag[0] != "auth" or not all(
isinstance(part, str) for part in auth_tag
):
raise ValueError("BUZZ_AUTH_TAG must be a four-string auth tag")
tags.append(auth_tag)
tags.append(parse_auth_tag(auth_tag_json, "BUZZ_AUTH_TAG"))
pubkey = public_key_hex(private_key)
timestamp = int(time.time()) if created_at is None else int(created_at)
serialized = json.dumps([0, pubkey, timestamp, 22242, tags, ""], separators=(",", ":"), ensure_ascii=False).encode()
event_id = hashlib.sha256(serialized).digest()
return {
"id": event_id.hex(),
"pubkey": pubkey,
"created_at": timestamp,
"kind": 22242,
"tags": tags,
"content": "",
"id": event_id.hex(), "pubkey": pubkey, "created_at": timestamp, "kind": 22242, "tags": tags, "content": "",
"sig": schnorr_sign(event_id, private_key, auxiliary_randomness=auxiliary_randomness).hex(),
}
File diff suppressed because it is too large Load Diff
+56 -127
View File
@@ -19,11 +19,7 @@ EXT_MAP = {
}
# rich-text runtime type → (media_types entry, MessageType promotion when still TEXT)
_RICH_MEDIA = {
"image": ("image", MessageType.PHOTO),
"video": ("video", MessageType.VIDEO),
"file": ("application/octet-stream", MessageType.DOCUMENT),
}
_RICH_MEDIA = {"image": ("image", MessageType.PHOTO), "video": ("video", MessageType.VIDEO), "file": ("application/octet-stream", MessageType.DOCUMENT)}
def _extensions(message: Any) -> Any:
@@ -45,46 +41,33 @@ def _rich_list(message: Any) -> Optional[list]:
return rich_list if isinstance(rich_list, list) else None
def _url_of(raw: Any) -> str:
return (raw.get("url", "") or raw.get("docUrl", "")) if isinstance(raw, dict) else ""
def _card_text(message: Any) -> str:
"""msgtype='card' (钉钉文档分享卡片 / link card): title + doc URL from ``extensions['card']``."""
extensions = _extensions(message)
content = ""
card = extensions.get("card", {})
content = ""
if isinstance(card, dict):
title = card.get("title", "")
raw_content = card.get("content", "")
doc_url = ""
if isinstance(raw_content, dict):
doc_url = raw_content.get("url", "") or raw_content.get("docUrl", "")
elif isinstance(raw_content, str) and raw_content.strip():
title, raw_content = card.get("title", ""), card.get("content", "")
doc_url = _url_of(raw_content)
if isinstance(raw_content, str) and raw_content.strip():
try:
parsed = json.loads(raw_content.strip())
if isinstance(parsed, dict):
doc_url = parsed.get("url", "") or parsed.get("docUrl", "")
doc_url = _url_of(json.loads(raw_content.strip()))
except (ValueError, TypeError):
doc_url = raw_content
parts = ([f"[文档] {title}"] if title else []) + ([doc_url] if doc_url else [])
if parts:
content = " ".join(parts)
if not content:
# Last-resort: raw text field from extensions (if present)
ext_text = extensions.get("text", {})
if isinstance(ext_text, dict):
content = (ext_text.get("content", "") or "").strip()
return content
content = " ".join(([f"[文档] {title}"] if title else []) + ([doc_url] if doc_url else []))
ext_text = extensions.get("text", {}) # last-resort: raw text field from extensions
return content or ((ext_text.get("content", "") or "").strip() if isinstance(ext_text, dict) else "")
def _interactive_card_text(message: Any) -> str:
"""msgtype='interactiveCard': ``extensions['content']`` carries title + biz_custom_action_url."""
ext_content = _ext_content(message)
if not ext_content:
return ""
doc_url = ext_content.get("biz_custom_action_url", "")
title = ext_content.get("title", "")
if not (doc_url or title):
return ""
parts = [f"[文档卡片] {title}" if title else "[文档卡片]"] + ([doc_url] if doc_url else [])
return " ".join(parts)
ext_content = _ext_content(message) or {}
doc_url, title = ext_content.get("biz_custom_action_url", ""), ext_content.get("title", "")
return " ".join([f"[文档卡片] {title}" if title else "[文档卡片]"] + ([doc_url] if doc_url else [])) if (doc_url or title) else ""
def _ext_field(message: Any, field: str) -> Any:
@@ -92,65 +75,34 @@ def _ext_field(message: Any, field: str) -> Any:
return ext_content.get(field, "") if ext_content else ""
def _audio_text(message: Any) -> str:
"""msgtype='audio': DingTalk-provided speech recognition text."""
recognition = _ext_field(message, "recognition")
return recognition.strip() if recognition else ""
# Fallbacks by msgtype when no plain/rich text was found (types are exclusive): audio -> DingTalk speech
# recognition text; file -> fileName; card / interactiveCard -> title + doc URL.
_EMPTY_TEXT_FALLBACKS = {
"audio": lambda m: (_ext_field(m, "recognition") or "").strip(),
"file": lambda m: f"[文件] {_ext_field(m, 'fileName')}" if _ext_field(m, "fileName") else "",
"card": _card_text,
"interactiveCard": _interactive_card_text,
}
def _file_text(message: Any) -> str:
"""msgtype='file': use fileName as text."""
fname = _ext_field(message, "fileName")
return f"[文件] {fname}" if fname else ""
# Fallbacks by msgtype when no plain/rich text was found (types are exclusive).
_EMPTY_TEXT_FALLBACKS = (
("audio", _audio_text),
("file", _file_text),
("card", _card_text),
("interactiveCard", _interactive_card_text),
)
def _rich_item_text(item: Any) -> str:
if isinstance(item, dict):
return item.get("text") or item.get("content") or ""
return getattr(item, "text", "") or ""
def extract_text(message: Any) -> str:
"""Extract plain text from a DingTalk chatbot message.
Handles both SDK payload shapes: legacy ``message.text`` dict ``{"content": ...}`` and
>= 0.20 ``TextContent`` (whose ``__str__`` is ``"TextContent(content=...)"`` — always read
``.content`` first); rich text via ``rich_text_content.rich_text_list`` or legacy ``rich_text``.
"""
"""Extract plain text from a DingTalk chatbot message. Handles both SDK shapes: legacy ``message.text`` dict ``{"content": ...}`` and >= 0.20 ``TextContent``
(whose ``__str__`` is ``"TextContent(content=...)"`` — always read ``.content`` first)."""
text = getattr(message, "text", None) or ""
if hasattr(text, "content"):
content = (text.content or "").strip()
elif isinstance(text, dict):
content = text.get("content", "").strip()
else:
content = str(text).strip()
if not content:
rich_list = _rich_list(message)
if rich_list is not None:
parts = []
for item in rich_list:
if isinstance(item, dict):
t = item.get("text") or item.get("content") or ""
if t:
parts.append(t)
elif hasattr(item, "text") and item.text:
parts.append(item.text)
content = " ".join(parts).strip()
if not content:
msg_type = getattr(message, "message_type", "")
for kind, fallback in _EMPTY_TEXT_FALLBACKS:
if msg_type == kind:
content = fallback(message)
break
content = (text.content or "" if hasattr(text, "content") else text.get("content", "") if isinstance(text, dict) else str(text)).strip()
rich_list = _rich_list(message) if not content else None
if rich_list is not None:
content = " ".join(t for t in map(_rich_item_text, rich_list) if t).strip()
fallback = _EMPTY_TEXT_FALLBACKS.get(getattr(message, "message_type", "")) if not content else None
# Do NOT strip "@bot": the mention is routed structurally (callback ``isInAtList``), and
# regex-stripping @handles would damage e-mails, SSH URLs and literal "@openai" references.
return content
return fallback(message) if fallback else content
def extract_media(message: Any) -> Tuple[MessageType, List[str], List[str]]:
@@ -158,35 +110,25 @@ def extract_media(message: Any) -> Tuple[MessageType, List[str], List[str]]:
msg_type = MessageType.TEXT
media_urls: List[str] = []
media_types: List[str] = []
image_content = getattr(message, "image_content", None)
if image_content:
download_code = getattr(image_content, "download_code", None)
if download_code:
media_urls.append(download_code)
media_types.append("image")
msg_type = MessageType.PHOTO
if download_code := (getattr(image_content, "download_code", None) if image_content else None):
media_urls.append(download_code)
media_types.append("image")
msg_type = MessageType.PHOTO
for item in _rich_list(message) or ():
if not isinstance(item, dict):
continue
dl_code = item.get("downloadCode") or item.get("download_code") or ""
item_type = item.get("type", "")
if not dl_code:
continue
item_type = item.get("type", "")
mapped = DINGTALK_TYPE_MAPPING.get(item_type, "file")
# "voice" items are native voice notes → STT (VOICE); "audio" file uploads stay AUDIO.
mime, promoted = ("audio", MessageType.VOICE if item_type == "voice" else MessageType.AUDIO) if mapped == "audio" else _RICH_MEDIA[mapped]
media_urls.append(dl_code)
if mapped == "audio":
media_types.append("audio")
if msg_type == MessageType.TEXT:
# "voice" items are native voice notes → STT (VOICE); "audio" file uploads stay AUDIO.
msg_type = MessageType.VOICE if item_type == "voice" else MessageType.AUDIO
else:
mime, promoted = _RICH_MEDIA[mapped]
media_types.append(mime)
if msg_type == MessageType.TEXT:
msg_type = promoted
media_types.append(mime)
if msg_type == MessageType.TEXT:
msg_type = promoted
msg_type_str = getattr(message, "message_type", "") or ""
if msg_type_str == "picture" and not media_urls:
msg_type = MessageType.PHOTO
@@ -201,24 +143,14 @@ def extract_media(message: Any) -> Tuple[MessageType, List[str], List[str]]:
if msg_type == MessageType.TEXT:
msg_type = MessageType.VOICE
elif msg_type_str in ("file", "image"):
ext_content = _ext_content(message)
if ext_content:
dl_code = ext_content.get("downloadCode") or ""
ext_content = _ext_content(message) or {}
if dl_code := ext_content.get("downloadCode") or "":
fname = ext_content.get("fileName", "")
if dl_code:
media_urls.append(dl_code)
mime = "application/octet-stream"
if fname:
ext = fname.rsplit(".", 1)[-1].lower() if "." in fname else ""
mime = EXT_MAP.get(ext, mime)
media_types.append(mime)
if msg_type == MessageType.TEXT:
# Image messages, and files with image MIME (a .png sent as attachment), → PHOTO.
if msg_type_str == "image" or mime.startswith("image/"):
msg_type = MessageType.PHOTO
else:
msg_type = MessageType.DOCUMENT
mime = EXT_MAP.get(fname.rsplit(".", 1)[-1].lower() if fname and "." in fname else "", "application/octet-stream")
media_urls.append(dl_code)
media_types.append(mime)
if msg_type == MessageType.TEXT: # image messages, and files with image MIME (a .png attachment) → PHOTO
msg_type = MessageType.PHOTO if (msg_type_str == "image" or mime.startswith("image/")) else MessageType.DOCUMENT
return msg_type, media_urls, media_types
@@ -229,12 +161,9 @@ def collect_download_codes(message: Any) -> List[Tuple[Any, str]]:
if img_content and getattr(img_content, "download_code", None):
codes.append((img_content, "download_code"))
rich_text = getattr(message, "rich_text_content", None)
if rich_text:
for item in getattr(rich_text, "rich_text_list", []) or []:
if isinstance(item, dict):
for key in ("downloadCode", "pictureDownloadCode", "download_code"):
if item.get(key):
codes.append((item, key))
for item in (getattr(rich_text, "rich_text_list", []) or []) if rich_text else []:
if isinstance(item, dict):
codes.extend((item, key) for key in ("downloadCode", "pictureDownloadCode", "download_code") if item.get(key))
if (getattr(message, "message_type", "") or "") in ("file", "image"):
ext_content = _ext_content(message)
if ext_content and ext_content.get("downloadCode"):
+9 -17
View File
@@ -1,11 +1,6 @@
"""Shared ffmpeg executable discovery for Discord voice paths.
Discovery itself is owned by ``tools.transcription_tools`` (the same helper
the STT pipeline uses — PATH plus common Homebrew/local prefixes); this module
only layers the Discord-voice-specific extras on top: an explicit
``FFMPEG_PATH`` override and a Windows winget fallback for installs that
never touch PATH.
"""
"""ffmpeg discovery for Discord voice: ``tools.transcription_tools`` owns the shared lookup
(PATH + Homebrew/local prefixes); this layers an explicit ``FFMPEG_PATH`` override and a
Windows winget fallback (installs that never touch PATH) on top."""
from __future__ import annotations
@@ -25,16 +20,13 @@ def _shared_find_ffmpeg():
def resolve_ffmpeg_executable() -> str:
"""Return an ffmpeg command that also covers common Windows installs."""
explicit = os.getenv("FFMPEG_PATH")
if explicit and explicit.strip():
return os.path.expandvars(os.path.expanduser(explicit.strip()))
discovered = _shared_find_ffmpeg()
if discovered:
explicit = (os.getenv("FFMPEG_PATH") or "").strip()
if explicit:
return os.path.expandvars(os.path.expanduser(explicit))
if discovered := _shared_find_ffmpeg():
return discovered
local_appdata = os.getenv("LOCALAPPDATA")
if local_appdata:
packages_dir = Path(local_appdata) / "Microsoft" / "WinGet" / "Packages"
candidates = sorted(packages_dir.glob("Gyan.FFmpeg_*/*/bin/ffmpeg.exe"))
if local_appdata := os.getenv("LOCALAPPDATA"):
candidates = sorted((Path(local_appdata) / "Microsoft" / "WinGet" / "Packages").glob("Gyan.FFmpeg_*/*/bin/ffmpeg.exe"))
if candidates:
return str(candidates[-1])
return "ffmpeg"
+15 -41
View File
@@ -55,52 +55,26 @@ class DiscordRecoveryStore:
def _initialize(self, conn: sqlite3.Connection) -> None:
from hermes_state import apply_wal_with_fallback
apply_wal_with_fallback(conn, db_label="discord_recovery.db")
conn.execute("""
conn.executescript("""
CREATE TABLE IF NOT EXISTS discord_messages (
message_id TEXT PRIMARY KEY,
channel_id TEXT,
thread_id TEXT,
parent_channel_id TEXT,
author_id TEXT,
created_at TEXT,
status TEXT NOT NULL,
replied INTEGER NOT NULL DEFAULT 0,
emoji_ack INTEGER NOT NULL DEFAULT 0,
outage_response INTEGER NOT NULL DEFAULT 0,
response_message_id TEXT,
attempts INTEGER NOT NULL DEFAULT 0,
last_attempt_at TEXT,
last_error TEXT,
message_id TEXT PRIMARY KEY, channel_id TEXT, thread_id TEXT, parent_channel_id TEXT,
author_id TEXT, created_at TEXT, status TEXT NOT NULL,
replied INTEGER NOT NULL DEFAULT 0, emoji_ack INTEGER NOT NULL DEFAULT 0,
outage_response INTEGER NOT NULL DEFAULT 0, response_message_id TEXT,
attempts INTEGER NOT NULL DEFAULT 0, last_attempt_at TEXT, last_error TEXT,
updated_at TEXT NOT NULL
)
""")
conn.execute("""
);
CREATE TABLE IF NOT EXISTS discord_recovery_scans (
scan_id TEXT PRIMARY KEY,
started_at TEXT NOT NULL,
completed_at TEXT,
status TEXT NOT NULL,
channels TEXT NOT NULL,
window_seconds REAL NOT NULL,
limit_count INTEGER NOT NULL,
scanned INTEGER NOT NULL DEFAULT 0,
missed INTEGER NOT NULL DEFAULT 0,
dispatched INTEGER NOT NULL DEFAULT 0,
error TEXT
)
""")
conn.execute("""
scan_id TEXT PRIMARY KEY, started_at TEXT NOT NULL, completed_at TEXT, status TEXT NOT NULL,
channels TEXT NOT NULL, window_seconds REAL NOT NULL, limit_count INTEGER NOT NULL,
scanned INTEGER NOT NULL DEFAULT 0, missed INTEGER NOT NULL DEFAULT 0,
dispatched INTEGER NOT NULL DEFAULT 0, error TEXT
);
CREATE TABLE IF NOT EXISTS discord_recovery_cursors (
channel_id TEXT PRIMARY KEY,
last_message_id TEXT NOT NULL,
updated_at TEXT NOT NULL
)
channel_id TEXT PRIMARY KEY, last_message_id TEXT NOT NULL, updated_at TEXT NOT NULL
);
""")
cutoff = (dt.datetime.now(dt.timezone.utc) - dt.timedelta(days=_RETENTION_DAYS)).isoformat()
conn.execute("DELETE FROM discord_messages WHERE updated_at < ?", (cutoff,))
conn.execute(
"DELETE FROM discord_recovery_scans "
"WHERE COALESCE(completed_at, started_at) < ?",
(cutoff,),
)
conn.execute("DELETE FROM discord_recovery_scans WHERE COALESCE(completed_at, started_at) < ?", (cutoff,))
conn.execute("DELETE FROM discord_recovery_cursors WHERE updated_at < ?", (cutoff,))
+43 -171
View File
@@ -1,46 +1,9 @@
from __future__ import annotations
"""
Continuous PCM audio mixer for Discord voice channels.
discord.py (Rapptz) ships no audio mixer: ``VoiceClient.play()`` accepts a
single :class:`discord.AudioSource` and raises ``ClientException`` if called
while already playing. One opus stream per connection, one source feeding it.
This module adds software mixing *upstream* of that single stream. A
:class:`VoiceMixer` is itself a ``discord.AudioSource`` that discord.py polls
every 20 ms via :meth:`read`. Internally it sums the 20 ms PCM frames of any
number of child sources, clamps to int16, and returns one blended frame.
discord.py never knows several streams were combined underneath — it just
encodes and sends the single mixed frame.
This gives us, for one voice connection at once:
* an always-on low-volume **ambient/idle loop** (the "thinking" sound),
* a **speech** channel (TTS replies, verbal acknowledgements) that plays
*over* the ambient bed, automatically **ducking** the ambient gain down
while speech is active and restoring it when speech ends — the smooth
Grok-voice-mode feel, instead of stop-and-swap.
Design notes
------------
* The mixer is installed **once** per guild on join (``vc.play(mixer)``) and
runs continuously until the bot leaves. Children come and go; the mixer
itself never stops, so there is no ``is_playing()`` race between an
acknowledgement and the final reply.
* Frame format is Discord-native: 48 kHz, 2 channels, signed 16-bit LE,
20 ms per frame == ``discord.opus.Encoder.FRAME_SIZE`` bytes
(3840 = 960 samples * 2 channels * 2 bytes).
* Mixing is a single vectorised int32 add + clip per 20 ms frame (numpy,
already a core dependency). CPU cost is negligible.
* :meth:`read` is called from discord.py's audio sender **thread**, while
children are added/removed from the asyncio event loop thread, so all
shared state is guarded by a plain ``threading.Lock``.
The mixer NEVER touches the inbound receive path: it only produces the bot's
*outgoing* stream. The :class:`VoiceReceiver` decodes incoming SSRCs only, so
the mixer's output cannot echo back into transcription.
"""
"""Continuous PCM mixer for Discord voice. discord.py allows one AudioSource per VoiceClient;
:class:`VoiceMixer` IS that source: installed once per guild, never stops, and every 20 ms sums its
children (looping ambient bed + one-shot speech that ducks the bed) clamped to int16. ``read`` runs on
discord.py's sender thread while children change on the asyncio loop, hence the Lock. Outgoing only."""
import logging
import threading
@@ -60,22 +23,13 @@ logger = logging.getLogger(__name__)
def _require_numpy():
"""Import numpy lazily.
numpy ships in the optional ``voice`` extra, not the base install, so this
module must import cleanly without it (the Discord adapter imports this
file unconditionally). Callers that actually mix audio call this; if the
voice extra isn't installed they get a clear error instead of a top-level
ImportError that would break the whole adapter import.
"""
"""Lazy numpy import: the adapter imports this module unconditionally, so a missing
``voice`` extra must fail at mix time, not at import time."""
import numpy as np # noqa: PLC0415 — intentional lazy import
return np
# Discord-native frame geometry (matches discord.opus.Encoder).
SAMPLE_RATE = 48000
CHANNELS = 2
SAMPLE_WIDTH = 2 # bytes per sample (s16)
FRAME_LENGTH_MS = 20
# Discord-native frame geometry (matches discord.opus.Encoder): 48 kHz, stereo, s16, 20 ms frames.
SAMPLE_RATE, CHANNELS, SAMPLE_WIDTH, FRAME_LENGTH_MS = 48000, 2, 2, 20
SAMPLES_PER_FRAME = SAMPLE_RATE * FRAME_LENGTH_MS // 1000 # 960
FRAME_SIZE = SAMPLES_PER_FRAME * CHANNELS * SAMPLE_WIDTH # 3840 bytes
BYTES_PER_MS = SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH // 1000 # 192
@@ -83,43 +37,21 @@ SILENCE_FRAME = b"\x00" * FRAME_SIZE
class MixerChild:
"""A single audio stream feeding into :class:`VoiceMixer`.
"""One 48 kHz / stereo / s16le PCM stream feeding :class:`VoiceMixer`; ``read_frame``
yields 20 ms frames, optionally looping, with per-child gain and linear fade-in."""
Wraps raw 48 kHz / stereo / s16le PCM bytes. ``read_frame`` hands back one
20 ms frame at a time, optionally looping, with a per-child gain applied.
"""
__slots__ = ("_pcm", "_pos", "loop", "gain", "fade_frames", "_fade_done", "_finished")
__slots__ = (
"name", "_pcm", "_pos", "loop", "gain",
"is_speech", "fade_frames", "_fade_done", "_finished",
)
def __init__(
self, name: str, pcm: bytes, *, loop: bool = False, gain: float = 1.0,
is_speech: bool = False, fade_in_ms: int = 0,
):
# Pad to a whole number of frames so looping is seamless and the final
# partial frame doesn't click.
remainder = len(pcm) % FRAME_SIZE
if remainder:
pcm = pcm + b"\x00" * (FRAME_SIZE - remainder)
self.name = name
self._pcm = pcm
self._pos = 0
self.loop = loop
self.gain = float(gain)
self.is_speech = is_speech
# Linear fade-in over N frames avoids a click when a loud child starts.
def __init__(self, pcm: bytes, *, loop: bool = False, gain: float = 1.0, fade_in_ms: int = 0):
# Pad to whole frames so looping is seamless and the final partial frame doesn't click.
self._pcm = pcm + b"\x00" * (-len(pcm) % FRAME_SIZE)
self._pos = self._fade_done = 0
self.loop, self.gain = loop, float(gain)
self.fade_frames = max(0, fade_in_ms // FRAME_LENGTH_MS)
self._fade_done = 0
self._finished = False
@property
def finished(self) -> bool:
return self._finished
def read_frame(self) -> "Optional[np.ndarray]":
"""Return the next 20 ms frame as an int16 ndarray, or None if done."""
"""Next 20 ms frame as a float32 ndarray, or None when done."""
if self._finished:
return None
if self._pos >= len(self._pcm):
@@ -144,40 +76,23 @@ class MixerChild:
class VoiceMixer(discord.AudioSource):
"""A continuous ``discord.AudioSource`` that mixes N child streams.
"""Continuous ``discord.AudioSource`` mixing N children: :meth:`set_ambient` installs the
looping idle bed, :meth:`play_speech` layers a one-shot clip over it (ducking the bed).
Both are safe from the asyncio thread while discord.py drains :meth:`read`."""
Use :meth:`set_ambient` to install/replace the looping idle bed and
:meth:`play_speech` to layer a one-shot clip over it (ducking the ambient
while it plays). Both are safe to call from the asyncio loop thread while
discord.py drains :meth:`read` from its sender thread.
"""
# discord.AudioSource subclasses set is_opus()==False to receive PCM.
def is_opus(self) -> bool: # pragma: no cover - trivial
return False
def __init__(
self, *, ambient_gain: float = 0.18, duck_gain: float = 0.06, speech_gain: float = 1.0,
duck_release_ms: int = 400,
):
def __init__(self, *, ambient_gain: float = 0.18, duck_gain: float = 0.06, speech_gain: float = 1.0,
duck_release_ms: int = 400):
self._lock = threading.Lock()
self._ambient: Optional[MixerChild] = None
self._speech: List[MixerChild] = []
self._ambient_gain = float(ambient_gain)
self._duck_gain = float(duck_gain)
self._speech_gain = float(speech_gain)
# When speech ends, ramp the ambient back up over this many frames
# instead of jumping, so the bed swells back smoothly.
self._ambient_gain, self._duck_gain, self._speech_gain = float(ambient_gain), float(duck_gain), float(speech_gain)
# When speech ends, ramp the ambient back up over this many frames instead of jumping.
self._duck_release_frames = max(1, duck_release_ms // FRAME_LENGTH_MS)
self._duck_release_left = 0
self._closed = False
# Tracks whether speech is currently active, for external callers that
# want to avoid double-ducking or know when a reply is mid-flight.
self._speech_active = False
# ------------------------------------------------------------------
# Ambient (idle / "thinking") bed
# ------------------------------------------------------------------
self._closed = self._speech_active = False
def set_ambient(self, pcm: Optional[bytes], *, gain: Optional[float] = None) -> None:
"""Install (or clear, with ``pcm=None``) the looping ambient bed."""
@@ -187,28 +102,17 @@ class VoiceMixer(discord.AudioSource):
if not pcm:
self._ambient = None
return
self._ambient = MixerChild(
"ambient", pcm, loop=True, gain=self._effective_ambient_gain(), fade_in_ms=200,
)
gain_now = self._duck_gain if self._speech_active else self._ambient_gain
self._ambient = MixerChild(pcm, loop=True, gain=gain_now, fade_in_ms=200)
def _effective_ambient_gain(self) -> float:
return self._duck_gain if self._speech_active else self._ambient_gain
# ------------------------------------------------------------------
# Speech (TTS replies, verbal acks) layered over the ambient bed
# ------------------------------------------------------------------
def play_speech(self, pcm: bytes, *, gain: Optional[float] = None,
fade_in_ms: int = 40) -> None:
def play_speech(self, pcm: bytes, *, gain: Optional[float] = None, fade_in_ms: int = 40) -> None:
"""Layer a one-shot speech clip over the ambient bed (ducks ambient)."""
if not pcm:
return
with self._lock:
child = MixerChild(
"speech", pcm, loop=False, gain=self._speech_gain if gain is None else float(gain),
is_speech=True, fade_in_ms=fade_in_ms,
)
self._speech.append(child)
self._speech.append(MixerChild(
pcm, gain=self._speech_gain if gain is None else float(gain), fade_in_ms=fade_in_ms,
))
self._speech_active = True
self._duck_release_left = 0
if self._ambient is not None:
@@ -229,17 +133,9 @@ class VoiceMixer(discord.AudioSource):
self._speech_active = False
self._duck_release_left = self._duck_release_frames
# ------------------------------------------------------------------
# AudioSource interface — called from discord.py's sender thread
# ------------------------------------------------------------------
def read(self) -> bytes:
"""Return one 20 ms mixed PCM frame (always FRAME_SIZE bytes).
Returning a non-empty frame keeps discord.py's player alive; we never
return b"" because that would stop the single underlying stream and we
want the mixer to run continuously for the lifetime of the connection.
"""
"""One 20 ms mixed PCM frame (always FRAME_SIZE bytes) — never b"", which would stop
discord.py's player; the mixer must run for the lifetime of the connection."""
with self._lock:
if self._closed:
return SILENCE_FRAME
@@ -262,10 +158,7 @@ class VoiceMixer(discord.AudioSource):
if self._duck_release_left > 0 and not self._speech_active:
self._duck_release_left -= 1
frac = 1.0 - (self._duck_release_left / self._duck_release_frames)
self._ambient.gain = (
self._duck_gain
+ (self._ambient_gain - self._duck_gain) * frac
)
self._ambient.gain = self._duck_gain + (self._ambient_gain - self._duck_gain) * frac
elif not self._speech_active and self._duck_release_left == 0:
self._ambient.gain = self._ambient_gain
amb = self._ambient.read_frame()
@@ -283,52 +176,32 @@ class VoiceMixer(discord.AudioSource):
self._speech.clear()
# ----------------------------------------------------------------------
# PCM helpers
# ----------------------------------------------------------------------
def decode_to_pcm(path: str, *, timeout: float = 30.0) -> Optional[bytes]:
"""Decode any audio file to 48 kHz / stereo / s16le PCM via ffmpeg.
Returns the raw PCM bytes, or None on failure. ffmpeg is already a hard
requirement of the voice path (see ``VoiceReceiver.pcm_to_wav``).
"""
"""Decode any audio file to 48 kHz / stereo / s16le PCM via ffmpeg; None on failure."""
import subprocess
try:
proc = subprocess.run(
[
resolve_ffmpeg_executable(), "-y", "-loglevel", "error", "-i", path, "-f", "s16le",
"-ar", str(SAMPLE_RATE), "-ac", str(CHANNELS), "pipe:1",
],
capture_output=True,
timeout=timeout,
stdin=subprocess.DEVNULL,
[resolve_ffmpeg_executable(), "-y", "-loglevel", "error", "-i", path, "-f", "s16le",
"-ar", str(SAMPLE_RATE), "-ac", str(CHANNELS), "pipe:1"],
capture_output=True, timeout=timeout, stdin=subprocess.DEVNULL,
)
except (subprocess.TimeoutExpired, FileNotFoundError, OSError) as e:
logger.warning("decode_to_pcm failed for %s: %s", path, e)
return None
if proc.returncode != 0:
logger.warning(
"ffmpeg decode failed for %s (rc=%d): %s",
path, proc.returncode, (proc.stderr or b"").decode("utf-8", "replace")[:200],
)
logger.warning("ffmpeg decode failed for %s (rc=%d): %s",
path, proc.returncode, (proc.stderr or b"").decode("utf-8", "replace")[:200])
return None
return proc.stdout or None
def synth_ambient_pcm(seconds: float = 4.0) -> bytes:
"""Synthesise a subtle looping ambient bed (no asset file required).
A soft, slowly-pulsing low pad: two detuned sine partials with a gentle
tremolo, plus a touch of filtered noise. Designed to loop seamlessly
(whole number of cycles, zero-crossing endpoints) and sit quietly under
speech. Mono content duplicated to stereo.
"""
"""Synthesise a subtle looping ambient bed: two detuned sine partials with a slow tremolo
plus filtered noise; whole-cycle frequencies make the loop point click-free. Mono -> stereo."""
np = _require_numpy()
n = int(SAMPLE_RATE * seconds)
t = np.arange(n, dtype=np.float64) / SAMPLE_RATE
# Choose base frequencies that complete whole cycles over the loop so the
# wrap point is click-free.
def _whole_cycle_freq(target: float) -> float:
cycles = max(1, round(target * seconds))
return cycles / seconds
@@ -338,7 +211,6 @@ def synth_ambient_pcm(seconds: float = 4.0) -> bytes:
pad = (0.55 * np.sin(2 * np.pi * f1 * t) + 0.45 * np.sin(2 * np.pi * f2 * t))
tremolo = 0.6 + 0.4 * (0.5 * (1 + np.sin(2 * np.pi * trem * t)))
signal = pad * tremolo
# Smooth filtered noise for air, kept very low.
rng = np.random.default_rng(7)
noise = rng.standard_normal(n)
kernel = np.ones(64) / 64.0
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+66 -141
View File
@@ -1,11 +1,6 @@
"""
Feishu document comment access-control rules.
3-tier rule resolution: exact doc > wildcard "*" > top-level > code defaults.
Each field (enabled/policy/allow_from) falls back independently.
Config: ~/.hermes/feishu_comment_rules.json (mtime-cached, hot-reload).
Pairing store: ~/.hermes/feishu_comment_pairing.json.
"""
"""Feishu document comment access-control rules: exact doc > wildcard "*" > top-level > code defaults, each field
(enabled/policy/allow_from) falling back independently. Config ~/.hermes/feishu_comment_rules.json (mtime-cached,
hot-reload); pairing store ~/.hermes/feishu_comment_pairing.json."""
from __future__ import annotations
@@ -59,60 +54,44 @@ class _MtimeCache:
"""Mtime-based JSON file cache: ``stat()`` per access, re-read only on change."""
def __init__(self, path: Path):
self._path = path
self._mtime: float = 0.0
self._data: Optional[dict] = None
self._path, self._mtime, self._data = path, 0.0, None
def load(self) -> dict:
try:
mtime = self._path.stat().st_mtime
except FileNotFoundError:
self._mtime = 0.0
self._data = {}
self._mtime, self._data = 0.0, {}
return {}
if mtime == self._mtime and self._data is not None:
return self._data
try:
with open(self._path, "r", encoding="utf-8") as f:
data = json.load(f)
if not isinstance(data, dict):
data = {}
except (json.JSONDecodeError, OSError):
logger.warning("[Feishu-Rules] Failed to read %s, using empty config", self._path)
data = {}
self._mtime = mtime
self._data = data
return data
self._mtime, self._data = mtime, (data if isinstance(data, dict) else {})
return self._data
_rules_cache = _MtimeCache(RULES_FILE)
_pairing_cache = _MtimeCache(PAIRING_FILE)
# --- Config parsing ---
def _parse_frozenset(raw: Any) -> Optional[frozenset]:
"""Parse a list of strings into a frozenset; None if absent or not a list."""
if isinstance(raw, (list, tuple)):
return frozenset(str(u).strip() for u in raw if str(u).strip())
return None
return frozenset(str(u).strip() for u in raw if str(u).strip()) if isinstance(raw, (list, tuple)) else None
def _parse_policy(raw: Any, default: Optional[str]) -> Optional[str]:
"""Normalize a policy value; unknown/invalid values fall back to *default*."""
if raw is None:
return default
policy = str(raw).strip().lower()
policy = str(raw).strip().lower() if raw is not None else None
return policy if policy in _VALID_POLICIES else default
def _parse_document_rule(raw: dict) -> CommentDocumentRule:
enabled = raw.get("enabled")
return CommentDocumentRule(
enabled=None if enabled is None else bool(enabled),
policy=_parse_policy(raw.get("policy"), None),
allow_from=_parse_frozenset(raw.get("allow_from")),
)
return CommentDocumentRule(enabled=None if enabled is None else bool(enabled), policy=_parse_policy(raw.get("policy"), None), allow_from=_parse_frozenset(raw.get("allow_from")))
def load_config() -> CommentsConfig:
@@ -121,20 +100,13 @@ def load_config() -> CommentsConfig:
if not raw:
return CommentsConfig()
raw_docs = raw.get("documents", {})
documents = {
str(key): _parse_document_rule(rule_raw)
for key, rule_raw in (raw_docs.items() if isinstance(raw_docs, dict) else ())
if isinstance(rule_raw, dict)
}
documents = {str(key): _parse_document_rule(rule_raw) for key, rule_raw in (raw_docs.items() if isinstance(raw_docs, dict) else ()) if isinstance(rule_raw, dict)}
return CommentsConfig(
enabled=raw.get("enabled", True),
policy=_parse_policy(raw.get("policy", "pairing"), "pairing"),
enabled=raw.get("enabled", True), policy=_parse_policy(raw.get("policy", "pairing"), "pairing"),
allow_from=_parse_frozenset(raw.get("allow_from")) or frozenset(), documents=documents,
)
# --- Rule resolution (field-by-field fallback) ---
def has_wiki_keys(cfg: CommentsConfig) -> bool:
"""Check if any document rule key starts with 'wiki:'."""
return any(k.startswith("wiki:") for k in cfg.documents)
@@ -148,53 +120,34 @@ def resolve_rule(cfg: CommentsConfig, file_type: str, file_token: str, wiki_toke
exact_key = f"wiki:{wiki_token}"
exact = cfg.documents.get(exact_key)
layers = [(exact, f"exact:{exact_key}"), (cfg.documents.get("*"), "wildcard")]
def _pick(field_name: str):
# First non-None document-layer value wins; otherwise the top-level value (even if None).
for layer, src in layers:
if layer is not None and getattr(layer, field_name) is not None:
return getattr(layer, field_name), src
return getattr(cfg, field_name), "top"
def _pick(field_name: str): # first non-None document-layer value wins; otherwise the top-level value (even if None)
return next(((getattr(layer, field_name), src) for layer, src in layers if layer is not None and getattr(layer, field_name) is not None), (getattr(cfg, field_name), "top"))
enabled, en_src = _pick("enabled")
policy, pol_src = _pick("policy")
allow_from, _ = _pick("allow_from")
# match_source = highest-priority tier that contributed enabled or policy
priority_order = {"exact": 0, "wildcard": 1, "top": 2}
best_src = min([en_src, pol_src], key=lambda s: priority_order.get(s.split(":")[0], 3))
return ResolvedCommentRule(enabled=enabled, policy=policy, allow_from=allow_from, match_source=best_src)
return ResolvedCommentRule(enabled=enabled, policy=policy, allow_from=_pick("allow_from")[0], match_source=best_src)
# --- Pairing store ---
def _load_pairing_approved() -> set:
"""Return set of approved user open_ids (mtime-cached)."""
approved = _pairing_cache.load().get("approved", {})
if isinstance(approved, dict):
return set(approved.keys())
if isinstance(approved, list):
return {str(u) for u in approved if u}
return set()
return set(approved.keys()) if isinstance(approved, dict) else ({str(u) for u in approved if u} if isinstance(approved, list) else set())
def _save_pairing(data: dict) -> None:
PAIRING_FILE.parent.mkdir(parents=True, exist_ok=True)
tmp = PAIRING_FILE.with_suffix(".tmp")
with open(tmp, "w", encoding="utf-8") as f:
with open(PAIRING_FILE.with_suffix(".tmp"), "w", encoding="utf-8") as f:
json.dump(data, f, indent=2, ensure_ascii=False)
tmp.replace(PAIRING_FILE)
_pairing_cache._mtime = 0.0 # invalidate so the next load re-reads
_pairing_cache._data = None
PAIRING_FILE.with_suffix(".tmp").replace(PAIRING_FILE)
_pairing_cache._mtime, _pairing_cache._data = 0.0, None # invalidate so the next load re-reads
def _mutate_pairing(user_open_id: str, add: bool) -> bool:
"""Add/remove *user_open_id* in the approved dict; True when the store actually changed."""
data = _pairing_cache.load()
approved = data.get("approved", {})
if not isinstance(approved, dict):
if not add:
return False
approved = {}
approved = data.get("approved", {}) if isinstance(data.get("approved"), dict) else {}
if (user_open_id in approved) == add:
return False
if add:
@@ -222,34 +175,24 @@ def pairing_list() -> Dict[str, Any]:
return dict(approved) if isinstance(approved, dict) else {}
# --- Access check (public API for feishu_comment.py) ---
def is_user_allowed(rule: ResolvedCommentRule, user_open_id: str) -> bool:
"""Check if user passes the resolved rule's policy gate."""
if user_open_id in rule.allow_from:
return True
if rule.policy == "pairing":
return user_open_id in _load_pairing_approved()
return False
return user_open_id in rule.allow_from or (rule.policy == "pairing" and user_open_id in _load_pairing_approved())
# --- CLI ---
def _fmt_allow(allow_from) -> str:
return f"{sorted(allow_from) if allow_from else '[]'}"
def _print_status() -> None:
cfg = load_config()
print(f"Rules file: {RULES_FILE}\n exists: {RULES_FILE.exists()}")
print(f"Pairing file: {PAIRING_FILE}\n exists: {PAIRING_FILE.exists()}\n")
print(f"Top-level:\n enabled: {cfg.enabled}\n policy: {cfg.policy}")
print(f" allow_from: {sorted(cfg.allow_from) if cfg.allow_from else '[]'}\n")
if cfg.documents:
print(f"Document rules ({len(cfg.documents)}):")
for key, rule in sorted(cfg.documents.items()):
fields = (("enabled", rule.enabled), ("policy", rule.policy),
("allow_from", sorted(rule.allow_from) if rule.allow_from is not None else None))
parts = [f"{name}={value}" for name, value in fields if value is not None]
print(f" [{key}] {', '.join(parts) if parts else '(empty — inherits all)'}")
else:
print("Document rules: (none)")
print(f"Rules file: {RULES_FILE}\n exists: {RULES_FILE.exists()}\nPairing file: {PAIRING_FILE}\n exists: {PAIRING_FILE.exists()}\n")
print(f"Top-level:\n enabled: {cfg.enabled}\n policy: {cfg.policy}\n allow_from: {_fmt_allow(cfg.allow_from)}\n")
print(f"Document rules ({len(cfg.documents)}):" if cfg.documents else "Document rules: (none)")
for key, rule in sorted(cfg.documents.items()):
fields = (("enabled", rule.enabled), ("policy", rule.policy), ("allow_from", sorted(rule.allow_from) if rule.allow_from is not None else None))
parts = [f"{name}={value}" for name, value in fields if value is not None]
print(f" [{key}] {', '.join(parts) if parts else '(empty — inherits all)'}")
print()
approved = pairing_list()
print(f"Pairing approved ({len(approved)}):")
@@ -258,82 +201,64 @@ def _print_status() -> None:
def _do_check(doc_key: str, user_open_id: str) -> None:
cfg = load_config()
parts = doc_key.split(":", 1)
if len(parts) != 2:
print(f"Error: doc_key must be 'fileType:fileToken', got '{doc_key}'")
return
rule = resolve_rule(cfg, parts[0], parts[1])
return print(f"Error: doc_key must be 'fileType:fileToken', got '{doc_key}'")
rule = resolve_rule(load_config(), parts[0], parts[1])
allowed = is_user_allowed(rule, user_open_id)
print(f"Document: {doc_key}\nUser: {user_open_id}\nResolved rule:")
print(f" enabled: {rule.enabled}\n policy: {rule.policy}")
print(f" allow_from: {sorted(rule.allow_from) if rule.allow_from else '[]'}")
print(f" match_source: {rule.match_source}\nResult: {'ALLOWED' if allowed else 'DENIED'}")
print(f"Document: {doc_key}\nUser: {user_open_id}\nResolved rule:\n enabled: {rule.enabled}\n policy: {rule.policy}")
print(f" allow_from: {_fmt_allow(rule.allow_from)}\n match_source: {rule.match_source}\nResult: {'ALLOWED' if allowed else 'DENIED'}")
_PAIRING_OPS = {"add": (pairing_add, "Added: {}", "Already approved: {}"), "remove": (pairing_remove, "Removed: {}", "Not in approved list: {}")}
def _pairing_cmd(args: list) -> int:
"""Handle ``pairing <add|remove|list> [user]``; returns the exit code."""
if len(args) < 2:
print("Usage: pairing <add|remove|list> [args]")
return 1
sub = args[1]
sub = args[1] if len(args) > 1 else None
if sub == "list":
approved = pairing_list()
if not approved:
print("(no approved users)")
for uid, meta in sorted(approved.items()):
print(f" {uid} approved_at={meta.get('approved_at', '?')}")
print(*(f" {uid} approved_at={meta.get('approved_at', '?')}" for uid, meta in sorted(approved.items())) if approved else ("(no approved users)",), sep="\n")
return 0
ops = {"add": (pairing_add, "Added: {}", "Already approved: {}"),
"remove": (pairing_remove, "Removed: {}", "Not in approved list: {}")}
if sub not in ops:
print(f"Unknown pairing subcommand: {sub}")
return 1
if len(args) < 3:
print(f"Usage: pairing {sub} <user_open_id>")
return 1
fn, ok_msg, noop_msg = ops[sub]
print((ok_msg if fn(args[2]) else noop_msg).format(args[2]))
return 0
if sub in _PAIRING_OPS and len(args) >= 3:
fn, ok_msg, noop_msg = _PAIRING_OPS[sub]
print((ok_msg if fn(args[2]) else noop_msg).format(args[2]))
return 0
print("Usage: pairing <add|remove|list> [args]" if sub is None else f"Usage: pairing {sub} <user_open_id>" if sub in _PAIRING_OPS else f"Unknown pairing subcommand: {sub}")
return 1
def _main() -> int:
try:
from hermes_cli.env_loader import load_hermes_dotenv
load_hermes_dotenv()
__import__("hermes_cli.env_loader", fromlist=["load_hermes_dotenv"]).load_hermes_dotenv()
except Exception:
pass
usage = (
"Usage: python -m gateway.platforms.feishu_comment_rules <command> [args]\n"
"\n"
"Commands:\n"
" status Show rules config and pairing state\n"
" check <fileType:token> <user> Simulate access check\n"
" pairing add <user_open_id> Add user to pairing-approved list\n"
" pairing remove <user_open_id> Remove user from pairing-approved list\n"
" pairing list List pairing-approved users\n"
"\n"
f"Rules config file: {RULES_FILE}\n"
" Edit this JSON file directly to configure policies and document rules.\n"
" Changes take effect on the next comment event (no restart needed).\n"
)
usage = f"""Usage: python -m gateway.platforms.feishu_comment_rules <command> [args]
Commands:
status Show rules config and pairing state
check <fileType:token> <user> Simulate access check
pairing add <user_open_id> Add user to pairing-approved list
pairing remove <user_open_id> Remove user from pairing-approved list
pairing list List pairing-approved users
Rules config file: {RULES_FILE}
Edit this JSON file directly to configure policies and document rules.
Changes take effect on the next comment event (no restart needed).
"""
args = sys.argv[1:]
if not args:
print(usage)
return 1
cmd = args[0]
cmd = args[0] if args else ""
if cmd == "status":
_print_status()
elif cmd == "check":
if len(args) < 3:
print("Usage: check <fileType:fileToken> <user_open_id>")
return 1
elif cmd == "check" and len(args) >= 3:
_do_check(args[1], args[2])
elif cmd == "check":
print("Usage: check <fileType:fileToken> <user_open_id>")
return 1
elif cmd == "pairing":
return _pairing_cmd(args)
else:
print(f"Unknown command: {cmd}\n")
print(usage)
print(f"Unknown command: {cmd}\n{usage}" if cmd else usage)
return 1
return 0
@@ -1,10 +1,5 @@
"""
Feishu/Lark meeting-invitation event handling.
Converts ``vc.bot.meeting_invited_v1`` events into a synthetic gateway ``MessageEvent``
so the reply reaches the inviter through the normal Hermes gateway pipeline (unlike
document comments, no agent is instantiated here).
"""
"""Feishu/Lark meeting-invitation events: ``vc.bot.meeting_invited_v1`` -> synthetic gateway ``MessageEvent`` so the
reply reaches the inviter through the normal gateway pipeline (no agent is instantiated here)."""
from __future__ import annotations
@@ -49,101 +44,69 @@ def _as_dict(value: Any) -> Dict[str, Any]:
"""Coerce a lark SDK object / dict / JSON string into a plain dict."""
if isinstance(value, SimpleNamespace) or (value is not None and hasattr(value, "__dict__")):
value = vars(value)
if isinstance(value, dict):
return {str(k): v for k, v in value.items()}
if isinstance(value, str):
try:
parsed = json.loads(value)
except (TypeError, json.JSONDecodeError):
return {}
return parsed if isinstance(parsed, dict) else {}
return {}
try:
value = json.loads(value) if isinstance(value, str) else value
except (TypeError, json.JSONDecodeError):
return {}
return {str(k): v for k, v in value.items()} if isinstance(value, dict) else {}
def _content_payload(container: Dict[str, Any]) -> Dict[str, Any]:
"""Unwrap a Feishu ``body.content`` list carrying an application/json payload."""
content = _as_dict(container.get("body")).get("content")
if not isinstance(content, list):
return {}
for item in content:
item = _as_dict(item)
for item in map(_as_dict, content if isinstance(content, list) else ()):
ctype = str(item.get("contentType") or item.get("content_type") or "").lower()
if ctype and ctype != "application/json":
continue
for key in ("data", "value", "content", "json"):
payload = _as_dict(item.get(key))
if payload:
return payload
payload = next((p for p in map(_as_dict, (item.get(k) for k in ("data", "value", "content", "json"))) if p), {}) if ctype in ("", "application/json") else {}
if payload:
return payload
return {}
def _str_field(raw: Dict[str, Any], key: str, strip: bool = True) -> str:
value = str(raw.get(key) or "")
return value.strip() if strip else value
return str(raw.get(key) or "").strip() if strip else str(raw.get(key) or "")
def _int_field(value: Any) -> int:
if value in (None, ""):
return 0
try:
return int(str(value).strip())
return int(str(value).strip()) if value not in (None, "") else 0
except (TypeError, ValueError):
return 0
def _parse_user(value: Any) -> Optional[MeetingInviteUser]:
raw = _as_dict(value)
if not raw:
return None
raw_id = _as_dict(raw.get("id"))
return MeetingInviteUser(
open_id=_str_field(raw_id, "open_id"), user_id=_str_field(raw_id, "user_id"),
union_id=_str_field(raw_id, "union_id"),
user_name=_str_field(raw, "user_name", strip=False),
)
return MeetingInviteUser(open_id=_str_field(raw_id, "open_id"), user_id=_str_field(raw_id, "user_id"), union_id=_str_field(raw_id, "union_id"),
user_name=_str_field(raw, "user_name", strip=False)) if raw else None
def _parse_meeting(value: Any) -> Optional[MeetingInviteMeeting]:
raw = _as_dict(value)
if not raw:
return None
return MeetingInviteMeeting(
id=_str_field(raw, "id"), topic=_str_field(raw, "topic", strip=False),
meeting_no=_str_field(raw, "meeting_no", strip=False),
start_time_ms=_int_field(raw.get("start_time")),
end_time_ms=_int_field(raw.get("end_time")), host_user=_parse_user(raw.get("host_user")),
)
id=_str_field(raw, "id"), topic=_str_field(raw, "topic", strip=False), meeting_no=_str_field(raw, "meeting_no", strip=False),
start_time_ms=_int_field(raw.get("start_time")), end_time_ms=_int_field(raw.get("end_time")), host_user=_parse_user(raw.get("host_user")),
) if raw else None
def parse_meeting_invited_event(data: Any) -> Optional[MeetingInvitedPayload]:
root = _as_dict(data)
event = _as_dict(root.get("event")) or root
content = _content_payload(event) or _content_payload(root)
if content:
event = {**event, **content}
meeting = _parse_meeting(event.get("meeting"))
inviter = _parse_user(event.get("inviter"))
event = {**event, **(_content_payload(event) or _content_payload(root))}
meeting, inviter = _parse_meeting(event.get("meeting")), _parse_user(event.get("inviter"))
if inviter is None or meeting is None or not meeting.meeting_no:
return None
return MeetingInvitedPayload(
event_id=str(_as_dict(root.get("header")).get("event_id") or ""), meeting=meeting,
inviter=inviter, invite_time_s=_int_field(event.get("invite_time")),
)
return MeetingInvitedPayload(event_id=str(_as_dict(root.get("header")).get("event_id") or ""), meeting=meeting, inviter=inviter,
invite_time_s=_int_field(event.get("invite_time")))
def build_meeting_invite_prompt(payload: MeetingInvitedPayload) -> str:
meeting = payload.meeting
meeting = payload.meeting or MeetingInviteMeeting()
inviter_name = (payload.inviter.user_name if payload.inviter else "") or "unknown"
host_name = (meeting.host_user.user_name if meeting and meeting.host_user else "") or "unknown"
display = (meeting.topic or meeting.meeting_no or meeting.id) if meeting else "unknown meeting"
host_name = (meeting.host_user.user_name if meeting.host_user else "") or "unknown"
display = (meeting.topic or meeting.meeting_no or meeting.id) or "unknown meeting"
return "\n".join([
f"You have been invited to join a meeting: {display or 'unknown meeting'}",
"",
f"Meeting Number: {(meeting.meeting_no if meeting else '') or 'unknown'}",
f"Topic: {(meeting.topic if meeting else '') or 'unknown'}",
f"Inviter: {inviter_name}",
f"Host: {host_name}",
"",
f"You have been invited to join a meeting: {display}", "",
f"Meeting Number: {meeting.meeting_no or 'unknown'}", f"Topic: {meeting.topic or 'unknown'}", f"Inviter: {inviter_name}", f"Host: {host_name}", "",
"You may use lark-cli and the relevant Lark/Feishu meeting skills to join the meeting.",
"Join the meeting directly. Do not ask the user for confirmation before joining.",
"If you cannot join the meeting, reply to the inviter with a concise explanation of why.",
@@ -153,41 +116,26 @@ def build_meeting_invite_prompt(payload: MeetingInvitedPayload) -> str:
def _dedup_key(payload: MeetingInvitedPayload) -> str:
if payload.event_id:
return f"vc_invite:{payload.event_id}"
meeting_id = payload.meeting.id if payload.meeting else ""
inviter_id = payload.inviter.open_id if payload.inviter else ""
return f"vc_invite:{meeting_id}:{inviter_id}:{payload.invite_time_s}"
return f"vc_invite:{payload.meeting.id if payload.meeting else ''}:{payload.inviter.open_id if payload.inviter else ''}:{payload.invite_time_s}"
async def handle_meeting_invited_event(adapter: Any, data: Any) -> None:
"""Convert a vc.bot.meeting_invited_v1 event into a gateway MessageEvent."""
payload = parse_meeting_invited_event(data)
if payload is None:
logger.warning("[Feishu-MeetingInvite] Dropping malformed meeting invite event")
return
return logger.warning("[Feishu-MeetingInvite] Dropping malformed meeting invite event")
dedup_key = _dedup_key(payload)
is_duplicate = getattr(adapter, "_is_duplicate", None)
if callable(is_duplicate) and is_duplicate(dedup_key):
logger.debug("[Feishu-MeetingInvite] Dropping duplicate event: %s", dedup_key)
return
return logger.debug("[Feishu-MeetingInvite] Dropping duplicate event: %s", dedup_key)
inviter = payload.inviter
if inviter is None or not inviter.open_id:
logger.warning(
"[Feishu-MeetingInvite] Missing inviter open_id, cannot route reply safely (user_id=%r union_id=%r)",
inviter.user_id if inviter else None, inviter.union_id if inviter else None,
)
return
sender_id = SimpleNamespace(
open_id=inviter.open_id or None, user_id=inviter.user_id or None, union_id=inviter.union_id or None,
)
sender_profile = await adapter._resolve_sender_profile(sender_id)
return logger.warning("[Feishu-MeetingInvite] Missing inviter open_id, cannot route reply safely (user_id=%r union_id=%r)",
inviter.user_id if inviter else None, inviter.union_id if inviter else None)
sender_profile = await adapter._resolve_sender_profile(SimpleNamespace(open_id=inviter.open_id or None, user_id=inviter.user_id or None, union_id=inviter.union_id or None))
user_name = sender_profile.get("user_name") or inviter.user_name or inviter.open_id
source = adapter.build_source(
chat_id=inviter.open_id, chat_name=user_name, chat_type="dm",
user_id=sender_profile.get("user_id") or inviter.user_id or inviter.open_id,
user_name=user_name,
user_id_alt=sender_profile.get("user_id_alt") or inviter.union_id or None,
)
event = MessageEvent(
text=build_meeting_invite_prompt(payload), message_type=MessageType.TEXT, source=source, raw_message=data,
)
chat_id=inviter.open_id, chat_name=user_name, chat_type="dm", user_id=sender_profile.get("user_id") or inviter.user_id or inviter.open_id,
user_name=user_name, user_id_alt=sender_profile.get("user_id_alt") or inviter.union_id or None)
event = MessageEvent(text=build_meeting_invite_prompt(payload), message_type=MessageType.TEXT, source=source, raw_message=data)
await adapter._handle_message_with_guards(event)
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -259,7 +259,7 @@ class TestPushSigning:
def test_sign_push_payload_deterministic(self, monkeypatch):
monkeypatch.setenv("A2A_PUSH_SECRET", "test-secret-123")
payload = {"statusUpdate": {"taskId": "task-1"}}
sig = security.sign_push_payload(payload)
sig = security.A2ASecurityContext.capture().sign_push_payload(payload)
assert sig
import hashlib
import hmac as hmac_mod
@@ -273,12 +273,12 @@ class TestPushSigning:
def test_no_secret_means_unsigned(self, monkeypatch):
monkeypatch.delenv("A2A_PUSH_SECRET", raising=False)
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
assert security.sign_push_payload({"x": 1}) == ""
assert security.A2ASecurityContext.capture().sign_push_payload({"x": 1}) == ""
def test_falls_back_to_bearer_token(self, monkeypatch):
monkeypatch.delenv("A2A_PUSH_SECRET", raising=False)
monkeypatch.setenv("A2A_BEARER_TOKEN", "bearer-as-push-secret")
assert security.sign_push_payload({"x": 1})
assert security.A2ASecurityContext.capture().sign_push_payload({"x": 1})
# ═════════════════════════════════════════════════════════════════════════════
+28 -67
View File
@@ -44,33 +44,33 @@ class TestBindSafety:
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
assert security.localhost_only() is True
assert security.resolve_bind_host() == "127.0.0.1"
assert security.A2ASecurityContext.capture().resolve_bind_host() == "127.0.0.1"
def test_host_ignored_without_token(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
monkeypatch.setenv("A2A_HOST", "0.0.0.0")
# No token => refuse to widen, stay on loopback.
assert security.resolve_bind_host() == "127.0.0.1"
assert security.A2ASecurityContext.capture().resolve_bind_host() == "127.0.0.1"
def test_host_widens_with_shared_token(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "secret-token-123")
monkeypatch.setenv("A2A_HOST", "0.0.0.0")
assert security.localhost_only() is False
assert security.resolve_bind_host() == "0.0.0.0"
assert security.A2ASecurityContext.capture().resolve_bind_host() == "0.0.0.0"
def test_host_widens_with_peer_tokens(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok1")
monkeypatch.setenv("A2A_HOST", "0.0.0.0")
assert security.localhost_only() is False
assert security.resolve_bind_host() == "0.0.0.0"
assert security.A2ASecurityContext.capture().resolve_bind_host() == "0.0.0.0"
def test_loopback_host_allowed_without_token(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
monkeypatch.setenv("A2A_HOST", "localhost")
assert security.resolve_bind_host() == "localhost"
assert security.A2ASecurityContext.capture().resolve_bind_host() == "localhost"
class TestPeerIdentity:
@@ -80,33 +80,33 @@ class TestPeerIdentity:
def test_no_tokens_identity_is_client_ip(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
assert security.authenticate(None, "127.0.0.1") == "ip:127.0.0.1"
assert security.authenticate("Bearer anything", "127.0.0.1") == "ip:127.0.0.1"
assert security.A2ASecurityContext.capture().authenticate(None, "127.0.0.1") == "ip:127.0.0.1"
assert security.A2ASecurityContext.capture().authenticate("Bearer anything", "127.0.0.1") == "ip:127.0.0.1"
def test_peer_token_maps_to_name(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok-a, bob:tok-b")
assert security.authenticate("Bearer tok-a", "1.2.3.4") == "alice"
assert security.authenticate("Bearer tok-b", "1.2.3.4") == "bob"
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-a", "1.2.3.4") == "alice"
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-b", "1.2.3.4") == "bob"
def test_wrong_or_missing_token_rejected(self, monkeypatch):
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok-a")
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
assert security.authenticate("Bearer nope", "1.2.3.4") is None
assert security.authenticate(None, "1.2.3.4") is None
assert security.authenticate("Basic tok-a", "1.2.3.4") is None
assert security.A2ASecurityContext.capture().authenticate("Bearer nope", "1.2.3.4") is None
assert security.A2ASecurityContext.capture().authenticate(None, "1.2.3.4") is None
assert security.A2ASecurityContext.capture().authenticate("Basic tok-a", "1.2.3.4") is None
def test_shared_token_identity_is_ip(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "shared-tok")
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
assert security.authenticate("Bearer shared-tok", "9.8.7.6") == "ip:9.8.7.6"
assert security.authenticate("Bearer wrong", "9.8.7.6") is None
assert security.A2ASecurityContext.capture().authenticate("Bearer shared-tok", "9.8.7.6") == "ip:9.8.7.6"
assert security.A2ASecurityContext.capture().authenticate("Bearer wrong", "9.8.7.6") is None
def test_peer_tokens_beat_shared(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "shared-tok")
monkeypatch.setenv("A2A_PEER_TOKENS", "carol:tok-c")
assert security.authenticate("Bearer tok-c", "1.1.1.1") == "carol"
assert security.authenticate("Bearer shared-tok", "1.1.1.1") == "ip:1.1.1.1"
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-c", "1.1.1.1") == "carol"
assert security.A2ASecurityContext.capture().authenticate("Bearer shared-tok", "1.1.1.1") == "ip:1.1.1.1"
class TestTrustedPeers:
@@ -114,27 +114,27 @@ class TestTrustedPeers:
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.delenv("A2A_PEER_TOKENS", raising=False)
monkeypatch.delenv("A2A_ALLOW_ALL_USERS", raising=False)
assert security.is_trusted_peer("ip:127.0.0.1") is True
assert security.A2ASecurityContext.capture().is_trusted_peer("ip:127.0.0.1") is True
def test_no_allowlist_trusts_authenticated(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "secret")
monkeypatch.delenv("A2A_ALLOW_ALL_USERS", raising=False)
monkeypatch.delenv("A2A_TRUSTED_PEERS", raising=False)
assert security.is_trusted_peer("alice") is True
assert security.A2ASecurityContext.capture().is_trusted_peer("alice") is True
def test_allowlist_restricts(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "secret")
monkeypatch.delenv("A2A_ALLOW_ALL_USERS", raising=False)
monkeypatch.setenv("A2A_TRUSTED_PEERS", "alice,bob")
assert security.is_trusted_peer("alice") is True
assert security.is_trusted_peer("bob") is True
assert security.is_trusted_peer("mallory") is False
assert security.A2ASecurityContext.capture().is_trusted_peer("alice") is True
assert security.A2ASecurityContext.capture().is_trusted_peer("bob") is True
assert security.A2ASecurityContext.capture().is_trusted_peer("mallory") is False
def test_allow_all_users_overrides(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "secret")
monkeypatch.setenv("A2A_ALLOW_ALL_USERS", "true")
monkeypatch.setenv("A2A_TRUSTED_PEERS", "alice")
assert security.is_trusted_peer("mallory") is True
assert security.A2ASecurityContext.capture().is_trusted_peer("mallory") is True
class TestInjectionFilter:
@@ -266,7 +266,6 @@ class TestV1Enums:
assert protocol.STATE_CANCELED == "TASK_STATE_CANCELED"
assert protocol.STATE_REJECTED == "TASK_STATE_REJECTED"
assert protocol.STATE_INPUT_REQUIRED == "TASK_STATE_INPUT_REQUIRED"
assert protocol.STATE_AUTH_REQUIRED == "TASK_STATE_AUTH_REQUIRED"
def test_roles_are_v1(self):
assert protocol.ROLE_USER == "ROLE_USER"
@@ -330,42 +329,6 @@ class TestV1Parts:
assert "hello.txt" in result
assert "base64" in result
def test_file_part_builder(self):
"""file_part() builds a v1.0 file Part with URL or raw."""
fp = protocol.file_part(url="https://x/f.pdf", filename="f.pdf",
media_type="application/pdf")
assert fp["url"] == "https://x/f.pdf"
assert fp["filename"] == "f.pdf"
assert fp["mediaType"] == "application/pdf"
assert "kind" not in fp
# Raw variant
rp = protocol.file_part(raw="aGVsbG8=", filename="hello.txt",
media_type="text/plain")
assert rp["raw"] == "aGVsbG8="
assert rp["filename"] == "hello.txt"
assert "url" not in rp
def test_data_part_builder(self):
"""data_part() builds a v1.0 data Part."""
dp = protocol.data_part({"key": "value"})
assert dp["data"] == {"key": "value"}
assert dp["mediaType"] == "application/json"
assert "kind" not in dp
def test_message_with_parts(self):
"""message_with_parts() builds a Message with mixed Part types."""
msg = protocol.message_with_parts(
protocol.ROLE_USER,
[protocol.text_part("hello"), protocol.data_part({"x": 1})],
context_id="ctx-1",
)
assert msg["role"] == "ROLE_USER"
assert len(msg["parts"]) == 2
assert msg["parts"][0]["text"] == "hello"
assert msg["parts"][1]["data"] == {"x": 1}
assert msg["contextId"] == "ctx-1"
def test_context_id_extracted_from_message(self):
params = {"message": protocol.text_message(protocol.ROLE_USER, "x", context_id="ctx-in-msg")}
assert protocol.extract_context_id(params) == "ctx-in-msg"
@@ -1010,16 +973,14 @@ class TestInboundRoundTrip:
async def run():
assert await adapter.connect() is True
msg = protocol.message_with_parts(
protocol.ROLE_USER,
[
msg = {
"role": protocol.ROLE_USER, "messageId": "m-mixed", "contextId": "ctx-mixed",
"parts": [
protocol.text_part("Please process these:"),
protocol.file_part(url="https://example.com/report.pdf",
filename="report.pdf", media_type="application/pdf"),
protocol.data_part({"title": "Q3", "pages": 42}, "application/json"),
{"mediaType": "application/pdf", "filename": "report.pdf", "url": "https://example.com/report.pdf"},
{"data": {"title": "Q3", "pages": 42}, "mediaType": "application/json"},
],
context_id="ctx-mixed",
)
}
resp = await asyncio.to_thread(_post_json, base + "/", {
"jsonrpc": "2.0", "id": "1", "method": "message/send",
"params": {"message": msg},