Merge pull request #38656 from aaaaaandrew/litellm_preserve_stream_selected_model

fix(streaming): preserve provider model for cost calculation
This commit is contained in:
Mateo Wang 2026-08-28 13:00:35 -07:00 committed by GitHub
commit d4f6b4491d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 452 additions and 13 deletions

View file

@ -2,7 +2,7 @@
## File for 'response_cost' calculation in Logging
import logging
import time
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal, cast
@ -739,6 +739,13 @@ def _get_provider_for_cost_calc(
return custom_llm_provider
def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None:
if not isinstance(hidden_params, Mapping):
return None
value: Final[object] = hidden_params.get(key)
return value if isinstance(value, str) and value else None
def _select_model_name_for_cost_calc(
model: str | None,
completion_response: object | None,
@ -755,7 +762,6 @@ def _select_model_name_for_cost_calc(
"""
return_model: str | None = None
region_name: str | None = None
custom_llm_provider = _get_provider_for_cost_calc(model=model, custom_llm_provider=custom_llm_provider)
completion_response_model: str | None = None
@ -765,6 +771,14 @@ def _select_model_name_for_cost_calc(
elif isinstance(completion_response, dict):
completion_response_model = completion_response.get("model", None)
hidden_params: Final[dict | None] = getattr(completion_response, "_hidden_params", None)
provider_response_model: Final = _get_hidden_str_for_cost_calc(hidden_params, "provider_response_model")
explicit_pricing: Final = custom_pricing is True or base_model is not None
priced_from_response: Final = provider_response_model is not None or completion_response_model is not None
region_name: Final = (
_get_hidden_str_for_cost_calc(hidden_params, "region_name")
if not explicit_pricing and priced_from_response
else None
)
if custom_pricing is True:
if router_model_id is not None and router_model_id in litellm.model_cost:
@ -780,14 +794,12 @@ def _select_model_name_for_cost_calc(
else:
return_model = model
elif base_model is not None:
return_model = base_model
elif base_model is not None or provider_response_model is not None:
return_model = base_model if base_model is not None else provider_response_model
elif completion_response_model is None and hidden_params is not None:
if hidden_params.get("model", None) is not None and len(hidden_params["model"]) > 0:
return_model = hidden_params.get("model", model)
elif hidden_params is not None and hidden_params.get("region_name", None) is not None:
region_name = hidden_params.get("region_name", None)
if return_model is None and completion_response_model is not None:
return_model = completion_response_model

View file

@ -239,6 +239,22 @@ class ChunkProcessor:
model_response._hidden_params = chunk.get("_hidden_params", {})
return model_response
@staticmethod
def _get_provider_response_model(
chunks: Sequence["_BaseChunk"],
first_chunk_model: str,
) -> str | None:
models: Final = tuple(
model
for chunk in chunks
if isinstance((hidden_params := chunk.get("_hidden_params")), Mapping)
if isinstance((model := hidden_params.get("provider_response_model")), str) and model
)
return next(
(model for model in models if model != first_chunk_model),
models[0] if models else None,
)
@staticmethod
def apply_provider_assembled_streaming_metadata(
response: ModelResponse,
@ -360,6 +376,15 @@ class ChunkProcessor:
)
response = self.update_model_response_with_hidden_params(model_response=response, chunk=chunk)
provider_response_model: Final = self._get_provider_response_model(
chunks,
first_chunk_model,
)
if provider_response_model is not None:
response._hidden_params = dict( # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter
response._hidden_params, # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params getter
provider_response_model=provider_response_model,
)
return response
@staticmethod

View file

@ -187,17 +187,42 @@ class _ParsedChunkHiddenParams(BaseModel):
provider_specific_fields: Mapping[str, object] | None = None
def _provider_hidden_params(chunk: object) -> Mapping[str, object] | None:
hidden: Final[object] = getattr(chunk, "_hidden_params", None)
def _provider_response_model(chunk: object) -> str | None:
model: Final[object] = chunk.get("model") if isinstance(chunk, Mapping) else getattr(chunk, "model", None)
return model if isinstance(model, str) and model else None
def _parsed_provider_hidden_params(hidden: object) -> _ParsedChunkHiddenParams | None:
if not isinstance(hidden, dict):
return None
try:
parsed: Final = _ParsedChunkHiddenParams.model_validate(hidden)
return _ParsedChunkHiddenParams.model_validate(hidden)
except ValidationError:
return None
if not parsed.provider_specific_fields:
return None
return MappingProxyType({"provider_specific_fields": dict(parsed.provider_specific_fields)})
def _provider_hidden_params(
chunk: object,
provider_response_model: str | None,
) -> Mapping[str, object] | None:
hidden: Final[object] = getattr(chunk, "_hidden_params", None)
parsed: Final = _parsed_provider_hidden_params(hidden)
provider_specific_fields: Final[object | None] = (
dict(parsed.provider_specific_fields) # mutable-ok: stream assembly merges provider metadata into this dict
if parsed is not None and parsed.provider_specific_fields
else None
)
params: Final[Mapping[str, object]] = MappingProxyType(
{
key: value
for key, value in (
("provider_response_model", provider_response_model),
("provider_specific_fields", provider_specific_fields),
)
if value is not None
}
)
return params or None
class CustomStreamWrapper:
@ -229,6 +254,7 @@ class CustomStreamWrapper:
self.thinking_content = ""
self.system_fingerprint: str | None = None
self._provider_response_model: str | None = None
self.received_finish_reason: str | None = None
self.intermittent_finish_reason: str | None = None # finish reasons that show up mid-stream
self.special_tokens = [
@ -1524,7 +1550,12 @@ class CustomStreamWrapper:
def chunk_creator(self, chunk: Any):
if hasattr(chunk, "id"):
self.response_id = chunk.id
model_response = self.model_response_creator(hidden_params=_provider_hidden_params(chunk))
provider_response_model: Final = _provider_response_model(chunk)
if provider_response_model is not None:
self._provider_response_model = provider_response_model
model_response = self.model_response_creator(
hidden_params=_provider_hidden_params(chunk, self._provider_response_model)
)
response_obj: dict[str, Any] = {}
try:
# return this for all models

View file

@ -4478,6 +4478,175 @@ def test_chunk_creator_preserves_hidden_provider_specific_fields_from_parsed_chu
assert result is not None
assert result._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"}
assembled = litellm.stream_chunk_builder(chunks=[result])
assert assembled is not None
assert assembled._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"}
def test_chunk_creator_keeps_provider_model_private_across_stream():
from litellm.router_utils.add_retry_fallback_headers import (
get_hidden_params_dict,
)
wrapper = CustomStreamWrapper(
completion_stream=None,
model="requested-route",
logging_obj=MagicMock(),
custom_llm_provider="openai",
)
selected_chunk = ModelResponseStream(
id="chunk-1",
model="selected-model",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content="hello"),
)
],
)
terminal_chunk = ModelResponseStream(
id="chunk-1",
model=None,
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
],
)
first_result = wrapper.chunk_creator(chunk=selected_chunk)
terminal_result = wrapper.chunk_creator(chunk=terminal_chunk)
assert first_result is not None
assert terminal_result is not None
assert first_result.model == "requested-route"
assert terminal_result.model == "requested-route"
assert (
get_hidden_params_dict(first_result)["provider_response_model"]
== "selected-model"
)
assert (
get_hidden_params_dict(terminal_result)["provider_response_model"]
== "selected-model"
)
assembled = litellm.stream_chunk_builder(chunks=[first_result, terminal_result])
assert assembled is not None
assert assembled.model == "requested-route"
assert (
get_hidden_params_dict(assembled)["provider_response_model"]
== "selected-model"
)
def test_assembled_stream_uses_later_provider_model_for_cost(
monkeypatch: pytest.MonkeyPatch,
):
from litellm.router_utils.add_retry_fallback_headers import (
get_hidden_params_dict,
)
selected_model_info = {
"input_cost_per_token": 0.000002,
"output_cost_per_token": 0.000004,
"litellm_provider": "azure",
}
monkeypatch.setitem(
litellm.model_cost,
"azure/gpt-4.1-nano-2025-04-14",
selected_model_info,
)
monkeypatch.setitem(
litellm.model_cost,
"azure/azure-model-router",
{
"input_cost_per_token": 0.00002,
"output_cost_per_token": 0.00004,
"litellm_provider": "azure",
},
)
logging_obj = MagicMock()
logging_obj.model_call_details = {"custom_llm_provider": "azure"}
wrapper = CustomStreamWrapper(
completion_stream=None,
model="azure-model-router",
logging_obj=logging_obj,
custom_llm_provider="azure",
)
router_chunk = ModelResponseStream(
id="chunk-1",
model="azure-model-router",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content="hello "),
)
],
)
selected_chunk = ModelResponseStream(
id="chunk-1",
model="gpt-4.1-nano-2025-04-14",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content="world"),
)
],
)
terminal_chunk = ModelResponseStream(
id="chunk-1",
model="azure-model-router",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(),
)
],
)
router_result = wrapper.chunk_creator(chunk=router_chunk)
selected_result = wrapper.chunk_creator(chunk=selected_chunk)
terminal_result = wrapper.chunk_creator(chunk=terminal_chunk)
assert router_result is not None
assert selected_result is not None
assert terminal_result is not None
assert (
get_hidden_params_dict(router_result)["provider_response_model"]
== "azure-model-router"
)
assert (
get_hidden_params_dict(selected_result)["provider_response_model"]
== "gpt-4.1-nano-2025-04-14"
)
assert (
get_hidden_params_dict(terminal_result)["provider_response_model"]
== "azure-model-router"
)
assembled = litellm.stream_chunk_builder(
chunks=[router_result, selected_result, terminal_result]
)
assert assembled is not None
assert assembled.model == "gpt-4.1-nano-2025-04-14"
assert (
get_hidden_params_dict(assembled)["provider_response_model"]
== "gpt-4.1-nano-2025-04-14"
)
assembled.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
assert litellm.completion_cost(
completion_response=assembled,
custom_llm_provider="azure",
) == pytest.approx(
10 * selected_model_info["input_cost_per_token"]
+ 5 * selected_model_info["output_cost_per_token"]
)
@pytest.mark.asyncio

View file

@ -1719,3 +1719,82 @@ def test_in_schema_unsupported_params_still_raise():
store=True,
)
assert "store" not in optional_params
def test_streaming_preserves_selected_model_for_private_accounting():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
requested_route = (
"accounts/fireworks/routers/firerouter/"
"kimi-k3/deepseek-v4-pro-0813/deepseek-v4-flash-0731"
)
selected_model = "deepseek-v4-flash-0731"
sse_lines = [
"data: "
+ json.dumps(
{
"id": "stream-1",
"object": "chat.completion.chunk",
"created": 1,
"model": selected_model,
"choices": [
{
"index": 0,
"delta": {"role": "assistant", "content": "Hi"},
}
],
}
),
"data: "
+ json.dumps(
{
"id": "stream-1",
"object": "chat.completion.chunk",
"created": 1,
"model": selected_model,
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 1,
"total_tokens": 6,
},
}
),
"data: [DONE]",
]
raw_response = MagicMock()
raw_response.status_code = 200
raw_response.headers = {}
raw_response.iter_lines = lambda: iter(sse_lines)
client = HTTPHandler()
with patch.object(client, "post", return_value=raw_response):
stream = litellm.completion(
model=f"fireworks_ai/{requested_route}",
messages=[{"role": "user", "content": "hi"}],
stream=True,
api_key="test-key",
client=client,
)
chunks = list(stream)
assert chunks
assert {chunk.model for chunk in chunks} == {requested_route}
assert {
chunk._hidden_params.get("provider_response_model") for chunk in chunks
} == {selected_model}
assembled = litellm.stream_chunk_builder(chunks=chunks)
assert assembled is not None
assert assembled.model == requested_route
assert assembled._hidden_params["provider_response_model"] == selected_model
selected_model_info = litellm.model_cost[f"fireworks_ai/{selected_model}"]
expected_cost = (
5 * selected_model_info["input_cost_per_token"]
+ selected_model_info["output_cost_per_token"]
)
assert litellm.completion_cost(
completion_response=assembled,
custom_llm_provider="fireworks_ai",
) == pytest.approx(expected_cost)

View file

@ -4106,6 +4106,53 @@ def test_select_model_name_strips_duplicated_region_segment(_local_model_cost_ma
assert selected == "bedrock/us-east-1/anthropic.claude-v2:1"
def _bedrock_response_with_private_model(model: str, region_name: str) -> litellm.ModelResponse:
response = litellm.ModelResponse(
id="x",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
model=model,
)
response._hidden_params = {"provider_response_model": model, "region_name": region_name}
return response
def test_select_model_name_applies_region_to_private_provider_response_model(_local_model_cost_map):
"""A Bedrock stream carries its requested model as the private provider model and must keep the
request's region in the cost key, exactly as the same request does without streaming."""
from litellm.cost_calculator import _select_model_name_for_cost_calc
selected = _select_model_name_for_cost_calc(
model=None,
completion_response=_bedrock_response_with_private_model("anthropic.claude-v2:1", "us-east-1"),
custom_llm_provider="bedrock",
)
assert selected == "bedrock/us-east-1/anthropic.claude-v2:1"
def test_select_model_name_keeps_base_model_free_of_region(_local_model_cost_map):
"""An explicit base_model keeps pricing on that model's own key even when the request carries a
region with different regional rates, so the private provider model never widens region pricing."""
from litellm.cost_calculator import _select_model_name_for_cost_calc
selected = _select_model_name_for_cost_calc(
model="my-bedrock-deployment",
completion_response=_bedrock_response_with_private_model("moonshotai.kimi-k2.5", "ap-northeast-1"),
base_model="moonshotai.kimi-k2.5",
custom_llm_provider="bedrock",
)
assert selected == "bedrock/moonshotai.kimi-k2.5"
def test_completion_cost_nonzero_for_slash_alias_model_name(_local_model_cost_map):
"""End-to-end cost through a "/"-containing alias must price above zero (#38069)."""
@ -4350,3 +4397,79 @@ def test_realtime_explicitly_free_session_model_still_bills_zero(
)
assert cost == 0.0
def test_completion_cost_prefers_private_provider_response_model(
_local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setitem(
litellm.model_cost,
"openai/selected-cost-model",
{
"input_cost_per_token": 0.000002,
"output_cost_per_token": 0.000004,
"litellm_provider": "openai",
},
)
response = litellm.ModelResponse(
id="x",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
model="requested-route",
)
response._hidden_params = {
"custom_llm_provider": "openai",
"provider_response_model": "selected-cost-model",
}
response.usage = litellm.Usage(prompt_tokens=100, completion_tokens=50)
cost = litellm.completion_cost(
completion_response=response,
custom_llm_provider="openai",
)
assert response.model == "requested-route"
assert cost == pytest.approx(100 * 0.000002 + 50 * 0.000004)
@pytest.mark.parametrize(
("base_model", "custom_pricing", "expected"),
[
("openai/base-model", False, "openai/base-model"),
(None, True, "openai/requested-route"),
],
)
def test_explicit_pricing_precedes_private_provider_response_model(
base_model: str | None,
custom_pricing: bool,
expected: str,
) -> None:
from litellm.cost_calculator import _select_model_name_for_cost_calc
response = litellm.ModelResponse(
id="x",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
model="requested-route",
)
response._hidden_params = {"provider_response_model": "selected-cost-model"}
selected = _select_model_name_for_cost_calc(
model="requested-route",
completion_response=response,
base_model=base_model,
custom_pricing=custom_pricing,
custom_llm_provider="openai",
)
assert selected == expected