fix(proxy): enforce lifetime limits for explicit expiration

This commit is contained in:
Emerson Gomes 2026-09-26 11:51:50 -05:00
parent 4956fc8404
commit acdccee0fe
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
2 changed files with 157 additions and 34 deletions

View file

@ -1109,6 +1109,20 @@ _BUDGET_NUMERIC_KEYS = frozenset(
)
def _key_expiration_exceeds_limit(
data: GenerateKeyRequest | UpdateKeyRequest, maximum_duration: str, fill_defaults: bool
) -> bool:
if data.duration:
return False
if not fill_defaults and data.duration is None and "duration" in data.model_fields_set:
return True
if data.expires is None:
return not fill_defaults and "expires" in data.model_fields_set
expires: Final = data.expires if data.expires.tzinfo is not None else data.expires.replace(tzinfo=timezone.utc)
maximum_expiration: Final = datetime.now(timezone.utc) + timedelta(seconds=duration_in_seconds(maximum_duration))
return expires > maximum_expiration
def _enforce_upperbound_key_params(
data: GenerateKeyRequest | UpdateKeyRequest,
fill_defaults: bool = True,
@ -1135,40 +1149,41 @@ def _enforce_upperbound_key_params(
if litellm.upperbound_key_generate_params is None:
return
maximum_duration: Final = litellm.upperbound_key_generate_params.duration
if maximum_duration is not None and _key_expiration_exceeds_limit(data, maximum_duration, fill_defaults):
raise HTTPException(
status_code=400,
detail={"error": f"expires is over max limit set in config - max_duration={maximum_duration}"},
)
for elem in data:
key, value = elem
upperbound_value = getattr(litellm.upperbound_key_generate_params, key, None)
if upperbound_value is not None:
if value is None:
if fill_defaults:
setattr(data, key, upperbound_value)
else:
if key in [
"max_budget",
"max_parallel_requests",
"tpm_limit",
"rpm_limit",
]:
if value > upperbound_value:
raise HTTPException(
status_code=400,
detail={
"error": f"{key} is over max limit set in config - user_value={value}; max_value={upperbound_value}"
},
)
elif key in ["budget_duration", "duration"]:
upperbound_duration = duration_in_seconds(duration=upperbound_value)
if value == "-1":
user_duration = float("inf")
else:
user_duration = duration_in_seconds(duration=value)
if user_duration > upperbound_duration:
raise HTTPException(
status_code=400,
detail={
"error": f"{key} is over max limit set in config - user_value={value}; max_value={upperbound_value}"
},
)
if (upperbound_value := getattr(litellm.upperbound_key_generate_params, key, None)) is None:
continue
if value is None:
if fill_defaults and not (key == "duration" and data.expires is not None):
setattr(data, key, upperbound_value)
continue
if key in ["max_budget", "max_parallel_requests", "tpm_limit", "rpm_limit"]:
if value > upperbound_value:
raise HTTPException(
status_code=400,
detail={
"error": f"{key} is over max limit set in config - user_value={value}; max_value={upperbound_value}"
},
)
elif key in ["budget_duration", "duration"]:
upperbound_duration, user_duration = (
duration_in_seconds(duration=upperbound_value),
float("inf") if value == "-1" else duration_in_seconds(duration=value),
)
if user_duration > upperbound_duration:
raise HTTPException(
status_code=400,
detail={
"error": f"{key} is over max limit set in config - user_value={value}; max_value={upperbound_value}"
},
)
async def _common_key_generation_helper(
@ -1242,6 +1257,7 @@ async def _common_key_generation_helper(
if (
value is None
and (key != "budget_duration" or key not in data.model_fields_set)
and (key != "duration" or data.expires is None)
and key
in [
"max_budget",

View file

@ -85,6 +85,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyResponse,
CustomKeyPolicyRequest,
)
from litellm.types.proxy.management_endpoints.ui_sso import LiteLLM_UpperboundKeyGenerateParams
client = TestClient(app)
@ -12722,6 +12723,87 @@ def test_enforce_upperbound_duration_over_limit(monkeypatch):
assert "duration" in str(exc_info.value.detail)
@pytest.mark.parametrize(
("request_type", "fill_defaults"),
((GenerateKeyRequest, True), (UpdateKeyRequest, False), (RegenerateKeyRequest, False)),
)
@pytest.mark.parametrize("with_timezone", (False, True))
def test_enforce_upperbound_rejects_expiration_beyond_limit(
monkeypatch: pytest.MonkeyPatch,
request_type: type[GenerateKeyRequest] | type[UpdateKeyRequest],
fill_defaults: bool,
with_timezone: bool,
) -> None:
monkeypatch.setattr(litellm, "upperbound_key_generate_params", LiteLLM_UpperboundKeyGenerateParams(duration="1d"))
future: Final = datetime.now(timezone.utc) + timedelta(days=2)
expires: Final = future if with_timezone else future.replace(tzinfo=None)
data: Final = request_type(key="sk-expiration-test", expires=expires)
with pytest.raises(HTTPException) as error:
_enforce_upperbound_key_params(data, fill_defaults=fill_defaults)
assert error.value.status_code == 400
assert "expires" in str(error.value.detail)
@pytest.mark.parametrize(
("request_type", "fill_defaults"),
((GenerateKeyRequest, True), (UpdateKeyRequest, False), (RegenerateKeyRequest, False)),
)
def test_enforce_upperbound_preserves_bounded_expiration(
monkeypatch: pytest.MonkeyPatch,
request_type: type[GenerateKeyRequest] | type[UpdateKeyRequest],
fill_defaults: bool,
) -> None:
monkeypatch.setattr(litellm, "upperbound_key_generate_params", LiteLLM_UpperboundKeyGenerateParams(duration="1d"))
expires: Final = datetime.now(timezone.utc) + timedelta(hours=1)
data: Final = request_type(key="sk-expiration-test", expires=expires)
_enforce_upperbound_key_params(data, fill_defaults=fill_defaults)
assert (data.duration, data.expires) == (None, expires)
@pytest.mark.parametrize("request_type", (UpdateKeyRequest, RegenerateKeyRequest))
@pytest.mark.parametrize("clear_duration", (False, True))
def test_enforce_upperbound_rejects_clearing_expiration(
monkeypatch: pytest.MonkeyPatch,
request_type: type[UpdateKeyRequest] | type[RegenerateKeyRequest],
clear_duration: bool,
) -> None:
monkeypatch.setattr(litellm, "upperbound_key_generate_params", LiteLLM_UpperboundKeyGenerateParams(duration="1d"))
expires: Final = datetime.now(timezone.utc) + timedelta(hours=1)
data: Final = request_type.model_validate(
{"key": "sk-expiration-test", "duration": None, "expires": expires}
if clear_duration
else {"key": "sk-expiration-test", "expires": None}
)
with pytest.raises(HTTPException) as error:
_enforce_upperbound_key_params(data, fill_defaults=False)
assert error.value.status_code == 400
assert "expires" in str(error.value.detail)
@pytest.mark.parametrize(
("request_type", "fill_defaults"),
((GenerateKeyRequest, True), (UpdateKeyRequest, False), (RegenerateKeyRequest, False)),
)
def test_enforce_upperbound_uses_explicit_duration_over_expiration(
monkeypatch: pytest.MonkeyPatch,
request_type: type[GenerateKeyRequest] | type[UpdateKeyRequest],
fill_defaults: bool,
) -> None:
monkeypatch.setattr(litellm, "upperbound_key_generate_params", LiteLLM_UpperboundKeyGenerateParams(duration="1d"))
expires: Final = datetime.now(timezone.utc) + timedelta(days=2)
data: Final = request_type(key="sk-expiration-test", duration="1h", expires=expires)
_enforce_upperbound_key_params(data, fill_defaults=fill_defaults)
assert (data.duration, data.expires) == ("1h", expires)
def test_enforce_upperbound_no_config_is_noop(monkeypatch):
"""Test that no enforcement happens when upperbound params are not configured."""
import litellm
@ -12783,7 +12865,10 @@ def _make_regenerate_existing_key():
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monkeypatch):
@pytest.mark.parametrize("use_absolute_expiration", (False, True))
async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(
monkeypatch: pytest.MonkeyPatch, use_absolute_expiration: bool
) -> None:
"""Regenerate must reject durations exceeding upperbound_key_generate_params.duration."""
from litellm.proxy._types import RegenerateKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -12801,7 +12886,11 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monk
),
)
existing_key = _make_regenerate_existing_key()
data = RegenerateKeyRequest(duration="2h")
data: Final = (
RegenerateKeyRequest(expires=datetime.now(timezone.utc) + timedelta(hours=2))
if use_absolute_expiration
else RegenerateKeyRequest(duration="2h")
)
user_api_key_dict = _make_regenerate_user_api_key_dict()
mock_prisma_client = _make_regenerate_mock_prisma()
@ -19310,6 +19399,24 @@ async def test_key_generate_omitted_budget_duration_still_takes_default_key_gene
assert key_row["budget_reset_at"] is not None
@pytest.mark.parametrize("with_lifetime_limit", (False, True))
async def test_key_generate_expiration_overrides_default_duration(
monkeypatch: pytest.MonkeyPatch, with_lifetime_limit: bool
) -> None:
monkeypatch.setattr(litellm, "default_key_generate_params", {"duration": "1d"})
monkeypatch.setattr(
litellm,
"upperbound_key_generate_params",
LiteLLM_UpperboundKeyGenerateParams(duration="1d") if with_lifetime_limit else None,
)
insert_data: Final = _wire_key_generation_prisma(monkeypatch)
expires: Final = datetime.now(timezone.utc) + timedelta(hours=1)
key_row: Final = await _generate_key_and_get_persisted_row(GenerateKeyRequest(expires=expires), insert_data)
assert key_row["expires"] == expires
@pytest.mark.asyncio
async def test_key_generate_explicit_null_budget_duration_cannot_bypass_upperbound(monkeypatch):
"""upperbound_key_generate_params is an admin ceiling: an explicit null must not mint an uncapped key,