diff --git a/litellm/router.py b/litellm/router.py index 9185e437a3a..189f37c91b3 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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: diff --git a/litellm/router_strategy/auto_router/auto_router.py b/litellm/router_strategy/auto_router/auto_router.py index 6a786115193..bef65a17d39 100644 --- a/litellm/router_strategy/auto_router/auto_router.py +++ b/litellm/router_strategy/auto_router/auto_router.py @@ -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 ) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 29bed360fab..09308685b63 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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( diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 2ca823f6a12..1a4a0287c01 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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."""