Merge pull request #34046 from BerriAI/litellm_openai_cache_write_tokens

fix(cost_tracking): map OpenAI cache_write_tokens for prompt cache creation billing
This commit is contained in:
Mateo Wang 2026-07-23 19:19:46 -07:00 • committed by GitHub
commit 15af874fcf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 378 additions and 12 deletions

View file

@ -2191,6 +2191,13 @@ def batch_cost_calculator(
return total_prompt_cost, total_completion_cost
def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> List[str]:
field_names = list(type(prompt_tokens_details).model_fields)
if getattr(prompt_tokens_details, "cache_write_tokens", None) is None:
return field_names
return [attr for attr in field_names if attr != "cache_creation_tokens"]
class BaseTokenUsageProcessor:
@staticmethod
def combine_usage_objects(usage_objects: List[Usage]) -> Usage:
@ -2225,7 +2232,7 @@ class BaseTokenUsageProcessor:
# Check what keys exist in the model's prompt_tokens_details
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
for attr in type(usage.prompt_tokens_details).model_fields:
for attr in _summable_prompt_token_fields(usage.prompt_tokens_details):
if (
hasattr(usage.prompt_tokens_details, attr)
and not attr.startswith("_")

View file

@ -457,7 +457,8 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
cache_creation_tokens = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
getattr(usage.prompt_tokens_details, "cache_write_tokens", 0)
or getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
)
or 0
)
@ -906,10 +907,6 @@ def get_token_type_cost_breakdown(
cache_read_tokens = prompt_tokens_details["cache_hit_tokens"]
cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"]
cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"]
# Some OpenAI-compatible providers (e.g. kimi-k2) report cache-write tokens
# under `cache_write_tokens`; mirror the total-cost normalization path.
if not cache_creation_tokens:
cache_creation_tokens = _coerce_token_count(getattr(usage.prompt_tokens_details, "cache_write_tokens", 0))
# Fall back to the private top-level counters the Usage constructor mirrors cache
# tokens onto, so providers/callers that bypass prompt_tokens_details are covered.
if not cache_read_tokens:

View file

@ -374,12 +374,22 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
if isinstance(v, BaseModel):
v = v.model_dump()
additional_usage_values.update({k: v})
if "cache_read_input_tokens" not in additional_usage_values:
prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
if isinstance(prompt_tokens_details, dict):
prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
if not isinstance(prompt_tokens_details, dict):
usage_object = clean_metadata.get("usage_object")
if isinstance(usage_object, dict):
prompt_tokens_details = usage_object.get("prompt_tokens_details")
if isinstance(prompt_tokens_details, dict):
if "cache_read_input_tokens" not in additional_usage_values:
cached_tokens = prompt_tokens_details.get("cached_tokens")
if isinstance(cached_tokens, int) and cached_tokens > 0:
additional_usage_values["cache_read_input_tokens"] = cached_tokens
if "cache_creation_input_tokens" not in additional_usage_values:
cache_write_tokens = prompt_tokens_details.get("cache_write_tokens") or prompt_tokens_details.get(
"cache_creation_tokens"
)
if isinstance(cache_write_tokens, int) and cache_write_tokens > 0:
additional_usage_values["cache_creation_input_tokens"] = cache_write_tokens
clean_metadata["additional_usage_values"] = additional_usage_values
if litellm.cache is not None:

View file

@ -1049,6 +1049,7 @@ class ResponseAPILoggingUtils:
audio_tokens=getattr(response_api_usage.input_tokens_details, "audio_tokens", None),
text_tokens=getattr(response_api_usage.input_tokens_details, "text_tokens", None),
image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None),
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
)
completion_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
output_tokens_details = getattr(response_api_usage, "output_tokens_details", None)

View file

