fix(xai): keep streamed and custom-priced billing inside the cost calculator

Restate xAI's usage.cost_in_usd_ticks as usage.cost on chat and responses
replies, streamed ones included, then let the cost calculator own the
figure: a deployment with its own input_cost_per_token and
output_cost_per_token keeps that price, cost margins apply on chat streams
as they already did on non-streamed calls, and only OpenRouter's usage
cost becomes the llm_provider-x-litellm-response-cost header, so xAI
streams no longer skip the calculator through the header or the
stream_chunk_builder hidden response_cost.
This commit is contained in:
mateo-berri 2026-09-02 16:13:31 -07:00
parent 2ba923e18c
commit 1a1d459701
11 changed files with 257 additions and 127 deletions

View file

@ -4,6 +4,7 @@ import logging
import time
from collections.abc import Mapping, Sequence
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from httpx import Response
@ -1164,6 +1165,12 @@ def _store_cost_breakdown_in_logging_obj(
# Don't fail the main cost calculation if breakdown storage fails
def _without_provider_stated_cost(usage: Usage | None) -> Usage | None:
if usage is None or getattr(usage, "cost", None) is None:
return usage
return usage.model_copy(update=MappingProxyType({"cost": None}))
def completion_cost(
completion_response: object | None = None,
model: str | None = None,
@ -1243,7 +1250,10 @@ def completion_cost(
cache_creation_input_tokens: int | None = None
cache_read_input_tokens: int | None = None
audio_transcription_file_duration: float = 0.0
cost_per_token_usage_object: Final[Usage | None] = _get_usage_object(completion_response=completion_response)
provider_usage_object: Final = _get_usage_object(completion_response=completion_response)
cost_per_token_usage_object: Final[Usage | None] = (
_without_provider_stated_cost(provider_usage_object) if custom_pricing else provider_usage_object
)
rerank_billed_units: RerankBilledUnits | None = None
# Extract service_tier from optional_params if not provided directly

View file

@ -54,6 +54,7 @@ FUNCTION_CALL_ATTRIBUTE: Final = "function_call"
_SYNC_ITER_EXHAUSTED: Final = object()
_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value})
def _next_sync_or_exhausted(it: Any) -> object:
@ -1884,8 +1885,8 @@ class CustomStreamWrapper:
@staticmethod
def _resolve_provider_reported_cost(usage_cost: object) -> float | None:
"""
Providers report usage.cost either as a number or, for Perplexity, as a
breakdown object whose total lives under ``total_cost``.
Providers report usage.cost either as a number or as a breakdown object
whose total lives under ``total_cost``.
"""
if isinstance(usage_cost, bool):
return None
@ -1898,12 +1899,10 @@ class CustomStreamWrapper:
@staticmethod
def _propagate_usage_cost_to_hidden_params(
response: "ModelResponse",
custom_llm_provider: str | None,
) -> None:
"""
If the assembled response carries a provider-reported cost on
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
if custom_llm_provider not in _USAGE_COST_HEADER_PROVIDERS:
return
_usage: Final[Usage | None] = getattr(response, "usage", None)
_cost: Final = CustomStreamWrapper._resolve_provider_reported_cost(getattr(_usage, "cost", None))
if _cost is not None:
@ -2018,7 +2017,7 @@ class CustomStreamWrapper:
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider)
setattr(
response,
@ -2268,7 +2267,7 @@ class CustomStreamWrapper:
response: Final = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider)
setattr(
response,

View file

@ -1,4 +1,5 @@
from collections.abc import AsyncIterator, Iterator, Mapping
from types import MappingProxyType
from typing import Any, Final
import httpx
@ -30,31 +31,11 @@ from ...openai.chat.gpt_transformation import (
)
def _adopt_cost_reported_by_xai(usage: Usage | dict[str, Any] | None) -> None: # mutable-ok: streaming dict write
"""Bill what xAI charged instead of repricing the request locally.
xAI reports the amount on ``cost_in_usd_ticks``; restate it in USD on ``cost``,
the field litellm already carries a provider stated cost in and the one
``llms/xai/cost_calculator.py`` prices from. When xAI reported nothing usable,
``cost`` is left alone and the request falls back to token pricing.
Accepts a ``Usage`` (non-streaming) or a raw usage ``dict`` (streaming chunk),
matching ``_fold_reasoning_tokens_into_completion``, so both paths stay in sync.
Streaming needs the dict form because chunk aggregation rebuilds usage from the
fields it models plus ``cost``, dropping everything else xAI sent.
"""
if usage is None:
return
if isinstance(usage, dict):
chunk_cost: Final = xai_reported_cost_in_usd(usage.get("cost_in_usd_ticks"))
if chunk_cost is not None:
usage["cost"] = chunk_cost
return
def _usage_restated_from_xai_ticks(usage: Usage | None) -> Usage | None:
reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None))
if reported_cost is not None:
usage.cost = reported_cost
if usage is None or reported_cost is None:
return None
return usage.model_copy(update=MappingProxyType({"cost": reported_cost}))
class XAIChatConfig(OpenAIGPTConfig):
@ -310,7 +291,9 @@ class XAIChatConfig(OpenAIGPTConfig):
self._fold_reasoning_tokens_into_completion(response)
self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None))
_adopt_cost_reported_by_xai(getattr(response, "usage", None))
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(response, "usage", None))
if restated_usage is not None:
response.usage = restated_usage
return response
@staticmethod
@ -438,6 +421,9 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
if "usage" in chunk and chunk["usage"] is not None:
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
_adopt_cost_reported_by_xai(chunk["usage"])
return super().chunk_parser(chunk)
parsed_chunk: Final = super().chunk_parser(chunk)
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(parsed_chunk, "usage", None))
if restated_usage is not None:
parsed_chunk.usage = restated_usage
return parsed_chunk

View file

@ -12,20 +12,7 @@ USD_TICKS_PER_DOLLAR: Final = 10_000_000_000
def xai_reported_cost_in_usd(cost_in_usd_ticks: object) -> float | None:
"""
Convert the amount xAI says it charged into USD, or None when it reported nothing usable.
xAI states what it billed in ``usage.cost_in_usd_ticks``, at ``USD_TICKS_PER_DOLLAR``
ticks to the dollar: https://docs.x.ai/developers/cost-tracking
That single figure covers the whole request, tokens and every server side tool
invocation together, so whoever bills from it must not add anything on top.
The value arrives on an untyped field of a response body that a caller able to set
api_base controls, so only the documented shape is accepted: a non-negative integer,
with bool refused since it is an int subclass. Anything else yields None and the
request is priced from tokens instead, which stops such an endpoint from reporting a
negative amount to subtract from its own recorded spend.
"""
"""xAI bills in ticks of a dollar: https://docs.x.ai/developers/cost-tracking"""
if not isinstance(cost_in_usd_ticks, int) or isinstance(cost_in_usd_ticks, bool):
return None
if cost_in_usd_ticks < 0:

View file

@ -39,24 +39,6 @@ def apply_server_side_tool_usage_details_to_usage(usage: Usage, details: Mapping
def _cost_reported_by_xai(usage: "Usage") -> float | None:
"""
Return what xAI billed for the request in USD, or None if it reported nothing usable.
The xAI transformations restate ``usage.cost_in_usd_ticks`` as ``usage.cost``, the
field litellm already carries a provider stated cost in and the same one
``llms/perplexity/cost_calculator.py`` bills from. That figure is the total for the
whole request, tokens and every server side tool invocation together, so nothing may
be added on top of it.
A negative amount is refused rather than billed: a caller who can point litellm at an
api_base they control also controls the response body, and a negative cost would
subtract from their own recorded spend. Those requests are priced from tokens instead.
NaN and the infinities are refused for the same reason and are the worse case, because
``Usage`` stores a provider supplied ``cost`` without validating it and NaN compares
false against every budget threshold. Billing one would disable spend enforcement for
the key rather than mispricing a single request.
"""
reported_cost: Final[object] = getattr(usage, "cost", None)
if not isinstance(reported_cost, (int, float)) or isinstance(reported_cost, bool):
return None
@ -69,11 +51,8 @@ def _cost_reported_by_xai(usage: "Usage") -> float | None:
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
"""
Prefers the amount xAI reported for the request, matching how the perplexity
calculator treats a provider-stated cost. That total is returned as completion
cost because xAI does not break it down by direction. Without one, falls back to
the generic cost calculator for all pricing logic, with XAI-specific reasoning
token handling.
Calculates the cost per token for a given XAI model, prompt tokens, and completion tokens.
Uses the generic cost calculator for all pricing logic, with XAI-specific reasoning token handling.
Input:
- model: str, the model name without provider prefix
@ -146,11 +125,9 @@ def cost_per_web_search_request(usage: "Usage", model_info: "ModelInfo") -> floa
"""
Calculate the cost of web search requests for X.AI models.
When xAI reports what it billed, that figure already covers the server-side
search calls and ``cost_per_token`` has returned it, so there is nothing to add
here. Otherwise price the invocations from
usage.server_side_tool_usage_details.web_search_calls at the per-call rate
(model_info.search_context_cost_per_query when set, else the default $5 / 1k).
Counts invocations from usage.server_side_tool_usage_details.web_search_calls.
Per-call rate comes from model_info.search_context_cost_per_query when set,
otherwise the default xAI tools rate ($5 / 1k calls).
"""
if _cost_reported_by_xai(usage) is not None:
return 0.0

View file

@ -1,3 +1,4 @@
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -10,8 +11,13 @@ from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfi
from litellm.llms.xai.common_utils import XAIModelInfo, xai_reported_cost_in_usd
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamingResponse,
)
from litellm.types.llms.xai import XAIWebSearchTool, XAIXSearchTool
from litellm.types.router import GenericLiteLLMParams
@ -27,6 +33,13 @@ else:
LiteLLMLoggingObj = Any
def _usage_restated_from_xai_ticks(usage: ResponseAPIUsage | None) -> ResponseAPIUsage | None:
reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None))
if usage is None or reported_cost is None:
return None
return usage.model_copy(update=MappingProxyType({"cost": reported_cost}))
class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
Configuration for XAI's Responses API.
@ -270,31 +283,35 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
"""
Bill what xAI charged instead of repricing the request locally.
xAI reports the amount on ``usage.cost_in_usd_ticks``; restate it in USD on
``usage.cost``, which ``ResponseAPILoggingUtils`` already copies onto the chat
Usage that ``llms/xai/cost_calculator.py`` prices, so /v1/responses bills the
reported figure the same way /v1/chat/completions does. When xAI reported nothing
usable, ``cost`` is left alone and the request falls back to token pricing.
"""
response: Final = super().transform_response_api_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
)
usage: Final = response.usage
if usage is None:
return response
reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None))
if reported_cost is not None:
usage.cost = reported_cost
restated_usage: Final = _usage_restated_from_xai_ticks(response.usage)
if restated_usage is not None:
response.usage = restated_usage
return response
def transform_streaming_response(
self,
model: str,
parsed_chunk: dict, # mutable-ok: overrides the base class signature
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIStreamingResponse:
event: Final = super().transform_streaming_response(
model=model,
parsed_chunk=parsed_chunk,
logging_obj=logging_obj,
)
if not isinstance(event, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent)):
return event
restated_usage: Final = _usage_restated_from_xai_ticks(event.response.usage)
if restated_usage is not None:
event.response.usage = restated_usage
return event
def supports_native_websocket(self) -> bool:
"""XAI does not support native WebSocket for Responses API"""
return False

View file

@ -8592,10 +8592,11 @@ def stream_chunk_builder_text_completion(chunks: list, messages: list | None = N
def _stream_builder_response_cost(response: ModelResponse, logging_obj: Optional["Logging"]) -> float | None:
usage_cost: Final = getattr(getattr(response, "usage", None), "cost", None)
if isinstance(usage_cost, (int, float)):
return float(usage_cost)
numeric_usage_cost: Final = float(usage_cost) if isinstance(usage_cost, (int, float)) else None
if logging_obj is not None:
return None
return numeric_usage_cost if litellm.include_cost_in_streaming_usage else None
if numeric_usage_cost is not None:
return numeric_usage_cost
provider_hint: Final = response._hidden_params.get( # pyright: ignore[reportPrivateUsage] # no public accessor
"custom_llm_provider"
)

View file

@ -21,6 +21,7 @@ from litellm.litellm_core_utils.streaming_handler import (
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
Delta,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
StandardLoggingPayload,
@ -1750,7 +1751,7 @@ def test_openrouter_streaming_cost_propagates_to_hidden_params():
assert complete_response.usage.cost == 0.00025
# Use the real propagation method from CustomStreamWrapper
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openrouter")
assert "additional_headers" in complete_response._hidden_params
assert (
@ -1769,14 +1770,12 @@ def test_openrouter_streaming_cost_propagates_to_hidden_params():
assert provider_cost == 0.00025
def test_perplexity_streaming_dict_cost_propagates_to_hidden_params():
"""
Regression: Perplexity reports usage.cost as a breakdown object, which used to
blow up the end of the stream with
`float() argument must be a string or a real number, not 'dict'`.
"""
def test_perplexity_streaming_dict_cost_bills_through_its_own_calculator():
import litellm
from litellm.cost_calculator import get_response_cost_from_hidden_params
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
chunks = [
ModelResponseStream(
@ -1828,13 +1827,81 @@ def test_perplexity_streaming_dict_cost_propagates_to_hidden_params():
assert complete_response is not None
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "perplexity")
assert (
get_response_cost_from_hidden_params(complete_response._hidden_params)
== 0.00503
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
assert response_cost_calculator(
response_object=complete_response,
model="perplexity/sonar",
custom_llm_provider="perplexity",
call_type="completion",
optional_params={},
) == pytest.approx(0.00503)
def test_openai_compatible_streaming_cost_is_priced_from_the_cost_map():
import litellm
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
model = "openai/streams-cost-in-nanodollars"
litellm.register_model(
{
model: {
"input_cost_per_token": 1e-6,
"output_cost_per_token": 2e-6,
"litellm_provider": "openai",
"mode": "chat",
}
}
)
complete_response = ModelResponse(
id="chatcmpl-openai-compatible",
model=model,
choices=[],
usage=Usage(completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=3_144_000),
)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openai")
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
assert response_cost_calculator(
response_object=complete_response,
model=model,
custom_llm_provider="openai",
call_type="completion",
optional_params={},
) == pytest.approx(2e-5)
def test_xai_streaming_reported_cost_still_takes_the_margin(monkeypatch):
import litellm
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
complete_response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[],
usage=Usage(completion_tokens=353, prompt_tokens=198, total_tokens=551, cost=0.0009956),
)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "xai")
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5})
assert response_cost_calculator(
response_object=complete_response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
call_type="completion",
optional_params={},
) == pytest.approx(0.0009956 * 1.5)
def test_provider_reported_cost_ignores_unusable_shapes():
assert CustomStreamWrapper._resolve_provider_reported_cost(None) is None

View file

@ -412,22 +412,22 @@ class TestXAIResponsesReportedCost:
"""
@staticmethod
def _transformed_usage(usage: dict) -> ResponseAPIUsage | None:
raw_response = httpx.Response(
status_code=200,
json={
"id": "resp_xai",
"object": "response",
"created_at": 0,
"model": "grok-4-latest",
"status": "completed",
"output": [],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"usage": usage,
},
)
def _response_body(usage: dict) -> dict:
return {
"id": "resp_xai",
"object": "response",
"created_at": 0,
"model": "grok-4-latest",
"status": "completed",
"output": [],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"usage": usage,
}
def _transformed_usage(self, usage: dict) -> ResponseAPIUsage | None:
raw_response = httpx.Response(status_code=200, json=self._response_body(usage))
response = XAIResponsesAPIConfig().transform_response_api_response(
model="grok-4-latest",
@ -451,6 +451,28 @@ class TestXAIResponsesReportedCost:
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756)
def test_streamed_reported_cost_reaches_the_cost_calculator(self):
event = XAIResponsesAPIConfig().transform_streaming_response(
model="grok-4-latest",
parsed_chunk={
"type": "response.completed",
"sequence_number": 7,
"response": self._response_body(
{
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": 37756000,
}
),
},
logging_obj=Mock(),
)
assert isinstance(event, ResponseCompletedEvent)
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage)
assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756)
def test_usage_without_a_reported_cost_is_left_alone(self):
usage = self._transformed_usage(
{"input_tokens": 100, "output_tokens": 200, "total_tokens": 300}

View file

@ -7,7 +7,10 @@ import os
import litellm
from litellm.types.utils import (
Choices,
CompletionTokensDetailsWrapper,
Message,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
@ -576,6 +579,48 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_custom_pricing_beats_the_reported_cost(self):
response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[Choices(index=0, message=Message(role="assistant", content="x"), finish_reason="stop")],
usage=Usage(prompt_tokens=198, completion_tokens=353, total_tokens=551, cost=0.0009956),
)
billed = litellm.completion_cost(
completion_response=response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
custom_cost_per_token={"input_cost_per_token": 0.001, "output_cost_per_token": 0.001},
custom_pricing=True,
)
assert math.isclose(billed, 0.551, rel_tol=1e-10)
def test_deployment_custom_pricing_beats_the_reported_cost(self, monkeypatch):
deployment_id = "xai-deployment-priced-by-the-operator"
monkeypatch.setitem(
litellm.model_cost,
deployment_id,
{"input_cost_per_token": 0.001, "output_cost_per_token": 0.001, "litellm_provider": "xai", "mode": "chat"},
)
response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[Choices(index=0, message=Message(role="assistant", content="x"), finish_reason="stop")],
usage=Usage(prompt_tokens=198, completion_tokens=353, total_tokens=551, cost=0.0009956),
)
billed = litellm.completion_cost(
completion_response=response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
custom_pricing=True,
router_model_id=deployment_id,
)
assert math.isclose(billed, 0.551, rel_tol=1e-10)
class TestXAIWebSearchCostHelpers:
"""Focused coverage for web_search / tool-usage helpers in cost_calculator.py."""

View file

@ -3181,3 +3181,22 @@ def test_stream_chunk_builder_defers_cost_to_logging_obj_when_usage_cost_absent(
assert response is not None
assert response._hidden_params.get("response_cost") is None
def test_stream_chunk_builder_defers_provider_reported_cost_to_logging_obj(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False)
usage_chunk: Final = _stream_builder_text_chunk("gpt-4o", "")
usage_chunk.usage = Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7, cost=0.42)
chunks: Final = [
_stream_builder_text_chunk("gpt-4o", "Hello "),
_stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"),
usage_chunk,
]
response: Final = litellm.stream_chunk_builder(
chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=_stream_builder_logging_obj()
)
assert response is not None
assert response.usage.cost == 0.42
assert response._hidden_params.get("response_cost") is None