This commit is contained in:
devin-ai-integration[bot] 2026-09-30 18:41:48 +00:00 • committed by GitHub
commit 05d7719fba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 849 additions and 79 deletions

View file

@ -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 <provider>/<model-pattern> 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

View file

@ -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)

View file

@ -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 <provider>/<model-pattern>, 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(

View file

@ -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 <provider>/<model-pattern> 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):
"""

View file

@ -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

View file

@ -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 <provider>/<model-pattern> 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 <provider>/<model-pattern> 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 <provider>/<region>/ 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/<region>/<model> 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:

View file

@ -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"
>
<AddProviderForm
discountConfig={{} as DiscountConfig}
selectedProvider="OpenAI"
newDiscount="5"
modelPattern=""
onProviderChange={vi.fn()}
onDiscountChange={vi.fn()}
onModelPatternChange={vi.fn()}
onAddProvider={onAddProvider}
/>
</form>

View file

@ -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(<AddProviderForm {...DEFAULT_PROPS} />);
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(<AddProviderForm {...DEFAULT_PROPS} onModelPatternChange={onModelPatternChange} />);
await user.type(screen.getByPlaceholderText("claude-*"), "g");
expect(onModelPatternChange).toHaveBeenCalledWith("g");
});
it("should show the percent sign next to the discount input", () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} />);
expect(screen.getByText("%")).toBeInTheDocument();

View file

@ -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<AddProviderFormProps> = ({
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<AddProviderFormProps> = ({
</Combobox>
</Field>
<Field>
<FieldLabel htmlFor={MODEL_PATTERN_FIELD_ID}>
{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",
)}
</FieldLabel>
<Input
id={MODEL_PATTERN_FIELD_ID}
placeholder="claude-*"
value={modelPattern}
onChange={(event) => onModelPatternChange(event.target.value)}
className="rounded-lg"
/>
</Field>
<Field>
<FieldLabel htmlFor={DISCOUNT_FIELD_ID}>
{labelWithHint("Discount Percentage", "Enter a percentage value (e.g., 5 for 5% discount)")}

View file

@ -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 () => {

View file

@ -58,6 +58,7 @@ const SectionHeader: React.FC<{ title: string; description: string }> = ({ title
const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, userRole, accessToken }) => {
const [selectedProvider, setSelectedProvider] = useState<string | undefined>(undefined);
const [newDiscount, setNewDiscount] = useState<string>("");
const [modelPattern, setModelPattern] = useState<string>("");
const [isFetching, setIsFetching] = useState(true);
const [isModalVisible, setIsModalVisible] = useState(false);
const [isMarginModalVisible, setIsMarginModalVisible] = useState(false);
@ -115,10 +116,11 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ 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<CostTrackingSettingsProps> = ({ userID, use
setIsModalVisible(false);
setSelectedProvider(undefined);
setNewDiscount("");
setModelPattern("");
};
const handleRemoveProvider = (provider: string, providerDisplayName: string) => {
@ -386,11 +389,12 @@ const CostTrackingSettings: React.FC<CostTrackingSettingsProps> = ({ userID, use
</p>
<form onSubmit={(event) => event.preventDefault()} className="space-y-6">
<AddProviderForm
discountConfig={discountConfig}
selectedProvider={selectedProvider}
newDiscount={newDiscount}
modelPattern={modelPattern}
onProviderChange={setSelectedProvider}
onDiscountChange={setNewDiscount}
onModelPatternChange={setModelPattern}
onAddProvider={handleAddProvider}
/>
</form>

View file

@ -295,6 +295,78 @@ describe("ProviderDiscountTable", () => {
expect(screen.getByRole("button", { name: "Cancel editing discount for OpenAI" })).toBeInTheDocument();
});
it("should render a Models column", () => {
renderWithProviders(
<ProviderDiscountTable
discountConfig={DEFAULT_DISCOUNT_CONFIG}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>,
);
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(
<ProviderDiscountTable
discountConfig={{ "vertex_ai/claude-*": 0.2, vertex_ai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>,
);
expect(screen.getByText("claude-*")).toBeInTheDocument();
expect(screen.getByText("All models")).toBeInTheDocument();
});
it("should split a nested pattern key on the first slash only", () => {
renderWithProviders(
<ProviderDiscountTable
discountConfig={{ "openrouter/anthropic/claude-*": 0.15 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>,
);
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(
<ProviderDiscountTable
discountConfig={{ "vertex_ai/claude-*": 0.2 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>,
);
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(
<ProviderDiscountTable
discountConfig={{ "vertex_ai/claude-*": 0.2 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>,
);
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(
<ProviderDiscountTable

View file

@ -14,52 +14,67 @@ interface ProviderDiscountTableProps {
}
interface ProviderDiscountRow {
key: string;
provider: string;
modelPattern: string | null;
discount: number;
}
const splitDiscountKey = (key: string): { provider: string; modelPattern: string | null } => {
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<ProviderDiscountTableProps> = ({
discountConfig,
onDiscountChange,
onRemoveProvider,
}) => {
const [editingProvider, setEditingProvider] = useState<string | null>(null);
const [editingKey, setEditingKey] = useState<string | null>(null);
const [editValue, setEditValue] = useState<string>("");
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<ProviderDiscountTableProps> = ({
);
},
},
{
header: "Models",
cell: (row) =>
row.modelPattern ? (
<span className="font-mono text-sm">{row.modelPattern}</span>
) : (
<span className="text-muted-foreground">All models</span>
),
},
{
header: "Discount Percentage",
numeric: true,
cell: (row) => {
const { displayName } = getProviderLogoAndName(row.provider);
const label = rowLabel(row);
return (
<div className="flex items-center justify-end gap-2">
{editingProvider === row.provider ? (
{editingKey === row.key ? (
<>
<Input
value={editValue}
onChange={(e) => 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<ProviderDiscountTableProps> = ({
<Button
variant="ghost"
size="icon-sm"
aria-label={`Save discount for ${displayName}`}
onClick={() => handleSaveEdit(row.provider)}
aria-label={`Save discount for ${label}`}
onClick={() => handleSaveEdit(row.key)}
className="cursor-pointer text-success hover:text-success/80"
>
<Check className="size-5" />
@ -108,7 +132,7 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
<Button
variant="ghost"
size="icon-sm"
aria-label={`Cancel editing discount for ${displayName}`}
aria-label={`Cancel editing discount for ${label}`}
onClick={handleCancelEdit}
className="cursor-pointer text-muted-foreground hover:text-foreground"
>
@ -121,8 +145,8 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
<Button
variant="ghost"
size="icon-sm"
aria-label={`Edit discount for ${displayName}`}
onClick={() => handleStartEdit(row.provider, row.discount)}
aria-label={`Edit discount for ${label}`}
onClick={() => handleStartEdit(row.key, row.discount)}
className="cursor-pointer text-info hover:text-info/80"
>
<SquarePen className="size-5" />
@ -137,13 +161,13 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
{
header: "Actions",
cell: (row) => {
const { displayName } = getProviderLogoAndName(row.provider);
const label = rowLabel(row);
return (
<Button
variant="ghost"
size="icon-sm"
aria-label={`Remove discount for ${displayName}`}
onClick={() => onRemoveProvider(row.provider, displayName)}
aria-label={`Remove discount for ${label}`}
onClick={() => onRemoveProvider(row.key, label)}
className="cursor-pointer hover:text-destructive"
>
<Trash2 className="size-5" />
@ -153,7 +177,7 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
width: "80px",
},
]}
getRowKey={(row) => row.provider}
getRowKey={(row) => row.key}
emptyMessage="No provider discounts configured"
/>
);

View file

@ -13,6 +13,8 @@ vi.mock("./provider_display_helpers", () => ({
const map: Record<string, string> = {
OpenAI: "openai",
Anthropic: "anthropic",
Vertex_AI: "vertex_ai",
Openrouter: "openrouter",
};
return map[enumKey] ?? null;
}),
@ -22,6 +24,8 @@ vi.mock("@/components/provider_info_helpers", () => ({
Providers: {
OpenAI: "OpenAI",
Anthropic: "Anthropic",
Vertex_AI: "Vertex AI",
Openrouter: "OpenRouter",
},
}));
@ -80,7 +84,7 @@ describe("useDiscountConfig", () => {
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider(undefined, "5");
success = await result.current.handleAddProvider(undefined, "5", "");
});
expect(success!).toBe(false);
@ -92,7 +96,7 @@ describe("useDiscountConfig", () => {
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "");
success = await result.current.handleAddProvider("OpenAI", "", "");
});
expect(success!).toBe(false);
@ -104,7 +108,7 @@ describe("useDiscountConfig", () => {
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "150");
success = await result.current.handleAddProvider("OpenAI", "150", "");
});
expect(success!).toBe(false);
@ -127,13 +131,37 @@ describe("useDiscountConfig", () => {
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "10");
success = await result.current.handleAddProvider("OpenAI", "10", "");
});
expect(success!).toBe(false);
expect(toast.fromError).toHaveBeenCalledWith(expect.stringMatching(/already exists/i));
});
it("should reject a duplicate even when the stored discount is 0", async () => {
vi.spyOn(global, "fetch").mockResolvedValueOnce({
ok: true,
json: async () => ({ values: { vertex_ai: 0 } }),
} as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
await act(async () => {
await result.current.fetchDiscountConfig();
});
vi.clearAllMocks();
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("Vertex_AI", "5", "");
});
expect(success!).toBe(false);
expect(toast.fromError).toHaveBeenCalledWith(expect.stringMatching(/already exists/i));
expect(global.fetch).not.toHaveBeenCalled();
});
it("should save the config and return true on a valid new provider", async () => {
vi.spyOn(global, "fetch")
.mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response)
@ -148,12 +176,121 @@ describe("useDiscountConfig", () => {
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "5");
success = await result.current.handleAddProvider("OpenAI", "5", "");
});
expect(success!).toBe(true);
expect(toast.success).toHaveBeenCalledWith("Discount configuration updated successfully");
});
it("should store a <provider>/<pattern> key when a model pattern is given", async () => {
vi.spyOn(global, "fetch")
.mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response)
.mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response)
.mockResolvedValueOnce({ ok: true, json: async () => ({ values: { "openai/gpt-*": 0.2 } }) } as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
await act(async () => {
await result.current.fetchDiscountConfig();
});
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "20", "gpt-*");
});
expect(success!).toBe(true);
const patchCall = vi.mocked(global.fetch).mock.calls.find(([, init]) => init?.method === "PATCH");
expect(patchCall).toBeDefined();
expect(JSON.parse(patchCall![1]!.body as string)).toEqual({ "openai/gpt-*": 0.2 });
expect(result.current.discountConfig).toHaveProperty("openai/gpt-*", 0.2);
});
it("should store the bare provider key when the pattern is empty or whitespace", async () => {
vi.spyOn(global, "fetch")
.mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response)
.mockResolvedValueOnce({ ok: true, json: async () => ({ values: { openai: 0.05 } }) } as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "5", " ");
});
expect(success!).toBe(true);
const patchCall = vi.mocked(global.fetch).mock.calls.find(([, init]) => init?.method === "PATCH");
expect(JSON.parse(patchCall![1]!.body as string)).toEqual({ openai: 0.05 });
});
it("should compose a nested pattern with slashes into the config key", async () => {
vi.spyOn(global, "fetch")
.mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response)
.mockResolvedValueOnce({
ok: true,
json: async () => ({ values: { "openrouter/anthropic/claude-*": 0.15 } }),
} as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("Openrouter", "15", "anthropic/claude-*");
});
expect(success!).toBe(true);
const patchCall = vi.mocked(global.fetch).mock.calls.find(([, init]) => init?.method === "PATCH");
expect(JSON.parse(patchCall![1]!.body as string)).toEqual({ "openrouter/anthropic/claude-*": 0.15 });
});
it("should let a provider with a bare discount also get a pattern entry", async () => {
vi.spyOn(global, "fetch")
.mockResolvedValueOnce({ ok: true, json: async () => ({ values: { openai: 0.05 } }) } as Response)
.mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response)
.mockResolvedValueOnce({
ok: true,
json: async () => ({ values: { openai: 0.05, "openai/gpt-*": 0.2 } }),
} as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
await act(async () => {
await result.current.fetchDiscountConfig();
});
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "20", "gpt-*");
});
expect(success!).toBe(true);
const patchCall = vi.mocked(global.fetch).mock.calls.find(([, init]) => init?.method === "PATCH");
expect(JSON.parse(patchCall![1]!.body as string)).toEqual({ openai: 0.05, "openai/gpt-*": 0.2 });
});
it("should reject a duplicate pattern key", async () => {
vi.spyOn(global, "fetch").mockResolvedValueOnce({
ok: true,
json: async () => ({ values: { "openai/gpt-*": 0.2 } }),
} as Response);
const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" }));
await act(async () => {
await result.current.fetchDiscountConfig();
});
let success: boolean;
await act(async () => {
success = await result.current.handleAddProvider("OpenAI", "10", "gpt-*");
});
expect(success!).toBe(false);
expect(toast.fromError).toHaveBeenCalledWith(
"Discount for OpenAI (gpt-*) already exists. Edit it in the table above.",
);
});
});
describe("handleRemoveProvider", () => {

View file

@ -14,7 +14,11 @@ export interface UseDiscountConfigReturn {
setDiscountConfig: React.Dispatch<React.SetStateAction<DiscountConfig>>;
fetchDiscountConfig: () => Promise<void>;
saveDiscountConfig: (config: DiscountConfig) => Promise<void>;
handleAddProvider: (selectedProvider: string | undefined, newDiscount: string) => Promise<boolean>;
handleAddProvider: (
selectedProvider: string | undefined,
newDiscount: string,
modelPattern: string,
) => Promise<boolean>;
handleRemoveProvider: (provider: string) => Promise<void>;
handleDiscountChange: (provider: string, value: string) => Promise<void>;
}
@ -79,7 +83,7 @@ export function useDiscountConfig({ accessToken }: UseDiscountConfigProps): UseD
);
const handleAddProvider = useCallback(
async (selectedProvider: string | undefined, newDiscount: string): Promise<boolean> => {
async (selectedProvider: string | undefined, newDiscount: string, modelPattern: string): Promise<boolean> => {
if (!selectedProvider || !newDiscount) {
toast.fromError("Please select a provider and enter discount percentage");
return false;
@ -91,6 +95,8 @@ export function useDiscountConfig({ accessToken }: UseDiscountConfigProps): UseD
return false;
}
const trimmedPattern = modelPattern.trim();
const providerValue = getProviderBackendValue(selectedProvider);
if (!providerValue) {
@ -98,9 +104,12 @@ export function useDiscountConfig({ accessToken }: UseDiscountConfigProps): UseD
return false;
}
if (discountConfig[providerValue]) {
const configKey = trimmedPattern ? `${providerValue}/${trimmedPattern}` : providerValue;
if (configKey in discountConfig) {
const displayName = Providers[selectedProvider as keyof typeof Providers];
toast.fromError(
`Discount for ${Providers[selectedProvider as keyof typeof Providers]} already exists. Edit it in the table above.`,
`Discount for ${displayName} (${trimmedPattern || "all models"}) already exists. Edit it in the table above.`,
);
return false;
}
@ -108,7 +117,7 @@ export function useDiscountConfig({ accessToken }: UseDiscountConfigProps): UseD
const discountValue = percentageValue / 100;
const updatedConfig = {
...discountConfig,
[providerValue]: discountValue,
[configKey]: discountValue,
};
setDiscountConfig(updatedConfig);

View file

@ -2880,11 +2880,14 @@ export interface paths {
*
* 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 <provider>/<model-pattern>, 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
* }