@ -1534,14 +1534,27 @@ class PromptTokensDetailsWrapper(
audio_length_seconds: Optional[float] = None
"""Length of audio sent to the model. Used for multimodal embeddings priced per audio-second."""
cache_write_tokens: Optional[int] = None
"""Number of cache write (creation) tokens sent to the model. OpenAI naming (prompt_tokens_details.cache_write_tokens); this is the canonical field."""
cache_creation_tokens: Optional[int] = None
"""Number of cache creation tokens sent to the model. Used for Anthropic prompt caching."""
"""Number of cache creation tokens sent to the model. Anthropic/Bedrock naming; kept in sync with cache_write_tokens (assigning either mirrors to the other)."""
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
"""Details of cache creation tokens sent to the model. Used for tracking 5m/1h cache creation tokens for Anthropic prompt caching."""
def __setattr__(self, name: str, value: object) -> None:
super().__setattr__(name, value)
if name == "cache_write_tokens":
super().__setattr__("cache_creation_tokens", value)
elif name == "cache_creation_tokens":
super().__setattr__("cache_write_tokens", value)
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.cache_write_tokens = (
self.cache_write_tokens if self.cache_write_tokens is not None else self.cache_creation_tokens
)
if self.character_count is None:
del self.character_count
if self.image_count is None:
@ -1554,6 +1567,8 @@ class PromptTokensDetailsWrapper(
del self.web_search_requests
if self.tool_use_tokens is None:
del self.tool_use_tokens
if self.cache_write_tokens is None:
del self.cache_write_tokens
if self.cache_creation_tokens is None:
del self.cache_creation_tokens
if self.cache_creation_token_details is None:
@ -1662,10 +1677,10 @@ class Usage(SafeAttributeModel, CompletionUsage):
if "cache_creation_input_tokens" in params and isinstance(params["cache_creation_input_tokens"], int):
if _prompt_tokens_details is None:
_prompt_tokens_details = PromptTokensDetailsWrapper(
cache_creation_tokens=params["cache_creation_input_tokens"]
cache_write_tokens=params["cache_creation_input_tokens"]
)
else:
_prompt_tokens_details.cache_creation_tokens = params["cache_creation_input_tokens"]
_prompt_tokens_details.cache_write_tokens = params["cache_creation_input_tokens"]
super().__init__(
prompt_tokens=prompt_tokens or 0,

View file

@ -2111,6 +2111,37 @@ def test_token_type_cost_breakdown_reads_cache_write_tokens():
)
def test_generic_cost_per_token_openai_cache_write_tokens_gpt_5_6():
"""
Regression: OpenAI gpt-5.6 reports cache-write tokens under
prompt_tokens_details.cache_write_tokens (not the Anthropic cache_creation_tokens
name). Those tokens must be billed at the cache-write rate rather than the plain
input rate. Customer report: cache creation tokens were never counted for the
GPT-5.6 series, so cost was undercounted on cache-write requests.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "gpt-5.6"
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
total_tokens=1010,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, cache_write_tokens=800),
)
assert usage.prompt_tokens_details.cache_write_tokens == 800
assert usage.prompt_tokens_details.cache_creation_tokens == 800
prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="openai")
info = litellm.get_model_info(model=model, custom_llm_provider="openai")
expected_prompt = (1000 - 800) * info["input_cost_per_token"] + 800 * info["cache_creation_input_token_cost"]
assert prompt_cost == pytest.approx(expected_prompt)
assert info["cache_creation_input_token_cost"] > info["input_cost_per_token"]
assert prompt_cost > 1000 * info["input_cost_per_token"]
def test_token_type_cost_breakdown_reconciles_with_generic_total():
"""
Both-ways check: the reasoning subset must sum with the remaining (text) output
@ -2166,6 +2197,65 @@ def test_token_type_cost_breakdown_zero_without_special_tokens():
)
@pytest.mark.parametrize(
"raw_usage, expect_read, expect_write",
[
(
{
"input_tokens": 5000,
"output_tokens": 10,
"total_tokens": 5010,
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 4012},
},
False,
True,
),
(
{
"input_tokens": 5000,
"output_tokens": 10,
"total_tokens": 5010,
"input_tokens_details": {"cached_tokens": 4012, "cache_write_tokens": 0},
},
True,
False,
),
],
)
def test_token_type_cost_breakdown_openai_responses_api_cache_write_read(
raw_usage, expect_read, expect_write
):
"""Regression for #34309: OpenAI Responses API reports cache tokens under
input_tokens_details.{cached_tokens, cache_write_tokens}, not the Anthropic-style
top-level cache_creation_input_tokens. The itemized breakdown must still populate
cache_read_cost / cache_creation_cost from the transformed usage."""
from litellm.responses.utils import ResponseAPILoggingUtils
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "gpt-5.6"
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage)
breakdown = get_token_type_cost_breakdown(
model=model, custom_llm_provider="openai", usage=usage
)
info = litellm.get_model_info(model=model, custom_llm_provider="openai")
if expect_write:
assert breakdown.cache_creation_cost == pytest.approx(
4012 * info["cache_creation_input_token_cost"]
)
assert breakdown.cache_creation_cost > 0
assert breakdown.cache_read_cost == 0.0
if expect_read:
assert breakdown.cache_read_cost == pytest.approx(
4012 * info["cache_read_input_token_cost"]
)
assert breakdown.cache_read_cost > 0
assert breakdown.cache_creation_cost == 0.0
def test_token_type_cost_breakdown_handles_unknown_model_gracefully():
"""A model with no pricing must yield zeros, never raise."""
breakdown = get_token_type_cost_breakdown(

View file

@ -109,6 +109,143 @@ def test_get_logging_payload_does_not_map_missing_or_zero_cached_tokens(prompt_t
assert "cache_read_input_tokens" not in additional_usage_values
def test_get_logging_payload_maps_openai_cache_write_tokens_to_cache_creation_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=1000,
completion_tokens=2,
total_tokens=1002,
prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800},
)
)
assert additional_usage_values["cache_creation_input_tokens"] == 800
assert additional_usage_values["prompt_tokens_details"]["cache_write_tokens"] == 800
def test_get_logging_payload_preserves_anthropic_cache_creation_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=1000,
completion_tokens=2,
total_tokens=1002,
cache_creation_input_tokens=300,
)
)
assert additional_usage_values["cache_creation_input_tokens"] == 300
@pytest.mark.parametrize(
"prompt_tokens_details",
[None, {"cached_tokens": 100}, {"cached_tokens": 100, "cache_write_tokens": 0}],
)
def test_get_logging_payload_does_not_map_missing_or_zero_cache_write_tokens(prompt_tokens_details):
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=10,
completion_tokens=2,
total_tokens=12,
prompt_tokens_details=prompt_tokens_details,
)
)
assert "cache_creation_input_tokens" not in additional_usage_values
def _make_standard_logging_payload_with_usage_object(usage_object: dict) -> StandardLoggingPayload:
return StandardLoggingPayload(
id="test-id-responses",
call_type="responses",
stream=False,
response_cost=0.02,
status="success",
total_tokens=1010,
prompt_tokens=1000,
completion_tokens=10,
startTime=1234567890.0,
endTime=1234567891.0,
completionStartTime=None,
model_map_information=StandardLoggingModelInformation(model_map_key="gpt-5.6", model_map_value=None),
model="gpt-5.6",
model_id="model-123",
model_group="openai",
custom_llm_provider="openai",
api_base="https://api.openai.com",
metadata=StandardLoggingMetadata(
user_api_key_hash="test_hash",
user_api_key_alias=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_user_id=None,
user_api_key_team_alias=None,
spend_logs_metadata=None,
requester_ip_address=None,
requester_metadata=None,
user_api_key_end_user_id=None,
usage_object=usage_object,
),
cache_hit=False,
cache_key=None,
saved_cache_cost=0.0,
request_tags=[],
end_user=None,
requester_ip_address=None,
messages=[],
response={},
error_str=None,
model_parameters={},
hidden_params=StandardLoggingHiddenParams(
model_id="model-123",
cache_key=None,
api_base="https://api.openai.com",
response_cost="0.02",
litellm_overhead_time_ms=None,
additional_headers=None,
batch_models=None,
litellm_model_name=None,
usage_object=None,
),
)
def test_get_logging_payload_maps_responses_api_cache_write_tokens_from_usage_object():
"""Responses API (/v1/responses) usage is not chat-Usage-shaped, so
additional_usage_values can't derive cache tokens from response_obj.usage.
The Admin UI Logs "Cache Creation Tokens" row reads
additional_usage_values.cache_creation_input_tokens, so it must be filled
from the normalized standard_logging usage_object (LIT-4633)."""
standard_logging_payload = _make_standard_logging_payload_with_usage_object(
usage_object={
"prompt_tokens": 1000,
"completion_tokens": 10,
"total_tokens": 1010,
"prompt_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800, "cache_creation_tokens": 800},
}
)
payload = get_logging_payload(
kwargs={
"model": "gpt-5.6",
"call_type": "responses",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
"standard_logging_object": standard_logging_payload,
},
response_obj={
"id": "resp-test",
"usage": {
"input_tokens": 1000,
"output_tokens": 10,
"total_tokens": 1010,
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800},
},
},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
additional_usage_values = json.loads(payload["metadata"])["additional_usage_values"]
assert additional_usage_values["cache_creation_input_tokens"] == 800
def test_sanitize_request_body_for_spend_logs_payload_basic():
request_body = {
"messages": [{"role": "user", "content": "Hello, how are you?"}],

View file

@ -369,6 +369,32 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details.image_tokens == 272
assert result.completion_tokens_details.text_tokens == 100
def test_transform_response_api_usage_maps_cache_write_tokens(self):
"""Responses API (/v1/responses) cache-write tokens must survive the usage transform.
gpt-5.6 returns usage.input_tokens_details.cache_write_tokens (an extra field
not typed on InputTokensDetails). Before the fix the transform rebuilt the token
details and dropped it, leaving the cache-creation metric empty (LIT-4633).
"""
usage = {
"input_tokens": 10062,
"output_tokens": 16,
"total_tokens": 10078,
"input_tokens_details": {
"cached_tokens": 0,
"cache_write_tokens": 10059,
},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cache_write_tokens == 10059
assert result.prompt_tokens_details.cache_creation_tokens == 10059
assert result.prompt_tokens_details.cached_tokens == 0
def test_transform_response_api_usage_mixed_details(self):
"""Test transformation handles mixed token details (cached + image + audio)."""
# Setup - hypothetical usage with mixed token types

View file

@ -12,6 +12,7 @@ from pydantic import BaseModel
import litellm
from litellm.cost_calculator import (
BaseTokenUsageProcessor,
RealtimeAPITokenUsageProcessor,
completion_cost,
cost_per_token,
@ -3479,3 +3480,32 @@ def test_batch_cost_calculator_cache_creation_falls_back_to_input_rate():
)
assert prompt_cost == pytest.approx((1000 * 3e-6 + 8000 * 3e-7 + 2000 * 3e-6) / 2)
def test_combine_usage_objects_sums_mirrored_cache_write_fields_once():
"""
cache_write_tokens and cache_creation_tokens mirror each other on
PromptTokensDetailsWrapper, so field-iterating aggregation must sum the pair
once: a single 50-token usage stays 50 and two combine to 100, not double.
"""
single = Usage(
prompt_tokens=100,
completion_tokens=10,
total_tokens=110,
prompt_tokens_details=PromptTokensDetailsWrapper(cache_write_tokens=50),
)
combined = BaseTokenUsageProcessor.combine_usage_objects([single])
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.cache_write_tokens == 50
assert combined.prompt_tokens_details.cache_creation_tokens == 50
anthropic_style = Usage(
prompt_tokens=100,
completion_tokens=10,
total_tokens=110,
cache_creation_input_tokens=50,
)
combined_pair = BaseTokenUsageProcessor.combine_usage_objects([anthropic_style, anthropic_style])
assert combined_pair.prompt_tokens_details is not None
assert combined_pair.prompt_tokens_details.cache_write_tokens == 100
assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100

View file

@ -17,7 +17,9 @@ from litellm.types.utils import (
Delta,
LlmProviders,
ModelResponseStream,
PromptTokensDetailsWrapper,
StreamingChoices,
Usage,
)
from litellm.utils import (
ProviderConfigManager,
@ -34,6 +36,57 @@ from litellm.utils import (
# Adds the parent directory to the system path
def test_usage_openai_cache_write_tokens_populates_both_names():
"""OpenAI reports cache-write tokens as prompt_tokens_details.cache_write_tokens.
The Usage constructor must expose it under both cache_write_tokens (canonical,
OpenAI naming) and cache_creation_tokens (legacy, Anthropic naming)."""
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
total_tokens=1010,
prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800},
)
assert usage.prompt_tokens_details.cache_write_tokens == 800
assert usage.prompt_tokens_details.cache_creation_tokens == 800
def test_usage_anthropic_cache_creation_maps_to_cache_write_tokens():
"""Anthropic/Bedrock report the top-level cache_creation_input_tokens field.
It must be normalized onto the OpenAI cache_write_tokens name as well as the
legacy cache_creation_tokens name."""
usage = Usage(
prompt_tokens=500,
completion_tokens=50,
total_tokens=550,
cache_creation_input_tokens=300,
cache_read_input_tokens=120,
)
assert usage.prompt_tokens_details.cache_write_tokens == 300
assert usage.prompt_tokens_details.cache_creation_tokens == 300
assert usage.prompt_tokens_details.cached_tokens == 120
def test_prompt_tokens_details_no_cache_write_tokens_when_absent():
"""A read-only cache hit (no cache write) must not surface cache-write fields."""
details = PromptTokensDetailsWrapper(cached_tokens=800)
assert details.cached_tokens == 800
assert not hasattr(details, "cache_write_tokens")
assert not hasattr(details, "cache_creation_tokens")
def test_prompt_tokens_details_cache_write_creation_stay_in_sync_on_assignment():
"""Assigning either name after construction must mirror to the other, so a
caller that sets only one field can't leave the pair silently out of sync."""
details = PromptTokensDetailsWrapper(cache_write_tokens=100)
assert details.cache_write_tokens == details.cache_creation_tokens == 100
details.cache_write_tokens = 250
assert details.cache_write_tokens == details.cache_creation_tokens == 250
details.cache_creation_tokens = 375
assert details.cache_write_tokens == details.cache_creation_tokens == 375
@pytest.fixture
def local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost