diff --git a/scripts/sync_cost_map.py b/scripts/sync_cost_map.py index cb24bc36f5d..e2eb7a0ee26 100644 --- a/scripts/sync_cost_map.py +++ b/scripts/sync_cost_map.py @@ -240,38 +240,55 @@ def _tier_threshold(boundary: int) -> int | None: return next((start // 1000 for start in (boundary, boundary - 1) if start > 0 and start % 1000 == 0), None) -def _tiered(name: str, base: float | None, tiers: Sequence[VercelTier] | None) -> Prices | Unmappable: +@dataclass(frozen=True, slots=True) +class TierLadder: + name: str + base: float + tiers: tuple[VercelTier, ...] + thresholds: frozenset[int] + + def price_above(self, thousands: int) -> float | None: + tokens: Final = thousands * 1000 + 1 + tier: Final = next( + (tier for tier in self.tiers if (tier.min or 0) <= tokens and (tier.max is None or tokens < tier.max)), None + ) + return _token_price(tier.cost) if tier is not None else self.base + + +def _ladder(name: str, base: float | None, tiers: Sequence[VercelTier] | None) -> TierLadder | Unmappable | None: if base is None or not tiers: - return NO_PRICES - ordered: Final = sorted(tiers, key=lambda tier: tier.min or 0) + return None + ordered: Final = tuple(sorted(tiers, key=lambda tier: tier.min or 0)) contiguous: Final = ordered[-1].max is None and all( lower.max == upper.min for lower, upper in zip(ordered, ordered[1:], strict=False) ) if not contiguous: return Unmappable(f"{name} tiers are not contiguous") - steps: Final = tuple((_tier_threshold(tier.min), _token_price(tier.cost)) for tier in ordered if tier.min) - prices: Final = { - f"{name}_above_{thousands}k_tokens": price - for thousands, price in steps - if thousands is not None and price is not None - } - if len(prices) != len(steps): + thresholds: Final = tuple(_tier_threshold(tier.min) for tier in ordered if tier.min) + if None in thresholds or any(_token_price(tier.cost) is None for tier in ordered): return Unmappable(f"{name} tiers have a boundary that is not a whole thousand or an unusable price") - return MappingProxyType(prices) + return TierLadder(name, base, ordered, frozenset(threshold for threshold in thresholds if threshold is not None)) def _vercel_tiers(pricing: VercelPricing, cache_read: float | None, cache_write: float | None) -> Prices | Unmappable: parts: Final = ( - _tiered("input_cost_per_token", _token_price(pricing.input), pricing.input_tiers), - _tiered("output_cost_per_token", _token_price(pricing.output), pricing.output_tiers), - _tiered("cache_read_input_token_cost", cache_read, pricing.input_cache_read_tiers), - _tiered("cache_creation_input_token_cost", cache_write, pricing.input_cache_write_tiers), + _ladder("input_cost_per_token", _token_price(pricing.input), pricing.input_tiers), + _ladder("output_cost_per_token", _token_price(pricing.output), pricing.output_tiers), + _ladder("cache_read_input_token_cost", cache_read, pricing.input_cache_read_tiers), + _ladder("cache_creation_input_token_cost", cache_write, pricing.input_cache_write_tiers), ) problem: Final = next((part for part in parts if isinstance(part, Unmappable)), None) if problem is not None: return problem + ladders: Final = tuple(part for part in parts if isinstance(part, TierLadder)) + thresholds: Final = sorted(frozenset().union(*(ladder.thresholds for ladder in ladders))) return MappingProxyType( - {name: price for part in parts if not isinstance(part, Unmappable) for name, price in part.items()} + { + f"{ladder.name}_above_{thousands}k_tokens": price + for ladder in ladders + for thousands in thresholds + if (price := ladder.price_above(thousands)) is not None + } ) diff --git a/tests/test_litellm/test_sync_cost_map.py b/tests/test_litellm/test_sync_cost_map.py index a9fcb22655d..c622d30caea 100644 --- a/tests/test_litellm/test_sync_cost_map.py +++ b/tests/test_litellm/test_sync_cost_map.py @@ -2,10 +2,13 @@ import importlib.util import json from pathlib import Path from types import ModuleType -from typing import Final +from typing import Final, cast import pytest +from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token +from litellm.types.utils import ModelInfo, Usage + REPO_ROOT: Final = Path(__file__).resolve().parents[2] SCRIPT_PATH: Final = REPO_ROOT / "scripts" / "sync_cost_map.py" FIXTURES: Final = Path(__file__).parent / "fixtures" / "cost_map_sync" @@ -595,6 +598,23 @@ def test_vercel_long_context_tiers_map_to_above_threshold_prices(sync: ModuleTyp assert entry["cache_read_input_token_cost_above_256k_tokens"] == 6e-7 assert entry["cache_creation_input_token_cost_above_200k_tokens"] == 4e-6 assert not any(key.endswith("_above_0k_tokens") for key in entry) + assert entry["input_cost_per_token_above_200k_tokens"] == 4.5e-6 + assert entry["output_cost_per_token_above_128k_tokens"] == 7.5e-6 + assert entry["cache_read_input_token_cost_above_200k_tokens"] == 3e-7 + billed: Final = { + prompt_tokens: generic_cost_per_token( + model="acme/long", + usage=Usage(prompt_tokens=prompt_tokens, completion_tokens=1000, total_tokens=prompt_tokens + 1000), + custom_llm_provider="vercel_ai_gateway", + model_info=cast(ModelInfo, dict(entry)), + ) + for prompt_tokens in (100_000, 250_000, 300_000) + } + assert billed == { + 100_000: (pytest.approx(0.27), pytest.approx(0.0075)), + 250_000: (pytest.approx(1.125), pytest.approx(0.01125)), + 300_000: (pytest.approx(1.35), pytest.approx(0.01125)), + } @pytest.mark.parametrize(