refactor(agent/tool_guardrails): drop redundant defensive branches, fold config defaults, compact comments

This commit is contained in:
Teknium
2026-09-02 18:22:33 -07:00
parent 113f04616b
commit b47a5fedab
+109 -176
View File
@@ -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 (``<vendor>_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")