Merge pull request #13759 from kankute-sameer/litellm_feat_correct_cost_calculations

Add long context support for claude-4-sonnet
This commit is contained in:
Krish Dholakia 2025-08-19 22:30:25 -07:00 • committed by GitHub
commit be30bc68ae
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 164 additions and 21 deletions

View file

@ -113,15 +113,20 @@ def _generic_cost_per_character(
return prompt_cost, completion_cost
def _get_token_base_cost(model_info: ModelInfo, usage: Usage) -> Tuple[float, float]:
def _get_token_base_cost(model_info: ModelInfo, usage: Usage) -> Tuple[float, float, float, float]:
"""
Return prompt cost for a given model and usage.
Return prompt cost, completion cost, and cache costs for a given model and usage.
If input_tokens > threshold and `input_cost_per_token_above_[x]k_tokens` or `input_cost_per_token_above_[x]_tokens` is set,
then we use the corresponding threshold cost.
then we use the corresponding threshold cost for all token types.
Returns:
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
"""
prompt_base_cost = cast(float, _get_cost_per_unit(model_info, "input_cost_per_token"))
completion_base_cost = cast(float, _get_cost_per_unit(model_info, "output_cost_per_token"))
cache_creation_cost = cast(float, _get_cost_per_unit(model_info, "cache_creation_input_token_cost"))
cache_read_cost = cast(float, _get_cost_per_unit(model_info, "cache_read_input_token_cost"))
## CHECK IF ABOVE THRESHOLD
threshold: Optional[float] = None
@ -141,13 +146,28 @@ def _get_token_base_cost(model_info: ModelInfo, usage: Usage) -> Tuple[float, fl
f"output_cost_per_token_above_{threshold_str}_tokens",
completion_base_cost,
))
# Apply tiered pricing to cache costs
cache_creation_tiered_key = f"cache_creation_input_token_cost_above_{threshold_str}_tokens"
cache_read_tiered_key = f"cache_read_input_token_cost_above_{threshold_str}_tokens"
if cache_creation_tiered_key in model_info:
cache_creation_cost = cast(float, _get_cost_per_unit(
model_info, cache_creation_tiered_key, cache_creation_cost
))
if cache_read_tiered_key in model_info:
cache_read_cost = cast(float, _get_cost_per_unit(
model_info, cache_read_tiered_key, cache_read_cost
))
break
except (IndexError, ValueError):
continue
except Exception:
continue
return prompt_base_cost, completion_base_cost
return prompt_base_cost, completion_base_cost, cache_creation_cost, cache_read_cost
def calculate_cost_component(
@ -262,28 +282,22 @@ def generic_cost_per_token(
if text_tokens == 0:
text_tokens = usage.prompt_tokens - cache_hit_tokens - audio_tokens
prompt_base_cost, completion_base_cost = _get_token_base_cost(
prompt_base_cost, completion_base_cost, cache_creation_cost, cache_read_cost = _get_token_base_cost(
model_info=model_info, usage=usage
)
prompt_cost = float(text_tokens) * prompt_base_cost
### CACHE READ COST
prompt_cost += calculate_cost_component(
model_info, "cache_read_input_token_cost", cache_hit_tokens
)
### CACHE READ COST - Now uses tiered pricing
prompt_cost += float(cache_hit_tokens) * cache_read_cost
### AUDIO COST
prompt_cost += calculate_cost_component(
model_info, "input_cost_per_audio_token", audio_tokens
)
### CACHE WRITING COST
prompt_cost += calculate_cost_component(
model_info,
"cache_creation_input_token_cost",
usage._cache_creation_input_tokens,
)
### CACHE WRITING COST - Now uses tiered pricing
prompt_cost += float(usage._cache_creation_input_tokens or 0) * cache_creation_cost
### CHARACTER COST

View file

@ -6437,11 +6437,13 @@
"supports_computer_use": true
},
"claude-4-sonnet-20250514": {
"max_tokens": 64000,
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3e-06,
"output_cost_per_token": 1.5e-05,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
@ -6449,6 +6451,8 @@
},
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "anthropic",
"mode": "chat",
"supports_function_calling": true,

View file

@ -6437,11 +6437,13 @@
"supports_computer_use": true
},
"claude-4-sonnet-20250514": {
"max_tokens": 64000,
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 1000000,
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"input_cost_per_token": 3e-06,
"output_cost_per_token": 1.5e-05,
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01,
@ -6449,6 +6451,8 @@
},
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"litellm_provider": "anthropic",
"mode": "chat",
"supports_function_calling": true,

View file

@ -482,6 +482,125 @@ def test_gemini_25_implicit_caching_cost():
print(f"✓ Gemini 2.5 implicit caching cost calculation is correct: ${result:.8f}")
def test_log_context_cost_calculation():
"""
Test that log context cost calculation works correctly with tiered pricing.
This test verifies that when using extended context (above 200k tokens),
the log context costs are calculated using the appropriate tiered rates.
"""
from litellm import completion_cost
from litellm.types.utils import (
Choices,
Message,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
# Create a mock response with extended context usage
extended_context_response = ModelResponse(
id="test-extended-context-response",
created=1750733889,
model="claude-4-sonnet-20250514",
object="chat.completion",
system_fingerprint=None,
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="This is a test response for extended context cost calculation.",
role="assistant",
tool_calls=None,
function_call=None,
),
)
],
usage=Usage(
total_tokens=350000, # Above 200k threshold
prompt_tokens=300000, # Above 200k threshold
completion_tokens=50000,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=300000,
cached_tokens=0, # No cache hits
audio_tokens=None,
image_tokens=None,
character_count=None,
video_length_seconds=None,
),
completion_tokens_details=None,
_cache_creation_input_tokens=1000, # Some tokens added to cache
),
)
# Calculate the cost using the extended context model
result = completion_cost(
completion_response=extended_context_response,
model="claude-4-sonnet-20250514",
custom_llm_provider="anthropic",
)
# Debug: Print the actual result
print(f"DEBUG: Actual cost result: ${result:.6f}")
# Get model info to understand the pricing
from litellm import get_model_info
model_info = get_model_info(model="claude-4-sonnet-20250514", custom_llm_provider="anthropic")
# Calculate expected cost based on actual model pricing
input_cost_per_token = model_info.get("input_cost_per_token", 0)
output_cost_per_token = model_info.get("output_cost_per_token", 0)
cache_creation_cost_per_token = model_info.get("cache_creation_input_token_cost", 0)
# Check if tiered pricing is applied
input_cost_above_200k = model_info.get("input_cost_per_token_above_200k_tokens", input_cost_per_token)
output_cost_above_200k = model_info.get("output_cost_per_token_above_200k_tokens", output_cost_per_token)
cache_creation_above_200k = model_info.get("cache_creation_input_token_cost_above_200k_tokens", cache_creation_cost_per_token)
print(f"DEBUG: Base input cost per token: ${input_cost_per_token:.2e}")
print(f"DEBUG: Base output cost per token: ${output_cost_per_token:.2e}")
print(f"DEBUG: Base cache creation cost per token: ${cache_creation_cost_per_token:.2e}")
# Handle tiered pricing - if not available, use base pricing
if input_cost_above_200k is not None:
print(f"DEBUG: Tiered input cost per token (>200k): ${input_cost_above_200k:.2e}")
else:
print(f"DEBUG: No tiered input pricing available, using base pricing")
input_cost_above_200k = input_cost_per_token
if output_cost_above_200k is not None:
print(f"DEBUG: Tiered output cost per token (>200k): ${output_cost_above_200k:.2e}")
else:
print(f"DEBUG: No tiered output pricing available, using base pricing")
output_cost_above_200k = output_cost_per_token
if cache_creation_above_200k is not None:
print(f"DEBUG: Tiered cache creation cost per token (>200k): ${cache_creation_above_200k:.2e}")
else:
print(f"DEBUG: No tiered cache creation pricing available, using base pricing")
cache_creation_above_200k = cache_creation_cost_per_token
# Since we're above 200k tokens, we should use tiered pricing if available
expected_input_cost = 300000 * input_cost_above_200k
expected_output_cost = 50000 * output_cost_above_200k
expected_cache_cost = 1000 * cache_creation_above_200k
expected_total = expected_input_cost + expected_output_cost + expected_cache_cost
print(f"DEBUG: Expected total: ${expected_total:.6f}")
# Allow for small floating point differences
assert (
abs(result - expected_total) < 1e-6
), f"Expected cost ${expected_total:.6f}, but got ${result:.6f}"
print(f"✓ Log context cost calculation with tiered pricing is correct: ${result:.6f}")
print(f" - Input tokens (300k): ${expected_input_cost:.6f}")
print(f" - Output tokens (50k): ${expected_output_cost:.6f}")
print(f" - Cache creation (1k): ${expected_cache_cost:.6f}")
print(f" - Total: ${result:.6f}")
def test_gemini_25_explicit_caching_cost_direct_usage():
"""
Test that Gemini 2.5 models correctly calculate costs with explicit caching.

View file

@ -497,7 +497,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_computer_use": {"type": "boolean"},
"cache_creation_input_audio_token_cost": {"type": "number"},
"cache_creation_input_token_cost": {"type": "number"},
"cache_creation_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_read_input_token_cost": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_read_input_audio_token_cost": {"type": "number"},
"deprecation_date": {"type": "string"},
"input_cost_per_audio_per_second": {"type": "number"},