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:
@@ -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
@@ -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
@@ -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"))
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
+401
-802
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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"):
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,))
|
||||
|
||||
@@ -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
|
||||
|
||||
+227
-555
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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})
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
@@ -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},
|
||||
|
||||
Reference in New Issue
Block a user