"""Middleware to detect and break repetitive tool call loops. P0 safety: prevents the agent from calling the same tool with the same arguments indefinitely until the recursion limit kills the run. Detection strategy: 1. After each model response, hash the tool calls (name + args). 2. Track recent hashes in a sliding window. 3. If the same hash appears >= warn_threshold times, queue a "you are repeating yourself — wrap up" warning for the current thread/run. The warning is **injected at the next model call** (in ``wrap_model_call``) as a ``HumanMessage`` appended to the message list, *after* all ToolMessage responses to the previous AIMessage(tool_calls). 4. If it appears >= hard_limit times, strip all tool_calls from the response so the agent is forced to produce a final text answer. Why the warning is injected at ``wrap_model_call`` instead of ``after_model``: ``after_model`` fires immediately after the model emits an ``AIMessage`` that may carry ``tool_calls``. The tools node has not run yet, so no matching ``ToolMessage`` exists in the history. Any message we add here lands *between* the assistant's tool_calls and their responses. OpenAI/Moonshot reject the next request with ``"tool_call_ids did not have response messages"`` because their validators require the assistant's tool_calls to be followed immediately by tool messages. Anthropic also disallows mid-stream ``SystemMessage``. By deferring the warning to ``wrap_model_call``, every prior ToolMessage is already present in the request's message list and the warning is appended at the end — pairing intact, no ``AIMessage`` semantics are mutated. Queued warnings are intentionally transient. If a run ends before the next model request drains a queued warning, ``after_agent`` drops it instead of carrying it into a later invocation for the same thread. The hard-stop path still forces termination when the configured safety limit is reached. Detection histories and warning-suppression state are scoped by ``(thread_id, run_id)`` because one compiled graph can serve many runs for the same conversation. They deliberately survive ``after_agent``: a single Gateway run may re-enter that graph for hidden goal continuations, and those continuations share one loop budget. A later user run receives a fresh budget even when it reuses the graph. Standalone library invocations that omit ``run_id`` receive an opaque fallback ID anchored to LangGraph's run-scoped ``Runtime.control`` object, so replacement ``Runtime`` wrappers share one budget within an invocation while a later invocation starts fresh. Stop-reason surfacing (#3875 Phase 2): Like the token-budget guard, the loop hard stop does NOT raise — it strips ``tool_calls`` so the agent loop terminates naturally with a final answer. To let the caller (the subagent executor) distinguish a loop-capped completion from a clean one, the run that triggered the hard stop is recorded in ``_stop_reason`` and exposed via :meth:`consume_stop_reason`. The executor collects that reason alongside the token-budget guard's so a loop-capped run surfaces as ``completed + loop_capped`` and the lead/ledger can tell it was capped without parsing result text. """ from __future__ import annotations import hashlib import json import logging import threading import uuid from collections import Counter, OrderedDict, defaultdict, deque from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import TYPE_CHECKING, Literal, override from langchain.agents import AgentState from langchain.agents.middleware import AgentMiddleware from langchain.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse from langchain_core.messages import HumanMessage from langgraph.runtime import Runtime from deerflow.agents.middlewares._bounded_dict import BoundedDict from deerflow.agents.middlewares.audit_context import ( LOOP_DETECTION_RECORDER_CONTEXT_KEY, resolve_audit_recorder, ) from deerflow.agents.middlewares.tool_call_metadata import clone_ai_message_with_tool_calls from deerflow.runtime.events.catalog import MIDDLEWARE_LOOP_DETECTION_TAG if TYPE_CHECKING: from deerflow.config.loop_detection_config import LoopDetectionConfig logger = logging.getLogger(__name__) # Defaults — can be overridden via constructor _DEFAULT_WARN_THRESHOLD = 3 # inject warning after 3 identical calls _DEFAULT_HARD_LIMIT = 5 # force-stop after 5 identical calls _DEFAULT_WINDOW_SIZE = 20 # track last N tool calls _DEFAULT_MAX_TRACKED_THREADS = 100 # LRU limit for tracked thread/run scopes _DEFAULT_TOOL_FREQ_WARN = 30 # warn after 30 calls to the same tool type _DEFAULT_TOOL_FREQ_HARD_LIMIT = 50 # force-stop after 50 calls to the same tool type _MAX_PENDING_WARNINGS_PER_RUN = 4 type _RunScopeKey = tuple[str, str | None] def _normalize_tool_call_args(raw_args: object) -> tuple[dict, str | None]: """Normalize tool call args to a dict plus an optional fallback key. Some providers serialize ``args`` as a JSON string instead of a dict. We defensively parse those cases so loop detection does not crash while still preserving a stable fallback key for non-dict payloads. """ if isinstance(raw_args, dict): return raw_args, None if isinstance(raw_args, str): try: parsed = json.loads(raw_args) except (TypeError, ValueError, json.JSONDecodeError): return {}, raw_args if isinstance(parsed, dict): return parsed, None return {}, json.dumps(parsed, sort_keys=True, default=str) if raw_args is None: return {}, None return {}, json.dumps(raw_args, sort_keys=True, default=str) def _stable_tool_key(name: str, args: dict, fallback_key: str | None) -> str: """Derive a stable key from salient args without overfitting to noise.""" if name == "read_file" and fallback_key is None: path = args.get("path") or "" start_line = args.get("start_line") end_line = args.get("end_line") bucket_size = 200 try: start_line = int(start_line) if start_line is not None else 1 except (TypeError, ValueError): start_line = 1 try: end_line = int(end_line) if end_line is not None else start_line except (TypeError, ValueError): end_line = start_line start_line, end_line = sorted((start_line, end_line)) bucket_start = max(start_line, 1) bucket_end = max(end_line, 1) bucket_start = (bucket_start - 1) // bucket_size bucket_end = (bucket_end - 1) // bucket_size return f"{path}:{bucket_start}-{bucket_end}" # write_file / str_replace are content-sensitive: same path may be updated # with different payloads during iteration. Using only salient fields (path) # can collapse distinct calls, so we hash full args to reduce false positives. if name in {"write_file", "str_replace"}: if fallback_key is not None: return fallback_key return json.dumps(args, sort_keys=True, default=str) salient_fields = ("path", "url", "query", "command", "pattern", "glob", "cmd") stable_args = {field: args[field] for field in salient_fields if args.get(field) is not None} if stable_args: return json.dumps(stable_args, sort_keys=True, default=str) if fallback_key is not None: return fallback_key return json.dumps(args, sort_keys=True, default=str) def _hash_tool_calls(tool_calls: list[dict]) -> str: """Deterministic hash of a set of tool calls (name + stable key). This is intended to be order-independent: the same multiset of tool calls should always produce the same hash, regardless of their input order. """ # Normalize each tool call to a stable (name, key) structure. normalized: list[str] = [] for tc in tool_calls: name = tc.get("name", "") args, fallback_key = _normalize_tool_call_args(tc.get("args", {})) key = _stable_tool_key(name, args, fallback_key) normalized.append(f"{name}:{key}") # Sort so permutations of the same multiset of calls yield the same ordering. normalized.sort() blob = json.dumps(normalized, sort_keys=True, default=str) return hashlib.md5(blob.encode()).hexdigest()[:12] _WARNING_MSG = "[LOOP DETECTED] You are repeating the same tool calls. Stop calling tools and produce your final answer now. If you cannot complete the task, summarize what you accomplished so far." _TOOL_FREQ_WARNING_MSG = ( "[LOOP DETECTED] You have called {tool_name} {count} times without producing a final answer. Stop calling tools and produce your final answer now. If you cannot complete the task, summarize what you accomplished so far." ) _HARD_STOP_MSG = "[FORCED STOP] Repeated tool calls exceeded the safety limit. Producing final answer with results collected so far." _TOOL_FREQ_HARD_STOP_MSG = "[FORCED STOP] Tool {tool_name} called {count} times — exceeded the per-tool safety limit. Producing final answer with results collected so far." @dataclass(frozen=True) class _LoopDecision: """A loop-detection transition that may be persisted for audit.""" message: str action: Literal["warn", "hard_stop"] detection_layer: Literal["identical_call_set", "tool_frequency"] tool_names: tuple[str, ...] count: int threshold: int @property def hard_stop(self) -> bool: return self.action == "hard_stop" class LoopDetectionMiddleware(AgentMiddleware[AgentState]): """Detects and breaks repetitive tool call loops. Threshold parameters are validated upstream by :class:`LoopDetectionConfig`; construct via :meth:`from_config` to ensure values pass Pydantic validation. Args: warn_threshold: Number of identical tool call sets before injecting a warning message. Default: 3. hard_limit: Number of identical tool call sets before stripping tool_calls entirely. Default: 5. window_size: Size of the sliding window for tracking calls. Default: 20. max_tracked_threads: Maximum number of thread/run scopes to track before evicting the least recently used. The configuration name is retained for compatibility. Default: 100. tool_freq_warn: Maximum number of same-tool-type calls within a sliding window of ``_tool_freq_window`` before injecting a frequency warning. Catches cross-file read loops that hash-based detection misses. Default: 30 (within a window of 50). tool_freq_hard_limit: Maximum number of same-tool-type calls within a sliding window of ``_tool_freq_window`` before forcing a stop. Default: 50 (within a window of 50). tool_freq_overrides: Per-tool overrides for frequency thresholds, keyed by tool name. Each value is a ``(warn, hard_limit)`` tuple that replaces ``tool_freq_warn`` / ``tool_freq_hard_limit`` for that specific tool. Tools not listed here fall back to the global thresholds. Useful for raising limits on intentionally high-frequency tools (e.g. ``bash`` in batch pipelines) without weakening protection on all other tools. Default: ``None`` (no overrides). """ def __init__( self, warn_threshold: int = _DEFAULT_WARN_THRESHOLD, hard_limit: int = _DEFAULT_HARD_LIMIT, window_size: int = _DEFAULT_WINDOW_SIZE, max_tracked_threads: int = _DEFAULT_MAX_TRACKED_THREADS, tool_freq_warn: int = _DEFAULT_TOOL_FREQ_WARN, tool_freq_hard_limit: int = _DEFAULT_TOOL_FREQ_HARD_LIMIT, tool_freq_overrides: dict[str, tuple[int, int]] | None = None, ): super().__init__() self.warn_threshold = warn_threshold self.hard_limit = hard_limit self.window_size = window_size self.max_tracked_threads = max_tracked_threads self.tool_freq_warn = tool_freq_warn self.tool_freq_hard_limit = tool_freq_hard_limit self._default_tool_freq_thresholds = (tool_freq_warn, tool_freq_hard_limit) self._tool_freq_overrides: dict[str, tuple[int, int]] = tool_freq_overrides or {} # Layer 2's windowed frequency count can never exceed the deque length, # so the deque MUST be at least as long as the largest hard limit it is # compared against — otherwise the hard-stop branch is dead code. Do NOT # reuse Layer 1's ``window_size`` (which is unrelated and defaults below # the freq thresholds, e.g. 20 < hard 50); size the frequency window to # the largest hard limit in play (global + every per-tool override) so a # tight burst can actually reach it while spread-out calls still decay # out of the window. Warn thresholds are intentionally excluded: a sane # config enforces warn <= hard (covered by sizing to hard), and a misconfig # with warn > hard would hard-stop first anyway, so an unreachable warn # is harmless and must not inflate the window. self._tool_freq_window = max( self.window_size, self.tool_freq_hard_limit, *(hard for _, hard in self._tool_freq_overrides.values()), ) self._lock = threading.Lock() # LangGraph replaces Runtime per graph node but retains one RunControl # object for the whole invocation. Map that stable anchor to an opaque # token when embedders omit run_id. Keeping the anchor strongly referenced # also prevents CPython from reusing its address while the mapping is live; # after_agent releases normal invocations and the cap bounds abnormal ones. self._fallback_run_ids: OrderedDict[int, tuple[object, str]] = OrderedDict() self._max_fallback_run_ids = max(1, self.max_tracked_threads * 2) self._history: OrderedDict[_RunScopeKey, list[str]] = OrderedDict() self._warned: dict[_RunScopeKey, set[str]] = defaultdict(set) # Windowed per-tool-type frequency: recent tool names per run scope, # trimmed to ``window_size`` so the count decays instead of growing # monotonically (replaces the old monotonic ``_tool_freq`` integer). self._tool_name_history: defaultdict[_RunScopeKey, deque[str]] = defaultdict(deque) # Per-run Counter mirroring the deque so freq_count is O(1) instead # of scanning the whole window on every tool call. A single high # per-tool override (e.g. bash: {hard_limit: 1000}) inflates the window # globally, so the scan would cost 1000 per call for every tool; Counter # increments on append and decrements on popleft. self._tool_name_counter: defaultdict[_RunScopeKey, Counter[str]] = defaultdict(Counter) # Per-run set of tool names already warned about in Layer 2, so a # frequency warning is enqueued once rather than on every subsequent # call. Cleared per name when the windowed count decays back below the # warn threshold, mirroring the hash-layer ``_warned`` pruning. self._tool_freq_warned: dict[_RunScopeKey, set[str]] = defaultdict(set) # Per-thread/run queue of warnings to inject at the next model call. # Populated by ``after_model`` (detection) and drained by # ``wrap_model_call`` (injection); see module docstring. self._pending_warnings: dict[_RunScopeKey, list[str]] = defaultdict(list) self._pending_warning_touch_order: OrderedDict[_RunScopeKey, None] = OrderedDict() self._max_pending_warning_keys = max(1, self.max_tracked_threads * 2) # Stop reason set when a hard-stop fires (#3875 Phase 2). Keyed by run_id # (matching ``TokenBudgetMiddleware``) and bounded — the lead agent's # middleware instance is long-lived across many runs, so without a cap # an entry would accumulate for every looped lead run. Intentionally NOT # cleared by ``after_agent``/``_clear_current_run_pending_warnings`` so # the subagent executor can consume it after the run returns; ``reset()`` # still drops it. The parallel bounded owner map lets # ``reset(thread_id)`` remove only that thread's unconsumed reasons; # both maps are written, evicted, and popped together under ``_lock``. self._stop_reason: BoundedDict[str | None, str] = BoundedDict(1000) self._stop_reason_thread_id: BoundedDict[str | None, str] = BoundedDict(1000) def release_policy_parameters(self) -> dict[str, object]: return { "warn_threshold": self.warn_threshold, "hard_limit": self.hard_limit, "window_size": self.window_size, "max_tracked_threads": self.max_tracked_threads, "tool_freq_warn": self.tool_freq_warn, "tool_freq_hard_limit": self.tool_freq_hard_limit, "tool_freq_overrides": self._tool_freq_overrides, } @classmethod def from_config(cls, config: LoopDetectionConfig) -> LoopDetectionMiddleware: """Construct from a Pydantic-validated config, trusting its validation.""" return cls( warn_threshold=config.warn_threshold, hard_limit=config.hard_limit, window_size=config.window_size, max_tracked_threads=config.max_tracked_threads, tool_freq_warn=config.tool_freq_warn, tool_freq_hard_limit=config.tool_freq_hard_limit, tool_freq_overrides={name: (o.warn, o.hard_limit) for name, o in config.tool_freq_overrides.items()}, ) def _get_thread_id(self, runtime: Runtime) -> str: """Extract thread_id from runtime context for per-thread tracking.""" thread_id = runtime.context.get("thread_id") if runtime.context else None if thread_id: return str(thread_id) return "default" def _get_run_id(self, runtime: Runtime) -> str | None: """Extract run_id from runtime context for per-run warning scoping. Context presence is authoritative, including an explicit ``None``: ``SubagentExecutor`` later consumes the stop reason with its raw, possibly-None run_id, so normalizing that value would lose the signal. A RunnableConfig run_id exposed through ``Runtime.execution_info`` is the next-best stable identifier. If neither source provides one, use LangGraph's run-scoped ``Runtime.control`` object as the invocation anchor. LangGraph creates replacement Runtime wrappers per graph node, but preserves that control object across the invocation. The anchor is mapped to an opaque generated token instead of embedding ``id(anchor)`` in the key, because CPython may reuse an address after garbage collection. The bounded map keeps a strong reference while active and is released by ``after_agent`` on the normal completion path. """ ctx = getattr(runtime, "context", None) if isinstance(ctx, dict) and "run_id" in ctx: return ctx["run_id"] execution_info = getattr(runtime, "execution_info", None) execution_run_id = getattr(execution_info, "run_id", None) if execution_run_id is not None: return str(execution_run_id) control = getattr(runtime, "control", None) anchor = control if control is not None else runtime anchor_id = id(anchor) with self._lock: existing = self._fallback_run_ids.get(anchor_id) if existing is not None and existing[0] is anchor: self._fallback_run_ids.move_to_end(anchor_id) return existing[1] fallback_run_id = f"__invocation__:{uuid.uuid4().hex}" self._fallback_run_ids[anchor_id] = (anchor, fallback_run_id) self._fallback_run_ids.move_to_end(anchor_id) while len(self._fallback_run_ids) > self._max_fallback_run_ids: self._fallback_run_ids.popitem(last=False) return fallback_run_id def _release_fallback_run_id(self, runtime: Runtime) -> None: """Release a completed invocation's fallback anchor, if it used one.""" ctx = getattr(runtime, "context", None) if isinstance(ctx, dict) and "run_id" in ctx: return execution_info = getattr(runtime, "execution_info", None) if getattr(execution_info, "run_id", None) is not None: return control = getattr(runtime, "control", None) anchor = control if control is not None else runtime anchor_id = id(anchor) with self._lock: existing = self._fallback_run_ids.get(anchor_id) if existing is not None and existing[0] is anchor: self._fallback_run_ids.pop(anchor_id, None) def consume_stop_reason(self, run_id: str | None) -> str | None: """Pop and return the stop reason the hard-stop set for this run. Returns ``"loop_capped"`` when a repeated tool-call loop tripped the hard stop during the run — the run still completed with a forced final answer (the hard stop strips ``tool_calls`` rather than raising). The subagent executor calls this after the run returns so a loop-capped completion carries ``stop_reason=loop_capped`` to the lead instead of looking like a clean ``completed``. Mirrors ``TokenBudgetMiddleware.consume_stop_reason``; popping keeps the dict from accumulating on a reused instance. """ with self._lock: reason = self._stop_reason.pop(run_id, None) self._stop_reason_thread_id.pop(run_id, None) return reason def _run_scope_key(self, runtime: Runtime) -> _RunScopeKey: """Return the shared tracking key for the current thread/run.""" return self._get_thread_id(runtime), self._get_run_id(runtime) def _pending_key(self, runtime: Runtime) -> _RunScopeKey: """Return the pending-warning key for the current thread/run.""" return self._run_scope_key(runtime) def _evict_if_needed(self) -> None: """Evict least recently used thread/run scopes if over the limit. Must be called while holding self._lock. """ while len(self._history) > self.max_tracked_threads: evicted_key, _ = self._history.popitem(last=False) self._warned.pop(evicted_key, None) self._tool_name_history.pop(evicted_key, None) self._tool_name_counter.pop(evicted_key, None) self._tool_freq_warned.pop(evicted_key, None) self._drop_pending_warning_key_locked(evicted_key) logger.debug( "Evicted loop tracking for thread/run scope (LRU)", extra={"thread_id": evicted_key[0], "run_id": evicted_key[1]}, ) def _drop_pending_warning_key_locked(self, key: _RunScopeKey) -> None: """Drop all pending-warning bookkeeping for one thread/run key. Must be called while holding self._lock. """ self._pending_warnings.pop(key, None) self._pending_warning_touch_order.pop(key, None) def _touch_pending_warning_key_locked(self, key: _RunScopeKey) -> None: """Mark a pending-warning key as recently used. Must be called while holding self._lock. """ self._pending_warning_touch_order[key] = None self._pending_warning_touch_order.move_to_end(key) def _prune_pending_warning_state_locked(self, protected_key: _RunScopeKey) -> None: """Cap pending-warning state across abnormal or concurrent runs. Must be called while holding self._lock. """ overflow = len(self._pending_warning_touch_order) - self._max_pending_warning_keys if overflow <= 0: return candidates = [key for key in self._pending_warning_touch_order if key != protected_key] for key in candidates[:overflow]: self._drop_pending_warning_key_locked(key) def _queue_pending_warning(self, runtime: Runtime, warning: str) -> None: """Queue one transient warning for the current thread/run with caps.""" pending_key = self._pending_key(runtime) with self._lock: warnings = self._pending_warnings[pending_key] if warning not in warnings: warnings.append(warning) if len(warnings) > _MAX_PENDING_WARNINGS_PER_RUN: del warnings[: len(warnings) - _MAX_PENDING_WARNINGS_PER_RUN] self._touch_pending_warning_key_locked(pending_key) self._prune_pending_warning_state_locked(protected_key=pending_key) def _track_and_check( self, state: AgentState, runtime: Runtime, ) -> _LoopDecision | None: """Track tool calls and check for loops. Two detection layers: 1. **Hash-based** (existing): catches identical tool call sets. 2. **Frequency-based** (new): catches the same *tool type* being called many times with varying arguments (e.g. ``read_file`` on 40 different files). Returns: A structured decision when a warning or hard stop is triggered. """ messages = state.get("messages", []) if not messages: return None last_msg = messages[-1] if getattr(last_msg, "type", None) != "ai": return None tool_calls = getattr(last_msg, "tool_calls", None) if not tool_calls: return None scope_key = self._run_scope_key(runtime) thread_id, run_id = scope_key call_hash = _hash_tool_calls(tool_calls) with self._lock: # Touch / create entry (move to end for LRU) if scope_key in self._history: self._history.move_to_end(scope_key) else: self._history[scope_key] = [] self._evict_if_needed() history = self._history[scope_key] history.append(call_hash) if len(history) > self.window_size: history[:] = history[-self.window_size :] warned_hashes = self._warned.get(scope_key) if warned_hashes is not None: warned_hashes.intersection_update(history) if not warned_hashes: self._warned.pop(scope_key, None) count = history.count(call_hash) tool_names = [str(tc.get("name") or "?") for tc in tool_calls] # --- Layer 1: hash-based (identical call sets) --- if count >= self.hard_limit: logger.error( "Loop hard limit reached — forcing stop", extra={ "thread_id": thread_id, "run_id": run_id, "call_hash": call_hash, "count": count, "tools": tool_names, }, ) return _LoopDecision( message=_HARD_STOP_MSG, action="hard_stop", detection_layer="identical_call_set", tool_names=tuple(tool_names), count=count, threshold=self.hard_limit, ) # Warnings admit the whole batch, so they must not skip frequency # accounting or hide a later hard limit. Keep one candidate (hash # warnings retain priority over frequency warnings) until every # admitted call has been checked. Only the selected warning is # marked/logged; a hard stop may supersede it below. warning: _LoopDecision | None = None if count >= self.warn_threshold and call_hash not in self._warned.get(scope_key, set()): warning = _LoopDecision( message=_WARNING_MSG, action="warn", detection_layer="identical_call_set", tool_names=tuple(tool_names), count=count, threshold=self.warn_threshold, ) # --- Layer 2: per-tool-type frequency (windowed) --- tool_name_history = self._tool_name_history[scope_key] name_counter = self._tool_name_counter[scope_key] for tc in tool_calls: name = tc.get("name", "") if not name: continue # Windowed counting: append the name and trim to the frequency # window (>= the largest threshold) so the count can reach the # warn/hard limits on a tight burst yet still decay for # spread-out calls. A mirrored Counter gives O(1) freq_count # even when a per-tool override inflates the window globally. tool_name_history.append(name) name_counter[name] += 1 while len(tool_name_history) > self._tool_freq_window: old = tool_name_history.popleft() c = name_counter[old] - 1 if c <= 0: del name_counter[old] else: name_counter[old] = c old_warn = self._tool_freq_overrides.get(old, self._default_tool_freq_thresholds)[0] if c < old_warn: # Any tool can evict an older name from the shared # window. Rearm that name as soon as its burst decays, # even when the current call belongs to another tool. self._tool_freq_warned[scope_key].discard(old) freq_count = name_counter.get(name, 0) eff_warn, eff_hard = self._tool_freq_overrides.get(name, self._default_tool_freq_thresholds) if freq_count >= eff_hard: logger.error( "Tool frequency hard limit reached — forcing stop", extra={ "thread_id": thread_id, "run_id": run_id, "tool_name": name, "count": freq_count, }, ) return _LoopDecision( message=_TOOL_FREQ_HARD_STOP_MSG.format(tool_name=name, count=freq_count), action="hard_stop", detection_layer="tool_frequency", tool_names=(name,), count=freq_count, threshold=eff_hard, ) if freq_count >= eff_warn: freq_warned = self._tool_freq_warned[scope_key] if warning is None and name not in freq_warned: warning = _LoopDecision( message=_TOOL_FREQ_WARNING_MSG.format(tool_name=name, count=freq_count), action="warn", detection_layer="tool_frequency", tool_names=(name,), count=freq_count, threshold=eff_warn, ) else: # Windowed count decayed below the warn threshold; allow a # future burst of this tool to warn again. self._tool_freq_warned[scope_key].discard(name) if warning is not None: if warning.detection_layer == "identical_call_set": self._warned[scope_key].add(call_hash) logger.warning( "Repetitive tool calls detected — injecting warning", extra={ "thread_id": thread_id, "run_id": run_id, "call_hash": call_hash, "count": warning.count, "tools": list(warning.tool_names), }, ) else: warned_name = warning.tool_names[0] # Later calls in this batch may already have decayed this # burst. Do not suppress the next burst with a stale mark. if name_counter.get(warned_name, 0) >= warning.threshold: self._tool_freq_warned[scope_key].add(warned_name) logger.warning( "Tool frequency warning — too many calls to same tool type", extra={ "thread_id": thread_id, "run_id": run_id, "tool_name": warned_name, "count": warning.count, }, ) return warning @staticmethod def _append_text(content: str | list | None, text: str) -> str | list: """Append *text* to AIMessage content, handling str, list, and None. When content is a list of content blocks (e.g. Anthropic thinking mode), we append a new ``{"type": "text", ...}`` block instead of concatenating a string to a list, which would raise ``TypeError``. """ if content is None: return text if isinstance(content, list): return [*content, {"type": "text", "text": f"\n\n{text}"}] if isinstance(content, str): return content + f"\n\n{text}" # Fallback: coerce unexpected types to str to avoid TypeError return str(content) + f"\n\n{text}" def _record_audit_event( self, decision: _LoopDecision, runtime: Runtime, ) -> None: """Persist a loop-detection transition without sensitive tool data.""" recorder, is_subagent, agent_id = resolve_audit_recorder( getattr(runtime, "context", None), recorder_key=LOOP_DETECTION_RECORDER_CONTEXT_KEY, ) if recorder is None: return try: recorder.record_middleware( tag=MIDDLEWARE_LOOP_DETECTION_TAG, name=type(self).__name__, hook="after_model", action=decision.action, changes={ "is_subagent": is_subagent, "agent_id": agent_id, "detection_layer": decision.detection_layer, "tool_names": list(decision.tool_names), "count": decision.count, "threshold": decision.threshold, }, ) except Exception: # noqa: BLE001 # Audit persistence must never break the agent run. logger.warning( "Failed to record middleware:loop_detection event", exc_info=True, ) def _apply(self, state: AgentState, runtime: Runtime) -> dict | None: decision = self._track_and_check(state, runtime) if decision is None: return None # Keep the shared loop-detection lock's critical section bounded; the # journal append is cheap and non-blocking but does not belong under it. self._record_audit_event(decision, runtime) warning = decision.message if decision.hard_stop: # Record the stop reason so the executor can surface # ``stop_reason=loop_capped`` after the run returns (#3875 Phase 2). # The hard stop does not raise — it strips tool_calls and lets the # run finish with a forced final answer — so without this the caller # would see a clean ``completed``. See ``consume_stop_reason``. # Written under the lock to match ``TokenBudgetMiddleware``: the lead # agent's middleware instance is shared across concurrent Gateway # threads, so the bounded-dict write needs the same guard. thread_id, run_id = self._run_scope_key(runtime) with self._lock: self._stop_reason[run_id] = "loop_capped" self._stop_reason_thread_id[run_id] = thread_id # Also write to runtime.context so the lead worker can read it # without needing a reference to this middleware instance (#4176). ctx = getattr(runtime, "context", None) if isinstance(ctx, dict): ctx["stop_reason"] = "loop_capped" # Strip tool calls from every provider surface of the last # AIMessage (structured, raw, and content blocks) to force text # output. With no call left on any surface, the AIMessage no # longer requires matching ToolMessage responses, so replacing it # here is safe for strict provider pairing validators. messages = state.get("messages", []) last_msg = messages[-1] content = self._append_text(last_msg.content, warning or _HARD_STOP_MSG) stripped_msg = clone_ai_message_with_tool_calls(last_msg, [], content=content) return {"messages": [stripped_msg]} if decision.action == "warn": # Defer injection to the next model call. We must NOT alter the # AIMessage(tool_calls=...) here (would put framework words in # the model's mouth, polluting downstream consumers like # MemoryMiddleware), nor insert a separate non-tool message # (would break OpenAI/Moonshot tool-call pairing because the # tools node has not produced ToolMessage responses yet). The # warning is delivered via ``wrap_model_call`` below. self._queue_pending_warning(runtime, warning) return None return None def _clear_current_run_pending_warnings(self, runtime: Runtime) -> None: """Drop pending warnings owned by the current thread/run.""" pending_key = self._pending_key(runtime) with self._lock: self._drop_pending_warning_key_locked(pending_key) @staticmethod def _format_warning_message(warnings: list[str]) -> str: """Merge pending warnings into one prompt message.""" deduped = list(dict.fromkeys(warnings)) return "\n\n".join(deduped) @override def before_agent(self, state: AgentState, runtime: Runtime) -> dict | None: # Keep this hook in the compiled graph topology. Pending warnings are # already run-scoped; touching sibling runs here would break overlapping # invocations, while after_agent and the bounded queue own cleanup. return None @override async def abefore_agent(self, state: AgentState, runtime: Runtime) -> dict | None: # Async topology must mirror the sync hook above. return None @override def after_model(self, state: AgentState, runtime: Runtime) -> dict | None: return self._apply(state, runtime) @override async def aafter_model(self, state: AgentState, runtime: Runtime) -> dict | None: return self._apply(state, runtime) @override def after_agent(self, state: AgentState, runtime: Runtime) -> dict | None: self._clear_current_run_pending_warnings(runtime) self._release_fallback_run_id(runtime) return None @override async def aafter_agent(self, state: AgentState, runtime: Runtime) -> dict | None: self._clear_current_run_pending_warnings(runtime) self._release_fallback_run_id(runtime) return None def _drain_pending_warnings(self, runtime: Runtime) -> list[str]: """Pop and return all queued warnings for *runtime*'s thread/run.""" pending_key = self._pending_key(runtime) with self._lock: warnings = self._pending_warnings.pop(pending_key, []) self._pending_warning_touch_order.pop(pending_key, None) return warnings def _restore_pending_warnings(self, runtime: Runtime, warnings: list[str]) -> None: """Requeue warnings taken for a model call that raised. LLMErrorHandlingMiddleware sits outside this middleware and retries a failed call by running this wrap again, so the retry must still find the warning. It would not be queued again: it is already marked warned. """ if not warnings: return pending_key = self._pending_key(runtime) with self._lock: queued = self._pending_warnings[pending_key] queued[:0] = [warning for warning in warnings if warning not in queued] # Keep the restored warnings at the front; trim what came after them. del queued[_MAX_PENDING_WARNINGS_PER_RUN:] self._touch_pending_warning_key_locked(pending_key) self._prune_pending_warning_state_locked(protected_key=pending_key) def _inject_warnings(self, request: ModelRequest, warnings: list[str]) -> ModelRequest: """Append *warnings* to the outgoing message list. The warning is placed *after* every existing message, including the ToolMessage responses to the previous AIMessage(tool_calls). This keeps ``assistant tool_calls -> tool_messages`` pairing intact for OpenAI/Moonshot, avoids the Anthropic mid-stream SystemMessage restriction (we use HumanMessage), and never mutates an existing AIMessage. """ if not warnings: return request new_messages = [ *request.messages, HumanMessage(content=self._format_warning_message(warnings), name="loop_warning"), ] return request.override(messages=new_messages) @override def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelCallResult: warnings = self._drain_pending_warnings(request.runtime) try: return handler(self._inject_warnings(request, warnings)) except Exception: self._restore_pending_warnings(request.runtime, warnings) raise @override async def awrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelCallResult: warnings = self._drain_pending_warnings(request.runtime) try: return await handler(self._inject_warnings(request, warnings)) except Exception: self._restore_pending_warnings(request.runtime, warnings) raise def reset(self, thread_id: str | None = None) -> None: """Clear tracking state. If thread_id given, clear only that thread.""" with self._lock: if thread_id: for mapping in ( self._history, self._warned, self._tool_name_history, self._tool_name_counter, self._tool_freq_warned, ): for key in list(mapping): if key[0] == thread_id: mapping.pop(key, None) pending_keys = set(self._pending_warnings) | set(self._pending_warning_touch_order) for key in pending_keys: if key[0] == thread_id: self._drop_pending_warning_key_locked(key) for run_id, owner_thread_id in list(self._stop_reason_thread_id.items()): if owner_thread_id == thread_id: self._stop_reason.pop(run_id, None) self._stop_reason_thread_id.pop(run_id, None) else: self._history.clear() self._warned.clear() self._tool_name_history.clear() self._tool_name_counter.clear() self._tool_freq_warned.clear() self._pending_warnings.clear() self._pending_warning_touch_order.clear() self._stop_reason.clear() self._stop_reason_thread_id.clear() self._fallback_run_ids.clear()