From cc45d18e9c41b19d9eabe077288f7fcc6a080b11 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 18 Jul 2026 19:24:56 -0700 Subject: [PATCH] feat(complexity-router): add return_raw_model_name toggle for response model field (#33875) * feat(complexity-router): optionally return raw model name Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): restore asyncio import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(tests): preserve staging asyncio import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): drop unused local asyncio import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(dashboard): add complexity router raw model toggle Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(complexity-router): move metadata key constant to constants.py Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(proxy-tests): preserve module spacing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 1 + litellm/proxy/common_request_processing.py | 12 +++++++- litellm/proxy/proxy_server.py | 4 +++ .../complexity_router/complexity_router.py | 7 +++++ .../complexity_router/config.py | 8 +++++ .../proxy_server/test_streaming_helpers.py | 16 ++++++++++ .../proxy/test_common_request_processing.py | 29 +++++++++++++++++-- .../router_strategy/test_complexity_router.py | 24 +++++++++++++++ .../add_model/ComplexityRouterConfig.test.tsx | 14 +++++++++ .../add_model/ComplexityRouterConfig.tsx | 25 +++++++++++++++- .../add_model/add_auto_router_tab.tsx | 2 ++ .../build_complexity_router_config.test.ts | 11 +++++++ .../build_complexity_router_config.ts | 4 +++ .../edit_auto_router_modal.test.ts | 10 +++++++ .../edit_auto_router_modal.tsx | 3 ++ 15 files changed, 166 insertions(+), 4 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 6432e2176c7..05944c81ea2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1292,6 +1292,7 @@ MAXIMUM_TRACEBACK_LINES_TO_LOG = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", X_LITELLM_DISABLE_CALLBACKS = "x-litellm-disable-callbacks" LITELLM_METADATA_FIELD = "litellm_metadata" OLD_LITELLM_METADATA_FIELD = "metadata" +RETURN_RAW_MODEL_NAME_METADATA_KEY = "_complexity_router_return_raw_model_name" LITELLM_TRUNCATED_PAYLOAD_FIELD = "litellm_truncated" LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE = ( "Truncation is a DB storage safeguard. " diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 1dc0ee3f947..3f9929f81da 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -33,6 +33,7 @@ from litellm.constants import ( LITELLM_DETAILED_TIMING, LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED, MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG, + RETURN_RAW_MODEL_NAME_METADATA_KEY, STREAM_SSE_DATA_PREFIX, ) from litellm.integrations.custom_guardrail import CustomGuardrail @@ -91,6 +92,13 @@ _CLIENT_DISCONNECTED_ERROR_INFORMATION: StandardLoggingPayloadErrorInformation = } +def _should_return_raw_model_name(request_data: dict[str, object]) -> bool: + return any( + isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True + for metadata in (request_data.get("metadata"), request_data.get("litellm_metadata")) + ) + + def _apply_client_disconnect_metadata(target_metadata: Optional[dict[str, object]]) -> None: if target_metadata is None: return @@ -672,6 +680,7 @@ def _override_openai_response_model( response_obj: Any, requested_model: str, log_context: str, + return_raw_model_name: bool = False, ) -> None: """ Force the OpenAI-compatible `model` field in the response to match what the client requested. @@ -695,7 +704,7 @@ def _override_openai_response_model( 3. If this was a fastest_response batch completion, use the winning model's model group name instead of the comma-separated list the client sent. """ - if not requested_model: + if return_raw_model_name or not requested_model: return hidden_params = get_hidden_params_dict(response_obj) @@ -1938,6 +1947,7 @@ class ProxyBaseLLMRequestProcessing: response_obj=response, requested_model=requested_model_from_client, log_context=f"litellm_call_id={logging_obj.litellm_call_id}", + return_raw_model_name=_should_return_raw_model_name(self.data), ) hidden_params = get_hidden_params_dict(response) # get any updated response headers diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index aed345c5db4..3b40abed19e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -291,6 +291,7 @@ from litellm.proxy.caching_routes import router as caching_router from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, _is_azure_model_router_request, + _should_return_raw_model_name, create_response, ) from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_proxy @@ -7076,6 +7077,9 @@ def _restamp_streaming_chunk_model( fallback_was_attempted: bool = False, fallback_model_from_metadata: str | None = None, ) -> tuple[Any, bool]: + if _should_return_raw_model_name(request_data): + return chunk, model_mismatch_logged + target_model = fallback_model_from_metadata if fallback_was_attempted else requested_model_from_client # Always return the client-requested model name (not provider-prefixed internal identifiers) # on streaming chunks. diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 695d8b8aeaa..e5268b5107b 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -23,6 +23,7 @@ from typing import TYPE_CHECKING, Any, Literal, Union, cast from pydantic import BaseModel from litellm._logging import verbose_router_logger +from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import ModelResponse @@ -956,6 +957,12 @@ class ComplexityRouter(CustomLogger): """ from litellm.types.router import PreRoutingHookResponse + if self.config.return_raw_model_name: + metadata_key = "litellm_metadata" if "litellm_metadata" in request_kwargs else "metadata" + metadata = request_kwargs.setdefault(metadata_key, {}) + if isinstance(metadata, dict): + metadata[RETURN_RAW_MODEL_NAME_METADATA_KEY] = True + use_session_affinity = self.config.session_affinity and not self.config.plugins session_id = self._get_session_id_from_request_kwargs(request_kwargs) if use_session_affinity else None cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 17c2c287dde..7437138fbb7 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -311,6 +311,14 @@ class ComplexityRouterConfig(BaseModel): description="Default model to use if tier cannot be determined", ) + return_raw_model_name: bool = Field( + default=False, + description=( + "Return the resolved raw model name in the response model field instead of " + "the client-requested complexity-router alias" + ), + ) + # Classifier strategy classifier_type: Literal["heuristic", "llm"] = Field( default="heuristic", diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index 699606b5277..f7e2d276a2e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -19,6 +19,7 @@ import json import pytest +from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY import litellm.proxy.proxy_server as ps from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ( @@ -272,6 +273,21 @@ def test_restamp_streaming_chunk_model_overrides_model_on_basemodel(): assert snapshot == {"model": "gpt-4", "logged": True, "same_object": True} +@pytest.mark.parametrize("return_raw_model_name", [False, True]) +def test_restamp_streaming_chunk_model_respects_raw_model_name_toggle(return_raw_model_name): + chunk = _simple_chunk(model="gpt-4o-mini") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="auto_router/complexity_router", + request_data={"metadata": {RETURN_RAW_MODEL_NAME_METADATA_KEY: return_raw_model_name}}, + model_mismatch_logged=False, + ) + + expected_model = "gpt-4o-mini" if return_raw_model_name else "auto_router/complexity_router" + assert new_chunk.model == expected_model + assert logged is (not return_raw_model_name) + + def test_restamp_streaming_chunk_model_overrides_model_on_dict(): chunk = {"model": "internal", "choices": []} new_chunk, logged = _restamp_streaming_chunk_model( diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index ebfbb46053d..58f81cdad35 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -11,6 +11,7 @@ from fastapi.responses import JSONResponse, StreamingResponse import litellm from litellm._uuid import uuid +from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.opentelemetry import UserAPIKeyAuth from litellm.proxy.common_request_processing import ( @@ -27,6 +28,7 @@ from litellm.proxy.common_request_processing import ( _is_azure_model_router_request, _override_openai_response_model, _parse_event_data_for_error, + _should_return_raw_model_name, _UpstreamClosingStreamingResponse, create_response, ) @@ -1675,6 +1677,31 @@ class TestExtractErrorFromSSEChunk: class TestOverrideOpenAIResponseModel: """Tests for _override_openai_response_model function""" + @pytest.mark.parametrize("return_raw_model_name", [False, True]) + def test_raw_model_name_toggle(self, return_raw_model_name): + response_obj = {"model": "gpt-4o-mini"} + + _override_openai_response_model( + response_obj=response_obj, + requested_model="auto_router/complexity_router", + log_context="test_context", + return_raw_model_name=return_raw_model_name, + ) + + expected_model = "gpt-4o-mini" if return_raw_model_name else "auto_router/complexity_router" + assert response_obj["model"] == expected_model + + @pytest.mark.parametrize( + "request_data, expected", + [ + ({"metadata": {}}, False), + ({"metadata": {RETURN_RAW_MODEL_NAME_METADATA_KEY: True}}, True), + ({"litellm_metadata": {RETURN_RAW_MODEL_NAME_METADATA_KEY: True}}, True), + ], + ) + def test_raw_model_name_toggle_metadata(self, request_data, expected): + assert _should_return_raw_model_name(request_data) is expected + def test_override_model_preserves_fallback_model_when_fallback_occurred_object( self, ): @@ -3203,8 +3230,6 @@ class TestDisconnectGatherCleanup: async def test_base_process_llm_request_preserves_llm_error_after_gather( self, monkeypatch ): - import asyncio - import litellm.proxy.common_request_processing as cpr from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 280a0fe072a..ef70687bd97 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -20,6 +20,7 @@ import litellm from litellm import Router from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache +from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.router_strategy.complexity_router.complexity_router import ( ComplexityRouter, DimensionScore, @@ -125,6 +126,29 @@ class TestComplexityRouterInit: ) assert router.config.default_model == "fallback-model" + @pytest.mark.asyncio + @pytest.mark.parametrize("return_raw_model_name", [False, True]) + async def test_pre_routing_hook_propagates_raw_model_response_setting( + self, mock_router_instance, basic_config, return_raw_model_name + ): + config = {**basic_config, "return_raw_model_name": return_raw_model_name} + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config=config, + ) + request_kwargs = {} + + result = await router.async_pre_routing_hook( + model="test-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "Hello"}], + ) + + assert result is not None + metadata = request_kwargs.get("metadata", {}) + assert metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY, False) is return_raw_model_name + class TestTokenScoring: """Test token count scoring.""" diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index a2e2ca21d00..6b3c1961468 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -77,6 +77,20 @@ describe("ComplexityRouterConfig", () => { expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument(); }); + it("should toggle returning the raw model name", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + renderWithProviders(); + + await user.click(screen.getByText("Advanced: Response Format")); + await user.click(screen.getByRole("switch")); + + expect(onChange).toHaveBeenCalledWith({ + ...defaultValue, + return_raw_model_name: true, + }); + }); + it("should reveal classifier model and timeout fields when llm is selected", () => { const onChange = vi.fn(); renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 8008012a95c..1f2edf697a9 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -1,5 +1,5 @@ import { InfoCircleOutlined } from "@ant-design/icons"; -import { Select as AntdSelect, Card, Collapse, Divider, Space, Tooltip, Typography } from "antd"; +import { Select as AntdSelect, Card, Collapse, Divider, Space, Switch, Tooltip, Typography } from "antd"; import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; @@ -44,6 +44,7 @@ export interface ComplexityRouterConfigValue { adaptive_weights?: AdaptiveRouterWeights; tier_distance_penalty?: number; adaptive_eligible?: AdaptiveEligible; + return_raw_model_name?: boolean; } interface ComplexityRouterConfigProps { @@ -218,6 +219,28 @@ const ComplexityRouterConfig: React.FC = ({ ), children: , }, + { + key: "response", + label: ( + + Advanced: Response Format + + ), + children: ( + <> +
+ onChange({ ...value, return_raw_model_name: returnRawModelName })} + /> + Return raw model name +
+ + Return the resolved underlying model name in responses instead of the autorouter alias. + + + ), + }, ...(onEscalationKeywordsChange ? [ { diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 6e7bc49afce..e7826e09dce 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -100,6 +100,7 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc adaptive_weights: adaptiveWeights = DEFAULT_ADAPTIVE_WEIGHTS, tier_distance_penalty: tierDistancePenalty = DEFAULT_TIER_DISTANCE_PENALTY, adaptive_eligible: adaptiveEligible = "all", + return_raw_model_name: returnRawModelName = false, } = complexityRouterConfig; const missingTiersError = getMissingTiersError(tiers); @@ -148,6 +149,7 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc adaptiveWeights, tierDistancePenalty, adaptiveEligible, + returnRawModelName, }; const submitValues = { diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 0c9c19d1286..b5973bf7101 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -26,6 +26,7 @@ const baseParams: BuildComplexityRouterConfigParams = { adaptiveWeights: { quality: 0.3, cost: 0.7 }, tierDistancePenalty: 0.5, adaptiveEligible: "all", + returnRawModelName: false, }; describe("buildComplexityRouterConfig", () => { @@ -164,6 +165,16 @@ describe("buildComplexityRouterConfig", () => { expect(config.adaptive_eligible).toBeUndefined(); }); + it("omits return_raw_model_name when disabled", () => { + const config = buildComplexityRouterConfig({ ...baseParams, returnRawModelName: false }); + expect(config.return_raw_model_name).toBeUndefined(); + }); + + it("includes return_raw_model_name when enabled", () => { + const config = buildComplexityRouterConfig({ ...baseParams, returnRawModelName: true }); + expect(config.return_raw_model_name).toBe(true); + }); + it("includes tier_distance_penalty when adaptive is enabled with eligible='all'", () => { const config = buildComplexityRouterConfig({ ...baseParams, diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 0b92dc1b02d..3b41b916611 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -21,6 +21,7 @@ export interface BuildComplexityRouterConfigParams { adaptiveWeights: AdaptiveRouterWeights; tierDistancePenalty: number; adaptiveEligible: AdaptiveEligible; + returnRawModelName: boolean; } export interface ComplexityRouterConfigPayload { @@ -37,6 +38,7 @@ export interface ComplexityRouterConfigPayload { adaptive_weights?: AdaptiveRouterWeights; tier_distance_penalty?: number; adaptive_eligible?: AdaptiveEligible; + return_raw_model_name?: boolean; } const TIER_KEYS: Array = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]; @@ -76,6 +78,7 @@ export const buildComplexityRouterConfig = ({ adaptiveWeights, tierDistancePenalty, adaptiveEligible, + returnRawModelName, }: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => { const cleanedEscalationKeywords = escalationKeywords.map((keyword) => keyword.trim()).filter(Boolean); // Trim keywords and drop empty ones; drop any rule left with no keywords. Clicking @@ -104,5 +107,6 @@ export const buildComplexityRouterConfig = ({ ...(adaptiveEligible === "all" && { tier_distance_penalty: tierDistancePenalty }), adaptive_eligible: adaptiveEligible, }), + ...(returnRawModelName && { return_raw_model_name: true }), }; }; diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.ts index cd8093928d5..17fa810b529 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.test.ts @@ -18,6 +18,7 @@ const storedConfigValue = { adaptive_weights: { quality: 0.3, cost: 0.7 }, tier_distance_penalty: 0.8, adaptive_eligible: "all", + return_raw_model_name: true, }; const storedConfig = JSON.stringify(storedConfigValue); @@ -80,6 +81,15 @@ describe("buildUpdatedComplexityRouterConfig", () => { expect(updatedConfig).toEqual(expectedAdaptiveDisabledConfig); }); + it("includes return_raw_model_name only when enabled", () => { + const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, { + ...classifiedTierValue, + return_raw_model_name: true, + }); + + expect(updatedConfig.return_raw_model_name).toBe(true); + }); + it("updates custom technical keywords when they are edited", () => { const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue, ["postgres"]); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index f85cd16486a..46d7d41d9b3 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -38,6 +38,7 @@ const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "adaptive_weights", "tier_distance_penalty", "adaptive_eligible", + "return_raw_model_name", ]); const toRecord = (value: unknown): Record => { @@ -78,6 +79,7 @@ export const buildUpdatedComplexityRouterConfig = ( }), adaptive_eligible: adaptiveEligible, }), + ...(value.return_raw_model_name && { return_raw_model_name: true }), }; }; @@ -158,6 +160,7 @@ const EditAutoRouterModal: React.FC = ({ adaptive_weights: parsedConfig.adaptive_weights, tier_distance_penalty: parsedConfig.tier_distance_penalty, adaptive_eligible: parsedConfig.adaptive_eligible || "all", + return_raw_model_name: parsedConfig.return_raw_model_name || false, }); setCustomTechnicalKeywords( Array.isArray(parsedConfig.custom_technical_keywords) ? parsedConfig.custom_technical_keywords : [],