fix(fireworks_ai): bill prompt-cache hits at cache_read rate (#33714)

Co-authored-by: Krrish Dholakia <krrishdholakia@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-07-17 10:35:16 -07:00 • committed by GitHub
parent b0a0f11b09
commit 0e88b57ec2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 81 additions and 2 deletions

View file

@ -75,10 +75,23 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
model_info = get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
## CALCULATE INPUT COST
prompt_tokens_details = usage.prompt_tokens_details
cached_tokens: int = (
prompt_tokens_details.cached_tokens
if prompt_tokens_details is not None and prompt_tokens_details.cached_tokens is not None
else 0
)
input_cost_per_token: float = model_info["input_cost_per_token"] or 0.0
cache_read_input_token_cost = model_info.get("cache_read_input_token_cost")
cache_read_cost_per_token: float = (
cache_read_input_token_cost if cache_read_input_token_cost is not None else input_cost_per_token
)
non_cached_prompt_tokens: int = max(usage.prompt_tokens - cached_tokens, 0)
prompt_cost: float = usage["prompt_tokens"] * model_info["input_cost_per_token"]
prompt_cost: float = non_cached_prompt_tokens * input_cost_per_token + cached_tokens * cache_read_cost_per_token
## CALCULATE OUTPUT COST
completion_cost = usage["completion_tokens"] * model_info["output_cost_per_token"]
output_cost_per_token: float = model_info["output_cost_per_token"] or 0.0
completion_cost: float = usage.completion_tokens * output_cost_per_token
return prompt_cost, completion_cost

View file

@ -0,0 +1,66 @@
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
MODEL = "accounts/fireworks/models/glm-5p2"
INPUT_COST = 1.4e-06
CACHE_READ_COST = 2.6e-07
OUTPUT_COST = 4.4e-06
def _usage(prompt_tokens: int, cached_tokens: int, completion_tokens: int) -> Usage:
return Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
)
def test_cached_prompt_tokens_billed_at_cache_read_rate():
prompt_tokens = 7036
cached_tokens = 7020
completion_tokens = 8
prompt_cost, completion_cost = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, cached_tokens, completion_tokens)
)
expected_prompt_cost = (prompt_tokens - cached_tokens) * INPUT_COST + cached_tokens * CACHE_READ_COST
assert prompt_cost == pytest.approx(expected_prompt_cost)
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
full_rate_cost = prompt_tokens * INPUT_COST
assert prompt_cost < full_rate_cost
def test_warm_call_cheaper_than_cold_call():
prompt_tokens = 7036
completion_tokens = 8
cold_prompt_cost, _ = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 16, completion_tokens)
)
warm_prompt_cost, _ = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 7020, completion_tokens)
)
assert warm_prompt_cost < cold_prompt_cost
def test_no_cached_tokens_matches_full_input_rate():
prompt_tokens = 100
completion_tokens = 10
prompt_cost, completion_cost = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 0, completion_tokens)
)
assert prompt_cost == pytest.approx(prompt_tokens * INPUT_COST)
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)