fix: fix auto router bug
Some checks failed
Unit Tests: Proxy DB Operations / proxy-db (key-generation, tests/proxy_unit_tests/test_key_generate_prisma.py, 30, 0) (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-db (auth-checks, tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py, 20, 8) (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-db (remaining, tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py, 30, 8) (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled

Fixes https://github.com/BerriAI/litellm/issues/14633
This commit is contained in:
Krrish Dholakia 2026-04-16 08:31:17 -07:00
parent 2f849e59a9
commit 5f8522176e
6 changed files with 367 additions and 8 deletions

View file

@ -239,7 +239,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:

View file

@ -9612,7 +9612,7 @@ class Router:
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:

View file

@ -81,11 +81,34 @@ class AutoRouter(CustomLogger):
)
return auto_router_routes
@staticmethod
def _extract_text_from_messages(messages: List[Dict[str, Any]]) -> str:
"""
Extract text content from the last user message for routing.
Handles tool-call conversations (where the last message may be an
assistant or tool message with non-string content) and multimodal
messages (where content is a list of content blocks).
"""
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content")
if content is None:
return ""
if isinstance(content, list):
return " ".join(
block.get("text", "")
for block in content
if isinstance(block, dict) and block.get("type") == "text"
)
return str(content)
return ""
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
@ -119,8 +142,7 @@ class AutoRouter(CustomLogger):
auto_sync=self.auto_sync_value,
)
user_message: Dict[str, str] = messages[-1]
message_content: str = user_message.get("content", "")
message_content = self._extract_text_from_messages(messages)
route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(
text=message_content
)

197
repro_14633.py Normal file
View file

@ -0,0 +1,197 @@
"""
Reproduction script for https://github.com/BerriAI/litellm/issues/14633
The auto-router fails when messages contain tool calls because:
1. Function signatures use Dict[str, str] but tool call messages have non-string values
(tool_calls is a list of dicts, content can be None)
2. The content extraction in auto_router.py doesn't handle None content from
assistant messages with tool_calls
Run with: uv run python repro_14633.py
"""
import sys
from litellm.types.router import PreRoutingHookResponse
def test_pre_routing_hook_response_with_tool_calls():
"""
The PreRoutingHookResponse model was fixed to use Dict[str, Any].
This test confirms it no longer rejects messages with tool_calls.
"""
messages_with_tool_calls = [
{"role": "user", "content": "What's the weather in NYC?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "The weather in NYC is 72°F and sunny.",
},
{"role": "user", "content": "Now tell me about London"},
]
response = PreRoutingHookResponse(
model="test-model",
messages=messages_with_tool_calls,
)
print("[PASS] PreRoutingHookResponse accepts messages with tool_calls (Dict[str, Any])")
print(f" model={response.model}, num_messages={len(response.messages)}")
def test_content_extraction_fixed():
"""
Verifies the content extraction fix in auto_router.py.
Previously, when the last message was an assistant message with tool_calls
and content=None, the code would pass None to the semantic router.
Now _extract_text_from_messages finds the last *user* message and handles
None content and multimodal (list) content correctly.
"""
from litellm.router_strategy.auto_router.auto_router import AutoRouter
messages_ending_with_assistant_tool_call = [
{"role": "user", "content": "What's the weather in NYC?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
]
message_content = AutoRouter._extract_text_from_messages(
messages_ending_with_assistant_tool_call
)
print(f"\n[CHECK] Content extraction from conversation with tool_calls:")
print(f" _extract_text_from_messages result = {repr(message_content)}")
if message_content == "What's the weather in NYC?":
print(" --> FIXED: correctly extracts last user message")
return False
else:
print(f" --> BUG: expected 'What's the weather in NYC?', got {repr(message_content)}")
return True
def test_type_annotation_mismatch():
"""
Shows that function signatures still declare messages as Dict[str, str]
even though messages with tool_calls have non-string values.
This is a type-hint issue (not a runtime error), but it causes static
analysis tools and IDEs to flag valid tool-call messages as type errors.
"""
import inspect
# Check the auto_router function signature
from litellm.integrations.custom_logger import CustomLogger
sig = inspect.signature(CustomLogger.async_pre_routing_hook)
messages_param = sig.parameters["messages"]
annotation = str(messages_param.annotation)
print(f"\n[INFO] Type annotation check:")
print(f" CustomLogger.async_pre_routing_hook 'messages' param: {annotation}")
if "Dict[str, str]" in annotation:
print(" --> ISSUE: Still uses Dict[str, str], should be Dict[str, Any]")
print(" This means tool_call messages (with list/dict values) violate the type contract")
return True
else:
print(" --> OK: Uses Dict[str, Any]")
return False
def test_multimodal_content_fixed():
"""
Verifies that multimodal messages (content as a list of blocks) are
handled correctly by _extract_text_from_messages.
"""
from litellm.router_strategy.auto_router.auto_router import AutoRouter
messages_with_image = [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/image.png"},
},
],
}
]
message_content = AutoRouter._extract_text_from_messages(messages_with_image)
print(f"\n[CHECK] Content extraction from multimodal message:")
print(f" type(message_content) = {type(message_content).__name__}")
print(f" message_content = {repr(message_content)}")
if isinstance(message_content, str) and message_content == "What's in this image?":
print(" --> FIXED: correctly extracts text from multimodal content")
return False
else:
print(f" --> BUG: expected string, got {type(message_content).__name__}")
return True
def main():
print("=" * 70)
print("Reproduction for GitHub Issue #14633")
print("AutoRouter fails with tool call messages")
print("=" * 70)
bugs_found = 0
# Test 1: PreRoutingHookResponse type (partially fixed)
test_pre_routing_hook_response_with_tool_calls()
# Test 2: Content extraction fix (None content)
if test_content_extraction_fixed():
bugs_found += 1
# Test 3: Type annotation fix
if test_type_annotation_mismatch():
bugs_found += 1
# Test 4: Multimodal content fix
if test_multimodal_content_fixed():
bugs_found += 1
print("\n" + "=" * 70)
if bugs_found == 0:
print("ALL CHECKS PASSED - all bugs have been fixed!")
else:
print(f"Results: {bugs_found} bug(s) remaining")
print("=" * 70)
sys.exit(1 if bugs_found > 0 else 0)
if __name__ == "__main__":
main()

View file

@ -12,7 +12,146 @@ sys.path.insert(
from litellm.router_strategy.auto_router.auto_router import AutoRouter
pytestmark = pytest.mark.skip(reason="Skipping auto router tests - beta feature")
pytestmark_skip_beta = pytest.mark.skip(reason="Skipping auto router tests - beta feature")
class TestExtractTextFromMessages:
"""Tests for AutoRouter._extract_text_from_messages (no semantic_router dependency)."""
def test_should_extract_content_from_simple_user_message(self):
messages = [{"role": "user", "content": "Hello world"}]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "Hello world"
def test_should_extract_last_user_message_from_tool_call_conversation(self):
messages = [
{"role": "user", "content": "What's the weather in NYC?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "72°F and sunny",
},
{"role": "user", "content": "Now tell me about London"},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "Now tell me about London"
def test_should_find_user_message_when_last_message_is_assistant_with_tool_calls(
self,
):
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's the weather?"
def test_should_find_user_message_when_last_message_is_tool_response(self):
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc",
"content": "72°F and sunny",
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's the weather?"
def test_should_handle_multimodal_content_list(self):
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/img.png"},
},
],
}
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's in this image?"
def test_should_handle_multimodal_content_with_multiple_text_blocks(self):
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "First part"},
{"type": "text", "text": "Second part"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/img.png"},
},
],
}
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "First part Second part"
def test_should_return_empty_string_when_user_content_is_none(self):
messages = [{"role": "user", "content": None}]
result = AutoRouter._extract_text_from_messages(messages)
assert result == ""
def test_should_return_empty_string_when_no_user_messages(self):
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == ""
def test_should_return_empty_string_for_empty_messages_list(self):
result = AutoRouter._extract_text_from_messages([])
assert result == ""
@pytest.fixture
@ -41,6 +180,7 @@ def mock_route_choice():
return mock_choice
@pytestmark_skip_beta
class TestAutoRouter:
"""Test class for AutoRouter methods."""

4
uv.lock generated
View file

@ -11,7 +11,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-04-08T16:01:27.663665Z"
exclude-newer = "2026-04-13T02:23:17.47714Z"
exclude-newer-span = "P3D"
[manifest]
@ -3602,7 +3602,7 @@ wheels = [
[[package]]
name = "litellm"
version = "1.83.6"
version = "1.83.8"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },