feat(cost): support deployment-level cost discounts

This commit is contained in:
silencedoctor 2026-07-29 18:13:02 +08:00
parent f005afa146
commit b748ba7ca3
8 changed files with 253 additions and 18 deletions

View file

@ -157,6 +157,7 @@ _SPEECH_CALL_TYPES: Final = frozenset(
}
)
_COST_DISCOUNT_FIELD: Final = "cost_discount"
_TRANSCRIPTION_CALL_TYPES: Final = frozenset(
{
CallTypes.atranscription.value,
@ -960,13 +961,17 @@ def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: An
def _apply_cost_discount(
base_cost: float,
custom_llm_provider: str | None,
model: str | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, float, float]:
"""
Apply provider-specific cost discount from module-level config.
Apply model/deployment-specific cost discount or provider-specific config.
Args:
base_cost: The base cost before discount
custom_llm_provider: The LLM provider name
model: Model cost-map key used for cost calculation
model_info: Optional deployment/model pricing metadata
Returns:
Tuple of (final_cost, discount_percent, discount_amount)
@ -975,14 +980,21 @@ def _apply_cost_discount(
discount_percent = 0.0
discount_amount = 0.0
if custom_llm_provider and custom_llm_provider in litellm.cost_discount_config:
model_cost_discount = _get_model_cost_discount(model=model, model_info=model_info)
discount_source = custom_llm_provider
if model_cost_discount is not None:
discount_percent = model_cost_discount
discount_source = _COST_DISCOUNT_FIELD
elif custom_llm_provider and custom_llm_provider in litellm.cost_discount_config:
discount_percent = litellm.cost_discount_config[custom_llm_provider]
if discount_percent:
discount_amount = original_cost * discount_percent
final_cost: Final = original_cost - discount_amount
if verbose_logger.isEnabledFor(logging.DEBUG):
verbose_logger.debug(
f"Applied {discount_percent * 100}% discount to {custom_llm_provider}: "
f"Applied {discount_percent * 100}% discount to {discount_source}: "
f"${original_cost:.6f} -> ${final_cost:.6f} (saved ${discount_amount:.6f})"
)
@ -991,6 +1003,28 @@ def _apply_cost_discount(
return base_cost, discount_percent, discount_amount
def _get_model_cost_discount(
model: str | None = None,
model_info: ModelInfo | None = None,
) -> float | None:
discount: Any = None
if model_info is not None:
discount = model_info.get(_COST_DISCOUNT_FIELD)
if discount is None and model is not None:
registered_model_info = litellm.model_cost.get(model)
if isinstance(registered_model_info, dict):
discount = registered_model_info.get(_COST_DISCOUNT_FIELD)
if discount is None:
return None
discount_float = float(discount)
if not 0 <= discount_float <= 1:
raise ValueError("cost_discount must be between 0 and 1")
return discount_float
def _apply_cost_margin(
base_cost: float,
custom_llm_provider: str | None,
@ -1477,6 +1511,7 @@ def completion_cost(
) = _apply_cost_discount(
base_cost=_final_cost,
custom_llm_provider=custom_llm_provider,
model=router_model_id or model,
)
# Apply margin from module-level config if configured
@ -1630,18 +1665,15 @@ def completion_cost(
_final_cost += sum(additional_costs.values())
original_cost = _final_cost
if litellm.cost_discount_config:
(
_final_cost,
discount_percent,
discount_amount,
) = _apply_cost_discount(
base_cost=_final_cost,
custom_llm_provider=custom_llm_provider,
)
else:
discount_percent = 0.0
discount_amount = 0.0
(
_final_cost,
discount_percent,
discount_amount,
) = _apply_cost_discount(
base_cost=_final_cost,
custom_llm_provider=custom_llm_provider,
model=router_model_id or model,
)
# Apply margin from module-level config if configured
if litellm.cost_margin_config:

View file

@ -5283,8 +5283,10 @@ def completion(
### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ###
if (
input_cost_per_token is not None and output_cost_per_token is not None
) or input_cost_per_second is not None:
(input_cost_per_token is not None and output_cost_per_token is not None)
or input_cost_per_second is not None
or kwargs.get("cost_discount") is not None
):
_register_custom_pricing_for_request(
model=model,
custom_llm_provider=custom_llm_provider,
@ -6139,7 +6141,11 @@ def embedding(
)
### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ###
if (input_cost_per_token is not None and output_cost_per_token is not None) or input_cost_per_second is not None:
if (
(input_cost_per_token is not None and output_cost_per_token is not None)
or input_cost_per_second is not None
or kwargs.get("cost_discount") is not None
):
_register_custom_pricing_for_request(
model=model,
custom_llm_provider=custom_llm_provider,

View file

@ -148,6 +148,7 @@ class ModelInfo(MirroredPricingParams):
base_model: str | None = None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking
tier: Literal["free", "paid"] | None = None
cost_discount: float | None = None
"""
Team Model Specific Fields
@ -208,6 +209,8 @@ class ModelInfo(MirroredPricingParams):
end: Final = _as_utc(self.ptu_effective_to)
if start is not None and end is not None and end <= start:
raise ValueError("ptu_effective_to must be after ptu_effective_from")
if self.cost_discount is not None and not 0 <= self.cost_discount <= 1:
raise ValueError("cost_discount must be between 0 and 1")
return self
model_config = ConfigDict(extra="allow")
@ -482,6 +485,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
output_cost_per_second: float | None
output_cost_per_second_1080p: float | None
num_retries: int | None
cost_discount: float | None
## MOCK RESPONSES ##
mock_response: str | ModelResponse | Exception | None

View file

@ -219,6 +219,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
cache_read_input_token_cost_above_272k_tokens_priority: float | None
cache_read_input_token_cost_above_272k_tokens_flex: float | None
cache_read_input_token_cost_above_512k_tokens: float | None
cost_discount: float | None
# Smallest prefix this model will actually cache, whatever caching mechanism its provider uses.
# Absent means the provider-agnostic default applies; see MINIMUM_PROMPT_CACHE_TOKEN_COUNT.
prompt_cache_min_tokens: int | None
@ -3333,6 +3334,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_second_1080p: float | None = None
input_cost_per_pixel: float | None = None
output_cost_per_pixel: float | None = None
cost_discount: float | None = None
# Include all ModelInfoBase fields as optional
# This allows any model_info parameter to be set in litellm_params
@ -3411,6 +3413,13 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
regional_processing_uplift_multiplier_us: float | None = None
regional_endpoint_uplift_multiplier: float | None = None
@field_validator("cost_discount")
@classmethod
def validate_cost_discount(cls, value: float | None) -> float | None:
if value is not None and not 0 <= value <= 1:
raise ValueError("cost_discount must be between 0 and 1")
return value
@classmethod
def strip_custom_pricing_fields(cls, model_info: dict[str, Any]) -> dict[str, Any]:
"""Return a copy of ``model_info`` without per-deployment custom pricing fields.

View file

@ -63,6 +63,7 @@ class TestStripClientPricingOverrides:
)
# Sanity: the obvious top-level pricing fields are in the set.
for field in (
"cost_discount",
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_second",

View file

@ -1885,6 +1885,66 @@ def test_cost_discount_not_applied_to_other_providers(monkeypatch):
print(f" - Cost remains unchanged: ${cost_with_selective_discount:.6f}")
def test_deployment_cost_discount_from_router_model_id_overrides_provider_discount():
"""
Test that deployment-level discounts can be configured without overriding
the base model prices, and take precedence over provider-level discounts.
"""
original_discount_config = litellm.cost_discount_config.copy()
deployment_model_id = "test-deployment-cost-discount"
original_deployment_entry = litellm.model_cost.get(deployment_model_id)
def _response() -> ModelResponse:
return ModelResponse(
id="test-id",
choices=[],
created=1234567890,
model="gpt-4o-mini",
object="chat.completion",
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
)
try:
litellm.cost_discount_config = {}
base_cost = completion_cost(
completion_response=_response(),
model="gpt-4o-mini",
custom_llm_provider="openai",
)
litellm.model_cost[deployment_model_id] = {"cost_discount": 0.8}
litellm.cost_discount_config = {"openai": 0.05}
discounted_cost = completion_cost(
completion_response=_response(),
model="gpt-4o-mini",
custom_llm_provider="openai",
custom_pricing=True,
router_model_id=deployment_model_id,
)
assert discounted_cost == pytest.approx(base_cost * 0.2, rel=1e-9)
finally:
litellm.cost_discount_config = original_discount_config
if original_deployment_entry is None:
litellm.model_cost.pop(deployment_model_id, None)
else:
litellm.model_cost[deployment_model_id] = original_deployment_entry
def test_cost_discount_validation():
from litellm.types.router import ModelInfo as RouterModelInfo
from litellm.types.utils import CustomPricingLiteLLMParams
CustomPricingLiteLLMParams(cost_discount=0.8)
RouterModelInfo(cost_discount=0.8)
with pytest.raises(ValueError, match="cost_discount must be between 0 and 1"):
CustomPricingLiteLLMParams(cost_discount=1.1)
with pytest.raises(ValueError, match="cost_discount must be between 0 and 1"):
RouterModelInfo(cost_discount=1.1)
def test_cost_margin_percentage(monkeypatch):
"""
Test that percentage-based cost margin is applied correctly

View file

@ -17,6 +17,7 @@ import pytest
import litellm
from litellm.main import _build_custom_pricing_entry
from litellm.types.router import ModelInfo
from litellm.utils import _invalidate_model_cost_lowercase_map
@ -37,6 +38,7 @@ def test_build_custom_pricing_entry_includes_all_kwargs_fields():
"""All CustomPricingLiteLLMParams fields present in kwargs should be
included in the resulting entry dict."""
kwargs = {
"cost_discount": 0.8,
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
"cache_read_input_token_cost": 0.00025,
@ -52,6 +54,7 @@ def test_build_custom_pricing_entry_includes_all_kwargs_fields():
)
assert entry["litellm_provider"] == "openai"
assert entry["cost_discount"] == 0.8
assert entry["input_cost_per_token"] == 0.001
assert entry["output_cost_per_token"] == 0.002
assert entry["cache_read_input_token_cost"] == 0.00025
@ -61,6 +64,120 @@ def test_build_custom_pricing_entry_includes_all_kwargs_fields():
assert "unrelated_kwarg" not in entry
def test_model_info_validates_cost_discount():
ModelInfo(cost_discount=0.8)
with pytest.raises(ValueError, match="cost_discount must be between 0 and 1"):
ModelInfo(cost_discount=1.1)
def test_router_registers_deployment_cost_discount_and_litellm_params_override():
model_info_only_id = "deployment-cost-discount-model-info"
litellm_params_only_id = "deployment-cost-discount-litellm-params"
litellm_params_override_id = "deployment-cost-discount-override"
shared_backend_key = "openai/gpt-4o-mini"
snapshot = _snapshot_model_cost_entries(
[
model_info_only_id,
litellm_params_only_id,
litellm_params_override_id,
shared_backend_key,
]
)
try:
litellm.Router(
model_list=[
{
"model_name": "deployment-cost-discount-test",
"litellm_params": {
"model": shared_backend_key,
"api_key": "test-api-key",
},
"model_info": {
"id": model_info_only_id,
"cost_discount": 0.6,
},
},
{
"model_name": "deployment-cost-discount-test",
"litellm_params": {
"model": shared_backend_key,
"api_key": "test-api-key",
"cost_discount": 0.7,
},
"model_info": {
"id": litellm_params_only_id,
},
},
{
"model_name": "deployment-cost-discount-test",
"litellm_params": {
"model": shared_backend_key,
"api_key": "test-api-key",
"cost_discount": 0.8,
},
"model_info": {
"id": litellm_params_override_id,
"cost_discount": 0.4,
},
},
]
)
assert litellm.model_cost[model_info_only_id]["cost_discount"] == 0.6
assert litellm.model_cost[litellm_params_only_id]["cost_discount"] == 0.7
assert litellm.model_cost[litellm_params_override_id]["cost_discount"] == 0.8
shared_backend_entry = litellm.model_cost.get(shared_backend_key)
if shared_backend_entry is not None:
assert "cost_discount" not in shared_backend_entry
finally:
_restore_model_cost_entries(snapshot)
def test_router_completion_applies_deployment_cost_discount_without_custom_base_prices():
shared_backend_key = "openai/gpt-4o-mini"
deployment_id = "deployment-cost-discount-runtime"
snapshot = _snapshot_model_cost_entries([deployment_id, shared_backend_key])
try:
router = litellm.Router(
model_list=[
{
"model_name": "deployment-runtime-discount-test",
"litellm_params": {
"model": shared_backend_key,
"api_key": "test-api-key",
},
"model_info": {
"id": deployment_id,
"cost_discount": 0.8,
},
},
]
)
messages = [{"role": "user", "content": "hello"}]
discounted_response = router.completion(
model="deployment-runtime-discount-test",
messages=messages,
mock_response="ok",
max_tokens=20,
)
undiscounted_response = litellm.completion(
model=shared_backend_key,
messages=messages,
mock_response="ok",
max_tokens=20,
)
assert discounted_response._hidden_params["response_cost"] == pytest.approx(
undiscounted_response._hidden_params["response_cost"] * 0.2,
rel=1e-9,
)
finally:
_restore_model_cost_entries(snapshot)
def test_build_custom_pricing_entry_merges_model_info_metadata():
"""Fields from model_info (mode, supports_prompt_caching, max_tokens)
should be merged into the entry when present."""

View file

@ -27493,6 +27493,8 @@ export interface components {
complexity_router_default_model?: string | null;
/** Configurable Clientside Auth Params */
configurable_clientside_auth_params?: (string | components["schemas"]["ConfigurableClientsideParamsCustomAuth-Input"])[] | null;
/** Cost Discount */
cost_discount?: number | null;
/** Custom Llm Provider */
custom_llm_provider?: string | null;
/** Default Api Key Rpm Limit */
@ -36534,6 +36536,8 @@ export interface components {
cache_read_input_token_cost?: number | null;
/** Cost Per Ptu Per Hour */
cost_per_ptu_per_hour?: number | null;
/** Cost Discount */
cost_discount?: number | null;
/** Created At */
created_at?: string | null;
/** Created By */
@ -36702,6 +36706,8 @@ export interface components {
complexity_router_default_model?: string | null;
/** Configurable Clientside Auth Params */
configurable_clientside_auth_params?: (string | components["schemas"]["ConfigurableClientsideParamsCustomAuth-Input"])[] | null;
/** Cost Discount */
cost_discount?: number | null;
/** Custom Llm Provider */
custom_llm_provider?: string | null;
/** Default Api Key Rpm Limit */