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:
shivam 2026-09-15 22:42:58 +00:00
parent 226b1e1bb9
commit 65132ef2e9
5 changed files with 219 additions and 9 deletions

View file

@ -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

View file

@ -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 (

View file

@ -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")

View file

@ -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)."""

View file

@ -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"))