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 <krrishdholakia@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-07-18 19:24:56 -07:00 • committed by GitHub
parent f2e340cf2b
commit cc45d18e9c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 166 additions and 4 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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(<ComplexityRouterConfig {...baseProps} onChange={onChange} />);
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(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={onChange} />);

View file

@ -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<ComplexityRouterConfigProps> = ({
),
children: <AdaptiveRoutingConfig value={value} onChange={onChange} />,
},
{
key: "response",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Response Format
</Text>
),
children: (
<>
<div className="flex items-center gap-2 mb-2">
<Switch
checked={value.return_raw_model_name ?? false}
onChange={(returnRawModelName) => onChange({ ...value, return_raw_model_name: returnRawModelName })}
/>
<Text strong>Return raw model name</Text>
</div>
<Text type="secondary" style={{ display: "block", fontSize: 12 }}>
Return the resolved underlying model name in responses instead of the autorouter alias.
</Text>
</>
),
},
...(onEscalationKeywordsChange
? [
{

View file

@ -100,6 +100,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ 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<AddAutoRouterTabProps> = ({ form, handleOk, acc
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
returnRawModelName,
};
const submitValues = {

View file

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

View file

@ -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<keyof ComplexityTiers> = ["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 }),
};
};

View file

@ -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"]);

View file

@ -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<string, unknown> => {
@ -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<EditAutoRouterModalProps> = ({
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 : [],