fix(streaming): preserve provider model for cost calculation

This commit is contained in:
Andrew Mattie 2026-08-27 23:14:47 -05:00
parent 3300fc3a96
commit 134a4cd9fd
5 changed files with 259 additions and 14 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_provider_response_model_for_cost_calc(hidden_params: object) -> str | None:
if not isinstance(hidden_params, Mapping):
return None
model: Final[object] = hidden_params.get("provider_response_model")
return model if isinstance(model, str) and model 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,9 @@ 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_provider_response_model_for_cost_calc(hidden_params)
region_name_value: Final[object] = hidden_params.get("region_name") if hidden_params is not None else None
region_name: str | None = region_name_value if isinstance(region_name_value, str) else None
if custom_pricing is True:
if router_model_id is not None and router_model_id in litellm.model_cost:
@ -780,14 +789,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

@ -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 = [
@ -1522,7 +1548,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,55 @@ 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():
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 first_result._hidden_params["provider_response_model"] == "selected-model"
assert terminal_result._hidden_params["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 assembled._hidden_params["provider_response_model"] == "selected-model"
@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

@ -4095,7 +4095,10 @@ def test_select_model_name_strips_duplicated_region_segment(_local_model_cost_ma
],
model="us-east-1/anthropic.claude-v2:1",
)
response._hidden_params = {"region_name": "us-east-1"}
response._hidden_params = {
"provider_response_model": "anthropic.claude-v2:1",
"region_name": "us-east-1",
}
selected = _select_model_name_for_cost_calc(
model=None,
@ -4350,3 +4353,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