mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
Merge pull request #41917 from BerriAI/litellm_fireworks_cache_read_default_discount
fix(cost_calc): default fireworks cached input to the documented 50% discount when the map has no cache-read rate
This commit is contained in:
commit
c4ddea7cfc
7 changed files with 269 additions and 49 deletions
|
|
@ -578,6 +578,9 @@ FIREWORKS_AI_176_B_MOE: Final = int(os.getenv("FIREWORKS_AI_176_B_MOE", 176))
|
|||
FIREWORKS_AI_4_B: Final = int(os.getenv("FIREWORKS_AI_4_B", 4))
|
||||
FIREWORKS_AI_16_B: Final = int(os.getenv("FIREWORKS_AI_16_B", 16))
|
||||
FIREWORKS_AI_80_B: Final = int(os.getenv("FIREWORKS_AI_80_B", 80))
|
||||
# https://docs.fireworks.ai/guides/prompt-caching (accessed 2026-09-19): serverless cached prompt tokens
|
||||
# default to a 50% discount off the input rate
|
||||
FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO: Final = 0.5
|
||||
#### Logging callback constants ####
|
||||
REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM"
|
||||
MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import (
|
|||
select_tier_for_input,
|
||||
tier_rate,
|
||||
)
|
||||
from litellm.llms.fireworks_ai.cache_pricing import with_default_cache_read_rate
|
||||
from litellm.types.utils import (
|
||||
CacheCreationTokenDetails,
|
||||
CallTypes,
|
||||
|
|
@ -72,6 +73,12 @@ def _uses_inclusive_token_thresholds(custom_llm_provider: str | None) -> bool:
|
|||
return custom_llm_provider in _INCLUSIVE_THRESHOLD_PROVIDERS
|
||||
|
||||
|
||||
def apply_provider_cache_read_default(model_info: ModelInfo, custom_llm_provider: str | None) -> ModelInfo:
|
||||
if custom_llm_provider == "fireworks_ai":
|
||||
return with_default_cache_read_rate(model_info)
|
||||
return model_info
|
||||
|
||||
|
||||
def _get_token_detail_value(details: object, key: str) -> int | None:
|
||||
if isinstance(details, dict):
|
||||
value = details.get(key)
|
||||
|
|
@ -1170,8 +1177,10 @@ def generic_cost_per_token(
|
|||
# rather than handing back a name for this to re-resolve. A name cannot express a
|
||||
# per-deployment override: those are registered under the deployment id and kept off
|
||||
# the shared model-name key, so resolving from the name here reads the public rate.
|
||||
if model_info is None:
|
||||
model_info = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
resolved_model_info: Final = apply_provider_cache_read_default(
|
||||
get_model_info(model=model, custom_llm_provider=custom_llm_provider) if model_info is None else model_info,
|
||||
custom_llm_provider,
|
||||
)
|
||||
|
||||
## CALCULATE INPUT COST
|
||||
### Cost of processing (non-cache hit + cache hit) + Cost of cache-writing (cache writing)
|
||||
|
|
@ -1236,7 +1245,7 @@ def generic_cost_per_token(
|
|||
cache_creation_cost_above_1hr,
|
||||
cache_read_cost,
|
||||
) = _get_token_base_cost(
|
||||
model_info=model_info,
|
||||
model_info=resolved_model_info,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
current_time=billing_time,
|
||||
|
|
@ -1245,7 +1254,7 @@ def generic_cost_per_token(
|
|||
|
||||
prompt_cost = _calculate_input_cost(
|
||||
prompt_tokens_details=prompt_tokens_details,
|
||||
model_info=model_info,
|
||||
model_info=resolved_model_info,
|
||||
prompt_base_cost=prompt_base_cost,
|
||||
cache_read_cost=cache_read_cost,
|
||||
cache_creation_cost=cache_creation_cost,
|
||||
|
|
@ -1290,7 +1299,7 @@ def generic_cost_per_token(
|
|||
|
||||
## AUDIO COST
|
||||
if not is_text_tokens_total and audio_tokens is not None and audio_tokens > 0:
|
||||
_output_cost_per_audio_token = _get_cost_per_unit(model_info, "output_cost_per_audio_token", None)
|
||||
_output_cost_per_audio_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_audio_token", None)
|
||||
_output_cost_per_audio_token = (
|
||||
_output_cost_per_audio_token if _output_cost_per_audio_token is not None else completion_base_cost
|
||||
)
|
||||
|
|
@ -1299,7 +1308,7 @@ def generic_cost_per_token(
|
|||
## REASONING COST
|
||||
if not is_text_tokens_total and reasoning_tokens and reasoning_tokens > 0:
|
||||
completion_cost += float(reasoning_tokens) * _resolve_billed_reasoning_rate(
|
||||
model_info=model_info,
|
||||
model_info=resolved_model_info,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
completion_base_cost=completion_base_cost,
|
||||
|
|
@ -1308,7 +1317,7 @@ def generic_cost_per_token(
|
|||
|
||||
## IMAGE COST
|
||||
if not is_text_tokens_total and image_tokens and image_tokens > 0:
|
||||
_output_cost_per_image_token = _get_cost_per_unit(model_info, "output_cost_per_image_token", None)
|
||||
_output_cost_per_image_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_image_token", None)
|
||||
_output_cost_per_image_token = (
|
||||
_output_cost_per_image_token if _output_cost_per_image_token is not None else completion_base_cost
|
||||
)
|
||||
|
|
@ -1316,7 +1325,7 @@ def generic_cost_per_token(
|
|||
|
||||
## VIDEO COST
|
||||
if not is_text_tokens_total and video_tokens and video_tokens > 0:
|
||||
_output_cost_per_video_token = _get_cost_per_unit(model_info, "output_cost_per_video_token", None)
|
||||
_output_cost_per_video_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_video_token", None)
|
||||
_output_cost_per_video_token = (
|
||||
_output_cost_per_video_token if _output_cost_per_video_token is not None else completion_base_cost
|
||||
)
|
||||
|
|
@ -1325,12 +1334,12 @@ def generic_cost_per_token(
|
|||
## REGIONAL DATA-RESIDENCY UPLIFT
|
||||
# Applied as a flat multiplier across all token costs for the request
|
||||
# when the upstream is a regionalized OpenAI host (eu./us.api.openai.com).
|
||||
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency)
|
||||
uplift: Final = _get_regional_uplift_multiplier(resolved_model_info, data_residency)
|
||||
if uplift != 1.0:
|
||||
prompt_cost *= uplift
|
||||
completion_cost *= uplift
|
||||
|
||||
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
|
||||
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(resolved_model_info, vertex_location)
|
||||
if vertex_uplift != 1.0:
|
||||
prompt_cost *= vertex_uplift
|
||||
completion_cost *= vertex_uplift
|
||||
|
|
@ -1487,7 +1496,10 @@ def get_billed_token_rates(
|
|||
if custom_cost_per_token is not None:
|
||||
return _custom_pricing_rates(custom_cost_per_token)
|
||||
try:
|
||||
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
model_info: Final = apply_provider_cache_read_default(
|
||||
get_model_info(model=model, custom_llm_provider=custom_llm_provider),
|
||||
custom_llm_provider,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # get_model_info raises a bare Exception for an unmapped model: no rates
|
||||
return None
|
||||
return _cost_map_billed_rates(
|
||||
|
|
@ -1578,8 +1590,9 @@ def calculate_prompt_caching_savings(
|
|||
``billed_at`` is the request's completion time, so off-peak windows resolve as the
|
||||
biller saw them rather than at the later spend write.
|
||||
"""
|
||||
model_info_with_cache_read_default: Final = apply_provider_cache_read_default(model_info, custom_llm_provider)
|
||||
prompt_base_cost, _, cache_creation_cost, cache_creation_cost_above_1hr, cache_read_cost = _get_token_base_cost(
|
||||
model_info=model_info,
|
||||
model_info=model_info_with_cache_read_default,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
current_time=billed_at,
|
||||
|
|
|
|||
42
litellm/llms/fireworks_ai/cache_pricing.py
Normal file
42
litellm/llms/fireworks_ai/cache_pricing.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # the derived entry is a dict copy of a ReadOnly TypedDict; no cast-free way to retype it
|
||||
)
|
||||
|
||||
from litellm.constants import FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO
|
||||
from litellm.types.utils import ModelInfo
|
||||
|
||||
|
||||
def _as_rate(value: object) -> float | None:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def with_default_cache_read_rate(model_info: ModelInfo) -> ModelInfo:
|
||||
input_rate: Final = _as_rate(model_info.get("input_cost_per_token"))
|
||||
if model_info.get("cache_read_input_token_cost") is not None or input_rate is None:
|
||||
return model_info
|
||||
cache_read_rate: Final = input_rate * FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO
|
||||
off_peak: Final = model_info.get("off_peak_pricing")
|
||||
if off_peak is None or "cache_read_input_token_cost" in off_peak:
|
||||
return cast(ModelInfo, {**model_info, "cache_read_input_token_cost": cache_read_rate})
|
||||
off_peak_input_rate: Final = _as_rate(off_peak.get("input_cost_per_token"))
|
||||
return cast(
|
||||
ModelInfo,
|
||||
{
|
||||
**model_info,
|
||||
"cache_read_input_token_cost": cache_read_rate,
|
||||
"off_peak_pricing": {
|
||||
**off_peak,
|
||||
"cache_read_input_token_cost": (
|
||||
off_peak_input_rate * FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO
|
||||
if off_peak_input_rate is not None
|
||||
else cache_read_rate
|
||||
),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
|
@ -3,10 +3,7 @@ For calculating cost of fireworks ai serverless inference models.
|
|||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # the fallback entry is a dict copy of a ReadOnly TypedDict; no cast-free way to retype it
|
||||
)
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
FIREWORKS_AI_4_B,
|
||||
|
|
@ -67,28 +64,6 @@ def _resolve_model_info(model: str) -> ModelInfo:
|
|||
return get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
|
||||
|
||||
|
||||
def _with_cache_read_fallback(model_info: ModelInfo) -> ModelInfo:
|
||||
"""Entries without a cache-read rate keep the previous calculator's input-rate fallback for cached
|
||||
reads (LIT-7845 tracks the documented discount); the shared map is never mutated, so a copy carries it."""
|
||||
input_rate: Final = model_info.get("input_cost_per_token")
|
||||
if model_info.get("cache_read_input_token_cost") is not None or input_rate is None:
|
||||
return model_info
|
||||
off_peak: Final = model_info.get("off_peak_pricing")
|
||||
if off_peak is None or "cache_read_input_token_cost" in off_peak:
|
||||
return cast(ModelInfo, {**model_info, "cache_read_input_token_cost": input_rate})
|
||||
return cast(
|
||||
ModelInfo,
|
||||
{
|
||||
**model_info,
|
||||
"cache_read_input_token_cost": input_rate,
|
||||
"off_peak_pricing": {
|
||||
**off_peak,
|
||||
"cache_read_input_token_cost": off_peak.get("input_cost_per_token", input_rate),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def cost_per_token(model: str, usage: Usage, current_time: datetime | None = None) -> tuple[float, float]:
|
||||
"""
|
||||
Calculates the cost per token for a given model, prompt tokens, and completion tokens,
|
||||
|
|
@ -102,7 +77,7 @@ def cost_per_token(model: str, usage: Usage, current_time: datetime | None = Non
|
|||
Returns:
|
||||
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
|
||||
"""
|
||||
model_info: Final = _with_cache_read_fallback(_resolve_model_info(model))
|
||||
model_info: Final = _resolve_model_info(model)
|
||||
return generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
|||
_is_off_peak,
|
||||
_is_within_off_peak_window,
|
||||
apply_off_peak_pricing,
|
||||
apply_provider_cache_read_default,
|
||||
calculate_cache_writing_cost,
|
||||
generic_cost_per_token,
|
||||
get_billed_token_rates,
|
||||
|
|
@ -96,6 +97,19 @@ def test_generic_cost_per_token_bills_cache_reads_at_input_rate_when_no_cache_re
|
|||
assert completion_cost == pytest.approx(380 * 9.7e-7)
|
||||
|
||||
|
||||
def test_apply_provider_cache_read_default_only_derives_a_rate_for_fireworks() -> None:
|
||||
openai_info: ModelInfo = {"input_cost_per_token": 2e-6}
|
||||
fireworks_info: ModelInfo = {"input_cost_per_token": 2e-6}
|
||||
|
||||
assert apply_provider_cache_read_default(openai_info, "openai") is openai_info
|
||||
assert apply_provider_cache_read_default(openai_info, None) is openai_info
|
||||
|
||||
processed_fireworks_info = apply_provider_cache_read_default(fireworks_info, "fireworks_ai")
|
||||
|
||||
assert processed_fireworks_info is not fireworks_info
|
||||
assert processed_fireworks_info["cache_read_input_token_cost"] == pytest.approx(2e-6 * 0.5)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_prefers_audio_per_second_rate() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"key": "gemini-embedding-2",
|
||||
|
|
@ -239,9 +253,7 @@ def test_reasoning_tokens_no_price_set(_local_model_cost_map):
|
|||
model_cost_map["input_cost_per_token"] * usage.prompt_tokens,
|
||||
10,
|
||||
)
|
||||
print(f"completion_cost: {completion_cost}")
|
||||
expected_completion_cost = model_cost_map["output_cost_per_token"] * usage.completion_tokens
|
||||
print(f"expected_completion_cost: {expected_completion_cost}")
|
||||
assert round(completion_cost, 10) == round(
|
||||
expected_completion_cost,
|
||||
10,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,65 @@
|
|||
from copy import deepcopy
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.constants import FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO
|
||||
from litellm.llms.fireworks_ai.cache_pricing import with_default_cache_read_rate
|
||||
from litellm.types.utils import ModelInfo
|
||||
|
||||
|
||||
def test_explicit_cache_read_rate_and_missing_input_rate_keep_the_entry_untouched() -> None:
|
||||
explicit_info: ModelInfo = {"input_cost_per_token": 2e-6, "cache_read_input_token_cost": 1e-6}
|
||||
no_input_rate_info: ModelInfo = {"output_cost_per_token": 3e-6}
|
||||
|
||||
assert with_default_cache_read_rate(explicit_info) is explicit_info
|
||||
assert with_default_cache_read_rate(no_input_rate_info) is no_input_rate_info
|
||||
|
||||
|
||||
def test_missing_cache_read_rate_is_derived_for_standard_and_off_peak_without_mutating_the_entry() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"input_cost_per_token": 2e-6,
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "14:00-00:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 3e-6,
|
||||
},
|
||||
}
|
||||
original: ModelInfo = deepcopy(model_info)
|
||||
|
||||
derived = with_default_cache_read_rate(model_info)
|
||||
|
||||
assert model_info == original
|
||||
assert derived is not model_info
|
||||
assert derived["cache_read_input_token_cost"] == pytest.approx(2e-6 * FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO)
|
||||
assert derived["off_peak_pricing"] == {
|
||||
"hours_utc": "14:00-00:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 3e-6,
|
||||
"cache_read_input_token_cost": 1e-6 * FIREWORKS_AI_DEFAULT_CACHE_READ_RATE_RATIO,
|
||||
}
|
||||
|
||||
|
||||
def test_off_peak_window_without_its_own_input_rate_reuses_the_standard_derived_rate() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"input_cost_per_token": 2e-6,
|
||||
"off_peak_pricing": {"hours_utc": "14:00-00:00", "output_cost_per_token": 3e-6},
|
||||
}
|
||||
|
||||
derived = with_default_cache_read_rate(model_info)
|
||||
|
||||
assert derived["off_peak_pricing"]["cache_read_input_token_cost"] == derived["cache_read_input_token_cost"]
|
||||
|
||||
|
||||
def test_string_rates_from_config_are_coerced_before_the_discount_is_applied() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"input_cost_per_token": "2e-6",
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "14:00-00:00",
|
||||
"input_cost_per_token": "1e-6",
|
||||
},
|
||||
}
|
||||
|
||||
derived = with_default_cache_read_rate(model_info)
|
||||
|
||||
assert derived["cache_read_input_token_cost"] == pytest.approx(1e-6)
|
||||
assert derived["off_peak_pricing"]["cache_read_input_token_cost"] == pytest.approx(5e-7)
|
||||
|
|
@ -5,6 +5,11 @@ from typing import Final
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
calculate_prompt_caching_savings,
|
||||
generic_cost_per_token,
|
||||
get_token_type_cost_breakdown,
|
||||
)
|
||||
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
|
|
@ -48,11 +53,13 @@ STANDARD_CACHE_READ_COST = 1.5e-08
|
|||
|
||||
|
||||
def _register_off_peak_model(
|
||||
off_peak_pricing: OffPeakPricing, cache_read_cost: float | None = STANDARD_CACHE_READ_COST
|
||||
off_peak_pricing: OffPeakPricing,
|
||||
cache_read_cost: float | None = STANDARD_CACHE_READ_COST,
|
||||
model: str = OFF_PEAK_MODEL,
|
||||
) -> None:
|
||||
litellm.model_cost = { # test-quality-ok: the save/restore conftest returns litellm.model_cost to the original object after each test, so replacing the map for this entry leaks nothing
|
||||
litellm.model_cost = { # test-quality-ok: conftest restores litellm.model_cost after each test
|
||||
**litellm.model_cost,
|
||||
f"fireworks_ai/{OFF_PEAK_MODEL}": {
|
||||
f"fireworks_ai/{model}": {
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": STANDARD_INPUT_COST,
|
||||
|
|
@ -103,9 +110,8 @@ def test_off_peak_rates_left_unset_keep_the_standard_rates():
|
|||
assert math.isclose(completion_cost, 200 * STANDARD_OUTPUT_COST, rel_tol=1e-10)
|
||||
|
||||
|
||||
def test_off_peak_window_bills_cached_tokens_at_the_off_peak_input_rate_without_a_cache_read_rate():
|
||||
"""Most fireworks_ai price-map entries carry no cache_read_input_token_cost, so cached tokens
|
||||
fall back to the input rate, and inside the window that has to be the off-peak one."""
|
||||
def test_off_peak_window_bills_cached_tokens_at_the_discounted_off_peak_input_rate_without_a_cache_read_rate():
|
||||
"""Entries without a cache-read rate use Fireworks' documented 50% cached-token discount."""
|
||||
_register_off_peak_model(
|
||||
{"hours_utc": OFF_PEAK_WINDOW, "input_cost_per_token": 1e-08, "output_cost_per_token": 2e-08},
|
||||
cache_read_cost=None,
|
||||
|
|
@ -114,12 +120,116 @@ def test_off_peak_window_bills_cached_tokens_at_the_off_peak_input_rate_without_
|
|||
|
||||
prompt_cost, completion_cost = cost_per_token(model=OFF_PEAK_MODEL, usage=usage, current_time=INSIDE_WINDOW)
|
||||
|
||||
assert math.isclose(prompt_cost, 1000 * 1e-08, rel_tol=1e-10)
|
||||
assert math.isclose(prompt_cost, (700 * 1e-08) + (300 * 1e-08 * 0.5), rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, 200 * 2e-08, rel_tol=1e-10)
|
||||
|
||||
peak_prompt_cost, _ = cost_per_token(model=OFF_PEAK_MODEL, usage=usage, current_time=OUTSIDE_WINDOW)
|
||||
|
||||
assert math.isclose(peak_prompt_cost, 1000 * STANDARD_INPUT_COST, rel_tol=1e-10)
|
||||
assert math.isclose(
|
||||
peak_prompt_cost,
|
||||
(700 * STANDARD_INPUT_COST) + (300 * STANDARD_INPUT_COST * 0.5),
|
||||
rel_tol=1e-10,
|
||||
)
|
||||
|
||||
no_input_rate_model = "accounts/fireworks/models/off-peak-no-input-rate-test"
|
||||
_register_off_peak_model(
|
||||
{"hours_utc": OFF_PEAK_WINDOW, "output_cost_per_token": 2e-08},
|
||||
cache_read_cost=None,
|
||||
model=no_input_rate_model,
|
||||
)
|
||||
|
||||
standard_cache_prompt_cost, _ = cost_per_token(model=no_input_rate_model, usage=usage, current_time=INSIDE_WINDOW)
|
||||
|
||||
assert math.isclose(
|
||||
standard_cache_prompt_cost,
|
||||
(700 * STANDARD_INPUT_COST) + (300 * STANDARD_INPUT_COST * 0.5),
|
||||
rel_tol=1e-10,
|
||||
)
|
||||
|
||||
|
||||
def test_an_entry_without_a_cache_read_rate_bills_cached_tokens_at_the_documented_default_discount():
|
||||
"""Fireworks documents a default 50% cached-token discount for serverless models:
|
||||
https://docs.fireworks.ai/guides/prompt-caching, accessed 2026-09-19."""
|
||||
model = "accounts/fireworks/models/default-cache-read-test"
|
||||
litellm.model_cost = { # test-quality-ok: conftest restores litellm.model_cost after each test
|
||||
**litellm.model_cost,
|
||||
f"fireworks_ai/{model}": {
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": INPUT_COST,
|
||||
"output_cost_per_token": OUTPUT_COST,
|
||||
},
|
||||
}
|
||||
usage = _usage(prompt_tokens=1000, cached_tokens=300, completion_tokens=200)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model=model, usage=usage)
|
||||
|
||||
assert math.isclose(prompt_cost, (700 * INPUT_COST) + (300 * INPUT_COST * 0.5), rel_tol=1e-10)
|
||||
assert prompt_cost < 1000 * INPUT_COST
|
||||
assert math.isclose(completion_cost, 200 * OUTPUT_COST, rel_tol=1e-10)
|
||||
|
||||
|
||||
def test_fireworks_cache_read_rates_match_breakdown_and_caching_savings():
|
||||
model = "accounts/fireworks/models/breakdown-cache-read-test"
|
||||
litellm.model_cost = { # test-quality-ok: the save/restore conftest returns litellm.model_cost to the original object after each test, so replacing the map for this entry leaks nothing
|
||||
**litellm.model_cost,
|
||||
f"fireworks_ai/{model}": {
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": INPUT_COST,
|
||||
"output_cost_per_token": OUTPUT_COST,
|
||||
},
|
||||
}
|
||||
usage = _usage(prompt_tokens=1000, cached_tokens=300, completion_tokens=200)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model=model,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
usage=usage,
|
||||
)
|
||||
prompt_cost, _ = cost_per_token(model=model, usage=usage)
|
||||
savings = calculate_prompt_caching_savings(
|
||||
model_info=litellm.get_model_info(model=model, custom_llm_provider="fireworks_ai"),
|
||||
usage=usage,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
)
|
||||
|
||||
assert math.isclose(breakdown.cache_read_cost, 300 * INPUT_COST * 0.5, rel_tol=1e-10)
|
||||
assert math.isclose(breakdown.rates.cache_read_input_token_cost, INPUT_COST * 0.5, rel_tol=1e-10)
|
||||
assert math.isclose(
|
||||
(700 * breakdown.rates.input_cost_per_token) + breakdown.cache_read_cost, prompt_cost, rel_tol=1e-10
|
||||
)
|
||||
assert math.isclose(savings, 300 * INPUT_COST * 0.5, rel_tol=1e-10)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_applies_fireworks_cache_read_default_with_or_without_model_info():
|
||||
model = "accounts/fireworks/models/generic-cache-read-test"
|
||||
litellm.model_cost = { # test-quality-ok: conftest restores litellm.model_cost after each test
|
||||
**litellm.model_cost,
|
||||
f"fireworks_ai/{model}": {
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat",
|
||||
"input_cost_per_token": INPUT_COST,
|
||||
"output_cost_per_token": OUTPUT_COST,
|
||||
},
|
||||
}
|
||||
usage = _usage(prompt_tokens=1000, cached_tokens=300, completion_tokens=200)
|
||||
expected_prompt_cost = (700 * INPUT_COST) + (300 * INPUT_COST * 0.5)
|
||||
|
||||
implicit_model_info_cost, _ = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
)
|
||||
explicit_model_info_cost, _ = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="fireworks_ai",
|
||||
model_info=litellm.get_model_info(model=model, custom_llm_provider="fireworks_ai"),
|
||||
)
|
||||
|
||||
assert math.isclose(implicit_model_info_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(explicit_model_info_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
|
||||
|
||||
def test_off_peak_defaults_to_the_current_time():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue