mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): enforce lifetime limits for explicit expiration
This commit is contained in:
parent
4956fc8404
commit
acdccee0fe
2 changed files with 157 additions and 34 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue