fix(router): keep per-chunk response_cost None after a streaming fallback

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-10-01 16:43:08 +00:00
parent 6f123b7083
commit 9d416e2958
3 changed files with 89 additions and 1 deletions

View file

@ -2883,6 +2883,7 @@ class Router:
cast(_HiddenParamsHost, fallback_item)._hidden_params = {
**item_hidden_params,
**fallback_hidden_params,
"response_cost": item_hidden_params.get("response_cost"),
"additional_headers": {**item_headers, **fallback_headers},
}

View file

@ -3,7 +3,7 @@ import json
import threading
import uuid
from pathlib import Path
from typing import Final
from typing import Final, Literal
import pytest
import yaml
@ -15,6 +15,10 @@ from integration._support.process import owned_proxy
from integration._support.wire import Reply, wire_server
from openai import OpenAI
from litellm import Router
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
from litellm.types.utils import ModelResponseStream
def frame(identity: str, delta: dict, *, finish: str | None = None) -> bytes:
value: Final = {
@ -45,6 +49,59 @@ def text_stream(identity: str) -> tuple[bytes, ...]:
)
@pytest.mark.parametrize("requested_model", ("primary", "backup"))
def test_chunks_after_a_streaming_fallback_carry_no_response_cost_and_final_usage_is_priced(
requested_model: Literal["primary", "backup"],
) -> None:
identity: Final = "stream-fallback-response-cost"
messages: Final = [{"role": "user", "content": identity}]
primary_chunks: Final = (
b'data: {"error": {"message": "overloaded", "type": "server_error", "code": 500}}\n\n',
b"data: [DONE]\n\n",
)
with (
wire_server(lambda request: Reply(content_type="text/event-stream", chunks=primary_chunks)) as primary,
wire_server(lambda request: Reply(content_type="text/event-stream", chunks=text_stream(identity))) as backup,
):
router: Final = Router(
model_list=[
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "synthetic-fallback-key",
"api_base": server.url + "/v1",
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
},
}
for name, server in (("primary", primary), ("backup", backup))
],
fallbacks=[{"primary": ["backup"]}],
num_retries=0,
)
async def collect_async_chunks() -> tuple[ModelResponseStream, ...]:
stream: Final = await router.acompletion(
model=requested_model,
messages=messages,
stream=True,
stream_options={"include_usage": True},
)
return tuple([chunk async for chunk in stream])
chunks: Final = asyncio.run(collect_async_chunks())
observed_costs: Final = tuple(get_hidden_params_dict(chunk).get("response_cost") for chunk in chunks)
expected_costs: Final = (None, None, None, None)
assert observed_costs == expected_costs, f"observed per-chunk response costs: {observed_costs}"
assert chunks[-1].usage.cost == pytest.approx(11 * 0.001 + 4 * 0.002)
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "Hello 雪 café"
assert len(primary.drain()) == (1 if requested_model == "primary" else 0)
assert len(backup.drain()) == 1
@pytest.mark.covers("other.streaming.byte_partitions.preserve_text_identity_and_usage")
def test_generated_tcp_partitions_preserve_unicode_text_identity_and_final_usage() -> None:
import litellm

View file

@ -2966,6 +2966,36 @@ def test_completion_streaming_iterator_preserves_response_headers():
assert result._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req-provider-sync"
@pytest.mark.parametrize("response_cost", (None, 0.00015))
def test_apply_fallback_hidden_params_to_item_keeps_the_chunks_own_response_cost(
response_cost: float | None,
) -> None:
chunk: Final = litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": "chunk"}}])
chunk._hidden_params = {
"response_cost": response_cost,
"model_id": "failed-deployment",
"additional_headers": {"x-a": "1"},
}
Router._apply_fallback_hidden_params_to_item(
chunk,
(
{
"response_cost": 0.0,
"model_id": "fallback-deployment",
"additional_headers": {"x-b": "2"},
},
{"x-b": "2"},
),
)
assert chunk._hidden_params == {
"response_cost": response_cost,
"model_id": "fallback-deployment",
"additional_headers": {"x-a": "1", "x-b": "2"},
}
def test_adopt_fallback_response_headers_replaces_rather_than_merges():
"""LIT-6767: direct unit for FallbackAwareStreamWrapper.adopt_fallback_response_headers.