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 auto router
This commit is contained in:
parent
72a461ba4a
commit
1cc387bc6c
43 changed files with 524 additions and 55 deletions
|
|
@ -34,6 +34,7 @@ from litellm.types.llms.anthropic import (
|
|||
)
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionRequest,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
|
|
@ -67,6 +68,34 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
super().__init__()
|
||||
self.adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
def _translate_to_openai(self, data: dict) -> ChatCompletionRequest:
|
||||
"""Translate Anthropic request to OpenAI chat completion format."""
|
||||
(
|
||||
chat_completion_compatible_request,
|
||||
_tool_name_mapping,
|
||||
) = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
|
||||
)
|
||||
return chat_completion_compatible_request
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert Anthropic messages request data to OpenAI-spec structured messages.
|
||||
|
||||
Uses the Anthropic-to-OpenAI adapter to translate message format.
|
||||
"""
|
||||
messages = data.get("messages")
|
||||
if messages is None:
|
||||
return None
|
||||
chat_completion_compatible_request = self._translate_to_openai(data)
|
||||
result = cast(
|
||||
List[AllMessageValues],
|
||||
chat_completion_compatible_request.get("messages", []),
|
||||
)
|
||||
return result if result else None
|
||||
|
||||
async def process_input_messages(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -82,13 +111,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
|
||||
(
|
||||
chat_completion_compatible_request,
|
||||
_tool_name_mapping,
|
||||
) = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
# Use a shallow copy to avoid mutating request data (pop on litellm_metadata).
|
||||
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
|
||||
)
|
||||
chat_completion_compatible_request = self._translate_to_openai(data)
|
||||
|
||||
structured_messages = cast(
|
||||
List[AllMessageValues],
|
||||
|
|
@ -103,8 +126,6 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
ChatCompletionToolParam
|
||||
] = chat_completion_compatible_request.get("tools", [])
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (message_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
for msg_idx, message in enumerate(messages):
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ if TYPE_CHECKING:
|
|||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
class BaseTranslation(ABC):
|
||||
|
|
@ -101,6 +102,18 @@ class BaseTranslation(ABC):
|
|||
"""
|
||||
return responses_so_far
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List["AllMessageValues"]]:
|
||||
"""
|
||||
Convert request data to OpenAI-spec structured messages.
|
||||
|
||||
Override in subclasses for format-specific conversion.
|
||||
|
||||
Returns None if no convertible content is found.
|
||||
"""
|
||||
return None
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""
|
||||
Extract tool names from the request body for allowlist/policy checks.
|
||||
|
|
|
|||
|
|
@ -48,6 +48,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert chat completions request data to OpenAI-spec structured messages.
|
||||
|
||||
Messages are already in OpenAI format, so this is a simple extraction.
|
||||
"""
|
||||
messages = data.get("messages")
|
||||
if messages is None:
|
||||
return None
|
||||
return cast(List[AllMessageValues], messages)
|
||||
|
||||
async def process_input_messages(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -68,9 +81,6 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
tool_calls_to_check: List[ChatCompletionToolParam] = []
|
||||
text_task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
tool_call_task_mappings: List[Tuple[int, int]] = []
|
||||
# text_task_mappings: Track (message_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
# tool_call_task_mappings: Track (message_index, tool_call_index) for each tool call
|
||||
|
||||
# Step 1: Extract all text content, images, and tool calls
|
||||
for msg_idx, message in enumerate(messages):
|
||||
|
|
@ -92,12 +102,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check # type: ignore
|
||||
if messages:
|
||||
msg_list = cast(List[AllMessageValues], messages)
|
||||
structured_messages = self.get_structured_messages(data)
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = (
|
||||
openai_messages_without_system(msg_list)
|
||||
openai_messages_without_system(structured_messages)
|
||||
if skip_system
|
||||
else msg_list
|
||||
else structured_messages
|
||||
)
|
||||
# Pass tools (function definitions) to the guardrail
|
||||
tools = data.get("tools")
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ from litellm.responses.litellm_completion_transformation.transformation import (
|
|||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
|
|
@ -70,6 +71,24 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
def get_structured_messages(
|
||||
self, data: dict
|
||||
) -> Optional[List[AllMessageValues]]:
|
||||
"""
|
||||
Convert Responses API request data to OpenAI-spec structured messages.
|
||||
|
||||
Transforms `input` (string or ResponseInputParam) and optional
|
||||
`instructions` into chat completion messages.
|
||||
"""
|
||||
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,
|
||||
)
|
||||
return cast(List[AllMessageValues], messages) if messages else None
|
||||
|
||||
async def process_input_messages(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -86,12 +105,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
if input_data is None:
|
||||
return data
|
||||
|
||||
structured_messages = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_data,
|
||||
responses_api_request=data,
|
||||
)
|
||||
)
|
||||
structured_messages = self.get_structured_messages(data)
|
||||
|
||||
# Handle simple string input
|
||||
if isinstance(input_data, str):
|
||||
|
|
|
|||
|
|
@ -22,11 +22,21 @@ model_list:
|
|||
output_cost_per_token: 10 # 100x standard ($10.00/1M = $0.00001)
|
||||
|
||||
# Anthropic model for /v1/messages test — 100x custom pricing
|
||||
- model_name: "claude-sonnet-4-20250514"
|
||||
- model_name: "claude-sonnet-4-6"
|
||||
litellm_params:
|
||||
model: anthropic/claude-sonnet-4-20250514
|
||||
model: anthropic/claude-sonnet-4-6
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
model_info:
|
||||
id: claude-sonnet-4-custom-pricing
|
||||
input_cost_per_token: 0.0003 # 100x standard ($0.000003)
|
||||
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
|
||||
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
|
||||
- model_name: my-auto
|
||||
litellm_params:
|
||||
model: auto_router/complexity_router
|
||||
complexity_router_config:
|
||||
tiers:
|
||||
SIMPLE: "gpt-4.1-mini"
|
||||
COMPLEX: claude-sonnet-4-6
|
||||
tier_boundaries:
|
||||
simple_medium: 0.30
|
||||
complexity_router_default_model: small-model
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ No external API calls - all scoring is local and <1ms.
|
|||
|
||||
Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
|
|
@ -331,45 +332,73 @@ class ComplexityRouter(CustomLogger):
|
|||
f"No model configured for tier {tier_key} and no default_model set"
|
||||
)
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
def _resolve_messages(
|
||||
self,
|
||||
model: str,
|
||||
messages: Optional[List[Dict[str, Any]]],
|
||||
request_kwargs: Dict,
|
||||
messages: Optional[List[Dict[str, Any]]] = None,
|
||||
input: Optional[Union[str, List]] = None,
|
||||
specific_deployment: Optional[bool] = False,
|
||||
) -> Optional["PreRoutingHookResponse"]:
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""
|
||||
Pre-routing hook called before the routing decision.
|
||||
Resolve messages from the request, converting from other formats if needed.
|
||||
|
||||
Classifies the request by complexity and returns the appropriate model.
|
||||
|
||||
Args:
|
||||
model: The original model name requested.
|
||||
request_kwargs: The request kwargs.
|
||||
messages: The messages in the request.
|
||||
input: Optional input for embeddings.
|
||||
specific_deployment: Whether a specific deployment was requested.
|
||||
|
||||
Returns:
|
||||
PreRoutingHookResponse with the routed model, or None if no routing needed.
|
||||
Uses the guardrail translation handler dispatch to convert Responses API
|
||||
``input`` (or other non-chat-completions formats) into OpenAI-spec messages.
|
||||
"""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
if messages:
|
||||
return messages
|
||||
|
||||
if messages is None or len(messages) == 0:
|
||||
verbose_router_logger.debug(
|
||||
"ComplexityRouter: No messages provided, skipping routing"
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import (
|
||||
get_call_types_for_route,
|
||||
)
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
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"
|
||||
)
|
||||
)
|
||||
if route:
|
||||
call_types_list = get_call_types_for_route(route)
|
||||
if call_types_list:
|
||||
for ct in call_types_list:
|
||||
if ct in mappings:
|
||||
call_type = ct
|
||||
break
|
||||
|
||||
# 2. Fall back to data-shape inference
|
||||
if call_type is None:
|
||||
if "input" in request_kwargs:
|
||||
call_type = CallTypes.responses
|
||||
else:
|
||||
return None
|
||||
|
||||
if call_type not in mappings:
|
||||
return None
|
||||
|
||||
# Extract the last user message and the last system prompt
|
||||
handler = mappings[call_type]()
|
||||
structured = handler.get_structured_messages(request_kwargs)
|
||||
if structured:
|
||||
return [
|
||||
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
|
||||
for msg in structured
|
||||
]
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_user_message_and_system_prompt(
|
||||
messages: List[Dict[str, Any]],
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Extract the last user message text and last system prompt from messages."""
|
||||
user_message: Optional[str] = None
|
||||
system_prompt: Optional[str] = None
|
||||
|
||||
for msg in reversed(messages):
|
||||
role = msg.get("role", "")
|
||||
content = msg.get("content") or ""
|
||||
# content may be a list of content parts (e.g. [{"type": "text", "text": "..."}])
|
||||
if isinstance(content, list):
|
||||
text_parts = [
|
||||
part.get("text", "")
|
||||
|
|
@ -382,6 +411,52 @@ class ComplexityRouter(CustomLogger):
|
|||
user_message = content
|
||||
elif role == "system" and system_prompt is None:
|
||||
system_prompt = content
|
||||
if user_message is not None and system_prompt is not None:
|
||||
break
|
||||
|
||||
return user_message, system_prompt
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: Dict,
|
||||
messages: Optional[List[Dict[str, Any]]] = None,
|
||||
input: Optional[Union[str, List]] = None,
|
||||
specific_deployment: Optional[bool] = False,
|
||||
) -> Optional["PreRoutingHookResponse"]:
|
||||
"""
|
||||
Pre-routing hook called before the routing decision.
|
||||
|
||||
Classifies the request by complexity and returns the appropriate model.
|
||||
Supports chat completions (messages), Responses API (input), and other
|
||||
formats via the guardrail translation handler dispatch.
|
||||
|
||||
Args:
|
||||
model: The original model name requested.
|
||||
request_kwargs: The request kwargs.
|
||||
messages: The messages in the request.
|
||||
input: Optional input for Responses API or embeddings.
|
||||
specific_deployment: Whether a specific deployment was requested.
|
||||
|
||||
Returns:
|
||||
PreRoutingHookResponse with the routed model, or None if no routing needed.
|
||||
"""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
resolved_messages = self._resolve_messages(messages, request_kwargs)
|
||||
|
||||
if not resolved_messages:
|
||||
verbose_router_logger.debug(
|
||||
"ComplexityRouter: No messages could be resolved, skipping routing"
|
||||
)
|
||||
return None
|
||||
|
||||
# Determine whether the original request used messages directly
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
|
||||
user_message, system_prompt = self._extract_user_message_and_system_prompt(
|
||||
resolved_messages
|
||||
)
|
||||
|
||||
if user_message is None:
|
||||
verbose_router_logger.debug(
|
||||
|
|
@ -390,13 +465,10 @@ class ComplexityRouter(CustomLogger):
|
|||
return PreRoutingHookResponse(
|
||||
model=self.config.default_model
|
||||
or self.get_model_for_tier(ComplexityTier.MEDIUM),
|
||||
messages=messages,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
# Classify the request
|
||||
tier, score, signals = self.classify(user_message, system_prompt)
|
||||
|
||||
# Get the model for this tier
|
||||
routed_model = self.get_model_for_tier(tier)
|
||||
|
||||
verbose_router_logger.info(
|
||||
|
|
@ -406,5 +478,5 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -881,6 +881,58 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
assert result == responses_so_far
|
||||
|
||||
|
||||
class TestGetStructuredMessages:
|
||||
"""Test the get_structured_messages method."""
|
||||
|
||||
def test_should_return_messages_from_chat_completions_request(self):
|
||||
"""Test that messages are returned from a chat completions request."""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are helpful."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is not None
|
||||
assert len(result) == 2
|
||||
assert result[0]["role"] == "system"
|
||||
assert result[1]["role"] == "user"
|
||||
|
||||
def test_should_return_none_when_no_messages(self):
|
||||
"""Test that None is returned when no messages key exists."""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
data = {"model": "gpt-4"}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is None
|
||||
|
||||
def test_should_return_none_for_none_messages(self):
|
||||
"""Test that None is returned when messages is explicitly None."""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
data = {"messages": None}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is None
|
||||
|
||||
def test_should_handle_multimodal_content(self):
|
||||
"""Test that messages with multimodal content are returned."""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What's in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/image.png"}},
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is not None
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0]["content"], list)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run the tests
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
|
|||
|
|
@ -995,3 +995,64 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
|
|||
|
||||
# Should return the responses
|
||||
assert result == responses_so_far
|
||||
|
||||
|
||||
class TestGetStructuredMessages:
|
||||
"""Test the get_structured_messages method for Responses API handler."""
|
||||
|
||||
def test_should_convert_string_input_to_messages(self):
|
||||
"""Test that a simple string input is converted to OpenAI messages."""
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {"input": "What is the capital of France?"}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is not None
|
||||
assert len(result) >= 1
|
||||
found_user = False
|
||||
for msg in result:
|
||||
if isinstance(msg, dict) and msg.get("role") == "user":
|
||||
found_user = True
|
||||
break
|
||||
assert found_user, f"Expected a user message, got: {result}"
|
||||
|
||||
def test_should_convert_list_input_to_messages(self):
|
||||
"""Test that list input (ResponseInputParam) is converted to OpenAI messages."""
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"input": [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi there!"},
|
||||
{"role": "user", "content": "How are you?"},
|
||||
]
|
||||
}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is not None
|
||||
assert len(result) >= 3
|
||||
|
||||
def test_should_include_instructions_as_system_message(self):
|
||||
"""Test that instructions are included as a system message."""
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"input": "Roll a d20",
|
||||
"instructions": "You are a helpful dungeon master.",
|
||||
}
|
||||
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
|
||||
)
|
||||
assert has_system, f"Expected system message from instructions, got: {result}"
|
||||
|
||||
def test_should_return_none_when_no_input(self):
|
||||
"""Test that None is returned when input key is missing."""
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {"model": "gpt-4o"}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is None
|
||||
|
||||
def test_should_return_none_for_none_input(self):
|
||||
"""Test that None is returned when input is explicitly None."""
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {"input": None}
|
||||
result = handler.get_structured_messages(data)
|
||||
assert result is None
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Tests the rule-based complexity scoring and tier assignment logic.
|
|||
import os
|
||||
import sys
|
||||
from typing import Dict, List
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -761,3 +761,219 @@ class TestRouterComplexityDeploymentMethods:
|
|||
)
|
||||
router.init_complexity_router_deployment(deployment)
|
||||
assert "auto_router/complexity_router/test-router" in router.complexity_routers
|
||||
|
||||
|
||||
class TestAsyncPreRoutingHookMultiFormat:
|
||||
"""Test async_pre_routing_hook with multiple input formats."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_route_with_chat_completions_messages(self, complexity_router):
|
||||
"""Test routing with standard chat completions messages."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model is not None
|
||||
assert result.messages is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_route_with_responses_api_string_input(
|
||||
self, complexity_router
|
||||
):
|
||||
"""Test routing with Responses API string input via handler dispatch."""
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.load_guardrail_translation_mappings",
|
||||
return_value=mock_mappings,
|
||||
):
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={"input": "What is the capital of France?"},
|
||||
messages=None,
|
||||
input="What is the capital of France?",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.model is not None
|
||||
# messages should be None since the original request didn't have messages
|
||||
assert result.messages is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_route_with_responses_api_list_input(self, complexity_router):
|
||||
"""Test routing with Responses API list input via handler dispatch."""
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
||||
|
||||
list_input = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi there!"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Write a Python function to sort a list using merge sort",
|
||||
},
|
||||
]
|
||||
|
||||
with patch(
|
||||
"litellm.llms.load_guardrail_translation_mappings",
|
||||
return_value=mock_mappings,
|
||||
):
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={"input": list_input},
|
||||
messages=None,
|
||||
input=list_input,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.model is not None
|
||||
assert result.messages is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_use_route_based_inference(self, complexity_router):
|
||||
"""Test that route-based call type inference is used when available."""
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.load_guardrail_translation_mappings",
|
||||
return_value=mock_mappings,
|
||||
):
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={
|
||||
"input": "Roll 2d4+1",
|
||||
"litellm_metadata": {
|
||||
"user_api_key_request_route": "/v1/responses",
|
||||
},
|
||||
},
|
||||
messages=None,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.model is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_return_none_when_no_messages_or_input(
|
||||
self, complexity_router
|
||||
):
|
||||
"""Test that None is returned when neither messages nor input is available."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=None,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_prefer_original_messages_over_conversion(
|
||||
self, complexity_router
|
||||
):
|
||||
"""Test that original messages are used when both messages and input are available."""
|
||||
messages = [{"role": "user", "content": "What is 2+2?"}]
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={"input": "This should be ignored"},
|
||||
messages=messages,
|
||||
)
|
||||
assert result is not None
|
||||
assert result.messages == messages
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_include_instructions_in_classification(
|
||||
self, complexity_router
|
||||
):
|
||||
"""Test that Responses API instructions influence classification via system message."""
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.load_guardrail_translation_mappings",
|
||||
return_value=mock_mappings,
|
||||
):
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={
|
||||
"input": "Write merge sort",
|
||||
"instructions": "You are an expert Python developer. Use advanced algorithms and optimize for performance.",
|
||||
},
|
||||
messages=None,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.model is not None
|
||||
|
||||
|
||||
class TestExtractUserMessageAndSystemPrompt:
|
||||
"""Test the _extract_user_message_and_system_prompt static method."""
|
||||
|
||||
def test_should_extract_user_message(self):
|
||||
"""Test extraction of the last user message."""
|
||||
messages = [
|
||||
{"role": "system", "content": "You are helpful."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "Hi!"},
|
||||
{"role": "user", "content": "How are you?"},
|
||||
]
|
||||
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
||||
messages
|
||||
)
|
||||
assert user_msg == "How are you?"
|
||||
assert sys_prompt == "You are helpful."
|
||||
|
||||
def test_should_handle_no_user_message(self):
|
||||
"""Test when there is no user message."""
|
||||
messages = [
|
||||
{"role": "system", "content": "You are helpful."},
|
||||
{"role": "assistant", "content": "Hi!"},
|
||||
]
|
||||
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
||||
messages
|
||||
)
|
||||
assert user_msg is None
|
||||
assert sys_prompt == "You are helpful."
|
||||
|
||||
def test_should_handle_multipart_content(self):
|
||||
"""Test extraction from multipart content messages."""
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Describe this image"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/img.png"}},
|
||||
],
|
||||
}
|
||||
]
|
||||
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
||||
messages
|
||||
)
|
||||
assert user_msg == "Describe this image"
|
||||
assert sys_prompt is None
|
||||
|
||||
def test_should_handle_empty_messages(self):
|
||||
"""Test with empty messages list."""
|
||||
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
|
||||
[]
|
||||
)
|
||||
assert user_msg is None
|
||||
assert sys_prompt is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue