mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(cost_calculator): route Fireworks through generic cache-aware cost calculator
The Fireworks cost_per_token function was charging all prompt tokens at the full input_cost_per_token rate, ignoring cache_read_input_token_cost even when usage.prompt_tokens_details.cached_tokens was present. Delegate to generic_cost_per_token (like DeepSeek, xAI, etc.) which already handles cache-read token discounts, tiered pricing, and other shared cost logic. Preserve the fallback to parameter-size-based pricing for models not in the cost map. Fixes #31714
This commit is contained in:
parent
cca71a07c2
commit
282f73b0cd
2 changed files with 123 additions and 25 deletions
|
|
@ -10,15 +10,13 @@ from litellm.constants import (
|
|||
FIREWORKS_AI_56_B_MOE,
|
||||
FIREWORKS_AI_176_B_MOE,
|
||||
)
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.types.utils import Usage
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
||||
# Extract the number of billion parameters from the model name
|
||||
# only used for together_computer LLMs
|
||||
def get_base_model_for_pricing(model_name: str) -> str:
|
||||
"""
|
||||
Helper function for calculating together ai pricing.
|
||||
Resolves a Fireworks model name to its pricing category based on parameter count.
|
||||
|
||||
Returns:
|
||||
- str: model pricing category if mapped else received model name
|
||||
|
|
@ -27,7 +25,6 @@ def get_base_model_for_pricing(model_name: str) -> str:
|
|||
|
||||
model_name = model_name.lower()
|
||||
|
||||
# Check for MoE models in the form <number>x<number>b
|
||||
moe_match = re.search(r"(\d+)x(\d+)b", model_name)
|
||||
if moe_match:
|
||||
total_billion = int(moe_match.group(1)) * int(moe_match.group(2))
|
||||
|
|
@ -36,13 +33,10 @@ def get_base_model_for_pricing(model_name: str) -> str:
|
|||
elif total_billion <= FIREWORKS_AI_176_B_MOE:
|
||||
return "fireworks-ai-56b-to-176b"
|
||||
|
||||
# Check for standard models in the form <number>b
|
||||
re_params_match = re.search(r"(\d+)b", model_name)
|
||||
if re_params_match is not None:
|
||||
params_match = str(re_params_match.group(1))
|
||||
params_billion = float(params_match)
|
||||
params_billion = float(re_params_match.group(1))
|
||||
|
||||
# Determine the category based on the number of parameters
|
||||
if params_billion <= FIREWORKS_AI_4_B:
|
||||
return "fireworks-ai-up-to-4b"
|
||||
elif params_billion <= FIREWORKS_AI_16_B:
|
||||
|
|
@ -50,7 +44,6 @@ def get_base_model_for_pricing(model_name: str) -> str:
|
|||
elif params_billion > FIREWORKS_AI_16_B:
|
||||
return "fireworks-ai-above-16b"
|
||||
|
||||
# If no matches, return the original model_name
|
||||
return "fireworks-ai-default"
|
||||
|
||||
|
||||
|
|
@ -60,25 +53,13 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
|
|||
|
||||
Input:
|
||||
- model: str, the model name without provider prefix
|
||||
- usage: LiteLLM Usage block, containing anthropic caching information
|
||||
- usage: LiteLLM Usage block, containing prompt caching information
|
||||
|
||||
Returns:
|
||||
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
|
||||
"""
|
||||
## check if model mapped, else use default pricing
|
||||
try:
|
||||
model_info = get_model_info(model=model, custom_llm_provider="fireworks_ai")
|
||||
return generic_cost_per_token(model=model, usage=usage, custom_llm_provider="fireworks_ai")
|
||||
except Exception:
|
||||
base_model = get_base_model_for_pricing(model_name=model)
|
||||
|
||||
## GET MODEL INFO
|
||||
model_info = get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
|
||||
|
||||
## CALCULATE INPUT COST
|
||||
|
||||
prompt_cost: float = usage["prompt_tokens"] * model_info["input_cost_per_token"]
|
||||
|
||||
## CALCULATE OUTPUT COST
|
||||
completion_cost = usage["completion_tokens"] * model_info["output_cost_per_token"]
|
||||
|
||||
return prompt_cost, completion_cost
|
||||
return generic_cost_per_token(model=base_model, usage=usage, custom_llm_provider="fireworks_ai")
|
||||
|
|
|
|||
|
|
@ -0,0 +1,117 @@
|
|||
"""
|
||||
Regression tests for litellm/llms/fireworks_ai/cost_calculator.py
|
||||
|
||||
Covers:
|
||||
- Cache-read tokens billed at cache_read_input_token_cost (issue #31714)
|
||||
- Fallback to parameter-size-based pricing for unknown models
|
||||
- Basic non-cached cost calculation still works
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.llms.fireworks_ai.cost_calculator import (
|
||||
cost_per_token,
|
||||
get_base_model_for_pricing,
|
||||
)
|
||||
from litellm.types.utils import PromptTokensDetails, Usage
|
||||
|
||||
|
||||
class TestCacheReadTokenCost:
|
||||
"""Regression: cached prompt tokens must be billed at cache_read_input_token_cost."""
|
||||
|
||||
def test_cached_tokens_billed_at_reduced_rate(self):
|
||||
"""
|
||||
Given a model with cache_read_input_token_cost defined and a usage object
|
||||
containing cached_tokens in prompt_tokens_details, the prompt cost must
|
||||
reflect the reduced rate for cached tokens rather than charging all
|
||||
prompt tokens at the full input rate.
|
||||
|
||||
Reproduces issue #31714.
|
||||
"""
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=100,
|
||||
total_tokens=1100,
|
||||
prompt_tokens_details=PromptTokensDetails(cached_tokens=800),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="accounts/fireworks/models/deepseek-v4-flash",
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
model_info = get_model_info(
|
||||
model="accounts/fireworks/models/deepseek-v4-flash",
|
||||
custom_llm_provider="fireworks_ai",
|
||||
)
|
||||
input_rate = model_info["input_cost_per_token"]
|
||||
cache_rate = model_info["cache_read_input_token_cost"]
|
||||
|
||||
non_cached_tokens = 1000 - 800
|
||||
expected_prompt_cost = non_cached_tokens * input_rate + 800 * cache_rate
|
||||
expected_completion_cost = 100 * model_info["output_cost_per_token"]
|
||||
|
||||
assert prompt_cost == pytest.approx(expected_prompt_cost)
|
||||
assert completion_cost == pytest.approx(expected_completion_cost)
|
||||
|
||||
naive_full_rate_cost = 1000 * input_rate
|
||||
assert prompt_cost < naive_full_rate_cost
|
||||
|
||||
def test_no_cached_tokens_charges_full_rate(self):
|
||||
"""Without cached tokens, all prompt tokens are billed at the full input rate."""
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=100,
|
||||
total_tokens=1100,
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="accounts/fireworks/models/deepseek-v4-flash",
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
model_info = get_model_info(
|
||||
model="accounts/fireworks/models/deepseek-v4-flash",
|
||||
custom_llm_provider="fireworks_ai",
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(1000 * model_info["input_cost_per_token"])
|
||||
assert completion_cost == pytest.approx(100 * model_info["output_cost_per_token"])
|
||||
|
||||
|
||||
class TestBaseModelFallback:
|
||||
"""Verify that unknown models fall back to parameter-size-based pricing."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name,expected_category",
|
||||
[
|
||||
("custom-3b-chat", "fireworks-ai-up-to-4b"),
|
||||
("custom-8b-instruct", "fireworks-ai-4.1b-to-16b"),
|
||||
("custom-70b-instruct", "fireworks-ai-above-16b"),
|
||||
("custom-8x7b-moe", "fireworks-ai-moe-up-to-56b"),
|
||||
("custom-8x22b-moe", "fireworks-ai-56b-to-176b"),
|
||||
],
|
||||
)
|
||||
def test_base_model_resolution(self, model_name: str, expected_category: str):
|
||||
assert get_base_model_for_pricing(model_name) == expected_category
|
||||
|
||||
def test_unknown_model_still_calculates_cost(self):
|
||||
"""An unregistered model with a recognizable param count yields a valid cost."""
|
||||
usage = Usage(
|
||||
prompt_tokens=500,
|
||||
completion_tokens=50,
|
||||
total_tokens=550,
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="my-custom-70b-chat",
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
assert prompt_cost > 0
|
||||
assert completion_cost > 0
|
||||
Loading…
Add table
Reference in a new issue