From b47a5fedab151f1fd4bf4951696a6ecde1e257e0 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:22:33 -0700 Subject: [PATCH] refactor(agent/tool_guardrails): drop redundant defensive branches, fold config defaults, compact comments --- agent/tool_guardrails.py | 285 +++++++++++++++------------------------ 1 file changed, 109 insertions(+), 176 deletions(-) diff --git a/agent/tool_guardrails.py b/agent/tool_guardrails.py index 6ec3b289ae..b557de38e9 100644 --- a/agent/tool_guardrails.py +++ b/agent/tool_guardrails.py @@ -9,7 +9,7 @@ from __future__ import annotations import hashlib import json -from dataclasses import dataclass, field +from dataclasses import asdict, dataclass, field from typing import Any, Mapping from utils import safe_json_loads @@ -31,22 +31,14 @@ MUTATING_TOOL_NAMES = frozenset({ "send_message", "cronjob_manage", "delegate_task", "process_manage", }) -# Tools legitimately re-invoked with identical args while waiting on external -# progress (pollers). The identical-call NOTICE never fires for these. +# Pollers: legitimately re-invoked with identical args; the identical-call NOTICE never fires. STALL_GUARD_REPEATABLE_TOOLS = frozenset({"process_manage"}) -# Poller naming conventions on generated / MCP surfaces (``_get_result``). -_STALL_GUARD_REPEATABLE_SUFFIXES = ("_get_result", "_poll") - -# 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_REPEATABLE_SUFFIXES = ("_get_result", "_poll") # generated / MCP poller conventions +# Nth consecutive identical (tool, args, result) call that fires the notice; 3 tolerates one double-check. STALL_GUARD_IDENTICAL_CALL_THRESHOLD = 3 - -# 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. +# From the 2nd byte-identical repeat the duplicate payload becomes a reference stub; smaller results +# aren't worth it, errors never are. The args preview keeps WHAT was called if compression evicts the original. IDENTICAL_RESULT_STUB_MIN_CHARS = 512 -# 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 @@ -64,12 +56,59 @@ PROGRESS_RESET_TOOL_NAMES = frozenset({ "send_message", "cronjob", "cronjob_manage", "todo", "todo_list", "memory", "skill_manage", }) +# 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). +_DEFAULT_MAX_WEB_SEARCHES_PER_TURN = 50 +_DEFAULT_MAX_SUBAGENTS_PER_TURN = 50 + +_INTERACTIVE_PLATFORMS = frozenset({"cli", "tui", "desktop", "acp"}) +# Bounded supervised task loops (subagent stopped by its parent; api_server has a live +# client) doing real edit -> re-run work keep the interactive warn-only default. +_SUPERVISED_TASK_PLATFORMS = frozenset({"subagent", "api_server"}) + def is_stall_guard_repeatable(tool_name: str) -> bool: """Whether a tool is exempt from the identical-call loop notice.""" - return tool_name in STALL_GUARD_REPEATABLE_TOOLS or tool_name.endswith( - _STALL_GUARD_REPEATABLE_SUFFIXES - ) + return tool_name in STALL_GUARD_REPEATABLE_TOOLS or tool_name.endswith(_STALL_GUARD_REPEATABLE_SUFFIXES) + + +def _is_non_interactive_platform(platform: str | None) -> bool: + """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() + return key not in _INTERACTIVE_PLATFORMS and key not in _SUPERVISED_TASK_PLATFORMS + + +@dataclass(frozen=True) +class LoopCapConfig: + """Per-turn hard ceilings on web_search calls / subagent spawns. + + Unlike the loop detector 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 + max_subagents: int = _DEFAULT_MAX_SUBAGENTS_PER_TURN + + @classmethod + def from_mapping(cls, data: Mapping[str, Any] | None) -> "LoopCapConfig": + """Build config from the ``tool_loop_guardrails.loop_caps`` section.""" + if not isinstance(data, Mapping): + return cls() + return cls( + max_web_searches=_int_at_least(data.get("max_web_searches"), _DEFAULT_MAX_WEB_SEARCHES_PER_TURN, 0), + max_subagents=_int_at_least(data.get("max_subagents"), _DEFAULT_MAX_SUBAGENTS_PER_TURN, 0), + ) @dataclass(frozen=True) @@ -92,14 +131,11 @@ class ToolCallGuardrailConfig: no_progress_block_after: int = 5 idempotent_tools: frozenset[str] = field(default_factory=lambda: IDEMPOTENT_TOOL_NAMES) mutating_tools: frozenset[str] = field(default_factory=lambda: MUTATING_TOOL_NAMES) - loop_caps: "LoopCapConfig" = field(default_factory=lambda: LoopCapConfig()) + loop_caps: LoopCapConfig = field(default_factory=LoopCapConfig) @classmethod def from_mapping( - cls, - data: Mapping[str, Any] | None, - *, - platform: str | None = None, + cls, data: Mapping[str, Any] | None, *, platform: str | None = None, ) -> "ToolCallGuardrailConfig": """Build config from the `tool_loop_guardrails` config.yaml section. @@ -109,10 +145,10 @@ class ToolCallGuardrailConfig: data = {} d = cls() hard_stop_enabled = _as_bool(data.get("hard_stop_enabled"), d.hard_stop_enabled) - non_interactive_hard_stop_enabled = _as_bool( + non_interactive = _as_bool( 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: + if non_interactive and _is_non_interactive_platform(platform): hard_stop_enabled = True thresholds: dict[str, int] = {} @@ -127,73 +163,15 @@ class ToolCallGuardrailConfig: return cls( 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, + non_interactive_hard_stop_enabled=non_interactive, loop_caps=LoopCapConfig.from_mapping(data.get("loop_caps")), **thresholds, ) -# 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 hard ceilings on web_search calls / subagent spawns. - - 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 - max_subagents: int = _DEFAULT_MAX_SUBAGENTS_PER_TURN - - @classmethod - def from_mapping(cls, data: Mapping[str, Any] | None) -> "LoopCapConfig": - """Build config from the ``tool_loop_guardrails.loop_caps`` section.""" - if not isinstance(data, Mapping): - return cls() - defaults = cls() - return cls( - 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"}) -# 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: - """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() - return key not in _INTERACTIVE_PLATFORMS and key not in _SUPERVISED_TASK_PLATFORMS - - @dataclass(frozen=True) class IdenticalCallObservation: - """Outcome of observing one completed call: ``notice`` is appended after the - result, ``stub`` replaces a byte-identical duplicate result. Both may be set.""" + """``notice`` is appended after the result, ``stub`` replaces a byte-identical duplicate result.""" notice: str | None = None stub: str | None = None @@ -235,15 +213,9 @@ class ToolGuardrailDecision: return self.action in {"block", "halt"} def to_metadata(self) -> dict[str, Any]: - data: dict[str, Any] = { - "action": self.action, - "code": self.code, - "message": self.message, - "tool_name": self.tool_name, - "count": self.count, - } - if self.signature is not None: - data["signature"] = self.signature.to_metadata() + data = asdict(self) + if data["signature"] is None: + del data["signature"] return data @@ -261,19 +233,16 @@ def canonical_tool_args(args: Mapping[str, Any]) -> str: def classify_tool_failure(tool_name: str, result: str | None) -> tuple[bool, str]: """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. + Mirrors ``agent.display._detect_tool_failure`` so the guardrail never + disagrees with the CLI's user-visible ``[error]`` tag. """ if result is None or file_mutation_result_landed(tool_name, result): return False, "" if tool_name == "terminal": data = safe_json_loads(result) - if isinstance(data, dict): - exit_code = data.get("exit_code") - if exit_code is not None and exit_code != 0: - return True, f" [exit {exit_code}]" - return False, "" + exit_code = data.get("exit_code") if isinstance(data, dict) else None + return (True, f" [exit {exit_code}]") if exit_code is not None and exit_code != 0 else (False, "") if tool_name == "memory": data = safe_json_loads(result) @@ -301,16 +270,14 @@ 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 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. + # 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. first_call_id lets a stub point at the full payload. self._identical_streak_sig: ToolCallSignature | None = None self._identical_streak_result_hash: str = "" self._identical_streak_count: int = 0 - # 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 path, so a stub referencing a result that - # entered context only as a persisted-output preview can't dangle. + # tool_call_id -> spillover path, so a stub referencing a persisted-output preview can't dangle. self._persisted_result_paths: dict[str, str] = {} self._turn_web_search_count = 0 self._turn_subagent_count = 0 @@ -343,11 +310,8 @@ class ToolCallGuardrailController: if not self.config.hard_stop_enabled: return allow - exact_count = self._exact_failure_counts.get(signature, 0) - if self._progress_since_failure.get(signature): - # Something landed since this call last failed — let it run; the - # streak restarts in after_call if it fails again. - exact_count = 0 + # A mutation since this call last failed makes the retry a new experiment. + exact_count = 0 if self._progress_since_failure.get(signature) else self._exact_failure_counts.get(signature, 0) if exact_count >= self.config.exact_failure_block_after: return self._halt( "block", "repeated_exact_failure_block", @@ -372,12 +336,8 @@ class ToolCallGuardrailController: return allow def after_call( - self, - tool_name: str, - args: Mapping[str, Any] | None, - result: str | None, - *, - failed: bool | None = None, + self, tool_name: str, args: Mapping[str, Any] | None, result: str | None, + *, failed: bool | None = None, ) -> ToolGuardrailDecision: args = _coerce_args(args) signature = ToolCallSignature.from_call(tool_name, args) @@ -391,9 +351,8 @@ class ToolCallGuardrailController: ) if failed: - # An identical failing call is only a REPLAY if nothing landed in - # between; a mutation since the last identical failure makes the - # retry a new experiment, so the exact-args streak restarts. + # An identical failing call is only a REPLAY if nothing landed in between; + # a mutation since the last identical failure restarts the exact-args streak. 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 @@ -403,9 +362,8 @@ 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 - # failure-tolerant tools a run of distinct red commands is diagnosis, - # not a loop — warn, never halt (exact-args replay still applies). + # 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. if ( self.config.hard_stop_enabled and tool_name not in FAILURE_TOLERANT_TOOL_NAMES @@ -439,9 +397,8 @@ class ToolCallGuardrailController: self._exact_failure_counts.pop(signature, None) self._same_tool_failure_counts.pop(tool_name, None) - # A successful mutation is progress for every failing signature still - # counted this turn (next identical retry runs against changed state). - # Pure loops never mutate between attempts, so the replay detector keeps its teeth. + # A successful mutation is progress for every failing signature still counted + # this turn. 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 @@ -471,25 +428,18 @@ class ToolCallGuardrailController: return tool_name not in self.config.mutating_tools and tool_name in self.config.idempotent_tools def observe_call( - self, - tool_name: str, - args: Mapping[str, Any] | None, - result: str | None, - *, - tool_call_id: str = "", - failed: bool = False, - ) -> "IdenticalCallObservation": + self, tool_name: str, args: Mapping[str, Any] | None, result: str | None, + *, tool_call_id: str = "", failed: bool = False, + ) -> IdenticalCallObservation: """Track consecutive identical calls; return notice + dedupe stub info. - ``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. + ``notice`` fires from the ``STALL_GUARD_IDENTICAL_CALL_THRESHOLD``-th consecutive + identical (tool, args, result) call; 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 most. + Short, failed and non-string results are never stubbed. """ is_plain_str = isinstance(result, str) signature = ToolCallSignature.from_call(tool_name, _coerce_args(args)) @@ -519,9 +469,9 @@ class ToolCallGuardrailController: "Do not repeat it — change arguments, use a different tool, or " "proceed with what you have.]" ) - # 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). + # 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 @@ -549,10 +499,7 @@ class ToolCallGuardrailController: 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: - args_preview = "{}" + args_preview = canonical_tool_args(_coerce_args(args)) if len(args_preview) > _RESULT_STUB_ARGS_PREVIEW_CHARS: args_preview = args_preview[:_RESULT_STUB_ARGS_PREVIEW_CHARS] + "…" first_id = self._identical_streak_first_call_id @@ -571,10 +518,7 @@ class ToolCallGuardrailController: return stub def _check_loop_cap( - self, - tool_name: str, - args: Mapping[str, Any], - signature: ToolCallSignature, + self, tool_name: str, args: Mapping[str, Any], signature: ToolCallSignature, ) -> ToolGuardrailDecision | None: """Block once a per-turn cap is reached, else advance the counter and return None. @@ -616,10 +560,7 @@ class ToolCallGuardrailController: 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()}, - ensure_ascii=False, - ) + return json.dumps({"error": decision.message, "guardrail": decision.to_metadata()}, ensure_ascii=False) def append_toolguard_guidance(result: str, decision: ToolGuardrailDecision) -> str: @@ -656,36 +597,28 @@ def _coerce_args(args: Mapping[str, Any] | None) -> Mapping[str, Any]: def _result_hash(result: str | None) -> str: parsed = safe_json_loads(result or "") - if parsed is not None: - try: - canonical = _canonical_json(parsed) - except TypeError: - canonical = str(parsed) - else: - canonical = result or "" + canonical = _canonical_json(parsed) if parsed is not None else (result or "") return _sha256(canonical) +_TRUE_WORDS = frozenset({"1", "true", "yes", "on", "enabled"}) +_FALSE_WORDS = frozenset({"0", "false", "no", "off", "disabled"}) + + def _as_bool(value: Any, default: bool) -> bool: - if value is None: - return default - if isinstance(value, bool): - return value - if isinstance(value, (int, float)): + if isinstance(value, (bool, int, float)): return bool(value) if isinstance(value, str): lowered = value.strip().lower() - if lowered in {"1", "true", "yes", "on", "enabled"}: + if lowered in _TRUE_WORDS: return True - if lowered in {"0", "false", "no", "off", "disabled"}: + if lowered in _FALSE_WORDS: return False return default 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): @@ -694,8 +627,8 @@ def _int_at_least(value: Any, default: int, minimum: int) -> int: def _subagent_spawn_count(args: Mapping[str, Any]) -> int: - """Subagents one delegate_task call spawns: batch size when ``tasks`` is a - non-empty list, else 1; control actions (list/steer/stop) spawn 0.""" + """Subagents one delegate_task call spawns: ``len(tasks)`` for a non-empty batch, + 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")