fix(cost): pick each batch price component's tier from its own keys

The batch rate picker crossed one threshold for every component, so a
deployment declaring only an output tier also moved its input, cached, and
cache-write rates to that cutoff. Each component now crosses its own
*_above_<N>k_tokens_batches keys and falls back to its flat key.

The JSON schema is regenerated with the generator as it is on main:
cost-map-guard renders the PR's cost map with the base branch's generator,
so the descriptions for the new batch cache keys move to a follow-up.
This commit is contained in:
mateo-berri 2026-09-19 03:21:14 -07:00
parent 1631d2cef8
commit 9d59fb4071
4 changed files with 55 additions and 36 deletions

View file

@ -222,10 +222,6 @@ COST_DESCRIPTIONS: dict[str, str] = {
"cache_read_input_token_cost": "USD per prompt token served from the provider's prompt cache.",
"input_cost_per_token_batches": "USD per prompt token via the provider's batch API.",
"output_cost_per_token_batches": "USD per generated token via the provider's batch API.",
"cache_read_input_token_cost_batches": "USD per cached prompt token via the provider's batch API.",
"cache_creation_input_token_cost_batches": (
"USD per token written to the provider's prompt cache via its batch API."
),
}
@ -236,8 +232,6 @@ def cost_description(key: str) -> Optional[str]:
return "Flex service-tier rate for the same-named base field."
if key.endswith("_priority"):
return "Priority service-tier rate for the same-named base field."
if "_above_" in key and key.endswith("_batches"):
return "Batch API rate applied once the prompt exceeds the token threshold in the field name."
if "_above_" in key:
return "Rate applied once the prompt exceeds the token threshold in the field name."
return None

View file

@ -74,7 +74,7 @@ _BATCH_RATE_PREFIXES: Final = (
"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}$")
_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)
@ -294,36 +294,42 @@ 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:
def _batch_tier_thresholds(model_info: ModelInfo, prefix: str) -> frozenset[str]:
return frozenset(
tier.group(2)
for key, value in model_info.items()
if value is not None and (tier := _BATCH_TIER_KEY.match(key)) is not None and tier.group(1) == prefix
)
def _crossed_batch_tier(model_info: ModelInfo, prefix: str, usage: Usage, inclusive: bool) -> str | None:
return next(
(
threshold
for threshold in sorted(
_batch_tier_thresholds(model_info, prefix), key=_parse_token_threshold, reverse=True
)
if _prompt_exceeds_threshold(usage.prompt_tokens, _parse_token_threshold(threshold), inclusive)
),
None,
)
def _batch_rate_for_prefix(model_info: ModelInfo, prefix: str, usage: Usage, inclusive: bool) -> float | None:
flat_key: Final = f"{prefix}{_BATCH_KEY_SUFFIX}"
threshold: Final = _crossed_batch_tier(model_info, prefix, usage, inclusive)
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)
crossed_threshold: Final = next(
(
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,
)
return BatchCostRates(
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),
input=_batch_rate_for_prefix(model_info, "input_cost_per_token", usage, inclusive),
output=_batch_rate_for_prefix(model_info, "output_cost_per_token", usage, inclusive),
cache_read=_batch_rate_for_prefix(model_info, "cache_read_input_token_cost", usage, inclusive),
cache_creation=_batch_rate_for_prefix(model_info, "cache_creation_input_token_cost", usage, inclusive),
)

View file

@ -112,7 +112,7 @@
"cache_creation_input_token_cost_above_272k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"cache_creation_input_token_cost_above_272k_tokens_flex": {
"type": "number",
@ -126,8 +126,7 @@
},
"cache_creation_input_token_cost_batches": {
"type": "number",
"minimum": 0,
"description": "USD per token written to the provider's prompt cache via its batch API."
"minimum": 0
},
"cache_creation_input_token_cost_flex": {
"type": "number",
@ -176,7 +175,7 @@
"cache_read_input_token_cost_above_272k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"cache_read_input_token_cost_above_272k_tokens_flex": {
"type": "number",
@ -195,8 +194,7 @@
},
"cache_read_input_token_cost_batches": {
"type": "number",
"minimum": 0,
"description": "USD per cached prompt token via the provider's batch API."
"minimum": 0
},
"cache_read_input_token_cost_flex": {
"type": "number",
@ -353,7 +351,7 @@
"input_cost_per_token_above_272k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"input_cost_per_token_above_272k_tokens_flex": {
"type": "number",
@ -557,7 +555,7 @@
"output_cost_per_token_above_272k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Batch API rate applied once the prompt exceeds the token threshold in the field name."
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"output_cost_per_token_above_272k_tokens_flex": {
"type": "number",

View file

@ -3743,6 +3743,27 @@ def test_get_batch_cost_rates_crosses_a_tier_declared_without_an_input_tier_key(
assert rates.input == 1e-6
@pytest.mark.parametrize(
("prompt_tokens", "expected_input", "expected_output"),
[(200_000, 1e-6, 4e-6), (250_000, 1e-6, 5e-6), (300_000, 2e-6, 5e-6)],
)
def test_get_batch_cost_rates_crosses_each_components_own_tier(prompt_tokens, expected_input, expected_output):
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,
input_cost_per_token_above_272k_tokens_batches=2e-6,
output_cost_per_token_batches=4e-6,
output_cost_per_token_above_200k_tokens_batches=5e-6,
),
Usage(prompt_tokens=prompt_tokens, completion_tokens=1, total_tokens=prompt_tokens + 1),
"openai",
)
assert (rates.input, rates.output) == (expected_input, expected_output)
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