From ed7f85dbcee72c15f9722997aec8608e2c07134f Mon Sep 17 00:00:00 2001
From: yassin
Date: Wed, 30 Sep 2026 18:41:42 +0000
Subject: [PATCH] feat(cost): support provider/model-pattern keys in
cost_discount_config
cost_discount_config keys may now be / with fnmatch globs, so a discount can target a subset of a provider's models. An exact pattern beats a glob, the glob with the longest literal wins, and the bare provider key is the fallback. PATCH /config/cost_discount_config validates the provider half only, and the admin UI cost tracking page gains a Model pattern field and a Models column
Resolves LIT-8465
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/cost_calculator.py | 21 ++-
litellm/litellm_core_utils/cost_discount.py | 97 ++++++++++
.../cost_tracking_settings.py | 18 +-
.../test_cost_tracking_settings.py | 74 ++++++++
.../litellm_core_utils/test_cost_discount.py | 175 ++++++++++++++++++
tests/unit/test_cost_calculator.py | 148 +++++++++++++--
.../add_provider_form.integration.test.tsx | 4 +-
.../_components/add_provider_form.test.tsx | 18 +-
.../_components/add_provider_form.tsx | 36 ++--
...ost_tracking_settings.integration.test.tsx | 18 +-
.../_components/cost_tracking_settings.tsx | 8 +-
.../provider_discount_table.test.tsx | 72 +++++++
.../_components/provider_discount_table.tsx | 70 ++++---
.../_components/use_discount_config.test.ts | 147 ++++++++++++++-
.../_components/use_discount_config.ts | 19 +-
ui/litellm-dashboard/src/lib/http/schema.d.ts | 3 +
16 files changed, 849 insertions(+), 79 deletions(-)
create mode 100644 litellm/litellm_core_utils/cost_discount.py
create mode 100644 tests/unit/litellm_core_utils/test_cost_discount.py
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
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
= ({