mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge pull request #38656 from aaaaaandrew/litellm_preserve_stream_selected_model
fix(streaming): preserve provider model for cost calculation
This commit is contained in:
commit
d4f6b4491d
6 changed files with 452 additions and 13 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue