From 288f51bda4a7fe8a90c1f8769295569be53c72bf Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 04:21:30 -0700 Subject: [PATCH] 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. --- litellm/litellm_core_utils/litellm_logging.py | 8 ++- .../litellm_core_utils/llm_cost_calc/utils.py | 66 ++++++++++--------- .../llm_cost_calc/test_llm_cost_calc_utils.py | 27 ++++++++ .../test_litellm_logging.py | 21 ++++++ 4 files changed, 88 insertions(+), 34 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 8a36ce609d3..53d1c93bb01 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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 diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 5eefe88a2e2..cf6caaa469f 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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), ) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index acabc1e1291..726bde1a558 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 4103945a183..829fd107624 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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: