mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(cost): price native Responses WebSocket turns at their returned service_tier
The logging object built for aresponses_websocket sessions dropped the service_tier carried by each billable response.completed/response.incomplete event, so sessions were priced at the default tier. The logging object now carries the tier when the session is single-tier, and completion_cost splits mixed-tier sessions into one object per tier before pricing. Fixes #41299 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
226b1e1bb9
commit
65132ef2e9
5 changed files with 219 additions and 9 deletions
|
|
@ -1203,6 +1203,24 @@ def _without_provider_stated_cost(usage: Usage | None) -> Usage | None:
|
|||
return usage.model_copy(update=MappingProxyType({"cost": None}))
|
||||
|
||||
|
||||
def _split_responses_ws_logging_object_by_service_tier(
|
||||
completion_response: LiteLLMRealtimeStreamLoggingObject,
|
||||
) -> tuple[LiteLLMRealtimeStreamLoggingObject, ...] | None:
|
||||
partition: Final = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(
|
||||
cast(Sequence[Mapping[str, object]], completion_response.results)
|
||||
)
|
||||
if len(partition) <= 1:
|
||||
return None
|
||||
return tuple(
|
||||
LiteLLMRealtimeStreamLoggingObject(
|
||||
results=cast(OpenAIRealtimeStreamList, list(group)),
|
||||
usage=ResponsesWebSocketTokenUsageProcessor.collect_and_combine_usage_from_responses_ws_results(group),
|
||||
service_tier=tier,
|
||||
)
|
||||
for tier, group in partition.items()
|
||||
)
|
||||
|
||||
|
||||
def completion_cost(
|
||||
completion_response: object | None = None,
|
||||
model: str | None = None,
|
||||
|
|
@ -1266,6 +1284,41 @@ def completion_cost(
|
|||
try:
|
||||
call_type = _infer_call_type(call_type, completion_response) or "completion"
|
||||
|
||||
if call_type == CallTypes.aresponses_websocket.value and isinstance(
|
||||
completion_response, LiteLLMRealtimeStreamLoggingObject
|
||||
):
|
||||
ws_tier_parts: Final = _split_responses_ws_logging_object_by_service_tier(completion_response)
|
||||
if ws_tier_parts is not None:
|
||||
return sum(
|
||||
completion_cost(
|
||||
completion_response=part,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
messages=messages,
|
||||
completion=completion,
|
||||
total_time=total_time,
|
||||
call_type=call_type,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
region_name=region_name,
|
||||
size=size,
|
||||
quality=quality,
|
||||
n=n,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
custom_cost_per_second=custom_cost_per_second,
|
||||
optional_params=optional_params,
|
||||
custom_pricing=custom_pricing,
|
||||
base_model=base_model,
|
||||
standard_built_in_tools_params=standard_built_in_tools_params,
|
||||
litellm_model_name=litellm_model_name,
|
||||
router_model_id=router_model_id,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
for part in ws_tier_parts
|
||||
)
|
||||
|
||||
if (
|
||||
(call_type == "aimage_generation" or call_type == "image_generation")
|
||||
and model is not None
|
||||
|
|
@ -2558,6 +2611,7 @@ _RESPONSES_WS_BILLABLE_EVENT_TYPES: Final = frozenset({"response.completed", "re
|
|||
|
||||
class _ResponsesWsEventResponse(BaseModel):
|
||||
usage: Mapping[str, object] | None = None
|
||||
service_tier: str | None = None
|
||||
|
||||
|
||||
class _ResponsesWsEvent(BaseModel):
|
||||
|
|
@ -2565,20 +2619,39 @@ class _ResponsesWsEvent(BaseModel):
|
|||
response: _ResponsesWsEventResponse | None = None
|
||||
|
||||
|
||||
def _billable_responses_ws_events(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> tuple[tuple[Mapping[str, object], _ResponsesWsEventResponse], ...]:
|
||||
return tuple(
|
||||
(result, event.response)
|
||||
for result in results
|
||||
if (event := _ResponsesWsEvent.model_validate(result)).type in _RESPONSES_WS_BILLABLE_EVENT_TYPES
|
||||
and event.response is not None
|
||||
and event.response.usage is not None
|
||||
)
|
||||
|
||||
|
||||
class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
|
||||
@staticmethod
|
||||
def collect_usage_from_responses_ws_results(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> tuple[Usage, ...]:
|
||||
events: Final = tuple(_ResponsesWsEvent.model_validate(result) for result in results)
|
||||
return tuple(
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # same shared transform the realtime processor uses
|
||||
event.response.usage
|
||||
response.usage
|
||||
)
|
||||
for event in events
|
||||
if event.type in _RESPONSES_WS_BILLABLE_EVENT_TYPES
|
||||
and event.response is not None
|
||||
and event.response.usage is not None
|
||||
for _, response in _billable_responses_ws_events(results)
|
||||
if response.usage is not None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def partition_results_by_service_tier(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> Mapping[str | None, tuple[Mapping[str, object], ...]]:
|
||||
billable: Final = _billable_responses_ws_events(results)
|
||||
tiers: Final = dict.fromkeys(response.service_tier for _, response in billable)
|
||||
return MappingProxyType(
|
||||
{tier: tuple(result for result, response in billable if response.service_tier == tier) for tier in tiers}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -2101,9 +2101,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
results=result # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
)
|
||||
)
|
||||
ws_tier_partition: Final = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(
|
||||
results=result # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
)
|
||||
ws_service_tier: Final = next(iter(ws_tier_partition)) if len(ws_tier_partition) == 1 else None
|
||||
logging_result = LiteLLMRealtimeStreamLoggingObject(
|
||||
usage=combined_ws_usage,
|
||||
results=result, # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
service_tier=ws_service_tier,
|
||||
)
|
||||
|
||||
elif (
|
||||
|
|
|
|||
|
|
@ -4275,6 +4275,7 @@ class LiteLLMRealtimeStreamLoggingObject(LiteLLMPydanticObjectBase):
|
|||
# rate_limits.updated), blocks the event loop, and discards the session usage.
|
||||
results: SkipValidation[OpenAIRealtimeStreamList]
|
||||
usage: Usage
|
||||
service_tier: str | None = None
|
||||
_hidden_params: dict = {}
|
||||
|
||||
@field_serializer("results")
|
||||
|
|
|
|||
|
|
@ -6554,9 +6554,9 @@ async def test_prompt_hook_injection_marker_recorded_for_every_surface(logging_o
|
|||
assert pre_choice["metadata"]["litellm_gateway_injected_cache"] == ""
|
||||
|
||||
|
||||
def _responses_ws_logging_obj() -> LitellmLogging:
|
||||
def _responses_ws_logging_obj(model: str = "gpt-4o") -> LitellmLogging:
|
||||
return LitellmLogging(
|
||||
model="gpt-4o",
|
||||
model=model,
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
|
|
@ -6638,6 +6638,62 @@ def test_normalize_logging_result_bills_incomplete_responses_websocket_turns():
|
|||
assert normalized.usage.total_tokens == 75
|
||||
|
||||
|
||||
def test_normalize_logging_result_prices_responses_websocket_at_returned_service_tier():
|
||||
"""Issue #41299: a WebSocket turn billed at priority tier reported it on
|
||||
response.completed.response.service_tier, but the logging object dropped it and the
|
||||
session was priced at the default tier."""
|
||||
events = [
|
||||
{"type": "response.created", "response": {}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"service_tier": "priority",
|
||||
"usage": {"input_tokens": 100, "output_tokens": 40, "total_tokens": 140},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
normalized = _responses_ws_logging_obj(model="gpt-5.4").normalize_logging_result(result=events)
|
||||
|
||||
assert isinstance(normalized, LiteLLMRealtimeStreamLoggingObject)
|
||||
assert normalized.service_tier == "priority"
|
||||
|
||||
usage = ResponseAPIUsage(input_tokens=100, output_tokens=40, total_tokens=140)
|
||||
ws_cost = litellm.completion_cost(
|
||||
completion_response=normalized,
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
priority_http_cost = litellm.completion_cost(
|
||||
completion_response=ResponsesAPIResponse(
|
||||
id="resp-priority",
|
||||
created_at=1700000000,
|
||||
output=[],
|
||||
service_tier="priority",
|
||||
usage=usage,
|
||||
),
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
default_http_cost = litellm.completion_cost(
|
||||
completion_response=ResponsesAPIResponse(
|
||||
id="resp-default",
|
||||
created_at=1700000000,
|
||||
output=[],
|
||||
service_tier="default",
|
||||
usage=usage,
|
||||
),
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert ws_cost == priority_http_cost
|
||||
assert priority_http_cost > default_http_cost
|
||||
|
||||
|
||||
def test_get_standard_logging_object_payload_reads_overhead_from_logging_obj_for_dict_results(logging_obj):
|
||||
"""LIT-5466: /v1/messages returns a plain dict with no _hidden_params, so the overhead
|
||||
recorded on the logging object must reach hidden_params.litellm_overhead_time_ms (SpendLogs)."""
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
|
||||
import time
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -9,6 +10,7 @@ import litellm
|
|||
from litellm.cost_calculator import (
|
||||
BaseTokenUsageProcessor,
|
||||
RealtimeAPITokenUsageProcessor,
|
||||
ResponsesWebSocketTokenUsageProcessor,
|
||||
completion_cost,
|
||||
cost_per_token,
|
||||
handle_realtime_stream_cost_calculation,
|
||||
|
|
@ -17,10 +19,12 @@ from litellm.cost_calculator import (
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse, OCRUsageInfo
|
||||
from litellm.types.llms.base import CachedTokensDetails
|
||||
from litellm.types.llms.openai import OpenAIRealtimeStreamList
|
||||
from litellm.types.llms.openai import OpenAIRealtimeStreamList, ResponseAPIUsage, ResponsesAPIResponse
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.utils import (
|
||||
CacheCreationTokenDetails,
|
||||
CallTypes,
|
||||
LiteLLMRealtimeStreamLoggingObject,
|
||||
ModelInfo,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
|
|
@ -5243,3 +5247,74 @@ def test_completion_cost_ocr_ignores_deployment_pricing_without_custom_pricing_f
|
|||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
assert cost == 0.0
|
||||
|
||||
|
||||
def test_completion_cost_prices_responses_websocket_turns_per_service_tier():
|
||||
"""Issue #41299: a session mixing default and priority turns must price each turn at
|
||||
its own returned service_tier, not the summed usage at a single tier."""
|
||||
events = [
|
||||
{"type": "response.created", "response": {}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"service_tier": "default",
|
||||
"usage": {"input_tokens": 100, "output_tokens": 40, "total_tokens": 140},
|
||||
},
|
||||
},
|
||||
{"type": "rate_limits.updated", "rate_limits": {}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"service_tier": "priority",
|
||||
"usage": {"input_tokens": 60, "output_tokens": 10, "total_tokens": 70},
|
||||
},
|
||||
},
|
||||
{"type": "response.failed", "response": {"usage": None}},
|
||||
]
|
||||
|
||||
partition = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(events)
|
||||
assert tuple(partition.keys()) == ("default", "priority")
|
||||
assert len(partition["default"]) == 1
|
||||
assert len(partition["priority"]) == 1
|
||||
|
||||
logging_obj = Logging(
|
||||
model="gpt-5.4",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
start_time=time.time(),
|
||||
litellm_call_id="responses-ws-tier-test",
|
||||
function_id="responses-ws-tier-test",
|
||||
)
|
||||
normalized = logging_obj.normalize_logging_result(result=events)
|
||||
assert isinstance(normalized, LiteLLMRealtimeStreamLoggingObject)
|
||||
assert normalized.service_tier is None
|
||||
|
||||
def _http_cost(input_tokens: int, output_tokens: int, service_tier: str) -> float:
|
||||
return completion_cost(
|
||||
completion_response=ResponsesAPIResponse(
|
||||
id=f"resp-{service_tier}",
|
||||
created_at=1700000000,
|
||||
output=[],
|
||||
service_tier=service_tier,
|
||||
usage=ResponseAPIUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_tokens + output_tokens,
|
||||
),
|
||||
),
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
ws_cost = completion_cost(
|
||||
completion_response=normalized,
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert ws_cost == pytest.approx(_http_cost(100, 40, "default") + _http_cost(60, 10, "priority"))
|
||||
assert ws_cost != pytest.approx(_http_cost(160, 50, "default"))
|
||||
assert ws_cost != pytest.approx(_http_cost(160, 50, "priority"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue