From 36ce7590b78727e7f428fcef484044140e253a30 Mon Sep 17 00:00:00 2001 From: PeaceMaker-best Date: Fri, 11 Sep 2026 13:45:52 +0800 Subject: [PATCH] fix(agents): isolate loop detection state by run (#5344) * fix(agents): scope loop detection state per run * fix(agents): harden loop scope fallback * docs: move loop lifecycle detail out of inherited guidance --------- Co-authored-by: PeaceMaker-best <221849497+PeaceMaker-best@users.noreply.github.com> Co-authored-by: Willem Jiang --- backend/docs/LOOP_DETECTION.md | 38 ++ .../deerflow/agents/middlewares/AGENTS.md | 16 +- .../middlewares/loop_detection_middleware.py | 243 +++++++---- .../middlewares/tool_progress_middleware.py | 18 +- .../deerflow/config/loop_detection_config.py | 2 +- .../tests/test_loop_detection_middleware.py | 406 ++++++++++++++++-- config.example.yaml | 1 + 7 files changed, 581 insertions(+), 143 deletions(-) create mode 100644 backend/docs/LOOP_DETECTION.md diff --git a/backend/docs/LOOP_DETECTION.md b/backend/docs/LOOP_DETECTION.md new file mode 100644 index 000000000..ca9035c1c --- /dev/null +++ b/backend/docs/LOOP_DETECTION.md @@ -0,0 +1,38 @@ +# Loop detection lifecycle + +`LoopDetectionMiddleware` owns call-pattern detection. Its place in the chain +and audit contract are documented in the +[middleware guide](../packages/harness/deerflow/agents/middlewares/AGENTS.md). + +## Run-scoped state + +Hash histories, frequency windows/counters, and warning-suppression sets +share a `(thread_id, run_id)` key. This gives a fresh budget to each user +run when a compiled graph is cached and reused, while keeping one budget +across repeated graph entries belonging to the same Gateway run (including +hidden goal continuations). `after_agent` clears only transient pending +warnings for its own scope, not those histories; the sync/async +`before_agent` hooks remain topology-preserving no-ops and must not delete a +sibling run's pending warning. Direct LangGraph embedders may omit +`context.run_id`; that fallback is anchored to the invocation's shared +`Runtime.control` object and mapped to an opaque generated ID, because +LangGraph replaces `Runtime` per node and CPython can reuse freed object +addresses. `after_agent` releases the anchor mapping, while the bounded map +covers abnormal exits. The compatibility-named +`max_tracked_threads` limit bounds run scopes, and `reset(thread_id)` clears +every retained run scope for that thread. + +## Decision ordering + +Loop decisions are severity-first across both detection layers: a warning +candidate never short-circuits frequency accounting for the remaining calls +in an admitted batch. A hard limit can stop scanning immediately because it +rejects the entire batch. Only the selected warning is marked and logged; +hash warnings still precede frequency warnings when neither layer stops the +run. Among simultaneous frequency-warning candidates, the first crossing in +model tool-call order remains selected for compatibility; later calls are +still counted and can warn in a later batch. A frequency warning whose burst +decays within the batch must not leave a stale suppression mark. + +`backend/tests/test_loop_detection_middleware.py` covers mixed-tool batches, +window decay, overrides, and sync/async compiled-graph execution. diff --git a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md index e25588cfb..767ee69d5 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md +++ b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md @@ -97,17 +97,11 @@ Before changing a later authorization phase, read the [authorization RFC](../../ 27. **SystemMessageCoalescingMiddleware** - Merges every SystemMessage into a single leading SystemMessage per request; provider-agnostic fix for strict backends (vLLM/SGLang/Qwen/Anthropic) that reject non-leading system messages. Touches the per-request payload only (checkpoint state unchanged); on midnight crossings only the latest `dynamic_context_reminder` SystemMessage survives. The subagent builder places its date-only context middleware immediately before this coalescer, so the built-in subagent prompt and hidden date reminder still reach providers as one leading system block 28. **SubagentLimitMiddleware** - *(optional, if `subagent_enabled`)* Truncates excess ordinary `task` tool calls to enforce both the per-response concurrency limit (`max_concurrent_subagents`, resolved against startup `subagent_runtime.max_running` and the 1-64 safety range before construction) and the per-run total delegation cap (`max_total_subagents` runtime override or `subagents.max_total_per_run`, default 6, clamped to 1-50). The total cap counts current-run entries in the durable delegation ledger (entries are tagged with `run_id` when captured), so repeated planning checkpoints in one run cannot keep launching legal-sized batches indefinitely, while later user turns in the same thread get a fresh run budget. Explicit durable `batch_task` calls are a separate mode with persisted total/live/running limits and are not rewritten into ordinary ledger entries. If the ordinary cap is exhausted, the middleware strips remaining `task` calls, forces `finish_reason="stop"`, and appends a visible limit note so the run can synthesize existing results instead of ending with an empty tool-call response. 29. **LoopDetectionMiddleware** - *(optional, if `loop_detection.enabled`)* Detects repeated tool-call loops; hard-stop clears both structured `tool_calls` and raw provider tool-call metadata before forcing a final text answer; stamps `loop_capped` via `consume_stop_reason` (#3875 Phase 2), symmetric to `TokenBudgetMiddleware`; persists warned-state transitions (first per call hash or per tool-frequency burst) and hard stops as `middleware:loop_detection`, attributed with `is_subagent` and the optional `agent_id`, without tool arguments, message content, tool results, or argument-derived hashes. Ordinary task subagents get dedicated recorder keys through a parent-loop proxy; never pass `RunJournal` into their isolated loop. Durable batch subagents have no parent run journal and do not persist these transitions - Loop decisions are severity-first across both detection layers: a warning - candidate never short-circuits frequency accounting for the remaining calls - in an admitted batch. A hard limit can stop scanning immediately because it - rejects the entire batch. Only the selected warning is marked and logged; - hash warnings still precede frequency warnings when neither layer stops the - run. Among simultaneous frequency-warning candidates, the first crossing in - model tool-call order remains selected for compatibility; later calls are - still counted and can warn in a later batch. A frequency warning whose burst - decays within the batch must not leave a stale suppression mark. - `tests/test_loop_detection_middleware.py` covers mixed-tool batches, window - decay, overrides, and sync/async compiled-graph execution. + State is run-scoped: new user runs get fresh budgets; same-run goal + continuations share history. Keep sibling warnings isolated and lifecycle + hooks topology-stable. Before changing this guard, read + [Loop detection lifecycle](../../../../../docs/LOOP_DETECTION.md) for + fallback identity, cleanup/LRU/reset, severity ordering, and test invariants. 30. **TokenBudgetMiddleware** - *(optional, if `token_budget.enabled`)* Enforces per-run token limits 31. **Custom middlewares** - *(optional)* Any `custom_middlewares` passed to `build_middlewares` are injected here, before config-declared extensions and the terminal-response/safety/clarification tail 32. **Configured extension middlewares** - *(optional, if `extensions.middlewares` is set in `config.yaml` or `extensions_config.json`)* Zero-argument `AgentMiddleware` classes loaded from `module.path:ClassName` entries via `deerflow.reflection.resolve_class`. Missing packages, invalid classes, and broken modules fail loudly at agent creation. These run after built-ins/programmatic custom middleware and after the lead/subagent loop/token guards, but before the terminal-response/safety/clarification tail; subagents receive the same configured extension middleware class list before their safety tail. Treat these files as trusted operator config because middleware paths instantiate arbitrary code. Gateway skill/MCP toggle endpoints preserve this field through `to_file_dict()` but must not add a write path for `extensions.middlewares` without an explicit trust-boundary review. Lead-only vs subagent-only middleware lists and per-context constructor parameters are not expressible in this MVP. diff --git a/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py index 061af7f2b..a0af866ac 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/loop_detection_middleware.py @@ -37,6 +37,16 @@ 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 @@ -55,6 +65,7 @@ import hashlib import json import logging import threading +import uuid from collections import Counter, OrderedDict, defaultdict, deque from collections.abc import Awaitable, Callable from copy import deepcopy @@ -80,11 +91,13 @@ logger = logging.getLogger(__name__) _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 eviction limit +_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. @@ -216,8 +229,9 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): tool_calls entirely. Default: 5. window_size: Size of the sliding window for tracking calls. Default: 20. - max_tracked_threads: Maximum number of threads to track before - evicting the least recently used. Default: 100. + 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 @@ -272,28 +286,35 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): *(hard for _, hard in self._tool_freq_overrides.values()), ) self._lock = threading.Lock() - self._history: OrderedDict[str, list[str]] = OrderedDict() - self._warned: dict[str, set[str]] = defaultdict(set) - # Windowed per-tool-type frequency: recent tool names per thread, + # 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[str, deque[str]] = defaultdict(deque) - # Per-thread Counter mirroring the deque so freq_count is O(1) instead + 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[str, Counter[str]] = defaultdict(Counter) - # Per-thread set of tool names already warned about in Layer 2, so a + 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[str, set[str]] = defaultdict(set) + 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[tuple[str, str], list[str]] = defaultdict(list) - self._pending_warning_touch_order: OrderedDict[tuple[str, str], None] = OrderedDict() + 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 @@ -301,8 +322,11 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): # 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. - self._stop_reason: BoundedDict[str, str] = BoundedDict(1000) + # 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 { @@ -335,31 +359,65 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): return str(thread_id) return "default" - def _get_run_id(self, runtime: Runtime) -> str: + def _get_run_id(self, runtime: Runtime) -> str | None: """Extract run_id from runtime context for per-run warning scoping. - Keyed by presence, not truthiness: ``SubagentExecutor`` sets - ``context["run_id"] = self.run_id`` unconditionally (no truthiness - guard), so an embedded/TUI-dispatched subagent — whose ``run_id`` is - never assigned per ``AGENTS.md``'s description of the embedded - ``DeerFlowClient`` — runs with a context that legitimately carries - ``run_id=None`` (the key is *present*, not absent). The executor - later reads the stop reason back with the raw attribute, - ``consume_stop_reason(self.run_id)``, so this must return exactly - that value (``None`` included) when the key is present, rather than - collapsing it to a shared fallback indistinguishable from an absent - key. A truthiness check (``if run_id:``) previously conflated - "present but None/falsy" with "absent", both mapping to the same - literal ``"default"`` — so a genuine ``run_id=None`` hard-stop was - recorded under ``"default"`` here but looked up under ``None`` by - the executor, silently losing the ``loop_capped`` stop reason. - Mirrors ``TokenBudgetMiddleware._get_run_id``. + 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"] - # Fallback to runtime object ID to prevent collisions across embedded client runs - return str(id(runtime)) + + 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. @@ -373,29 +431,36 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): popping keeps the dict from accumulating on a reused instance. """ with self._lock: - return self._stop_reason.pop(run_id, None) + reason = self._stop_reason.pop(run_id, None) + self._stop_reason_thread_id.pop(run_id, None) + return reason - def _pending_key(self, runtime: Runtime) -> tuple[str, str]: - """Return the pending-warning key for the current thread/run.""" + 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 threads if over the limit. + """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_id, _ = self._history.popitem(last=False) - self._warned.pop(evicted_id, None) - self._tool_name_history.pop(evicted_id, None) - self._tool_name_counter.pop(evicted_id, None) - self._tool_freq_warned.pop(evicted_id, None) - for key in list(self._pending_warnings): - if key[0] == evicted_id: - self._drop_pending_warning_key_locked(key) - logger.debug("Evicted loop tracking for thread %s (LRU)", evicted_id) + 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: tuple[str, str]) -> None: + 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. @@ -403,7 +468,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): self._pending_warnings.pop(key, None) self._pending_warning_touch_order.pop(key, None) - def _touch_pending_warning_key_locked(self, key: tuple[str, str]) -> 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. @@ -411,7 +476,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): 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: tuple[str, str]) -> None: + 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. @@ -464,27 +529,28 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): if not tool_calls: return None - thread_id = self._get_thread_id(runtime) + 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 thread_id in self._history: - self._history.move_to_end(thread_id) + if scope_key in self._history: + self._history.move_to_end(scope_key) else: - self._history[thread_id] = [] + self._history[scope_key] = [] self._evict_if_needed() - history = self._history[thread_id] + 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(thread_id) + 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(thread_id, None) + self._warned.pop(scope_key, None) count = history.count(call_hash) tool_names = [str(tc.get("name") or "?") for tc in tool_calls] @@ -495,6 +561,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): "Loop hard limit reached — forcing stop", extra={ "thread_id": thread_id, + "run_id": run_id, "call_hash": call_hash, "count": count, "tools": tool_names, @@ -515,7 +582,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): # 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(thread_id, set()): + if count >= self.warn_threshold and call_hash not in self._warned.get(scope_key, set()): warning = _LoopDecision( message=_WARNING_MSG, action="warn", @@ -526,8 +593,8 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): ) # --- Layer 2: per-tool-type frequency (windowed) --- - tool_name_history = self._tool_name_history[thread_id] - name_counter = self._tool_name_counter[thread_id] + 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: @@ -551,7 +618,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): # 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[thread_id].discard(old) + 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) @@ -561,6 +628,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): "Tool frequency hard limit reached — forcing stop", extra={ "thread_id": thread_id, + "run_id": run_id, "tool_name": name, "count": freq_count, }, @@ -575,7 +643,7 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): ) if freq_count >= eff_warn: - freq_warned = self._tool_freq_warned[thread_id] + 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), @@ -588,15 +656,16 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): else: # Windowed count decayed below the warn threshold; allow a # future burst of this tool to warn again. - self._tool_freq_warned[thread_id].discard(name) + self._tool_freq_warned[scope_key].discard(name) if warning is not None: if warning.detection_layer == "identical_call_set": - self._warned[thread_id].add(call_hash) + 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), @@ -607,11 +676,12 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): # 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[thread_id].add(warned_name) + 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, }, @@ -712,9 +782,10 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): # 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. - run_id = self._get_run_id(runtime) + 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) @@ -743,14 +814,6 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): return None - def _clear_other_run_pending_warnings(self, runtime: Runtime) -> None: - """Drop stale pending warnings for previous runs in this thread.""" - thread_id, current_run_id = self._pending_key(runtime) - with self._lock: - for key in list(self._pending_warnings): - if key[0] == thread_id and key[1] != current_run_id: - self._drop_pending_warning_key_locked(key) - 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) @@ -765,12 +828,14 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): @override def before_agent(self, state: AgentState, runtime: Runtime) -> dict | None: - self._clear_other_run_pending_warnings(runtime) + # 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: - self._clear_other_run_pending_warnings(runtime) + # Async topology must mirror the sync hook above. return None @override @@ -784,11 +849,13 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): @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]: @@ -838,14 +905,24 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): """Clear tracking state. If thread_id given, clear only that thread.""" with self._lock: if thread_id: - self._history.pop(thread_id, None) - self._warned.pop(thread_id, None) - self._tool_name_history.pop(thread_id, None) - self._tool_name_counter.pop(thread_id, None) - self._tool_freq_warned.pop(thread_id, None) - for key in list(self._pending_warnings): + 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() @@ -855,3 +932,5 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]): 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() diff --git a/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py index 03fb41fc1..92c921a9d 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/tool_progress_middleware.py @@ -473,16 +473,14 @@ class ToolProgressMiddleware(AgentMiddleware[AgentState]): the previous run are also cleared so a single first-call problem in the new run cannot falsely trip WARNED against stale context from a run the model no longer sees. - **Cross-run scoping vs LoopDetectionMiddleware**: this per-run reset is an intentional - policy choice, not an oversight. Errors like ``rate_limited`` and ``transient`` are - time-bound: their root cause may resolve between user turns, so carrying a stale - counter forward risks a false-positive BLOCKED on calls that would now succeed. - LoopDetectionMiddleware takes the opposite stance — it retains ``_history`` across - runs (only clearing other-run *pending* warnings at ``before_agent``), because - call-pattern loops are time-invariant: a model that keeps issuing the same tool_calls - regardless of results does so regardless of when the run started. The two middlewares - therefore guard different failure modes (result quality vs. call pattern) and their - cross-run scoping policies intentionally differ as a consequence. + **Graph-entry scoping vs LoopDetectionMiddleware**: this reset at every + ``before_agent`` is an intentional policy choice, not an oversight. Errors like + ``rate_limited`` and ``transient`` are time-bound, so carrying a stale counter into a + later graph entry risks a false-positive BLOCKED on calls that would now succeed. + LoopDetectionMiddleware instead keys call-pattern state by ``(thread_id, run_id)``: + separate user runs are isolated even on a cached graph, while repeated graph entries + in one Gateway run (including hidden goal continuations) share a loop budget. The two + middlewares therefore guard different failure modes and use different lifetimes. """ thread_id = self._thread_id(runtime) with self._lock: diff --git a/backend/packages/harness/deerflow/config/loop_detection_config.py b/backend/packages/harness/deerflow/config/loop_detection_config.py index 3eb9a2f37..684e9649b 100644 --- a/backend/packages/harness/deerflow/config/loop_detection_config.py +++ b/backend/packages/harness/deerflow/config/loop_detection_config.py @@ -46,7 +46,7 @@ class LoopDetectionConfig(BaseModel): max_tracked_threads: int = Field( default=100, ge=1, - description="Maximum number of thread histories to keep in memory", + description="Maximum number of thread/run histories to keep in memory (legacy field name)", ) tool_freq_warn: int = Field( default=30, diff --git a/backend/tests/test_loop_detection_middleware.py b/backend/tests/test_loop_detection_middleware.py index de893fb7a..2d5710528 100644 --- a/backend/tests/test_loop_detection_middleware.py +++ b/backend/tests/test_loop_detection_middleware.py @@ -14,6 +14,7 @@ from langchain_core.runnables import Runnable from langchain_core.tools import tool as as_tool from pydantic import PrivateAttr +from deerflow.agents.middlewares import loop_detection_middleware as loop_detection_module from deerflow.agents.middlewares.loop_detection_middleware import ( _HARD_STOP_MSG, _MAX_PENDING_WARNINGS_PER_RUN, @@ -33,6 +34,10 @@ def _pending_key(thread_id="test-thread", run_id="test-run"): return (thread_id, run_id) +def _scope_key(thread_id="test-thread", run_id="test-run"): + return (thread_id, run_id) + + def _make_request(messages, runtime): """Build a minimal ModelRequest stand-in for wrap_model_call tests.""" request = MagicMock() @@ -307,21 +312,20 @@ class TestLoopDetection: def test_missing_run_id_uses_per_runtime_pending_scope(self): """When runtime.context has no ``run_id`` key at all, warning handling - falls back to a key scoped to the runtime object's identity — - mirroring ``TokenBudgetMiddleware._get_run_id``'s fallback — instead - of a shared literal like the old ``"default"``, which would collide - across concurrent runs that both lack a run_id (the ``_stop_reason`` - dict this same key derivation feeds is keyed by run_id alone, with - no thread scoping).""" + falls back to a key scoped to the LangGraph invocation instead of a + shared literal like the old ``"default"``.""" mw = LoopDetectionMiddleware(warn_threshold=3, hard_limit=10) - runtime = MagicMock() - runtime.context = {"thread_id": "test-thread"} + runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=object(), + execution_info=None, + ) call = [_bash_call("ls")] for _ in range(3): mw._apply(_make_state(tool_calls=call), runtime) - fallback_run_id = str(id(runtime)) + fallback_run_id = mw._get_run_id(runtime) assert mw._pending_warnings.get(_pending_key(run_id=fallback_run_id)) request = _make_request([AIMessage(content="hi")], runtime) @@ -333,8 +337,59 @@ class TestLoopDetection: assert "LOOP DETECTED" in loop_warnings[0].content assert not mw._pending_warnings.get(_pending_key(run_id=fallback_run_id)) - def test_before_agent_clears_stale_pending_warnings_for_thread(self): - """Starting a new run drops stale warnings from prior runs in the same thread.""" + def test_missing_run_id_shares_scope_across_runtime_wrappers(self): + """LangGraph replaces ``Runtime`` per node but preserves ``control``.""" + mw = LoopDetectionMiddleware() + invocation_control = object() + first_runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=invocation_control, + execution_info=None, + ) + later_runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=invocation_control, + execution_info=None, + ) + + assert mw._get_run_id(first_runtime) == mw._get_run_id(later_runtime) + + def test_missing_run_id_fallback_survives_reused_object_address(self, monkeypatch): + """A later invocation must not inherit a freed anchor's fallback key.""" + mw = LoopDetectionMiddleware() + monkeypatch.setattr(loop_detection_module, "id", lambda _value: 42, raising=False) + first_runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=object(), + execution_info=None, + ) + first_run_id = mw._get_run_id(first_runtime) + mw.after_agent({"messages": []}, first_runtime) + assert not mw._fallback_run_ids + + later_runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=object(), + execution_info=None, + ) + + assert mw._get_run_id(later_runtime) != first_run_id + + def test_missing_run_id_fallback_map_is_bounded_on_abnormal_exits(self): + mw = LoopDetectionMiddleware(max_tracked_threads=2) + + for _ in range(10): + runtime = SimpleNamespace( + context={"thread_id": "test-thread"}, + control=object(), + execution_info=None, + ) + mw._get_run_id(runtime) + + assert len(mw._fallback_run_ids) == mw._max_fallback_run_ids == 4 + + def test_before_agent_preserves_pending_warning_for_sibling_run(self): + """An overlapping run must not erase a warning owned by another run.""" mw = LoopDetectionMiddleware(warn_threshold=3, hard_limit=10) runtime_a = _make_runtime(run_id="run-A") runtime_b = _make_runtime(run_id="run-B") @@ -345,6 +400,13 @@ class TestLoopDetection: assert mw._pending_warnings.get(_pending_key(run_id="run-A")) mw.before_agent({"messages": []}, runtime_b) + assert mw._pending_warnings.get(_pending_key(run_id="run-A")) + + request = _make_request([AIMessage(content="hi")], runtime_a) + captured, handler = _capture_handler() + mw.wrap_model_call(request, handler) + loop_warnings = [message for message in captured[0].messages if isinstance(message, HumanMessage) and message.name == "loop_warning"] + assert len(loop_warnings) == 1 assert not mw._pending_warnings.get(_pending_key(run_id="run-A")) def test_after_agent_clears_current_run_pending_warnings(self): @@ -580,9 +642,9 @@ class TestLoopDetection: runtime_new = _make_runtime("thread-new") mw._apply(_make_state(tool_calls=call), runtime_new) - assert "thread-0" not in mw._history - assert "thread-0" not in mw._tool_name_history - assert "thread-new" in mw._history + assert _scope_key("thread-0") not in mw._history + assert _scope_key("thread-0") not in mw._tool_name_history + assert _scope_key("thread-new") in mw._history assert len(mw._history) == 3 def test_warned_hashes_are_pruned_to_sliding_window(self): @@ -595,9 +657,9 @@ class TestLoopDetection: mw._apply(_make_state(tool_calls=call), runtime) mw._apply(_make_state(tool_calls=call), runtime) - assert len(mw._history["test-thread"]) <= 4 - assert set(mw._warned["test-thread"]).issubset(set(mw._history["test-thread"])) - assert len(mw._warned["test-thread"]) <= 4 + assert len(mw._history[_scope_key()]) <= 4 + assert set(mw._warned[_scope_key()]).issubset(set(mw._history[_scope_key()])) + assert len(mw._warned[_scope_key()]) <= 4 def test_pending_warning_keys_are_capped(self): """Abnormal same-thread runs cannot grow pending-warning keys forever.""" @@ -644,12 +706,190 @@ class TestLoopDetection: def test_fallback_thread_id_when_missing(self): """When runtime context has no thread_id, should use 'default'.""" mw = LoopDetectionMiddleware(warn_threshold=2) - runtime = MagicMock() - runtime.context = {} + runtime = SimpleNamespace(context={}, control=object(), execution_info=None) call = [_bash_call("ls")] mw._apply(_make_state(tool_calls=call), runtime) - assert "default" in mw._history + assert ("default", mw._get_run_id(runtime)) in mw._history + + +class TestRunScopedTracking: + def test_identical_call_history_isolated_between_runs_on_same_thread(self): + """Separate user runs on a cached graph must not share repeat counts.""" + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=3, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + first_run = _make_runtime("shared-thread", "run-1") + second_run = _make_runtime("shared-thread", "run-2") + state = _make_state(tool_calls=[_bash_call("pwd")]) + + assert mw._apply(state, first_run) is None + mw.after_agent({"messages": []}, first_run) + assert mw._apply(state, second_run) is None + + assert not mw._pending_warnings.get(_pending_key("shared-thread", "run-2")) + assert mw._history[_scope_key("shared-thread", "run-1")] == mw._history[_scope_key("shared-thread", "run-2")] + assert len(mw._history[_scope_key("shared-thread", "run-1")]) == 1 + + def test_same_run_accumulates_across_agent_hook_cycles_for_goal_continuation(self): + """One Gateway run may re-enter the graph for hidden goal continuations.""" + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=10, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + runtime = _make_runtime("goal-thread", "goal-run") + state = _make_state(tool_calls=[_bash_call("pwd")]) + + assert mw._apply(state, runtime) is None + mw.after_agent({"messages": []}, runtime) + assert mw._apply(state, runtime) is None + + assert mw._pending_warnings.get(_pending_key("goal-thread", "goal-run")) + assert len(mw._history[_scope_key("goal-thread", "goal-run")]) == 2 + + def test_tool_frequency_history_isolated_between_runs_on_same_thread(self): + mw = LoopDetectionMiddleware( + warn_threshold=100, + hard_limit=200, + tool_freq_warn=2, + tool_freq_hard_limit=3, + ) + first_run = _make_runtime("shared-thread", "run-1") + second_run = _make_runtime("shared-thread", "run-2") + + assert mw._apply(_make_state(tool_calls=[_bash_call("first")]), first_run) is None + mw.after_agent({"messages": []}, first_run) + assert mw._apply(_make_state(tool_calls=[_bash_call("second")]), second_run) is None + + assert not mw._pending_warnings.get(_pending_key("shared-thread", "run-2")) + assert mw._tool_name_counter[_scope_key("shared-thread", "run-1")]["bash"] == 1 + assert mw._tool_name_counter[_scope_key("shared-thread", "run-2")]["bash"] == 1 + + def test_warned_hash_is_run_scoped(self): + """A warning in one run must not suppress the same warning in a later run.""" + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=100, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + state = _make_state(tool_calls=[_bash_call("pwd")]) + + for run_id in ("run-1", "run-2"): + runtime = _make_runtime("shared-thread", run_id) + assert mw._apply(state, runtime) is None + assert mw._apply(state, runtime) is None + assert mw._pending_warnings.get(_pending_key("shared-thread", run_id)) + mw.after_agent({"messages": []}, runtime) + + assert len(mw._warned[_scope_key("shared-thread", "run-1")]) == 1 + assert len(mw._warned[_scope_key("shared-thread", "run-2")]) == 1 + + def test_tool_frequency_warning_suppression_is_run_scoped(self): + mw = LoopDetectionMiddleware( + warn_threshold=100, + hard_limit=200, + tool_freq_warn=2, + tool_freq_hard_limit=100, + ) + + for run_id in ("run-1", "run-2"): + runtime = _make_runtime("shared-thread", run_id) + for suffix in ("first", "second"): + assert mw._apply(_make_state(tool_calls=[_bash_call(f"{run_id}-{suffix}")]), runtime) is None + assert mw._pending_warnings.get(_pending_key("shared-thread", run_id)) + mw.after_agent({"messages": []}, runtime) + + assert mw._tool_freq_warned[_scope_key("shared-thread", "run-1")] == {"bash"} + assert mw._tool_freq_warned[_scope_key("shared-thread", "run-2")] == {"bash"} + + def test_reset_thread_clears_every_run_scope_and_preserves_other_threads(self): + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=100, + tool_freq_warn=3, + tool_freq_hard_limit=100, + ) + target_runs = [ + _make_runtime("thread-A", "run-A1"), + _make_runtime("thread-A", "run-A2"), + ] + other_run = _make_runtime("thread-B", "run-B1") + state = _make_state(tool_calls=[_bash_call("pwd")]) + for runtime in [*target_runs, other_run]: + mw._apply(state, runtime) + mw._apply(state, runtime) + + mw.reset(thread_id="thread-A") + + expected_keys = {_scope_key("thread-B", "run-B1")} + for mapping in ( + mw._history, + mw._warned, + mw._tool_name_history, + mw._tool_name_counter, + mw._tool_freq_warned, + mw._pending_warnings, + mw._pending_warning_touch_order, + ): + assert set(mapping) == expected_keys + + def test_reset_thread_clears_its_unconsumed_stop_reasons_only(self): + mw = LoopDetectionMiddleware( + warn_threshold=1, + hard_limit=2, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + runtimes = [ + _make_runtime("thread-A", "run-A1"), + _make_runtime("thread-A", "run-A2"), + _make_runtime("thread-B", "run-B1"), + ] + state = _make_state(tool_calls=[_bash_call("pwd")]) + for runtime in runtimes: + assert mw._apply(state, runtime) is None + result = mw._apply(state, runtime) + assert result is not None + + mw.reset(thread_id="thread-A") + + assert mw.consume_stop_reason("run-A1") is None + assert mw.consume_stop_reason("run-A2") is None + assert mw.consume_stop_reason("run-B1") == "loop_capped" + + def test_lru_evicts_one_run_scope_without_dropping_sibling_run_warning(self): + mw = LoopDetectionMiddleware( + warn_threshold=100, + hard_limit=200, + max_tracked_threads=2, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + run_1 = _make_runtime("shared-thread", "run-1") + run_2 = _make_runtime("shared-thread", "run-2") + run_3 = _make_runtime("shared-thread", "run-3") + + mw._apply(_make_state(tool_calls=[_bash_call("one")]), run_1) + mw._queue_pending_warning(run_1, "run-1 warning") + mw._apply(_make_state(tool_calls=[_bash_call("two")]), run_2) + mw._queue_pending_warning(run_2, "run-2 warning") + mw._apply(_make_state(tool_calls=[_bash_call("one-again")]), run_1) + mw._apply(_make_state(tool_calls=[_bash_call("three")]), run_3) + + assert list(mw._history) == [ + _scope_key("shared-thread", "run-1"), + _scope_key("shared-thread", "run-3"), + ] + assert _scope_key("shared-thread", "run-2") not in mw._tool_name_history + assert _scope_key("shared-thread", "run-2") not in mw._tool_name_counter + assert _pending_key("shared-thread", "run-1") in mw._pending_warnings + assert _pending_key("shared-thread", "run-2") not in mw._pending_warnings class TestLoopDetectionRunEvents: @@ -864,6 +1104,94 @@ class TestLoopDetectionRunEvents: class TestLoopDetectionAgentGraphIntegration: + def test_reused_agent_graph_without_run_id_gets_one_budget_per_invocation(self): + """Library embedders may omit run_id while reusing one compiled graph.""" + + @as_tool + def bash(command: str) -> str: + """Run a fake shell command.""" + return f"ran: {command}" + + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=3, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + model = _CapturingFakeMessagesListChatModel( + responses=[ + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-1-a", "args": {"command": "pwd"}}]), + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-1-b", "args": {"command": "pwd"}}]), + AIMessage(content="first final answer"), + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-2-a", "args": {"command": "pwd"}}]), + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-2-b", "args": {"command": "pwd"}}]), + AIMessage(content="second final answer"), + ], + ) + graph = create_agent(model=model, tools=[bash], middleware=[mw]) + shared_context = {"thread_id": "library-thread"} + + first = graph.invoke( + {"messages": [("user", "where am I?")]}, + context=shared_context, + config={"recursion_limit": 15}, + ) + second = graph.invoke( + {"messages": [("user", "where am I now?")]}, + context=shared_context, + config={"recursion_limit": 15}, + ) + + assert first["messages"][-1].content == "first final answer" + assert second["messages"][-1].content == "second final answer" + loop_warnings_by_call = [[message for message in messages if isinstance(message, HumanMessage) and message.name == "loop_warning"] for messages in model.seen_messages] + assert [len(warnings) for warnings in loop_warnings_by_call] == [0, 0, 1, 0, 0, 1] + assert not mw._fallback_run_ids + + def test_reused_agent_graph_isolates_loop_history_between_runs(self): + """A cached graph must give each new run a fresh loop-detection budget.""" + + @as_tool + def bash(command: str) -> str: + """Run a fake shell command.""" + return f"ran: {command}" + + mw = LoopDetectionMiddleware( + warn_threshold=2, + hard_limit=3, + tool_freq_warn=100, + tool_freq_hard_limit=200, + ) + model = _CapturingFakeMessagesListChatModel( + responses=[ + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-1-call", "args": {"command": "pwd"}}]), + AIMessage(content="first final answer"), + AIMessage(content="", tool_calls=[{"name": "bash", "id": "run-2-call", "args": {"command": "pwd"}}]), + AIMessage(content="second final answer"), + ], + ) + graph = create_agent(model=model, tools=[bash], middleware=[mw]) + + first = graph.invoke( + {"messages": [("user", "where am I?")]}, + context={"thread_id": "cached-thread", "run_id": "run-1"}, + config={"recursion_limit": 10}, + ) + second = graph.invoke( + {"messages": [("user", "where am I now?")]}, + context={"thread_id": "cached-thread", "run_id": "run-2"}, + config={"recursion_limit": 10}, + ) + + assert first["messages"][-1].content == "first final answer" + assert second["messages"][-1].content == "second final answer" + assert len(model.seen_messages) == 4 + assert not any(isinstance(message, HumanMessage) and message.name == "loop_warning" for request_messages in model.seen_messages for message in request_messages) + assert set(mw._history) == { + _scope_key("cached-thread", "run-1"), + _scope_key("cached-thread", "run-2"), + } + def test_loop_warning_is_transient_in_real_agent_graph(self): """after_model queues the warning; wrap_model_call injects it request-only.""" @@ -1263,7 +1591,7 @@ class TestToolFrequencyDetection: # Reset only thread-A mw.reset(thread_id="thread-A") - assert "thread-A" not in mw._tool_name_history + assert _scope_key("thread-A") not in mw._tool_name_history # thread-B state should still be intact — 3rd call queues a warn. result = mw._apply(_make_state(tool_calls=[self._read_call("/b_2.py")]), runtime_b) @@ -1310,8 +1638,8 @@ class TestToolFrequencyDetection: # Two other threads push it out of the LRU window (max_tracked_threads=2). mw._apply(_make_state(tool_calls=[_bash_call("ls")]), _make_runtime("thread-a")) mw._apply(_make_state(tool_calls=[_bash_call("ls")]), _make_runtime("thread-b")) - assert "thread-evicted" not in mw._tool_name_history - assert "thread-evicted" not in mw._tool_name_counter + assert _scope_key("thread-evicted") not in mw._tool_name_history + assert _scope_key("thread-evicted") not in mw._tool_name_counter # Thread id reused: its first fresh read_file must not force-stop. result = mw._apply(_make_state(tool_calls=[self._read_call("/fresh.py")]), evicted) @@ -1327,7 +1655,7 @@ class TestToolFrequencyDetection: for i in range(2): mw._apply(_make_state(tool_calls=[self._read_call(f"/file_{i}.py")]), runtime) mw.reset(thread_id="thread-A") - assert "thread-A" not in mw._tool_name_counter + assert _scope_key("thread-A") not in mw._tool_name_counter result = mw._apply(_make_state(tool_calls=[self._read_call("/fresh.py")]), runtime) assert result is None @@ -1444,7 +1772,7 @@ class TestToolCallBatchDecisions: assert result["messages"][0].tool_calls == [] assert mw.consume_stop_reason("test-run") == "loop_capped" assert not mw._pending_warnings - assert not mw._tool_freq_warned.get("test-thread") + assert not mw._tool_freq_warned.get(_scope_key()) journal.record_middleware.assert_called_once() recorded = journal.record_middleware.call_args.kwargs assert recorded["action"] == "hard_stop" @@ -1466,7 +1794,7 @@ class TestToolCallBatchDecisions: assert decision is not None and decision.hard_stop assert decision.tool_names == ("bash",) assert decision.count == 3 - assert "read_file" not in mw._tool_freq_warned["test-thread"] + assert "read_file" not in mw._tool_freq_warned[_scope_key()] def test_identical_call_warning_cannot_mask_frequency_hard_stop(self): mw = LoopDetectionMiddleware(warn_threshold=2, hard_limit=5, tool_freq_warn=2, tool_freq_hard_limit=3) @@ -1479,7 +1807,7 @@ class TestToolCallBatchDecisions: assert decision is not None and decision.hard_stop assert decision.detection_layer == "tool_frequency" assert decision.count == 3 - assert not mw._warned.get("test-thread") + assert not mw._warned.get(_scope_key()) def test_warning_batch_counts_all_calls_and_only_marks_selected_warning(self): mw = LoopDetectionMiddleware(tool_freq_warn=2, tool_freq_hard_limit=10) @@ -1490,9 +1818,9 @@ class TestToolCallBatchDecisions: assert first is not None and first.action == "warn" assert first.tool_names == ("read_file",) - assert list(mw._tool_name_history["test-thread"]) == ["read_file"] * 3 + ["bash"] * 2 - assert dict(mw._tool_name_counter["test-thread"]) == {"read_file": 3, "bash": 2} - assert mw._tool_freq_warned["test-thread"] == {"read_file"} + assert list(mw._tool_name_history[_scope_key()]) == ["read_file"] * 3 + ["bash"] * 2 + assert dict(mw._tool_name_counter[_scope_key()]) == {"read_file": 3, "bash": 2} + assert mw._tool_freq_warned[_scope_key()] == {"read_file"} second = mw._track_and_check(_make_state(tool_calls=[self._call("bash", 2)]), runtime) assert second is not None and second.action == "warn" @@ -1508,8 +1836,8 @@ class TestToolCallBatchDecisions: hash_warning = mw._track_and_check(_make_state(tool_calls=calls), runtime) assert hash_warning is not None and hash_warning.detection_layer == "identical_call_set" - assert mw._tool_name_counter["test-thread"]["read_file"] == 2 - assert not mw._tool_freq_warned.get("test-thread") + assert mw._tool_name_counter[_scope_key()]["read_file"] == 2 + assert not mw._tool_freq_warned.get(_scope_key()) freq_warning = mw._track_and_check(_make_state(tool_calls=[self._call("read_file", 1)]), runtime) assert freq_warning is not None and freq_warning.detection_layer == "tool_frequency" @@ -1523,8 +1851,8 @@ class TestToolCallBatchDecisions: first = mw._track_and_check(_make_state(tool_calls=calls), runtime) assert first is not None and first.tool_names == ("a",) - assert mw._tool_name_counter["test-thread"]["a"] == 1 - assert "a" not in mw._tool_freq_warned["test-thread"] + assert mw._tool_name_counter[_scope_key()]["a"] == 1 + assert "a" not in mw._tool_freq_warned[_scope_key()] second = mw._track_and_check(_make_state(tool_calls=[self._call("a", 5)]), runtime) assert second is not None and second.action == "warn" assert second.tool_names == ("a",) @@ -1541,8 +1869,8 @@ class TestToolCallBatchDecisions: second = mw._track_and_check(_make_state(tool_calls=[self._call("b", 3)]), runtime) assert second is not None and second.tool_names == ("b",) - assert mw._tool_name_counter["test-thread"]["a"] == 1 - assert "a" not in mw._tool_freq_warned["test-thread"] + assert mw._tool_name_counter[_scope_key()]["a"] == 1 + assert "a" not in mw._tool_freq_warned[_scope_key()] third = mw._track_and_check(_make_state(tool_calls=[self._call("a", 4)]), runtime) assert third is not None and third.tool_names == ("a",) @@ -1561,15 +1889,15 @@ class TestToolCallBatchDecisions: runtime, ) assert first is not None and first.tool_names == ("bash",) - assert mw._tool_freq_warned["test-thread"] == {"bash"} + assert mw._tool_freq_warned[_scope_key()] == {"bash"} second = mw._track_and_check( _make_state(tool_calls=[self._call("read_file", i) for i in range(3)]), runtime, ) assert second is not None and second.tool_names == ("read_file",) - assert mw._tool_name_counter["test-thread"]["bash"] == 2 - assert "bash" not in mw._tool_freq_warned["test-thread"] + assert mw._tool_name_counter[_scope_key()]["bash"] == 2 + assert "bash" not in mw._tool_freq_warned[_scope_key()] third = mw._track_and_check( _make_state(tool_calls=[self._call("bash", i) for i in range(3, 6)]), diff --git a/config.example.yaml b/config.example.yaml index aee154483..4d6e924c4 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -1236,6 +1236,7 @@ loop_detection: warn_threshold: 3 hard_limit: 5 window_size: 20 + # Compatibility name: bounds retained (thread_id, run_id) histories, not whole threads. max_tracked_threads: 100 tool_freq_warn: 30 tool_freq_hard_limit: 50