diff --git a/agent/outbound_webhooks.py b/agent/outbound_webhooks.py index cd600832cb..7fcef908f1 100644 --- a/agent/outbound_webhooks.py +++ b/agent/outbound_webhooks.py @@ -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=`` 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=`` 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: - X-Hermes-Delivery: - X-Hermes-Signature-256: sha256= # 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( diff --git a/agent/shell_hooks.py b/agent/shell_hooks.py index 49dfa9301e..93667be8c8 100644 --- a/agent/shell_hooks.py +++ b/agent/shell_hooks.py @@ -1,135 +1,23 @@ -""" -Shell-script hooks bridge. +"""Shell-script hooks bridge. -Reads the ``hooks:`` block from ``cli-config.yaml``, prompts the user for -consent on first use of each ``(event, command)`` pair, and registers -callbacks on the existing plugin hook manager so every existing -``invoke_hook()`` site dispatches to the configured shell scripts — with -zero changes to call sites. +Reads the ``hooks:`` block from config, prompts for first-use consent per +``(event, command)`` pair, and registers callbacks on the plugin hook manager +so every existing ``invoke_hook()`` site dispatches to the configured scripts. -Design notes ------------- -* Python plugins and shell hooks compose naturally: both flow through - :func:`hermes_cli.plugins.invoke_hook` and its aggregators. Python - plugins are registered first (via ``discover_and_load()``) so their - block decisions win ties over shell-hook blocks. -* Subprocess execution uses ``shlex.split(os.path.expanduser(command))`` - with ``shell=False`` — no shell injection footguns. Users that need - pipes/redirection wrap their logic in a script. -* First-use consent is gated by the allowlist under - ``~/.hermes/shell-hooks-allowlist.json``. Non-TTY callers must pass - ``accept_hooks=True`` (resolved from ``--accept-hooks``, - ``HERMES_ACCEPT_HOOKS``, or ``hooks_auto_accept: true`` in config) - for registration to succeed without a prompt. -* Registration is idempotent — safe to invoke from both the CLI entry - point (``hermes_cli/main.py``) and the gateway entry point - (``gateway/run.py``). +Wire protocol — stdin JSON:: -Wire protocol -------------- -**stdin** (JSON, piped to the script):: + {"hook_event_name": ..., "tool_name": ..., "tool_input": {...}, + "session_id": ..., "cwd": ..., "extra": {...event-specific kwargs}} - { - "hook_event_name": "pre_tool_call", - "tool_name": "terminal", - "tool_input": {"command": "rm -rf /"}, - "session_id": "sess_abc123", - "cwd": "/home/user/project", - "extra": {...} # event-specific kwargs - } - -**stdout** (JSON, optional — anything else is ignored):: - - # Block a pre_tool_call (either shape accepted; normalised internally): - {"decision": "block", "reason": "Forbidden command"} # Claude-Code-style - {"action": "block", "message": "Forbidden command"} # Hermes-canonical - - # Inject context for pre_llm_call: - {"context": "Today is Friday"} - - # Modify tool input for pre_tool_call (Hermes-canonical): - {"action": "modify", "args": {"new_string": "fixed content"}} - - # Modify tool input for pre_tool_call (Claude-Code-style): - {"decision": "modify", "tool_input": {"new_string": "fixed content"}} - - # Silent no-op: - - -**exit codes** - -Exit code 2 from a ``pre_tool_call`` hook blocks the tool call even when -stdout carries no block JSON (Claude-Code / Cursor compatible). The block -message is taken from the stdout block JSON when present, then the first -400 characters of stderr, then a generic default. For events whose block -directive is not honored, exit 2 is logged at warning like any other -non-zero exit. All other non-zero exits log a warning and stdout is still -parsed normally. - -**failure semantics** - -Hooks fail *open* by default: a spawn error, timeout, or unparseable -stdout logs a warning and contributes nothing. A ``pre_tool_call`` entry -can opt into fail-*closed* semantics with ``fail_closed: true`` -(``failClosed`` also accepted for Cursor/Claude-Code config compat) — -spawn errors, timeouts, and malformed stdout then BLOCK the tool call -with ``hook failed closed: ``. Use this for -security-gating hooks (secret scanners, policy checks) where a crashed -hook must not silently allow the action. On non-blocking events -``fail_closed`` is ignored with a warning. - -Per-event ``extra`` keys -~~~~~~~~~~~~~~~~~~~~~~~~ - -The ``extra`` object contains every kwarg that is **not** one of the -top-level payload keys (``tool_name``, ``args``, ``session_id``, -``parent_session_id``). The tables below list the ``extra`` keys -emitted by each built-in hook site. - -``post_tool_call`` (emitted from ``model_tools.py``):: - - result – tool return value (serialised string) - status – "ok" | "error" | "blocked" - error_type – error category (e.g. "ValueError"), or None - error_message – human-readable error text, or None - duration_ms – wall-clock time in milliseconds - task_id – current task id (empty string if none) - tool_call_id – provider tool-call id - turn_id – current turn id - api_request_id – current API request id - middleware_trace – list of dicts from tool middleware chain - -``pre_tool_call`` (emitted from ``model_tools.py``):: - - task_id – current task id (empty string if none) - tool_call_id – provider tool-call id - turn_id – current turn id - api_request_id – current API request id - middleware_trace – list of dicts from tool middleware chain - -``on_session_start`` (emitted from ``agent/conversation_loop.py``):: - - model – model name (e.g. "claude-sonnet-4-20250514") - platform – platform identifier (e.g. "cli", "whatsapp") - -``on_session_end`` (emitted from ``agent/turn_finalizer.py``):: - - task_id – current task id - turn_id – current turn id - completed – bool, True when the turn produced a final response - interrupted – bool, True when the user interrupted - model – model name - platform – platform identifier - -``subagent_stop`` (emitted from ``tools/delegate_tool.py``):: - - parent_turn_id – parent agent's current turn id - child_session_id – child (subagent) session id - child_role – role string of the child agent - child_summary – summary of the child's work - child_status – exit status string (e.g. "success", "error") - tool_call_history – redacted tool name/input summary/byte counts/status list - duration_ms – wall-clock time of the child run in milliseconds +stdout JSON (optional): ``{"decision"|"action": "block", "reason"|"message": ...}`` +blocks a ``pre_tool_call``; ``{"action": "modify", "args": {...}}`` / +``{"decision": "modify", "tool_input": {...}}`` rewrites tool args; +``{"context": "..."}`` injects context for ``pre_llm_call``; ``pre_verify`` +accepts ``continue``/``block`` with a message. Exit code 2 blocks a +``pre_tool_call`` even without block JSON (Claude-Code / Cursor compatible). +Hooks fail open unless ``fail_closed: true`` (``failClosed`` accepted) on a +blocking-capable event, in which case spawn errors, timeouts and unparseable +stdout block with ``hook failed closed: ``. """ from __future__ import annotations @@ -139,13 +27,12 @@ import json import logging import os import re -import shlex import subprocess import sys import tempfile import threading import time -from contextlib import contextmanager +from contextlib import ExitStack, contextmanager from dataclasses import dataclass, field from datetime import datetime, timezone from pathlib import Path @@ -167,68 +54,95 @@ DEFAULT_TIMEOUT_SECONDS = 60 MAX_TIMEOUT_SECONDS = 300 ALLOWLIST_FILENAME = "shell-hooks-allowlist.json" _DEFAULT_BLOCK_MESSAGE = "Blocked by shell hook." - -# Exit code that signals "block this action" from a hook script, independent -# of stdout content. Claude Code / Cursor compatible. +# Exit code that signals "block this action" independent of stdout (Claude Code / Cursor). BLOCK_EXIT_CODE = 2 - -# Events whose block directive is actually honored downstream (see -# hermes_cli.plugins.get_pre_tool_call_block_message / _get_pre_tool_call_ -# directive_details). Exit-code-2 blocking and ``fail_closed`` only make -# sense for these. +# Events whose block directive is honored downstream; exit-2 blocking and fail_closed only apply here. _BLOCKING_EVENTS = frozenset({"pre_tool_call"}) - -# Cap on stderr excerpt reused as a block message. +_TOOL_EVENTS = frozenset({"pre_tool_call", "post_tool_call"}) _STDERR_MESSAGE_LIMIT = 400 +_TRUTHY = {"1", "true", "yes", "on"} +# kwargs promoted to top-level payload keys; everything else lands under ``extra``. +_TOP_LEVEL_PAYLOAD_KEYS = {"tool_name", "args", "session_id", "parent_session_id"} - -# (home, event, matcher, command) tuples that have been wired to the plugin -# manager in the current process. Matcher is part of the key because -# the same script can legitimately register for different matchers under -# the same event (e.g. one entry per tool the user wants to gate). 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 hook triples without the first profile's registration -# silently shadowing the rest. -# Second registration attempts for the exact same tuple become no-ops -# so the CLI and gateway can both call register_from_config() safely. +# (home, event, matcher, command) tuples wired to the plugin manager in this process. +# Matcher is in the key: one script may legitimately register per-tool under one event. +# Home is part of the key so multiplexed-gateway profiles (each with their own +# plugin manager) can register identical triples without shadowing each other. _registered: Set[Tuple[str, str, Optional[str], str]] = set() _registered_lock = threading.Lock() - -# Intra-process lock for allowlist read-modify-write on platforms that -# lack ``fcntl`` (non-POSIX). Kept separate from ``_registered_lock`` -# because ``register_from_config`` already holds ``_registered_lock`` when -# it triggers ``_record_approval`` — reusing it here would self-deadlock -# (``threading.Lock`` is non-reentrant). POSIX callers use the sibling -# ``.lock`` file via ``fcntl.flock`` and bypass this. +# Non-POSIX fallback for allowlist read-modify-write. Separate from +# _registered_lock, which register_from_config already holds when it triggers +# _record_approval (threading.Lock is non-reentrant). _allowlist_write_lock = threading.Lock() -@dataclass -class ShellHookSpec: - """Parsed and validated representation of a single ``hooks:`` entry.""" +def _home_key() -> str: + return str(get_hermes_home().expanduser().resolve()) - event: str - command: str - matcher: Optional[str] = None - timeout: int = DEFAULT_TIMEOUT_SECONDS - fail_closed: bool = False - compiled_matcher: Optional[re.Pattern] = field(default=None, repr=False) + +def _forget_home_registrations(registry: Set[tuple], lock: threading.Lock) -> None: + """Drop the current home's idempotence keys only (shared with outbound webhooks). + + A force-reload in profile A must never drop profile B's live registration. + """ + home_key = _home_key() + with lock: + registry.difference_update({k for k in registry if k[0] == home_key}) + + +def _split(command: str) -> List[str]: + # Windows-safe: plain shlex.split eats backslashes in paths. + from hermes_cli._subprocess_compat import split_command_line + + return split_command_line(command) + + +def _entry_matches(e: Any, event: Optional[str], command: str) -> bool: + return ( + isinstance(e, dict) + and (event is None or e.get("event") == event) + and e.get("command") == command + ) + + +def _utc_now_iso() -> str: + return datetime.now(tz=timezone.utc).isoformat().replace("+00:00", "Z") + + +def _payload_fields(kwargs: Dict[str, Any]) -> Dict[str, Any]: + """Common stdin/POST payload fields (shared with outbound webhooks); key order is wire order.""" + 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 = "" + return { + "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, + } + + +class _ToolMatcherMixin: + """``matcher`` regex handling shared by shell-hook specs and outbound webhook targets.""" + + _MATCHER_KIND = "shell hook" + matcher: Optional[str] + compiled_matcher: Optional[re.Pattern] def __post_init__(self) -> None: - # Strip whitespace introduced by YAML quirks (e.g. multi-line string - # folding) — a matcher of " terminal" would otherwise silently fail - # to match "terminal" without any diagnostic. + # Strip YAML folding whitespace — " terminal" would silently never match. if isinstance(self.matcher, str): - stripped = self.matcher.strip() - self.matcher = stripped if stripped else None + self.matcher = self.matcher.strip() or None if self.matcher: try: self.compiled_matcher = re.compile(self.matcher) except re.error as exc: logger.warning( - "shell hook matcher %r is invalid (%s) — treating as " - "literal equality", self.matcher, exc, + "%s matcher %r is invalid (%s) — treating as " + "literal equality", self._MATCHER_KIND, self.matcher, exc, ) self.compiled_matcher = None @@ -239,42 +153,37 @@ class ShellHookSpec: return False if self.compiled_matcher is not None: return self.compiled_matcher.fullmatch(tool_name) is not None - # compiled_matcher is None only when the regex failed to compile, - # in which case we already warned and fall back to literal equality. - return tool_name == self.matcher + return tool_name == self.matcher # regex failed to compile: literal fallback -# --------------------------------------------------------------------------- -# Public API -# --------------------------------------------------------------------------- +@dataclass +class ShellHookSpec(_ToolMatcherMixin): + """Parsed and validated representation of a single ``hooks:`` entry.""" + + event: str + command: str + matcher: Optional[str] = None + timeout: int = DEFAULT_TIMEOUT_SECONDS + fail_closed: bool = False + compiled_matcher: Optional[re.Pattern] = field(default=None, repr=False) + + +# --- Public API --- def register_from_config( cfg: Optional[Dict[str, Any]], *, accept_hooks: bool = False, ) -> List[ShellHookSpec]: - """Register every configured shell hook on the plugin manager. + """Register every configured shell hook on the plugin manager; idempotent. - ``cfg`` is the full parsed config dict (``hermes_cli.config.load_config`` - output). The ``hooks:`` key is read out of it. Missing, empty, or - non-dict ``hooks`` is treated as zero configured hooks. - - ``accept_hooks=True`` skips the TTY consent prompt — the caller is - promising that the user has opted in via a flag, env var, or config - setting. ``HERMES_ACCEPT_HOOKS=1`` and ``hooks_auto_accept: true`` are - also honored inside this function so either CLI or gateway call sites - pick them up. - - Returns the list of :class:`ShellHookSpec` entries that ended up wired - up on the plugin manager. Skipped entries (unknown events, malformed, - not allowlisted, already registered) are logged but not returned. + Returns the specs that were newly wired up. Skipped entries (unknown + events, malformed, not allowlisted, already registered) are logged only. """ if not isinstance(cfg, dict): return [] - # Safe mode (--safe-mode / HERMES_SAFE_MODE=1): shell hooks are user - # customizations too — skip registration entirely so a troubleshooting - # run fires zero user-configured code (plugins, MCP, AND hooks). + # Safe mode: hooks are user customizations too — fire zero user-configured code. from utils import env_var_enabled if env_var_enabled("HERMES_SAFE_MODE"): @@ -282,23 +191,18 @@ def register_from_config( return [] effective_accept = _resolve_effective_accept(cfg, accept_hooks) - specs = _parse_hooks_block(cfg.get("hooks")) if not specs: return [] registered: List[ShellHookSpec] = [] - - # Import lazily — avoids circular imports at module-load time. - from hermes_cli.plugins import get_plugin_manager + from hermes_cli.plugins import get_plugin_manager # lazy: avoids import cycle manager = get_plugin_manager() - home_key = str(get_hermes_home().expanduser().resolve()) + home_key = _home_key() - # Idempotence + allowlist read happen under the lock; the TTY - # prompt runs outside so other threads aren't parked on a blocking - # input(). Mutation re-takes the lock with a defensive idempotence - # re-check in case two callers ever race through the prompt. + # Idempotence + allowlist read happen under the lock; the TTY prompt runs + # outside it; mutation re-takes the lock and re-checks in case two callers raced. for spec in specs: key = (home_key, spec.event, spec.matcher, spec.command) with _registered_lock: @@ -306,18 +210,17 @@ def register_from_config( continue already_allowlisted = _is_allowlisted(spec.event, spec.command) - if not already_allowlisted: - if not _prompt_and_record( - spec.event, spec.command, accept_hooks=effective_accept, - ): - logger.warning( - "shell hook for %s (%s) not allowlisted — skipped. " - "Use --accept-hooks / HERMES_ACCEPT_HOOKS=1 / " - "hooks_auto_accept: true, or approve at the TTY " - "prompt next run.", - spec.event, spec.command, - ) - continue + if not already_allowlisted and not _prompt_and_record( + spec.event, spec.command, accept_hooks=effective_accept, + ): + logger.warning( + "shell hook for %s (%s) not allowlisted — skipped. " + "Use --accept-hooks / HERMES_ACCEPT_HOOKS=1 / " + "hooks_auto_accept: true, or approve at the TTY " + "prompt next run.", + spec.event, spec.command, + ) + continue with _registered_lock: if key in _registered: @@ -336,38 +239,20 @@ def register_from_config( def iter_configured_hooks(cfg: Optional[Dict[str, Any]]) -> List[ShellHookSpec]: - """Return the parsed ``ShellHookSpec`` entries from config without - registering anything. Used by ``hermes hooks list`` and ``doctor``.""" + """Parse config hooks without registering (``hermes hooks list`` / doctor).""" if not isinstance(cfg, dict): return [] return _parse_hooks_block(cfg.get("hooks")) def re_register_config_hooks() -> None: - """Re-register shell hooks from config after a plugin force-reload. + """Re-register config hooks after a plugin force-reload cleared the manager's hooks. - ``PluginManager.discover_and_load(force=True)`` unloads via the ownership - ledger and clears the manager's ``_hooks`` dict, which silently drops - shell hooks that were registered from ``config.yaml`` at startup (they - are config-owned, not plugin-owned, so the ledger cannot restore them). - Clear the idempotence set and re-run ``register_from_config()`` so hooks - are wired again (#60036 / PR #60267; tracking #64178 — salvaged from - PR #64188). - - Only the idempotence keys for the *current* Hermes home are cleared — - ``discover_and_load(force=True)`` only unloads the manager scoped to - that one home, so clearing every home's keys would make a force-reload - in profile A drop profile B's still-live registration from the ledger - and duplicate it on B's next registration call (#92682 review). - - Commands already allowlisted stay allowlisted, so this never re-prompts - at a TTY for hooks the user previously approved. + Only the current home's idempotence keys are cleared so a force-reload in + profile A never drops profile B's live registration. Allowlisted commands + stay allowlisted, so this never re-prompts. """ - 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) from hermes_cli.config import load_config register_from_config(load_config()) @@ -379,34 +264,22 @@ def reset_for_tests() -> None: _registered.clear() -# --------------------------------------------------------------------------- -# Config parsing -# --------------------------------------------------------------------------- +# --- Config parsing --- def _parse_hooks_block(hooks_cfg: Any) -> List[ShellHookSpec]: - """Normalise the ``hooks:`` dict into a flat list of ``ShellHookSpec``. - - Malformed entries warn-and-skip — we never raise from config parsing - because a broken hook must not crash the agent. - """ + """Normalise ``hooks:`` into specs; malformed entries warn-and-skip, never raise.""" from hermes_cli.plugins import SHELL_UNSUPPORTED_HOOKS, VALID_HOOKS if not isinstance(hooks_cfg, dict): return [] specs: List[ShellHookSpec] = [] - for event_name, entries in hooks_cfg.items(): - # Reserved sub-keys that aren't event names — skip silently. These - # are config sub-sections nested under `hooks:` for related - # functionality (e.g. output-spill budgets, outbound webhooks — - # the latter parsed by agent/outbound_webhooks.py). + # Reserved non-event sub-sections nested under `hooks:`. if event_name in ("output_spill", "outbound"): continue if event_name in SHELL_UNSUPPORTED_HOOKS: - # Registering would "succeed" while the hook's return value is - # silently dropped (_parse_response has no channel for these - # events' directives) — refuse loudly instead. + # _parse_response has no channel for these directives — refuse loudly. logger.warning( "hook event %r is Python-plugin-only: shell hooks cannot " "return its directive, so this registration is refused " @@ -415,31 +288,20 @@ def _parse_hooks_block(hooks_cfg: Any) -> List[ShellHookSpec]: ) continue if event_name not in VALID_HOOKS: - suggestion = difflib.get_close_matches( - str(event_name), VALID_HOOKS, n=1, cutoff=0.6, - ) + suggestion = difflib.get_close_matches(str(event_name), VALID_HOOKS, n=1, cutoff=0.6) if suggestion: - logger.warning( - "unknown hook event %r in hooks: config — did you mean %r?", - event_name, suggestion[0], - ) + logger.warning("unknown hook event %r in hooks: config — did you mean %r?", event_name, suggestion[0]) else: - logger.warning( - "unknown hook event %r in hooks: config (valid: %s)", - event_name, ", ".join(sorted(VALID_HOOKS)), - ) + logger.warning("unknown hook event %r in hooks: config (valid: %s)", event_name, ", ".join(sorted(VALID_HOOKS))) continue - if entries is None: continue - if not isinstance(entries, list): logger.warning( "hooks.%s must be a list of hook definitions; got %s", event_name, type(entries).__name__, ) continue - for i, raw in enumerate(entries): spec = _parse_single_entry(event_name, i, raw) if spec is not None: @@ -448,38 +310,29 @@ def _parse_hooks_block(hooks_cfg: Any) -> List[ShellHookSpec]: return specs -def _parse_single_entry( - event: str, index: int, raw: Any, -) -> Optional[ShellHookSpec]: +def _parse_single_entry(event: str, index: int, raw: Any) -> Optional[ShellHookSpec]: + def warn(msg: str, *args: Any) -> None: + logger.warning("hooks.%s[%d]" + msg, event, index, *args) + if not isinstance(raw, dict): - logger.warning( - "hooks.%s[%d] must be a mapping with a 'command' key; got %s", - event, index, type(raw).__name__, - ) + warn(" must be a mapping with a 'command' key; got %s", type(raw).__name__) return None command = raw.get("command") if not isinstance(command, str) or not command.strip(): - logger.warning( - "hooks.%s[%d] is missing a non-empty 'command' field", - event, index, - ) + warn(" is missing a non-empty 'command' field") return None matcher = raw.get("matcher") if matcher is not None and not isinstance(matcher, str): - logger.warning( - "hooks.%s[%d].matcher must be a string regex; ignoring", - event, index, - ) + warn(".matcher must be a string regex; ignoring") matcher = None - - if matcher is not None and event not in {"pre_tool_call", "post_tool_call"}: - logger.warning( - "hooks.%s[%d].matcher=%r will be ignored at runtime — the " + if matcher is not None and event not in _TOOL_EVENTS: + warn( + ".matcher=%r will be ignored at runtime — the " "matcher field is only honored for pre_tool_call / " "post_tool_call. The hook will fire on every %s event.", - event, index, matcher, event, + matcher, event, ) matcher = None @@ -487,86 +340,48 @@ def _parse_single_entry( try: timeout = int(timeout_raw) except (TypeError, ValueError): - logger.warning( - "hooks.%s[%d].timeout must be an int (got %r); using default %ds", - event, 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 - if timeout < 1: - logger.warning( - "hooks.%s[%d].timeout must be >=1; using default %ds", - event, index, DEFAULT_TIMEOUT_SECONDS, - ) + warn(".timeout must be >=1; using default %ds", DEFAULT_TIMEOUT_SECONDS) timeout = DEFAULT_TIMEOUT_SECONDS - if timeout > MAX_TIMEOUT_SECONDS: - logger.warning( - "hooks.%s[%d].timeout=%ds exceeds max %ds; clamping", - event, index, timeout, MAX_TIMEOUT_SECONDS, - ) + warn(".timeout=%ds exceeds max %ds; clamping", timeout, MAX_TIMEOUT_SECONDS) timeout = MAX_TIMEOUT_SECONDS - # ``fail_closed`` (canonical) / ``failClosed`` (Cursor/Claude-Code - # config compat). Canonical spelling wins when both are present. - fail_closed_raw = raw.get("fail_closed", raw.get("failClosed", False)) - if not isinstance(fail_closed_raw, bool): - logger.warning( - "hooks.%s[%d].fail_closed must be a boolean (got %r); " - "using default false (fail open)", - event, index, fail_closed_raw, - ) - fail_closed_raw = False - fail_closed = fail_closed_raw - + # ``fail_closed`` (canonical) wins over ``failClosed`` (Cursor/Claude-Code compat). + fail_closed = raw.get("fail_closed", raw.get("failClosed", False)) + if not isinstance(fail_closed, bool): + warn(".fail_closed must be a boolean (got %r); using default false (fail open)", fail_closed) + fail_closed = False if fail_closed and event not in _BLOCKING_EVENTS: - logger.warning( - "hooks.%s[%d].fail_closed=true will be ignored at runtime — " + warn( + ".fail_closed=true will be ignored at runtime — " "fail_closed only applies to blocking-capable events (%s). " "The hook will fail open on %s like any other hook.", - event, index, ", ".join(sorted(_BLOCKING_EVENTS)), event, + ", ".join(sorted(_BLOCKING_EVENTS)), event, ) fail_closed = False return ShellHookSpec( - event=event, - command=command.strip(), - matcher=matcher, - timeout=timeout, - fail_closed=fail_closed, + event=event, command=command.strip(), matcher=matcher, + timeout=timeout, fail_closed=fail_closed, ) -# --------------------------------------------------------------------------- -# Subprocess callback -# --------------------------------------------------------------------------- - -_TOP_LEVEL_PAYLOAD_KEYS = {"tool_name", "args", "session_id", "parent_session_id"} - +# --- Subprocess callback --- def _spawn(spec: ShellHookSpec, stdin_json: str) -> Dict[str, Any]: - """Run ``spec.command`` as a subprocess with ``stdin_json`` on stdin. + """Run ``spec.command`` with ``stdin_json`` on stdin; the single subprocess site. - Returns a diagnostic dict with the same keys for every outcome - (``returncode``, ``stdout``, ``stderr``, ``timed_out``, - ``elapsed_seconds``, ``error``). This is the single place the - subprocess is actually invoked — both the live callback path - (:func:`_make_callback`) and the CLI test helper (:func:`run_once`) - go through it. + Returns a diagnostic dict with the same keys for every outcome. """ result: Dict[str, Any] = { - "returncode": None, - "stdout": "", - "stderr": "", - "timed_out": False, - "elapsed_seconds": 0.0, - "error": None, + "returncode": None, "stdout": "", "stderr": "", + "timed_out": False, "elapsed_seconds": 0.0, "error": None, } try: - # Windows-safe: plain shlex.split eats backslashes in paths (#78293). - from hermes_cli._subprocess_compat import split_command_line - - argv = split_command_line(os.path.expanduser(spec.command)) + argv = _split(os.path.expanduser(spec.command)) except ValueError as exc: result["error"] = f"command {spec.command!r} cannot be parsed: {exc}" return result @@ -575,25 +390,18 @@ def _spawn(spec: ShellHookSpec, stdin_json: str) -> Dict[str, Any]: return result t0 = time.monotonic() - # Spawn the hook in its own process group on POSIX (``process_group=0``, - # Python ≥3.11) so a timed-out hook's descendants can be reaped with the - # hook itself. Windows keeps the hidden-window flags; tree cleanup there - # goes through ``taskkill /T`` in ``kill_process_tree``. Hooks that - # complete in time keep their descendants — an intentionally detached - # helper (``some-daemon &``) survives a successful run. Ported from - # openai/codex#37527 ("Terminate timed-out hook process trees"). - _popen_kwargs: Dict[str, Any] = ( + # Own process group on POSIX so a timed-out hook's descendants are reaped + # with it; Windows tree cleanup goes through kill_process_tree (taskkill /T). + # Hooks that finish in time keep detached helpers alive. + popen_kwargs: Dict[str, Any] = ( {"creationflags": windows_hide_flags()} if IS_WINDOWS else {"process_group": 0} ) try: proc = subprocess.Popen( argv, - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, encoding='utf-8', errors='replace', - shell=False, - **_popen_kwargs, + stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE, + text=True, encoding='utf-8', errors='replace', shell=False, + **popen_kwargs, ) except FileNotFoundError: result["error"] = "command not found" @@ -607,25 +415,18 @@ def _spawn(spec: ShellHookSpec, stdin_json: str) -> Dict[str, Any]: try: stdout, stderr = proc.communicate(input=stdin_json, timeout=spec.timeout) - except subprocess.TimeoutExpired: - # Take down the whole process tree, not just the direct child — - # otherwise a hook that forked helpers leaves them running (and, - # holding the pipe write ends, they'd stall the drain below). + except Exception as exc: + # Kill the whole tree — forked helpers holding the pipes would stall the drain. kill_process_tree(proc) try: proc.communicate(timeout=1) except Exception: pass - result["timed_out"] = True - result["elapsed_seconds"] = round(time.monotonic() - t0, 3) - return result - except Exception as exc: # pragma: no cover — defensive - kill_process_tree(proc) - try: - proc.communicate(timeout=1) - except Exception: - pass - result["error"] = str(exc) + if isinstance(exc, subprocess.TimeoutExpired): + result["timed_out"] = True + result["elapsed_seconds"] = round(time.monotonic() - t0, 3) + else: # pragma: no cover — defensive + result["error"] = str(exc) return result result["returncode"] = proc.returncode @@ -639,13 +440,9 @@ def _make_callback(spec: ShellHookSpec) -> Callable[..., Optional[Dict[str, Any] """Build the closure that ``invoke_hook()`` will call per firing.""" def _callback(**kwargs: Any) -> Optional[Dict[str, Any]]: - # Matcher gate — only meaningful for tool-scoped events. - if spec.event in {"pre_tool_call", "post_tool_call"}: - if not spec.matches_tool(kwargs.get("tool_name")): - return None - - r = _spawn(spec, _serialize_payload(spec.event, kwargs)) - return _evaluate_result(spec, r) + if spec.event in _TOOL_EVENTS and not spec.matches_tool(kwargs.get("tool_name")): + return None + return _evaluate_result(spec, _spawn(spec, _serialize_payload(spec.event, kwargs))) _callback.__name__ = f"shell_hook[{spec.event}:{spec.command}]" _callback.__qualname__ = _callback.__name__ @@ -653,33 +450,17 @@ def _make_callback(spec: ShellHookSpec) -> Callable[..., Optional[Dict[str, Any] def _fail_closed_block(spec: ShellHookSpec, reason: str) -> Dict[str, Any]: - """Canonical block shape for a ``fail_closed`` hook that failed.""" - return { - "action": "block", - "message": f"hook {spec.command} failed closed: {reason}", - } + return {"action": "block", "message": f"hook {spec.command} failed closed: {reason}"} -def _evaluate_result( - spec: ShellHookSpec, r: Dict[str, Any], -) -> Optional[Dict[str, Any]]: - """Turn a :func:`_spawn` diagnostic dict into the hook's contribution. +def _evaluate_result(spec: ShellHookSpec, r: Dict[str, Any]) -> Optional[Dict[str, Any]]: + """Turn a ``_spawn`` diagnostic dict into the hook's contribution. - Single place that encodes the failure semantics: - - * spawn error / timeout — fail open (log + ``None``) unless the spec - is ``fail_closed`` on a blocking-capable event, in which case a - canonical block shape is returned; - * exit code 2 on a blocking-capable event — block, with the message - taken from stdout block JSON, then stderr, then a default - (Claude-Code / Cursor compatible); - * other non-zero exits — warn, then parse stdout normally; - * non-JSON / unparseable stdout on a ``fail_closed`` blocking hook — - block instead of silently contributing nothing. - - Shared by the live callback path (:func:`_make_callback`) and the CLI - test helper (:func:`run_once`) so ``hermes hooks test`` reflects - production behaviour exactly. + Encodes the failure semantics once (shared by the live callback and + ``run_once``): spawn error/timeout fail open unless fail_closed; exit 2 on + a blocking event blocks (message from stdout JSON, then stderr, then + default); other non-zero exits warn then parse stdout; unparseable stdout + on a fail_closed hook blocks. """ blocking_event = spec.event in _BLOCKING_EVENTS fail_closed = spec.fail_closed and blocking_event @@ -689,19 +470,13 @@ def _evaluate_result( "shell hook failed (event=%s command=%s): %s", spec.event, spec.command, r["error"], ) - if fail_closed: - return _fail_closed_block(spec, r["error"]) - return None + return _fail_closed_block(spec, r["error"]) if fail_closed else None if r["timed_out"]: logger.warning( "shell hook timed out after %.2fs (event=%s command=%s)", r["elapsed_seconds"], spec.event, spec.command, ) - if fail_closed: - return _fail_closed_block( - spec, f"timed out after {spec.timeout}s", - ) - return None + return _fail_closed_block(spec, f"timed out after {spec.timeout}s") if fail_closed else None stderr = r["stderr"].strip() if stderr: @@ -710,9 +485,6 @@ def _evaluate_result( spec.event, spec.command, stderr[:_STDERR_MESSAGE_LIMIT], ) - # Exit code 2 = block (Claude-Code / Cursor compatible), for events - # whose block directive is honored downstream. stdout block JSON - # still wins for the message; otherwise stderr, then a default. if r["returncode"] == BLOCK_EXIT_CODE and blocking_event: parsed = _parse_response(spec.event, r["stdout"]) if isinstance(parsed, dict) and parsed.get("action") == "block": @@ -724,8 +496,7 @@ def _evaluate_result( ) return {"action": "block", "message": message} - # Other non-zero exits: log but still parse stdout so scripts that - # signal failure via exit code can also return a block directive. + # Other non-zero exits: still parse stdout so exit-code failures can carry a block directive. if r["returncode"] != 0: logger.warning( "shell hook exited %d (event=%s command=%s); stderr=%s", @@ -737,127 +508,82 @@ def _evaluate_result( parsed = _parse_response(spec.event, stdout) if parsed is None and fail_closed and stdout: - # The hook produced output we could not turn into a directive. - # A fail-closed gate must not silently allow the action on - # garbage output (e.g. a stack trace on stdout). + # A fail-closed gate must not silently allow on garbage stdout (e.g. a stack trace). try: - data = json.loads(stdout) - valid_json = isinstance(data, dict) + valid_json = isinstance(json.loads(stdout), dict) except json.JSONDecodeError: valid_json = False if not valid_json: - return _fail_closed_block( - spec, "unparseable stdout (expected a JSON object)", - ) + return _fail_closed_block(spec, "unparseable stdout (expected a JSON object)") return parsed def _serialize_payload(event: str, kwargs: Dict[str, Any]) -> str: - """Render the stdin JSON payload. Unserialisable values are - stringified via ``default=str`` rather than dropped.""" - 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 = "" - payload = { - "hook_event_name": event, - "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, - } + """Render the stdin JSON payload; unserialisable values are stringified.""" + payload = {"hook_event_name": event, **_payload_fields(kwargs)} return json.dumps(payload, ensure_ascii=False, default=str) def _block_message(primary: Any, secondary: Any) -> str: - """Return a validated string block message, falling back to the default. - - Accepts two candidate fields (primary wins over secondary) so callers - can express field-priority differences between the two hook wire formats - without duplicating the type-check logic. - """ + """Validated string block message (primary wins), falling back to the default.""" raw = primary or secondary return raw if isinstance(raw, str) and raw else _DEFAULT_BLOCK_MESSAGE -def _parse_response(event: str, stdout: str) -> Optional[Dict[str, Any]]: - """Translate stdout JSON into a Hermes wire-shape dict. - - For ``pre_tool_call`` the Claude-Code-style ``{"decision": "block", - "reason": "..."}`` payload is translated into the canonical Hermes - ``{"action": "block", "message": "..."}`` shape expected by - :func:`hermes_cli.plugins.get_pre_tool_call_block_message`. This is - the single most important correctness invariant in this module — - skipping the translation silently breaks every ``pre_tool_call`` - block directive. - - For ``pre_tool_call`` the ``modify`` action (canonical: ``{"action": - "modify", "args": {...}}``, Claude-Code-style: ``{"decision": - "modify", "tool_input": {...}}``) is translated to - ``{"action": "modify", "args": {...}}`` so callers can merge the - returned fields into the tool's ``args`` before dispatch. - - For ``pre_llm_call``, ``{"context": "..."}`` is passed through - unchanged to match the existing plugin-hook contract. - - Anything else returns ``None``. - """ - stdout = (stdout or "").strip() - if not stdout: - return None - - try: - data = json.loads(stdout) - except json.JSONDecodeError: - logger.warning( - "shell hook stdout was not valid JSON (event=%s): %s", - event, stdout[:200], - ) - return None - - if not isinstance(data, dict): - return None - - if event == "pre_tool_call": - if data.get("action") == "block": - return {"action": "block", "message": _block_message(data.get("message"), data.get("reason"))} - if data.get("decision") == "block": - return {"action": "block", "message": _block_message(data.get("reason"), data.get("message"))} - # "modify" action — transform tool_input before dispatch - if data.get("action") == "modify": - new_args = data.get("args") - if isinstance(new_args, dict): - return {"action": "modify", "args": new_args} - if data.get("decision") == "modify": - new_args = data.get("tool_input") - if isinstance(new_args, dict): - return {"action": "modify", "args": new_args} - return None - - if event == "pre_verify": - # "continue" (Hermes) / "block" (Claude-Code Stop: block the stop) both - # mean keep going; the message/reason is the follow-up for the model. A - # continue with no message is a no-op — let the turn finish. - action = str(data.get("action") or data.get("decision") or "").strip().lower() - if action in {"continue", "block"}: - message = data.get("message") or data.get("reason") - if isinstance(message, str) and message.strip(): - return {"action": "continue", "message": message.strip()} - return None - - context = data.get("context") - if isinstance(context, str) and context.strip(): - return {"context": context} - +def _parse_pre_tool_call(data: Dict[str, Any]) -> Optional[Dict[str, Any]]: + # Claude-Code-style {"decision": ..., "reason"/"tool_input": ...} is translated to + # the canonical Hermes shape expected by get_pre_tool_call_block_message — + # skipping this silently breaks every pre_tool_call block directive. + if data.get("action") == "block": + return {"action": "block", "message": _block_message(data.get("message"), data.get("reason"))} + if data.get("decision") == "block": + return {"action": "block", "message": _block_message(data.get("reason"), data.get("message"))} + if data.get("action") == "modify" and isinstance(data.get("args"), dict): + return {"action": "modify", "args": data["args"]} + if data.get("decision") == "modify" and isinstance(data.get("tool_input"), dict): + return {"action": "modify", "args": data["tool_input"]} return None -# --------------------------------------------------------------------------- -# Allowlist / consent -# --------------------------------------------------------------------------- +def _parse_pre_verify(data: Dict[str, Any]) -> Optional[Dict[str, Any]]: + # "continue" (Hermes) / "block" (Claude-Code Stop) both mean keep going; + # a continue with no message is a no-op. + action = str(data.get("action") or data.get("decision") or "").strip().lower() + if action in {"continue", "block"}: + message = data.get("message") or data.get("reason") + if isinstance(message, str) and message.strip(): + return {"action": "continue", "message": message.strip()} + return None + + +def _parse_context(data: Dict[str, Any]) -> Optional[Dict[str, Any]]: + context = data.get("context") + if isinstance(context, str) and context.strip(): + return {"context": context} + return None + + +_RESPONSE_PARSERS: Dict[str, Callable[[Dict[str, Any]], Optional[Dict[str, Any]]]] = { + "pre_tool_call": _parse_pre_tool_call, + "pre_verify": _parse_pre_verify, +} + + +def _parse_response(event: str, stdout: str) -> Optional[Dict[str, Any]]: + """Translate stdout JSON into a Hermes wire-shape dict, or ``None``.""" + stdout = (stdout or "").strip() + if not stdout: + return None + try: + data = json.loads(stdout) + except json.JSONDecodeError: + logger.warning("shell hook stdout was not valid JSON (event=%s): %s", event, stdout[:200]) + return None + return _RESPONSE_PARSERS.get(event, _parse_context)(data) if isinstance(data, dict) else None + + +# --- Allowlist / consent --- def allowlist_path() -> Path: """Path to the per-user shell-hook allowlist file.""" @@ -869,27 +595,20 @@ def load_allowlist() -> Dict[str, Any]: try: raw = json.loads(allowlist_path().read_text(encoding="utf-8")) except (FileNotFoundError, json.JSONDecodeError, OSError): - return {"approvals": []} + raw = None if not isinstance(raw, dict): return {"approvals": []} - approvals = raw.get("approvals") - if not isinstance(approvals, list): + if not isinstance(raw.get("approvals"), list): raw["approvals"] = [] return raw def save_allowlist(data: Dict[str, Any]) -> None: - """Atomically persist the allowlist via per-process ``mkstemp`` + - ``os.replace``. Cross-process read-modify-write races are handled - by :func:`_locked_update_approvals` (``fcntl.flock``). On OSError - the failure is logged; the in-process hook still registers but - the approval won't survive across runs.""" + """Atomically persist the allowlist; on OSError log and keep the in-process approval.""" p = allowlist_path() try: p.parent.mkdir(parents=True, exist_ok=True) - fd, tmp_path = tempfile.mkstemp( - prefix=f"{p.name}.", suffix=".tmp", dir=str(p.parent), - ) + fd, tmp_path = tempfile.mkstemp(prefix=f"{p.name}.", suffix=".tmp", dir=str(p.parent)) try: with os.fdopen(fd, "w", encoding="utf-8") as fh: fh.write(json.dumps(data, indent=2, sort_keys=True)) @@ -911,55 +630,35 @@ def save_allowlist(data: Dict[str, Any]) -> None: def _is_allowlisted(event: str, command: str) -> bool: - data = load_allowlist() - return any( - isinstance(e, dict) - and e.get("event") == event - and e.get("command") == command - for e in data.get("approvals", []) - ) + return any(_entry_matches(e, event, command) for e in load_allowlist().get("approvals", [])) @contextmanager def _locked_update_approvals() -> Iterator[Dict[str, Any]]: - """Serialise read-modify-write on the allowlist across processes. - - Holds an exclusive ``flock`` on a sibling lock file for the duration - of the update so concurrent ``_record_approval``/``revoke`` callers - cannot clobber each other's changes (the race Codex reproduced with - 20–50 simultaneous writers). Falls back to an in-process lock on - platforms without ``fcntl``. - """ + """Serialise allowlist read-modify-write across processes via a sibling flock file.""" p = allowlist_path() p.parent.mkdir(parents=True, exist_ok=True) - lock_path = p.with_suffix(p.suffix + ".lock") - - if fcntl is None: # pragma: no cover — non-POSIX fallback - with _allowlist_write_lock: - data = load_allowlist() - yield data - save_allowlist(data) - return - - with open(lock_path, "a+", encoding="utf-8") as lock_fh: - fcntl.flock(lock_fh.fileno(), fcntl.LOCK_EX) - try: - data = load_allowlist() - yield data - save_allowlist(data) - finally: - try: - fcntl.flock(lock_fh.fileno(), fcntl.LOCK_UN) - except (OSError, IOError): - pass + with ExitStack() as stack: + if fcntl is None: # pragma: no cover — non-POSIX fallback + stack.enter_context(_allowlist_write_lock) + else: + lock_fh = stack.enter_context(open(p.with_suffix(p.suffix + ".lock"), "a+", encoding="utf-8")) + fcntl.flock(lock_fh.fileno(), fcntl.LOCK_EX) + stack.callback(_flock_unlock, lock_fh) + data = load_allowlist() + yield data + save_allowlist(data) -def _prompt_and_record( - event: str, command: str, *, accept_hooks: bool, -) -> bool: - """Decide whether to approve an unseen ``(event, command)`` pair. - Returns ``True`` iff the approval was granted and recorded. - """ +def _flock_unlock(lock_fh: Any) -> None: + try: + fcntl.flock(lock_fh.fileno(), fcntl.LOCK_UN) + except (OSError, IOError): + pass + + +def _prompt_and_record(event: str, command: str, *, accept_hooks: bool) -> bool: + """Approve an unseen ``(event, command)`` pair; True iff granted and recorded.""" if accept_hooks: _record_approval(event, command) logger.info( @@ -988,7 +687,6 @@ def _prompt_and_record( if answer in {"y", "yes"}: _record_approval(event, command) return True - return False @@ -1001,31 +699,19 @@ def _record_approval(event: str, command: str) -> None: } with _locked_update_approvals() as data: data["approvals"] = [ - e for e in data.get("approvals", []) - if not ( - isinstance(e, dict) - and e.get("event") == event - and e.get("command") == command - ) + e for e in data.get("approvals", []) if not _entry_matches(e, event, command) ] + [entry] -def _utc_now_iso() -> str: - return datetime.now(tz=timezone.utc).isoformat().replace("+00:00", "Z") - - def revoke(command: str) -> int: - """Remove every allowlist entry matching ``command``. + """Remove every allowlist entry matching ``command``; returns the count removed. - Returns the number of entries removed. Does not unregister any - callbacks that are already live on the plugin manager in the current - process — restart the CLI / gateway to drop them. + Live callbacks in the current process stay registered until restart. """ with _locked_update_approvals() as data: before = len(data.get("approvals", [])) data["approvals"] = [ - e for e in data.get("approvals", []) - if not (isinstance(e, dict) and e.get("command") == command) + e for e in data.get("approvals", []) if not _entry_matches(e, None, command) ] after = len(data["approvals"]) return before - after @@ -1040,96 +726,56 @@ _SCRIPT_EXTENSIONS: Tuple[str, ...] = ( def _command_script_path(command: str) -> str: - """Return the script path from ``command`` for doctor / drift checks. - - Prefers a token ending in a known script extension, then a token - containing ``/`` or leading ``~``, then the first token. Handles - ``python3 /path/hook.py``, ``/usr/bin/env bash hook.sh``, and the - common bare-path form. - """ + """Script path from ``command``: first token with a script extension, then a + path-like token, then the first token (``python3 /p/hook.py``, ``/usr/bin/env bash x.sh``).""" try: - from hermes_cli._subprocess_compat import split_command_line - - parts = split_command_line(command) + parts = _split(command) except ValueError: return command if not parts: return command - for part in parts: - if part.lower().endswith(_SCRIPT_EXTENSIONS): - return part - for part in parts: - if "/" in part or part.startswith("~"): - return part - return parts[0] + return ( + next((p for p in parts if p.lower().endswith(_SCRIPT_EXTENSIONS)), None) + or next((p for p in parts if "/" in p or p.startswith("~")), None) + or parts[0] + ) -# --------------------------------------------------------------------------- -# Helpers for accept-hooks resolution -# --------------------------------------------------------------------------- - -def _resolve_effective_accept( - cfg: Dict[str, Any], accept_hooks_arg: bool, -) -> bool: - """Combine all three opt-in channels into a single boolean. - - Precedence (any truthy source flips us on): - 1. ``--accept-hooks`` flag (CLI) / explicit argument - 2. ``HERMES_ACCEPT_HOOKS`` env var - 3. ``hooks_auto_accept: true`` in ``cli-config.yaml`` - """ - if accept_hooks_arg: - return True - env = os.environ.get("HERMES_ACCEPT_HOOKS", "").strip().lower() - if env in {"1", "true", "yes", "on"}: +def _resolve_effective_accept(cfg: Dict[str, Any], accept_hooks_arg: bool) -> bool: + """Any truthy opt-in channel wins: explicit arg, HERMES_ACCEPT_HOOKS, hooks_auto_accept.""" + if accept_hooks_arg or os.environ.get("HERMES_ACCEPT_HOOKS", "").strip().lower() in _TRUTHY: return True cfg_val = cfg.get("hooks_auto_accept", False) if isinstance(cfg_val, bool): return cfg_val if isinstance(cfg_val, str): - return cfg_val.strip().lower() in {"1", "true", "yes", "on"} + return cfg_val.strip().lower() in _TRUTHY return False -# --------------------------------------------------------------------------- -# Introspection (used by `hermes hooks` CLI) -# --------------------------------------------------------------------------- +# --- Introspection (used by `hermes hooks` CLI) --- def allowlist_entry_for(event: str, command: str) -> Optional[Dict[str, Any]]: """Return the allowlist record for this pair, if any.""" - for e in load_allowlist().get("approvals", []): - if ( - isinstance(e, dict) - and e.get("event") == event - and e.get("command") == command - ): - return e - return None + return next((e for e in load_allowlist().get("approvals", []) if _entry_matches(e, event, command)), None) def script_mtime_iso(command: str) -> Optional[str]: - """ISO-8601 mtime of the resolved script path, or ``None`` if the - script is missing.""" + """ISO-8601 mtime of the resolved script path, or ``None`` if missing.""" path = _command_script_path(command) if not path: return None try: - expanded = os.path.expanduser(path) return datetime.fromtimestamp( - os.path.getmtime(expanded), tz=timezone.utc, + os.path.getmtime(os.path.expanduser(path)), tz=timezone.utc, ).isoformat().replace("+00:00", "Z") except OSError: return None def script_is_executable(command: str) -> bool: - """Return ``True`` iff ``command`` is runnable as configured. - - For a bare invocation (``/path/hook.sh``) the script itself must be - executable. For interpreter-prefixed commands (``python3 - /path/hook.py``, ``/usr/bin/env bash hook.sh``) the script just has - to be readable — the interpreter doesn't care about the ``X_OK`` - bit. Mirrors what ``_spawn`` would actually do at runtime.""" + """True iff ``command`` is runnable as configured: a bare script needs X_OK, + an interpreter-prefixed one only R_OK (mirrors what ``_spawn`` does).""" path = _command_script_path(command) if not path: return False @@ -1137,33 +783,19 @@ def script_is_executable(command: str) -> bool: if not os.path.isfile(expanded): return False try: - from hermes_cli._subprocess_compat import split_command_line - - argv = split_command_line(command) + argv = _split(command) except ValueError: return False is_bare_invocation = bool(argv) and argv[0] == path - required = os.X_OK if is_bare_invocation else os.R_OK - return os.access(expanded, required) + return os.access(expanded, os.X_OK if is_bare_invocation else os.R_OK) -def run_once( - spec: ShellHookSpec, kwargs: Dict[str, Any], -) -> Dict[str, Any]: - """Fire a single shell-hook invocation with a synthetic payload. - Used by ``hermes hooks test`` and ``hermes hooks doctor``. +def run_once(spec: ShellHookSpec, kwargs: Dict[str, Any]) -> Dict[str, Any]: + """Fire one hook with a synthetic payload (``hermes hooks test`` / doctor). - ``kwargs`` is the same dict that :func:`hermes_cli.plugins.invoke_hook` - would pass at runtime. It is routed through :func:`_serialize_payload` - so the synthetic stdin exactly matches what a real hook firing would - produce — otherwise scripts tested via ``hermes hooks test`` could - diverge silently from production behaviour. - - Returns the :func:`_spawn` diagnostic dict plus a ``parsed`` field - holding the canonical Hermes-wire-shape response — including exit-code-2 - blocking and ``fail_closed`` semantics, so what ``hermes hooks test`` - prints is exactly what the dispatcher would receive.""" - stdin_json = _serialize_payload(spec.event, kwargs) - result = _spawn(spec, stdin_json) + Routes through ``_serialize_payload`` and ``_evaluate_result`` so the + result (``_spawn`` dict + ``parsed``) matches production exactly. + """ + result = _spawn(spec, _serialize_payload(spec.event, kwargs)) result["parsed"] = _evaluate_result(spec, result) return result diff --git a/agent/tool_dispatch_helpers.py b/agent/tool_dispatch_helpers.py index 41dc4066b9..54bedab1ae 100644 --- a/agent/tool_dispatch_helpers.py +++ b/agent/tool_dispatch_helpers.py @@ -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 ```` 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. ````). +# 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 ```` 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 diff --git a/agent/tool_guardrails.py b/agent/tool_guardrails.py index b432c21d9f..6ec3b289ae 100644 --- a/agent/tool_guardrails.py +++ b/agent/tool_guardrails.py @@ -1,9 +1,8 @@ """Pure tool-call loop guardrail primitives. -The controller in this module is intentionally side-effect free: it tracks -per-turn tool-call observations and returns decisions. Runtime code owns whether -those decisions become warning guidance, synthetic tool results, or controlled -turn halts. +The controller is side-effect free: it tracks per-turn tool-call observations +and returns decisions. Runtime code decides whether a decision becomes warning +guidance, a synthetic tool result, or a controlled turn halt. """ from __future__ import annotations @@ -17,147 +16,69 @@ from utils import safe_json_loads from agent.tool_result_classification import file_mutation_result_landed -IDEMPOTENT_TOOL_NAMES = frozenset( - { - "read_file", - "search_files", - "web_search", - "web_extract", - "session_search", - "skill_view", - "skills_list", - "browser_snapshot", - "browser_console", - "browser_get_images", - "mcp_filesystem_read_file", - "mcp_filesystem_read_text_file", - "mcp_filesystem_read_multiple_files", - "mcp_filesystem_list_directory", - "mcp_filesystem_list_directory_with_sizes", - "mcp_filesystem_directory_tree", - "mcp_filesystem_get_file_info", - "mcp_filesystem_search_files", - } -) +IDEMPOTENT_TOOL_NAMES = frozenset({ + "read_file", "search_files", "web_search", "web_extract", "session_search", "skill_view", + "skills_list", "browser_snapshot", "browser_console", "browser_get_images", + "mcp_filesystem_read_file", "mcp_filesystem_read_text_file", + "mcp_filesystem_read_multiple_files", "mcp_filesystem_list_directory", + "mcp_filesystem_list_directory_with_sizes", "mcp_filesystem_directory_tree", + "mcp_filesystem_get_file_info", "mcp_filesystem_search_files", +}) -MUTATING_TOOL_NAMES = frozenset( - { - "terminal", - "execute_code", - "write_file", - "patch", - "todo_list", - "memory", - "skill_manage", - "browser_click", - "browser_type", - "browser_press", - "browser_scroll", - "browser_navigate", - "send_message", - "cronjob_manage", - "delegate_task", - "process_manage", - } -) +MUTATING_TOOL_NAMES = frozenset({ + "terminal", "execute_code", "write_file", "patch", "todo_list", "memory", "skill_manage", + "browser_click", "browser_type", "browser_press", "browser_scroll", "browser_navigate", + "send_message", "cronjob_manage", "delegate_task", "process_manage", +}) -# Tools that are legitimately re-invoked with identical arguments and may -# legitimately return an unchanged result while waiting on external progress — -# background-process management and job pollers. The identical-call loop -# notice (agent.stall_guards) never fires for these, so polling patterns like -# ``process(action="poll")`` or repeatedly checking a generation job stay -# unannotated. -STALL_GUARD_REPEATABLE_TOOLS = frozenset( - { - "process_manage", - } -) +# Tools legitimately re-invoked with identical args while waiting on external +# progress (pollers). The identical-call NOTICE never fires for these. +STALL_GUARD_REPEATABLE_TOOLS = frozenset({"process_manage"}) +# Poller naming conventions on generated / MCP surfaces (``_get_result``). +_STALL_GUARD_REPEATABLE_SUFFIXES = ("_get_result", "_poll") -# Poller naming conventions (e.g. ``_get_result``) used by generated / -# MCP tool surfaces. Matched as suffixes so vendor-prefixed pollers are exempt -# without enumerating every vendor. -_STALL_GUARD_REPEATABLE_SUFFIXES = ( - "_get_result", - "_poll", -) - -# The notice fires on the Nth consecutive identical call (same tool, same -# canonical args, same result). 3 tolerates one legitimate double-check while -# catching the observed re-issue loops (3x/4x identical calls in eval traces). +# Notice fires on the Nth consecutive identical (tool, args, result) call; 3 +# tolerates one legitimate double-check while catching observed re-issue loops. STALL_GUARD_IDENTICAL_CALL_THRESHOLD = 3 -# Result-reference stubbing (agent.stall_guards): from the 2nd consecutive -# identical call whose FRESH result is byte-identical to the previous one, -# the duplicate payload is replaced in context by a short reference stub. -# Results under this size aren't worth stubbing (the stub itself plus the -# lost locality outweigh the savings), and error results are never stubbed -# (the model must see every fresh error verbatim). +# Result-reference stubbing: from the 2nd consecutive identical call whose fresh +# result is byte-identical, the duplicate payload is replaced by a reference +# stub. Results under this size aren't worth stubbing; errors are never stubbed. IDENTICAL_RESULT_STUB_MIN_CHARS = 512 - -# How much of the canonical args JSON the stub carries so the model still -# knows WHAT the referenced call was even if context compression later -# evicts the referenced result (cheap dangling-reference mitigation). +# Canonical-args preview kept in the stub so the model still knows WHAT the call +# was if compression later evicts the referenced result. _RESULT_STUB_ARGS_PREVIEW_CHARS = 120 +# Tools whose "failure" is normal work output (red test run, empty grep, page +# timeout). same_tool_failure (DIFFERENT commands) never halts these; only an +# exact-args replay with no intervening change, or an identical-result streak, can. +FAILURE_TOLERANT_TOOL_NAMES = frozenset({ + "terminal", "execute_code", "process_manage", "process", "browser_navigate", "web_extract", +}) -# Tools whose "failure" is a normal, informative outcome of legitimate work: -# a red test run, a grep with no matches, a failing build during a fix loop, a -# page that times out. Hard stops never fire on these from failure counts of -# DIFFERENT commands (same_tool_failure) — only an exact-args replay with NO -# intervening change, or an identical-result streak, can halt them. -FAILURE_TOLERANT_TOOL_NAMES = frozenset( - { - "terminal", - "execute_code", - "process_manage", - "process", - "browser_navigate", - "web_extract", - } -) - -# A landed mutation between two attempts means the retry is a NEW experiment -# (edit -> re-run) rather than a replay. A successful call to one of these -# marks progress for every failing signature still being counted this turn. -PROGRESS_RESET_TOOL_NAMES = frozenset( - { - "write_file", - "patch", - "terminal", - "execute_code", - "browser_click", - "browser_type", - "browser_press", - "browser_navigate", - "process_manage", - "process", - "delegate_task", - "send_message", - "cronjob", - "cronjob_manage", - "todo", - "todo_list", - "memory", - "skill_manage", - } -) +# A successful call to one of these marks progress for every failing signature +# still counted this turn: the next retry is a new experiment (edit -> re-run), not a replay. +PROGRESS_RESET_TOOL_NAMES = frozenset({ + "write_file", "patch", "terminal", "execute_code", "browser_click", "browser_type", + "browser_press", "browser_navigate", "process_manage", "process", "delegate_task", + "send_message", "cronjob", "cronjob_manage", "todo", "todo_list", "memory", "skill_manage", +}) def is_stall_guard_repeatable(tool_name: str) -> bool: """Whether a tool is exempt from the identical-call loop notice.""" - if tool_name in STALL_GUARD_REPEATABLE_TOOLS: - return True - return tool_name.endswith(_STALL_GUARD_REPEATABLE_SUFFIXES) + return tool_name in STALL_GUARD_REPEATABLE_TOOLS or tool_name.endswith( + _STALL_GUARD_REPEATABLE_SUFFIXES + ) @dataclass(frozen=True) class ToolCallGuardrailConfig: """Thresholds for per-turn tool-call loop detection. - Warnings are enabled by default and never prevent tool execution. Hard stops - stay opt-in for interactive CLI/TUI/Desktop/ACP sessions, but default on for - non-interactive gateway/cron platforms where nobody is present to interrupt - a model that ignores loop warnings. + Warnings never prevent execution. Hard stops are opt-in for interactive + platforms but default on for unattended gateway/cron platforms where nobody + can interrupt a model that ignores loop warnings. """ warnings_enabled: bool = True @@ -180,84 +101,63 @@ class ToolCallGuardrailConfig: *, platform: str | None = None, ) -> "ToolCallGuardrailConfig": - """Build config from the `tool_loop_guardrails` config.yaml section.""" + """Build config from the `tool_loop_guardrails` config.yaml section. + + Nested ``warn_after`` / ``hard_stop_after`` keys win over the flat legacy keys. + """ if not isinstance(data, Mapping): data = {} - - warn_after = data.get("warn_after") - if not isinstance(warn_after, Mapping): - warn_after = {} - hard_stop_after = data.get("hard_stop_after") - if not isinstance(hard_stop_after, Mapping): - hard_stop_after = {} - - defaults = cls() - hard_stop_enabled = _as_bool(data.get("hard_stop_enabled"), defaults.hard_stop_enabled) + d = cls() + hard_stop_enabled = _as_bool(data.get("hard_stop_enabled"), d.hard_stop_enabled) non_interactive_hard_stop_enabled = _as_bool( - data.get("non_interactive_hard_stop_enabled"), - defaults.non_interactive_hard_stop_enabled, + data.get("non_interactive_hard_stop_enabled"), d.non_interactive_hard_stop_enabled, ) if _is_non_interactive_platform(platform) and non_interactive_hard_stop_enabled: hard_stop_enabled = True + thresholds: dict[str, int] = {} + for field_name, (section_name, key) in _THRESHOLD_SOURCES.items(): + section = data.get(section_name) + if not isinstance(section, Mapping): + section = {} + thresholds[field_name] = _int_at_least( + section.get(key, data.get(field_name)), getattr(d, field_name), 1, + ) + return cls( - warnings_enabled=_as_bool(data.get("warnings_enabled"), defaults.warnings_enabled), + warnings_enabled=_as_bool(data.get("warnings_enabled"), d.warnings_enabled), hard_stop_enabled=hard_stop_enabled, non_interactive_hard_stop_enabled=non_interactive_hard_stop_enabled, - exact_failure_warn_after=_positive_int( - warn_after.get("exact_failure", data.get("exact_failure_warn_after")), - defaults.exact_failure_warn_after, - ), - same_tool_failure_warn_after=_positive_int( - warn_after.get("same_tool_failure", data.get("same_tool_failure_warn_after")), - defaults.same_tool_failure_warn_after, - ), - no_progress_warn_after=_positive_int( - warn_after.get("idempotent_no_progress", data.get("no_progress_warn_after")), - defaults.no_progress_warn_after, - ), - exact_failure_block_after=_positive_int( - hard_stop_after.get("exact_failure", data.get("exact_failure_block_after")), - defaults.exact_failure_block_after, - ), - same_tool_failure_halt_after=_positive_int( - hard_stop_after.get("same_tool_failure", data.get("same_tool_failure_halt_after")), - defaults.same_tool_failure_halt_after, - ), - no_progress_block_after=_positive_int( - hard_stop_after.get("idempotent_no_progress", data.get("no_progress_block_after")), - defaults.no_progress_block_after, - ), loop_caps=LoopCapConfig.from_mapping(data.get("loop_caps")), + **thresholds, ) -# Default session-wide caps, matching Claude Code's v2.1.212 runaway-loop -# Per-turn (per-agent-loop) caps on runaway-prone tool calls. Counts reset at -# the start of every agent loop (reset_for_turn), so the limit is "within a -# single turn" rather than cumulative over the whole session. A single loop -# issuing dozens of web searches or spawning dozens of subagents is already -# pathological, so the defaults are deliberately low. +# Threshold field -> (nested section, nested key). The flat legacy key is the field name itself. +_THRESHOLD_SOURCES: dict[str, tuple[str, str]] = { + "exact_failure_warn_after": ("warn_after", "exact_failure"), + "same_tool_failure_warn_after": ("warn_after", "same_tool_failure"), + "no_progress_warn_after": ("warn_after", "idempotent_no_progress"), + "exact_failure_block_after": ("hard_stop_after", "exact_failure"), + "same_tool_failure_halt_after": ("hard_stop_after", "same_tool_failure"), + "no_progress_block_after": ("hard_stop_after", "idempotent_no_progress"), +} + + +# Per-turn caps on runaway-prone tools; counters reset in reset_for_turn at the +# start of every agent loop, so the limit is per turn, not per session. Dozens +# of searches / subagent spawns in one loop is already pathological. _DEFAULT_MAX_WEB_SEARCHES_PER_TURN = 50 _DEFAULT_MAX_SUBAGENTS_PER_TURN = 50 @dataclass(frozen=True) class LoopCapConfig: - """Per-turn caps on runaway-prone tool calls. + """Per-turn hard ceilings on web_search calls / subagent spawns. - Inspired by Claude Code v2.1.212 (Week 29, July 2026), which added caps on - WebSearch calls and subagent spawns to stop runaway search / delegation - loops. Here the caps count *within a single agent loop* (one turn): the - counters reset in ``reset_for_turn`` at the start of every - ``run_conversation``, so a legitimate multi-turn session is never starved, - but a single turn that spirals into an unbounded search / delegation loop - is stopped. - - Semantics differ from the per-turn loop *detector* above (which keys on - repeated identical/failing calls): these caps are a hard ceiling on the - total count of a tool within the turn and fire regardless of - ``hard_stop_enabled``. A value of ``0`` disables the cap (unlimited). + Unlike the loop detector (keyed on repeated identical/failing calls) these + count total calls within the turn and fire regardless of + ``hard_stop_enabled``. ``0`` disables a cap. """ max_web_searches: int = _DEFAULT_MAX_WEB_SEARCHES_PER_TURN @@ -270,44 +170,30 @@ class LoopCapConfig: return cls() defaults = cls() return cls( - max_web_searches=_non_negative_int( - data.get("max_web_searches"), defaults.max_web_searches - ), - max_subagents=_non_negative_int( - data.get("max_subagents"), defaults.max_subagents - ), + max_web_searches=_int_at_least(data.get("max_web_searches"), defaults.max_web_searches, 0), + max_subagents=_int_at_least(data.get("max_subagents"), defaults.max_subagents, 0), ) _INTERACTIVE_PLATFORMS = frozenset({"cli", "tui", "desktop", "acp"}) - -# Platforms that are not chat gateways but whose work is a bounded, supervised -# task loop: a subagent inherits its parent's budget and is stopped by the -# parent; api_server runs have a live client holding the request. Both do -# real edit -> re-run work, so they keep the interactive (warn-only) default. +# Not chat gateways, but bounded supervised task loops (a subagent is stopped by +# its parent; api_server has a live client). Both do real edit -> re-run work, +# so they keep the interactive warn-only default. _SUPERVISED_TASK_PLATFORMS = frozenset({"subagent", "api_server"}) def _is_non_interactive_platform(platform: str | None) -> bool: - """Return true for gateway/cron sessions where tool loops are unattended.""" + """True for gateway/cron sessions where tool loops are unattended.""" if not isinstance(platform, str) or not platform.strip(): return False key = platform.strip().lower() - if key in _INTERACTIVE_PLATFORMS or key in _SUPERVISED_TASK_PLATFORMS: - return False - return True + return key not in _INTERACTIVE_PLATFORMS and key not in _SUPERVISED_TASK_PLATFORMS @dataclass(frozen=True) class IdenticalCallObservation: - """Outcome of observing one completed tool call for the stall guards. - - ``notice`` is the identical-call loop-breaker notice (appended after the - result). ``stub`` is the result-reference replacement for a byte-identical - duplicate result (replaces the result content). Both may be set on the - same call (3rd+ identical call): the stub replaces the payload and the - notice is appended after it. - """ + """Outcome of observing one completed call: ``notice`` is appended after the + result, ``stub`` replaces a byte-identical duplicate result. Both may be set.""" notice: str | None = None stub: str | None = None @@ -322,8 +208,7 @@ class ToolCallSignature: @classmethod def from_call(cls, tool_name: str, args: Mapping[str, Any] | None) -> "ToolCallSignature": - canonical = canonical_tool_args(args or {}) - return cls(tool_name=tool_name, args_hash=_sha256(canonical)) + return cls(tool_name=tool_name, args_hash=_sha256(canonical_tool_args(args or {}))) def to_metadata(self) -> dict[str, str]: """Return public metadata without raw argument values.""" @@ -362,31 +247,24 @@ class ToolGuardrailDecision: return data +def _canonical_json(value: Any) -> str: + return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str) + + def canonical_tool_args(args: Mapping[str, Any]) -> str: """Return sorted compact JSON for parsed tool arguments.""" if not isinstance(args, Mapping): raise TypeError(f"tool args must be a mapping, got {type(args).__name__}") - return json.dumps( - args, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - default=str, - ) + return _canonical_json(args) def classify_tool_failure(tool_name: str, result: str | None) -> tuple[bool, str]: - """Safety-fallback classifier used only when callers don't pass ``failed``. + """Fallback classifier used only when callers don't pass ``failed``. Mirrors ``agent.display._detect_tool_failure`` exactly so the guardrail - never disagrees with the CLI's user-visible ``[error]`` tag. Production - callers in ``run_agent.py`` always pass an explicit ``failed=`` derived - from ``_detect_tool_failure``; this function exists so standalone callers - (tests, tooling) still get consistent behavior. + never disagrees with the CLI's user-visible ``[error]`` tag. """ - if result is None: - return False, "" - if file_mutation_result_landed(tool_name, result): + if result is None or file_mutation_result_landed(tool_name, result): return False, "" if tool_name == "terminal": @@ -399,9 +277,8 @@ def classify_tool_failure(tool_name: str, result: str | None) -> tuple[bool, str if tool_name == "memory": data = safe_json_loads(result) - if isinstance(data, dict): - if data.get("success") is False and "exceed the limit" in data.get("error", ""): - return True, " [full]" + if isinstance(data, dict) and data.get("success") is False and "exceed the limit" in data.get("error", ""): + return True, " [full]" lower = result[:500].lower() if '"error"' in lower or '"failed"' in lower or result.startswith("Error"): @@ -424,29 +301,17 @@ class ToolCallGuardrailController: self._progress_since_failure: dict[ToolCallSignature, bool] = {} self._no_progress: dict[ToolCallSignature, tuple[str, int]] = {} self._halt_decision: ToolGuardrailDecision | None = None - # Identical-call loop-breaker state (agent.stall_guards): tracks the - # CONSECUTIVE streak of identical (tool, canonical args) calls whose - # results were also identical. Any different call — or a different - # result — resets the streak, so legitimate re-reads after edits and - # varied polling are never flagged. Per-turn, like everything else here. - # NOTE: open PR #85352 (patrykkopycinski) tracks no-progress loops - # ACROSS turns via a detection window — a different mechanism from - # this per-turn consecutive streak. Coordinate future work there. + # Identical-call streak: CONSECUTIVE identical (tool, args) calls with + # identical results. Any different call or result resets it, so re-reads + # after edits and varied polling are never flagged. self._identical_streak_sig: ToolCallSignature | None = None self._identical_streak_result_hash: str = "" self._identical_streak_count: int = 0 - # tool_call_id of the FIRST call in the current streak, so a - # result-reference stub can point at the message that carries the - # full payload. + # tool_call_id of the streak's FIRST call, so a stub can point at the full payload. self._identical_streak_first_call_id: str = "" - # tool_call_id -> spillover file path for results that were persisted - # out of context (persisted-output preview). Lets a reference stub - # carry the file path so the reference can't dangle when the first - # occurrence entered context as a preview. + # tool_call_id -> spillover path, so a stub referencing a result that + # entered context only as a persisted-output preview can't dangle. self._persisted_result_paths: dict[str, str] = {} - # Per-turn runaway-loop cap counters. Reset every turn (this method - # runs at the start of each run_conversation), so the caps bound a - # single agent loop rather than accumulating across the session. self._turn_web_search_count = 0 self._turn_subagent_count = 0 @@ -454,21 +319,29 @@ class ToolCallGuardrailController: def halt_decision(self) -> ToolGuardrailDecision | None: return self._halt_decision - def before_call(self, tool_name: str, args: Mapping[str, Any] | None) -> ToolGuardrailDecision: - signature = ToolCallSignature.from_call(tool_name, _coerce_args(args)) + def _halt( + self, action: str, code: str, message: str, + tool_name: str, count: int, signature: ToolCallSignature, + ) -> ToolGuardrailDecision: + """Build a block/halt decision and record it as the turn's halt decision.""" + self._halt_decision = ToolGuardrailDecision( + action=action, code=code, message=message, + tool_name=tool_name, count=count, signature=signature, + ) + return self._halt_decision - # ── Per-turn runaway-loop caps ────────────────────────────────── - # These are hard ceilings on how many times a runaway-prone tool may - # be called within a single agent loop (turn). They apply regardless - # of hard_stop_enabled (which only governs the per-turn loop detector). - # We block BEFORE the call runs once the count is already at the cap, - # then increment for an allowed call so the (cap+1)-th is refused. - cap_block = self._check_loop_cap(tool_name, _coerce_args(args), signature) + def before_call(self, tool_name: str, args: Mapping[str, Any] | None) -> ToolGuardrailDecision: + args = _coerce_args(args) + signature = ToolCallSignature.from_call(tool_name, args) + allow = ToolGuardrailDecision(tool_name=tool_name, signature=signature) + + # Loop caps apply regardless of hard_stop_enabled (which only governs the detector). + cap_block = self._check_loop_cap(tool_name, args, signature) if cap_block is not None: return cap_block if not self.config.hard_stop_enabled: - return ToolGuardrailDecision(tool_name=tool_name, signature=signature) + return allow exact_count = self._exact_failure_counts.get(signature, 0) if self._progress_since_failure.get(signature): @@ -476,42 +349,27 @@ class ToolCallGuardrailController: # streak restarts in after_call if it fails again. exact_count = 0 if exact_count >= self.config.exact_failure_block_after: - decision = ToolGuardrailDecision( - action="block", - code="repeated_exact_failure_block", - message=( - f"Blocked {tool_name}: the same tool call failed {exact_count} " - "times with identical arguments. Stop retrying it unchanged; " - "change strategy or explain the blocker." - ), - tool_name=tool_name, - count=exact_count, - signature=signature, + return self._halt( + "block", "repeated_exact_failure_block", + f"Blocked {tool_name}: the same tool call failed {exact_count} " + "times with identical arguments. Stop retrying it unchanged; " + "change strategy or explain the blocker.", + tool_name, exact_count, signature, ) - self._halt_decision = decision - return decision if self._is_idempotent(tool_name): record = self._no_progress.get(signature) - if record is not None: - _result_hash, repeat_count = record - if repeat_count >= self.config.no_progress_block_after: - decision = ToolGuardrailDecision( - action="block", - code="idempotent_no_progress_block", - message=( - f"Blocked {tool_name}: this read-only call returned the same " - f"result {repeat_count} times. Stop repeating it unchanged; " - "use the result already provided or try a different query." - ), - tool_name=tool_name, - count=repeat_count, - signature=signature, - ) - self._halt_decision = decision - return decision + if record is not None and record[1] >= self.config.no_progress_block_after: + repeat_count = record[1] + return self._halt( + "block", "idempotent_no_progress_block", + f"Blocked {tool_name}: this read-only call returned the same " + f"result {repeat_count} times. Stop repeating it unchanged; " + "use the result already provided or try a different query.", + tool_name, repeat_count, signature, + ) - return ToolGuardrailDecision(tool_name=tool_name, signature=signature) + return allow def after_call( self, @@ -526,11 +384,16 @@ class ToolCallGuardrailController: if failed is None: failed, _ = classify_tool_failure(tool_name, result) + def warn(code: str, message: str, count: int) -> ToolGuardrailDecision: + return ToolGuardrailDecision( + action="warn", code=code, message=message, + tool_name=tool_name, count=count, signature=signature, + ) + if failed: # An identical failing call is only a REPLAY if nothing landed in - # between. If any mutating call succeeded since the previous - # identical failure (edit -> re-run pytest, click -> re-snapshot), - # the retry is a new experiment: restart the exact-args streak. + # between; a mutation since the last identical failure makes the + # retry a new experiment, so the exact-args streak restarts. if self._progress_since_failure.pop(signature, False): self._exact_failure_counts.pop(signature, None) exact_count = self._exact_failure_counts.get(signature, 0) + 1 @@ -540,53 +403,35 @@ class ToolCallGuardrailController: same_count = self._same_tool_failure_counts.get(tool_name, 0) + 1 self._same_tool_failure_counts[tool_name] = same_count - # same_tool_failure counts DIFFERENT args on one tool. For tools - # whose non-zero exit is ordinary work output (terminal, - # execute_code, pollers) a run of distinct red commands is - # diagnosis, not a loop — warn, never halt. The exact-args replay - # path still applies to them. - same_tool_halt_eligible = tool_name not in FAILURE_TOLERANT_TOOL_NAMES + # same_tool_failure counts DIFFERENT args on one tool; for + # failure-tolerant tools a run of distinct red commands is diagnosis, + # not a loop — warn, never halt (exact-args replay still applies). if ( self.config.hard_stop_enabled - and same_tool_halt_eligible + and tool_name not in FAILURE_TOLERANT_TOOL_NAMES and same_count >= self.config.same_tool_failure_halt_after ): - decision = ToolGuardrailDecision( - action="halt", - code="same_tool_failure_halt", - message=( - f"Stopped {tool_name}: it failed {same_count} times this turn. " - "Stop retrying the same failing tool path and choose a different approach." - ), - tool_name=tool_name, - count=same_count, - signature=signature, + return self._halt( + "halt", "same_tool_failure_halt", + f"Stopped {tool_name}: it failed {same_count} times this turn. " + "Stop retrying the same failing tool path and choose a different approach.", + tool_name, same_count, signature, ) - self._halt_decision = decision - return decision if self.config.warnings_enabled and exact_count >= self.config.exact_failure_warn_after: - return ToolGuardrailDecision( - action="warn", - code="repeated_exact_failure_warning", - message=( - f"{tool_name} has failed {exact_count} times with identical arguments. " - "This looks like a loop; inspect the error and change strategy " - "instead of retrying it unchanged." - ), - tool_name=tool_name, - count=exact_count, - signature=signature, + return warn( + "repeated_exact_failure_warning", + f"{tool_name} has failed {exact_count} times with identical arguments. " + "This looks like a loop; inspect the error and change strategy " + "instead of retrying it unchanged.", + exact_count, ) if self.config.warnings_enabled and same_count >= self.config.same_tool_failure_warn_after: - return ToolGuardrailDecision( - action="warn", - code="same_tool_failure_warning", - message=_tool_failure_recovery_hint(tool_name, same_count), - tool_name=tool_name, - count=same_count, - signature=signature, + return warn( + "same_tool_failure_warning", + _tool_failure_recovery_hint(tool_name, same_count), + same_count, ) return ToolGuardrailDecision(tool_name=tool_name, count=exact_count, signature=signature) @@ -595,10 +440,8 @@ class ToolCallGuardrailController: self._same_tool_failure_counts.pop(tool_name, None) # A successful mutation is progress for every failing signature still - # being counted this turn: the next identical retry runs against - # changed state, so it is a fresh attempt rather than a replay. Pure - # loops never mutate anything between attempts, so the replay detector - # keeps its teeth. + # counted this turn (next identical retry runs against changed state). + # Pure loops never mutate between attempts, so the replay detector keeps its teeth. if tool_name in PROGRESS_RESET_TOOL_NAMES or file_mutation_result_landed(tool_name, result): for sig in list(self._exact_failure_counts): self._progress_since_failure[sig] = True @@ -610,44 +453,22 @@ class ToolCallGuardrailController: result_hash = _result_hash(result) previous = self._no_progress.get(signature) - repeat_count = 1 - if previous is not None and previous[0] == result_hash: - repeat_count = previous[1] + 1 + repeat_count = previous[1] + 1 if previous is not None and previous[0] == result_hash else 1 self._no_progress[signature] = (result_hash, repeat_count) if self.config.warnings_enabled and repeat_count >= self.config.no_progress_warn_after: - return ToolGuardrailDecision( - action="warn", - code="idempotent_no_progress_warning", - message=( - f"{tool_name} returned the same result {repeat_count} times. " - "Use the result already provided or change the query instead of " - "repeating it unchanged." - ), - tool_name=tool_name, - count=repeat_count, - signature=signature, + return warn( + "idempotent_no_progress_warning", + f"{tool_name} returned the same result {repeat_count} times. " + "Use the result already provided or change the query instead of " + "repeating it unchanged.", + repeat_count, ) return ToolGuardrailDecision(tool_name=tool_name, count=repeat_count, signature=signature) def _is_idempotent(self, tool_name: str) -> bool: - if tool_name in self.config.mutating_tools: - return False - return tool_name in self.config.idempotent_tools - - def observe_identical_call( - self, - tool_name: str, - args: Mapping[str, Any] | None, - result: str | None, - ) -> str | None: - """Track consecutive identical calls; return a loop-breaker notice or None. - - Back-compat wrapper around :meth:`observe_call` for callers that only - care about the loop-breaker notice. - """ - return self.observe_call(tool_name, args, result).notice + return tool_name not in self.config.mutating_tools and tool_name in self.config.idempotent_tools def observe_call( self, @@ -660,29 +481,15 @@ class ToolCallGuardrailController: ) -> "IdenticalCallObservation": """Track consecutive identical calls; return notice + dedupe stub info. - Two independent outputs from the same consecutive-streak tracker: - - - ``notice``: the compact loop-breaker notice, fired when the SAME - tool is called with identical canonical arguments AND returns an - identical result for the ``STALL_GUARD_IDENTICAL_CALL_THRESHOLD``-th - (and every subsequent) consecutive time within the turn. Purely - observational — never blocks the call. Allowlisted pollers - (``is_stall_guard_repeatable``) are exempt from the NOTICE. - - ``stub``: a short reference replacement for the CURRENT result, - produced from the 2nd consecutive identical call whose fresh result - is byte-identical to the previous one. The tool still executed — - only the context representation is deduplicated, so polling - semantics are preserved (a changed result flows through whole and - resets the streak). Pollers are NOT exempt from stubbing: for a - poller, an identical result means nothing changed, which is exactly - when the stub saves the most context and loses nothing. Results - under ``IDENTICAL_RESULT_STUB_MIN_CHARS`` and failed/error results - are never stubbed, and only plain-string results are considered. - - Any intervening different call or changed result resets the streak. - Callers substitute/append at tool RESULT construction time, which is - cache-safe: tool results are append-only and never mutate - already-sent context. + ``notice`` fires from the ``STALL_GUARD_IDENTICAL_CALL_THRESHOLD``-th + consecutive identical (tool, args, result) call; purely observational, + pollers exempt. ``stub`` replaces the CURRENT result from the 2nd + byte-identical repeat — the tool still executed, only the context + representation is deduplicated, so polling semantics survive (a changed + result flows through whole and resets the streak). Pollers are NOT + exempt from stubbing: an unchanged poll is exactly when the stub saves + the most. Short results, failed results and non-string results are never + stubbed. Callers substitute at result construction time, which is cache-safe. """ is_plain_str = isinstance(result, str) signature = ToolCallSignature.from_call(tool_name, _coerce_args(args)) @@ -695,8 +502,7 @@ class ToolCallGuardrailController: ): self._identical_streak_count += 1 else: - # New streak (or non-string result, which never forms a streak — - # multimodal content lists pass through untouched). + # New streak; non-string (multimodal) results never form one. self._identical_streak_sig = signature if is_plain_str else None self._identical_streak_result_hash = result_hash self._identical_streak_count = 1 if is_plain_str else 0 @@ -705,10 +511,7 @@ class ToolCallGuardrailController: count = self._identical_streak_count notice = None - if ( - not is_stall_guard_repeatable(tool_name) - and count >= STALL_GUARD_IDENTICAL_CALL_THRESHOLD - ): + if not is_stall_guard_repeatable(tool_name) and count >= STALL_GUARD_IDENTICAL_CALL_THRESHOLD: ordinal = f"{count}{'th' if 11 <= count % 100 <= 13 else {1: 'st', 2: 'nd', 3: 'rd'}.get(count % 10, 'th')}" notice = ( f"[hermes note: this is the {ordinal} consecutive identical call to " @@ -716,62 +519,36 @@ class ToolCallGuardrailController: "Do not repeat it — change arguments, use a different tool, or " "proceed with what you have.]" ) - # Hard-stop widening (#89069 / #100849 bundle): the per-turn - # no-progress BLOCK above only covers tools in idempotent_tools, so - # a model replaying the same successful `terminal`/`skill_view` - # call with a byte-identical result ran until the iteration budget. - # The consecutive-identical streak is tool-agnostic; when hard - # stops are enabled, halt at the same idempotent_no_progress - # threshold. Pollers stay exempt (an unchanged poll is progress). + # The no-progress BLOCK in before_call only covers idempotent_tools; this + # streak is tool-agnostic, so with hard stops on, halt at the same threshold + # (a model replaying a successful `terminal` call otherwise runs to the budget). if ( self.config.hard_stop_enabled and count >= self.config.no_progress_block_after and self._halt_decision is None ): - self._halt_decision = ToolGuardrailDecision( - action="halt", - code="identical_call_streak_halt", - message=( - f"Stopped {tool_name}: the same call with identical arguments " - f"returned the same result {count} times in a row. Stop " - "repeating it unchanged; use the result already provided or " - "change strategy." - ), - tool_name=tool_name, - count=count, - signature=signature, + self._halt( + "halt", "identical_call_streak_halt", + f"Stopped {tool_name}: the same call with identical arguments " + f"returned the same result {count} times in a row. Stop " + "repeating it unchanged; use the result already provided or " + "change strategy.", + tool_name, count, signature, ) stub = None - if ( - is_plain_str - and count >= 2 - and not failed - and len(result) >= IDENTICAL_RESULT_STUB_MIN_CHARS - ): + if is_plain_str and count >= 2 and not failed and len(result) >= IDENTICAL_RESULT_STUB_MIN_CHARS: stub = self._build_result_reference_stub(tool_name, args) return IdenticalCallObservation(notice=notice, stub=stub) def record_persisted_result(self, tool_call_id: str, file_path: str) -> None: - """Remember the spillover path a persisted result was saved to. - - When the first occurrence of a result entered context as a - persisted-output preview, a later reference stub must carry the - spillover file path so the reference can't dangle. - """ + """Remember the spillover path a persisted result was saved to.""" if tool_call_id and file_path: self._persisted_result_paths[tool_call_id] = file_path - def _build_result_reference_stub( - self, tool_name: str, args: Mapping[str, Any] | None - ) -> str: - """Build the reference stub replacing a byte-identical duplicate result. - - Carries the tool name + a canonical-args preview so that even if - context compression later evicts the referenced result, the model - still knows WHAT the call was (cheap dangling-reference mitigation). - """ + def _build_result_reference_stub(self, tool_name: str, args: Mapping[str, Any] | None) -> str: + """Reference stub for a byte-identical duplicate result (tool + args preview).""" try: args_preview = canonical_tool_args(_coerce_args(args)) except TypeError: @@ -799,75 +576,48 @@ class ToolCallGuardrailController: args: Mapping[str, Any], signature: ToolCallSignature, ) -> ToolGuardrailDecision | None: - """Enforce and advance the per-turn runaway-loop counters. + """Block once a per-turn cap is reached, else advance the counter and return None. - Returns a ``block`` decision when the cap is already reached, otherwise - increments the relevant counter for the allowed call and returns - ``None``. A cap of 0 disables that limit entirely. Counters reset each - turn via ``reset_for_turn``. + A cap of 0 disables that limit. Blocking happens BEFORE the call when + the count is already at the cap, so the (cap+1)-th call is refused. """ caps = self.config.loop_caps - if tool_name == "web_search": - cap = caps.max_web_searches - if cap and self._turn_web_search_count >= cap: - decision = ToolGuardrailDecision( - action="block", - code="loop_web_search_cap", - message=( - f"Blocked web_search: this turn has already made {cap} " - "web searches, the per-turn limit. This looks like a " - "runaway search loop. Work with the results you already " - "have and give the user your answer." - ), - tool_name=tool_name, - count=self._turn_web_search_count, - signature=signature, - ) - self._halt_decision = decision - return decision - self._turn_web_search_count += 1 - return None - - if tool_name == "delegate_task": - cap = caps.max_subagents - if not cap: - return None - spawn_count = _subagent_spawn_count(args) - if spawn_count == 0: - # Control action (list/steer/stop) — spawns nothing. Never - # block: once the spawn cap is hit, steering/stopping the - # existing children is exactly what should still work. - return None - if self._turn_subagent_count >= cap: - decision = ToolGuardrailDecision( - action="block", - code="loop_subagent_cap", - message=( - f"Blocked delegate_task: this turn has already spawned " - f"{self._turn_subagent_count} subagents (limit {cap}). " - "This looks like a runaway delegation loop. Finish the " - "work with the results you have and answer the user." - ), - tool_name=tool_name, - count=self._turn_subagent_count, - signature=signature, - ) - self._halt_decision = decision - return decision - self._turn_subagent_count += spawn_count + cap, count, increment = caps.max_web_searches, self._turn_web_search_count, 1 + code, message = "loop_web_search_cap", ( + f"Blocked web_search: this turn has already made {cap} " + "web searches, the per-turn limit. This looks like a " + "runaway search loop. Work with the results you already " + "have and give the user your answer." + ) + elif tool_name == "delegate_task": + cap, count = caps.max_subagents, self._turn_subagent_count + # Control actions (list/steer/stop) spawn nothing and must keep working after the cap is hit. + increment = _subagent_spawn_count(args) if cap else 0 + if increment == 0: + return None + code, message = "loop_subagent_cap", ( + f"Blocked delegate_task: this turn has already spawned " + f"{count} subagents (limit {cap}). " + "This looks like a runaway delegation loop. Finish the " + "work with the results you have and answer the user." + ) + else: return None + if cap and count >= cap: + return self._halt("block", code, message, tool_name, count, signature) + if tool_name == "web_search": + self._turn_web_search_count += increment + else: + self._turn_subagent_count += increment return None def toolguard_synthetic_result(decision: ToolGuardrailDecision) -> str: """Build a synthetic role=tool content string for a blocked tool call.""" return json.dumps( - { - "error": decision.message, - "guardrail": decision.to_metadata(), - }, + {"error": decision.message, "guardrail": decision.to_metadata()}, ensure_ascii=False, ) @@ -877,11 +627,7 @@ def append_toolguard_guidance(result: str, decision: ToolGuardrailDecision) -> s if decision.action not in {"warn", "halt"} or not decision.message: return result label = "Tool loop hard stop" if decision.action == "halt" else "Tool loop warning" - suffix = ( - f"\n\n[{label}: " - f"{decision.code}; count={decision.count}; {decision.message}]" - ) - return (result or "") + suffix + return (result or "") + f"\n\n[{label}: {decision.code}; count={decision.count}; {decision.message}]" def _tool_failure_recovery_hint(tool_name: str, count: int) -> str: @@ -912,13 +658,7 @@ def _result_hash(result: str | None) -> str: parsed = safe_json_loads(result or "") if parsed is not None: try: - canonical = json.dumps( - parsed, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - default=str, - ) + canonical = _canonical_json(parsed) except TypeError: canonical = str(parsed) else: @@ -942,50 +682,27 @@ def _as_bool(value: Any, default: bool) -> bool: return default -def _positive_int(value: Any, default: int) -> int: +def _int_at_least(value: Any, default: int, minimum: int) -> int: + """Int parser: junk/None/below-minimum fall back to default (caps use minimum 0 so 0 = disabled).""" if value is None: return default try: parsed = int(value) except (TypeError, ValueError): return default - return parsed if parsed >= 1 else default - - -def _non_negative_int(value: Any, default: int) -> int: - """Parse a session-cap value. 0 is a valid (disable) value; negatives and - junk fall back to the default.""" - if value is None: - return default - try: - parsed = int(value) - except (TypeError, ValueError): - return default - return parsed if parsed >= 0 else default + return parsed if parsed >= minimum else default def _subagent_spawn_count(args: Mapping[str, Any]) -> int: - """How many subagents a single delegate_task call spawns. - - delegate_task runs in one of two modes: a batch (``tasks`` is a non-empty - list, one child per item) or a single task (``goal``). Count the batch size - when present, otherwise 1, so the session subagent cap reflects real spawns - rather than delegate_task invocations. Control actions (list/steer/stop) - spawn nothing and must not consume the cap. - """ - if isinstance(args, Mapping): - action = str(args.get("action") or "").strip().lower() - if action in ("list", "steer", "stop"): - return 0 - tasks = args.get("tasks") if isinstance(args, Mapping) else None - if isinstance(tasks, list) and tasks: - return len(tasks) - return 1 + """Subagents one delegate_task call spawns: batch size when ``tasks`` is a + non-empty list, else 1; control actions (list/steer/stop) spawn 0.""" + if str(args.get("action") or "").strip().lower() in ("list", "steer", "stop"): + return 0 + tasks = args.get("tasks") + return len(tasks) if isinstance(tasks, list) and tasks else 1 def _sha256(value: str) -> str: - # surrogatepass: tool results scraped from the web can carry unpaired - # UTF-16 surrogates (e.g. half of a mathematical-bold pair); a strict - # encode raises and takes down the whole conversation loop. The hash only - # needs deterministic bytes, not valid UTF-8. + # surrogatepass: web-scraped results can carry unpaired UTF-16 surrogates; a + # strict encode would raise and take down the conversation loop. return hashlib.sha256(value.encode("utf-8", "surrogatepass")).hexdigest() diff --git a/tests/agent/test_stall_guards.py b/tests/agent/test_stall_guards.py index 013f8b83f8..a0909dffb6 100644 --- a/tests/agent/test_stall_guards.py +++ b/tests/agent/test_stall_guards.py @@ -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 # block that maybe_persist_tool_result builds — assert the round trip.