mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
parent
8e4f2abb40
commit
dab7f6a86a
6 changed files with 487 additions and 14 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue