mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
fba736ca3c
commit
0cfcec68e9
2 changed files with 106 additions and 1 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue