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:
committed by
Teknium
parent
c6b0e3a80e
commit
37481dccf4
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
@@ -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"])
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user