mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(cost-map-sync): write every Vercel tier list at the row's shared breakpoints
The cost calculator walks the input_cost_per_token_above_* thresholds and reads the output and cache prices at the same threshold, so an output or cache tier that broke at a breakpoint no input tier had was stored and never billed. Each list is now written at the union of the row's breakpoints, priced from the tier that covers that breakpoint. Today's 41 tiered catalog rows are aligned, so the synced map is byte-identical; the regression test bills a mismatched row through generic_cost_per_token
This commit is contained in:
parent
9478a32d22
commit
26f74e662d
2 changed files with 54 additions and 17 deletions
|
|
@ -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
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue