feat(router): auto-escalate stalled complexity-router tasks

Adds stall_escalation_enabled to the complexity router: when the assistant's
own recent tool calls look stuck (identical repeats, or repeated tool
errors on a surface that reports one), the request is bumped one
configured tier higher, the automatic counterpart to escalation_keywords.

Detection is stateless: it rereads the last stall_escalation_window tool
calls from that request's own message list on every classified turn, so
the bump lasts only as long as the recent calls still look stuck and
lifts on its own once they don't, and evidence survives a plain follow-up
like "try again" instead of resetting on the newest human ask.

Off by default. Rejected together with session_affinity and
classification_mode='user_turn', which both replay a held routing
decision instead of classifying most turns, and with tier_definitions,
for the same reason escalation_keywords is: both rely on the built-in
tier severity order a custom tier set does not define.

Dashboard controls for this are not included; config.yaml and the
management API accept it today through ComplexityRouterConfig.
This commit is contained in:
moe-berri 2026-09-04 14:15:39 -07:00
parent 300d335255
commit 7c6638e5c3
6 changed files with 492 additions and 20 deletions

View file

@ -247,6 +247,52 @@ unless `modality_routing` is also on.
`session_affinity_ttl_seconds` is the idle window for both the model pin selected by session affinity and the deployment pin. Every request that reuses a pin refreshes its TTL, so a session actively sending requests stays pinned. After the window passes with no pin reuse, the next request classifies again and creates a fresh pin. Omit the setting to track the default of 3600 seconds.
### Mid-task stall escalation
A weak model working an agentic task can get stuck: it keeps calling the same tool with the
same arguments, or the same call keeps erroring, when a stronger model would have broken the
loop. `stall_escalation_enabled: true` catches this and bumps the request one tier higher, the
automatic counterpart to a user typing an escalation keyword:
```yaml
model_list:
- model_name: smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
stall_escalation_enabled: true
stall_escalation_window: 6
stall_escalation_repeat_threshold: 3
tiers:
SIMPLE: gpt-4o-mini
MEDIUM: gpt-4o
COMPLEX: claude-sonnet-4
REASONING: o1-preview
```
Detection looks at the assistant's own tool calls, not the human's messages: of the last
`stall_escalation_window` tool calls, if `stall_escalation_repeat_threshold` or more are
identical (same tool, same arguments) or came back as errors, the task counts as stalled and the
classified tier is bumped one step by the same `_escalate_tier` ladder `escalation_keywords`
uses, capped at the highest configured tier. It reads both tool-call shapes: Anthropic Messages
`tool_use`/`tool_result` blocks (including `is_error`) and chat-completions `tool_calls`/`tool`
messages (which carry no standard error flag, so those calls are judged on repetition alone).
There is no state to expire or leak: detection reruns on every classified turn from that
request's own message list, so the bump lasts only as long as the recent tool calls still look
stuck and lifts on its own the moment they don't. This also means it reads the whole
conversation rather than only the turns since the newest human ask, so a plain follow-up like
"try again" does not discard evidence from before it. Escalation records `stall_escalation` in
`routing_decision.signals`; unlike `escalation_keywords`, it does not set the
`escalated`/`escalation_keyword` pair, which is reserved for the keyword mechanism specifically.
`stall_escalation_enabled` cannot be combined with `session_affinity` or
`classification_mode: user_turn`: both replay a held routing decision on most turns instead of
classifying, so detection would never see the tool calls it needs to look at. It is also
rejected together with `tier_definitions`, for the same reason `escalation_keywords` is: both
rely on the built-in tier severity order, which a custom tier set does not define. Off by
default.
### Heuristic-first chaining
`classifier_type: heuristic_first` runs the local scorer on every request and only calls the LLM

View file

@ -70,6 +70,7 @@ from .config import (
ComplexityTier,
TierDefinition,
)
from .stall_detector import detect_stalled_task
if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
@ -3135,6 +3136,17 @@ class ComplexityRouter(CustomLogger):
escalated: Final = tier != classified_tier
if escalated:
signals = (*signals, "escalation")
# Recomputed from this request's own tool calls, not remembered from a prior turn: the
# bump lasts only as long as the recent tool calls still look stuck, and lifts itself
# the moment they don't, with nothing to expire or leak past the task that earned it.
stalled: Final = self.config.stall_escalation_enabled and detect_stalled_task(
resolved_messages,
window=self.config.stall_escalation_window,
repeat_threshold=self.config.stall_escalation_repeat_threshold,
)
if stalled:
tier = self._escalate_tier(tier)
signals = (*signals, "stall_escalation")
pre_floor_tier: Final = tier
if plan_floor is not None:
tier = self._apply_plan_mode_floor(tier)

View file

@ -809,6 +809,42 @@ class ComplexityRouterConfig(BaseModel):
description="Rules that force a specific tier when their keywords match the prompt",
)
stall_escalation_enabled: bool = Field(
default=False,
description=(
"Escalate mid-task to the next-higher configured tier when the assistant's own recent "
"tool calls look stuck: stall_escalation_repeat_threshold or more of the last "
"stall_escalation_window tool calls are identical repeats (same tool, same arguments) "
"or came back as errors. One tier at most, on the same ladder escalation_keywords bumps "
"along, and never above the highest configured tier. Detection re-runs on every "
"classified turn from the tool calls visible in that request, so it needs no state and "
"nothing survives past the task: once the recent tool calls stop looking stuck, the "
"next classified turn routes normally again. Mutually exclusive with session_affinity "
"and classification_mode='user_turn', which both replay a held routing decision instead "
"of classifying most turns, so this would never see the tool calls to look at. Off by "
"default."
),
)
stall_escalation_window: int = Field(
default=6,
gt=0,
description=(
"How many of the assistant's most recent tool calls stall detection looks at, oldest "
"ones dropped as new calls happen. Counted across the whole visible conversation "
"rather than reset at the newest human ask, so evidence from before a plain follow-up "
"message like 'try again' is still visible on the turn after it."
),
)
stall_escalation_repeat_threshold: int = Field(
default=3,
ge=2,
description=(
"How many of the last stall_escalation_window tool calls must be identical repeats, or "
"error results, before the task counts as stalled. Must not exceed "
"stall_escalation_window, or the condition could never be reached."
),
)
plan_mode_min_tier: str | None = Field(
default=None,
description=(
@ -1246,6 +1282,7 @@ class ComplexityRouterConfig(BaseModel):
("adaptive", self.adaptive),
("session_affinity", self.session_affinity),
("escalation_keywords", bool(self.escalation_keywords)),
("stall_escalation_enabled", self.stall_escalation_enabled),
("plugins", bool(self.plugins)),
)
if enabled
@ -1422,6 +1459,25 @@ class ComplexityRouterConfig(BaseModel):
)
return self
@model_validator(mode="after")
def _validate_stall_escalation(self) -> "ComplexityRouterConfig":
if not self.stall_escalation_enabled:
return self
if self.session_affinity or self.classification_mode == "user_turn":
raise ValueError(
"stall_escalation_enabled cannot be combined with session_affinity or "
"classification_mode='user_turn': both replay a held routing decision on most "
"turns instead of classifying, so stall detection would never see the tool calls "
"of the turns it needs to look at. Disable one or the other."
)
if self.stall_escalation_repeat_threshold > self.stall_escalation_window:
raise ValueError(
"stall_escalation_repeat_threshold "
f"({self.stall_escalation_repeat_threshold}) cannot exceed stall_escalation_window "
f"({self.stall_escalation_window}); the condition could never be reached."
)
return self
@model_validator(mode="after")
def _validate_tier_param_placement(self) -> "ComplexityRouterConfig":
"""Reject a router setting written into a tier entry's request params.

View file

@ -0,0 +1,126 @@
"""
Mid-task stall detection for the Complexity Router.
Looks at the assistant's own recent tool calls -- visible on every request an agentic
client resends, since each turn carries the whole conversation so far -- for a tight loop
of identical calls or repeated tool errors. No LLM call, no state: the same fixed-size
window is rescanned on every classified turn, so a stall reads the same way whether it
started one turn ago or ten, and stops reading as a stall the moment the recent calls
change.
Assistant tool calls appear in two shapes depending on the API surface, and this module
reads both without translating one into the other:
- Anthropic Messages: assistant `content` blocks of type "tool_use" (id, name, input),
answered by a later user-turn `content` block of type "tool_result" (tool_use_id,
is_error).
- Chat completions: assistant `tool_calls` entries (id, function.name, function.arguments
as a JSON string), answered by a later `role: "tool"` message. Chat completions has no
standard error flag on that message, so those calls are judged on repetition alone.
"""
from __future__ import annotations
import json
from collections import Counter
from collections.abc import Iterator, Mapping, Sequence
from itertools import islice
from typing import Final, NamedTuple
_ARGUMENTS_PARSE_FAILED: Final = object()
class _ToolCallEvent(NamedTuple):
signature: tuple[str, str]
is_error: bool | None
"""None when the surface carries no structured error signal for this call. Never
treated as an error: a call this module cannot judge must not count toward the tally."""
def _json_arguments(raw: str) -> object:
try:
return json.loads(raw)
except (TypeError, ValueError):
return _ARGUMENTS_PARSE_FAILED
def _tool_call_signature(name: str, raw_arguments: object) -> tuple[str, str]:
"""A (name, canonical-arguments) pair that compares equal across both surfaces'
argument shapes: a dict (Anthropic `input`) and a JSON-encoded string (chat
completions `function.arguments`) representing the same call must match."""
parsed: Final = _json_arguments(raw_arguments) if isinstance(raw_arguments, str) else raw_arguments
arguments: Final = raw_arguments if parsed is _ARGUMENTS_PARSE_FAILED else parsed
try:
return name, json.dumps(arguments, sort_keys=True, default=str)
except (TypeError, ValueError):
return name, str(arguments)
def _iter_tool_result_error_pairs(messages: Sequence[Mapping[str, object]]) -> Iterator[tuple[str, bool]]:
"""(call id, whether that call's result was an error), read only where the surface
reports one: an Anthropic Messages `tool_result` content block's `is_error`."""
for msg in messages:
content = msg.get("content")
if msg.get("role") != "user" or not isinstance(content, list):
continue
for part in content:
if isinstance(part, Mapping) and part.get("type") == "tool_result":
call_id = part.get("tool_use_id")
if isinstance(call_id, str):
yield call_id, bool(part.get("is_error", False))
def _iter_tool_call_events_newest_first(messages: Sequence[Mapping[str, object]]) -> Iterator[_ToolCallEvent]:
"""Every tool call the assistant made, newest first, paired with its result's error
status where the surface reports one."""
error_by_call_id: Final = dict(_iter_tool_result_error_pairs(messages))
for msg in reversed(messages):
if msg.get("role") != "assistant":
continue
content = msg.get("content")
if isinstance(content, list):
for part in reversed(content):
if not (isinstance(part, Mapping) and part.get("type") == "tool_use"):
continue
name = part.get("name")
if isinstance(name, str):
call_id = part.get("id")
yield _ToolCallEvent(
signature=_tool_call_signature(name, part.get("input")),
is_error=error_by_call_id.get(call_id) if isinstance(call_id, str) else None,
)
tool_calls = msg.get("tool_calls")
if not isinstance(tool_calls, list):
continue
for call in reversed(tool_calls):
function = call.get("function") if isinstance(call, Mapping) else None
name = function.get("name") if isinstance(function, Mapping) else None
if isinstance(name, str):
yield _ToolCallEvent(
signature=_tool_call_signature(name, function.get("arguments") if function else None),
is_error=None,
)
def detect_stalled_task(
messages: Sequence[Mapping[str, object]] | None,
*,
window: int,
repeat_threshold: int,
) -> bool:
"""Whether the assistant's recent tool-call activity looks stuck: repeat_threshold or
more of the last `window` tool calls share an identical signature, or resolved to an
error on a surface that reports one.
Reads the whole message list rather than only the turns since the newest human ask,
so a follow-up like "try again" does not discard the evidence that came before it.
"""
if not messages or repeat_threshold <= 0:
return False
recent: Final = tuple(islice(_iter_tool_call_events_newest_first(messages), window))
if len(recent) < repeat_threshold:
return False
_, most_common_count = Counter(event.signature for event in recent).most_common(1)[0]
if most_common_count >= repeat_threshold:
return True
error_count: Final = sum(1 for event in recent if event.is_error)
return error_count >= repeat_threshold

View file

@ -2690,9 +2690,7 @@ class TestRouterPreRoutingAliasOverrides:
import time
monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path))
(tmp_path / "api-key.json").write_text(
json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600})
)
(tmp_path / "api-key.json").write_text(json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600}))
router = Router(
model_list=[
{
@ -2717,7 +2715,9 @@ class TestRouterPreRoutingAliasOverrides:
copilot_resolutions: List = []
def _guarded(*args, **kwargs):
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
kwargs.get("custom_llm_provider") or ""
)
if "github_copilot" in target:
copilot_resolutions.append(target)
raise RuntimeError("routing must not resolve an authenticating provider")
@ -5887,6 +5887,123 @@ class TestEscalationKeywords:
assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c
def _stalled_tool_history(repeats: int = 3) -> List[Dict]:
"""`repeats` identical bash tool calls in a row, the automatic counterpart to a user
typing an escalation keyword: the assistant, not the human, is the one stuck."""
return [
turn
for i in range(repeats)
for turn in (
{
"role": "assistant",
"content": [{"type": "tool_use", "id": f"call-{i}", "name": "bash", "input": {"cmd": "pytest"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": f"call-{i}", "is_error": True, "content": "fail"}],
},
)
]
class TestStallEscalation:
"""Mid-task auto-escalation when the assistant's own recent tool calls look stuck: the
automatic counterpart to escalation_keywords, gated by stall_escalation_enabled and off
by default."""
@pytest.mark.asyncio
async def test_repeated_tool_calls_escalate_the_classified_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o" # SIMPLE bumped to MEDIUM
@pytest.mark.asyncio
async def test_varied_tool_calls_do_not_escalate(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [
{
"role": "assistant",
"content": [{"type": "tool_use", "id": "c1", "name": "bash", "input": {"cmd": "ls"}}],
},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "c1", "is_error": False, "content": "ok"}],
},
{"role": "user", "content": "Hello there!"},
]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o-mini" # not escalated
@pytest.mark.asyncio
async def test_disabled_by_default_ignores_repeated_tool_calls(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=basic_config,
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o-mini" # stall_escalation_enabled defaults False
@pytest.mark.asyncio
async def test_signals_record_stall_escalation(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert "stall_escalation" in result.routing_decision["signals"]
@pytest.mark.asyncio
async def test_stall_escalation_caps_at_highest_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [
*_stalled_tool_history(),
{"role": "user", "content": "Let's think step by step and reason through this carefully."},
]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "o1-preview" # already REASONING, stays there
@pytest.mark.asyncio
async def test_stall_escalation_stacks_with_keyword_escalation(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "LITELLM ESCALATE Hello there!"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "claude-sonnet-4-20250514" # SIMPLE -> MEDIUM (keyword) -> COMPLEX (stall)
@pytest.mark.asyncio
async def test_evidence_survives_a_new_human_ask(self, mock_router_instance, basic_config):
"""A plain follow-up like 'try again' must not erase the stall evidence that came
before it: escalation still fires on the turn carrying that follow-up."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "stall_escalation_enabled": True},
)
messages = [*_stalled_tool_history(), {"role": "user", "content": "try again"}]
result = await router.async_pre_routing_hook(model="test-model", request_kwargs={}, messages=messages)
assert result.model == "gpt-4o" # SIMPLE ("try again" carries no signal) bumped to MEDIUM
class TestRoutingDecisionContents:
"""Every routing path must return a PreRoutingHookResponse carrying a routing_decision
that names the mechanism that actually decided, with the facts of that path only."""
@ -7781,7 +7898,6 @@ class TestClientHousekeepingCalls:
assert result is not None
assert result.model == "claude-sonnet-4-20250514"
@pytest.mark.asyncio
async def test_a_classifier_plugin_still_decides_its_own_routers(self, mock_router_instance):
"""A plugin is where an operator encodes policy the tier ladder cannot express.
@ -7816,9 +7932,7 @@ class TestClientHousekeepingCalls:
assert result.model == "o1-preview"
assert result.routing_decision["cause"] == "classifier_plugin"
def _adaptive_router(
self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None
) -> ComplexityRouter:
def _adaptive_router(self, tier_distance_penalty: float, plan_mode_min_tier: str | None = None) -> ComplexityRouter:
adaptive_instance = MagicMock()
adaptive_instance.model_list = [
{
@ -7855,9 +7969,7 @@ class TestClientHousekeepingCalls:
return router
@pytest.mark.asyncio
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(
self, mock_router_instance
):
async def test_the_bandit_cannot_route_a_housekeeping_call_above_the_cheapest_tier(self, mock_router_instance):
"""The tier here is what the request IS, not how hard it is, so the bandit has nothing to win.
Without a ceiling the tier distance penalty is the only thing holding the tier, so a
@ -7890,7 +8002,6 @@ class TestClientHousekeepingCalls:
assert result is not None
assert result.model == "premium"
@pytest.mark.asyncio
async def test_a_housekeeping_call_never_becomes_the_session_pin(self, mock_router_instance):
"""Pinning this is the most expensive mistake of the transient causes.
@ -7932,9 +8043,7 @@ class TestClientHousekeepingCalls:
assert work_turn.routing_decision["cause"] == "llm_classifier"
@pytest.mark.asyncio
async def test_the_decision_records_which_sentinel_matched(
self, mock_router_instance, llm_classifier_config
):
async def test_the_decision_records_which_sentinel_matched(self, mock_router_instance, llm_classifier_config):
"""The cause's contract says the sentinel rides in matched_keyword, so it has to be there.
Without it an operator reading the logs can see that a call was treated as housekeeping but
@ -7955,7 +8064,6 @@ class TestClientHousekeepingCalls:
"Write the title in the predominant language of the session"
)
@pytest.mark.asyncio
async def test_the_plan_mode_floor_raises_a_housekeeping_call_under_adaptive(self, mock_router_instance):
"""Floor and ceiling must not contradict each other on the same request.
@ -9054,6 +9162,7 @@ class TestTierDefinitions:
({"adaptive": True}, "severity order"),
({"session_affinity": True}, "severity order"),
({"escalation_keywords": ["GO UP"]}, "severity order"),
({"stall_escalation_enabled": True}, "severity order"),
(
{"classifier_llm_config": {"model": "haiku-classifier", "system_prompt": "grade it"}},
"system_prompt",
@ -10340,9 +10449,7 @@ class TestHeuristicFirst:
# Scores 0.175 with one signal, so it sits 0.025 from simple_medium: the pair of tiers either side of
# that boundary are different model pools, and a hair's difference in score picks the other one.
NEAR_BOUNDARY_PROMPT = (
"design a distributed cache with consistent hashing, then explain the failure modes step by step"
)
NEAR_BOUNDARY_PROMPT = "design a distributed cache with consistent hashing, then explain the failure modes step by step"
# Scores 0.075 with signals, the far side of any margin under 0.075: the scorer is decided here.
CLEAR_OF_BOUNDARY_PROMPT = "explain step by step how consistent hashing rebalances keys"
@ -10784,6 +10891,7 @@ class TestContextWindowEscalation:
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
def session_kwargs() -> dict[str, object]:
return {"metadata": {"session_id": "s-1", "user_api_key_hash": "k-1"}}
@ -10808,6 +10916,7 @@ class TestContextWindowEscalation:
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
def session_kwargs() -> dict[str, object]:
return {"metadata": {"session_id": "s-2", "user_api_key_hash": "k-2"}}
@ -10885,7 +10994,9 @@ class TestContextWindowEscalation:
copilot_resolutions: List = []
def _guarded(*args, **kwargs):
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
target = str(kwargs.get("model") or (args[0] if args else "")) + str(
kwargs.get("custom_llm_provider") or ""
)
if "github_copilot" in target:
copilot_resolutions.append(target)
raise RuntimeError("the gate must not resolve an authenticating provider")

View file

@ -0,0 +1,121 @@
"""
Tests for mid-task stall detection: repeated identical tool calls or repeated tool
errors, read from both Anthropic Messages and chat-completions tool-call shapes.
"""
from litellm.router_strategy.complexity_router.stall_detector import detect_stalled_task
def _anthropic_call(call_id: str, name: str, arguments: dict, *, is_error: bool) -> list[dict]:
return [
{"role": "assistant", "content": [{"type": "tool_use", "id": call_id, "name": name, "input": arguments}]},
{
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": call_id, "is_error": is_error, "content": "result"}],
},
]
def _chat_completions_call(call_id: str, name: str, arguments_json: str) -> list[dict]:
return [
{
"role": "assistant",
"tool_calls": [
{"id": call_id, "type": "function", "function": {"name": name, "arguments": arguments_json}}
],
},
{"role": "tool", "tool_call_id": call_id, "content": "result"},
]
class TestDetectStalledTask:
def test_repeated_identical_anthropic_calls_are_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_repeated_errors_are_stalled_even_with_varied_arguments(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest tests/a.py"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest tests/b.py"}, is_error=True),
*_anthropic_call("t3", "bash", {"cmd": "pytest tests/c.py"}, is_error=True),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_varied_successful_calls_are_not_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "ls"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "grep", {"pattern": "x"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_chat_completions_repeats_are_stalled(self):
messages = [
*_chat_completions_call("c1", "bash", '{"cmd": "pytest"}'),
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
*_chat_completions_call("c3", "bash", '{"cmd": "pytest"}'),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_chat_completions_has_no_structured_error_signal(self):
"""A chat-completions tool message carries no standard error flag, so varied calls
whose content happens to read like failures still aren't flagged on error alone."""
messages = [
*_chat_completions_call("c1", "bash", '{"cmd": "a"}'),
*_chat_completions_call("c2", "bash", '{"cmd": "b"}'),
*_chat_completions_call("c3", "bash", '{"cmd": "c"}'),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_dict_and_json_string_arguments_compare_equal_across_surfaces(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_chat_completions_call("c2", "bash", '{"cmd": "pytest"}'),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_below_repeat_threshold_is_not_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is False
def test_evidence_older_than_the_window_does_not_count(self):
"""Only the most recent `window` tool calls are considered, so a stall the model
already recovered from does not keep re-triggering forever."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t4", "grep", {"pattern": "a"}, is_error=False),
*_anthropic_call("t5", "grep", {"pattern": "b"}, is_error=False),
]
assert detect_stalled_task(messages, window=2, repeat_threshold=2) is False
def test_evidence_survives_a_new_human_ask(self):
"""A follow-up like 'try again' must not erase evidence from before it: detection
reads the whole message list, not just the turns since the newest human ask."""
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=False),
*_anthropic_call("t3", "bash", {"cmd": "pytest"}, is_error=False),
{"role": "user", "content": [{"type": "text", "text": "try again"}]},
]
assert detect_stalled_task(messages, window=6, repeat_threshold=3) is True
def test_no_messages_is_not_stalled(self):
assert detect_stalled_task(None, window=6, repeat_threshold=3) is False
assert detect_stalled_task([], window=6, repeat_threshold=3) is False
def test_zero_threshold_never_flags_stalled(self):
messages = [
*_anthropic_call("t1", "bash", {"cmd": "pytest"}, is_error=True),
*_anthropic_call("t2", "bash", {"cmd": "pytest"}, is_error=True),
]
assert detect_stalled_task(messages, window=6, repeat_threshold=0) is False