mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(utils.py): log cache_creation_tokens in prompt token details
Closes LIT-907
This commit is contained in:
parent
1162c52692
commit
e488312873
5 changed files with 64 additions and 7 deletions
|
|
@ -278,6 +278,13 @@ def generic_cost_per_token(
|
|||
)
|
||||
or 0
|
||||
)
|
||||
cache_creation_tokens = (
|
||||
cast(
|
||||
Optional[int],
|
||||
getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
|
||||
)
|
||||
or 0
|
||||
)
|
||||
text_tokens = (
|
||||
cast(
|
||||
Optional[int], getattr(usage.prompt_tokens_details, "text_tokens", None)
|
||||
|
|
@ -307,9 +314,8 @@ def generic_cost_per_token(
|
|||
or 0
|
||||
)
|
||||
|
||||
if getattr(usage, "_cache_creation_input_tokens", 0) is not None:
|
||||
cache_creation_tokens = usage._cache_creation_input_tokens
|
||||
## EDGE CASE - text tokens not set inside PromptTokensDetails
|
||||
|
||||
if text_tokens == 0:
|
||||
text_tokens = (
|
||||
usage.prompt_tokens
|
||||
|
|
@ -333,7 +339,7 @@ def generic_cost_per_token(
|
|||
)
|
||||
|
||||
### CACHE WRITING COST - Now uses tiered pricing
|
||||
prompt_cost += float(usage._cache_creation_input_tokens or 0) * cache_creation_cost
|
||||
prompt_cost += float(cache_creation_tokens) * cache_creation_cost
|
||||
|
||||
### CHARACTER COST
|
||||
|
||||
|
|
|
|||
|
|
@ -162,7 +162,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
SearchContextCostPerQuery
|
||||
] # Cost for using web search tool
|
||||
citation_cost_per_token: Optional[float] # Cost per citation token for Perplexity
|
||||
tiered_pricing: Optional[List[Dict[str, Any]]] # Tiered pricing structure for models like Dashscope
|
||||
tiered_pricing: Optional[
|
||||
List[Dict[str, Any]]
|
||||
] # Tiered pricing structure for models like Dashscope
|
||||
litellm_provider: Required[str]
|
||||
mode: Required[
|
||||
Literal[
|
||||
|
|
@ -880,6 +882,9 @@ class PromptTokensDetailsWrapper(
|
|||
video_length_seconds: Optional[float] = None
|
||||
"""Length of videos sent to the model. Used for Vertex AI multimodal embeddings."""
|
||||
|
||||
cache_creation_tokens: Optional[int] = None
|
||||
"""Number of cache creation tokens sent to the model. Used for Anthropic prompt caching."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if self.character_count is None:
|
||||
|
|
@ -890,6 +895,8 @@ class PromptTokensDetailsWrapper(
|
|||
del self.video_length_seconds
|
||||
if self.web_search_requests is None:
|
||||
del self.web_search_requests
|
||||
if self.cache_creation_tokens is None:
|
||||
del self.cache_creation_tokens
|
||||
|
||||
|
||||
class ServerToolUse(BaseModel):
|
||||
|
|
@ -951,6 +958,7 @@ class Usage(CompletionUsage):
|
|||
# handle prompt_tokens_details
|
||||
_prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
|
||||
|
||||
# guarantee prompt_token_details is always a PromptTokensDetailsWrapper
|
||||
if prompt_tokens_details:
|
||||
if isinstance(prompt_tokens_details, dict):
|
||||
_prompt_tokens_details = PromptTokensDetailsWrapper(
|
||||
|
|
@ -985,6 +993,18 @@ class Usage(CompletionUsage):
|
|||
else:
|
||||
_prompt_tokens_details.cached_tokens = params["cache_read_input_tokens"]
|
||||
|
||||
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"]
|
||||
)
|
||||
else:
|
||||
_prompt_tokens_details.cache_creation_tokens = params[
|
||||
"cache_creation_input_tokens"
|
||||
]
|
||||
|
||||
super().__init__(
|
||||
prompt_tokens=prompt_tokens or 0,
|
||||
completion_tokens=completion_tokens or 0,
|
||||
|
|
|
|||
|
|
@ -954,6 +954,7 @@ class BaseLLMChatTest(ABC):
|
|||
|
||||
@pytest.mark.flaky(retries=4, delay=1)
|
||||
def test_prompt_caching(self):
|
||||
print("test_prompt_caching")
|
||||
litellm.set_verbose = True
|
||||
from litellm.utils import supports_prompt_caching
|
||||
|
||||
|
|
@ -1049,8 +1050,8 @@ class BaseLLMChatTest(ABC):
|
|||
assert (
|
||||
response.usage.prompt_tokens_details.cached_tokens > 0
|
||||
), f"cached_tokens={response.usage.prompt_tokens_details.cached_tokens} should be greater than 0. Got usage={response.usage}"
|
||||
except litellm.InternalServerError:
|
||||
pass
|
||||
except litellm.InternalServerError as e:
|
||||
print("InternalServerError", e)
|
||||
|
||||
@pytest.fixture
|
||||
def pdf_messages(self):
|
||||
|
|
|
|||
|
|
@ -250,7 +250,7 @@ async def test_anthropic_api_prompt_caching_basic():
|
|||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement"
|
||||
* 400,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
: {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
},
|
||||
|
|
@ -510,6 +510,7 @@ async def test_anthropic_api_prompt_caching_streaming():
|
|||
if hasattr(chunk, "usage") and hasattr(
|
||||
chunk.usage, "cache_creation_input_tokens"
|
||||
):
|
||||
print("chunk.usage", chunk.usage)
|
||||
is_cache_creation_input_tokens_in_usage = True
|
||||
|
||||
idx += 1
|
||||
|
|
|
|||
|
|
@ -174,6 +174,35 @@ def test_generic_cost_per_token_anthropic_prompt_caching():
|
|||
assert prompt_cost < 0.085
|
||||
|
||||
|
||||
def test_generic_cost_per_token_anthropic_prompt_caching_with_cache_creation():
|
||||
model = "claude-3-5-haiku-20241022"
|
||||
usage = Usage(
|
||||
completion_tokens=90,
|
||||
prompt_tokens=28436,
|
||||
total_tokens=28526,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=None,
|
||||
audio_tokens=None,
|
||||
reasoning_tokens=0,
|
||||
rejected_prediction_tokens=None,
|
||||
text_tokens=None,
|
||||
),
|
||||
prompt_tokens_details=None,
|
||||
cache_creation_input_tokens=2000,
|
||||
)
|
||||
|
||||
custom_llm_provider = "anthropic"
|
||||
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
print(f"prompt_cost: {prompt_cost}")
|
||||
assert round(prompt_cost, 3) == 0.023
|
||||
|
||||
|
||||
def test_string_cost_values():
|
||||
"""Test that cost values defined as strings are properly converted to floats."""
|
||||
from unittest.mock import patch
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue