refactor(agent/runtime): hooks/guardrails/dispatch — unify shell-hook and webhook plumbing, table-driven guardrail thresholds
- shell_hooks is the shared home: _ToolMatcherMixin (matcher compile + matches_tool), _payload_fields, _forget_home_registrations, _home_key, _utc_now_iso now serve outbound_webhooks too (copies deleted; every log string byte-identical). - shell_hooks: response parsing is a per-event dispatch table; _spawn diagnostic dict + _evaluate_result shared by the live callback and run_once; _locked_update_approvals POSIX/non-POSIX bodies merged via ExitStack. - tool_guardrails: ToolCallGuardrailConfig thresholds from a _THRESHOLD_SOURCES table (nested-wins-over-flat preserved); _int_at_least replaces _positive_int/_non_negative_int; observe_identical_call (0 refs) folded into observe_call; _halt helper for hard-stop decisions. - tool_dispatch_helpers: _plan_tool_batch_segments split into _batch_admission + close/extend helpers with the post-hoc normalization merged in. - Comment/docstring compaction keeping every stated rule.
This commit is contained in:
@@ -1,67 +1,31 @@
|
||||
"""
|
||||
Outbound webhook notifications.
|
||||
"""Outbound webhook notifications.
|
||||
|
||||
Reads the ``hooks.outbound:`` list from ``config.yaml`` and registers
|
||||
notify-only callbacks on the existing plugin hook manager, so every
|
||||
``invoke_hook()`` site can push lifecycle events to external HTTP
|
||||
endpoints — CI systems, dashboards, other agents — with zero changes to
|
||||
call sites and zero polling on the receiving end.
|
||||
Reads ``hooks.outbound:`` from config.yaml and registers notify-only callbacks on
|
||||
the plugin hook manager, so every ``invoke_hook()`` site can push lifecycle events
|
||||
to external HTTP endpoints. Outbound mirror of ``gateway/platforms/webhook.py``.
|
||||
|
||||
This is the outbound mirror of the inbound webhook platform
|
||||
(``gateway/platforms/webhook.py``): inbound wakes Hermes when the world
|
||||
changes; outbound tells the world when Hermes does something.
|
||||
* Delivery is fire-and-forget through a bounded queue and one daemon worker
|
||||
thread; callbacks serialize, enqueue, and return ``None`` immediately, so a
|
||||
target can never block a tool call or influence agent flow.
|
||||
* Payloads are HMAC-SHA256 signed (``X-Hermes-Signature-256: sha256=<hex>`` over
|
||||
the raw body) when a secret is configured.
|
||||
* No consent prompt (no code runs on this machine); ``HERMES_SAFE_MODE=1`` still
|
||||
skips registration. Registration is idempotent.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
* Delivery is fire-and-forget through a bounded in-process queue and a
|
||||
single daemon worker thread. ``invoke_hook()`` runs inside the agent
|
||||
loop, so callbacks must never block on network I/O — they serialize,
|
||||
enqueue, and return ``None`` immediately. Outbound targets can never
|
||||
block a tool call, inject context, or otherwise influence agent flow.
|
||||
* Payloads are signed with HMAC-SHA256 (GitHub-style
|
||||
``X-Hermes-Signature-256: sha256=<hexdigest>`` over the raw body) when
|
||||
a secret is configured. Receivers verify exactly like they verify
|
||||
GitHub webhooks.
|
||||
* No consent prompt: unlike shell hooks, an outbound target executes no
|
||||
code on this machine — it POSTs JSON to a URL the user themselves put
|
||||
in config. ``HERMES_SAFE_MODE=1`` still skips registration, matching
|
||||
plugins / MCP / shell hooks.
|
||||
* Registration is idempotent — safe to invoke from both the CLI entry
|
||||
point and the gateway entry point.
|
||||
|
||||
Config schema (``~/.hermes/config.yaml``)::
|
||||
Config::
|
||||
|
||||
hooks:
|
||||
outbound:
|
||||
- url: https://ci.example.com/hermes-events
|
||||
events: [on_session_end, subagent_stop]
|
||||
# secret literal (discouraged) or env var name (preferred):
|
||||
secret_env: HERMES_OUTBOUND_WEBHOOK_SECRET
|
||||
# optional regex, honored for pre/post_tool_call only:
|
||||
matcher: "terminal|delegate_task"
|
||||
timeout: 10 # per-attempt seconds, clamped to [1, 60]
|
||||
name: ci-notify # optional label for logs / `hermes hooks list`
|
||||
secret_env: HERMES_OUTBOUND_WEBHOOK_SECRET # or inline ``secret``
|
||||
matcher: "terminal|delegate_task" # pre/post_tool_call only
|
||||
timeout: 10 # seconds, clamped to [1, 60]
|
||||
name: ci-notify
|
||||
|
||||
Wire format (POST body)::
|
||||
|
||||
{
|
||||
"hook_event_name": "on_session_end",
|
||||
"tool_name": null,
|
||||
"tool_input": null,
|
||||
"session_id": "sess_abc123",
|
||||
"cwd": "/home/user/project",
|
||||
"extra": {...}, # event-specific kwargs
|
||||
"delivery_id": "3f2c...", # uuid4, unique per POST
|
||||
"timestamp": "2026-07-22T14:00:00Z"
|
||||
}
|
||||
|
||||
Headers::
|
||||
|
||||
Content-Type: application/json
|
||||
User-Agent: Hermes-Agent-Outbound-Webhook
|
||||
X-Hermes-Event: <hook event name>
|
||||
X-Hermes-Delivery: <delivery_id>
|
||||
X-Hermes-Signature-256: sha256=<hmac hexdigest> # only when secret set
|
||||
POST body: ``{hook_event_name, profile, tool_name, tool_input, session_id, cwd,
|
||||
extra, delivery_id, timestamp}``. Headers: ``Content-Type``, ``User-Agent``,
|
||||
``X-Hermes-Event``, ``X-Hermes-Delivery``, ``X-Hermes-Signature-256`` (if secret).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -78,12 +42,19 @@ import threading
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple
|
||||
from urllib import error as urlerror
|
||||
from urllib import request as urlrequest
|
||||
|
||||
from agent.shell_hooks import (
|
||||
_TOOL_EVENTS as _TOOL_SCOPED_EVENTS,
|
||||
_ToolMatcherMixin,
|
||||
_forget_home_registrations,
|
||||
_home_key,
|
||||
_payload_fields,
|
||||
_utc_now_iso,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT_SECONDS = 10
|
||||
@@ -92,17 +63,8 @@ MAX_DELIVERY_ATTEMPTS = 2
|
||||
RETRY_BACKOFF_SECONDS = 1.0
|
||||
QUEUE_MAX_SIZE = 256
|
||||
|
||||
# Events whose ``matcher`` field is honored (mirrors shell hooks).
|
||||
_TOOL_SCOPED_EVENTS = {"pre_tool_call", "post_tool_call"}
|
||||
|
||||
# kwargs promoted to top-level payload keys (mirrors shell hooks wire).
|
||||
_TOP_LEVEL_PAYLOAD_KEYS = {"tool_name", "args", "session_id", "parent_session_id"}
|
||||
|
||||
# (home, event, url) triples already wired to the plugin manager in this
|
||||
# process. Home is part of the key so a multiplexed gateway's secondary
|
||||
# profiles — each with their own plugin manager (see
|
||||
# hermes_cli.plugins.get_plugin_manager) — can register identical webhook
|
||||
# targets without the first profile's registration shadowing the rest.
|
||||
# (home, event, url) triples already wired in this process. Home is part of the key so a
|
||||
# multiplexed gateway's secondary profiles (own plugin managers) can register identical targets.
|
||||
_registered: Set[Tuple[str, str, str]] = set()
|
||||
_registered_lock = threading.Lock()
|
||||
|
||||
@@ -114,9 +76,11 @@ _worker: Optional[threading.Thread] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class WebhookTarget:
|
||||
class WebhookTarget(_ToolMatcherMixin):
|
||||
"""Parsed and validated representation of one ``hooks.outbound`` entry."""
|
||||
|
||||
_MATCHER_KIND = "outbound webhook"
|
||||
|
||||
url: str
|
||||
events: List[str]
|
||||
name: str = ""
|
||||
@@ -125,47 +89,18 @@ class WebhookTarget:
|
||||
timeout: int = DEFAULT_TIMEOUT_SECONDS
|
||||
compiled_matcher: Optional[re.Pattern] = field(default=None, repr=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if isinstance(self.matcher, str):
|
||||
stripped = self.matcher.strip()
|
||||
self.matcher = stripped if stripped else None
|
||||
if self.matcher:
|
||||
try:
|
||||
self.compiled_matcher = re.compile(self.matcher)
|
||||
except re.error as exc:
|
||||
logger.warning(
|
||||
"outbound webhook matcher %r is invalid (%s) — treating "
|
||||
"as literal equality", self.matcher, exc,
|
||||
)
|
||||
self.compiled_matcher = None
|
||||
|
||||
@property
|
||||
def label(self) -> str:
|
||||
return self.name or self.url
|
||||
|
||||
def matches_tool(self, tool_name: Optional[str]) -> bool:
|
||||
if not self.matcher:
|
||||
return True
|
||||
if tool_name is None:
|
||||
return False
|
||||
if self.compiled_matcher is not None:
|
||||
return self.compiled_matcher.fullmatch(tool_name) is not None
|
||||
return tool_name == self.matcher
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Public API -----------------------------------------------------------------
|
||||
|
||||
def register_from_config(cfg: Optional[Dict[str, Any]]) -> List[WebhookTarget]:
|
||||
"""Register every configured outbound webhook on the plugin manager.
|
||||
|
||||
``cfg`` is the full parsed config dict. Missing, empty, or malformed
|
||||
``hooks.outbound`` is treated as zero targets — config parsing never
|
||||
raises, because a broken webhook entry must not crash the agent.
|
||||
|
||||
Returns the targets that ended up wired (deduplicated across repeat
|
||||
calls, so the CLI and gateway can both invoke this safely).
|
||||
Malformed ``hooks.outbound`` means zero targets — never raises. Returns the
|
||||
targets that ended up wired (deduplicated across repeat calls).
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
return []
|
||||
@@ -176,18 +111,14 @@ def register_from_config(cfg: Optional[Dict[str, Any]]) -> List[WebhookTarget]:
|
||||
logger.info("HERMES_SAFE_MODE=1 — outbound webhook registration skipped")
|
||||
return []
|
||||
|
||||
hooks_cfg = cfg.get("hooks")
|
||||
targets = _parse_outbound_block(
|
||||
hooks_cfg.get("outbound") if isinstance(hooks_cfg, dict) else None
|
||||
)
|
||||
targets = iter_configured_targets(cfg)
|
||||
if not targets:
|
||||
return []
|
||||
|
||||
from hermes_cli.plugins import get_plugin_manager
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
manager = get_plugin_manager()
|
||||
home_key = str(get_hermes_home().expanduser().resolve())
|
||||
home_key = _home_key()
|
||||
|
||||
registered: List[WebhookTarget] = []
|
||||
with _registered_lock:
|
||||
@@ -214,8 +145,7 @@ def register_from_config(cfg: Optional[Dict[str, Any]]) -> List[WebhookTarget]:
|
||||
|
||||
|
||||
def iter_configured_targets(cfg: Optional[Dict[str, Any]]) -> List[WebhookTarget]:
|
||||
"""Parse ``hooks.outbound`` without registering anything.
|
||||
Used by ``hermes hooks list``."""
|
||||
"""Parse ``hooks.outbound`` without registering anything (``hermes hooks list``)."""
|
||||
if not isinstance(cfg, dict):
|
||||
return []
|
||||
hooks_cfg = cfg.get("hooks")
|
||||
@@ -225,8 +155,7 @@ def iter_configured_targets(cfg: Optional[Dict[str, Any]]) -> List[WebhookTarget
|
||||
|
||||
|
||||
def flush(timeout: float = 5.0) -> bool:
|
||||
"""Block until all queued deliveries are done (or *timeout* elapses).
|
||||
Returns ``True`` when the queue fully drained. Test/shutdown helper."""
|
||||
"""Block until all queued deliveries are done (or *timeout* elapses); True if drained."""
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
with _delivery_queue.all_tasks_done:
|
||||
@@ -238,25 +167,14 @@ def flush(timeout: float = 5.0) -> bool:
|
||||
|
||||
|
||||
def re_register_config_hooks() -> None:
|
||||
"""Re-register outbound webhooks from config after a plugin force-reload.
|
||||
"""Re-register outbound webhooks after a plugin force-reload cleared ``_hooks``.
|
||||
|
||||
Mirrors ``agent.shell_hooks.re_register_config_hooks``: config-owned
|
||||
outbound-webhook callbacks live in the same ``_hooks`` dict that
|
||||
``PluginManager.discover_and_load(force=True)`` clears via ``unload()``,
|
||||
so without this the force-reloaded profile's outbound webhooks go
|
||||
silently inert (#92682 review). Only the current home's idempotence
|
||||
keys are cleared so a force-reload in one profile cannot invalidate
|
||||
another profile's still-live registration.
|
||||
Only the current home's idempotence keys are cleared so a force-reload in one
|
||||
profile cannot invalidate another profile's still-live registration.
|
||||
"""
|
||||
from hermes_cli.config import load_config
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
home_key = str(get_hermes_home().expanduser().resolve())
|
||||
with _registered_lock:
|
||||
_registered.difference_update(
|
||||
{key for key in _registered if key[0] == home_key}
|
||||
)
|
||||
|
||||
_forget_home_registrations(_registered, _registered_lock)
|
||||
register_from_config(load_config())
|
||||
|
||||
|
||||
@@ -272,9 +190,7 @@ def reset_for_tests() -> None:
|
||||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Config parsing -------------------------------------------------------------
|
||||
|
||||
def _parse_outbound_block(raw: Any) -> List[WebhookTarget]:
|
||||
if raw is None:
|
||||
@@ -285,100 +201,68 @@ def _parse_outbound_block(raw: Any) -> List[WebhookTarget]:
|
||||
type(raw).__name__,
|
||||
)
|
||||
return []
|
||||
|
||||
targets: List[WebhookTarget] = []
|
||||
for i, entry in enumerate(raw):
|
||||
target = _parse_single_target(i, entry)
|
||||
if target is not None:
|
||||
targets.append(target)
|
||||
return targets
|
||||
targets = (_parse_single_target(i, entry) for i, entry in enumerate(raw))
|
||||
return [t for t in targets if t is not None]
|
||||
|
||||
|
||||
def _parse_single_target(index: int, raw: Any) -> Optional[WebhookTarget]:
|
||||
from hermes_cli.plugins import VALID_HOOKS
|
||||
|
||||
def warn(msg: str, *args: Any) -> None:
|
||||
logger.warning("hooks.outbound[%d]" + msg, index, *args)
|
||||
|
||||
if not isinstance(raw, dict):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d] must be a mapping with 'url' and 'events' "
|
||||
"keys; got %s", index, type(raw).__name__,
|
||||
)
|
||||
warn(" must be a mapping with 'url' and 'events' keys; got %s", type(raw).__name__)
|
||||
return None
|
||||
|
||||
url = raw.get("url")
|
||||
if not isinstance(url, str) or not url.strip():
|
||||
logger.warning("hooks.outbound[%d] is missing a non-empty 'url'", index)
|
||||
warn(" is missing a non-empty 'url'")
|
||||
return None
|
||||
url = url.strip()
|
||||
if not url.lower().startswith(("http://", "https://")):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d].url must be http(s); got %r — skipped",
|
||||
index, url,
|
||||
)
|
||||
warn(".url must be http(s); got %r — skipped", url)
|
||||
return None
|
||||
if url.lower().startswith("http://"):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d].url uses plain http:// — payloads (including "
|
||||
"tool inputs) travel unencrypted. Prefer https.", index,
|
||||
)
|
||||
warn(".url uses plain http:// — payloads (including tool inputs) travel unencrypted. Prefer https.")
|
||||
|
||||
events_raw = raw.get("events")
|
||||
valid_list = ", ".join(sorted(VALID_HOOKS))
|
||||
if not isinstance(events_raw, list) or not events_raw:
|
||||
logger.warning(
|
||||
"hooks.outbound[%d] needs a non-empty 'events' list (valid: %s)",
|
||||
index, ", ".join(sorted(VALID_HOOKS)),
|
||||
)
|
||||
warn(" needs a non-empty 'events' list (valid: %s)", valid_list)
|
||||
return None
|
||||
events: List[str] = []
|
||||
for ev in events_raw:
|
||||
if ev in VALID_HOOKS:
|
||||
events.append(ev)
|
||||
else:
|
||||
logger.warning(
|
||||
"hooks.outbound[%d]: unknown event %r ignored (valid: %s)",
|
||||
index, ev, ", ".join(sorted(VALID_HOOKS)),
|
||||
)
|
||||
warn(": unknown event %r ignored (valid: %s)", ev, valid_list)
|
||||
if not events:
|
||||
logger.warning(
|
||||
"hooks.outbound[%d] has no valid events — skipped", index,
|
||||
)
|
||||
warn(" has no valid events — skipped")
|
||||
return None
|
||||
|
||||
matcher = raw.get("matcher")
|
||||
if matcher is not None and not isinstance(matcher, str):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d].matcher must be a string regex; ignoring",
|
||||
index,
|
||||
)
|
||||
warn(".matcher must be a string regex; ignoring")
|
||||
matcher = None
|
||||
if matcher is not None and not any(e in _TOOL_SCOPED_EVENTS for e in events):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d].matcher=%r will be ignored — matcher is only "
|
||||
"honored for pre_tool_call / post_tool_call.", index, matcher,
|
||||
)
|
||||
warn(".matcher=%r will be ignored — matcher is only honored for pre_tool_call / post_tool_call.", matcher)
|
||||
matcher = None
|
||||
|
||||
timeout_raw = raw.get("timeout", DEFAULT_TIMEOUT_SECONDS)
|
||||
try:
|
||||
timeout = int(timeout_raw)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"hooks.outbound[%d].timeout must be an int (got %r); using "
|
||||
"default %ds", index, timeout_raw, DEFAULT_TIMEOUT_SECONDS,
|
||||
)
|
||||
warn(".timeout must be an int (got %r); using default %ds", timeout_raw, DEFAULT_TIMEOUT_SECONDS)
|
||||
timeout = DEFAULT_TIMEOUT_SECONDS
|
||||
timeout = max(1, min(timeout, MAX_TIMEOUT_SECONDS))
|
||||
|
||||
secret = _resolve_secret(index, raw)
|
||||
|
||||
name = raw.get("name")
|
||||
if not isinstance(name, str):
|
||||
name = ""
|
||||
|
||||
return WebhookTarget(
|
||||
url=url,
|
||||
events=events,
|
||||
name=name.strip(),
|
||||
secret=secret,
|
||||
name=name.strip() if isinstance(name, str) else "",
|
||||
secret=_resolve_secret(index, raw),
|
||||
matcher=matcher,
|
||||
timeout=timeout,
|
||||
)
|
||||
@@ -397,26 +281,21 @@ def _resolve_secret(index: int, raw: Dict[str, Any]) -> Optional[str]:
|
||||
)
|
||||
return None
|
||||
secret = raw.get("secret")
|
||||
if isinstance(secret, str) and secret:
|
||||
return secret
|
||||
return None
|
||||
return secret if isinstance(secret, str) and secret else None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Callback + delivery
|
||||
# ---------------------------------------------------------------------------
|
||||
# --- Callback + delivery --------------------------------------------------------
|
||||
|
||||
def _make_callback(event: str, target: WebhookTarget):
|
||||
"""Build the notify-only closure ``invoke_hook()`` calls per firing."""
|
||||
|
||||
def _callback(**kwargs: Any) -> None:
|
||||
if event in _TOOL_SCOPED_EVENTS:
|
||||
if not target.matches_tool(kwargs.get("tool_name")):
|
||||
return None
|
||||
if event in _TOOL_SCOPED_EVENTS and not target.matches_tool(kwargs.get("tool_name")):
|
||||
return None
|
||||
delivery_id = uuid.uuid4().hex
|
||||
try:
|
||||
body = _serialize_payload(event, kwargs, delivery_id)
|
||||
except Exception: # defensive — a bad payload must not hurt the loop
|
||||
except Exception: # a bad payload must not hurt the loop
|
||||
logger.warning(
|
||||
"outbound webhook payload serialization failed (event=%s "
|
||||
"target=%s)", event, target.label, exc_info=True,
|
||||
@@ -433,34 +312,20 @@ def _make_callback(event: str, target: WebhookTarget):
|
||||
def _serialize_payload(
|
||||
event: str, kwargs: Dict[str, Any], delivery_id: str,
|
||||
) -> bytes:
|
||||
"""Render the POST body. Same top-level shape as shell hooks' stdin
|
||||
(documented in :mod:`agent.shell_hooks`), plus delivery metadata.
|
||||
"""Render the POST body: shell-hooks stdin shape plus delivery metadata.
|
||||
|
||||
``delivery_id`` is shared with the ``X-Hermes-Delivery`` header so
|
||||
receivers can dedupe on either — and since it (plus ``timestamp``)
|
||||
lives inside the HMAC-signed body, it doubles as replay protection.
|
||||
``delivery_id`` (also the ``X-Hermes-Delivery`` header) and ``timestamp`` live
|
||||
inside the HMAC-signed body, so they double as replay protection.
|
||||
"""
|
||||
extras = {k: v for k, v in kwargs.items() if k not in _TOP_LEVEL_PAYLOAD_KEYS}
|
||||
try:
|
||||
cwd = str(Path.cwd())
|
||||
except OSError:
|
||||
cwd = ""
|
||||
# Resolved at fire time from the bound home so a multiplexed gateway's
|
||||
# receivers can tell which profile emitted the event (#92674).
|
||||
# Profile resolved at fire time so a multiplexed gateway's receivers can tell which profile emitted.
|
||||
from hermes_cli.profiles import get_active_profile_name
|
||||
|
||||
payload = {
|
||||
"hook_event_name": event,
|
||||
"profile": get_active_profile_name(),
|
||||
"tool_name": kwargs.get("tool_name"),
|
||||
"tool_input": kwargs.get("args") if isinstance(kwargs.get("args"), dict) else None,
|
||||
"session_id": kwargs.get("session_id") or kwargs.get("parent_session_id") or "",
|
||||
"cwd": cwd,
|
||||
"extra": extras,
|
||||
**_payload_fields(kwargs),
|
||||
"delivery_id": delivery_id,
|
||||
"timestamp": datetime.now(tz=timezone.utc)
|
||||
.isoformat()
|
||||
.replace("+00:00", "Z"),
|
||||
"timestamp": _utc_now_iso(),
|
||||
}
|
||||
return json.dumps(payload, ensure_ascii=False, default=str).encode("utf-8")
|
||||
|
||||
@@ -511,11 +376,9 @@ def _ensure_worker() -> None:
|
||||
target=_worker_loop, name="outbound-webhooks", daemon=True,
|
||||
)
|
||||
_worker.start()
|
||||
# The worker is a daemon thread, so a short-lived process (a `-q`
|
||||
# CLI run, a cron session) can exit right after enqueuing the
|
||||
# final events — silently dropping on_session_end, the headline
|
||||
# use case. Drain the queue at interpreter shutdown, bounded so
|
||||
# a dead endpoint can only delay exit, never hang it.
|
||||
# Daemon worker: a short-lived process could exit right after enqueuing
|
||||
# on_session_end. Drain at interpreter shutdown, bounded so a dead
|
||||
# endpoint can only delay exit, never hang it.
|
||||
atexit.register(flush, timeout=5.0)
|
||||
|
||||
|
||||
@@ -536,13 +399,8 @@ def _worker_loop() -> None:
|
||||
|
||||
|
||||
class _NoRedirectHandler(urlrequest.HTTPRedirectHandler):
|
||||
"""Refuse to follow redirects.
|
||||
|
||||
urllib's default handler converts a redirected POST into a body-less
|
||||
GET — the signed payload would be silently dropped and the headers
|
||||
re-sent to a location the user never configured. Treat any 3xx as a
|
||||
delivery failure instead (surfaced as HTTPError by returning None).
|
||||
"""
|
||||
"""Refuse redirects: urllib would turn a redirected POST into a body-less GET,
|
||||
silently dropping the signed payload. Any 3xx surfaces as HTTPError instead."""
|
||||
|
||||
def redirect_request(self, req, fp, code, msg, headers, newurl): # noqa: D102
|
||||
return None
|
||||
@@ -552,9 +410,7 @@ _opener = urlrequest.build_opener(_NoRedirectHandler)
|
||||
|
||||
|
||||
def _deliver(delivery: Dict[str, Any]) -> None:
|
||||
"""POST with bounded retries. Retries on connection errors and 5xx;
|
||||
4xx is the receiver telling us the request itself is wrong — no retry.
|
||||
3xx redirects are never followed (misconfiguration — fix the URL)."""
|
||||
"""POST with bounded retries: retry on connection errors and 5xx; 4xx and 3xx are final."""
|
||||
last_error = ""
|
||||
for attempt in range(1, MAX_DELIVERY_ATTEMPTS + 1):
|
||||
req = urlrequest.Request(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,26 +1,12 @@
|
||||
"""Tool-dispatch helpers — parallelism gating, multimodal envelopes, mutation tracking.
|
||||
|
||||
Pure module-level utilities extracted from ``run_agent.py``:
|
||||
|
||||
* ``_is_destructive_command`` — terminal-command heuristic used to gate
|
||||
parallel batch dispatch.
|
||||
* ``_should_parallelize_tool_batch`` / ``_extract_parallel_scope_paths`` /
|
||||
``_extract_parallel_scope_path`` / ``_paths_overlap`` — the rules engine
|
||||
deciding when a multi-tool batch can run concurrently (V4A patch scope
|
||||
uses patch-body file headers, not a decoy ``path=``).
|
||||
* ``_is_multimodal_tool_result`` / ``_multimodal_text_summary`` /
|
||||
``_append_subdir_hint_to_multimodal`` — envelope helpers for the
|
||||
``{"_multimodal": True, "content": [...], "text_summary": ...}`` dict
|
||||
shape returned by tools like ``computer_use``.
|
||||
* ``_extract_file_mutation_targets`` / ``_extract_landed_file_mutation_paths`` /
|
||||
``_extract_error_preview`` —
|
||||
per-turn file-mutation verifier inputs.
|
||||
* ``_trajectory_normalize_msg`` — strip image blobs from a message for
|
||||
trajectory saving.
|
||||
|
||||
All helpers are stateless. ``run_agent`` re-exports each name so existing
|
||||
``from run_agent import ...`` imports in tests and other modules keep
|
||||
working unchanged.
|
||||
Stateless module-level utilities extracted from ``run_agent.py``, which
|
||||
re-exports each name so existing ``from run_agent import ...`` imports keep
|
||||
working. Groups: batch-parallelism planner (path-overlap admission; V4A patch
|
||||
scope comes from patch-body headers, not a decoy ``path=``), multimodal
|
||||
``{"_multimodal": True, "content": [...], "text_summary": ...}`` envelope
|
||||
helpers, per-turn file-mutation verifier inputs, trajectory normalisation, and
|
||||
the tool-result message constructor with its untrusted-content wrapping.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -40,8 +26,7 @@ from tools.threat_patterns import scan_for_threats
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Tools that must never run concurrently (interactive / user-facing).
|
||||
# When any of these appear in a batch, we fall back to sequential execution.
|
||||
# Interactive / user-facing tools never run concurrently: any of these in a batch is a barrier.
|
||||
_NEVER_PARALLEL_TOOLS = frozenset({"clarify"})
|
||||
|
||||
# Read-only tools with no shared mutable session state.
|
||||
@@ -60,16 +45,12 @@ _PARALLEL_SAFE_TOOLS = frozenset({
|
||||
"web_search",
|
||||
})
|
||||
|
||||
# Filesystem tools whose parallel admission is decided by path overlap.
|
||||
# Readers may share a subtree with other readers; a writer conflicts with
|
||||
# ANY overlapping reservation (reader or writer). This is what keeps a
|
||||
# batched ``search_files``/``read_file`` from observing pre-mutation file
|
||||
# state when the model batches it alongside the ``patch``/``write_file``
|
||||
# it depends on (the classic same-block write→read race).
|
||||
# Filesystem tools admitted by path overlap. Readers may share a subtree; a
|
||||
# writer conflicts with ANY overlapping reservation. This keeps a batched
|
||||
# read_file/search_files from observing pre-mutation state when the model
|
||||
# batches it alongside the patch/write_file it depends on.
|
||||
_PATH_SCOPED_READERS = frozenset({"read_file", "search_files"})
|
||||
_PATH_SCOPED_WRITERS = frozenset({"write_file", "patch"})
|
||||
|
||||
# File tools can run concurrently when they target independent paths.
|
||||
_PATH_SCOPED_TOOLS = _PATH_SCOPED_READERS | _PATH_SCOPED_WRITERS
|
||||
|
||||
# Patterns that indicate a terminal command may modify/delete files.
|
||||
@@ -92,47 +73,29 @@ _REDIRECT_OVERWRITE = re.compile(r'[^>]>[^>]|^>[^>]')
|
||||
|
||||
def _is_destructive_command(cmd: str) -> bool:
|
||||
"""Heuristic: does this terminal command look like it modifies/deletes files?"""
|
||||
if not cmd:
|
||||
return False
|
||||
if _DESTRUCTIVE_PATTERNS.search(cmd):
|
||||
return True
|
||||
if _REDIRECT_OVERWRITE.search(cmd):
|
||||
return True
|
||||
return False
|
||||
return bool(cmd) and bool(_DESTRUCTIVE_PATTERNS.search(cmd) or _REDIRECT_OVERWRITE.search(cmd))
|
||||
|
||||
|
||||
def _is_mcp_tool_parallel_safe(tool_name: str) -> bool:
|
||||
"""Check if an MCP tool comes from a server with parallel tool calls enabled.
|
||||
|
||||
Lazy-imports from ``tools.mcp_tool`` to avoid circular dependencies.
|
||||
Returns False if the MCP module is not available.
|
||||
"""
|
||||
"""Whether an MCP tool's server opted into parallel calls; False if MCP is unavailable."""
|
||||
try:
|
||||
from tools.mcp_tool import is_mcp_tool_parallel_safe
|
||||
from tools.mcp_tool import is_mcp_tool_parallel_safe # lazy: avoids import cycle
|
||||
return is_mcp_tool_parallel_safe(tool_name)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# Read-only bridge lookups: dispatch_tool_search / dispatch_tool_describe are
|
||||
# stateless catalog reads (the catalog is rebuilt from the current tool-defs
|
||||
# list on every call), so a batch of them can run concurrently.
|
||||
# Stateless catalog reads (rebuilt from the current tool-defs on every call) — parallel-safe.
|
||||
_PARALLEL_SAFE_BRIDGE_LOOKUPS = frozenset({"tool_search", "tool_describe"})
|
||||
|
||||
|
||||
def _peel_bridge_call(tool_name: str, function_args: dict) -> tuple[str, dict]:
|
||||
"""Resolve a ``tool_call`` bridge invocation to its underlying tool.
|
||||
"""Resolve a ``tool_call`` bridge invocation to ``(underlying_name, underlying_args)``.
|
||||
|
||||
The batch planner admits calls to a parallel run by tool NAME, but when
|
||||
tool search is active the model emits the literal name ``tool_call`` for
|
||||
every deferred tool — so a server opted in via
|
||||
``supports_parallel_tool_calls: true`` silently lost concurrency the
|
||||
moment the bridge activated. Peel the wrapper here so admission is
|
||||
decided on the underlying tool, exactly like the executors' unwrap.
|
||||
|
||||
Returns ``(underlying_name, underlying_args)`` when the wrapper parses
|
||||
cleanly, else ``(tool_name, function_args)`` unchanged — an unparseable
|
||||
bridge call stays a sequential barrier and fails at dispatch as before.
|
||||
With tool search active the model emits the literal name ``tool_call`` for
|
||||
every deferred tool, so admission must be decided on the underlying tool
|
||||
(as the executors' unwrap does). An unparseable bridge call is returned
|
||||
unchanged: it stays a sequential barrier and fails at dispatch as before.
|
||||
"""
|
||||
try:
|
||||
from tools.tool_search import TOOL_CALL_NAME, resolve_underlying_call
|
||||
@@ -146,147 +109,97 @@ def _peel_bridge_call(tool_name: str, function_args: dict) -> tuple[str, dict]:
|
||||
return tool_name, function_args
|
||||
|
||||
|
||||
def _batch_admission(tool_call, execution_cwd: Optional[Path]) -> tuple[str, List[Path], bool] | None:
|
||||
"""Classify one call for the planner: ``None`` = sequential barrier, else
|
||||
``(effective_name, scoped_paths, is_writer)`` (empty paths = unscoped parallel-safe)."""
|
||||
tool_name = tool_call.function.name
|
||||
if tool_name in _NEVER_PARALLEL_TOOLS:
|
||||
return None
|
||||
try:
|
||||
function_args = json.loads(tool_call.function.arguments)
|
||||
except Exception:
|
||||
_raw = tool_call.function.arguments
|
||||
logging.debug(
|
||||
"Could not parse args for %s — treating as sequential barrier; raw=%s",
|
||||
tool_name,
|
||||
_raw[:200] if isinstance(_raw, str) else repr(_raw)[:200],
|
||||
)
|
||||
return None
|
||||
if not isinstance(function_args, dict):
|
||||
logging.debug(
|
||||
"Non-dict args for %s (%s) — treating as sequential barrier",
|
||||
tool_name,
|
||||
type(function_args).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
name, args = _peel_bridge_call(tool_name, function_args)
|
||||
if name in _NEVER_PARALLEL_TOOLS:
|
||||
return None
|
||||
if name in _PATH_SCOPED_TOOLS:
|
||||
scoped = _extract_parallel_scope_paths(name, args, execution_cwd=execution_cwd)
|
||||
return (name, scoped, name in _PATH_SCOPED_WRITERS) if scoped else None
|
||||
if name in _PARALLEL_SAFE_TOOLS or name in _PARALLEL_SAFE_BRIDGE_LOOKUPS or _is_mcp_tool_parallel_safe(name):
|
||||
return name, [], False
|
||||
return None
|
||||
|
||||
|
||||
def _plan_tool_batch_segments(tool_calls, *, execution_cwd: Optional[Path] = None) -> List[tuple]:
|
||||
"""Split a tool-call batch into ordered ``(kind, calls)`` segments.
|
||||
"""Split a tool-call batch into ordered ``("parallel"|"sequential", calls)`` segments.
|
||||
|
||||
``kind`` is ``"parallel"`` (a maximal contiguous run of parallel-safe
|
||||
calls) or ``"sequential"`` (one or more barrier calls that must run
|
||||
in-order on the sequential path). Segments preserve the model's
|
||||
original call order exactly — a later call never crosses an earlier
|
||||
barrier — so tool-result ordering and side-effect boundaries are
|
||||
identical to fully-sequential execution. The per-call safety rules
|
||||
are the same ones the old all-or-nothing gate applied to the whole
|
||||
batch:
|
||||
|
||||
* ``_NEVER_PARALLEL_TOOLS`` (interactive tools) → barrier.
|
||||
* Unparseable / non-dict arguments → barrier.
|
||||
* Path-scoped tools (``read_file``/``search_files``/``write_file``/
|
||||
``patch``) join a parallel run only when their target path(s) do not
|
||||
CONFLICT with a path already reserved in the same run. Reservations
|
||||
carry a reader/writer role: reader↔reader overlap is harmless (two
|
||||
reads of the same file commute) and stays parallel; any overlap
|
||||
involving a writer closes the run so the conflicting call starts a
|
||||
NEW run after the first completes. ``search_files`` reserves its
|
||||
search root (default ``.``) as a reader — a search batched after a
|
||||
write into the searched subtree is ordered behind that write instead
|
||||
of racing it. For V4A ``patch(mode="patch")`` the reserved paths are
|
||||
the file headers in the patch body, not a possibly-stale ``path=``
|
||||
argument.
|
||||
* Anything not in ``_PARALLEL_SAFE_TOOLS`` and not an opted-in MCP
|
||||
tool → barrier.
|
||||
|
||||
Parallel runs shorter than two calls are demoted to sequential (no
|
||||
concurrency win, and the sequential executor owns the richer inline
|
||||
dispatch), and adjacent sequential segments are merged.
|
||||
Segments preserve the model's call order exactly — a later call never
|
||||
crosses an earlier barrier — so result ordering and side-effect boundaries
|
||||
match fully-sequential execution. Barriers: ``_NEVER_PARALLEL_TOOLS``,
|
||||
unparseable/non-dict args, and anything not parallel-safe (built-in list,
|
||||
bridge lookups, opted-in MCP tools). Path-scoped tools join a run only when
|
||||
their paths don't conflict with the run's reservations: reader↔reader
|
||||
overlap commutes and stays parallel; any overlap involving a writer closes
|
||||
the run so the call starts a NEW run after the conflicting one lands.
|
||||
``search_files`` reserves its root (default ``.``) as a reader. Parallel
|
||||
runs shorter than two calls demote to sequential (the sequential executor
|
||||
owns the richer inline dispatch); adjacent sequential segments merge.
|
||||
"""
|
||||
segments: list[list] = [] # [kind, calls] pairs, normalized to tuples on return
|
||||
segments: List[tuple] = []
|
||||
current: list = []
|
||||
# (canonical_path, is_writer) reservations for the current parallel run.
|
||||
reserved_paths: list[tuple[Path, bool]] = []
|
||||
reserved_paths: list[tuple[Path, bool]] = [] # (canonical_path, is_writer) for the current run
|
||||
|
||||
def _close_parallel() -> None:
|
||||
nonlocal current, reserved_paths
|
||||
if current:
|
||||
segments.append(["parallel", current])
|
||||
current = []
|
||||
reserved_paths = []
|
||||
if len(current) >= 2:
|
||||
segments.append(("parallel", current))
|
||||
elif current:
|
||||
_extend_sequential(current)
|
||||
current = []
|
||||
reserved_paths = []
|
||||
|
||||
def _add_sequential(tc) -> None:
|
||||
_close_parallel()
|
||||
def _extend_sequential(calls: list) -> None:
|
||||
if segments and segments[-1][0] == "sequential":
|
||||
segments[-1][1].append(tc)
|
||||
segments[-1][1].extend(calls)
|
||||
else:
|
||||
segments.append(["sequential", [tc]])
|
||||
segments.append(("sequential", list(calls)))
|
||||
|
||||
for tool_call in tool_calls:
|
||||
tool_name = tool_call.function.name
|
||||
|
||||
if tool_name in _NEVER_PARALLEL_TOOLS:
|
||||
_add_sequential(tool_call)
|
||||
admission = _batch_admission(tool_call, execution_cwd)
|
||||
if admission is None:
|
||||
_close_parallel()
|
||||
_extend_sequential([tool_call])
|
||||
continue
|
||||
|
||||
try:
|
||||
function_args = json.loads(tool_call.function.arguments)
|
||||
except Exception:
|
||||
_raw = tool_call.function.arguments
|
||||
logging.debug(
|
||||
"Could not parse args for %s — treating as sequential barrier; raw=%s",
|
||||
tool_name,
|
||||
_raw[:200] if isinstance(_raw, str) else repr(_raw)[:200],
|
||||
)
|
||||
_add_sequential(tool_call)
|
||||
continue
|
||||
if not isinstance(function_args, dict):
|
||||
logging.debug(
|
||||
"Non-dict args for %s (%s) — treating as sequential barrier",
|
||||
tool_name,
|
||||
type(function_args).__name__,
|
||||
)
|
||||
_add_sequential(tool_call)
|
||||
continue
|
||||
|
||||
# Bridge unwrap: admission is decided on the UNDERLYING tool, not on
|
||||
# the literal wrapper name the model emitted. Read-only bridge
|
||||
# lookups (tool_search / tool_describe) are parallel-safe as-is.
|
||||
effective_name, effective_args = _peel_bridge_call(tool_name, function_args)
|
||||
|
||||
if effective_name in _NEVER_PARALLEL_TOOLS:
|
||||
_add_sequential(tool_call)
|
||||
continue
|
||||
|
||||
if effective_name in _PATH_SCOPED_TOOLS:
|
||||
scoped_paths = _extract_parallel_scope_paths(
|
||||
effective_name, effective_args, execution_cwd=execution_cwd
|
||||
)
|
||||
if not scoped_paths:
|
||||
_add_sequential(tool_call)
|
||||
continue
|
||||
is_writer = effective_name in _PATH_SCOPED_WRITERS
|
||||
if any(
|
||||
(is_writer or existing_is_writer)
|
||||
and _paths_overlap(scoped_path, existing)
|
||||
for scoped_path in scoped_paths
|
||||
for existing, existing_is_writer in reserved_paths
|
||||
):
|
||||
# Same-subtree conflict inside this run: close it so this
|
||||
# call starts a fresh run AFTER the conflicting one lands.
|
||||
# Reader↔reader overlap never conflicts — concurrent reads
|
||||
# of the same subtree commute.
|
||||
_close_parallel()
|
||||
reserved_paths.extend((p, is_writer) for p in scoped_paths)
|
||||
current.append(tool_call)
|
||||
continue
|
||||
|
||||
if (
|
||||
effective_name in _PARALLEL_SAFE_TOOLS
|
||||
or effective_name in _PARALLEL_SAFE_BRIDGE_LOOKUPS
|
||||
or _is_mcp_tool_parallel_safe(effective_name)
|
||||
_name, scoped_paths, is_writer = admission
|
||||
if any(
|
||||
(is_writer or existing_is_writer) and _paths_overlap(scoped_path, existing)
|
||||
for scoped_path in scoped_paths
|
||||
for existing, existing_is_writer in reserved_paths
|
||||
):
|
||||
current.append(tool_call)
|
||||
continue
|
||||
|
||||
_add_sequential(tool_call)
|
||||
_close_parallel()
|
||||
reserved_paths.extend((p, is_writer) for p in scoped_paths)
|
||||
current.append(tool_call)
|
||||
|
||||
_close_parallel()
|
||||
|
||||
normalized: list[list] = []
|
||||
for kind, calls in segments:
|
||||
if kind == "parallel" and len(calls) < 2:
|
||||
kind = "sequential"
|
||||
if normalized and normalized[-1][0] == "sequential" and kind == "sequential":
|
||||
normalized[-1][1].extend(calls)
|
||||
else:
|
||||
normalized.append([kind, calls])
|
||||
return [(kind, calls) for kind, calls in normalized]
|
||||
return segments
|
||||
|
||||
|
||||
def _should_parallelize_tool_batch(tool_calls) -> bool:
|
||||
"""Return True when the WHOLE tool-call batch is safe to run concurrently.
|
||||
|
||||
Thin view over ``_plan_tool_batch_segments`` kept for callers/tests that
|
||||
only care about the homogeneous case: True iff the planner produces a
|
||||
single all-parallel segment.
|
||||
"""
|
||||
"""True iff the planner yields a single all-parallel segment for the WHOLE batch."""
|
||||
if len(tool_calls) <= 1:
|
||||
return False
|
||||
segments = _plan_tool_batch_segments(tool_calls)
|
||||
@@ -294,19 +207,13 @@ def _should_parallelize_tool_batch(tool_calls) -> bool:
|
||||
|
||||
|
||||
def _canonical_path(raw_path: str, execution_cwd: Optional[Path] = None) -> Path:
|
||||
"""Return a canonical, OS-aware path for overlap detection.
|
||||
|
||||
Uses ``os.path.realpath`` to resolve symlinks on existing path components
|
||||
and ``os.path.normcase`` for case-insensitive platforms (Windows).
|
||||
Falls back to ``Path.cwd()`` when *execution_cwd* is not supplied.
|
||||
"""
|
||||
"""Canonical, OS-aware path for overlap detection (realpath for symlinks on
|
||||
existing components, normcase for case-insensitive platforms); relative
|
||||
paths resolve against *execution_cwd* or ``Path.cwd()``."""
|
||||
expanded = Path(raw_path).expanduser()
|
||||
base = execution_cwd if execution_cwd is not None else Path.cwd()
|
||||
candidate = expanded if expanded.is_absolute() else base / expanded
|
||||
# realpath resolves symlinks on path components that exist; for
|
||||
# not-yet-created files it canonicalises as far as possible.
|
||||
resolved = os.path.normcase(os.path.realpath(os.path.abspath(str(candidate))))
|
||||
return Path(resolved)
|
||||
return Path(os.path.normcase(os.path.realpath(os.path.abspath(str(candidate)))))
|
||||
|
||||
|
||||
def _extract_parallel_scope_paths(
|
||||
@@ -314,17 +221,12 @@ def _extract_parallel_scope_paths(
|
||||
function_args: dict,
|
||||
execution_cwd: Optional[Path] = None,
|
||||
) -> List[Path]:
|
||||
"""Return every canonical path this call reserves for overlap checks.
|
||||
"""Every canonical path this call reserves for overlap checks.
|
||||
|
||||
*execution_cwd* should be the working directory that the tool will
|
||||
actually use at runtime. When omitted the process cwd is used,
|
||||
which may differ from the tool execution environment on some
|
||||
platforms (e.g. WSL, sandboxed sub-processes).
|
||||
|
||||
For ``patch`` in V4A ``mode=patch``, scope comes from patch-body
|
||||
``*** Update/Add/Delete/Move File:`` headers (not a possibly-decoy
|
||||
``path=``). An empty result means the planner cannot determine the
|
||||
scope and must treat the call as a sequential barrier.
|
||||
*execution_cwd* should be the cwd the tool will actually use (may differ
|
||||
from the process cwd on WSL / sandboxed backends). For V4A ``patch`` the
|
||||
scope comes from patch-body file headers. An empty result means the scope
|
||||
is unknown and the planner must treat the call as a sequential barrier.
|
||||
"""
|
||||
if tool_name not in _PATH_SCOPED_TOOLS:
|
||||
return []
|
||||
@@ -337,25 +239,15 @@ def _extract_parallel_scope_paths(
|
||||
if isinstance(raw_path, str) and raw_path.strip():
|
||||
raw_paths.append(raw_path)
|
||||
elif tool_name == "search_files":
|
||||
# ``search_files`` defaults its search root to the cwd when
|
||||
# ``path`` is omitted — reserve that root rather than falling
|
||||
# back to a sequential barrier (an empty result here would
|
||||
# demote every bare search to a barrier and destroy read
|
||||
# parallelism).
|
||||
# search_files defaults its root to the cwd; reserve that rather than
|
||||
# demoting every bare search to a barrier.
|
||||
raw_paths.append(".")
|
||||
|
||||
scoped: List[Path] = []
|
||||
seen: set[str] = set()
|
||||
for raw in raw_paths:
|
||||
if not isinstance(raw, str) or not raw.strip():
|
||||
continue
|
||||
canonical = _canonical_path(raw, execution_cwd)
|
||||
key = str(canonical)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
scoped.append(canonical)
|
||||
return scoped
|
||||
# dict.fromkeys dedupes while preserving first-seen order.
|
||||
return list(dict.fromkeys(
|
||||
_canonical_path(raw, execution_cwd)
|
||||
for raw in raw_paths if isinstance(raw, str) and raw.strip()
|
||||
))
|
||||
|
||||
|
||||
def _extract_parallel_scope_path(
|
||||
@@ -363,41 +255,24 @@ def _extract_parallel_scope_path(
|
||||
function_args: dict,
|
||||
execution_cwd: Optional[Path] = None,
|
||||
) -> Optional[Path]:
|
||||
"""Return the primary canonical file target for path-scoped tools.
|
||||
|
||||
Thin view over ``_extract_parallel_scope_paths`` kept for callers/tests
|
||||
that only need a single representative path. For multi-file V4A
|
||||
patches this is the first header target.
|
||||
"""
|
||||
scoped = _extract_parallel_scope_paths(
|
||||
tool_name, function_args, execution_cwd=execution_cwd
|
||||
)
|
||||
"""Primary canonical target (first header target for multi-file V4A patches), or None."""
|
||||
scoped = _extract_parallel_scope_paths(tool_name, function_args, execution_cwd=execution_cwd)
|
||||
return scoped[0] if scoped else None
|
||||
|
||||
|
||||
def _paths_overlap(left: Path, right: Path) -> bool:
|
||||
"""Return True when two paths may refer to the same subtree.
|
||||
|
||||
Both *left* and *right* must already be canonical (as returned by
|
||||
``_extract_parallel_scope_paths`` / ``_canonical_path``) so that
|
||||
symlink aliases and case differences are already normalised.
|
||||
"""
|
||||
"""True when two already-canonical paths may refer to the same subtree."""
|
||||
left_parts = left.parts
|
||||
right_parts = right.parts
|
||||
if not left_parts or not right_parts:
|
||||
# Empty paths shouldn't reach here (guarded upstream), but be safe.
|
||||
# Empty paths are guarded upstream; only two non-empty equal prefixes overlap.
|
||||
return bool(left_parts) == bool(right_parts) and bool(left_parts)
|
||||
common_len = min(len(left_parts), len(right_parts))
|
||||
return left_parts[:common_len] == right_parts[:common_len]
|
||||
|
||||
|
||||
def _is_multimodal_tool_result(value: Any) -> bool:
|
||||
"""True if the value is a multimodal tool result envelope.
|
||||
|
||||
Multimodal handlers (e.g. tools/computer_use) return a dict with
|
||||
`_multimodal=True`, a `content` key holding OpenAI-style content
|
||||
parts, and an optional `text_summary` for string-only fallbacks.
|
||||
"""
|
||||
"""True for the multimodal envelope: dict with ``_multimodal=True`` and a ``content`` list."""
|
||||
return (
|
||||
isinstance(value, dict)
|
||||
and value.get("_multimodal") is True
|
||||
@@ -405,23 +280,17 @@ def _is_multimodal_tool_result(value: Any) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _multimodal_text_summary(value: Any) -> str:
|
||||
"""Extract a plain text view of a multimodal tool result.
|
||||
def _is_text_part(p: Any) -> bool:
|
||||
return isinstance(p, dict) and p.get("type") == "text"
|
||||
|
||||
Used wherever downstream code needs a string — logging, previews,
|
||||
persistence size heuristics, fall-back content for providers that
|
||||
don't support multipart tool messages.
|
||||
"""
|
||||
|
||||
def _multimodal_text_summary(value: Any) -> str:
|
||||
"""Plain-text view of a tool result (logging, previews, string-only providers)."""
|
||||
if _is_multimodal_tool_result(value):
|
||||
if value.get("text_summary"):
|
||||
return str(value["text_summary"])
|
||||
parts = []
|
||||
for p in value.get("content") or []:
|
||||
if isinstance(p, dict) and p.get("type") == "text":
|
||||
parts.append(str(p.get("text", "")))
|
||||
if parts:
|
||||
return "\n".join(parts)
|
||||
return "[multimodal tool result]"
|
||||
parts = [str(p.get("text", "")) for p in value.get("content") or [] if _is_text_part(p)]
|
||||
return "\n".join(parts) if parts else "[multimodal tool result]"
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
@@ -431,17 +300,12 @@ def _multimodal_text_summary(value: Any) -> str:
|
||||
|
||||
|
||||
def _append_subdir_hint_to_multimodal(value: Dict[str, Any], hint: str) -> None:
|
||||
"""Mutate a multimodal tool-result envelope to append a subdir hint.
|
||||
|
||||
The hint is added to the first text part so the model sees it; image
|
||||
parts are left untouched. `text_summary` is also updated for
|
||||
string-fallback callers.
|
||||
"""
|
||||
"""Append a subdir hint to the envelope's first text part (and ``text_summary``) in place."""
|
||||
if not _is_multimodal_tool_result(value):
|
||||
return
|
||||
parts = value.get("content") or []
|
||||
for p in parts:
|
||||
if isinstance(p, dict) and p.get("type") == "text":
|
||||
if _is_text_part(p):
|
||||
p["text"] = str(p.get("text", "")) + hint
|
||||
break
|
||||
else:
|
||||
@@ -451,52 +315,34 @@ def _append_subdir_hint_to_multimodal(value: Dict[str, Any], hint: str) -> None:
|
||||
value["text_summary"] = value["text_summary"] + hint
|
||||
|
||||
|
||||
def _extract_file_mutation_targets(tool_name: str, args: Dict[str, Any]) -> List[str]:
|
||||
"""Return the file paths a ``write_file`` or ``patch`` call is targeting.
|
||||
# ``\s*`` (not ``\s+``) after ``***`` matches patch_parser / file_tools, which
|
||||
# accept ``***Update File:`` with no space.
|
||||
_V4A_FILE_HEADER = re.compile(r'^\*\*\*\s*(?:Update|Add|Delete)\s+File:\s*(.+)$', re.MULTILINE)
|
||||
_V4A_MOVE_HEADER = re.compile(r'^\*\*\*\s*Move\s+File:\s*(.+?)\s*->\s*(.+)$', re.MULTILINE)
|
||||
|
||||
For ``write_file`` and ``patch`` in replace mode this is just ``args["path"]``.
|
||||
For ``patch`` in V4A patch mode we parse the patch content for
|
||||
``*** Update File:`` / ``*** Add File:`` / ``*** Delete File:`` headers so
|
||||
the verifier can track each file in a multi-file patch separately.
|
||||
|
||||
def _extract_file_mutation_targets(tool_name: str, args: Dict[str, Any]) -> List[str]:
|
||||
"""File paths a ``write_file`` / ``patch`` call targets.
|
||||
|
||||
Replace mode uses ``args["path"]``; V4A patch mode parses the
|
||||
``*** Update/Add/Delete/Move File:`` headers so each file in a multi-file
|
||||
patch is tracked separately.
|
||||
"""
|
||||
if tool_name not in _FILE_MUTATING_TOOLS:
|
||||
return []
|
||||
if tool_name == "write_file":
|
||||
p = args.get("path")
|
||||
return [str(p)] if p else []
|
||||
# tool_name == "patch"
|
||||
mode = args.get("mode") or "replace"
|
||||
mode = "replace" if tool_name == "write_file" else (args.get("mode") or "replace")
|
||||
if mode == "replace":
|
||||
p = args.get("path")
|
||||
return [str(p)] if p else []
|
||||
if mode == "patch":
|
||||
body = args.get("patch") or ""
|
||||
if not isinstance(body, str) or not body:
|
||||
return []
|
||||
paths: List[str] = []
|
||||
# ``\s*`` (not ``\s+``) after ``***`` matches patch_parser / file_tools:
|
||||
# they accept ``***Update File:`` with no space after the asterisks.
|
||||
for _m in re.finditer(
|
||||
r'^\*\*\*\s*(?:Update|Add|Delete)\s+File:\s*(.+)$',
|
||||
body,
|
||||
re.MULTILINE,
|
||||
):
|
||||
p = _m.group(1).strip()
|
||||
if p:
|
||||
paths.append(p)
|
||||
for _m in re.finditer(
|
||||
r'^\*\*\*\s*Move\s+File:\s*(.+?)\s*->\s*(.+)$',
|
||||
body,
|
||||
re.MULTILINE,
|
||||
):
|
||||
src = _m.group(1).strip()
|
||||
dst = _m.group(2).strip()
|
||||
if src:
|
||||
paths.append(src)
|
||||
if dst:
|
||||
paths.append(dst)
|
||||
return paths
|
||||
return []
|
||||
if mode != "patch":
|
||||
return []
|
||||
body = args.get("patch") or ""
|
||||
if not isinstance(body, str) or not body:
|
||||
return []
|
||||
paths = [m.group(1).strip() for m in _V4A_FILE_HEADER.finditer(body)]
|
||||
for m in _V4A_MOVE_HEADER.finditer(body):
|
||||
paths.extend((m.group(1).strip(), m.group(2).strip()))
|
||||
return [p for p in paths if p]
|
||||
|
||||
|
||||
def _extract_landed_file_mutation_paths(
|
||||
@@ -504,7 +350,8 @@ def _extract_landed_file_mutation_paths(
|
||||
args: Dict[str, Any],
|
||||
result: Any,
|
||||
) -> List[str]:
|
||||
"""Return the concrete file paths a successful mutation reports."""
|
||||
"""Concrete file paths a successful mutation reports (``files_modified`` /
|
||||
``resolved_path`` in the JSON result), falling back to the declared targets."""
|
||||
targets = _extract_file_mutation_targets(tool_name, args)
|
||||
if tool_name not in _FILE_MUTATING_TOOLS or not isinstance(result, str):
|
||||
return targets
|
||||
@@ -516,28 +363,17 @@ def _extract_landed_file_mutation_paths(
|
||||
return targets
|
||||
|
||||
files = data.get("files_modified")
|
||||
if isinstance(files, list):
|
||||
landed = [str(p) for p in files if p]
|
||||
if landed:
|
||||
return landed
|
||||
|
||||
landed = [str(p) for p in files if p] if isinstance(files, list) else []
|
||||
if landed:
|
||||
return landed
|
||||
resolved = data.get("resolved_path")
|
||||
if resolved:
|
||||
return [str(resolved)]
|
||||
|
||||
return targets
|
||||
return [str(resolved)] if resolved else targets
|
||||
|
||||
|
||||
def _extract_error_preview(result: Any, max_len: int = 180) -> str:
|
||||
"""Pull a one-line error summary out of a tool result for footer display."""
|
||||
"""One-line error summary of a tool result for footer display."""
|
||||
text = _multimodal_text_summary(result) if result is not None else ""
|
||||
if not isinstance(text, str):
|
||||
try:
|
||||
text = str(text)
|
||||
except Exception:
|
||||
return ""
|
||||
# Try to parse JSON and pull the ``error`` field — tool handlers return
|
||||
# ``{"success": false, "error": "..."}``; raw string wins if parse fails.
|
||||
# Handlers return {"success": false, "error": "..."}; the raw string wins if parse fails.
|
||||
stripped = text.strip()
|
||||
if stripped.startswith("{"):
|
||||
try:
|
||||
@@ -546,7 +382,6 @@ def _extract_error_preview(result: Any, max_len: int = 180) -> str:
|
||||
text = data["error"]
|
||||
except Exception:
|
||||
pass
|
||||
# Collapse whitespace, trim to max_len.
|
||||
text = " ".join(text.split())
|
||||
if len(text) > max_len:
|
||||
text = text[: max_len - 1] + "…"
|
||||
@@ -554,25 +389,20 @@ def _extract_error_preview(result: Any, max_len: int = 180) -> str:
|
||||
|
||||
|
||||
def _trajectory_normalize_msg(msg: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Strip image blobs from a message for trajectory saving.
|
||||
|
||||
Returns a shallow copy with multimodal tool results replaced by their
|
||||
text_summary, and image parts in content lists replaced by
|
||||
`[screenshot]` placeholders. Keeps the message schema otherwise intact.
|
||||
"""
|
||||
"""Shallow copy with image blobs stripped for trajectory saving: multimodal
|
||||
results become their text summary, image parts become ``[screenshot]``."""
|
||||
if not isinstance(msg, dict):
|
||||
return msg
|
||||
content = msg.get("content")
|
||||
if _is_multimodal_tool_result(content):
|
||||
return {**msg, "content": _multimodal_text_summary(content)}
|
||||
if isinstance(content, list):
|
||||
cleaned = []
|
||||
for p in content:
|
||||
if isinstance(p, dict) and p.get("type") in {"image", "image_url", "input_image"}:
|
||||
cleaned.append({"type": "text", "text": "[screenshot]"})
|
||||
else:
|
||||
cleaned.append(p)
|
||||
return {**msg, "content": cleaned}
|
||||
return {**msg, "content": [
|
||||
{"type": "text", "text": "[screenshot]"}
|
||||
if isinstance(p, dict) and p.get("type") in {"image", "image_url", "input_image"}
|
||||
else p
|
||||
for p in content
|
||||
]}
|
||||
return msg
|
||||
|
||||
|
||||
@@ -590,33 +420,19 @@ def make_tool_result_message(
|
||||
*,
|
||||
effect_disposition: str | None = None,
|
||||
) -> dict:
|
||||
"""Build a tool-result message dict with both the OpenAI-format ``name``
|
||||
field (required by the wire format and provider adapters) and the internal
|
||||
``tool_name`` field (written to the session DB messages table).
|
||||
"""Build a tool-result message with the OpenAI ``name`` field (wire format)
|
||||
and the internal ``tool_name`` field (session DB).
|
||||
|
||||
Content from high-risk tools (``web_extract``, ``web_search``, ``browser_*``,
|
||||
``mcp_*``) gets wrapped in semantic delimiters telling the model the content
|
||||
is untrusted data, not instructions. This is the architectural defense
|
||||
against indirect prompt injection from poisoned web pages, GitHub issues,
|
||||
and MCP responses — it changes how the model interprets the content rather
|
||||
than relying on regex pattern matching catching every payload.
|
||||
|
||||
Wrapping applies to plain string content and to multimodal content
|
||||
lists (``[{"type": "text", "text": "..."}, {"type": "image_url", ...}]``):
|
||||
each text-type part is wrapped individually using the same rules as plain
|
||||
string content (short text passes through unchanged; longer text is
|
||||
neutralized and framed). Non-text parts (e.g. image_url) are preserved.
|
||||
The outer list itself is rebuilt rather than returned by identity, so
|
||||
callers should compare by value, not by ``is``.
|
||||
Content from high-risk tools (web_extract, web_search, browser_*, mcp_*) is
|
||||
wrapped in untrusted-data delimiters — the architectural defense against
|
||||
indirect prompt injection; see ``_maybe_wrap_untrusted``.
|
||||
"""
|
||||
# Keep the constructor safe for every caller, including replay recovery
|
||||
# paths that do not go through the live executor's canonical-id helper.
|
||||
# Replay-recovery callers bypass the executor's canonical-id helper, so normalize here too.
|
||||
tool_call_id = _normalize_tool_call_id(tool_call_id)
|
||||
|
||||
# Order matters: detect provider-side elision on the RAW content and
|
||||
# append the notice first, THEN wrap — so the notice lives inside the
|
||||
# untrusted block next to the data it describes, appended exactly once
|
||||
# at construction time (cache-safe).
|
||||
# Order matters: detect elision on the RAW content and append the notice
|
||||
# first, THEN wrap, so the notice sits inside the untrusted block next to
|
||||
# the data it describes — once, at construction time (cache-safe).
|
||||
wrapped = _maybe_wrap_untrusted(name, _maybe_append_elision_notice(name, content))
|
||||
message = stamp_message_timestamp({
|
||||
"role": "tool",
|
||||
@@ -637,63 +453,41 @@ def make_tool_result_message(
|
||||
return message
|
||||
|
||||
|
||||
# Tools whose results carry attacker-controllable content. Wrapping their
|
||||
# string output in ``<untrusted_tool_result>`` delimiters tells the model the
|
||||
# payload is data, not instructions — the architectural piece of the
|
||||
# promptware defense. Skipped for short outputs (under 32 chars) where the
|
||||
# overhead of the wrapper outweighs any indirect-injection risk.
|
||||
_UNTRUSTED_TOOL_NAMES = frozenset({
|
||||
"web_extract",
|
||||
"web_search",
|
||||
})
|
||||
|
||||
_UNTRUSTED_TOOL_PREFIXES = (
|
||||
"browser_",
|
||||
"mcp_",
|
||||
)
|
||||
|
||||
# Tools whose results carry attacker-controllable content. Short outputs
|
||||
# (under 32 chars) skip wrapping: the overhead outweighs any injection risk.
|
||||
_UNTRUSTED_TOOL_NAMES = frozenset({"web_extract", "web_search"})
|
||||
_UNTRUSTED_TOOL_PREFIXES = ("browser_", "mcp_")
|
||||
_UNTRUSTED_WRAP_MIN_CHARS = 32
|
||||
|
||||
# Matches the delimiter token in any case so attacker content can't forge or
|
||||
# prematurely close the boundary with a differently-cased variant the model
|
||||
# would still read as a tag (e.g. ``</UNTRUSTED_TOOL_RESULT>``).
|
||||
# Case-insensitive so attacker content can't forge or prematurely close the
|
||||
# boundary with a differently-cased tag the model would still read as one.
|
||||
_DELIMITER_TOKEN_RE = re.compile(r"untrusted_tool_result", re.IGNORECASE)
|
||||
|
||||
|
||||
def _is_untrusted_tool(name: Optional[str]) -> bool:
|
||||
if not name:
|
||||
return False
|
||||
if name in _UNTRUSTED_TOOL_NAMES:
|
||||
return True
|
||||
return any(name.startswith(p) for p in _UNTRUSTED_TOOL_PREFIXES)
|
||||
return bool(name) and (name in _UNTRUSTED_TOOL_NAMES or name.startswith(_UNTRUSTED_TOOL_PREFIXES))
|
||||
|
||||
|
||||
# --- Upstream-elision detection --------------------------------------------
|
||||
#
|
||||
# Some MCP servers elide data SERVER-SIDE and mark the elision inside the
|
||||
# payload itself (e.g. Composio: '...13 more items' inside a JSON array,
|
||||
# '"has_more": true', 'Complete response was large (N tokens). Full data
|
||||
# saved to sandbox in /mnt/files/...', 'data_preview' envelopes). Because the
|
||||
# result looks structurally complete, models treat the visible slice as the
|
||||
# whole dataset and falsely claim completeness. When one of these markers is
|
||||
# present, we append ONE compact notice at result-construction time — before
|
||||
# the message enters history, never mutated later, so prompt caching is safe.
|
||||
def _is_text_item(item: Any) -> bool:
|
||||
return _is_text_part(item) and isinstance(item.get("text"), str)
|
||||
|
||||
# Conservative patterns only: each one is an explicit provider-side "there is
|
||||
# more data than what you can see" signal, not a generic truncation heuristic.
|
||||
|
||||
# --- Upstream-elision detection ---
|
||||
# Some MCP servers elide data SERVER-SIDE and mark it inside the payload
|
||||
# ('...13 more items', '"has_more": true', 'saved to sandbox', 'data_preview'
|
||||
# envelopes). The result looks structurally complete, so models treat the
|
||||
# visible slice as the whole dataset. Conservative, explicit markers only —
|
||||
# not a generic truncation heuristic. One notice is appended at construction
|
||||
# time, never mutated later (prompt-cache safe).
|
||||
_UPSTREAM_ELISION_PATTERNS = (
|
||||
re.compile(r"\.\.\.\s*\d+\s+more\s+items?", re.IGNORECASE),
|
||||
re.compile(r'"has_more"\s*:\s*true', re.IGNORECASE),
|
||||
re.compile(r"saved to sandbox", re.IGNORECASE),
|
||||
re.compile(r"data_preview", re.IGNORECASE),
|
||||
)
|
||||
|
||||
# Results smaller than this can't meaningfully hide an elided enumeration —
|
||||
# skip the scan entirely so tiny results pay nothing.
|
||||
# Tiny results can't hide an elided enumeration; markers for the sizes that
|
||||
# matter (20-50K) are always inside the first 64KB.
|
||||
_ELISION_SCAN_MIN_CHARS = 1_000
|
||||
|
||||
# Bound the regex scan: markers appear near the elided structure, which for
|
||||
# the payload sizes that matter (20-50K) is always inside the first 64KB.
|
||||
_ELISION_SCAN_MAX_CHARS = 65_536
|
||||
|
||||
_UPSTREAM_ELISION_NOTICE = (
|
||||
@@ -704,53 +498,31 @@ _UPSTREAM_ELISION_NOTICE = (
|
||||
|
||||
|
||||
def _detect_upstream_elision(content: Any) -> bool:
|
||||
"""True when a string tool result carries provider-side elision markers.
|
||||
|
||||
Cheap and safe by construction: non-string content is never scanned,
|
||||
results under ``_ELISION_SCAN_MIN_CHARS`` short-circuit, and the regex
|
||||
scan is capped at the first ``_ELISION_SCAN_MAX_CHARS`` chars.
|
||||
"""
|
||||
if not isinstance(content, str):
|
||||
return False
|
||||
if len(content) < _ELISION_SCAN_MIN_CHARS:
|
||||
"""True when a string result carries provider-side elision markers (bounded scan)."""
|
||||
if not isinstance(content, str) or len(content) < _ELISION_SCAN_MIN_CHARS:
|
||||
return False
|
||||
window = content[:_ELISION_SCAN_MAX_CHARS]
|
||||
return any(p.search(window) for p in _UPSTREAM_ELISION_PATTERNS)
|
||||
|
||||
|
||||
def _maybe_append_elision_notice(name: str, content: Any) -> Any:
|
||||
"""Append the incompleteness notice to untrusted string results that
|
||||
embed upstream elision markers. Returns ``content`` unchanged otherwise.
|
||||
|
||||
Runs on the RAW result before untrusted-wrapping so the notice sits with
|
||||
the data it describes, and only at result-construction time (cache-safe).
|
||||
"""
|
||||
if not _is_untrusted_tool(name):
|
||||
return content
|
||||
if _detect_upstream_elision(content):
|
||||
"""Append the incompleteness notice to untrusted string results with elision markers."""
|
||||
if _is_untrusted_tool(name) and _detect_upstream_elision(content):
|
||||
return content + _UPSTREAM_ELISION_NOTICE
|
||||
return content
|
||||
|
||||
|
||||
def _tool_output_risk_metadata(name: str, content: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Classify textual attacker-controlled output without retaining a copy.
|
||||
"""Internal-only advisory classification of attacker-controlled output.
|
||||
|
||||
The advisory metadata is internal-only. It records deterministic finding
|
||||
identifiers, never blocks or redacts the normal result, and deliberately
|
||||
omits raw scanned text.
|
||||
Records deterministic finding ids, never blocks or redacts, and omits the scanned text.
|
||||
"""
|
||||
if not _is_untrusted_tool(name):
|
||||
return None
|
||||
if isinstance(content, str):
|
||||
text_parts = [content]
|
||||
elif isinstance(content, list):
|
||||
text_parts = [
|
||||
item["text"]
|
||||
for item in content
|
||||
if isinstance(item, dict)
|
||||
and item.get("type") == "text"
|
||||
and isinstance(item.get("text"), str)
|
||||
]
|
||||
text_parts = [item["text"] for item in content if _is_text_item(item)]
|
||||
if not text_parts:
|
||||
return None
|
||||
else:
|
||||
@@ -769,40 +541,20 @@ def _tool_output_risk_metadata(name: str, content: Any) -> Optional[Dict[str, An
|
||||
|
||||
|
||||
def _neutralize_delimiters(content: str) -> str:
|
||||
"""Defang any literal ``untrusted_tool_result`` delimiter embedded in
|
||||
attacker-controlled content so it can't break out of the wrapper.
|
||||
|
||||
Without this, a poisoned web page / GitHub issue / MCP response that
|
||||
contains ``</untrusted_tool_result>`` would close the trust boundary early
|
||||
— everything the attacker writes after it then reads as trusted instructions
|
||||
outside the block. Replacing the underscores with hyphens leaves the text
|
||||
readable but means it no longer matches the real (underscore) delimiter.
|
||||
"""
|
||||
"""Defang embedded ``untrusted_tool_result`` tokens so poisoned content
|
||||
can't close the trust boundary early (hyphens keep it readable but non-matching)."""
|
||||
return _DELIMITER_TOKEN_RE.sub("untrusted-tool-result", content)
|
||||
|
||||
|
||||
def _maybe_wrap_untrusted(name: str, content: Any) -> Any:
|
||||
"""Wrap content from high-risk tools in untrusted-data delimiters.
|
||||
"""Wrap high-risk tool content in untrusted-data delimiters.
|
||||
|
||||
Handles plain string content and multimodal content lists
|
||||
(``[{"type": "text", "text": "..."}, {"type": "image_url", ...}]``).
|
||||
Text parts inside a multimodal list are wrapped individually — the same
|
||||
rules as plain string content — so vision-capable adapters still receive
|
||||
a valid content list while an injection payload embedded in a text chunk
|
||||
is still marked as untrusted data. Non-text parts (image_url, etc.) are
|
||||
preserved unchanged. The outer list is rebuilt rather than returned by
|
||||
identity, so callers must compare by value, not by ``is``.
|
||||
|
||||
Returns ``content`` unchanged when:
|
||||
- the tool is not in the high-risk set
|
||||
- the content is neither a string nor a list (dict, None, …)
|
||||
- (string) the content is too short to be worth wrapping
|
||||
|
||||
Wrapped string content is always neutralized (any embedded delimiter token
|
||||
is defanged) and wrapped in exactly one well-formed block. There is no
|
||||
"already wrapped" fast-path: such a check is attacker-forgeable — content
|
||||
that merely starts with the opening tag would be returned with no data
|
||||
framing at all — so re-wrapping (harmlessly) is the safe choice.
|
||||
Strings are neutralized and wrapped in exactly one block; text parts of a
|
||||
multimodal list are wrapped individually (non-text parts preserved, outer
|
||||
list rebuilt — compare by value, not ``is``). Unchanged when the tool is
|
||||
not high-risk, the content is neither str nor list, or a string is too
|
||||
short. There is deliberately no "already wrapped" fast-path: it would be
|
||||
attacker-forgeable, so harmless re-wrapping is the safe choice.
|
||||
"""
|
||||
if not _is_untrusted_tool(name):
|
||||
return content
|
||||
@@ -821,11 +573,7 @@ def _maybe_wrap_untrusted(name: str, content: Any) -> Any:
|
||||
)
|
||||
if isinstance(content, list):
|
||||
return [
|
||||
{**item, "text": _maybe_wrap_untrusted(name, item["text"])}
|
||||
if isinstance(item, dict)
|
||||
and item.get("type") == "text"
|
||||
and isinstance(item.get("text"), str)
|
||||
else item
|
||||
{**item, "text": _maybe_wrap_untrusted(name, item["text"])} if _is_text_item(item) else item
|
||||
for item in content
|
||||
]
|
||||
return content
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,7 +2,7 @@
|
||||
|
||||
Two guards, both notice/re-prompt-only:
|
||||
|
||||
1. Identical-call loop breaker — ``ToolCallGuardrailController.observe_identical_call``
|
||||
1. Identical-call loop breaker — ``ToolCallGuardrailController.observe_call``
|
||||
appends a compact notice to the tool RESULT on the 3rd consecutive call
|
||||
with identical (tool, canonical args) AND an identical result. It never
|
||||
blocks execution, exempts legitimately-repeatable pollers, and resets on
|
||||
@@ -29,7 +29,7 @@ def _observe_n(controller, n, tool="web_search", args=None, result="same result"
|
||||
notices = []
|
||||
for _ in range(n):
|
||||
notices.append(
|
||||
controller.observe_identical_call(tool, args or {"query": "x"}, result)
|
||||
controller.observe_call(tool, args or {"query": "x"}, result).notice
|
||||
)
|
||||
return notices
|
||||
|
||||
@@ -58,18 +58,18 @@ def test_keeps_firing_past_threshold():
|
||||
def test_does_not_fire_when_arguments_differ():
|
||||
c = ToolCallGuardrailController()
|
||||
for i in range(5):
|
||||
notice = c.observe_identical_call(
|
||||
notice = c.observe_call(
|
||||
"web_search", {"query": f"q{i}"}, "same result"
|
||||
)
|
||||
).notice
|
||||
assert notice is None
|
||||
|
||||
|
||||
def test_does_not_fire_when_results_differ():
|
||||
c = ToolCallGuardrailController()
|
||||
for i in range(5):
|
||||
notice = c.observe_identical_call(
|
||||
notice = c.observe_call(
|
||||
"terminal", {"command": "poll-status"}, f"output {i}"
|
||||
)
|
||||
).notice
|
||||
assert notice is None
|
||||
|
||||
|
||||
@@ -77,7 +77,7 @@ def test_streak_resets_when_a_different_call_intervenes():
|
||||
c = ToolCallGuardrailController()
|
||||
assert _observe_n(c, 2)[-1] is None
|
||||
# Different tool breaks the consecutive streak.
|
||||
assert c.observe_identical_call("read_file", {"path": "/a"}, "data") is None
|
||||
assert c.observe_call("read_file", {"path": "/a"}, "data").notice is None
|
||||
# Two more of the original are a fresh streak of 2 — still no notice.
|
||||
assert all(n is None for n in _observe_n(c, 2))
|
||||
|
||||
@@ -85,16 +85,16 @@ def test_streak_resets_when_a_different_call_intervenes():
|
||||
def test_arg_canonicalization_ignores_key_order():
|
||||
c = ToolCallGuardrailController()
|
||||
r = "same"
|
||||
assert c.observe_identical_call("t", {"a": 1, "b": 2}, r) is None
|
||||
assert c.observe_identical_call("t", {"b": 2, "a": 1}, r) is None
|
||||
assert c.observe_identical_call("t", {"a": 1, "b": 2}, r) is not None
|
||||
assert c.observe_call("t", {"a": 1, "b": 2}, r).notice is None
|
||||
assert c.observe_call("t", {"b": 2, "a": 1}, r).notice is None
|
||||
assert c.observe_call("t", {"a": 1, "b": 2}, r).notice is not None
|
||||
|
||||
|
||||
def test_allowlisted_pollers_never_fire():
|
||||
c = ToolCallGuardrailController()
|
||||
for tool in ("process_manage", "vendor_get_result", "job_poll"):
|
||||
for _ in range(STALL_GUARD_IDENTICAL_CALL_THRESHOLD + 2):
|
||||
assert c.observe_identical_call(tool, {"id": "j1"}, "Generating") is None
|
||||
assert c.observe_call(tool, {"id": "j1"}, "Generating").notice is None
|
||||
|
||||
|
||||
def test_allowlist_membership_contract():
|
||||
@@ -328,13 +328,6 @@ def test_multimodal_content_never_stubbed_and_breaks_streak():
|
||||
assert c.observe_call("vision", args, _BIG, tool_call_id="c3").stub is None
|
||||
|
||||
|
||||
def test_observe_identical_call_backcompat_notice_still_fires():
|
||||
c = ToolCallGuardrailController()
|
||||
for _ in range(STALL_GUARD_IDENTICAL_CALL_THRESHOLD - 1):
|
||||
assert c.observe_identical_call("web_search", {"q": 1}, "r") is None
|
||||
assert c.observe_identical_call("web_search", {"q": 1}, "r") is not None
|
||||
|
||||
|
||||
def test_extract_persisted_path_round_trip():
|
||||
# The stub's spillover reference is parsed from the <persisted-output>
|
||||
# block that maybe_persist_tool_result builds — assert the round trip.
|
||||
|
||||
Reference in New Issue
Block a user