fix(adaptive_router/hooks): populate tool_results so failure signal fires

The post-call hook was hardcoding tool_results=[] on every Turn, so the
failure detector never saw tool errors and the bandit only learned from
satisfaction — never from negative tool outcomes.

Added _recent_tool_results(messages): walks the request messages from the
tail and collects the contiguous run of role=='tool' entries — those are
the results from the most recent assistant tool_calls round. Normalizes
each to {content, is_error}, the only fields signals._detect_failure /
_detect_exhaustion read.

Tests: 6 new covering empty input, trailing-run extraction, is_error
propagation, boundary at first non-tool message, no-trailing-tool case,
and the end-to-end path from hook -> Turn.tool_results.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-20 15:25:51 -07:00
parent fba736ca3c
commit 0cfcec68e9
2 changed files with 106 additions and 1 deletions

View file

@ -102,6 +102,36 @@ def _last_user_content(messages: Optional[List[Dict[str, Any]]]) -> Optional[str
return None
def _recent_tool_results(messages: Optional[List[Dict[str, Any]]]) -> List[Dict[str, Any]]:
"""Extract the current turn's tool result payloads from the request messages.
Tool results are `role == "tool"` messages that sit at the tail of the
conversation — i.e. after the most recent assistant message with
`tool_calls`, waiting for the model to produce a user-facing reply. Walk
backwards from the end and collect the contiguous run of tool messages;
stop at the first non-tool message.
Each result is normalized to `{content, is_error}` — the only fields
`signals._detect_failure` / `_detect_exhaustion` actually read.
"""
if not messages:
return []
results: List[Dict[str, Any]] = []
for msg in reversed(messages):
if not isinstance(msg, dict):
break
if msg.get("role") != "tool":
break
content = msg.get("content")
# Some providers (Anthropic-style) carry an explicit error flag; OpenAI
# tool results don't, so fall back to an empty/missing content heuristic
# inside `_detect_failure`.
is_error = bool(msg.get("is_error"))
results.append({"content": content, "is_error": is_error})
results.reverse()
return results
def _assistant_content_and_tool_calls(response_obj: Any) -> tuple:
"""Return (assistant_text, tool_calls_list) extracted from a ModelResponse-ish object."""
if response_obj is None:
@ -222,6 +252,7 @@ class AdaptiveRouterPostCallHook(CustomLogger):
user_text = _last_user_content(messages)
assistant_text, tool_calls = _assistant_content_and_tool_calls(response_obj)
tool_results = _recent_tool_results(messages)
request_type = classify_prompt(user_text or "")
turn = Turn(
@ -230,7 +261,7 @@ class AdaptiveRouterPostCallHook(CustomLogger):
assistant_text if isinstance(assistant_text, str) else None
),
tool_calls=tool_calls,
tool_results=[],
tool_results=tool_results,
response_status=response_status,
)
await self.adaptive_router.record_turn(

View file

@ -10,6 +10,7 @@ from litellm.router_strategy.adaptive_router.config import (
)
from litellm.router_strategy.adaptive_router.hooks import (
AdaptiveRouterPostCallHook,
_recent_tool_results,
_resolve_session_key,
)
from litellm.router_strategy.adaptive_router.signals import Turn
@ -220,6 +221,79 @@ async def test_hook_passes_tool_calls_through():
assert turn.tool_calls == [tc]
# ---- _recent_tool_results ------------------------------------------------
def test_recent_tool_results_empty_when_no_messages():
assert _recent_tool_results(None) == []
assert _recent_tool_results([]) == []
def test_recent_tool_results_collects_trailing_tool_messages():
"""Tool messages at the tail of the conversation are extracted in order."""
messages = [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": None, "tool_calls": [{"id": "t1"}]},
{"role": "tool", "tool_call_id": "t1", "content": "result A"},
{"role": "tool", "tool_call_id": "t2", "content": "result B"},
]
results = _recent_tool_results(messages)
assert [r["content"] for r in results] == ["result A", "result B"]
assert all(r["is_error"] is False for r in results)
def test_recent_tool_results_propagates_is_error_flag():
messages = [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": None, "tool_calls": [{"id": "t1"}]},
{"role": "tool", "content": "boom", "is_error": True},
]
results = _recent_tool_results(messages)
assert results == [{"content": "boom", "is_error": True}]
def test_recent_tool_results_stops_at_first_non_tool_message():
"""Only the trailing run of tool messages counts — prior rounds are
considered already attributed."""
messages = [
{"role": "user", "content": "hi"},
{"role": "tool", "content": "stale"}, # earlier round, ignored
{"role": "assistant", "content": "intermediate"},
{"role": "user", "content": "follow-up"},
{"role": "tool", "content": "current"},
]
results = _recent_tool_results(messages)
assert [r["content"] for r in results] == ["current"]
def test_recent_tool_results_empty_when_no_trailing_tool_message():
messages = [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
]
assert _recent_tool_results(messages) == []
@pytest.mark.asyncio
async def test_hook_passes_tool_results_to_turn_for_failure_detection():
"""A trailing tool message with `is_error` must reach `Turn.tool_results`
so the failure-signal path fires."""
hook = _make_hook()
messages = _long_messages()
messages.append(
{"role": "assistant", "content": None, "tool_calls": [{"id": "t1"}]}
)
messages.append(
{"role": "tool", "tool_call_id": "t1", "content": "500", "is_error": True}
)
kwargs = _kwargs(chosen="fast", messages=messages)
await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0)
turn: Turn = hook.adaptive_router.record_turn.await_args.kwargs["turn"]
assert turn.tool_results == [{"content": "500", "is_error": True}]
@pytest.mark.asyncio
async def test_hook_swallows_exceptions_from_record_turn():
hook = _make_hook()