fix: derive DeepSeek prompt tokens from cache buckets

This commit is contained in:
Tanisha-Katara 2026-08-17 16:33:29 +04:00
parent 973329e986
commit 9cc265dd87
No known key found for this signature in database
GPG key ID: 943687E0EACCC73E
2 changed files with 72 additions and 3 deletions

View file

@ -1650,6 +1650,12 @@ class PromptTokensDetailsWrapper(
del self.cache_creation_token_details
def _non_negative_token_count(value: object) -> int | None:
if type(value) is int and value >= 0:
return value
return None
class ServerToolUse(BaseModel):
web_search_requests: int | None = None
tool_search_requests: int | None = None
@ -1736,12 +1742,20 @@ class Usage(SafeAttributeModel, CompletionUsage):
elif isinstance(prompt_tokens_details, PromptTokensDetailsWrapper):
_prompt_tokens_details = prompt_tokens_details
deepseek_cache_hit_tokens = _non_negative_token_count(params.get("prompt_cache_hit_tokens"))
deepseek_cache_miss_tokens = _non_negative_token_count(params.get("prompt_cache_miss_tokens"))
## DEEPSEEK MAPPING ##
if "prompt_cache_hit_tokens" in params and isinstance(params["prompt_cache_hit_tokens"], int):
if deepseek_cache_hit_tokens is not None:
if _prompt_tokens_details is None:
_prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=params["prompt_cache_hit_tokens"])
_prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=deepseek_cache_hit_tokens)
else:
_prompt_tokens_details.cached_tokens = params["prompt_cache_hit_tokens"]
_prompt_tokens_details.cached_tokens = deepseek_cache_hit_tokens
if not prompt_tokens:
deepseek_prompt_tokens = (deepseek_cache_hit_tokens or 0) + (deepseek_cache_miss_tokens or 0)
if deepseek_prompt_tokens:
prompt_tokens = deepseek_prompt_tokens
## ANTHROPIC MAPPING ##
if "cache_read_input_tokens" in params and isinstance(params["cache_read_input_tokens"], int):

View file

@ -3330,6 +3330,61 @@ def test_extract_cache_creation_tokens_zero_when_missing():
)
def test_usage_maps_deepseek_cache_hit_and_miss_tokens():
usage = Usage(
completion_tokens=50,
prompt_cache_hit_tokens=75,
prompt_cache_miss_tokens=25,
)
assert usage.prompt_tokens == 100
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.cached_tokens == 75
assert usage._cache_read_input_tokens == 75
def test_usage_does_not_override_reported_prompt_tokens():
usage = Usage(
prompt_tokens=120,
completion_tokens=50,
prompt_cache_hit_tokens=75,
prompt_cache_miss_tokens=25,
)
assert usage.prompt_tokens == 120
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.cached_tokens == 75
def test_usage_ignores_malformed_deepseek_cache_bucket_tokens():
usage = Usage(
completion_tokens=50,
prompt_cache_hit_tokens="75",
prompt_cache_miss_tokens=True,
)
assert usage.prompt_tokens == 0
assert usage.prompt_tokens_details is None
assert usage._cache_read_input_tokens == 0
def test_usage_preserves_cache_read_and_write_mappings():
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
cache_read_input_tokens=30,
cache_creation_input_tokens=10,
)
assert usage.prompt_tokens == 100
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.cached_tokens == 30
assert usage.prompt_tokens_details.cache_write_tokens == 10
assert usage.prompt_tokens_details.cache_creation_tokens == 10
assert usage._cache_read_input_tokens == 30
assert usage._cache_creation_input_tokens == 10
def test_custom_pricing_anthropic_style_cache_tokens_not_double_counted():
"""
Anthropic providers report cache tokens at the top level of Usage, and