diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index a40a8e1389c..96aed20529f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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("_") diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 9636db4f4cd..276ee96ed65 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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