From f9354fb64bed6d99064b778d44b89da061fa3549 Mon Sep 17 00:00:00 2001 From: Tin Date: Fri, 11 Sep 2026 12:35:24 -0700 Subject: [PATCH] feat(proxy): expose complexity routing headers Co-Authored-By: Claude Code (cherry picked from commit c817faec7aaf001baec138ff748fb0583a6d0c4c) --- litellm/proxy/common_request_processing.py | 13 +- litellm/router.py | 24 ++- .../add_retry_fallback_headers.py | 67 +++++- .../proxy/test_common_request_processing.py | 77 +++++++ .../test_add_retry_fallback_headers.py | 121 +++++++++++ tests/test_litellm/test_router.py | 199 +++++++++++++++++- 6 files changed, 487 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e3a2b892721..9f58aaf24f1 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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", diff --git a/litellm/router.py b/litellm/router.py index 597bcfaa20f..8865543badd 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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) diff --git a/litellm/router_utils/add_retry_fallback_headers.py b/litellm/router_utils/add_retry_fallback_headers.py index 3251ea457cf..bc88feef7d2 100644 --- a/litellm/router_utils/add_retry_fallback_headers.py +++ b/litellm/router_utils/add_retry_fallback_headers.py @@ -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 diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 50b26577e5c..efbb5eedad4 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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.""" diff --git a/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py index 3a0deeb13d8..eb7490d76f8 100644 --- a/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py +++ b/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py @@ -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() diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 2a97e92396a..f5e9b2091a0 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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