feat(proxy): expose complexity routing headers (#40792)

(cherry picked from commit c817faec7a)

Co-authored-by: Tin <tin@berri.ai>
Co-authored-by: Claude Code <noreply@anthropic.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-11 17:14:40 -07:00 committed by GitHub
parent 8e4f2abb40
commit dab7f6a86a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 487 additions and 14 deletions

View file

@ -2437,6 +2437,7 @@ class ProxyBaseLLMRequestProcessing:
if self._is_streaming_request(
data=self.data, is_streaming_request=is_streaming_request
) or self._is_streaming_response(response): # use generate_responses to stream responses
selected_data_generator: AsyncGenerator[str, None] | None = None
# Call response headers hook for streaming success
stream_callback_headers: Final = await proxy_logging_obj.post_call_response_headers_hook(
data=self.data,
@ -2561,14 +2562,9 @@ class ProxyBaseLLMRequestProcessing:
None if _should_return_raw_model_name(self.data) else requested_model_from_client
),
)
return await create_response(
generator=wrap_sse_stream_with_keepalive_pings(
stream=selected_data_generator,
ping_interval_seconds=litellm.anthropic_sse_ping_interval_seconds,
),
media_type="text/event-stream",
headers=custom_headers,
request=request,
selected_data_generator = wrap_sse_stream_with_keepalive_pings(
stream=selected_data_generator,
ping_interval_seconds=litellm.anthropic_sse_ping_interval_seconds,
)
# Non-streaming response - fall through to normal response handling
elif select_data_generator:
@ -2595,6 +2591,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
)
)
if selected_data_generator is not None:
return await create_response(
generator=selected_data_generator,
media_type="text/event-stream",

View file

@ -124,9 +124,11 @@ from litellm.router_utils.add_retry_fallback_headers import (
add_retry_headers_to_response,
apply_quality_router_decision_headers,
apply_remaining_usage_headers,
complexity_router_decision_headers,
ensure_response_additional_headers,
get_hidden_params_dict,
prepare_response_for_header_attachment,
replace_complexity_router_headers,
response_in_flight_token_count,
)
from litellm.router_utils.auto_router_model_naming import (
@ -555,6 +557,7 @@ class FallbackAwareAnthropicMessagesStream:
def __init__(self, async_generator: AsyncGenerator[bytes, None], source_iterator: object) -> None:
self._async_generator = async_generator
self._source_iterator = source_iterator
self.fallback_headers_adopted = False
self._hidden_params = dict( # mutable-ok: mutated in place by merge_fallback_hidden_params
getattr(source_iterator, "_hidden_params", None) or {}
)
@ -565,6 +568,7 @@ class FallbackAwareAnthropicMessagesStream:
def adopt_fallback_source(self, fallback_response: object) -> None:
self._source_iterator = fallback_response
self.fallback_headers_adopted = True
def __aiter__(self) -> "FallbackAwareAnthropicMessagesStream":
return self
@ -593,7 +597,9 @@ class FallbackAwareAnthropicMessagesStream:
self._hidden_params = { # mutable-ok: matches _hidden_params' existing dict[str, object] shape
**self._hidden_params,
**fallback_hidden_params,
"additional_headers": {**existing_headers, **fallback_headers}, # mutable-ok: same shape
"additional_headers": dict( # mutable-ok: hidden params expect a writable header bag
replace_complexity_router_headers(existing_headers, fallback_headers)
),
}
@ -3128,6 +3134,8 @@ class Router:
async generator.
"""
fallback_headers_adopted: bool = False
def __init__(self, async_generator: AsyncGenerator):
import time
from datetime import datetime
@ -3179,6 +3187,12 @@ class Router:
# api_base, additional_headers) keep flowing.
self._hidden_params = dict(getattr(source_iterator, "_hidden_params", None) or {})
def adopt_fallback_headers(self, fallback_response: object) -> tuple[dict[str, object], dict[str, object]]:
prepared: Final = Router._prepare_fallback_hidden_params(fallback_response)
self._hidden_params = {**prepared[0], "additional_headers": prepared[1]} # mutable-ok: stream metadata
self.fallback_headers_adopted = True
return prepared
def __aiter__(self):
return self
@ -3269,8 +3283,8 @@ class Router:
include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True,
)
prepared_fallback_hidden_params = wrapper.adopt_fallback_headers(fallback_response)
if hasattr(fallback_response, "__aiter__"):
prepared_fallback_hidden_params = Router._prepare_fallback_hidden_params(fallback_response)
async for fallback_item in fallback_response:
Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params)
if partial_usage is not None:
@ -3305,7 +3319,8 @@ class Router:
exc,
)
return FallbackResponsesStreamWrapper(stream_with_fallbacks())
wrapper: Final = FallbackResponsesStreamWrapper(stream_with_fallbacks())
return wrapper
def _completion_streaming_iterator(
self,
@ -11108,7 +11123,7 @@ class Router:
self,
response: object,
model_group: str | None = None,
request_kwargs: dict | None = None,
request_kwargs: dict[str, object] | None = None,
) -> Any:
"""
Add the most accurate rate limit headers for a given model response.
@ -11124,6 +11139,7 @@ class Router:
additional_headers: Final = ensure_response_additional_headers(response)
additional_headers["x-litellm-model-group"] = model_group
apply_quality_router_decision_headers(additional_headers, request_kwargs)
additional_headers.update(complexity_router_decision_headers(request_kwargs))
if model_group is not None:
remaining_usage: Final = await self.get_remaining_model_group_usage(model_group)

View file

@ -1,7 +1,10 @@
import json
import math
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Protocol, TypedDict, cast
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter, ValidationError
class FallbackErrorInfo(TypedDict):
@ -15,6 +18,68 @@ class _HiddenParamsHost(Protocol):
_hidden_params: dict[str, object]
_EMPTY_OBJECT_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_ROUTING_HEADER_MAPPING: Final = TypeAdapter(Mapping[str, object])
_COMPLEXITY_ROUTER_HEADER_PREFIX: Final = "x-litellm-complexity-router-"
def _routing_header_mapping(value: object) -> Mapping[str, object]:
try:
mapping: Final[Mapping[str, object]] = _ROUTING_HEADER_MAPPING.validate_python(value, strict=True)
return mapping
except ValidationError:
return _EMPTY_OBJECT_MAPPING
def _header_string(value: object) -> str | None:
if not isinstance(value, str):
return None
normalized: Final = value.strip()
return normalized if normalized and all(" " <= character <= "~" for character in normalized) else None
def complexity_router_decision_headers(request_kwargs: object) -> Mapping[str, str]:
data: Final = _routing_header_mapping(request_kwargs)
metadata_key: Final = "litellm_metadata" if "litellm_metadata" in data else "metadata"
decision: Final = _routing_header_mapping(_routing_header_mapping(data.get(metadata_key)).get("routing_decision"))
if decision.get("router_type") != "complexity":
return MappingProxyType({})
score: Final = decision.get("score")
values: Final = (
("tier", decision.get("tier")),
("cause", decision.get("cause")),
(
"score",
str(score)
if isinstance(score, (int, float)) and not isinstance(score, bool) and math.isfinite(score)
else None,
),
(
"reasoning-effort",
_routing_header_mapping(decision.get("tier_litellm_params")).get("reasoning_effort"),
),
)
return MappingProxyType(
{
f"{_COMPLEXITY_ROUTER_HEADER_PREFIX}{key}": header_value
for key, value in values
if (header_value := _header_string(value)) is not None
}
)
def replace_complexity_router_headers(
existing_headers: Mapping[str, object], new_headers: Mapping[str, object]
) -> Mapping[str, object]:
return MappingProxyType(
{
key: value
for key, value in (*existing_headers.items(), *new_headers.items())
if key in new_headers or not key.startswith(_COMPLEXITY_ROUTER_HEADER_PREFIX)
}
)
class HiddenParamsAsyncIteratorWrapper:
"""
Wraps a bare async generator/iterator (e.g. a provider's raw SSE

View file

@ -8259,6 +8259,83 @@ class TestStreamingResponseHeadersFollowFallback:
assert result.headers["x-callback-header"] == "kept"
class _MessagesFallbackStream:
def __init__(self) -> None:
self.fallback_headers_adopted = False
self._hidden_params: dict[str, object] = {
"additional_headers": {
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-reasoning-effort": "xhigh",
}
}
self._chunks = iter(
(
b'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"type":"text_delta","text":"OK"}}\n\n',
)
)
def __aiter__(self) -> "_MessagesFallbackStream":
return self
async def __anext__(self) -> bytes:
self._hidden_params = {
"model_id": "fallback-deployment",
"additional_headers": {"x-fallback-only": "yes"},
}
self.fallback_headers_adopted = True
try:
return next(self._chunks)
except StopIteration:
raise StopAsyncIteration from None
async def aclose(self) -> None:
return None
@pytest.mark.asyncio
async def test_messages_http_headers_refresh_after_lazy_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.caching.caching import DualCache
stream = _MessagesFallbackStream()
logging_obj = MagicMock()
logging_obj.litellm_call_id = "messages-fallback-headers"
logging_obj._defer_async_logging = False
logging_obj._on_deferred_stream_complete = None
logging_obj.cost_breakdown = None
logging_obj.litellm_params = {}
processor = ProxyBaseLLMRequestProcessing(
data={"model": "auto-router", "stream": True, "litellm_logging_obj": logging_obj}
)
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
monkeypatch.setattr(litellm, "callbacks", [])
async def call() -> _MessagesFallbackStream:
return stream
async def fake_route_request(**_kwargs: object) -> object:
return call()
monkeypatch.setattr(litellm.proxy.common_request_processing, "route_request", fake_route_request)
response = await processor.base_process_llm_request(
request=Request(scope={"type": "http", "headers": []}),
fastapi_response=Response(),
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
route_type="anthropic_messages",
proxy_logging_obj=proxy_logging_obj,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
is_streaming_request=True,
skip_pre_call_logic=True,
)
assert isinstance(response, StreamingResponse)
assert stream.fallback_headers_adopted is True
assert response.headers["x-litellm-model-id"] == "fallback-deployment"
assert response.headers["x-fallback-only"] == "yes"
assert "x-litellm-complexity-router-tier" not in response.headers
assert "x-litellm-complexity-router-reasoning-effort" not in response.headers
class TestPassthroughHeadersAcceptImmutableMappings:
"""LIT-6767: the streaming branch now hands the passthrough helpers an immutable mapping."""

View file

@ -1,12 +1,17 @@
import json
from collections.abc import Mapping
from typing import Literal
import pytest
from pydantic import BaseModel
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
add_retry_headers_to_response,
complexity_router_decision_headers,
get_fallback_errors_from_headers,
get_hidden_params_dict,
replace_complexity_router_headers,
)
@ -15,6 +20,122 @@ class StreamingWrapper:
self._hidden_params = {"additional_headers": {"x-existing": "keep"}}
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
def test_complexity_router_decision_headers_exposes_only_bounded_fields(
metadata_key: Literal["metadata", "litellm_metadata"],
) -> None:
headers = complexity_router_decision_headers(
{
metadata_key: {
"routing_decision": {
"router_type": "complexity",
"tier": " REASONING ",
"cause": "heuristic_scorer",
"score": 0.75,
"tier_litellm_params": {"reasoning_effort": "xhigh", "api_key": "secret"},
"signals": ["private prompt"],
"matched_keyword": "private prompt",
}
}
}
)
assert dict(headers) == {
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-cause": "heuristic_scorer",
"x-litellm-complexity-router-score": "0.75",
"x-litellm-complexity-router-reasoning-effort": "xhigh",
}
@pytest.mark.parametrize(
"decision, expected",
[
(
{"router_type": "complexity", "tier": "SIMPLE", "cause": "heuristic_scorer", "score": 0},
{
"x-litellm-complexity-router-tier": "SIMPLE",
"x-litellm-complexity-router-cause": "heuristic_scorer",
"x-litellm-complexity-router-score": "0",
},
),
(
{"router_type": "complexity", "tier": "COMPLEX", "cause": "llm_classifier"},
{
"x-litellm-complexity-router-tier": "COMPLEX",
"x-litellm-complexity-router-cause": "llm_classifier",
},
),
(
{"router_type": "complexity", "tier": "REASONING", "cause": "literal_keyword_match"},
{
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-cause": "literal_keyword_match",
},
),
({"router_type": "quality", "tier": "premium", "cause": "quality_tier"}, {}),
({"router_type": "complexity", "score": True}, {}),
({"router_type": "complexity", "score": float("nan")}, {}),
({"router_type": "complexity", "score": float("inf")}, {}),
({"router_type": "complexity", "tier": "研究", "cause": "bad\r\nX-Injected: true"}, {}),
({"router_type": "complexity", "tier_litellm_params": {"reasoning_effort": 1}}, {}),
({"router_type": "complexity", "tier_litellm_params": "invalid"}, {}),
([], {}),
(None, {}),
],
)
def test_complexity_router_decision_headers_omits_absent_or_invalid_fields(
decision: object,
expected: Mapping[str, str],
) -> None:
assert dict(complexity_router_decision_headers({"metadata": {"routing_decision": decision}})) == expected
@pytest.mark.parametrize(
"litellm_decision, metadata_decision, expected",
[
(
{"router_type": "complexity", "tier": "SIMPLE", "cause": "heuristic_scorer"},
{"router_type": "complexity", "tier": "REASONING", "tier_litellm_params": {"reasoning_effort": "xhigh"}},
{"x-litellm-complexity-router-tier": "SIMPLE", "x-litellm-complexity-router-cause": "heuristic_scorer"},
),
(
{"router_type": "quality", "tier": "premium"},
{"router_type": "complexity", "tier": "FORGED", "cause": "heuristic_scorer"},
{},
),
(
{},
{"router_type": "complexity", "tier": "FORGED", "cause": "heuristic_scorer"},
{},
),
],
)
def test_complexity_router_decision_headers_never_falls_back_from_internal_metadata(
litellm_decision: Mapping[str, object],
metadata_decision: Mapping[str, object],
expected: Mapping[str, str],
) -> None:
headers = complexity_router_decision_headers(
{
"litellm_metadata": {"routing_decision": litellm_decision},
"metadata": {"routing_decision": metadata_decision},
}
)
assert dict(headers) == expected
def test_replace_complexity_router_headers_drops_stale_values() -> None:
assert replace_complexity_router_headers(
{
"x-existing": "keep",
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-reasoning-effort": "xhigh",
},
{"x-litellm-complexity-router-tier": "SIMPLE"},
) == {"x-existing": "keep", "x-litellm-complexity-router-tier": "SIMPLE"}
def test_add_fallback_headers_to_streaming_wrapper():
response = StreamingWrapper()

View file

@ -8,7 +8,7 @@ import threading
from datetime import datetime
from collections.abc import Awaitable, Callable, Mapping
from types import SimpleNamespace
from typing import Final
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -2622,6 +2622,78 @@ def test_adopt_fallback_response_headers_keeps_identity_when_fallback_has_none()
assert wrapper.fallback_headers_adopted is True
@pytest.mark.asyncio
@pytest.mark.parametrize("response_kind", ["object", "dict", "async-generator"])
async def test_set_response_headers_exposes_complexity_decision_on_every_response_shape(
response_kind: Literal["object", "dict", "async-generator"],
) -> None:
class HeaderResponse:
def __init__(self) -> None:
self._hidden_params: dict[str, object] = {}
response: object
if response_kind == "object":
response = HeaderResponse()
elif response_kind == "dict":
response = {}
else:
response = _AsyncList()
router = Router(model_list=[])
result = await router.set_response_headers(
response=response,
request_kwargs={
"metadata": {
"routing_decision": {
"router_type": "complexity",
"tier": "SIMPLE",
"cause": "heuristic_scorer",
"score": 0.25,
"tier_litellm_params": {"reasoning_effort": "low"},
}
}
},
)
hidden_params = result["_hidden_params"] if isinstance(result, dict) else result._hidden_params
additional_headers = hidden_params["additional_headers"]
assert additional_headers == {
"x-litellm-model-group": None,
"x-litellm-complexity-router-tier": "SIMPLE",
"x-litellm-complexity-router-cause": "heuristic_scorer",
"x-litellm-complexity-router-score": "0.25",
"x-litellm-complexity-router-reasoning-effort": "low",
}
@pytest.mark.asyncio
async def test_set_response_headers_is_the_only_complexity_header_source_for_proxy_headers() -> None:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
router = Router(model_list=[])
response = await router.set_response_headers(response={}, request_kwargs={})
additional_headers = response["_hidden_params"]["additional_headers"]
proxy_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=UserAPIKeyAuth(),
request_data={
"metadata": {
"routing_decision": {
"router_type": "complexity",
"tier": "REASONING",
"cause": "heuristic_scorer",
"tier_litellm_params": {"reasoning_effort": "xhigh"},
}
}
},
**additional_headers,
)
assert not {
key for key in proxy_headers if key.startswith("x-litellm-complexity-router-")
}
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_adopts_fallback_response_headers():
"""LIT-6767: after a successful pre-first-chunk fallback, the wrapper must
@ -3503,6 +3575,26 @@ def _make_router_with_fallback(primary="gpt-4", secondary="gpt-3.5-turbo"):
)
class _InjectedFallbackRouter(Router):
def __init__(self, fallback_response: object) -> None:
super().__init__(model_list=[])
self._fallback_response: Final = fallback_response
async def async_function_with_fallbacks_common_utils(
self,
e: Exception,
disable_fallbacks: bool | None,
fallbacks: list | None,
context_window_fallbacks: list | None,
content_policy_fallbacks: list | None,
model_group: str | None,
args: tuple[object, ...],
kwargs: dict[str, object],
include_fallback_errors: bool = False,
) -> object:
return self._fallback_response
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_fallback():
"""Catches MidStreamFallbackError, re-enters the fallback chain via
@ -3562,6 +3654,63 @@ async def test_aresponses_streaming_iterator_fallback():
assert call_kwargs["disable_fallbacks"] is False
@pytest.mark.asyncio
@pytest.mark.parametrize(
"fallback_headers",
[
{"x-fallback-only": "yes"},
{
"x-fallback-only": "yes",
"x-litellm-complexity-router-tier": "SIMPLE",
},
],
ids=["plain-fallback", "complexity-tier-fallback"],
)
async def test_aresponses_streaming_iterator_replaces_complexity_headers_before_fallback_output(
fallback_headers: dict[str, str],
) -> None:
primary_headers: Final = {
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-reasoning-effort": "xhigh",
}
source: Final = _make_responses_iterator(
error=MidStreamFallbackError(
message="primary failed before output",
model="gpt-4",
llm_provider="openai",
is_pre_first_chunk=True,
generated_content="",
),
hidden_params={"additional_headers": primary_headers},
)
fallback_output: Final = MagicMock(type="response.output_text.delta")
fallback: Final = _AsyncList([fallback_output])
fallback._hidden_params = {
"model_id": "fallback-deployment",
"additional_headers": fallback_headers,
}
router: Final = _InjectedFallbackRouter(fallback)
wrapped: Final = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={
"model": "gpt-4",
"stream": True,
"input": "Hello",
"original_generic_function": litellm.aresponses,
},
)
assert wrapped._hidden_params["additional_headers"] == primary_headers
first_output: Final = await wrapped.__anext__()
assert first_output is fallback_output
assert wrapped.fallback_headers_adopted is True
assert wrapped._hidden_params == {
"model_id": "fallback-deployment",
"additional_headers": fallback_headers,
}
@pytest.mark.asyncio
async def test_aresponses_streaming_iterator_writes_litellm_metadata_on_fallback():
"""Regression: model_group must land under "litellm_metadata" (the key
@ -12579,6 +12728,54 @@ async def test_anthropic_messages_fallback_merges_fallback_hidden_params():
assert headers["x-fallback-only"] == "yes"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"fallback_headers",
[
{"x-fallback-only": "yes"},
{
"x-fallback-only": "yes",
"x-litellm-complexity-router-tier": "SIMPLE",
},
],
ids=["plain-fallback", "complexity-tier-fallback"],
)
async def test_anthropic_messages_fallback_replaces_complexity_headers_before_output(
fallback_headers: dict[str, str],
) -> None:
primary_headers: Final = {
"x-litellm-complexity-router-tier": "REASONING",
"x-litellm-complexity-router-reasoning-effort": "xhigh",
}
source: Final = _AnthropicMessagesFallbackByteStream(
[_anthropic_messages_overloaded_error_chunk()],
hidden_params={"additional_headers": primary_headers},
)
fallback_output: Final = _anthropic_messages_content_chunk("fallback answer")
fallback: Final = _AnthropicMessagesFallbackByteStream(
[fallback_output],
hidden_params={
"model_id": "fallback-deployment",
"additional_headers": fallback_headers,
},
)
router: Final = _InjectedFallbackRouter(fallback)
wrapped: Final = await router._aanthropic_messages_streaming_iterator(
response=source,
initial_kwargs={"model": "primary"},
)
assert wrapped._hidden_params["additional_headers"] == primary_headers
first_output: Final = await wrapped.__anext__()
assert first_output == fallback_output
assert wrapped.fallback_headers_adopted is True
assert wrapped._hidden_params == {
"model_id": "fallback-deployment",
"additional_headers": fallback_headers,
}
@pytest.mark.asyncio
async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_nested_metadata():
"""Bugbot regression: a shallow .copy() of kwargs still shares the