mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
eaf61eb34b
commit
d3f5c6dbf6
2 changed files with 38 additions and 1 deletions
|
|
@ -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("_")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue