diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 238b7cc3fdd..34b7d6eddd8 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -18,6 +18,7 @@ from litellm.constants import ( DEFAULT_MAX_LRU_CACHE_SIZE, DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND, ) +from litellm.litellm_core_utils.cost_discount import resolve_cost_discount from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( StandardBuiltInToolCostTracking, ) @@ -1092,6 +1093,8 @@ def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: ob def _apply_cost_discount( base_cost: float, custom_llm_provider: str | None, + model: str | None = None, + region_name: str | None = None, ) -> tuple[float, float, float]: """ Apply provider-specific cost discount from module-level config. @@ -1099,17 +1102,20 @@ def _apply_cost_discount( Args: base_cost: The base cost before discount custom_llm_provider: The LLM provider name + model: The model name, used to resolve / discount keys + region_name: The region segment the cost key may carry, so patterns can match the region-stripped name Returns: Tuple of (final_cost, discount_percent, discount_amount) """ original_cost: Final = base_cost - discount_percent = 0.0 - discount_amount = 0.0 + resolved_discount: Final = resolve_cost_discount( + litellm.cost_discount_config, custom_llm_provider, model, region_name + ) + discount_percent: Final = 0.0 if resolved_discount is None else resolved_discount + discount_amount: Final = original_cost * discount_percent - if custom_llm_provider and custom_llm_provider in litellm.cost_discount_config: - discount_percent = litellm.cost_discount_config[custom_llm_provider] - discount_amount = original_cost * discount_percent + if resolved_discount is not None: final_cost: Final = original_cost - discount_amount if verbose_logger.isEnabledFor(logging.DEBUG): @@ -1683,6 +1689,8 @@ def completion_cost( ) = _apply_cost_discount( base_cost=_final_cost, custom_llm_provider=custom_llm_provider, + model=model, + region_name=region_name, ) # Apply margin from module-level config if configured @@ -1742,6 +1750,7 @@ def completion_cost( litellm_logging_obj=litellm_logging_obj, total_time=total_time, ) + model_for_discount = model # Calculate cost based on prompt_tokens, completion_tokens if ( "togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai" @@ -1850,6 +1859,8 @@ def completion_cost( ) = _apply_cost_discount( base_cost=_final_cost, custom_llm_provider=custom_llm_provider, + model=model_for_discount, + region_name=region_name, ) else: discount_percent = 0.0 diff --git a/litellm/litellm_core_utils/cost_discount.py b/litellm/litellm_core_utils/cost_discount.py new file mode 100644 index 00000000000..22430e9c11f --- /dev/null +++ b/litellm/litellm_core_utils/cost_discount.py @@ -0,0 +1,97 @@ +import fnmatch +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from typing import Final + +_GLOB_CHARS: Final = frozenset("*?[") + + +@dataclass(frozen=True, slots=True) +class CostDiscountKey: + provider: str + model_pattern: str | None + + +def parse_cost_discount_key(key: str) -> CostDiscountKey: + provider, separator, pattern = key.partition("/") + return CostDiscountKey(provider=provider, model_pattern=pattern if separator else None) + + +def _literal_length(pattern: str) -> int: + length = 0 + index = 0 + while index < len(pattern): + char = pattern[index] + if char in "*?": + index += 1 + continue + if char == "[": + search_start = index + 2 if pattern[index + 1 : index + 2] == "!" else index + 1 + if pattern[search_start : search_start + 1] == "]": + search_start += 1 + closing = pattern.find("]", search_start) + if closing >= 0: + index = closing + 1 + continue + length += 1 + index += 1 + return length + + +def _iter_model_names(model: str, provider_prefix: str) -> Iterator[str]: + if not model.startswith(provider_prefix): + yield model + return + stripped = model + while stripped.startswith(provider_prefix): + stripped = stripped[len(provider_prefix) :] + yield stripped + + +def resolve_cost_discount( + cost_discount_config: Mapping[str, float], + custom_llm_provider: str | None, + model: str | None, + region_name: str | None = None, +) -> float | None: + if not custom_llm_provider: + return None + + provider_prefix: Final = f"{custom_llm_provider}/" + region_prefix: Final = f"{provider_prefix}{region_name}/" if region_name else "" + model_names: Final = ( + () + if model is None + else tuple(_iter_model_names(model, provider_prefix)) + + ( + tuple(_iter_model_names(f"{provider_prefix}{model[len(region_prefix) :]}", provider_prefix)) + if region_prefix and model.startswith(region_prefix) + else () + ) + ) + patterns: Final = tuple( + parsed.model_pattern + for parsed in (parse_cost_discount_key(key) for key in cost_discount_config) + if parsed.provider == custom_llm_provider and parsed.model_pattern is not None + ) + + exact_patterns: Final = frozenset(pattern for pattern in patterns if _GLOB_CHARS.isdisjoint(pattern)) + exact: Final = next((name for name in model_names if name in exact_patterns), None) + if exact is not None: + return cost_discount_config[f"{custom_llm_provider}/{exact}"] + prefixed_markers: Final = (provider_prefix,) if not region_name else (provider_prefix, f"{region_name}/") + matches: Final = tuple( + pattern + for pattern in patterns + if not _GLOB_CHARS.isdisjoint(pattern) + and any( + fnmatch.fnmatchcase(name, pattern) + for name in model_names + if "/" in pattern or not name.startswith(prefixed_markers) + ) + ) + if matches: + best_match: Final = max(matches, key=_literal_length) + return cost_discount_config[f"{custom_llm_provider}/{best_match}"] + + return cost_discount_config.get(custom_llm_provider) diff --git a/litellm/proxy/management_endpoints/cost_tracking_settings.py b/litellm/proxy/management_endpoints/cost_tracking_settings.py index cb376f286ec..82874ac92af 100644 --- a/litellm/proxy/management_endpoints/cost_tracking_settings.py +++ b/litellm/proxy/management_endpoints/cost_tracking_settings.py @@ -21,6 +21,7 @@ import litellm from litellm._internal_context import current_billing_time, pinned_billing_time from litellm._logging import verbose_proxy_logger from litellm.cost_calculator import completion_cost +from litellm.litellm_core_utils.cost_discount import parse_cost_discount_key from litellm.proxy._types import ( CommonProxyErrors, CostEstimateRequest, @@ -238,11 +239,14 @@ async def update_cost_discount_config( Updates the cost_discount_config in litellm_settings. Discounts should be between 0 and 1 (e.g., 0.05 = 5% discount). + A key may also be /, where the pattern is an + fnmatch glob matched against the model name. Example: ```json { "vertex_ai": 0.05, + "vertex_ai/claude-*": 0.20, "gemini": 0.05, "openai": 0.01 } @@ -267,10 +271,16 @@ async def update_cost_discount_config( ) # Validate that all providers are valid LiteLLM providers - invalid_providers: Final = [] - for provider in cost_discount_config: - if provider not in LlmProvidersSet: - invalid_providers.append(provider) + parsed_keys: Final = tuple((key, parse_cost_discount_key(key)) for key in cost_discount_config) + empty_pattern_key: Final = next((key for key, parsed in parsed_keys if parsed.model_pattern == ""), None) + if empty_pattern_key is not None: + raise HTTPException( + status_code=400, + detail=f"Model pattern for {empty_pattern_key} cannot be empty", + ) + invalid_providers: Final = tuple( + parsed.provider for _, parsed in parsed_keys if parsed.provider not in LlmProvidersSet + ) if invalid_providers: raise HTTPException( diff --git a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py index 94d388fce60..cc7818c5803 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py +++ b/tests/test_litellm/proxy/management_endpoints/test_cost_tracking_settings.py @@ -166,6 +166,80 @@ class TestCostTrackingSettings: # Verify litellm.cost_discount_config was updated assert litellm.cost_discount_config == test_discount_config + @pytest.mark.asyncio + async def test_update_cost_discount_config_model_pattern_key(self): + """ + Test PATCH /config/cost_discount_config accepts / keys. + """ + mock_proxy_config = AsyncMock() + mock_proxy_config.get_config = AsyncMock(return_value={"litellm_settings": {}}) + mock_proxy_config.save_config = AsyncMock() + + mock_prisma_client = MagicMock() + + test_discount_config = {"vertex_ai/claude-*": 0.2} + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch.object(litellm, "cost_discount_config", {}), + ): + response = client.patch( + "/config/cost_discount_config", + json=test_discount_config, + headers={"Authorization": "Bearer sk-1234"}, + ) + + assert response.status_code == 200 + assert response.json()["values"]["vertex_ai/claude-*"] == 0.2 + mock_proxy_config.save_config.assert_called_once() + assert litellm.cost_discount_config == test_discount_config + + @pytest.mark.asyncio + async def test_update_cost_discount_config_invalid_provider_with_pattern(self): + """ + Test PATCH /config/cost_discount_config rejects invalid providers in pattern keys. + """ + mock_proxy_config = AsyncMock() + mock_prisma_client = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + response = client.patch( + "/config/cost_discount_config", + json={"notaprovider/claude-*": 0.2}, + headers={"Authorization": "Bearer sk-1234"}, + ) + + assert response.status_code == 400 + assert "notaprovider" in response.json()["detail"]["error"] + + @pytest.mark.asyncio + async def test_update_cost_discount_config_empty_model_pattern(self): + """ + Test PATCH /config/cost_discount_config rejects a key ending in a bare slash. + """ + mock_proxy_config = AsyncMock() + mock_prisma_client = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + response = client.patch( + "/config/cost_discount_config", + json={"vertex_ai/": 0.2}, + headers={"Authorization": "Bearer sk-1234"}, + ) + + assert response.status_code == 400 + assert "cannot be empty" in response.json()["detail"] + @pytest.mark.asyncio async def test_update_cost_discount_config_invalid_provider(self): """ diff --git a/tests/unit/litellm_core_utils/test_cost_discount.py b/tests/unit/litellm_core_utils/test_cost_discount.py new file mode 100644 index 00000000000..8ab071a950d --- /dev/null +++ b/tests/unit/litellm_core_utils/test_cost_discount.py @@ -0,0 +1,175 @@ +import pytest + +from litellm.litellm_core_utils.cost_discount import ( + parse_cost_discount_key, + resolve_cost_discount, +) + + +def test_parse_cost_discount_key_bare_provider(): + parsed = parse_cost_discount_key("vertex_ai") + assert parsed.provider == "vertex_ai" + assert parsed.model_pattern is None + + +def test_parse_cost_discount_key_splits_on_first_slash(): + parsed = parse_cost_discount_key("vertex_ai/claude-*") + assert parsed.provider == "vertex_ai" + assert parsed.model_pattern == "claude-*" + + +def test_parse_cost_discount_key_pattern_containing_slash_stays_in_pattern(): + parsed = parse_cost_discount_key("vertex_ai/a/b") + assert parsed.provider == "vertex_ai" + assert parsed.model_pattern == "a/b" + + +def test_parse_cost_discount_key_empty_pattern_after_slash(): + parsed = parse_cost_discount_key("vertex_ai/") + assert parsed.provider == "vertex_ai" + assert parsed.model_pattern == "" + + +def test_resolve_cost_discount_exact_key_beats_glob_and_bare(): + config = {"vertex_ai/claude-sonnet-4-5": 0.3, "vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", "claude-sonnet-4-5") == 0.3 + + +def test_resolve_cost_discount_glob_beats_bare_provider(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", "claude-sonnet-4-5") == 0.2 + + +def test_resolve_cost_discount_longest_literal_glob_wins(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai/claude-sonnet-*": 0.25, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", "claude-sonnet-4-5") == 0.25 + + +def test_resolve_cost_discount_strips_provider_prefix_from_model(): + config = {"vertex_ai/claude-*": 0.2} + assert resolve_cost_discount(config, "vertex_ai", "vertex_ai/claude-sonnet-4-5") == 0.2 + + +@pytest.mark.parametrize("model", ["openai/openai/gpt-4", "openai/openai/openai/gpt-4"]) +def test_resolve_cost_discount_collapses_repeated_provider_prefix(model): + config = {"openai": 0.05, "openai/gpt-*": 0.20} + assert resolve_cost_discount(config, "openai", model) == 0.20 + + +def test_resolve_cost_discount_prefixed_pattern_beats_stripped_glob(): + config = {"openrouter": 0.05, "openrouter/openrouter/aurora-*": 0.30, "openrouter/aurora-*": 0.20} + assert resolve_cost_discount(config, "openrouter", "openrouter/openrouter/aurora-alpha") == 0.30 + + +def test_resolve_cost_discount_fully_stripped_name_still_matches(): + config = {"openrouter": 0.05, "openrouter/aurora-*": 0.20} + assert resolve_cost_discount(config, "openrouter", "openrouter/openrouter/aurora-alpha") == 0.20 + + +def test_resolve_cost_discount_exact_prefixed_name_beats_stripped_exact(): + config = {"openrouter/openrouter/auto": 0.40, "openrouter/auto": 0.10} + assert resolve_cost_discount(config, "openrouter", "openrouter/openrouter/auto") == 0.40 + assert resolve_cost_discount(config, "openrouter", "openrouter/auto") == 0.10 + + +def test_resolve_cost_discount_longest_literal_across_candidate_names(): + config = {"openai/*": 0.10, "openai/gpt-*": 0.20} + assert resolve_cost_discount(config, "openai", "openai/openai/gpt-4") == 0.20 + + +def test_resolve_cost_discount_glob_skips_prefixed_candidate_names(): + config = {"openai/o*": 0.30, "openai/gpt-*": 0.20} + assert resolve_cost_discount(config, "openai", "openai/openai/gpt-4") == 0.20 + + +def test_resolve_cost_discount_prefixless_glob_only_sees_fully_stripped_name(): + config = {"openai": 0.05, "openai/o*": 0.30} + assert resolve_cost_discount(config, "openai", "openai/openai/gpt-4") == 0.05 + + +def test_resolve_cost_discount_slash_glob_sees_prefixed_candidate_names(): + config = {"openrouter": 0.05, "openrouter/*/auto": 0.25} + assert resolve_cost_discount(config, "openrouter", "openrouter/openrouter/auto") == 0.25 + + +def test_resolve_cost_discount_non_matching_model_falls_back_to_bare(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", "gemini-3-pro-preview") == 0.05 + + +def test_resolve_cost_discount_non_matching_model_no_bare_returns_none(): + config = {"vertex_ai/claude-*": 0.2} + assert resolve_cost_discount(config, "vertex_ai", "gemini-3-pro-preview") is None + + +def test_resolve_cost_discount_other_provider_untouched(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "openai", "claude-sonnet-4-5") is None + + +def test_resolve_cost_discount_none_provider_returns_none(): + config = {"vertex_ai": 0.05} + assert resolve_cost_discount(config, None, "claude-sonnet-4-5") is None + + +def test_resolve_cost_discount_empty_provider_returns_none(): + config = {"vertex_ai": 0.05} + assert resolve_cost_discount(config, "", "claude-sonnet-4-5") is None + + +def test_resolve_cost_discount_none_model_only_bare_matches(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", None) == 0.05 + assert resolve_cost_discount({"vertex_ai/claude-*": 0.2}, "vertex_ai", None) is None + + +def test_resolve_cost_discount_char_class_counts_as_wildcard(): + config = {"vertex_ai/claude-sonnet-*": 0.3, "vertex_ai/claude-[abcdefghijklmnopqrstuvwxyz]*": 0.1} + assert resolve_cost_discount(config, "vertex_ai", "claude-sonnet-4-5") == 0.3 + + +def test_resolve_cost_discount_unclosed_bracket_matches_literally(): + config = {"vertex_ai/weird[": 0.2} + assert resolve_cost_discount(config, "vertex_ai", "weird[") == 0.2 + assert resolve_cost_discount(config, "vertex_ai", "weirdx") is None + + +def test_resolve_cost_discount_char_class_with_leading_bracket(): + config = {"vertex_ai/a[]x]*": 0.1, "vertex_ai/a]*": 0.2} + assert resolve_cost_discount(config, "vertex_ai", "a]q") == 0.2 + assert resolve_cost_discount(config, "vertex_ai", "axq") == 0.1 + + +def test_resolve_cost_discount_glob_crosses_slash_in_model(): + config = {"bedrock/*anthropic.claude-*": 0.15, "bedrock": 0.05} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-east-1/anthropic.claude-v2:1") == 0.15 + + +def test_resolve_cost_discount_exact_pattern_key_does_not_count_as_exact(): + config = {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05} + assert resolve_cost_discount(config, "vertex_ai", "vertex_ai") == 0.05 + + +def test_resolve_cost_discount_region_stripped_name_matches_pattern(): + config = {"bedrock": 0.05, "bedrock/claude-*": 0.2} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-gov-west-1/claude-x", "us-gov-west-1") == 0.2 + + +def test_resolve_cost_discount_region_stripped_exact_beats_glob(): + config = {"bedrock/claude-x": 0.3, "bedrock/claude-*": 0.2} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-gov-west-1/claude-x", "us-gov-west-1") == 0.3 + + +def test_resolve_cost_discount_region_qualified_pattern_wins_longest_literal(): + config = {"bedrock": 0.05, "bedrock/us-gov-west-1/*": 0.4, "bedrock/claude-*": 0.2} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-gov-west-1/claude-x", "us-gov-west-1") == 0.4 + + +def test_resolve_cost_discount_prefixless_glob_cannot_see_region_prefix(): + config = {"bedrock": 0.05, "bedrock/us*": 0.9} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-gov-west-1/claude-x", "us-gov-west-1") == 0.05 + + +def test_resolve_cost_discount_no_region_name_keeps_qualified_name_hidden(): + config = {"bedrock": 0.05, "bedrock/claude-*": 0.2} + assert resolve_cost_discount(config, "bedrock", "bedrock/us-gov-west-1/claude-x") == 0.05 diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 36e188e82d6..0ba13097d64 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -186,7 +186,9 @@ def test_response_cost_calculator_keeps_optional_params_out_of_hidden_params(): assert optional_params["aws_session_token"] == "session-secret" -def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials(monkeypatch: pytest.MonkeyPatch) -> None: +def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials( + monkeypatch: pytest.MonkeyPatch, +) -> None: from litellm.proxy import proxy_server from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload @@ -235,10 +237,6 @@ def test_embedding_success_logging_and_spend_log_carry_no_forwarded_credentials( assert logging_obj.optional_params["extra_headers"] == {"x-goog-api-key": "goog-secret"} - - - - def test_realtime_stream_combines_text_and_audio_token_details(): """Realtime response.done usage with input_token_details / output_token_details.""" from litellm.cost_calculator import RealtimeAPITokenUsageProcessor @@ -1353,8 +1351,6 @@ def test_bedrock_cost_calculator_comparison_with_without_cache(): print(f"Cost with cache: {cost_with_cache}") - - def test_gemini_25_explicit_caching_cost_direct_usage(): """ Test that Gemini 2.5 models correctly calculate costs with explicit caching. @@ -1607,6 +1603,97 @@ def test_cost_discount_vertex_ai(monkeypatch): print(f" - Savings: ${cost_without_discount - cost_with_discount:.6f}") +def test_cost_discount_model_pattern_beats_bare_provider(monkeypatch): + """ + Test that a / discount key applies only to matching models + """ + from litellm import completion_cost + from litellm.types.utils import Usage + + claude_response = ModelResponse( + id="test-id", + choices=[], + created=1234567890, + model="claude-sonnet-4-5", + object="chat.completion", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + gemini_response = ModelResponse( + id="test-id", + choices=[], + created=1234567890, + model="gemini-3-pro-preview", + object="chat.completion", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + + monkeypatch.setattr(litellm, "cost_discount_config", {}) + claude_undiscounted = completion_cost( + completion_response=claude_response, + model="vertex_ai/claude-sonnet-4-5", + custom_llm_provider="vertex_ai", + ) + gemini_undiscounted = completion_cost( + completion_response=gemini_response, + model="vertex_ai/gemini-3-pro-preview", + custom_llm_provider="vertex_ai", + ) + + monkeypatch.setattr(litellm, "cost_discount_config", {"vertex_ai/claude-*": 0.2, "vertex_ai": 0.05}) + claude_discounted = completion_cost( + completion_response=claude_response, + model="vertex_ai/claude-sonnet-4-5", + custom_llm_provider="vertex_ai", + ) + gemini_discounted = completion_cost( + completion_response=gemini_response, + model="vertex_ai/gemini-3-pro-preview", + custom_llm_provider="vertex_ai", + ) + + assert claude_discounted == pytest.approx(claude_undiscounted * 0.8, rel=1e-9) + assert gemini_discounted == pytest.approx(gemini_undiscounted * 0.95, rel=1e-9) + + +def test_cost_discount_together_ai_pattern_matches_unrewritten_model(monkeypatch): + """ + Test that a / discount matches the request model, not the + pricing category the Together AI cost lookup rewrites it to + """ + from litellm import completion_cost + from litellm.types.utils import Usage + + response = ModelResponse( + id="test-id", + choices=[], + created=1234567890, + model="together_ai/my-org/Custom-70B-Instruct", + object="chat.completion", + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + + monkeypatch.setattr(litellm, "cost_discount_config", {}) + undiscounted = completion_cost( + completion_response=response, + model="together_ai/my-org/Custom-70B-Instruct", + custom_llm_provider="together_ai", + ) + assert undiscounted > 0 + + monkeypatch.setattr( + litellm, + "cost_discount_config", + {"together_ai": 0.05, "together_ai/my-org/Custom-*": 0.20}, + ) + discounted = completion_cost( + completion_response=response, + model="together_ai/my-org/Custom-70B-Instruct", + custom_llm_provider="together_ai", + ) + + assert discounted == pytest.approx(undiscounted * 0.8, rel=1e-9) + + def test_cost_discount_not_applied_to_other_providers(monkeypatch): """ Test that cost discount only applies to configured providers @@ -1923,8 +2010,6 @@ def test_cost_margin_with_discount(monkeypatch): print(f" - Expected: ${expected_cost:.6f}") - - def test_completion_cost_extracts_service_tier_from_response(_local_model_cost_map): """Test that completion_cost extracts service_tier from completion_response object.""" from litellm import completion_cost @@ -2674,8 +2759,6 @@ def test_gemini_without_cache_tokens_details(): print("✅ Gemini without cacheTokensDetails works correctly") - - def test_additional_costs_only_for_azure_ai(_local_model_cost_map): """ Test that _get_additional_costs is only called for azure_ai provider. @@ -3220,9 +3303,7 @@ def test_cost_per_token_resolves_per_second_rate_precedence( model: Final = "test-chat-per-second-rate-precedence" entry: Final = {**pricing_fields, "litellm_provider": "together_ai", "mode": "chat"} - litellm.register_model( - model_cost={model: entry} - ) + litellm.register_model(model_cost={model: entry}) assert cost_per_token( model=model, @@ -3774,6 +3855,41 @@ def test_completion_cost_region_name_prices_mantle_on_the_regional_row(_local_mo ) == pytest.approx(expected_flat) +def test_completion_cost_discount_matches_region_stripped_model(monkeypatch, _local_model_cost_map): + """A discount pattern written against the bare model name must still match when the + cost-selected name carries a // prefix from a priced regional row.""" + + response = litellm.ModelResponse( + id="x", + choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + model="xai.grok-4.3", + usage={"prompt_tokens": 38, "completion_tokens": 20, "total_tokens": 58}, + ) + + monkeypatch.setattr(litellm, "cost_discount_config", {}) + undiscounted = litellm.completion_cost( + completion_response=response, + model="xai.grok-4.3", + custom_llm_provider="bedrock_mantle", + region_name="us-gov-west-1", + ) + assert undiscounted > 0 + + monkeypatch.setattr( + litellm, + "cost_discount_config", + {"bedrock_mantle": 0.05, "bedrock_mantle/xai.grok-*": 0.20}, + ) + discounted = litellm.completion_cost( + completion_response=response, + model="xai.grok-4.3", + custom_llm_provider="bedrock_mantle", + region_name="us-gov-west-1", + ) + + assert discounted == pytest.approx(undiscounted * 0.8, rel=1e-9) + + def test_cost_per_token_region_name_applies_to_provider_prefixed_model(_local_model_cost_map): """A provider-prefixed model must still find its bedrock_mantle// row instead of composing the region key with the provider segment twice.""" @@ -4785,9 +4901,7 @@ def test_xai_batch_tier_discounts_the_long_context_rate_like_the_flat_batch_rate assert info[f"{prefix}_above_200k_tokens_batches"] < info[f"{prefix}_above_200k_tokens"] -@pytest.mark.parametrize( - ("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")] -) +@pytest.mark.parametrize(("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")]) def test_xai_batch_cost_calculator_bills_the_200k_batch_tier_inclusively( _local_model_cost_map: None, prompt_tokens: int, tier: str ) -> None: diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.integration.test.tsx index 4680a4d504a..ec19625bfc9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.integration.test.tsx @@ -4,7 +4,6 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../../../tests/test-utils"; import AddProviderForm from "./add_provider_form"; -import { DiscountConfig } from "./types"; const onAddProvider = vi.fn(); const onParentFinish = vi.fn(); @@ -18,11 +17,12 @@ const ParentOwnedForm = () => ( className="space-y-6" > diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx index 94213f923c9..e82c61adcd8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx @@ -4,15 +4,15 @@ import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../tests/test-utils"; import AddProviderForm from "./add_provider_form"; -import { DiscountConfig } from "./types"; import { Providers, providerLogoMap } from "@/components/provider_info_helpers"; const DEFAULT_PROPS = { - discountConfig: {} as DiscountConfig, selectedProvider: undefined, newDiscount: "", + modelPattern: "", onProviderChange: vi.fn(), onDiscountChange: vi.fn(), + onModelPatternChange: vi.fn(), onAddProvider: vi.fn(), }; @@ -71,6 +71,20 @@ describe("AddProviderForm", () => { expect(onDiscountChange).toHaveBeenCalledWith("15"); }); + it("should render the model pattern input field", () => { + renderWithProviders(); + expect(screen.getByPlaceholderText("claude-*")).toBeInTheDocument(); + }); + + it("should report the edited model pattern as the user types", async () => { + const onModelPatternChange = vi.fn(); + const user = userEvent.setup(); + renderWithProviders(); + + await user.type(screen.getByPlaceholderText("claude-*"), "g"); + expect(onModelPatternChange).toHaveBeenCalledWith("g"); + }); + it("should show the percent sign next to the discount input", () => { renderWithProviders(); expect(screen.getByText("%")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx index 05748652180..8f3813e1705 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx @@ -2,7 +2,7 @@ import React from "react"; import { CircleHelp } from "lucide-react"; import { Logo } from "@/components/molecules/logo/Logo"; -import { Providers, provider_map } from "@/components/provider_info_helpers"; +import { Providers } from "@/components/provider_info_helpers"; import { Field, FieldGroup, FieldLabel } from "@/components/ui/field"; import { Button } from "@/components/ui/button"; import { @@ -16,7 +16,6 @@ import { import { Input } from "@/components/ui/input"; import { InputGroupAddon } from "@/components/ui/input-group"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; -import { DiscountConfig } from "./types"; interface ProviderOption { value: string; @@ -24,24 +23,20 @@ interface ProviderOption { } interface AddProviderFormProps { - discountConfig: DiscountConfig; selectedProvider: string | undefined; newDiscount: string; + modelPattern: string; onProviderChange: (provider: string | undefined) => void; onDiscountChange: (discount: string) => void; + onModelPatternChange: (value: string) => void; onAddProvider: () => void; } const PROVIDER_FIELD_ID = "add-provider-discount-provider"; +const MODEL_PATTERN_FIELD_ID = "add-provider-discount-model-pattern"; const DISCOUNT_FIELD_ID = "add-provider-discount-percentage"; -const providerOptionsWithoutDiscount = (discountConfig: DiscountConfig): ProviderOption[] => - Object.entries(Providers) - .filter(([providerEnum]) => { - const providerValue = provider_map[providerEnum as keyof typeof provider_map]; - return !(providerValue && discountConfig[providerValue]); - }) - .map(([value, label]) => ({ value, label })); +const providerOptions = (): ProviderOption[] => Object.entries(Providers).map(([value, label]) => ({ value, label })); const selectedProviderOption = (selectedProvider: string | undefined): ProviderOption | null => { if (!selectedProvider) { @@ -62,14 +57,15 @@ const labelWithHint = (label: string, hint: string): React.ReactNode => ( ); const AddProviderForm: React.FC = ({ - discountConfig, selectedProvider, newDiscount, + modelPattern, onProviderChange, onDiscountChange, + onModelPatternChange, onAddProvider, }) => { - const options = providerOptionsWithoutDiscount(discountConfig); + const options = providerOptions(); const selectedOption = selectedProviderOption(selectedProvider); return ( @@ -110,6 +106,22 @@ const AddProviderForm: React.FC = ({ + + + {labelWithHint( + "Model pattern", + "Leave empty to discount every model from this provider. Use a glob such as claude-* or gemini-2* to discount only matching models", + )} + + onModelPatternChange(event.target.value)} + className="rounded-lg" + /> + + {labelWithHint("Discount Percentage", "Enter a percentage value (e.g., 5 for 5% discount)")} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.integration.test.tsx index 3de88fb6f57..75785d9ab87 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.integration.test.tsx @@ -92,7 +92,21 @@ describe("CostTrackingSettings submit paths", () => { await waitFor(() => expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalled()); expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledTimes(1); - expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledWith("OpenAI", "5"); + expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledWith("OpenAI", "5", ""); + }); + + it("passes the model pattern through to the add handler", async () => { + const user = userEvent.setup(); + await openDiscountModal(user); + + await user.click(screen.getAllByRole("combobox")[0]); + await user.click((await screen.findAllByRole("option"))[0]); + fireEvent.change(screen.getByLabelText(/Model pattern/i), { target: { value: "claude-*" } }); + fireEvent.change(screen.getByLabelText(/Discount Percentage/i), { target: { value: "20" } }); + await user.click(submitDiscount()); + + await waitFor(() => expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalled()); + expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledWith("OpenAI", "20", "claude-*"); }); it("requests the margin exactly once per click", async () => { @@ -127,7 +141,7 @@ describe("CostTrackingSettings submit paths", () => { await waitFor(() => expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalled()); expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledTimes(1); - expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledWith("OpenAI", "5"); + expect(stableDiscountCallbacks.handleAddProvider).toHaveBeenCalledWith("OpenAI", "5", ""); }); it("requests the margin exactly once when Enter is pressed in the percentage field", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx index b8e51939dd0..717e1b53d13 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx @@ -58,6 +58,7 @@ const SectionHeader: React.FC<{ title: string; description: string }> = ({ title const CostTrackingSettings: React.FC = ({ userID, userRole, accessToken }) => { const [selectedProvider, setSelectedProvider] = useState(undefined); const [newDiscount, setNewDiscount] = useState(""); + const [modelPattern, setModelPattern] = useState(""); const [isFetching, setIsFetching] = useState(true); const [isModalVisible, setIsModalVisible] = useState(false); const [isMarginModalVisible, setIsMarginModalVisible] = useState(false); @@ -115,10 +116,11 @@ const CostTrackingSettings: React.FC = ({ userID, use }, [accessToken, fetchDiscountConfig, fetchMarginConfig, fetchBlockUnpriced]); const handleAddProvider = async () => { - const success = await addProvider(selectedProvider, newDiscount); + const success = await addProvider(selectedProvider, newDiscount, modelPattern); if (success) { setSelectedProvider(undefined); setNewDiscount(""); + setModelPattern(""); setIsModalVisible(false); } }; @@ -127,6 +129,7 @@ const CostTrackingSettings: React.FC = ({ userID, use setIsModalVisible(false); setSelectedProvider(undefined); setNewDiscount(""); + setModelPattern(""); }; const handleRemoveProvider = (provider: string, providerDisplayName: string) => { @@ -386,11 +389,12 @@ const CostTrackingSettings: React.FC = ({ userID, use

