mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(router): support Responses API 'input' field in complexity and auto routers
- Pass 'input' parameter through _ageneric_api_call_with_fallbacks_helper - Implement _extract_text_from_input in Complexity and Auto routers - Handle various Responses API input schemas (string, list of dicts) - Fixes 'Unmapped LLM provider' error for /v1/responses endpoint Closes #25134
This commit is contained in:
parent
0169f140bf
commit
4c8143aa2d
4 changed files with 307 additions and 23 deletions
|
|
@ -3846,6 +3846,7 @@ class Router:
|
|||
model=model,
|
||||
request_kwargs=kwargs,
|
||||
messages=kwargs.get("messages", None),
|
||||
input=kwargs.get("input", None),
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -81,6 +81,44 @@ class AutoRouter(CustomLogger):
|
|||
)
|
||||
return auto_router_routes
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_input(input: Union[str, List]) -> Optional[str]:
|
||||
"""
|
||||
Extract plain text from a Responses API ``input`` field.
|
||||
|
||||
Handles bare strings, ``{type: "text", text: ...}`` items, and
|
||||
``{type: "message", content: ...}`` items (where ``content`` may
|
||||
itself be a string or a list of ``{type, text}`` parts).
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
return input.strip() or None
|
||||
|
||||
if not isinstance(input, list):
|
||||
return None
|
||||
|
||||
parts: List[str] = []
|
||||
for item in input:
|
||||
if isinstance(item, str):
|
||||
parts.append(item)
|
||||
elif isinstance(item, dict):
|
||||
item_type = item.get("type", "")
|
||||
if item_type == "text":
|
||||
text = item.get("text") or ""
|
||||
if text:
|
||||
parts.append(text)
|
||||
elif item_type == "message":
|
||||
content = item.get("content") or ""
|
||||
if isinstance(content, str):
|
||||
if content:
|
||||
parts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
t = part.get("text") or ""
|
||||
if t:
|
||||
parts.append(t)
|
||||
return " ".join(parts).strip() or None
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -102,7 +140,10 @@ class AutoRouter(CustomLogger):
|
|||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
if messages is None:
|
||||
has_messages = messages is not None and len(messages) > 0
|
||||
has_input = input is not None
|
||||
|
||||
if not has_messages and not has_input:
|
||||
# do nothing, return same inputs
|
||||
return None
|
||||
|
||||
|
|
@ -119,8 +160,12 @@ class AutoRouter(CustomLogger):
|
|||
auto_sync=self.auto_sync_value,
|
||||
)
|
||||
|
||||
user_message: Dict[str, str] = messages[-1]
|
||||
message_content: str = user_message.get("content", "")
|
||||
if has_messages:
|
||||
user_message_dict: Dict[str, str] = messages[-1] # type: ignore[index]
|
||||
message_content: str = user_message_dict.get("content", "")
|
||||
else:
|
||||
# Responses API: extract plain text from ``input``
|
||||
message_content = self._extract_text_from_input(input) or "" # type: ignore[arg-type]
|
||||
route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(
|
||||
text=message_content
|
||||
)
|
||||
|
|
|
|||
|
|
@ -331,6 +331,52 @@ class ComplexityRouter(CustomLogger):
|
|||
f"No model configured for tier {tier_key} and no default_model set"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_input(input: Union[str, List]) -> Optional[str]:
|
||||
"""
|
||||
Extract plain text from a Responses API ``input`` field.
|
||||
|
||||
The Responses API accepts either a bare string or a list of input
|
||||
items (``ResponseInputParam``). Each item may be:
|
||||
|
||||
* A plain ``str``.
|
||||
* A dict with ``type="text"`` and a ``text`` key.
|
||||
* A dict with ``type="message"`` whose ``content`` is itself a list
|
||||
of content parts (same ``{type, text}`` shape).
|
||||
|
||||
Returns the concatenated text, or ``None`` when nothing extractable
|
||||
is found.
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
return input.strip() or None
|
||||
|
||||
if not isinstance(input, list):
|
||||
return None
|
||||
|
||||
parts: List[str] = []
|
||||
for item in input:
|
||||
if isinstance(item, str):
|
||||
parts.append(item)
|
||||
elif isinstance(item, dict):
|
||||
item_type = item.get("type", "")
|
||||
if item_type == "text":
|
||||
text = item.get("text") or ""
|
||||
if text:
|
||||
parts.append(text)
|
||||
elif item_type == "message":
|
||||
# { type: "message", role: "...", content: str | list }
|
||||
content = item.get("content") or ""
|
||||
if isinstance(content, str):
|
||||
if content:
|
||||
parts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
t = part.get("text") or ""
|
||||
if t:
|
||||
parts.append(t)
|
||||
return " ".join(parts).strip() or None
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -343,12 +389,14 @@ class ComplexityRouter(CustomLogger):
|
|||
Pre-routing hook called before the routing decision.
|
||||
|
||||
Classifies the request by complexity and returns the appropriate model.
|
||||
Supports both the Chat Completions API (``messages``) and the Responses
|
||||
API (``input``).
|
||||
|
||||
Args:
|
||||
model: The original model name requested.
|
||||
request_kwargs: The request kwargs.
|
||||
messages: The messages in the request.
|
||||
input: Optional input for embeddings.
|
||||
messages: The messages in the request (Chat Completions API).
|
||||
input: The input field from a Responses API request.
|
||||
specific_deployment: Whether a specific deployment was requested.
|
||||
|
||||
Returns:
|
||||
|
|
@ -356,9 +404,12 @@ class ComplexityRouter(CustomLogger):
|
|||
"""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
if messages is None or len(messages) == 0:
|
||||
has_messages = messages is not None and len(messages) > 0
|
||||
has_input = input is not None
|
||||
|
||||
if not has_messages and not has_input:
|
||||
verbose_router_logger.debug(
|
||||
"ComplexityRouter: No messages provided, skipping routing"
|
||||
"ComplexityRouter: No messages or input provided, skipping routing"
|
||||
)
|
||||
return None
|
||||
|
||||
|
|
@ -366,22 +417,29 @@ class ComplexityRouter(CustomLogger):
|
|||
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", "")
|
||||
for part in content
|
||||
if isinstance(part, dict) and part.get("type") == "text"
|
||||
]
|
||||
content = " ".join(text_parts).strip()
|
||||
if isinstance(content, str) and content:
|
||||
if role == "user" and user_message is None:
|
||||
user_message = content
|
||||
elif role == "system" and system_prompt is None:
|
||||
system_prompt = content
|
||||
if has_messages:
|
||||
for msg in reversed(messages): # type: ignore[arg-type]
|
||||
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", "")
|
||||
for part in content
|
||||
if isinstance(part, dict) and part.get("type") == "text"
|
||||
]
|
||||
content = " ".join(text_parts).strip()
|
||||
if isinstance(content, str) and content:
|
||||
if role == "user" and user_message is None:
|
||||
user_message = content
|
||||
elif role == "system" and system_prompt is None:
|
||||
system_prompt = content
|
||||
elif has_input:
|
||||
# Responses API: extract text from the input field
|
||||
user_message = self._extract_text_from_input(input) # type: ignore[arg-type]
|
||||
verbose_router_logger.debug(
|
||||
f"ComplexityRouter: extracted text from Responses API input: {user_message!r}"
|
||||
)
|
||||
|
||||
if user_message is None:
|
||||
verbose_router_logger.debug(
|
||||
|
|
|
|||
|
|
@ -698,6 +698,186 @@ class TestEdgeCases:
|
|||
assert any("multi-step" in s.lower() for s in signals), f"Expected multi-step signal, got {signals}"
|
||||
|
||||
|
||||
class TestExtractTextFromInput:
|
||||
"""Tests for ComplexityRouter._extract_text_from_input."""
|
||||
|
||||
def test_plain_string(self, complexity_router):
|
||||
"""A plain string is returned as-is."""
|
||||
assert complexity_router._extract_text_from_input("Hello world") == "Hello world"
|
||||
|
||||
def test_empty_string(self, complexity_router):
|
||||
"""An empty / whitespace-only string returns None."""
|
||||
assert complexity_router._extract_text_from_input("") is None
|
||||
assert complexity_router._extract_text_from_input(" ") is None
|
||||
|
||||
def test_text_type_item(self, complexity_router):
|
||||
"""A list with a single {type: text, text: ...} item."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[{"type": "text", "text": "Hello world"}]
|
||||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_message_type_item_string_content(self, complexity_router):
|
||||
"""A list with a {type: message, content: str} item."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[{"type": "message", "role": "user", "content": "Hello world"}]
|
||||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_message_type_item_list_content(self, complexity_router):
|
||||
"""A {type: message} item whose content is a list of text parts."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{"type": "image_url", "image_url": {"url": "data:..."}},
|
||||
{"type": "text", "text": "world"},
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_multiple_items_concatenated(self, complexity_router):
|
||||
"""Multiple text items are joined with spaces."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{"type": "text", "text": "world"},
|
||||
]
|
||||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_empty_list(self, complexity_router):
|
||||
"""An empty list returns None."""
|
||||
assert complexity_router._extract_text_from_input([]) is None
|
||||
|
||||
def test_list_with_no_text(self, complexity_router):
|
||||
"""A list with only non-text items returns None."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[{"type": "image_url", "image_url": {"url": "data:..."}}]
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_non_list_non_string(self, complexity_router):
|
||||
"""Non-string, non-list input returns None."""
|
||||
assert complexity_router._extract_text_from_input(42) is None # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestPreRoutingHookResponsesAPI:
|
||||
"""Tests for async_pre_routing_hook with the Responses API ``input`` field."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_plain_string_simple(self, complexity_router):
|
||||
"""Responses API with a simple plain-string input routes to SIMPLE tier."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input="Hello!",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "gpt-4o-mini"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_plain_string_reasoning(self, complexity_router):
|
||||
"""Responses API routes a reasoning-heavy input to REASONING tier."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=(
|
||||
"Let's think step by step and reason through this: "
|
||||
"analyze the architecture carefully."
|
||||
),
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "o1-preview"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_text_type_item(self, complexity_router):
|
||||
"""Responses API with a list of text-type items routes correctly."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=[{"type": "text", "text": "Hello!"}],
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "gpt-4o-mini"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_message_type_item(self, complexity_router):
|
||||
"""Responses API with a message-type item (list content) routes correctly."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=[
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "text", "text": "Hello!"}],
|
||||
}
|
||||
],
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "gpt-4o-mini"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_no_messages_no_input_returns_none(
|
||||
self, complexity_router
|
||||
):
|
||||
"""When both messages and input are absent, the hook skips routing."""
|
||||
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_responses_api_empty_input_returns_default(self, complexity_router):
|
||||
"""An empty string input falls back to the default model."""
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input="",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model in [
|
||||
"gpt-4o-mini",
|
||||
"gpt-4o",
|
||||
"claude-sonnet-4-20250514",
|
||||
"o1-preview",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_messages_take_priority_over_input(self, complexity_router):
|
||||
"""When both messages and input are present, messages are used."""
|
||||
# messages says REASONING but input says SIMPLE — messages should win
|
||||
result = await complexity_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"Think step by step and reason through: design a distributed system."
|
||||
),
|
||||
}
|
||||
],
|
||||
input="Hello!",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "o1-preview"
|
||||
|
||||
|
||||
class TestRouterComplexityDeploymentMethods:
|
||||
"""Tests for Router._is_complexity_router_deployment and Router.init_complexity_router_deployment."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue