fix: keep hidden params writes on duck-typed responses

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-10-05 23:00:06 +00:00
parent 1168d3d362
commit 03752c17f9
14 changed files with 236 additions and 90 deletions

View file

@ -35,6 +35,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
)
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params
from litellm.types.caching import EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding
from litellm.types.integrations.custom_logger import converted_stream_requested
from litellm.types.llms.openai import ResponsesAPIResponse
@ -172,6 +173,15 @@ def _request_cache_key(request_kwargs: Mapping[str, Any]) -> str | None:
return request_kwargs.get("cache_key", None)
def _set_cached_hidden_param(response: object, key: str, value: object) -> None:
if isinstance(response, TextCompletionResponse):
setattr(response.hidden_params, key, value)
return
hidden_params: Final = get_hidden_params(response)
if hidden_params is not None:
hidden_params[key] = value
class _CachedEmbeddingRecord(BaseModel):
model_config = ConfigDict(frozen=True)
@ -331,8 +341,7 @@ class LLMCachingHandler:
or self.request_kwargs.get("cache_key")
or litellm.cache.get_cache_key(**self.request_kwargs)
)
if hasattr(cached_result, "_hidden_params"):
cached_result._hidden_params["cache_key"] = cache_key
_set_cached_hidden_param(cached_result, "cache_key", cache_key)
return CachingHandlerResponse(cached_result=cached_result)
elif (
call_type == CallTypes.aembedding.value
@ -448,8 +457,7 @@ class LLMCachingHandler:
or self.request_kwargs.get("cache_key")
or litellm.cache.get_cache_key(**self.request_kwargs)
)
if hasattr(cached_result, "_hidden_params"):
cached_result._hidden_params["cache_key"] = cache_key
_set_cached_hidden_param(cached_result, "cache_key", cache_key)
return CachingHandlerResponse(cached_result=cached_result)
return CachingHandlerResponse(cached_result=cached_result)
@ -969,12 +977,7 @@ class LLMCachingHandler:
)
response_obj: Final = ResponsesAPIResponse(**cached_result)
if (
hasattr(response_obj, "_hidden_params")
and response_obj.hidden_params is not None
and isinstance(response_obj.hidden_params, dict)
):
response_obj.hidden_params["cache_hit"] = True
_set_cached_hidden_param(response_obj, "cache_hit", True)
if _stream_replay_requested(kwargs):
cached_result = CachedResponsesAPIStreamingIterator(
@ -986,12 +989,7 @@ class LLMCachingHandler:
else:
cached_result = response_obj
if (
hasattr(cached_result, "_hidden_params")
and cached_result._hidden_params is not None
and isinstance(cached_result._hidden_params, dict)
):
cached_result._hidden_params["cache_hit"] = True
_set_cached_hidden_param(cached_result, "cache_hit", True)
#########################################################
# Add final timing metrics to the cached result

View file

@ -37,6 +37,7 @@ from litellm.responses.sse_output_recovery import (
record_output_text_chunk,
)
from litellm.responses.utils import ResponsesAPIRequestUtils, normalize_responses_api_stream_options
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionReasoningItem,
@ -994,18 +995,19 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
# which contain important provider information like x-request-id
raw_response_hidden_params: Final = getattr(raw_response, "_hidden_params", {})
if raw_response_hidden_params:
if not hasattr(model_response, "_hidden_params") or model_response.hidden_params is None:
model_response.hidden_params = {}
model_response_hidden_params: Final = get_or_create_hidden_params(model_response)
# Merge the raw_response hidden params with model_response hidden params
# Preserve existing keys in model_response but add/override with raw_response params
for key, value in raw_response_hidden_params.items():
if key == "additional_headers" and key in model_response.hidden_params:
# Merge additional_headers to preserve both sets
existing_additional_headers = model_response.hidden_params.get("additional_headers", {})
merged_headers = {**value, **existing_additional_headers}
model_response.hidden_params[key] = merged_headers
if key == "additional_headers" and key in model_response_hidden_params:
existing_additional_headers = model_response_hidden_params.get("additional_headers", {})
merged_headers = {
**cast("dict[str, object]", value),
**cast("dict[str, object]", existing_additional_headers),
}
model_response_hidden_params[key] = merged_headers
else:
model_response.hidden_params[key] = value
model_response_hidden_params[key] = value
return model_response

View file

@ -12,6 +12,7 @@ from litellm.llms.anthropic.pass_through.utils import (
is_reasoning_auto_summary_enabled,
prompt_cache_key_from_user_id,
)
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params
# OpenAI has a 64-character limit for function/tool names
# Anthropic does not have this limit, so we need to truncate long names
@ -1739,10 +1740,13 @@ class LiteLLMAnthropicMessagesAdapter:
)
if getattr(response, "usage", None) is not None:
litellm_usage_chunk: Usage | None = response.usage
elif hasattr(response, "_hidden_params") and "usage" in response.hidden_params:
litellm_usage_chunk = response.hidden_params["usage"]
else:
litellm_usage_chunk = None
response_hidden_params: Final = get_hidden_params(response)
litellm_usage_chunk = (
response_hidden_params["usage"]
if response_hidden_params is not None and "usage" in response_hidden_params
else None
)
if litellm_usage_chunk is not None:
usage_delta = self._translate_openai_usage_to_anthropic_usage_delta(litellm_usage_chunk)
else:

View file

@ -34,6 +34,7 @@ from litellm.llms.azure_ai.agents.transformation import (
AzureAIAgentsConfig,
AzureAIAgentsError,
)
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionAnnotationURLCitation,
@ -231,9 +232,8 @@ class AzureAIAgentsHandler:
model_response.model = model
# Store thread_id for conversation continuity
if not hasattr(model_response, "_hidden_params") or model_response.hidden_params is None:
model_response.hidden_params = {}
model_response.hidden_params["thread_id"] = thread_id
model_response_hidden_params: Final = get_or_create_hidden_params(model_response)
model_response_hidden_params["thread_id"] = thread_id
# Estimate token usage
try:

View file

@ -1,5 +1,5 @@
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Any, Final, Literal, cast
import httpx
from typing_extensions import ReadOnly, TypedDict
@ -9,6 +9,7 @@ from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.secret_managers.main import get_secret_str
from litellm.types.containers.main import (
ContainerCreateOptionalRequestParams,
@ -165,11 +166,10 @@ class OpenAIContainerConfig(BaseContainerConfig):
provider="openai",
)
if not hasattr(container_obj, "_hidden_params") or container_obj.hidden_params is None:
container_obj.hidden_params = {}
if "additional_headers" not in container_obj.hidden_params:
container_obj.hidden_params["additional_headers"] = {}
container_obj.hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = container_cost
container_hidden_params: Final = get_or_create_hidden_params(container_obj)
container_hidden_params.setdefault("additional_headers", {})
container_additional_headers: Final = cast("dict[str, object]", container_hidden_params["additional_headers"])
container_additional_headers["llm_provider-x-litellm-response-cost"] = container_cost
return container_obj

View file

@ -17,6 +17,7 @@ from litellm.llms.base_llm.files.storage_backend_factory import get_storage_back
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.types.llms.openai import OpenAIFileObject, OpenAIFilesPurpose
from litellm.types.utils import ExtractedFileData, SpecialEnums
@ -170,9 +171,8 @@ class StorageBackendFileService:
)
# Store storage metadata in hidden params
if not hasattr(file_object, "_hidden_params") or file_object.hidden_params is None:
file_object.hidden_params = {}
file_object.hidden_params.update(
file_object_hidden_params: Final = get_or_create_hidden_params(file_object)
file_object_hidden_params.update(
{
"storage_backend": target_storage,
"storage_url": storage_url,

View file

@ -17,6 +17,7 @@ from litellm.llms.cohere.common_utils import (
)
from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
@ -114,9 +115,8 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler):
)
# Set the calculated cost in _hidden_params to prevent recalculation
if not hasattr(litellm_model_response, "_hidden_params"):
litellm_model_response.hidden_params = {}
litellm_model_response.hidden_params["response_cost"] = response_cost
litellm_model_response_hidden_params: Final = get_or_create_hidden_params(litellm_model_response)
litellm_model_response_hidden_params["response_cost"] = response_cost
kwargs["response_cost"] = response_cost
kwargs["model"] = model

View file

@ -28,6 +28,7 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params, get_or_create_hidden_params
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
EndpointType,
@ -414,7 +415,9 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
model=model,
custom_llm_provider=custom_llm_provider,
)
litellm_model_response.hidden_params["response_cost"] = response_cost
embedding_hidden_params: Final = get_hidden_params(litellm_model_response)
if embedding_hidden_params is not None:
embedding_hidden_params["response_cost"] = response_cost
elif is_image_generation:
# Handle image generation cost calculation
response_cost = OpenAIPassthroughLoggingHandler._calculate_image_generation_cost(
@ -433,9 +436,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
model=model,
)
# Set the calculated cost in _hidden_params to prevent recalculation
if not hasattr(litellm_model_response, "_hidden_params"):
litellm_model_response.hidden_params = {}
litellm_model_response.hidden_params["response_cost"] = response_cost
image_generation_hidden_params: Final = get_or_create_hidden_params(litellm_model_response)
image_generation_hidden_params["response_cost"] = response_cost
elif is_image_editing:
# Handle image editing cost calculation
response_cost = OpenAIPassthroughLoggingHandler._calculate_image_editing_cost(
@ -454,9 +456,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
model=model,
)
# Set the calculated cost in _hidden_params to prevent recalculation
if not hasattr(litellm_model_response, "_hidden_params"):
litellm_model_response.hidden_params = {}
litellm_model_response.hidden_params["response_cost"] = response_cost
image_editing_hidden_params: Final = get_or_create_hidden_params(litellm_model_response)
image_editing_hidden_params["response_cost"] = response_cost
elif is_responses:
# Responses-API cost tracking — see
# `_build_responses_api_response_and_cost` for why this needs

View file

@ -17,6 +17,7 @@ from litellm.responses.litellm_completion_transformation.transformation import (
)
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.router_utils.add_retry_fallback_headers import get_or_create_hidden_params
from litellm.types.llms.openai import (
PART_UNION_TYPES,
BaseLiteLLMOpenAIResponseObject,
@ -649,9 +650,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
),
)
if response is not None and self._accumulated_provider_specific_fields:
if not hasattr(response, "_hidden_params") or response.hidden_params is None:
response.hidden_params = {}
response.hidden_params.setdefault("provider_specific_fields", {}).update(
response_hidden_params: Final = get_or_create_hidden_params(response)
response_hidden_params.setdefault("provider_specific_fields", {}).update(
self._accumulated_provider_specific_fields
)
return response

View file

@ -146,7 +146,6 @@ from litellm.router_strategy.tag_based_routing import (
)
from litellm.router_utils.access_windows import access_windows_config_error, filter_reserved_deployments
from litellm.router_utils.add_retry_fallback_headers import (
_HiddenParamsHost,
add_fallback_headers_to_response,
add_retry_headers_to_response,
apply_quality_router_decision_headers,
@ -154,10 +153,12 @@ from litellm.router_utils.add_retry_fallback_headers import (
apply_response_model_id,
complexity_router_decision_headers,
ensure_response_additional_headers,
get_hidden_params,
get_hidden_params_dict,
prepare_response_for_header_attachment,
replace_complexity_router_headers,
response_total_token_count,
set_hidden_params,
)
from litellm.router_utils.auto_router_model_naming import (
AUTO_ROUTER_MODEL_PREFIX,
@ -2900,20 +2901,25 @@ class Router:
fallback_item: object,
prepared_fallback_hidden_params: tuple[dict[str, object], dict[str, object]],
) -> None:
if fallback_item is None or not hasattr(fallback_item, "_hidden_params"):
if fallback_item is None:
return
fallback_hidden_params, fallback_headers = prepared_fallback_hidden_params
item_hidden_params: Final = get_hidden_params_dict(fallback_item)
item_hidden_params: Final = get_hidden_params(fallback_item)
if item_hidden_params is None:
return
item_headers = item_hidden_params.get("additional_headers")
if not isinstance(item_headers, dict):
item_headers = {}
cast(_HiddenParamsHost, fallback_item).hidden_params = {
**item_hidden_params,
**fallback_hidden_params,
"additional_headers": {**item_headers, **fallback_headers},
}
set_hidden_params(
fallback_item,
{
**item_hidden_params,
**fallback_hidden_params,
"additional_headers": {**item_headers, **fallback_headers},
},
)
async def _acompletion_streaming_iterator(
self,
@ -2949,8 +2955,9 @@ class Router:
if isinstance(inner_chunks, list):
self.chunks = inner_chunks
# Preserve hidden params (including litellm_overhead_time_ms) from original response
if hasattr(model_response, "_hidden_params"):
self._hidden_params = model_response._hidden_params.copy()
model_response_hidden_params: Final = get_hidden_params(model_response)
if model_response_hidden_params is not None:
self._hidden_params = model_response_hidden_params.copy()
def __aiter__(self):
return self
@ -3565,8 +3572,9 @@ class Router:
_response_headers=getattr(model_response, "_response_headers", None),
)
self._sync_generator = sync_generator
if hasattr(model_response, "_hidden_params"):
self._hidden_params = model_response._hidden_params.copy()
model_response_hidden_params: Final = get_hidden_params(model_response)
if model_response_hidden_params is not None:
self._hidden_params = model_response_hidden_params.copy()
def __iter__(self):
return self
@ -4393,7 +4401,8 @@ class Router:
if result is not None:
# Return the first successful result
result.hidden_params["fastest_response_batch_completion"] = True
if (result_hidden_params := get_hidden_params(result)) is not None:
result_hidden_params["fastest_response_batch_completion"] = True
return result
# If we exit the loop without returning, all tasks failed
@ -4462,8 +4471,10 @@ class Router:
if make_request:
try:
_response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs)
_response.hidden_params.setdefault("additional_headers", {})
_response.hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True})
response_hidden_params: Final = get_hidden_params(_response)
if response_hidden_params is not None:
response_hidden_params.setdefault("additional_headers", {})
response_hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True})
return _response
except Exception as e:
setattr(e, "priority", priority)
@ -4522,11 +4533,10 @@ class Router:
if make_request:
try:
_response: Final = await original_function(*args, **kwargs)
if isinstance(_response.hidden_params, dict):
_response.hidden_params.setdefault("additional_headers", {})
_response.hidden_params["additional_headers"].update(
{"x-litellm-request-prioritization-used": True}
)
response_hidden_params: Final = get_hidden_params(_response)
if response_hidden_params is not None:
response_hidden_params.setdefault("additional_headers", {})
response_hidden_params["additional_headers"].update({"x-litellm-request-prioritization-used": True})
return _response
except Exception as e:
setattr(e, "priority", priority)
@ -6151,7 +6161,9 @@ class Router:
healthy_deployments=healthy_deployments, responses=responses
)
returned_response: Final = cast(OpenAIFileObject, responses[0])
returned_response.hidden_params["model_file_id_mapping"] = model_file_id_mapping
returned_response_hidden_params: Final = get_hidden_params(returned_response)
if returned_response_hidden_params is not None:
returned_response_hidden_params["model_file_id_mapping"] = model_file_id_mapping
return returned_response
except Exception as e:
verbose_router_logger.exception(

View file

@ -2,7 +2,7 @@ import json
import math
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final, Protocol, TypedDict, cast
from typing import Any, Final, TypedDict, cast
from pydantic import BaseModel, TypeAdapter, ValidationError
@ -14,16 +14,7 @@ class FallbackErrorInfo(TypedDict):
code: str | None
class _HiddenParamsHost(Protocol):
_hidden_params: dict[str, object]
@property
def hidden_params(self) -> dict[str, object]: ... # mutable-ok: API requires mutation
@hidden_params.setter
def hidden_params(self, hidden_params: dict[str, object]) -> None: ... # mutable-ok: API requires mutation
_HIDDEN_PARAMS_ATTR: Final = "_hidden_params"
_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-"
@ -215,13 +206,32 @@ def get_hidden_params_dict(
return hidden_params
def get_hidden_params(obj: object) -> dict[str, object] | None:
hidden_params: Final = (
obj.get(_HIDDEN_PARAMS_ATTR) if isinstance(obj, dict) else getattr(obj, _HIDDEN_PARAMS_ATTR, None)
)
return hidden_params if isinstance(hidden_params, dict) else None
def set_hidden_params(obj: object, hidden_params: dict[str, object]) -> None:
if isinstance(obj, dict):
obj[_HIDDEN_PARAMS_ATTR] = hidden_params
else:
setattr(obj, _HIDDEN_PARAMS_ATTR, hidden_params)
def get_or_create_hidden_params(obj: object) -> dict[str, object]:
hidden_params: Final = get_hidden_params(obj)
if hidden_params is not None:
return hidden_params
created_hidden_params: Final[dict[str, object]] = {}
set_hidden_params(obj, created_hidden_params)
return created_hidden_params
def _write_hidden_params(response: object, hidden_params: dict[str, object]) -> None:
if isinstance(response, dict):
response["_hidden_params"] = hidden_params
elif hasattr(response, "_hidden_params"):
host: Final = cast(_HiddenParamsHost, response)
if get_hidden_params_dict(response) is not hidden_params:
host.hidden_params = hidden_params
if isinstance(response, dict) or hasattr(response, _HIDDEN_PARAMS_ATTR):
set_hidden_params(response, hidden_params)
def _ensure_additional_headers_dict(

View file

@ -2140,6 +2140,49 @@ async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monke
assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key
@pytest.mark.asyncio
async def test_text_completion_cache_hit_records_cache_key_in_hidden_params(monkeypatch):
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
async def atext_completion(**kwargs):
return None
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {"model": "gpt-5.4", "prompt": "hello", "caching": True}
await litellm.cache.async_add_cache(
litellm.TextCompletionResponse(
id="cached-text-response",
choices=[litellm.utils.TextChoices(text="cached")],
model="gpt-5.4",
),
**kwargs,
)
handler = LLMCachingHandler(
original_function=atext_completion,
request_kwargs=kwargs,
start_time=datetime.now(),
)
logging_obj = _build_logging_obj(CallTypes.atext_completion.value, stream=False)
logging_obj.async_success_handler = AsyncMock()
hit = await handler._async_get_cache(
model="gpt-5.4",
original_function=atext_completion,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.atext_completion.value,
kwargs=kwargs,
args=(),
)
assert hit.cached_result is not None
assert isinstance(hit.cached_result, litellm.TextCompletionResponse)
assert handler.preset_cache_key is not None
assert hit.cached_result.hidden_params.get("cache_key") == handler.preset_cache_key
@pytest.mark.asyncio
async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch):
import litellm

View file

@ -9,11 +9,15 @@ from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
add_retry_headers_to_response,
complexity_router_decision_headers,
ensure_response_additional_headers,
get_fallback_errors_from_headers,
get_hidden_params,
get_hidden_params_dict,
replace_complexity_router_headers,
set_hidden_params,
)
from litellm.types.decisions import DecisionsResponse
from litellm.types.utils import ModelResponse
class StreamingWrapper:
@ -219,6 +223,38 @@ def test_get_hidden_params_dict_with_no_hidden_params():
assert get_hidden_params_dict(PlainResponse()) == {}
def test_get_and_set_hidden_params_on_plain_object() -> None:
class PlainResponse:
def __init__(self) -> None:
self._hidden_params = {"existing": True}
response: Final = PlainResponse()
stored: Final = get_hidden_params(response)
assert stored is response._hidden_params
replacement: Final = {"replacement": True}
set_hidden_params(response, replacement)
assert response._hidden_params is replacement
assert get_hidden_params(response) is replacement
def test_get_hidden_params_preserves_model_response_identity() -> None:
response: Final = ModelResponse()
assert get_hidden_params(response) is response.hidden_params
def test_set_hidden_params_replaces_frozen_decisions_response_private_attr() -> None:
response: Final = DecisionsResponse(model="decider", answers={}, usage=None)
replacement: Final = {"replacement": True}
set_hidden_params(response, replacement)
assert response.hidden_params is replacement
assert response._hidden_params is replacement
def test_add_fallback_headers_when_no_existing_additional_headers():
class NoHeadersWrapper:
def __init__(self):
@ -240,6 +276,15 @@ def test_add_fallback_headers_to_frozen_decisions_response() -> None:
assert response.hidden_params["additional_headers"] == {"x-litellm-attempted-fallbacks": 1}
def test_ensure_response_additional_headers_updates_frozen_decisions_response() -> None:
response: Final = DecisionsResponse(model="decider", answers={}, usage=None)
additional_headers: Final = ensure_response_additional_headers(response)
assert additional_headers == {}
assert response.hidden_params["additional_headers"] is additional_headers
def test_add_fallback_headers_returns_none_when_response_is_none():
result = add_fallback_headers_to_response(response=None, attempted_fallbacks=1)
assert result is None

View file

@ -55,6 +55,37 @@ def test_apply_fallback_hidden_params_copies_from_fallback_response():
}
def test_apply_fallback_hidden_params_updates_plain_duck_chunk():
class PlainChunk:
def __init__(self) -> None:
self._hidden_params = {
"additional_headers": {"x-existing-chunk-header": "keep"},
"model_id": "chunk-model-id",
}
chunk: Final = PlainChunk()
fallback_response = {
"_hidden_params": {
"additional_headers": {
"x-litellm-attempted-fallbacks": 1,
},
"api_base": "https://fallback.example",
}
}
Router._apply_fallback_hidden_params_to_item(
fallback_item=chunk,
prepared_fallback_hidden_params=Router._prepare_fallback_hidden_params(fallback_response),
)
assert chunk._hidden_params["api_base"] == "https://fallback.example"
assert chunk._hidden_params["model_id"] == "chunk-model-id"
assert chunk._hidden_params["additional_headers"] == {
"x-existing-chunk-header": "keep",
"x-litellm-attempted-fallbacks": 1,
}
def _two_group_fallback_router() -> Router:
return litellm.Router(
model_list=[