mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
15af874fcf
10 changed files with 378 additions and 12 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("_")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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?"}],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue