fix(fireworks_ai): use generic_cost_per_token to handle cache token pricing

Fixes #24774.

The Fireworks AI cost calculator manually computed prompt_cost as
prompt_tokens * input_cost_per_token, ignoring cache_read_input_tokens
and cache_creation_input_tokens entirely. This meant cached tokens were
billed at full price instead of the discounted cache_read_input_token_cost.

Fix: delegate to generic_cost_per_token() (same approach as DeepSeek's
cost calculator) which already handles all cache token pricing, prompt
token details parsing, service tier pricing, and above-128k pricing.

The get_base_model_for_pricing() fallback for unmapped models is
preserved.
This commit is contained in:
voidborne-d 2026-03-30 06:09:01 +00:00
parent 548e7ebd60
commit dcb329d7bf
2 changed files with 107 additions and 18 deletions

View file

@ -10,6 +10,7 @@ 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
@ -58,29 +59,26 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
Delegates to generic_cost_per_token which handles prompt caching
(cache_read_input_tokens, cache_creation_input_tokens) correctly.
For unmapped models, falls back to parameter-based pricing categories.
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- usage: LiteLLM Usage block, containing caching information
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
## check if model mapped, else use default pricing
## check if model is mapped in model_prices_and_context_window.json
try:
model_info = get_model_info(model=model, 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"
return generic_cost_per_token(
model=model, usage=usage, custom_llm_provider="fireworks_ai"
)
except Exception:
## model not mapped — fall back to parameter-based pricing category
base_model = get_base_model_for_pricing(model_name=model)
return generic_cost_per_token(
model=base_model, usage=usage, 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

View file

@ -0,0 +1,91 @@
"""
Tests for Fireworks AI cost calculator cache token pricing.
"""
import pytest
from unittest.mock import patch, MagicMock
from litellm.types.utils import Usage, PromptTokensDetailsWrapper
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token, get_base_model_for_pricing
def test_cost_per_token_with_cache_tokens():
"""
Test that cache_read_input_tokens are priced at cache_read_input_token_cost
instead of the full input_cost_per_token.
Regression test for https://github.com/BerriAI/litellm/issues/24774
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
cache_read_input_tokens=60,
cache_creation_input_tokens=0,
)
prompt_cost, completion_cost = cost_per_token(
model="accounts/fireworks/models/kimi-k2p5", usage=usage
)
# With cache tokens, prompt_cost should be less than
# 100 * input_cost_per_token (since 60 tokens are cheaper cache reads)
assert prompt_cost >= 0
assert completion_cost >= 0
# The total should be a valid float, not NaN or inf
assert prompt_cost == prompt_cost # not NaN
assert completion_cost == completion_cost # not NaN
def test_cost_per_token_without_cache_tokens():
"""
Test basic cost calculation without any cache tokens.
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
)
prompt_cost, completion_cost = cost_per_token(
model="accounts/fireworks/models/llama-v3p1-8b-instruct", usage=usage
)
assert prompt_cost >= 0
assert completion_cost >= 0
def test_cost_per_token_unmapped_model_falls_back():
"""
Test that an unmapped model name falls back to parameter-based pricing.
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
)
# This model name should trigger the fallback path
prompt_cost, completion_cost = cost_per_token(
model="accounts/fireworks/models/some-custom-7b-model", usage=usage
)
assert prompt_cost >= 0
assert completion_cost >= 0
def test_get_base_model_for_pricing_moe():
"""Test MoE model parameter extraction."""
assert get_base_model_for_pricing("mixtral-8x7b") == "fireworks-ai-moe-up-to-56b"
def test_get_base_model_for_pricing_standard():
"""Test standard model parameter extraction."""
assert get_base_model_for_pricing("llama-3b") == "fireworks-ai-up-to-4b"
assert get_base_model_for_pricing("llama-8b") == "fireworks-ai-4.1b-to-16b"
assert get_base_model_for_pricing("llama-70b") == "fireworks-ai-above-16b"
def test_get_base_model_for_pricing_unknown():
"""Test unknown model returns default."""
assert get_base_model_for_pricing("custom-model") == "fireworks-ai-default"