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:
Teknium
2026-09-02 10:36:07 -07:00
parent 2ad284b473
commit 4bfc55fbf0
5 changed files with 847 additions and 1901 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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.