mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Fix allow strings in calculate cost (#12200)
* Allow strings in calculate cost Sometimes the cost per unit is a string (e.g.: If a value like "3e-7" was read from the config.yaml) * Add comprehensive tests for string cost value handling - Added test_string_cost_values() to test basic string cost conversion functionality - Added test_calculate_cost_component_with_string_values() to test the calculate_cost_component function directly - Added test_string_cost_values_edge_cases() to test mixed string/float costs and error handling - Added test_string_cost_values_with_threshold() to test string costs with threshold pricing - Enhanced _get_token_base_cost() to handle string-to-float conversion for base costs and threshold costs - Enhanced generic_cost_per_token() to handle string-to-float conversion for audio and reasoning token costs - All tests cover scientific notation (e.g., '3e-7'), decimal notation (e.g., '0.000001'), and error handling for invalid strings - Maintains backward compatibility with existing float cost values * Dry up code * Fixed case where number was an integer * Allowing None --------- Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
parent
0dd50ea336
commit
e50789730d
2 changed files with 220 additions and 21 deletions
|
|
@ -114,8 +114,8 @@ def _get_token_base_cost(model_info: ModelInfo, usage: Usage) -> Tuple[float, fl
|
|||
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.
|
||||
"""
|
||||
prompt_base_cost = model_info["input_cost_per_token"]
|
||||
completion_base_cost = model_info["output_cost_per_token"]
|
||||
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"))
|
||||
|
||||
## CHECK IF ABOVE THRESHOLD
|
||||
threshold: Optional[float] = None
|
||||
|
|
@ -128,17 +128,13 @@ def _get_token_base_cost(model_info: ModelInfo, usage: Usage) -> Tuple[float, fl
|
|||
1000 if "k" in threshold_str else 1
|
||||
)
|
||||
if usage.prompt_tokens > threshold:
|
||||
prompt_base_cost = cast(
|
||||
float,
|
||||
model_info.get(key, prompt_base_cost),
|
||||
)
|
||||
completion_base_cost = cast(
|
||||
float,
|
||||
model_info.get(
|
||||
f"output_cost_per_token_above_{threshold_str}_tokens",
|
||||
completion_base_cost,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_base_cost = cast(float, _get_cost_per_unit(model_info, key, prompt_base_cost))
|
||||
completion_base_cost = cast(float, _get_cost_per_unit(
|
||||
model_info,
|
||||
f"output_cost_per_token_above_{threshold_str}_tokens",
|
||||
completion_base_cost,
|
||||
))
|
||||
break
|
||||
except (IndexError, ValueError):
|
||||
continue
|
||||
|
|
@ -162,7 +158,7 @@ def calculate_cost_component(
|
|||
Returns:
|
||||
float: The calculated cost
|
||||
"""
|
||||
cost_per_unit = model_info.get(cost_key)
|
||||
cost_per_unit = _get_cost_per_unit(model_info, cost_key)
|
||||
if (
|
||||
cost_per_unit is not None
|
||||
and isinstance(cost_per_unit, float)
|
||||
|
|
@ -173,6 +169,24 @@ def calculate_cost_component(
|
|||
return 0.0
|
||||
|
||||
|
||||
def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: Optional[float] = 0.0) -> Optional[float]:
|
||||
# Sometimes the cost per unit is a string (e.g.: If a value like "3e-7" was read from the config.yaml)
|
||||
cost_per_unit = model_info.get(cost_key)
|
||||
if isinstance(cost_per_unit, float):
|
||||
return cost_per_unit
|
||||
if isinstance(cost_per_unit, int):
|
||||
return float(cost_per_unit)
|
||||
if isinstance(cost_per_unit, str):
|
||||
try:
|
||||
return float(cost_per_unit)
|
||||
except ValueError:
|
||||
verbose_logger.exception(
|
||||
f"litellm.litellm_core_utils.llm_cost_calc.utils.py::calculate_cost_per_component(): Exception occured - {cost_per_unit}\nDefaulting to 0.0"
|
||||
)
|
||||
return default_value
|
||||
|
||||
|
||||
|
||||
def generic_cost_per_token(
|
||||
model: str, usage: Usage, custom_llm_provider: str
|
||||
) -> Tuple[float, float]:
|
||||
|
|
@ -316,13 +330,8 @@ def generic_cost_per_token(
|
|||
## TEXT COST
|
||||
completion_cost = float(text_tokens) * completion_base_cost
|
||||
|
||||
_output_cost_per_audio_token: Optional[float] = model_info.get(
|
||||
"output_cost_per_audio_token"
|
||||
)
|
||||
|
||||
_output_cost_per_reasoning_token: Optional[float] = model_info.get(
|
||||
"output_cost_per_reasoning_token"
|
||||
)
|
||||
_output_cost_per_audio_token = _get_cost_per_unit(model_info, "output_cost_per_audio_token", None)
|
||||
_output_cost_per_reasoning_token = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
|
||||
|
||||
## AUDIO COST
|
||||
if not is_text_tokens_total and audio_tokens is not None and audio_tokens > 0:
|
||||
|
|
|
|||
|
|
@ -172,3 +172,193 @@ def test_generic_cost_per_token_anthropic_prompt_caching():
|
|||
|
||||
print(f"prompt_cost: {prompt_cost}")
|
||||
assert prompt_cost < 0.085
|
||||
|
||||
|
||||
def test_string_cost_values():
|
||||
"""Test that cost values defined as strings are properly converted to floats."""
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock model info with string cost values (as might be read from config.yaml)
|
||||
mock_model_info = {
|
||||
"input_cost_per_token": "3e-7", # String representation of scientific notation
|
||||
"output_cost_per_token": "6e-7", # String representation of scientific notation
|
||||
"input_cost_per_audio_token": "0.000001", # String representation of decimal
|
||||
"output_cost_per_audio_token": "0.000002", # String representation of decimal
|
||||
"cache_read_input_token_cost": "1.5e-8", # String representation of scientific notation
|
||||
"cache_creation_input_token_cost": "2.5e-8", # String representation of scientific notation
|
||||
}
|
||||
|
||||
# Test usage with various token types
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=500,
|
||||
total_tokens=1500,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
audio_tokens=100,
|
||||
cached_tokens=200,
|
||||
text_tokens=700,
|
||||
image_tokens=None
|
||||
),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
audio_tokens=50,
|
||||
reasoning_tokens=None,
|
||||
text_tokens=450,
|
||||
accepted_prediction_tokens=None,
|
||||
rejected_prediction_tokens=None
|
||||
),
|
||||
_cache_creation_input_tokens=150
|
||||
)
|
||||
|
||||
# Mock get_model_info to return our mock model info
|
||||
with patch('litellm.litellm_core_utils.llm_cost_calc.utils.get_model_info', return_value=mock_model_info):
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model="test-model",
|
||||
usage=usage,
|
||||
custom_llm_provider="test-provider"
|
||||
)
|
||||
|
||||
# Calculate expected costs manually
|
||||
# Prompt cost = text_tokens * input_cost + audio_tokens * audio_cost + cached_tokens * cache_read_cost + cache_creation_tokens * cache_creation_cost
|
||||
expected_prompt_cost = (
|
||||
700 * 3e-7 + # text tokens
|
||||
100 * 1e-6 + # audio tokens
|
||||
200 * 1.5e-8 + # cached tokens
|
||||
150 * 2.5e-8 # cache creation tokens
|
||||
)
|
||||
|
||||
# Completion cost = text_tokens * output_cost + audio_tokens * audio_output_cost
|
||||
expected_completion_cost = (
|
||||
450 * 6e-7 + # text tokens
|
||||
50 * 2e-6 # audio tokens
|
||||
)
|
||||
|
||||
# Assert costs are calculated correctly
|
||||
assert round(prompt_cost, 12) == round(expected_prompt_cost, 12)
|
||||
assert round(completion_cost, 12) == round(expected_completion_cost, 12)
|
||||
|
||||
|
||||
def test_calculate_cost_component_with_string_values():
|
||||
"""Test the calculate_cost_component function directly with string cost values."""
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import calculate_cost_component
|
||||
|
||||
# Test with valid string scientific notation
|
||||
model_info = {"input_cost_per_token": "3e-7"}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 1000)
|
||||
assert cost == 1000 * 3e-7
|
||||
|
||||
# Test with valid string decimal notation
|
||||
model_info = {"output_cost_per_token": "0.000001"}
|
||||
cost = calculate_cost_component(model_info, "output_cost_per_token", 500)
|
||||
assert cost == 500 * 0.000001
|
||||
|
||||
# Test with float value (should work as before)
|
||||
model_info = {"input_cost_per_token": 3e-7}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 1000)
|
||||
assert cost == 1000 * 3e-7
|
||||
|
||||
# Test with invalid string value (should return 0.0)
|
||||
model_info = {"input_cost_per_token": "invalid_number"}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 1000)
|
||||
assert cost == 0.0
|
||||
|
||||
# Test with None value (should return 0.0)
|
||||
model_info = {"input_cost_per_token": None}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 1000)
|
||||
assert cost == 0.0
|
||||
|
||||
# Test with missing key (should return 0.0)
|
||||
model_info = {}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 1000)
|
||||
assert cost == 0.0
|
||||
|
||||
# Test with zero usage (should return 0.0)
|
||||
model_info = {"input_cost_per_token": "3e-7"}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", 0)
|
||||
assert cost == 0.0
|
||||
|
||||
# Test with None usage (should return 0.0)
|
||||
model_info = {"input_cost_per_token": "3e-7"}
|
||||
cost = calculate_cost_component(model_info, "input_cost_per_token", None)
|
||||
assert cost == 0.0
|
||||
|
||||
|
||||
def test_string_cost_values_edge_cases():
|
||||
"""Test edge cases for string cost value handling."""
|
||||
from unittest.mock import patch
|
||||
|
||||
# Test with mixed string and float cost values
|
||||
mock_model_info = {
|
||||
"input_cost_per_token": "1e-6", # String
|
||||
"output_cost_per_token": 2e-6, # Float
|
||||
"input_cost_per_audio_token": "invalid", # Invalid string
|
||||
"output_cost_per_audio_token": None, # None value
|
||||
}
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=500,
|
||||
total_tokens=1500,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
audio_tokens=100,
|
||||
cached_tokens=0,
|
||||
text_tokens=1000,
|
||||
image_tokens=None
|
||||
),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
audio_tokens=50,
|
||||
reasoning_tokens=None,
|
||||
text_tokens=500,
|
||||
accepted_prediction_tokens=None,
|
||||
rejected_prediction_tokens=None
|
||||
)
|
||||
)
|
||||
|
||||
with patch('litellm.litellm_core_utils.llm_cost_calc.utils.get_model_info', return_value=mock_model_info):
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model="test-model",
|
||||
usage=usage,
|
||||
custom_llm_provider="test-provider"
|
||||
)
|
||||
|
||||
# Expected costs:
|
||||
# Prompt: 1000 * 1e-6 + 100 * 0 (invalid string becomes 0)
|
||||
# Completion: 500 * 2e-6 (text_tokens == completion_tokens, so is_text_tokens_total=True, no separate audio cost)
|
||||
expected_prompt_cost = 1000 * 1e-6
|
||||
expected_completion_cost = 500 * 2e-6
|
||||
|
||||
assert round(prompt_cost, 12) == round(expected_prompt_cost, 12)
|
||||
assert round(completion_cost, 12) == round(expected_completion_cost, 12)
|
||||
|
||||
|
||||
def test_string_cost_values_with_threshold():
|
||||
"""Test that string cost values work correctly with threshold pricing."""
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock model info with string cost values including threshold pricing
|
||||
mock_model_info = {
|
||||
"input_cost_per_token": "1e-6", # String base cost
|
||||
"output_cost_per_token": "2e-6", # String base cost
|
||||
"input_cost_per_token_above_200k_tokens": "5e-7", # String threshold cost (lower)
|
||||
"output_cost_per_token_above_200k_tokens": "1e-6", # String threshold cost (lower)
|
||||
}
|
||||
|
||||
# Test usage above threshold
|
||||
usage = Usage(
|
||||
prompt_tokens=250000, # Above 200k threshold
|
||||
completion_tokens=1000,
|
||||
total_tokens=251000,
|
||||
)
|
||||
|
||||
with patch('litellm.litellm_core_utils.llm_cost_calc.utils.get_model_info', return_value=mock_model_info):
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model="test-model",
|
||||
usage=usage,
|
||||
custom_llm_provider="test-provider"
|
||||
)
|
||||
|
||||
# Expected costs using threshold pricing (string values converted to float)
|
||||
expected_prompt_cost = 250000 * 5e-7 # threshold cost
|
||||
expected_completion_cost = 1000 * 1e-6 # threshold cost
|
||||
|
||||
assert round(prompt_cost, 12) == round(expected_prompt_cost, 12)
|
||||
assert round(completion_cost, 12) == round(expected_completion_cost, 12)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue