mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(complexity_router.py): fix /v1/responses routing for complexity router
This commit is contained in:
parent
1cc387bc6c
commit
2b816c33f0
8 changed files with 137 additions and 63 deletions
|
|
@ -78,9 +78,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
return chat_completion_compatible_request
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert Anthropic messages request data to OpenAI-spec structured messages.
|
||||
|
||||
|
|
@ -122,9 +120,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
tools_to_check: List[
|
||||
ChatCompletionToolParam
|
||||
] = chat_completion_compatible_request.get("tools", [])
|
||||
tools_to_check: List[ChatCompletionToolParam] = (
|
||||
chat_completion_compatible_request.get("tools", [])
|
||||
)
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
|
|
|
|||
|
|
@ -102,9 +102,7 @@ class BaseTranslation(ABC):
|
|||
"""
|
||||
return responses_so_far
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List["AllMessageValues"]]:
|
||||
def get_structured_messages(self, data: dict) -> Optional[List["AllMessageValues"]]:
|
||||
"""
|
||||
Convert request data to OpenAI-spec structured messages.
|
||||
|
||||
|
|
|
|||
|
|
@ -48,9 +48,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert chat completions request data to OpenAI-spec structured messages.
|
||||
|
||||
|
|
|
|||
|
|
@ -71,9 +71,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert Responses API request data to OpenAI-spec structured messages.
|
||||
|
||||
|
|
@ -83,9 +81,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
input_data = data.get("input")
|
||||
if input_data is None:
|
||||
return None
|
||||
messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_data,
|
||||
responses_api_request=data,
|
||||
messages = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_data,
|
||||
responses_api_request=data,
|
||||
)
|
||||
)
|
||||
return cast(List[AllMessageValues], messages) if messages else None
|
||||
|
||||
|
|
|
|||
|
|
@ -356,10 +356,8 @@ class ComplexityRouter(CustomLogger):
|
|||
call_type: Optional[CallTypes] = None
|
||||
|
||||
# 1. Try route-based inference from proxy metadata
|
||||
route = (
|
||||
request_kwargs.get("litellm_metadata", {}).get(
|
||||
"user_api_key_request_route"
|
||||
)
|
||||
route = request_kwargs.get("litellm_metadata", {}).get(
|
||||
"user_api_key_request_route"
|
||||
)
|
||||
if route:
|
||||
call_types_list = get_call_types_for_route(route)
|
||||
|
|
|
|||
|
|
@ -134,7 +134,9 @@ class TestOpenAIChatCompletionsHandlerToolsInput:
|
|||
tool = guardrail.last_inputs["tools"][0]
|
||||
assert tool["type"] == "function"
|
||||
assert tool["function"]["name"] == "get_weather"
|
||||
assert tool["function"]["description"] == "Get the current weather in a location"
|
||||
assert (
|
||||
tool["function"]["description"] == "Get the current weather in a location"
|
||||
)
|
||||
assert "parameters" in tool["function"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -189,7 +191,10 @@ class TestOpenAIChatCompletionsHandlerToolsInput:
|
|||
|
||||
assert guardrail.last_inputs is not None
|
||||
# tools should not be in inputs if not provided
|
||||
assert "tools" not in guardrail.last_inputs or guardrail.last_inputs.get("tools") is None
|
||||
assert (
|
||||
"tools" not in guardrail.last_inputs
|
||||
or guardrail.last_inputs.get("tools") is None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tools_and_tool_calls_both_passed(self):
|
||||
|
|
@ -220,7 +225,10 @@ class TestOpenAIChatCompletionsHandlerToolsInput:
|
|||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"location": {"type": "string"}}},
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
@ -833,7 +841,9 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
assert result == responses_so_far
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_output_streaming_response_mixed_empty_and_valid_choices_no_finish(self):
|
||||
async def test_process_output_streaming_response_mixed_empty_and_valid_choices_no_finish(
|
||||
self,
|
||||
):
|
||||
"""Test streaming response with mix of empty and valid choices chunks (stream not finished)
|
||||
|
||||
This tests the has_stream_ended check when iterating through chunks with mixed choices.
|
||||
|
|
@ -922,7 +932,10 @@ class TestGetStructuredMessages:
|
|||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What's in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/image.png"}},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1038,8 +1038,7 @@ class TestGetStructuredMessages:
|
|||
result = handler.get_structured_messages(data)
|
||||
assert result is not None
|
||||
has_system = any(
|
||||
isinstance(msg, dict) and msg.get("role") == "system"
|
||||
for msg in result
|
||||
isinstance(msg, dict) and msg.get("role") == "system" for msg in result
|
||||
)
|
||||
assert has_system, f"Expected system message from instructions, got: {result}"
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Tests for the ComplexityRouter.
|
|||
|
||||
Tests the rule-based complexity scoring and tier assignment logic.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from typing import Dict, List
|
||||
|
|
@ -123,7 +124,9 @@ class TestTokenScoring:
|
|||
tier, score, signals = complexity_router.classify("What is Python?")
|
||||
# Should be classified as SIMPLE due to short length and simple indicator
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
assert any("short" in s.lower() for s in signals) or any("simple" in s.lower() for s in signals)
|
||||
assert any("short" in s.lower() for s in signals) or any(
|
||||
"simple" in s.lower() for s in signals
|
||||
)
|
||||
|
||||
def test_long_prompt_positive_score(self, complexity_router):
|
||||
"""Long prompts should get positive scores (complex indicator)."""
|
||||
|
|
@ -134,7 +137,9 @@ class TestTokenScoring:
|
|||
tier, score, signals = complexity_router.classify(long_prompt)
|
||||
# Should have positive score and detect long token count or technical terms
|
||||
assert score > 0, f"Expected positive score for long prompt, got {score}"
|
||||
assert any("long" in s.lower() for s in signals) or any("technical" in s.lower() for s in signals)
|
||||
assert any("long" in s.lower() for s in signals) or any(
|
||||
"technical" in s.lower() for s in signals
|
||||
)
|
||||
|
||||
|
||||
class TestCodePresenceScoring:
|
||||
|
|
@ -209,7 +214,9 @@ class TestMultiStepPatterns:
|
|||
|
||||
def test_first_then_pattern(self, complexity_router):
|
||||
"""'First...then' patterns should increase complexity."""
|
||||
prompt = "First analyze the data, then create a visualization, then write a report"
|
||||
prompt = (
|
||||
"First analyze the data, then create a visualization, then write a report"
|
||||
)
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert any("multi-step" in s.lower() for s in signals)
|
||||
|
||||
|
|
@ -253,7 +260,9 @@ class TestTierAssignment:
|
|||
)
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# Should detect technical terms
|
||||
assert any("technical" in s.lower() for s in signals), f"Expected technical signals, got {signals}"
|
||||
assert any(
|
||||
"technical" in s.lower() for s in signals
|
||||
), f"Expected technical signals, got {signals}"
|
||||
# Score should be positive due to technical content
|
||||
assert score > 0, f"Expected positive score, got {score}"
|
||||
|
||||
|
|
@ -321,12 +330,15 @@ class TestPreRoutingHook:
|
|||
async def test_pre_routing_hook_complex_message(self, complexity_router):
|
||||
"""Test pre-routing hook with a message containing technical content."""
|
||||
messages = [
|
||||
{"role": "user", "content": (
|
||||
"Design a distributed microservice architecture with Kubernetes "
|
||||
"orchestration, implementing proper authentication, encryption, "
|
||||
"and database optimization for high throughput. Think step by step "
|
||||
"about the performance implications and scalability requirements."
|
||||
)}
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"Design a distributed microservice architecture with Kubernetes "
|
||||
"orchestration, implementing proper authentication, encryption, "
|
||||
"and database optimization for high throughput. Think step by step "
|
||||
"about the performance implications and scalability requirements."
|
||||
),
|
||||
}
|
||||
]
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
|
|
@ -335,7 +347,12 @@ class TestPreRoutingHook:
|
|||
)
|
||||
assert result is not None
|
||||
# Should return a valid model from the configured tiers
|
||||
assert result.model in ["gpt-4o-mini", "gpt-4o", "claude-sonnet-4-20250514", "o1-preview"]
|
||||
assert result.model in [
|
||||
"gpt-4o-mini",
|
||||
"gpt-4o",
|
||||
"claude-sonnet-4-20250514",
|
||||
"o1-preview",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_routing_hook_no_messages(self, complexity_router):
|
||||
|
|
@ -377,7 +394,10 @@ class TestPreRoutingHook:
|
|||
async def test_pre_routing_hook_reasoning_message(self, complexity_router):
|
||||
"""Test pre-routing hook with reasoning markers."""
|
||||
messages = [
|
||||
{"role": "user", "content": "Let's think step by step and reason through this problem carefully."}
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Let's think step by step and reason through this problem carefully.",
|
||||
}
|
||||
]
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
|
|
@ -416,7 +436,9 @@ class TestConfigOverrides:
|
|||
"Explain how HTTP works with REST APIs and distributed systems"
|
||||
)
|
||||
# With boundaries this low, should be at least MEDIUM (anything above -0.5)
|
||||
assert tier != ComplexityTier.SIMPLE, f"Expected non-SIMPLE tier, got {tier} with score {score}"
|
||||
assert (
|
||||
tier != ComplexityTier.SIMPLE
|
||||
), f"Expected non-SIMPLE tier, got {tier} with score {score}"
|
||||
|
||||
def test_custom_token_thresholds(self, mock_router_instance):
|
||||
"""Test custom token thresholds work correctly."""
|
||||
|
|
@ -441,7 +463,9 @@ class TestConfigOverrides:
|
|||
long_prompt = "This is a test prompt " * 30 # ~120 tokens
|
||||
tier, score, signals = router.classify(long_prompt)
|
||||
# Should get token length signal indicating "long"
|
||||
assert any("long" in s.lower() if s else False for s in signals), f"Expected 'long' signal, got {signals}"
|
||||
assert any(
|
||||
"long" in s.lower() if s else False for s in signals
|
||||
), f"Expected 'long' signal, got {signals}"
|
||||
|
||||
|
||||
class TestAsyncPreRoutingHookEdgeCases:
|
||||
|
|
@ -468,9 +492,15 @@ class TestAsyncPreRoutingHookEdgeCases:
|
|||
"""Test pre-routing hook uses the last user message for classification."""
|
||||
# Multiple user messages - should classify based on the LAST one
|
||||
messages = [
|
||||
{"role": "user", "content": "Design a complex distributed system"}, # Complex prompt
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Design a complex distributed system",
|
||||
}, # Complex prompt
|
||||
{"role": "assistant", "content": "I can help with that."},
|
||||
{"role": "user", "content": "Hello!"}, # Simple prompt - this should be used
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello!",
|
||||
}, # Simple prompt - this should be used
|
||||
]
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
|
|
@ -496,13 +526,21 @@ class TestAsyncPreRoutingHookEdgeCases:
|
|||
# Should return default model rather than None (None would cause
|
||||
# the complexity_router deployment itself to be selected, crashing)
|
||||
assert result is not None
|
||||
assert result.model in ["gpt-4o-mini", "gpt-4o", "claude-sonnet-4-20250514", "o1-preview"]
|
||||
assert result.model in [
|
||||
"gpt-4o-mini",
|
||||
"gpt-4o",
|
||||
"claude-sonnet-4-20250514",
|
||||
"o1-preview",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_routing_hook_list_content(self, complexity_router):
|
||||
"""Test pre-routing hook handles list-format message content (OpenAI multi-part format)."""
|
||||
messages = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "Hello, how are you?"}]},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "text", "text": "Hello, how are you?"}],
|
||||
},
|
||||
]
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
|
|
@ -520,8 +558,14 @@ class TestAsyncPreRoutingHookEdgeCases:
|
|||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Think step by step and reason through this: design a distributed system"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,abc"}},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Think step by step and reason through this: design a distributed system",
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,abc"},
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
|
@ -561,7 +605,12 @@ class TestAsyncPreRoutingHookEdgeCases:
|
|||
)
|
||||
# Empty string content → no extractable user message → routes to default model
|
||||
assert result is not None
|
||||
assert result.model in ["gpt-4o-mini", "gpt-4o", "claude-sonnet-4-20250514", "o1-preview"]
|
||||
assert result.model in [
|
||||
"gpt-4o-mini",
|
||||
"gpt-4o",
|
||||
"claude-sonnet-4-20250514",
|
||||
"o1-preview",
|
||||
]
|
||||
|
||||
|
||||
class TestSingletonMutation:
|
||||
|
|
@ -575,7 +624,7 @@ class TestSingletonMutation:
|
|||
|
||||
# Get original default
|
||||
original_default = ComplexityRouterConfig().default_model
|
||||
|
||||
|
||||
# Create router with empty config and custom default_model
|
||||
router1 = ComplexityRouter(
|
||||
model_name="test-router-1",
|
||||
|
|
@ -583,14 +632,14 @@ class TestSingletonMutation:
|
|||
complexity_router_config=None,
|
||||
default_model="custom-fallback",
|
||||
)
|
||||
|
||||
|
||||
# Create another router without config
|
||||
router2 = ComplexityRouter(
|
||||
model_name="test-router-2",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=None,
|
||||
)
|
||||
|
||||
|
||||
# Router2 should have fresh defaults, not router1's custom default_model
|
||||
# Create a fresh config to check
|
||||
fresh_config = ComplexityRouterConfig()
|
||||
|
|
@ -608,7 +657,9 @@ class TestKeywordFalsePositives:
|
|||
prompt = "What is the capital of France?"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# Should NOT detect code presence from 'api' in 'capital'
|
||||
assert not any("code" in s.lower() for s in signals), f"False positive: got code signal from 'capital'"
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'capital'"
|
||||
# Should be SIMPLE (definition question)
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
|
|
@ -617,7 +668,9 @@ class TestKeywordFalsePositives:
|
|||
prompt = "Explain digital marketing strategies"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# Should NOT detect code presence from 'git' in 'digital'
|
||||
assert not any("code" in s.lower() for s in signals), f"False positive: got code signal from 'digital'"
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'digital'"
|
||||
|
||||
def test_try_not_in_entry(self, complexity_router):
|
||||
"""'try' should not match in 'entry'."""
|
||||
|
|
@ -631,33 +684,43 @@ class TestKeywordFalsePositives:
|
|||
"""'error' should not match in 'terrorism'."""
|
||||
prompt = "The country is dealing with terrorism"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any("code" in s.lower() for s in signals), f"False positive: got code signal from 'terrorism'"
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'terrorism'"
|
||||
|
||||
def test_class_not_in_classical(self, complexity_router):
|
||||
"""'class' should not match in 'classical'."""
|
||||
prompt = "I enjoy listening to classical music"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any("code" in s.lower() for s in signals), f"False positive: got code signal from 'classical'"
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'classical'"
|
||||
|
||||
def test_merge_not_in_emerged(self, complexity_router):
|
||||
"""'merge' should not match in 'emerged'."""
|
||||
prompt = "A new leader emerged from the crowd"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
assert not any("code" in s.lower() for s in signals), f"False positive: got code signal from 'emerged'"
|
||||
assert not any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"False positive: got code signal from 'emerged'"
|
||||
|
||||
def test_actual_api_keyword_detected(self, complexity_router):
|
||||
"""Actual 'api' usage should be detected."""
|
||||
prompt = "How do I call the REST api endpoint?"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# Should detect code presence from actual 'api' usage
|
||||
assert any("code" in s.lower() for s in signals), f"Expected code signal for 'api', got {signals}"
|
||||
assert any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"Expected code signal for 'api', got {signals}"
|
||||
|
||||
def test_actual_git_keyword_detected(self, complexity_router):
|
||||
"""Actual 'git' usage should be detected."""
|
||||
prompt = "How do I use git to commit changes?"
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# Should detect code presence from actual 'git' usage
|
||||
assert any("code" in s.lower() for s in signals), f"Expected code signal for 'git', got {signals}"
|
||||
assert any(
|
||||
"code" in s.lower() for s in signals
|
||||
), f"Expected code signal for 'git', got {signals}"
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
|
|
@ -677,7 +740,9 @@ class TestEdgeCases:
|
|||
# Should have positive score due to length
|
||||
assert score > 0, f"Expected positive score for very long prompt, got {score}"
|
||||
# Should detect long token count
|
||||
assert any("long" in s.lower() for s in signals), f"Expected 'long' signal, got {signals}"
|
||||
assert any(
|
||||
"long" in s.lower() for s in signals
|
||||
), f"Expected 'long' signal, got {signals}"
|
||||
|
||||
def test_unicode_prompt(self, complexity_router):
|
||||
"""Test handling of unicode characters."""
|
||||
|
|
@ -695,7 +760,9 @@ class TestEdgeCases:
|
|||
"""
|
||||
tier, score, signals = complexity_router.classify(prompt)
|
||||
# The "step N" pattern should be detected
|
||||
assert any("multi-step" in s.lower() for s in signals), f"Expected multi-step signal, got {signals}"
|
||||
assert any(
|
||||
"multi-step" in s.lower() for s in signals
|
||||
), f"Expected multi-step signal, got {signals}"
|
||||
|
||||
|
||||
class TestRouterComplexityDeploymentMethods:
|
||||
|
|
@ -960,7 +1027,10 @@ class TestExtractUserMessageAndSystemPrompt:
|
|||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Describe this image"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/img.png"}},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/img.png"},
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue