fix(cost_calculator): sum mirrored cache token fields once in combine_usage_objects

combine_usage_objects iterates prompt_tokens_details model_fields and sums each;
with cache_write_tokens and cache_creation_tokens now mirroring each other via
__setattr__, the pair was summed twice, doubling cache creation counts for
Anthropic batch cost calc, mid-stream fallback usage merges, and realtime usage.
Collapse the mirrored pair to one representative before summing.
This commit is contained in:
mateo-berri 2026-07-23 18:47:19 -07:00
parent eaf61eb34b
commit d3f5c6dbf6
2 changed files with 38 additions and 1 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

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