Merge pull request #35114 from BerriAI/litellm_fix_messages_stream_cost_cache_tokens

fix(cost): match streamed Messages usage cost to the recorded spend
This commit is contained in:
Mateo Wang 2026-08-20 19:35:09 -07:00 • committed by GitHub
commit a4dd1be53b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 364 additions and 48 deletions

View file

@ -2216,7 +2216,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
def calculate_usage(
self,
usage_object: dict,
usage_object: Mapping[str, Any],
reasoning_content: str | None,
completion_response: dict | None = None,
speed: str | None = None,

View file

@ -158,7 +158,7 @@ ProxyRouteType: TypeAlias = Literal[
"acancel_run",
"adelete_run",
]
from litellm.types.utils import ServerToolUse
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
StreamChunkSerializer = Callable[[Any], str]
@ -3321,7 +3321,9 @@ class ProxyBaseLLMRequestProcessing:
str_so_far += str(chunk.get("content", ""))
model_name = request_data.get("model", "")
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, model_name)
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
chunk, model_name, request_data.get("litellm_logging_obj")
)
# Set before the yield: an async generator suspends at the yield,
# so a GeneratorExit on client disconnect is raised there and any
@ -3418,20 +3420,27 @@ class ProxyBaseLLMRequestProcessing:
@overload
@staticmethod
def _process_chunk_with_cost_injection(chunk: bytes, model_name: str) -> bytes: ...
def _process_chunk_with_cost_injection(
chunk: bytes, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> bytes: ...
@overload
@staticmethod
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object: ...
def _process_chunk_with_cost_injection(
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> object: ...
@staticmethod
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object:
def _process_chunk_with_cost_injection(
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> object:
"""
Process a streaming chunk and inject cost information if enabled.
Args:
chunk: The streaming chunk (dict, str, bytes, or bytearray)
model_name: Model name for cost calculation
litellm_logging_obj: The call's logging object, used for pricing
Returns:
The processed chunk with cost information injected if applicable
@ -3441,21 +3450,27 @@ class ProxyBaseLLMRequestProcessing:
try:
if isinstance(chunk, dict):
maybe_modified: Final = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(chunk, model_name)
maybe_modified: Final = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(
chunk, model_name, litellm_logging_obj
)
if maybe_modified is not None:
return maybe_modified
elif isinstance(chunk, (bytes, bytearray)):
try:
s: Final = chunk.decode("utf-8")
if s.endswith(("\n\n", "\r\n\r\n")):
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(
s, model_name, litellm_logging_obj
)
if maybe_mod is not None:
return maybe_mod.encode("utf-8")
except Exception:
pass
elif isinstance(chunk, str):
# Try to parse SSE frame and inject cost into the data line
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(chunk, model_name)
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(
chunk, model_name, litellm_logging_obj
)
if maybe_mod is not None:
# Ensure trailing frame separator
return maybe_mod if maybe_mod.endswith("\n\n") else (maybe_mod + "\n\n")
@ -3466,13 +3481,16 @@ class ProxyBaseLLMRequestProcessing:
return chunk
@staticmethod
def _inject_cost_into_sse_frame_str(frame_str: str, model_name: str) -> str | None:
def _inject_cost_into_sse_frame_str(
frame_str: str, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> str | None:
"""
Inject cost information into an SSE frame string by modifying the JSON in the 'data:' line.
Args:
frame_str: SSE frame string that may contain multiple lines
model_name: Model name for cost calculation
litellm_logging_obj: The call's logging object, forwarded for pricing
Returns:
Modified SSE frame string with cost injected, or None if no modification needed
@ -3486,7 +3504,9 @@ class ProxyBaseLLMRequestProcessing:
json_part = stripped_ln.split("data:", 1)[1].strip()
if json_part and json_part != "[DONE]":
obj = json.loads(json_part)
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(obj, model_name)
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(
obj, model_name, litellm_logging_obj
)
if maybe_modified is not None:
lines[idx] = "data: " + safe_dumps(maybe_modified) + ("\r" if ln.endswith("\r") else "")
return "\n".join(lines)
@ -3494,34 +3514,6 @@ class ProxyBaseLLMRequestProcessing:
except Exception:
return None
@staticmethod
def _anthropic_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
prompt_tokens: Final = int(usage.get("input_tokens", 0) or 0)
completion_tokens: Final = int(usage.get("output_tokens", 0) or 0)
total_tokens: Final = int(
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
)
web_search_requests: Final = usage.get("web_search_requests")
server_tool_use: Final = (
ServerToolUse(web_search_requests=web_search_requests) if web_search_requests is not None else None
)
return MappingProxyType(
{
key: value
for key, value in (
("prompt_tokens", prompt_tokens),
("completion_tokens", completion_tokens),
("total_tokens", total_tokens),
("completion_tokens_details", usage.get("completion_tokens_details")),
("prompt_tokens_details", usage.get("prompt_tokens_details")),
("cache_creation_input_tokens", usage.get("cache_creation_input_tokens")),
("cache_read_input_tokens", usage.get("cache_read_input_tokens")),
("server_tool_use", server_tool_use),
)
if value is not None
}
)
@staticmethod
def _openai_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
prompt_tokens: Final = int(usage.get("prompt_tokens", 0) or 0)
@ -3544,11 +3536,13 @@ class ProxyBaseLLMRequestProcessing:
)
@staticmethod
def _stream_usage_kwargs_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Mapping[str, Any] | None:
def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Usage | None:
# Anthropic reports input_tokens excluding cache tokens, so reuse the non-streaming
# transformation to total the prompt and keep the 5m/1h cache creation split
if obj.get("type") == "message_delta":
return ProxyBaseLLMRequestProcessing._anthropic_stream_usage_kwargs(usage)
return AnthropicConfig().calculate_usage(usage_object=usage, reasoning_content=None)
if obj.get("object") == "chat.completion.chunk":
return ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage)
return Usage(**ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage))
return None
@staticmethod
@ -3563,7 +3557,54 @@ class ProxyBaseLLMRequestProcessing:
return None
@staticmethod
def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> dict | None:
def _logging_obj_cost_or_none(
model_response: ModelResponse, litellm_logging_obj: LiteLLMLoggingObj
) -> float | None:
# Pricing a frame stamps cost_breakdown and, on failure, the cost-failure debug key onto
# the live logging object. The pass-through handlers never recompute either one, so a
# frame-derived breakdown would outlive the stream and land in the spend log. Snapshot
# both and put them back, so pricing here stays a read as far as the request is concerned
breakdown_before: Final = getattr(litellm_logging_obj, "cost_breakdown", None)
call_details: Final = getattr(litellm_logging_obj, "model_call_details", None)
debug_key: Final = "response_cost_failure_debug_information"
debug_missing: Final = object()
debug_before: Final = call_details.get(debug_key, debug_missing) if isinstance(call_details, dict) else None
try:
cost: Final = litellm_logging_obj._response_cost_calculator(result=model_response) # pyright: ignore[reportPrivateUsage] # reuse the call's own cost calc for pricing parity with the logging callback
except Exception: # noqa: BLE001 # a pricing failure falls back to model-name pricing instead of breaking the stream
return None
finally:
if hasattr(litellm_logging_obj, "cost_breakdown"):
litellm_logging_obj.cost_breakdown = breakdown_before
if isinstance(call_details, dict):
if debug_before is debug_missing:
call_details.pop(debug_key, None)
else:
call_details[debug_key] = debug_before
return float(cost) if isinstance(cost, (int, float)) and not isinstance(cost, bool) else None
@staticmethod
def _streamed_usage_cost(
model_response: ModelResponse,
model_name: str,
service_tier: str | None,
litellm_logging_obj: LiteLLMLoggingObj | None,
) -> float | None:
# Pricing via the logging object inherits the deployment's custom pricing, so the
# streamed cost matches what the logging callback records instead of sticker price
cost_from_logging_obj: Final = (
ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, litellm_logging_obj)
if litellm_logging_obj is not None
else None
)
if cost_from_logging_obj is not None:
return cost_from_logging_obj
return ProxyBaseLLMRequestProcessing._completion_cost_or_none(model_response, model_name, service_tier)
@staticmethod
def _inject_cost_into_usage_dict(
obj: dict, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> dict | None:
"""
Inject cost information into the usage object of a streamed usage event
(Anthropic ``message_delta`` or OpenAI ``chat.completion.chunk``).
@ -3571,6 +3612,7 @@ class ProxyBaseLLMRequestProcessing:
Args:
obj: Dictionary containing the SSE event data
model_name: Model name for cost calculation
litellm_logging_obj: The call's logging object, used for pricing
Returns:
Modified dictionary with cost injected, or None if no modification needed
@ -3578,14 +3620,15 @@ class ProxyBaseLLMRequestProcessing:
usage: Final = obj.get("usage")
if not isinstance(usage, dict):
return None
usage_kwargs: Final = ProxyBaseLLMRequestProcessing._stream_usage_kwargs_for_event(obj, usage)
if usage_kwargs is None:
stream_usage: Final = ProxyBaseLLMRequestProcessing._stream_usage_for_event(obj, usage)
if stream_usage is None:
return None
service_tier: Final = obj.get("service_tier")
cost_val: Final = ProxyBaseLLMRequestProcessing._completion_cost_or_none(
ModelResponse(usage=Usage(**usage_kwargs)),
cost_val: Final = ProxyBaseLLMRequestProcessing._streamed_usage_cost(
ModelResponse(usage=stream_usage),
model_name,
service_tier if isinstance(service_tier, str) else None,
litellm_logging_obj,
)
if cost_val is None:
return None

View file

@ -106,7 +106,7 @@ class PassThroughStreamingHandler:
) # rebind-ok: SSE frame reassembly buffer across transport chunks
if complete_frames:
yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
complete_frames, resolved_model_name
complete_frames, resolved_model_name, litellm_logging_obj
)
if pending:
yield pending

View file

@ -6082,6 +6082,254 @@ class TestInjectCostIntoUsageDict:
injected = json.loads(result.split("\n")[0].split("data:", 1)[1].strip())
assert injected["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
def test_message_delta_cost_charges_the_non_cached_input_tokens(self):
"""Anthropic reports ``input_tokens`` excluding cache tokens, so reading it as the whole
prompt total drops the non-cached input from the bill on every cache hit."""
model = "claude-haiku-4-5"
pricing = litellm.model_cost[model]
event = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {
"input_tokens": 14,
"output_tokens": 8,
"cache_read_input_tokens": 3202,
"cache_creation_input_tokens": 0,
},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model)
assert result is not None
expected = (
14 * pricing["input_cost_per_token"]
+ 3202 * pricing["cache_read_input_token_cost"]
+ 8 * pricing["output_cost_per_token"]
)
dropped_input = expected - 14 * pricing["input_cost_per_token"]
assert result["usage"]["cost"] == pytest.approx(expected)
assert result["usage"]["cost"] > dropped_input
def test_message_delta_prices_1h_cache_creation_above_the_5m_rate(self):
"""The ``cache_creation`` 5m/1h split has to survive into ``prompt_tokens_details``,
otherwise a 1h write is billed at the cheaper 5m rate."""
model = "claude-haiku-4-5"
pricing = litellm.model_cost[model]
event = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {
"input_tokens": 14,
"output_tokens": 8,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 2000,
"cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 2000},
},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model)
assert result is not None
base = 14 * pricing["input_cost_per_token"] + 8 * pricing["output_cost_per_token"]
expected_1h = base + 2000 * pricing["cache_creation_input_token_cost_above_1hr"]
flat_5m = base + 2000 * pricing["cache_creation_input_token_cost"]
assert expected_1h != pytest.approx(flat_5m)
assert result["usage"]["cost"] == pytest.approx(expected_1h)
def test_message_delta_prices_through_the_logging_obj_so_custom_pricing_applies(self):
"""Costing by model name alone yields sticker price, so a deployment with a negotiated
discount streamed a ``usage.cost`` that disagreed with the callback's ``response_cost``."""
class _StubLoggingObj:
def __init__(self, cost):
self._cost = cost
self.captured_result = None
def _response_cost_calculator(self, result):
self.captured_result = result
return self._cost
model = "claude-haiku-4-5"
discounted_cost = 0.00099
stub = _StubLoggingObj(discounted_cost)
event = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {
"input_tokens": 14,
"output_tokens": 8,
"cache_read_input_tokens": 3202,
"cache_creation_input_tokens": 500,
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 400},
},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, stub)
assert result is not None
assert result["usage"]["cost"] == discounted_cost
assert result["usage"]["cost"] != pytest.approx(self._expected_cost(model, 14 + 500 + 3202, 8))
usage = stub.captured_result.usage
assert usage.prompt_tokens == 14 + 500 + 3202
details = usage.prompt_tokens_details.cache_creation_token_details
assert details.ephemeral_5m_input_tokens == 100
assert details.ephemeral_1h_input_tokens == 400
def test_message_delta_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self):
class _StubLoggingObj:
def _response_cost_calculator(self, result):
return None
model = "claude-haiku-4-5"
pricing = litellm.model_cost[model]
event = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"input_tokens": 14, "output_tokens": 8, "cache_read_input_tokens": 3202},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, _StubLoggingObj())
assert result is not None
assert result["usage"]["cost"] == pytest.approx(
14 * pricing["input_cost_per_token"]
+ 3202 * pricing["cache_read_input_token_cost"]
+ 8 * pricing["output_cost_per_token"]
)
def test_message_delta_falls_back_to_model_pricing_when_the_logging_obj_raises(self):
"""A pricing failure mid-stream must not break the frame, so the raise falls back to
model-name pricing rather than propagating into the response body."""
class _StubLoggingObj:
def _response_cost_calculator(self, result):
raise ValueError("no pricing for this deployment")
model = "claude-haiku-4-5"
pricing = litellm.model_cost[model]
event = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"input_tokens": 14, "output_tokens": 8, "cache_read_input_tokens": 3202},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, _StubLoggingObj())
assert result is not None
assert result["usage"]["cost"] == pytest.approx(
14 * pricing["input_cost_per_token"]
+ 3202 * pricing["cache_read_input_token_cost"]
+ 8 * pricing["output_cost_per_token"]
)
def test_pricing_a_frame_leaves_the_real_logging_obj_unchanged(self):
"""Pricing runs against the live logging object, and the pass-through handlers never
recompute cost_breakdown, so a frame-derived breakdown would reach the spend log."""
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
from litellm.types.utils import ModelResponse, Usage
logging_obj = LiteLLMLoggingObj(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "test"}],
stream=True,
call_type="completion",
start_time=None,
litellm_call_id="lit4902-breakdown-test",
function_id="lit4902-breakdown-test",
)
logging_obj.update_environment_variables(litellm_params={}, optional_params={})
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
assert logging_obj.cost_breakdown is None
model_response = ModelResponse(
usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)
)
cost = ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj)
assert cost is not None and cost > 0
assert logging_obj.cost_breakdown is None
assert "response_cost_failure_debug_information" not in logging_obj.model_call_details
def test_pricing_a_frame_restores_a_breakdown_the_request_already_had(self):
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
from litellm.types.utils import ModelResponse, Usage
logging_obj = LiteLLMLoggingObj(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "test"}],
stream=True,
call_type="completion",
start_time=None,
litellm_call_id="lit4902-breakdown-restore",
function_id="lit4902-breakdown-restore",
)
logging_obj.update_environment_variables(litellm_params={}, optional_params={})
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
logging_obj.set_cost_breakdown(
input_cost=0.5, output_cost=0.25, total_cost=0.75, cost_for_built_in_tools_cost_usd_dollar=0.0
)
existing = logging_obj.cost_breakdown
model_response = ModelResponse(
usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)
)
ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj)
assert logging_obj.cost_breakdown is existing
assert logging_obj.cost_breakdown["total_cost"] == 0.75
def test_openai_chunk_prices_through_the_logging_obj_so_custom_pricing_applies(self):
"""The chat.completion.chunk path rides the same pricer, so a discounted deployment
streaming /v1/chat/completions gets its negotiated price instead of sticker."""
class _StubLoggingObj:
def __init__(self, cost):
self._cost = cost
self.captured_result = None
def _response_cost_calculator(self, result):
self.captured_result = result
return self._cost
discounted_cost = 0.00031
stub = _StubLoggingObj(discounted_cost)
event = {
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"choices": [],
"usage": {"prompt_tokens": 1000, "completion_tokens": 100, "total_tokens": 1100},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini", stub)
assert result is not None
assert result["usage"]["cost"] == discounted_cost
assert result["usage"]["cost"] != pytest.approx(self._expected_cost("gpt-4o-mini", 1000, 100))
usage = stub.captured_result.usage
assert usage.prompt_tokens == 1000
assert usage.completion_tokens == 100
def test_openai_chunk_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self):
class _StubLoggingObj:
def _response_cost_calculator(self, result):
return None
event = {
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"choices": [],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini", _StubLoggingObj())
assert result is not None
assert result["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
class TestProcessChunkWithCostInjection:
def test_complete_usage_frame_chunk_is_injected(self, monkeypatch):
@ -6116,6 +6364,31 @@ class TestProcessChunkWithCostInjection:
assert ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini") == chunk
def test_message_delta_frame_is_priced_with_the_logging_obj(self, monkeypatch):
"""Pins that the logging object reaches the pricer through the byte-frame entry point,
which is how the proxy actually calls this on a streamed Messages API request."""
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
class _StubLoggingObj:
def _response_cost_calculator(self, result):
return 0.00042
chunk = (
b"event: message_delta\n"
b'data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},'
b'"usage":{"input_tokens":14,"output_tokens":8,"cache_read_input_tokens":3202}}\n\n'
)
result = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
chunk, "claude-haiku-4-5", _StubLoggingObj()
)
assert result != chunk
data_line = next(ln for ln in result.decode("utf-8").splitlines() if ln.startswith("data:"))
payload = json.loads(data_line.split("data:", 1)[1].strip())
assert payload["usage"]["cost"] == 0.00042
assert payload["usage"]["cache_read_input_tokens"] == 3202
# ---------------------------------------------------------------------------
# SSE keepalive during the time-to-first-token (issue #34819)