mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
f2e340cf2b
commit
cc45d18e9c
15 changed files with 166 additions and 4 deletions
|
|
@ -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. "
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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} />);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
? [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 }),
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
|
||||
|
|
|
|||
|
|
@ -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 : [],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue