diff --git a/ci_cd/generate_model_prices_schema.py b/ci_cd/generate_model_prices_schema.py index 63c9edb5e53..8eec07dadda 100644 --- a/ci_cd/generate_model_prices_schema.py +++ b/ci_cd/generate_model_prices_schema.py @@ -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 diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index cbcb2dfc54d..78fc4c18895 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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), ) diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index edbb1c2699b..eaf7e9767d1 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -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", 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 2af3c9d171b..4b25c87d70f 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 @@ -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