mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(anthropic): preserve messages cache usage
This commit is contained in:
parent
69b0dd2da0
commit
cd5ad02d63
5 changed files with 324 additions and 74 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# =====================================================================
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue