mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(streaming): preserve provider model for cost calculation
This commit is contained in:
parent
3300fc3a96
commit
134a4cd9fd
5 changed files with 259 additions and 14 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_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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue