diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index bef65a17d39..d2795dc8d7b 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -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 ) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 09308685b63..5f1c6abffef 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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}" ) diff --git a/litellm/router_strategy/utils.py b/litellm/router_strategy/utils.py new file mode 100644 index 00000000000..d8352ac98f4 --- /dev/null +++ b/litellm/router_strategy/utils.py @@ -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 diff --git a/tests/test_litellm/router_strategy/test_auto_router_responses_api.py b/tests/test_litellm/router_strategy/test_auto_router_responses_api.py new file mode 100644 index 00000000000..82de9537514 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_auto_router_responses_api.py @@ -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" diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 1a4a0287c01..2df38a39063 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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: