fix(streaming): preserve 5m/1h cache creation split in combined usage

Anthropic-style streams only carry the cache creation TTL split
(usage.cache_creation.ephemeral_5m_input_tokens / ephemeral_1h_input_tokens)
in the message_start usage chunk; the closing message_delta chunk repeats
the totals without it. ChunkProcessor._calculate_usage_per_chunk overwrote
prompt_tokens_details with every successive usage chunk, so the split was
clobbered by the last chunk while the aggregate cache_creation_input_tokens
survived through its own keep-the-meaningful-chunk tracking

With the split gone, generic_cost_per_token had nothing to read from
usage.prompt_tokens_details.cache_creation_token_details and billed the
whole cache write at cache_creation_input_token_cost, the 5m rate. 1h-TTL
writes (for example Bedrock passthrough streaming with cache_control ttl
1h) were undercharged by the cache_creation_input_token_cost_above_1hr
delta

Track cache_creation_token_details across usage chunks the same way the
aggregate already is, expose it on UsagePerChunk, and reattach it to the
final usage's prompt_tokens_details in calculate_usage

Fixes #29432.
This commit is contained in:
Filippo Mattia Menghi 2026-06-10 09:32:15 +02:00
parent e15b37a18e
commit ce4c724e59
3 changed files with 115 additions and 1 deletions

View file

@ -7,6 +7,7 @@ from litellm.types.llms.openai import (
ChatCompletionAudioDelta,
)
from litellm.types.utils import (
CacheCreationTokenDetails,
ChatCompletionAudioResponse,
ChatCompletionMessageToolCall,
Choices,
@ -587,6 +588,7 @@ class ChunkProcessor:
completion_tokens = 0
## anthropic prompt caching information ##
cache_creation_input_tokens: Optional[int] = None
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
cache_read_input_tokens: Optional[int] = None
server_tool_use: Optional[ServerToolUse] = None
@ -651,6 +653,19 @@ class ChunkProcessor:
usage_chunk_dict["prompt_tokens_details"],
"web_search_requests",
)
if (
usage_chunk_dict["prompt_tokens_details"] is not None
and getattr(
usage_chunk_dict["prompt_tokens_details"],
"cache_creation_token_details",
None,
)
is not None
):
cache_creation_token_details = getattr(
usage_chunk_dict["prompt_tokens_details"],
"cache_creation_token_details",
)
prompt_tokens_details = usage_chunk_dict["prompt_tokens_details"]
@ -658,6 +673,7 @@ class ChunkProcessor:
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_creation_token_details=cache_creation_token_details,
cache_read_input_tokens=cache_read_input_tokens,
server_tool_use=server_tool_use,
web_search_requests=web_search_requests,
@ -686,6 +702,9 @@ class ChunkProcessor:
cache_creation_input_tokens: Optional[int] = calculated_usage_per_chunk[
"cache_creation_input_tokens"
]
cache_creation_token_details: Optional[CacheCreationTokenDetails] = (
calculated_usage_per_chunk["cache_creation_token_details"]
)
cache_read_input_tokens: Optional[int] = calculated_usage_per_chunk[
"cache_read_input_tokens"
]
@ -757,6 +776,15 @@ class ChunkProcessor:
)
if prompt_tokens_details is not None:
returned_usage.prompt_tokens_details = prompt_tokens_details
if cache_creation_token_details is not None:
if returned_usage.prompt_tokens_details is None:
returned_usage.prompt_tokens_details = PromptTokensDetailsWrapper(
cache_creation_token_details=cache_creation_token_details
)
else:
returned_usage.prompt_tokens_details.cache_creation_token_details = (
cache_creation_token_details
)
if server_tool_use is not None:
returned_usage.server_tool_use = server_tool_use

View file

@ -2,13 +2,19 @@ from typing import TYPE_CHECKING, Optional
from typing_extensions import TypedDict
from ..utils import CompletionTokensDetails, PromptTokensDetailsWrapper, ServerToolUse
from ..utils import (
CacheCreationTokenDetails,
CompletionTokensDetails,
PromptTokensDetailsWrapper,
ServerToolUse,
)
class UsagePerChunk(TypedDict):
prompt_tokens: int
completion_tokens: int
cache_creation_input_tokens: Optional[int]
cache_creation_token_details: Optional[CacheCreationTokenDetails]
cache_read_input_tokens: Optional[int]
server_tool_use: Optional[ServerToolUse]
web_search_requests: Optional[int]

View file

@ -11,12 +11,14 @@ sys.path.insert(
from litellm import stream_chunk_builder
from litellm.litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor
from litellm.types.utils import (
CacheCreationTokenDetails,
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Delta,
Function,
ModelResponseStream,
PromptTokensDetails,
PromptTokensDetailsWrapper,
ServerToolUse,
StreamingChoices,
Usage,
@ -325,6 +327,84 @@ def test_cache_read_input_tokens_retained():
assert usage.prompt_tokens_details.cached_tokens == 11775
def test_calculate_usage_retains_cache_creation_token_details():
"""
The 5m/1h cache creation split only arrives in the message_start usage
chunk; the final message_delta usage chunk must not clobber it, or 1h
cache writes get billed at the 5m rate.
"""
chunk1 = ModelResponseStream(
id="chatcmpl-c8889184-careful-fox-12345",
created=1745513206,
model="claude-opus-4-8",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content=""),
logprobs=None,
)
],
usage=Usage(
completion_tokens=1,
prompt_tokens=76,
total_tokens=77,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=31034,
cache_creation_tokens=362,
cache_creation_token_details=CacheCreationTokenDetails(
ephemeral_5m_input_tokens=288,
ephemeral_1h_input_tokens=74,
),
),
cache_creation_input_tokens=362,
cache_read_input_tokens=31034,
),
)
chunk2 = ModelResponseStream(
id="chatcmpl-c8889184-careful-fox-12345",
created=1745513207,
model="claude-opus-4-8",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content=None),
logprobs=None,
)
],
usage=Usage(
completion_tokens=259,
prompt_tokens=0,
total_tokens=259,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0),
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
),
)
chunks = [chunk1, chunk2]
processor = ChunkProcessor(chunks=chunks)
usage = processor.calculate_usage(
chunks=chunks,
model="claude-opus-4-8",
completion_output="",
)
assert usage.cache_creation_input_tokens == 362
assert usage.cache_read_input_tokens == 31034
cache_creation_token_details = (
usage.prompt_tokens_details.cache_creation_token_details
)
assert cache_creation_token_details is not None
assert cache_creation_token_details.ephemeral_5m_input_tokens == 288
assert cache_creation_token_details.ephemeral_1h_input_tokens == 74
def test_stream_chunk_builder_litellm_usage_chunks():
"""
Validate ChunkProcessor.calculate_usage uses provided usage fields from streaming chunks