fix(cost): select the batch long-context tier from any batch tier key

A deployment that declares its own flat standard input rate keeps every
published batch rate of the output direction, including the 272K output
tier, but the tier was only ever selected when an input tier key was also
present. Detect the crossed tier from any of the four batch tier keys so
the carried output, cache-read, and cache-write tiers bill at their tier
rate above 272K tokens.
This commit is contained in:
mateo-berri 2026-09-05 04:21:30 -07:00
parent d268d8179b
commit 288f51bda4
4 changed files with 88 additions and 34 deletions

View file

@ -373,9 +373,11 @@ def deployment_pricing_model_info(model_id: str | None, deployment_model: str |
may declare only one side of its pricing, so the side it leaves out keeps
the model's published standard rate and every published batch rate for
that direction (flat, long-context tier, cached, cache write) instead of
billing as zero. Ownership is per token direction: declaring either rate
for a direction takes that whole direction, so a published batch rate can
never displace a standard rate the deployment configured itself.
billing as zero. Ownership is per token direction: declaring the flat
standard or flat batch rate for a direction takes that whole direction, so
a published batch rate can never displace a standard rate the deployment
configured itself. A tier-only override keeps every published rate it left
out.
"""
if model_id is None:
return None

View file

@ -62,7 +62,13 @@ _SERVICE_TIER_TO_COST_KEY_SUFFIX: Final[Mapping[str, str]] = MappingProxyType(
_INCLUSIVE_THRESHOLD_PROVIDERS: Final = frozenset({"xai"})
_BATCH_KEY_SUFFIX: Final = "_batches"
_BATCH_TIER_INPUT_KEY: Final = re.compile(r"^input_cost_per_token_above_\d+k?_tokens_batches$")
_BATCH_RATE_PREFIXES: Final = (
"input_cost_per_token",
"output_cost_per_token",
"cache_read_input_token_cost",
"cache_creation_input_token_cost",
)
_BATCH_TIER_KEY: Final = re.compile(rf"^(?:{'|'.join(_BATCH_RATE_PREFIXES)})_above_(\d+k?)_tokens{_BATCH_KEY_SUFFIX}$")
_NON_STANDARD_THRESHOLD_SUFFIXES: Final = (*_SERVICE_TIER_SUFFIXES, _BATCH_KEY_SUFFIX)
@ -239,9 +245,12 @@ def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str:
return f"{base_key}_{suffix}"
def _parse_token_threshold(threshold: str) -> float:
return float(threshold.replace("k", "")) * (1000 if "k" in threshold else 1)
def _parse_above_token_threshold(key: str) -> float:
threshold_str: Final = key.split("_above_")[1].split("_tokens")[0]
return float(threshold_str.replace("k", "")) * (1000 if "k" in threshold_str else 1)
return _parse_token_threshold(key.split("_above_")[1].split("_tokens")[0])
def _prompt_exceeds_threshold(prompt_tokens: int, threshold: float, inclusive: bool) -> bool:
@ -273,41 +282,36 @@ def _batch_tier_rate(model_info: ModelInfo, tier_key: str, flat_key: str) -> flo
return _batch_rate(model_info, flat_key) if tier_rate is None else tier_rate
def _batch_rate_for_threshold(model_info: ModelInfo, prefix: str, threshold: str | None) -> float | None:
flat_key: Final = f"{prefix}{_BATCH_KEY_SUFFIX}"
if threshold is None:
return _batch_rate(model_info, flat_key)
return _batch_tier_rate(model_info, f"{prefix}_above_{threshold}_tokens{_BATCH_KEY_SUFFIX}", flat_key)
def _batch_tier_thresholds(model_info: ModelInfo) -> frozenset[str]:
return frozenset(
tier.group(1)
for key, value in model_info.items()
if value is not None and (tier := _BATCH_TIER_KEY.match(key)) is not None
)
def get_batch_cost_rates(model_info: ModelInfo, usage: Usage, custom_llm_provider: str | None) -> BatchCostRates:
inclusive: Final = _uses_inclusive_token_thresholds(custom_llm_provider)
tier_input_keys: Final = tuple(
key for key, value in model_info.items() if _BATCH_TIER_INPUT_KEY.match(key) and value is not None
)
crossed_input_key: Final = next(
crossed_threshold: Final = next(
(
key
for key in sorted(tier_input_keys, key=_parse_above_token_threshold, reverse=True)
if _prompt_exceeds_threshold(usage.prompt_tokens, _parse_above_token_threshold(key), inclusive)
threshold
for threshold in sorted(_batch_tier_thresholds(model_info), key=_parse_token_threshold, reverse=True)
if _prompt_exceeds_threshold(usage.prompt_tokens, _parse_token_threshold(threshold), inclusive)
),
None,
)
if crossed_input_key is None:
return BatchCostRates(
input=_batch_rate(model_info, "input_cost_per_token_batches"),
output=_batch_rate(model_info, "output_cost_per_token_batches"),
cache_read=_batch_rate(model_info, "cache_read_input_token_cost_batches"),
cache_creation=_batch_rate(model_info, "cache_creation_input_token_cost_batches"),
)
return BatchCostRates(
input=_batch_tier_rate(model_info, crossed_input_key, "input_cost_per_token_batches"),
output=_batch_tier_rate(
model_info, crossed_input_key.replace("input_", "output_", 1), "output_cost_per_token_batches"
),
cache_read=_batch_tier_rate(
model_info,
crossed_input_key.replace("input_cost_per_token", "cache_read_input_token_cost", 1),
"cache_read_input_token_cost_batches",
),
cache_creation=_batch_tier_rate(
model_info,
crossed_input_key.replace("input_cost_per_token", "cache_creation_input_token_cost", 1),
"cache_creation_input_token_cost_batches",
),
input=_batch_rate_for_threshold(model_info, "input_cost_per_token", crossed_threshold),
output=_batch_rate_for_threshold(model_info, "output_cost_per_token", crossed_threshold),
cache_read=_batch_rate_for_threshold(model_info, "cache_read_input_token_cost", crossed_threshold),
cache_creation=_batch_rate_for_threshold(model_info, "cache_creation_input_token_cost", crossed_threshold),
)

View file

@ -4848,6 +4848,33 @@ def test_get_batch_cost_rates_reads_the_batch_cache_write_rate_for_the_crossed_t
assert rates.cache_creation == expected
@pytest.mark.parametrize(
("tier_key", "attribute"),
[
("output_cost_per_token_above_272k_tokens_batches", "output"),
("cache_read_input_token_cost_above_272k_tokens_batches", "cache_read"),
("cache_creation_input_token_cost_above_272k_tokens_batches", "cache_creation"),
],
)
def test_get_batch_cost_rates_crosses_a_tier_declared_without_an_input_tier_key(tier_key, attribute):
from litellm.litellm_core_utils.llm_cost_calc.utils import get_batch_cost_rates
rates = get_batch_cost_rates(
_batch_rates_model_info(
input_cost_per_token_batches=1e-6,
output_cost_per_token_batches=4e-6,
cache_read_input_token_cost_batches=1e-7,
cache_creation_input_token_cost_batches=1.25e-7,
**{tier_key: 9e-6},
),
Usage(prompt_tokens=300_000, completion_tokens=1, total_tokens=300_001),
"openai",
)
assert getattr(rates, attribute) == 9e-6
assert rates.input == 1e-6
def test_get_batch_cost_rates_has_no_cache_write_rate_without_a_cache_write_batch_key():
from litellm.litellm_core_utils.llm_cost_calc.utils import get_batch_cost_rates

View file

@ -6337,6 +6337,27 @@ def test_deployment_pricing_model_info_carries_the_published_output_batch_tier_w
assert info["cache_creation_input_token_cost_batches"] is None
def test_batch_cost_calculator_bills_the_carried_output_tier_when_the_deployment_declares_its_own_input_rate(
_local_model_cost_map: None,
) -> None:
from litellm.cost_calculator import batch_cost_calculator
from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info
from litellm.types.utils import Usage
info = deployment_pricing_model_info(_luna_deployment_id({"input_cost_per_token": 2e-6}), "openai/gpt-5.6-luna")
assert info is not None
prompt_cost, completion_cost = batch_cost_calculator(
usage=Usage(prompt_tokens=300_000, completion_tokens=10, total_tokens=300_010),
model="openai/gpt-5.6-luna",
custom_llm_provider="openai",
model_info=info,
)
assert prompt_cost == pytest.approx(300_000 * 1e-6)
assert completion_cost == pytest.approx(10 * 9e-7)
def test_deployment_pricing_model_info_honors_a_tier_only_batch_override_over_the_published_flat_rates(
_local_model_cost_map: None,
) -> None: