fix(a2a): security hardening from code review

Critical fixes:
- SSRF protection: validate push notification callback URLs (block
  internal/private/loopback/metadata, enforce http/https only)
- Request body size limit: 1MB max (prevents memory exhaustion DoS)
- Thread safety: module-level locks for turn tracking, rate limiting,
  and pending task registry (was lazily initialized, racy)
- Peer identity: fall back to client IP when 'peer' field absent
  (prevents rate limiting collapse to single 'unknown' bucket)

Minor fixes:
- Watchdog survives reconnect: clear _watchdog_stop in connect()
- Redact error messages before sending to peers
- Remove dead _streaming_queues state
- Fix duplicate tags key in Agent Card skills
- Always send contextId in a2a_call (fixes client/server mismatch)
- Clear push_callbacks on disconnect
- SSE streaming cleanup via try/finally

16 new tests covering SSRF, body size, thread safety, watchdog
reconnect, error redaction, contextId consistency.
Tests: 97 passed, 3 deselected, 0 failed.
This commit is contained in:
Kevin (OpenClaw Bot)
2026-07-06 14:09:58 +10:00
committed by Teknium
parent c6b0e3a80e
commit 37481dccf4
5 changed files with 315 additions and 89 deletions
+49 -26
View File
@@ -51,6 +51,7 @@ _DEFAULT_PORT = 9900
_REPLY_TIMEOUT = 300 # seconds to wait for the agent to answer an inbound task
_ORPHAN_TIMEOUT = 300 # seconds before a pending task is considered orphaned
_WATCHDOG_INTERVAL = 60 # seconds between orphaned task watchdog runs
_MAX_BODY = 1_048_576 # 1MB max request body — prevents DoS via memory exhaustion
def _default_agent_name() -> str:
@@ -83,9 +84,6 @@ class A2AAdapter(BasePlatformAdapter):
# Per-context reply futures: an inbound HTTP request blocks on its
# future until adapter.send() resolves it with the agent's reply.
self._pending_replies: Dict[str, Future] = {}
# Per-context streaming queues: for message/stream, the handler writes
# SSE chunks and the send() method pushes intermediate results.
self._streaming_queues: Dict[str, list] = {}
self._pending_lock = threading.Lock()
# Push notification callback URLs per task
@@ -93,8 +91,8 @@ class A2AAdapter(BasePlatformAdapter):
self._push_lock = threading.Lock()
# Orphaned task watchdog
self._watchdog_thread: Optional[threading.Thread] = None
self._watchdog_stop = threading.Event()
self._watchdog_thread: Optional[threading.Thread] = None
@property
def name(self) -> str:
@@ -167,6 +165,9 @@ class A2AAdapter(BasePlatformAdapter):
return
try:
length = int(self.headers.get("Content-Length", 0))
if length > _MAX_BODY:
self._json(413, protocol.jsonrpc_error(None, -32700, "payload too large"))
return
raw = self.rfile.read(length) if length else b"{}"
req = json.loads(raw.decode("utf-8"))
except Exception:
@@ -177,8 +178,14 @@ class A2AAdapter(BasePlatformAdapter):
method = req.get("method", "")
params = req.get("params", {}) or {}
# Rate limit check
peer_id = str(params.get("peer") or (params.get("message", {}) or {}).get("from") or "unknown")
# Rate limit check — derive peer identity from authenticated source.
# A2A spec has no "peer" field, so we use the remote IP as fallback.
# This prevents rate limiting from collapsing to a single "unknown" bucket
# and makes the trusted-peer gate meaningful.
peer_id = str(params.get("peer") or (params.get("message", {}) or {}).get("from") or "")
if not peer_id:
# Fall back to client IP — the one thing we actually know
peer_id = self.client_address[0] if hasattr(self, "client_address") else "unknown"
if not protocol.rate_limit_allow(peer_id):
protocol.metrics.rate_limit_triggers += 1
self._json(429, protocol.jsonrpc_error(req_id, -32002, "rate limit exceeded"))
@@ -246,6 +253,9 @@ class A2AAdapter(BasePlatformAdapter):
)
self._server_thread.start()
# Reset watchdog state for reconnection (disconnect sets the event)
self._watchdog_stop.clear()
# Start orphaned task watchdog
self._watchdog_thread = threading.Thread(
target=self._watchdog_loop,
@@ -279,7 +289,8 @@ class A2AAdapter(BasePlatformAdapter):
if not fut.done():
fut.set_result("[agent shutting down]")
self._pending_replies.clear()
self._streaming_queues.clear()
with self._push_lock:
self._push_callbacks.clear()
# ── Orphaned task watchdog ─────────────────────────────────────────────
@@ -404,8 +415,8 @@ class A2AAdapter(BasePlatformAdapter):
except Exception as e:
with self._pending_lock:
self._pending_replies.pop(context_id, None)
protocol.complete_pending_task(task_id, protocol.STATE_FAILED, f"Dispatch failed: {e}")
return protocol.build_task(task_id, context_id, protocol.STATE_FAILED, f"Dispatch failed: {e}")
protocol.complete_pending_task(task_id, protocol.STATE_FAILED, security.redact_outbound(f"Dispatch failed: {e}"))
return protocol.build_task(task_id, context_id, protocol.STATE_FAILED, security.redact_outbound(f"Dispatch failed: {e}"))
try:
reply = fut.result(timeout=_REPLY_TIMEOUT)
@@ -535,7 +546,7 @@ class A2AAdapter(BasePlatformAdapter):
with self._pending_lock:
self._pending_replies.pop(context_id, None)
event = protocol.build_streaming_event("status", task_id, context_id, {
"status": {"state": protocol.STATE_FAILED, "message": f"Dispatch failed: {e}"},
"status": {"state": protocol.STATE_FAILED, "message": security.redact_outbound(f"Dispatch failed: {e}")},
})
handler.wfile.write(event.encode("utf-8"))
done = protocol.build_streaming_event("done", task_id, context_id)
@@ -545,23 +556,24 @@ class A2AAdapter(BasePlatformAdapter):
# 3. Wait for reply (with keepalive pings)
start = time.time()
reply = None
while True:
try:
reply = fut.result(timeout=5)
break
except TimeoutError:
# Send keepalive comment
handler.wfile.write(b": keepalive\n\n")
handler.wfile.flush()
if time.time() - start > _REPLY_TIMEOUT:
try:
while True:
try:
reply = fut.result(timeout=5)
break
except TimeoutError:
# Send keepalive comment
handler.wfile.write(b": keepalive\n\n")
handler.wfile.flush()
if time.time() - start > _REPLY_TIMEOUT:
reply = "[agent did not reply in time]"
break
except Exception:
reply = "[agent did not reply in time]"
break
except Exception:
reply = "[agent did not reply in time]"
break
with self._pending_lock:
self._pending_replies.pop(context_id, None)
finally:
with self._pending_lock:
self._pending_replies.pop(context_id, None)
reply = security.redact_outbound(reply or "")
protocol.persist_message(context_id, "agent", reply, task_id)
@@ -588,13 +600,24 @@ class A2AAdapter(BasePlatformAdapter):
# ── Push notifications ────────────────────────────────────────────────
def _send_push_notification(self, task_id: str, context_id: str, reply: str, state: str) -> None:
"""Send a push notification to the registered callback URL for this task."""
"""Send a push notification to the registered callback URL for this task.
Validates the callback URL to prevent SSRF — blocks internal/private
addresses (169.254.x.x metadata, loopback, RFC1918 private ranges)
unless we're in localhost-only mode (where internal access is expected).
"""
with self._push_lock:
callback_url = self._push_callbacks.pop(task_id, None)
if not callback_url:
return
# SSRF protection: validate the callback URL
if not security.is_safe_callback_url(callback_url):
logger.warning("A2A: push notification for task %s blocked — unsafe callback URL: %s", task_id, callback_url)
protocol.metrics.push_failed += 1
return
payload = {
"taskId": task_id,
"contextId": context_id,
+41 -53
View File
@@ -131,8 +131,7 @@ def skills_from_real_toolsets(toolset_registry: dict) -> list[dict]:
"id": f"toolset.{ts_name}",
"name": ts_name,
"description": desc,
"tags": [ts_name],
# Include tool names as tags for capability matching
# Include toolset name + tool names as tags for capability matching
"tags": [ts_name] + tool_names[:10],
})
if not skills:
@@ -233,6 +232,8 @@ def _now_iso() -> str:
# Anti-loop ping-pong protection
# --------------------------------------------------------------------------
import threading
# Track turns per context_id to prevent infinite agent-to-agent loops.
# A "turn" is one inbound message/send from a peer. When the count exceeds
# max_pingpong_turns(), we reject further messages for that context.
@@ -240,6 +241,7 @@ def _now_iso() -> str:
_turn_counts: dict[str, int] = defaultdict(int)
_turn_timestamps: dict[str, float] = {}
_turn_lock = threading.Lock()
# Clean up turn tracking for contexts older than 1 hour.
_TURN_TTL = 3600
@@ -251,27 +253,30 @@ def track_turn(context_id: str) -> int:
Returns the *new* count. Caller should reject if > max_pingpong_turns().
Also prunes stale entries to prevent unbounded growth.
"""
now = time.time()
# Prune stale entries
stale = [cid for cid, ts in _turn_timestamps.items() if now - ts > _TURN_TTL]
for cid in stale:
_turn_counts.pop(cid, None)
_turn_timestamps.pop(cid, None)
with _turn_lock:
now = time.time()
# Prune stale entries
stale = [cid for cid, ts in _turn_timestamps.items() if now - ts > _TURN_TTL]
for cid in stale:
_turn_counts.pop(cid, None)
_turn_timestamps.pop(cid, None)
_turn_counts[context_id] += 1
_turn_timestamps[context_id] = now
return _turn_counts[context_id]
_turn_counts[context_id] += 1
_turn_timestamps[context_id] = now
return _turn_counts[context_id]
def turn_count(context_id: str) -> int:
"""Return current turn count for a context (0 if unknown)."""
return _turn_counts.get(context_id, 0)
with _turn_lock:
return _turn_counts.get(context_id, 0)
def reset_turns(context_id: str) -> None:
"""Reset turn count for a context (e.g. after explicit cancel)."""
_turn_counts.pop(context_id, None)
_turn_timestamps.pop(context_id, None)
with _turn_lock:
_turn_counts.pop(context_id, None)
_turn_timestamps.pop(context_id, None)
# --------------------------------------------------------------------------
@@ -393,6 +398,7 @@ def list_conversations() -> list[str]:
_RATE_LIMIT_DEFAULT = 60 # requests per minute
_rate_buckets: dict[str, deque[float]] = defaultdict(deque)
_rate_lock = threading.Lock()
_RATE_WINDOW = 60.0 # seconds
@@ -405,25 +411,27 @@ def _rate_limit_per_minute() -> int:
def rate_limit_allow(peer: str) -> bool:
"""Check if peer is within rate limit. Returns True if allowed."""
limit = _rate_limit_per_minute()
now = time.time()
bucket = _rate_buckets[peer]
# Expire old entries
while bucket and now - bucket[0] > _RATE_WINDOW:
bucket.popleft()
if len(bucket) >= limit:
return False
bucket.append(now)
return True
with _rate_lock:
limit = _rate_limit_per_minute()
now = time.time()
bucket = _rate_buckets[peer]
# Expire old entries
while bucket and now - bucket[0] > _RATE_WINDOW:
bucket.popleft()
if len(bucket) >= limit:
return False
bucket.append(now)
return True
def rate_limit_status(peer: str) -> dict[str, Any]:
"""Return rate limit status for a peer."""
limit = _rate_limit_per_minute()
now = time.time()
bucket = _rate_buckets[peer]
# Count active entries
active = sum(1 for ts in bucket if now - ts <= _RATE_WINDOW)
with _rate_lock:
limit = _rate_limit_per_minute()
now = time.time()
bucket = _rate_buckets[peer]
# Count active entries
active = sum(1 for ts in bucket if now - ts <= _RATE_WINDOW)
return {
"peer": peer,
"limit_per_minute": limit,
@@ -441,15 +449,11 @@ def rate_limit_status(peer: str) -> dict[str, Any]:
# OpenClaw pattern: sessions_send (async durable messaging).
_pending_tasks: dict[str, dict[str, Any]] = {}
_pending_lock = None # Will be set by adapter on init
_pending_lock = threading.Lock()
def register_pending_task(task_id: str, context_id: str, peer: str, callback_url: str = "") -> None:
"""Register a task as pending (in-flight)."""
import threading
global _pending_lock
if _pending_lock is None:
_pending_lock = threading.Lock()
with _pending_lock:
_pending_tasks[task_id] = {
"context_id": context_id,
@@ -462,10 +466,6 @@ def register_pending_task(task_id: str, context_id: str, peer: str, callback_url
def complete_pending_task(task_id: str, state: str, reply: str = "") -> dict | None:
"""Mark a pending task as complete. Returns the task info if found."""
import threading
global _pending_lock
if _pending_lock is None:
_pending_lock = threading.Lock()
with _pending_lock:
info = _pending_tasks.pop(task_id, None)
if info:
@@ -477,10 +477,6 @@ def complete_pending_task(task_id: str, state: str, reply: str = "") -> dict | N
def pending_task_info(task_id: str) -> dict | None:
"""Get info about a pending task (for tasks/get)."""
import threading
global _pending_lock
if _pending_lock is None:
_pending_lock = threading.Lock()
with _pending_lock:
return _pending_tasks.get(task_id)
@@ -490,12 +486,8 @@ def orphaned_tasks(timeout_seconds: int = 300) -> list[dict]:
Used by the orphaned task watchdog to clean up stale tasks.
"""
import threading
global _pending_lock
if _pending_lock is None:
_pending_lock = threading.Lock()
now = time.time()
with _pending_lock:
now = time.time()
return [
{"task_id": tid, **info}
for tid, info in _pending_tasks.items()
@@ -505,13 +497,9 @@ def orphaned_tasks(timeout_seconds: int = 300) -> list[dict]:
def clear_orphaned_tasks(timeout_seconds: int = 300) -> list[str]:
"""Remove and return task_ids of tasks pending longer than timeout."""
import threading
global _pending_lock
if _pending_lock is None:
_pending_lock = threading.Lock()
now = time.time()
cleared = []
with _pending_lock:
now = time.time()
cleared = []
for tid in list(_pending_tasks.keys()):
info = _pending_tasks[tid]
if now - info.get("started_at", now) > timeout_seconds:
+71
View File
@@ -281,6 +281,77 @@ def verify_push_signature(payload: dict, signature: str) -> bool:
return hmac.compare_digest(signature, expected)
# --------------------------------------------------------------------------
# SSRF protection for push notification callback URLs
# --------------------------------------------------------------------------
import ipaddress
import urllib.parse
# Blocked IP ranges for push callback URLs (SSRF prevention).
# Even in localhost-only mode we block these — a remote peer shouldn't
# be able to make us probe internal services.
_BLOCKED_PREFIXES = (
"169.254.", # link-local / AWS metadata
"127.", # loopback
"10.", # RFC1918 private
"172.16.", "172.17.", "172.18.", "172.19.", "172.20.",
"172.21.", "172.22.", "172.23.", "172.24.", "172.25.",
"172.26.", "172.27.", "172.28.", "172.29.", "172.30.", "172.31.", # RFC1918 private
"192.168.", # RFC1918 private
"0.0.0.0", # unspecified
"::1", # IPv6 loopback
"fe80:", # IPv6 link-local
"fc00:", "fd00:", # IPv6 unique-local
)
def is_safe_callback_url(url: str) -> bool:
"""Check if a push notification callback URL is safe from SSRF.
Blocks internal/private/loopback/metadata addresses.
Only allows http:// and https:// schemes.
"""
if not url or not isinstance(url, str):
return False
try:
parsed = urllib.parse.urlparse(url)
except Exception:
return False
# Scheme check
if parsed.scheme not in ("http", "https"):
return False
hostname = parsed.hostname or ""
if not hostname:
return False
# Check for literal "localhost" hostname
hostname_lower = hostname.lower()
if hostname_lower == "localhost":
if localhost_only():
return True
return False
# Check against blocked prefixes
for prefix in _BLOCKED_PREFIXES:
if hostname_lower.startswith(prefix.lower()):
# Allow localhost in localhost-only mode (local testing)
if localhost_only() and prefix == "127.":
return True
if localhost_only() and prefix == "::1":
return True
return False
# Also check via ipaddress for numeric IPs
try:
ip = ipaddress.ip_address(hostname)
if ip.is_loopback or ip.is_link_local or ip.is_private or ip.is_reserved:
# Allow localhost in localhost-only mode
if localhost_only() and ip.is_loopback:
return True
return False
except ValueError:
pass # not an IP, it's a hostname — fine
return True
# --------------------------------------------------------------------------
# Audit log
# --------------------------------------------------------------------------
+12 -9
View File
@@ -172,12 +172,13 @@ def a2a_call(args: dict, **_: Any) -> str:
"jsonrpc": "2.0",
"id": protocol.new_task_id(),
"method": "message/send",
"params": {"message": protocol.text_message("user", safe_message)},
"params": {
"message": protocol.text_message("user", safe_message),
"contextId": ctx, # Always send contextId so server persists under same id
},
}
if context_id:
# A2A spec: contextId at top level of params (not just inside message)
rpc_body["params"]["contextId"] = context_id
rpc_body["params"]["message"]["contextId"] = context_id
# Also set contextId inside message for legacy callers
rpc_body["params"]["message"]["contextId"] = ctx
security.audit("outbound", agent, rpc_body["id"], safe_message)
protocol.persist_message(ctx, "user", safe_message, rpc_body["id"])
@@ -299,11 +300,13 @@ def _call_peer_sync(agent_name: str, peer_entry: dict, message: str, context_id:
"jsonrpc": "2.0",
"id": protocol.new_task_id(),
"method": "message/send",
"params": {"message": protocol.text_message("user", safe_message)},
"params": {
"message": protocol.text_message("user", safe_message),
"contextId": ctx, # Always send contextId so server persists under same id
},
}
if context_id:
rpc_body["params"]["contextId"] = context_id
rpc_body["params"]["message"]["contextId"] = context_id
# Also set contextId inside message for legacy callers
rpc_body["params"]["message"]["contextId"] = ctx
security.audit("outbound", agent_name, rpc_body["id"], safe_message)
protocol.persist_message(ctx, "user", safe_message, rpc_body["id"])
+142 -1
View File
@@ -449,4 +449,145 @@ class TestWatchdog:
protocol._pending_tasks["task-watchdog-1"]["started_at"] = time.time() - 600
cleared = protocol.clear_orphaned_tasks(timeout_seconds=300)
assert "task-watchdog-1" in cleared
assert protocol.pending_task_info("task-watchdog-1") is None
assert protocol.pending_task_info("task-watchdog-1") is None
# ═════════════════════════════════════════════════════════════════════════════
# Code review fixes: SSRF protection, body size limit, watchdog reconnect
# ═════════════════════════════════════════════════════════════════════════════
class TestSSRFProtection:
"""Push notification callback URL validation (SSRF prevention)."""
def test_safe_https_url_allowed(self):
assert security.is_safe_callback_url("https://example.com/webhook") is True
def test_safe_http_url_allowed(self):
assert security.is_safe_callback_url("http://example.com/webhook") is True
def test_localhost_blocked_in_remote_mode(self):
"""In remote mode (localhost_only=False), localhost should be blocked."""
# Temporarily disable localhost_only
original = security.localhost_only
security.localhost_only = lambda: False
try:
assert security.is_safe_callback_url("http://127.0.0.1:8080/hook") is False
assert security.is_safe_callback_url("http://localhost:8080/hook") is False
finally:
security.localhost_only = original
def test_aws_metadata_blocked(self):
"""169.254.x.x (AWS metadata) must be blocked."""
original = security.localhost_only
security.localhost_only = lambda: False
try:
assert security.is_safe_callback_url("http://169.254.169.254/latest/meta-data/") is False
finally:
security.localhost_only = original
def test_private_ranges_blocked(self):
"""RFC1918 private ranges must be blocked in remote mode."""
original = security.localhost_only
security.localhost_only = lambda: False
try:
assert security.is_safe_callback_url("http://10.0.0.1/hook") is False
assert security.is_safe_callback_url("http://192.168.1.1/hook") is False
assert security.is_safe_callback_url("http://172.16.0.1/hook") is False
finally:
security.localhost_only = original
def test_file_scheme_blocked(self):
"""file:// URLs must be blocked (urllib follows them)."""
assert security.is_safe_callback_url("file:///etc/passwd") is False
def test_ftp_scheme_blocked(self):
assert security.is_safe_callback_url("ftp://example.com/file") is False
def test_empty_url_blocked(self):
assert security.is_safe_callback_url("") is False
assert security.is_safe_callback_url(None) is False
class TestBodySizeLimit:
"""Request body size limit (DoS prevention)."""
def test_max_body_constant_exists(self):
"""adapter module should define _MAX_BODY."""
from plugins.platforms.a2a import adapter
assert hasattr(adapter, "_MAX_BODY")
assert adapter._MAX_BODY > 0
# Should be reasonable (1-10MB)
assert adapter._MAX_BODY <= 10_485_760
def test_max_body_imports_from_adapter(self):
"""The body size check should reference _MAX_BODY."""
from plugins.platforms.a2a import adapter
import inspect
source = inspect.getsource(adapter)
assert "_MAX_BODY" in source
assert "413" in source # HTTP 413 Payload Too Large
class TestWatchdogReconnect:
"""Watchdog should survive reconnection (disconnect → connect cycle)."""
def test_watchdog_stop_cleared_on_connect(self):
"""connect() should call _watchdog_stop.clear() to reset state."""
from plugins.platforms.a2a.adapter import A2AAdapter
import inspect
source = inspect.getsource(A2AAdapter.connect)
assert "_watchdog_stop.clear()" in source
class TestErrorRedaction:
"""Error messages should be redacted before sending to peers."""
def test_dispatch_failed_uses_redact_outbound(self):
"""adapter should call security.redact_outbound on error messages."""
from plugins.platforms.a2a import adapter
import inspect
source = inspect.getsource(adapter)
# All 'Dispatch failed' messages should go through redact_outbound
import re
dispatch_fails = re.findall(r'"Dispatch failed[^"]*"', source)
for match in dispatch_fails:
# Find the surrounding context (should contain redact_outbound)
idx = source.index(match)
context = source[max(0, idx-200):idx+100]
assert "redact_outbound" in context, f"Dispatch failed message not redacted: {match}"
class TestThreadSafety:
"""Verify thread-safe shared state access."""
def test_turn_tracking_is_thread_safe(self):
"""track_turn, turn_count, reset_turns should use a lock."""
import inspect
assert hasattr(protocol, "_turn_lock")
source = inspect.getsource(protocol.track_turn)
assert "_turn_lock" in source
def test_rate_limiting_is_thread_safe(self):
"""rate_limit_allow should use a lock."""
import inspect
assert hasattr(protocol, "_rate_lock")
source = inspect.getsource(protocol.rate_limit_allow)
assert "_rate_lock" in source
def test_pending_tasks_lock_initialized_at_import(self):
"""_pending_lock should be initialized at module level, not lazily."""
assert hasattr(protocol, "_pending_lock")
assert protocol._pending_lock is not None
# Should be a Lock instance, not None
assert hasattr(protocol._pending_lock, "acquire")
class TestContextIdConsistency:
"""a2a_call should always send contextId, even on first turn."""
def test_a2a_call_always_sends_context_id(self):
"""tools.a2a_call should include contextId in params even when not provided."""
import inspect
source = inspect.getsource(tools.a2a_call)
# The contextId should be set unconditionally, not inside an if block
assert "contextId" in source