fix(complexity_router.py): fix /v1/responses routing for complexity router

This commit is contained in:
Krrish Dholakia 2026-04-15 18:51:10 -07:00
parent 1cc387bc6c
commit 2b816c33f0
8 changed files with 137 additions and 63 deletions

View file

@ -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

View file

@ -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.

View file

@ -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.

View file

@ -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

View file

@ -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)

View file

@ -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"},
},
],
}
]

View file

@ -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}"

View file

@ -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"},
},
],
}
]