fix(anthropic): preserve messages cache usage

This commit is contained in:
Kannan Priyadharshan 2026-06-23 19:17:08 +08:00
parent 69b0dd2da0
commit cd5ad02d63
5 changed files with 324 additions and 74 deletions

View file

@ -207,32 +207,13 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if "delta" not in merged_chunk:
merged_chunk["delta"] = {}
uncached_input_tokens = chunk.usage.prompt_tokens or 0
if (
hasattr(chunk.usage, "prompt_tokens_details")
and chunk.usage.prompt_tokens_details
):
cached_tokens = (
getattr(chunk.usage.prompt_tokens_details, "cached_tokens", 0) or 0
)
uncached_input_tokens -= cached_tokens
from .transformation import LiteLLMAnthropicMessagesAdapter
usage_dict: UsageDelta = {
"input_tokens": uncached_input_tokens,
"output_tokens": chunk.usage.completion_tokens or 0,
}
if (
hasattr(chunk.usage, "_cache_creation_input_tokens")
and chunk.usage._cache_creation_input_tokens > 0
):
usage_dict["cache_creation_input_tokens"] = (
chunk.usage._cache_creation_input_tokens
usage_dict: UsageDelta = (
LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage_delta(
chunk.usage
)
if (
hasattr(chunk.usage, "_cache_read_input_tokens")
and chunk.usage._cache_read_input_tokens > 0
):
usage_dict["cache_read_input_tokens"] = chunk.usage._cache_read_input_tokens
)
merged_chunk["usage"] = usage_dict
if self.applied_edits and "context_management" not in merged_chunk:
merged_chunk["context_management"] = ContextManagementResponse(

View file

@ -1402,6 +1402,95 @@ class LiteLLMAnthropicMessagesAdapter:
return "tool_use"
return "end_turn"
@staticmethod
def _positive_int(value: object) -> int:
if isinstance(value, int) and value > 0:
return value
return 0
@classmethod
def _first_positive_usage_value(
cls, usage: Usage, field_names: Tuple[str, ...]
) -> int:
for field_name in field_names:
value = cls._positive_int(getattr(usage, field_name, None))
if value > 0:
return value
return 0
@classmethod
def _first_positive_prompt_tokens_detail_value(
cls, usage: Usage, field_names: Tuple[str, ...]
) -> int:
prompt_tokens_details = getattr(usage, "prompt_tokens_details", None)
if prompt_tokens_details is None:
return 0
for field_name in field_names:
if isinstance(prompt_tokens_details, dict):
value = cls._positive_int(prompt_tokens_details.get(field_name))
else:
value = cls._positive_int(
getattr(prompt_tokens_details, field_name, None)
)
if value > 0:
return value
return 0
@classmethod
def _get_cache_read_input_tokens(cls, usage: Usage) -> int:
explicit_value = cls._first_positive_usage_value(
usage, ("cache_read_input_tokens", "_cache_read_input_tokens")
)
if explicit_value > 0:
return explicit_value
return cls._first_positive_prompt_tokens_detail_value(
usage, ("cached_tokens",)
)
@classmethod
def _get_cache_creation_input_tokens(cls, usage: Usage) -> int:
explicit_value = cls._first_positive_usage_value(
usage, ("cache_creation_input_tokens", "_cache_creation_input_tokens")
)
if explicit_value > 0:
return explicit_value
return cls._first_positive_prompt_tokens_detail_value(
usage, ("cache_creation_tokens", "cache_write_tokens")
)
@classmethod
def _translate_openai_usage_to_anthropic_usage_delta(
cls, usage: Usage
) -> UsageDelta:
cache_read_input_tokens = cls._get_cache_read_input_tokens(usage)
cache_creation_input_tokens = cls._get_cache_creation_input_tokens(usage)
input_tokens = max(
(usage.prompt_tokens or 0)
- cache_read_input_tokens
- cache_creation_input_tokens,
0,
)
usage_delta = UsageDelta(
input_tokens=input_tokens,
output_tokens=usage.completion_tokens or 0,
)
if cache_creation_input_tokens > 0:
usage_delta["cache_creation_input_tokens"] = cache_creation_input_tokens
if cache_read_input_tokens > 0:
usage_delta["cache_read_input_tokens"] = cache_read_input_tokens
return usage_delta
@classmethod
def _translate_openai_usage_to_anthropic_usage(
cls, usage: Usage
) -> AnthropicUsage:
return cast(
AnthropicUsage,
cls._translate_openai_usage_to_anthropic_usage_delta(usage),
)
def translate_openai_response_to_anthropic(
self,
response: ModelResponse,
@ -1433,32 +1522,12 @@ class LiteLLMAnthropicMessagesAdapter:
)
# extract usage
usage: Usage = getattr(response, "usage")
uncached_input_tokens = usage.prompt_tokens or 0
cached_tokens = 0
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
cached_tokens = (
getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
)
uncached_input_tokens -= cached_tokens
anthropic_usage = AnthropicUsage(
input_tokens=uncached_input_tokens,
output_tokens=usage.completion_tokens or 0,
)
if (
hasattr(usage, "_cache_creation_input_tokens")
and usage._cache_creation_input_tokens > 0
):
anthropic_usage["cache_creation_input_tokens"] = (
usage._cache_creation_input_tokens
)
if cached_tokens > 0:
anthropic_usage["cache_read_input_tokens"] = cached_tokens
anthropic_usage = self._translate_openai_usage_to_anthropic_usage(usage)
if polyfill_result is not None and polyfill_result.iterations_usage is not None:
message_iteration: UsageIteration = {
"type": "message",
"input_tokens": uncached_input_tokens,
"input_tokens": anthropic_usage["input_tokens"],
"output_tokens": usage.completion_tokens or 0,
}
anthropic_usage["iterations"] = list(polyfill_result.iterations_usage) + [message_iteration] # type: ignore[typeddict-unknown-key]
@ -1647,35 +1716,9 @@ class LiteLLMAnthropicMessagesAdapter:
else:
litellm_usage_chunk = None
if litellm_usage_chunk is not None:
uncached_input_tokens = litellm_usage_chunk.prompt_tokens or 0
cached_tokens = 0
if (
hasattr(litellm_usage_chunk, "prompt_tokens_details")
and litellm_usage_chunk.prompt_tokens_details
):
cached_tokens = (
getattr(
litellm_usage_chunk.prompt_tokens_details,
"cached_tokens",
0,
)
or 0
)
uncached_input_tokens -= cached_tokens
usage_delta = UsageDelta(
input_tokens=uncached_input_tokens,
output_tokens=litellm_usage_chunk.completion_tokens or 0,
usage_delta = self._translate_openai_usage_to_anthropic_usage_delta(
litellm_usage_chunk
)
if (
hasattr(litellm_usage_chunk, "_cache_creation_input_tokens")
and litellm_usage_chunk._cache_creation_input_tokens > 0
):
usage_delta["cache_creation_input_tokens"] = (
litellm_usage_chunk._cache_creation_input_tokens
)
if cached_tokens > 0:
usage_delta["cache_read_input_tokens"] = cached_tokens
else:
usage_delta = UsageDelta(input_tokens=0, output_tokens=0)
message_block = MessageBlockDelta(

View file

@ -2146,6 +2146,148 @@ def test_translate_openai_response_to_anthropic_cache_tokens_from_prompt_tokens_
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
def test_translate_openai_response_to_anthropic_cache_creation_from_prompt_tokens_details():
from litellm.types.utils import PromptTokensDetailsWrapper
usage = Usage(
prompt_tokens=120,
completion_tokens=50,
total_tokens=170,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30,
cache_creation_tokens=20,
),
)
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="gpt-4o-2024-08-06",
usage=usage,
)
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
assert anthropic_response["usage"]["input_tokens"] == 70
assert anthropic_response["usage"]["output_tokens"] == 50
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
assert anthropic_response["usage"]["cache_creation_input_tokens"] == 20
def test_translate_openai_response_to_anthropic_cache_tokens_from_usage_fields():
usage = Usage(prompt_tokens=120, completion_tokens=50, total_tokens=170)
usage.cache_read_input_tokens = 30
usage.cache_creation_input_tokens = 20
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="claude-3-sonnet-20240229",
usage=usage,
)
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
assert anthropic_response["usage"]["input_tokens"] == 70
assert anthropic_response["usage"]["output_tokens"] == 50
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
assert anthropic_response["usage"]["cache_creation_input_tokens"] == 20
def test_translate_openai_response_to_anthropic_cache_tokens_from_private_usage_fields():
usage = Usage(prompt_tokens=120, completion_tokens=50, total_tokens=170)
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="claude-3-sonnet-20240229",
usage=usage,
)
response.usage._cache_read_input_tokens = 30
response.usage._cache_creation_input_tokens = 20
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
assert anthropic_response["usage"]["input_tokens"] == 70
assert anthropic_response["usage"]["output_tokens"] == 50
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
assert anthropic_response["usage"]["cache_creation_input_tokens"] == 20
def test_translate_streaming_openai_response_to_anthropic_cache_tokens_from_prompt_tokens_details():
from litellm.types.utils import PromptTokensDetailsWrapper
usage = Usage(
prompt_tokens=120,
completion_tokens=50,
total_tokens=170,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30,
cache_creation_tokens=20,
),
)
response = ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(),
finish_reason="stop",
)
],
usage=usage,
)
adapter = LiteLLMAnthropicMessagesAdapter()
message_delta = adapter.translate_streaming_openai_response_to_anthropic(
response=response,
current_content_block_index=0,
)
assert message_delta["usage"]["input_tokens"] == 70
assert message_delta["usage"]["output_tokens"] == 50
assert message_delta["usage"]["cache_read_input_tokens"] == 30
assert message_delta["usage"]["cache_creation_input_tokens"] == 20
# =====================================================================
# Web Search Tool Transformation Tests
# =====================================================================

View file

@ -24,6 +24,7 @@ from litellm.types.utils import (
Message,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
StreamingChoices,
Usage,
)
@ -88,6 +89,59 @@ def test_fake_stream_usage_preserved():
assert message_delta["usage"]["input_tokens"] == 10
def test_delayed_usage_chunk_preserves_cache_tokens():
usage = Usage(
prompt_tokens=120,
completion_tokens=5,
total_tokens=125,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30,
cache_creation_tokens=20,
),
)
chunks = [
ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(content="Two."),
finish_reason=None,
)
],
),
ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(),
finish_reason="stop",
)
],
),
ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(),
finish_reason=None,
)
],
usage=usage,
),
]
wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="gpt-4o")
events = list(wrapper)
message_delta = next(
event for event in events if event.get("type") == "message_delta"
)
assert message_delta["usage"]["input_tokens"] == 70
assert message_delta["usage"]["output_tokens"] == 5
assert message_delta["usage"]["cache_read_input_tokens"] == 30
assert message_delta["usage"]["cache_creation_input_tokens"] == 20
def test_splitter_passes_through_non_combined_chunks():
"""A chunk with content but no finish_reason is not split."""
chunk = ModelResponseStream(

View file

@ -30,7 +30,9 @@ from litellm.types.utils import (
ChatCompletionDeltaToolCall,
Delta,
Function,
PromptTokensDetailsWrapper,
StreamingChoices,
Usage,
)
@ -107,6 +109,34 @@ def _input_json_deltas(events: List[dict]) -> List[str]:
]
def test_held_stop_reason_usage_merge_preserves_openai_cache_token_details():
"""OpenAI-compatible usage chunks carry cache reads in prompt_tokens_details."""
wrapper = AnthropicStreamWrapper(completion_stream=iter([]), model="claude-x")
wrapper.holding_stop_reason_chunk = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"input_tokens": 0, "output_tokens": 0},
}
usage_chunk = MagicMock()
usage_chunk.usage = Usage(
prompt_tokens=120,
completion_tokens=50,
total_tokens=170,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30,
cache_creation_tokens=20,
),
)
merged_chunk = wrapper._merge_usage_into_held_stop_reason_chunk(usage_chunk)
assert merged_chunk["usage"]["input_tokens"] == 70
assert merged_chunk["usage"]["output_tokens"] == 50
assert merged_chunk["usage"]["cache_read_input_tokens"] == 30
assert merged_chunk["usage"]["cache_creation_input_tokens"] == 20
def test_first_text_delta_after_tool_use_is_not_dropped_sync():
"""A tool_use -> text transition (text resuming after a tool call) carries
the resumed text's first token in the trigger chunk. Without the fix it was