From 0b7d8b9a0d91f75a6cf6e525e312a9cb72e4daee Mon Sep 17 00:00:00 2001 From: Harshit28j Date: Sat, 28 Feb 2026 16:05:55 +0530 Subject: [PATCH] fix: edge case when key alias empty --- .../proxy/hooks/key_management_event_hooks.py | 2 +- .../key_management_endpoints.py | 42 +++++++ .../test_key_rotation_integration.py | 116 ++++++++++++++++++ .../test_key_management_endpoints.py | 35 ++++++ 4 files changed, 194 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index c07f30f8646..95c9c806120 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -153,7 +153,7 @@ class KeyManagementEventHooks: new_secret_name = ( response.key_alias or data.key_alias - or f"virtual-key-{response.token_id}" + or initial_secret_name ) verbose_proxy_logger.info( "Updating secret in secret manager: secret_name=%s", diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 5b56133f1ce..c5f83de0145 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -13,6 +13,7 @@ import asyncio import copy import inspect import json +import re import os import secrets import traceback @@ -136,6 +137,17 @@ def _set_key_rotation_fields( rotation_interval: The rotation interval string (required if auto_rotate is True) """ if auto_rotate and rotation_interval: + if ( + litellm._key_management_settings is not None + and litellm._key_management_settings.store_virtual_keys is True + and data.get("key_alias") is None + ): + raise ProxyException( + message="key_alias is required when auto_rotate=True and store_virtual_keys is enabled. This ensures stable secret naming during rotation.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=400, + ) data.update( { "auto_rotate": auto_rotate, @@ -625,6 +637,8 @@ async def _common_key_generation_helper( # noqa: PLR0915 prisma_client=prisma_client, ) + _validate_key_alias_format(key_alias=data_json.get("key_alias", None)) + await _enforce_unique_key_alias( key_alias=data_json.get("key_alias", None), prisma_client=prisma_client, @@ -1931,6 +1945,8 @@ async def update_key_fn( data=data, existing_key_row=existing_key_row ) + _validate_key_alias_format(key_alias=non_default_values.get("key_alias", None)) + await _enforce_unique_key_alias( key_alias=non_default_values.get("key_alias", None), prisma_client=prisma_client, @@ -3378,6 +3394,7 @@ async def _execute_virtual_key_regeneration( non_default_values = await prepare_key_update_data( data=data, existing_key_row=key_in_db ) + _validate_key_alias_format(key_alias=non_default_values.get("key_alias")) verbose_proxy_logger.debug("non_default_values: %s", non_default_values) update_data.update(non_default_values) update_data = prisma_client.jsonify_object(data=update_data) @@ -4952,6 +4969,31 @@ async def test_key_logging( ) +_KEY_ALIAS_PATTERN = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9_\-/\.]{0,253}[a-zA-Z0-9]$") + + +def _validate_key_alias_format(key_alias: Optional[str]) -> None: + """ + Validate the format of the key_alias. + + Rules: + - None is OK (no alias). + - Otherwise must be 2–255 chars + - start/end with alphanumeric + - only allow a-zA-Z0-9_-/. + """ + if key_alias is None: + return + + if not _KEY_ALIAS_PATTERN.match(key_alias): + raise ProxyException( + message="Invalid key_alias format. Must be 2-255 characters, start/end with alphanumeric, and only contain a-zA-Z0-9_-/.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=400, + ) + + async def _enforce_unique_key_alias( key_alias: Optional[str], prisma_client: Any, diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py index 308c8cdbce1..65c89a35f0b 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py @@ -142,3 +142,119 @@ class TestKeyRotationManagerPassesKeyAlias: assert ( captured_request.key_alias is None ), "key_alias should be None for keys without alias" + +class TestKeyRotationSecretNamingStability: + """ + Tests that the fallback secret name in the rotation hook remains stable + across rotations to prevent AWS secret sprawl. + + Couple this with the validation fix (Step 1-2) to ensure a stable + experience for secret management. + """ + + @pytest.mark.asyncio + async def test_rotation_hook_uses_initial_secret_name_fallback(self): + """ + GIVEN: A key WITHOUT an alias (has an initial_secret_name based on token ID) + WHEN: The key is rotated + THEN: The hook MUST reuse the existing secret name, NOT generate a new one + based on the new token ID. + """ + from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + + # 1. Existing key without alias + initial_token_hash = "hashed-initial-token" + existing_key = MagicMock(spec=LiteLLM_VerificationToken) + existing_key.token = initial_token_hash + existing_key.key_alias = None + initial_secret_name = f"virtual-key-{initial_token_hash}" + + # 2. Rotation response (new token ID) + new_token_id = "hashed-new-token" + response = GenerateKeyResponse( + key="sk-new-key", + token_id=new_token_id, + key_alias=None + ) + + # 3. Request data without alias + request_data = RegenerateKeyRequest( + key=initial_token_hash, + key_alias=None + ) + + with patch("litellm.proxy.hooks.key_management_event_hooks.KeyManagementEventHooks._rotate_virtual_key_in_secret_manager", new_callable=AsyncMock) as mock_rotate: + await KeyManagementEventHooks.async_key_rotated_hook( + data=request_data, + existing_key_row=existing_key, + response=response, + user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin", api_key="sk-1234", user_id="1234") + ) + + # ASSERT: The new_secret_name MUST be the same as initial_secret_name + # This ensures PutSecretValue instead of a new secret creation + mock_rotate.assert_called_once() + call_kwargs = mock_rotate.call_args.kwargs + assert call_kwargs["current_secret_name"] == initial_secret_name + assert call_kwargs["new_secret_name"] == initial_secret_name, \ + f"Secret name drift! Expected {initial_secret_name}, got {call_kwargs['new_secret_name']}. This causes secret sprawl." + + @pytest.mark.asyncio + async def test_rotation_hook_pre_rotation_alias_consistency(self): + """ + GIVEN: A key WITH an alias + WHEN: The key is rotated + THEN: The hook uses the alias for both current and new names. + """ + from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + + test_alias = "tenant1/stable-key" + existing_key = MagicMock(spec=LiteLLM_VerificationToken) + existing_key.token = "old-hash" + existing_key.key_alias = test_alias + + response = GenerateKeyResponse(token_id="new-hash", key="sk-new", key_alias=test_alias) + request_data = RegenerateKeyRequest(key="old-hash", key_alias=test_alias) + + with patch("litellm.proxy.hooks.key_management_event_hooks.KeyManagementEventHooks._rotate_virtual_key_in_secret_manager", new_callable=AsyncMock) as mock_rotate: + await KeyManagementEventHooks.async_key_rotated_hook( + data=request_data, + existing_key_row=existing_key, + response=response, + user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin", api_key="sk-123", user_id="1") + ) + mock_rotate.assert_called_once() + assert mock_rotate.call_args.kwargs["current_secret_name"] == test_alias + assert mock_rotate.call_args.kwargs["new_secret_name"] == test_alias + + @pytest.mark.asyncio + async def test_set_key_rotation_fields_requires_alias(self): + """ + Tests that _set_key_rotation_fields enforces key_alias requirement + when secret storage is enabled. + """ + import litellm + from litellm.proxy.management_endpoints.key_management_endpoints import _set_key_rotation_fields + from litellm.proxy._types import ProxyException + # Create a mock for settings + mock_settings = MagicMock() + mock_settings.store_virtual_keys = True + + # Mock settings: store_virtual_keys = True + with patch("litellm._key_management_settings", mock_settings): + data = {"auto_rotate": True} # Missing key_alias + + # Should raise ProxyException 400 + with pytest.raises(ProxyException) as exc: + _set_key_rotation_fields(data, auto_rotate=True, rotation_interval="30d") + + assert str(exc.value.code) == "400" + assert "key_alias is required" in str(exc.value.message) + + # Adding key_alias should work + data["key_alias"] = "valid-alias" + _set_key_rotation_fields(data, auto_rotate=True, rotation_interval="30d") + assert data["auto_rotate"] is True + assert "key_rotation_at" in data 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 7565e901ecd..e928df21d14 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 @@ -6305,3 +6305,38 @@ async def test_key_aliases_no_search_omits_ilike_filter(): assert "ILIKE" not in count_sql + +class TestValidateKeyAliasFormat: + def test_validate_key_alias_format_valid(self): + from litellm.proxy.management_endpoints.key_management_endpoints import _validate_key_alias_format + # Valid cases + _validate_key_alias_format(None) # OK + _validate_key_alias_format("valid-alias") + _validate_key_alias_format("valid_alias") + _validate_key_alias_format("valid.alias") + _validate_key_alias_format("valid/alias") + _validate_key_alias_format("a" * 255) + _validate_key_alias_format("my-key-123") + + def test_validate_key_alias_format_invalid(self): + from litellm.proxy.management_endpoints.key_management_endpoints import _validate_key_alias_format + from litellm.proxy._types import ProxyException + + invalid_aliases = [ + "", # empty + " ", # whitespace + "a", # too short (min 2) + "!", # special char + "-start", # non-alphanumeric start + "end-", # non-alphanumeric end + "invalid@char", # invalid char + "a" * 256, # too long + " leading", + "trailing ", + ] + + for alias in invalid_aliases: + with pytest.raises(ProxyException) as exc: + _validate_key_alias_format(alias) + assert str(exc.value.code) == "400" + assert "Invalid key_alias format" in str(exc.value.message)