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:
Jerry Wei 2026-04-11 16:46:24 -05:00
parent 4c8143aa2d
commit 35536bb145
5 changed files with 204 additions and 115 deletions

View file

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

View file

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

View 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

View file

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

View file

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