From acdccee0fe72c38ecc3d4b156086a541446e3854 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 26 Sep 2026 11:51:50 -0500 Subject: [PATCH] fix(proxy): enforce lifetime limits for explicit expiration --- .../key_management_endpoints.py | 80 ++++++++----- .../test_key_management_endpoints.py | 111 +++++++++++++++++- 2 files changed, 157 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 6e963b095af..db9eb38732b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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", diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index a5729e394c4..26d37c5e0e2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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,