mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge ed7f85dbce into 3930c5bab6
This commit is contained in:
commit
05d7719fba
16 changed files with 849 additions and 79 deletions
|
|
@ -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
|
||||
|
|
|
|||
97
litellm/litellm_core_utils/cost_discount.py
Normal file
97
litellm/litellm_core_utils/cost_discount.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
175
tests/unit/litellm_core_utils/test_cost_discount.py
Normal file
175
tests/unit/litellm_core_utils/test_cost_discount.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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)")}
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
/>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
3
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
3
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
* }
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue