Fix lost loop_capped stop reason when a subagent's run_id is None (#4059)

LoopDetectionMiddleware._get_run_id used a truthiness check that
collapsed a present-but-None run_id to the same "default" key as a
totally absent one. SubagentExecutor sets context["run_id"] =
self.run_id unconditionally, so run_id is genuinely None for an
embedded/TUI-dispatched subagent, and later reads the stop reason back
with that same raw attribute via consume_stop_reason(self.run_id). The
write (keyed "default") and the read (keyed None) disagreed, so a
genuine loop-detection hard-stop's loop_capped reason was silently
dropped instead of reaching the lead.

Align _get_run_id with TokenBudgetMiddleware's key-presence-based
version, which does not have this bug: return the context value as-is
when the key is present (None included), and fall back to a
per-runtime-unique key only when the key is absent.
This commit is contained in:
Daoyuan Li 2026-07-12 08:12:20 -07:00 committed by GitHub
parent d2ab5bb819
commit c4b651fb9d
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 70 additions and 9 deletions

View File

@ -274,11 +274,30 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]):
return "default" return "default"
def _get_run_id(self, runtime: Runtime) -> str: def _get_run_id(self, runtime: Runtime) -> str:
"""Extract run_id from runtime context for per-run warning scoping.""" """Extract run_id from runtime context for per-run warning scoping.
run_id = runtime.context.get("run_id") if runtime.context else None
if run_id: Keyed by presence, not truthiness: ``SubagentExecutor`` sets
return str(run_id) ``context["run_id"] = self.run_id`` unconditionally (no truthiness
return "default" 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``.
"""
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))
def consume_stop_reason(self, run_id: str | None) -> str | 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. """Pop and return the stop reason the hard-stop set for this run.

View File

@ -2,6 +2,7 @@
import copy import copy
from collections import OrderedDict from collections import OrderedDict
from types import SimpleNamespace
from typing import Any from typing import Any
from unittest.mock import MagicMock from unittest.mock import MagicMock
@ -304,8 +305,14 @@ class TestLoopDetection:
mw.wrap_model_call(request_a, handler) mw.wrap_model_call(request_a, handler)
assert any(isinstance(message, HumanMessage) and message.name == "loop_warning" for message in captured[1].messages) assert any(isinstance(message, HumanMessage) and message.name == "loop_warning" for message in captured[1].messages)
def test_missing_run_id_uses_default_pending_scope(self): def test_missing_run_id_uses_per_runtime_pending_scope(self):
"""When runtime has no run_id, warning handling falls back to the default run scope.""" """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)."""
mw = LoopDetectionMiddleware(warn_threshold=3, hard_limit=10) mw = LoopDetectionMiddleware(warn_threshold=3, hard_limit=10)
runtime = MagicMock() runtime = MagicMock()
runtime.context = {"thread_id": "test-thread"} runtime.context = {"thread_id": "test-thread"}
@ -314,7 +321,8 @@ class TestLoopDetection:
for _ in range(3): for _ in range(3):
mw._apply(_make_state(tool_calls=call), runtime) mw._apply(_make_state(tool_calls=call), runtime)
assert mw._pending_warnings.get(_pending_key(run_id="default")) fallback_run_id = str(id(runtime))
assert mw._pending_warnings.get(_pending_key(run_id=fallback_run_id))
request = _make_request([AIMessage(content="hi")], runtime) request = _make_request([AIMessage(content="hi")], runtime)
captured, handler = _capture_handler() captured, handler = _capture_handler()
@ -323,7 +331,7 @@ class TestLoopDetection:
loop_warnings = [message for message in captured[0].messages if isinstance(message, HumanMessage) and message.name == "loop_warning"] loop_warnings = [message for message in captured[0].messages if isinstance(message, HumanMessage) and message.name == "loop_warning"]
assert len(loop_warnings) == 1 assert len(loop_warnings) == 1
assert "LOOP DETECTED" in loop_warnings[0].content assert "LOOP DETECTED" in loop_warnings[0].content
assert not mw._pending_warnings.get(_pending_key(run_id="default")) assert not mw._pending_warnings.get(_pending_key(run_id=fallback_run_id))
def test_before_agent_clears_stale_pending_warnings_for_thread(self): 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.""" """Starting a new run drops stale warnings from prior runs in the same thread."""
@ -448,6 +456,40 @@ class TestLoopDetection:
assert mw.consume_stop_reason("test-run") == "loop_capped" assert mw.consume_stop_reason("test-run") == "loop_capped"
def test_hard_stop_stamps_loop_capped_with_explicit_none_run_id(self):
"""Regression: a subagent whose ``run_id`` is genuinely ``None`` must
still round-trip its ``loop_capped`` stop reason.
``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 reason back with the raw
attribute: ``consume_stop_reason(self.run_id)``, i.e.
``consume_stop_reason(None)``. Before the fix, ``_get_run_id`` used a
truthiness check (``if run_id:``) that collapsed this present-but-None
state to the same literal ``"default"`` key used for a totally absent
run_id, so the write (``self._stop_reason["default"] = "loop_capped"``)
and this read (keyed by the raw ``None``) disagreed and the signal was
silently lost. Mirrors ``TokenBudgetMiddleware``'s key-presence-based
``_get_run_id``, which does not have this bug."""
mw = LoopDetectionMiddleware(warn_threshold=2, hard_limit=4)
runtime = SimpleNamespace(context={"thread_id": "t", "run_id": None})
call = [_bash_call("ls")]
for _ in range(3):
mw._apply(_make_state(tool_calls=call), runtime)
hard_stop = mw._apply(_make_state(tool_calls=call), runtime)
assert hard_stop is not None
# Exactly what SubagentExecutor._consume_guard_stop_reason does:
# consume_stop_reason(self.run_id), where self.run_id is the raw,
# un-normalized (possibly-None) attribute value.
assert mw.consume_stop_reason(None) == "loop_capped"
# Popped on read — a second read is None (no double-report on reuse).
assert mw.consume_stop_reason(None) is None
def test_different_calls_dont_trigger(self): def test_different_calls_dont_trigger(self):
mw = LoopDetectionMiddleware(warn_threshold=2) mw = LoopDetectionMiddleware(warn_threshold=2)
runtime = _make_runtime() runtime = _make_runtime()