event.preventDefault()} className="space-y-6"> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx index 8d1ee100a2a..6713351face 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx @@ -295,6 +295,78 @@ describe("ProviderDiscountTable", () => { expect(screen.getByRole("button", { name: "Cancel editing discount for OpenAI" })).toBeInTheDocument(); }); + it("should render a Models column", () => { + renderWithProviders( + , + ); + expect(screen.getByRole("columnheader", { name: "Models" })).toBeInTheDocument(); + }); + + it("should show the model pattern for a provider/pattern key and All models for a bare key", () => { + renderWithProviders( + , + ); + expect(screen.getByText("claude-*")).toBeInTheDocument(); + expect(screen.getByText("All models")).toBeInTheDocument(); + }); + + it("should split a nested pattern key on the first slash only", () => { + renderWithProviders( + , + ); + expect(screen.getByText("anthropic/claude-*")).toBeInTheDocument(); + expect(screen.queryByText("All models")).not.toBeInTheDocument(); + }); + + it("should include the pattern in the row action labels for a pattern key", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await user.click(screen.getByRole("button", { name: /Remove discount for Vertex AI .* \(claude-\*\)/ })); + + expect(onRemoveProvider).toHaveBeenCalledWith( + "vertex_ai/claude-*", + "Vertex AI (Anthropic, Gemini, etc.) (claude-*)", + ); + }); + + it("should call onDiscountChange with the full config key when editing a pattern row", async () => { + const user = userEvent.setup(); + renderWithProviders( + , + ); + + await user.click(screen.getByRole("button", { name: /Edit discount for Vertex AI .* \(claude-\*\)/ })); + const input = screen.getByPlaceholderText("5"); + await user.clear(input); + fireEvent.change(input, { target: { value: "30" } }); + await user.click(screen.getByRole("button", { name: /Save discount for Vertex AI .* \(claude-\*\)/ })); + + expect(onDiscountChange).toHaveBeenCalledWith("vertex_ai/claude-*", "0.3"); + }); + it("should render the empty message when no discounts are configured", () => { renderWithProviders( { + const slashIndex = key.indexOf("/"); + if (slashIndex < 0) { + return { provider: key, modelPattern: null }; + } + return { provider: key.slice(0, slashIndex), modelPattern: key.slice(slashIndex + 1) }; +}; + const ProviderDiscountTable: React.FC = ({ discountConfig, onDiscountChange, onRemoveProvider, }) => { - const [editingProvider, setEditingProvider] = useState(null); + const [editingKey, setEditingKey] = useState(null); const [editValue, setEditValue] = useState(""); - const handleStartEdit = (provider: string, currentDiscount: number) => { - setEditingProvider(provider); + const handleStartEdit = (key: string, currentDiscount: number) => { + setEditingKey(key); setEditValue((currentDiscount * 100).toString()); }; - const handleSaveEdit = (provider: string) => { + const handleSaveEdit = (key: string) => { const percentValue = parseFloat(editValue); if (!isNaN(percentValue) && percentValue >= 0 && percentValue <= 100) { - onDiscountChange(provider, (percentValue / 100).toString()); + onDiscountChange(key, (percentValue / 100).toString()); } - setEditingProvider(null); + setEditingKey(null); setEditValue(""); }; const handleCancelEdit = () => { - setEditingProvider(null); + setEditingKey(null); setEditValue(""); }; - const handleKeyDown = (e: React.KeyboardEvent, provider: string) => { + const handleKeyDown = (e: React.KeyboardEvent, key: string) => { if (e.key === "Enter") { - handleSaveEdit(provider); + handleSaveEdit(key); } else if (e.key === "Escape") { handleCancelEdit(); } }; + const rowLabel = (row: ProviderDiscountRow): string => { + const { displayName } = getProviderLogoAndName(row.provider); + return row.modelPattern ? `${displayName} (${row.modelPattern})` : displayName; + }; + // Convert discount config to array and sort const data: ProviderDiscountRow[] = Object.entries(discountConfig) - .map(([provider, discount]) => ({ provider, discount })) + .map(([key, discount]) => ({ key, ...splitDiscountKey(key), discount })) .sort((a, b) => { const displayA = getProviderLogoAndName(a.provider).displayName; const displayB = getProviderLogoAndName(b.provider).displayName; - return displayA.localeCompare(displayB); + return displayA.localeCompare(displayB) || (a.modelPattern ?? "").localeCompare(b.modelPattern ?? ""); }); return ( @@ -78,19 +93,28 @@ const ProviderDiscountTable: React.FC = ({ ); }, }, + { + header: "Models", + cell: (row) => + row.modelPattern ? ( + {row.modelPattern} + ) : ( + All models + ), + }, { header: "Discount Percentage", numeric: true, cell: (row) => { - const { displayName } = getProviderLogoAndName(row.provider); + const label = rowLabel(row); return (
- {editingProvider === row.provider ? ( + {editingKey === row.key ? ( <> setEditValue(e.target.value)} - onKeyDown={(e) => handleKeyDown(e, row.provider)} + onKeyDown={(e) => handleKeyDown(e, row.key)} placeholder="5" className="w-20" autoFocus @@ -99,8 +123,8 @@ const ProviderDiscountTable: React.FC = ({