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: