mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor(router): extract shared extract_text_from_input utility and fix AutoRouter empty-input guard
- Move _extract_text_from_input to litellm/router_strategy/utils.py to eliminate duplication between ComplexityRouter and AutoRouter - Fix AutoRouter empty-string/empty-list input guard: fall back to default model instead of passing empty text to the semantic router - Move semantic_router imports after early-return checks so environments without semantic_router installed can still hit early-return paths - Add TestExtractTextFromInput tests against shared utils function - Add test_plain_string_item_in_list to cover previously missing branch - Add test_auto_router_responses_api.py for AutoRouter early-return paths (no semantic_router dependency required)
This commit is contained in:
parent
4c8143aa2d
commit
35536bb145
5 changed files with 204 additions and 115 deletions
|
|
@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
|||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_strategy.utils import extract_text_from_input
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from semantic_router.routers.base import Route
|
||||
|
|
@ -81,44 +82,6 @@ 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,
|
||||
|
|
@ -132,12 +95,6 @@ class AutoRouter(CustomLogger):
|
|||
|
||||
Used for the litellm auto-router to modify the request before the routing decision is made.
|
||||
"""
|
||||
from semantic_router.routers import SemanticRouter
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||
LiteLLMRouterEncoder,
|
||||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
has_messages = messages is not None and len(messages) > 0
|
||||
|
|
@ -147,6 +104,23 @@ class AutoRouter(CustomLogger):
|
|||
# do nothing, return same inputs
|
||||
return None
|
||||
|
||||
# Responses API: if input is present but yields no usable text, fall
|
||||
# back to the default model without touching the semantic router.
|
||||
if not has_messages and has_input:
|
||||
extracted_early = extract_text_from_input(input) # type: ignore[arg-type]
|
||||
if extracted_early is None:
|
||||
return PreRoutingHookResponse(
|
||||
model=self.default_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
from semantic_router.routers import SemanticRouter
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||
LiteLLMRouterEncoder,
|
||||
)
|
||||
|
||||
if self.routelayer is None:
|
||||
#######################
|
||||
# Create the route layer
|
||||
|
|
@ -164,8 +138,8 @@ class AutoRouter(CustomLogger):
|
|||
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]
|
||||
# extracted_early is guaranteed non-None here (None case returned above)
|
||||
message_content = extracted_early # type: ignore[assignment]
|
||||
route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(
|
||||
text=message_content
|
||||
)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
|||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_strategy.utils import extract_text_from_input
|
||||
|
||||
from .config import (
|
||||
DEFAULT_CODE_KEYWORDS,
|
||||
|
|
@ -331,52 +332,6 @@ 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,
|
||||
|
|
@ -436,7 +391,7 @@ class ComplexityRouter(CustomLogger):
|
|||
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]
|
||||
user_message = extract_text_from_input(input) # type: ignore[arg-type]
|
||||
verbose_router_logger.debug(
|
||||
f"ComplexityRouter: extracted text from Responses API input: {user_message!r}"
|
||||
)
|
||||
|
|
|
|||
47
litellm/router_strategy/utils.py
Normal file
47
litellm/router_strategy/utils.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
"""Shared utilities for router strategies."""
|
||||
from typing import List, Optional, Union
|
||||
|
||||
|
||||
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 string
|
||||
or 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":
|
||||
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
|
||||
|
|
@ -0,0 +1,109 @@
|
|||
"""
|
||||
Tests for AutoRouter Responses API (input field) support.
|
||||
|
||||
These tests cover the new ``has_messages / has_input`` branching logic and the
|
||||
empty-string guard in ``AutoRouter.async_pre_routing_hook``. They do NOT
|
||||
require ``semantic_router`` to be installed because they only exercise code
|
||||
paths that return before any SemanticRouter call.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
# semantic_router is an optional dependency (beta feature). Stub it out so
|
||||
# these tests can run in environments where it is not installed.
|
||||
_SEMANTIC_ROUTER_MOCK = MagicMock()
|
||||
_SEMANTIC_ROUTER_STUBS = {
|
||||
"semantic_router": _SEMANTIC_ROUTER_MOCK,
|
||||
"semantic_router.routers": _SEMANTIC_ROUTER_MOCK.routers,
|
||||
"semantic_router.schema": _SEMANTIC_ROUTER_MOCK.schema,
|
||||
"semantic_router.routers.base": _SEMANTIC_ROUTER_MOCK.routers.base,
|
||||
}
|
||||
|
||||
|
||||
def _make_auto_router(default_model: str = "default-model") -> "AutoRouter": # type: ignore[name-defined]
|
||||
"""Create an AutoRouter instance without requiring semantic_router."""
|
||||
with patch.dict(sys.modules, _SEMANTIC_ROUTER_STUBS):
|
||||
from litellm.router_strategy.auto_router.auto_router import AutoRouter
|
||||
|
||||
with patch.object(AutoRouter, "_load_semantic_routing_routes", return_value=[]):
|
||||
return AutoRouter(
|
||||
model_name="test-auto-router",
|
||||
default_model=default_model,
|
||||
embedding_model="text-embedding-model",
|
||||
litellm_router_instance=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
class TestAutoRouterResponsesAPIEarlyReturns:
|
||||
"""
|
||||
Tests for async_pre_routing_hook paths that return before calling the
|
||||
semantic router — safe to run without semantic_router installed.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_messages_no_input_returns_none(self):
|
||||
"""When both messages and input are absent, the hook skips routing."""
|
||||
auto_router = _make_auto_router()
|
||||
result = await auto_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=None,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_messages_no_input_returns_none(self):
|
||||
"""An empty messages list with no input also skips routing."""
|
||||
auto_router = _make_auto_router()
|
||||
result = await auto_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[],
|
||||
input=None,
|
||||
)
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_string_input_returns_default_model(self):
|
||||
"""An empty string input falls back to the default model (not routelayer)."""
|
||||
auto_router = _make_auto_router(default_model="my-default")
|
||||
result = await auto_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input="",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "my-default"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_whitespace_only_input_returns_default_model(self):
|
||||
"""A whitespace-only input is treated the same as empty — use default."""
|
||||
auto_router = _make_auto_router(default_model="my-default")
|
||||
result = await auto_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=" ",
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "my-default"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_list_input_returns_default_model(self):
|
||||
"""An empty list input falls back to the default model."""
|
||||
auto_router = _make_auto_router(default_model="my-default")
|
||||
result = await auto_router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=None,
|
||||
input=[],
|
||||
)
|
||||
assert result is not None
|
||||
assert result.model == "my-default"
|
||||
|
|
@ -24,6 +24,7 @@ from litellm.router_strategy.complexity_router.config import (
|
|||
ComplexityRouterConfig,
|
||||
ComplexityTier,
|
||||
)
|
||||
from litellm.router_strategy.utils import extract_text_from_input
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -699,34 +700,37 @@ class TestEdgeCases:
|
|||
|
||||
|
||||
class TestExtractTextFromInput:
|
||||
"""Tests for ComplexityRouter._extract_text_from_input."""
|
||||
"""Tests for the shared extract_text_from_input utility."""
|
||||
|
||||
def test_plain_string(self, complexity_router):
|
||||
def test_plain_string(self):
|
||||
"""A plain string is returned as-is."""
|
||||
assert complexity_router._extract_text_from_input("Hello world") == "Hello world"
|
||||
assert extract_text_from_input("Hello world") == "Hello world"
|
||||
|
||||
def test_empty_string(self, complexity_router):
|
||||
def test_empty_string(self):
|
||||
"""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
|
||||
assert extract_text_from_input("") is None
|
||||
assert extract_text_from_input(" ") is None
|
||||
|
||||
def test_text_type_item(self, complexity_router):
|
||||
def test_text_type_item(self):
|
||||
"""A list with a single {type: text, text: ...} item."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
[{"type": "text", "text": "Hello world"}]
|
||||
)
|
||||
result = extract_text_from_input([{"type": "text", "text": "Hello world"}])
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_message_type_item_string_content(self, complexity_router):
|
||||
def test_plain_string_item_in_list(self):
|
||||
"""A list containing a bare string item is appended directly."""
|
||||
result = extract_text_from_input(["Hello", "world"])
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_message_type_item_string_content(self):
|
||||
"""A list with a {type: message, content: str} item."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
result = 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):
|
||||
def test_message_type_item_list_content(self):
|
||||
"""A {type: message} item whose content is a list of text parts."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
result = extract_text_from_input(
|
||||
[
|
||||
{
|
||||
"type": "message",
|
||||
|
|
@ -741,9 +745,9 @@ class TestExtractTextFromInput:
|
|||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_multiple_items_concatenated(self, complexity_router):
|
||||
def test_multiple_items_concatenated(self):
|
||||
"""Multiple text items are joined with spaces."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
result = extract_text_from_input(
|
||||
[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{"type": "text", "text": "world"},
|
||||
|
|
@ -751,20 +755,20 @@ class TestExtractTextFromInput:
|
|||
)
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_empty_list(self, complexity_router):
|
||||
def test_empty_list(self):
|
||||
"""An empty list returns None."""
|
||||
assert complexity_router._extract_text_from_input([]) is None
|
||||
assert extract_text_from_input([]) is None
|
||||
|
||||
def test_list_with_no_text(self, complexity_router):
|
||||
def test_list_with_no_text(self):
|
||||
"""A list with only non-text items returns None."""
|
||||
result = complexity_router._extract_text_from_input(
|
||||
result = extract_text_from_input(
|
||||
[{"type": "image_url", "image_url": {"url": "data:..."}}]
|
||||
)
|
||||
assert result is None
|
||||
|
||||
def test_non_list_non_string(self, complexity_router):
|
||||
def test_non_list_non_string(self):
|
||||
"""Non-string, non-list input returns None."""
|
||||
assert complexity_router._extract_text_from_input(42) is None # type: ignore[arg-type]
|
||||
assert extract_text_from_input(42) is None # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestPreRoutingHookResponsesAPI:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue