fix: edge case when key alias empty

This commit is contained in:
Harshit28j 2026-02-28 16:05:55 +05:30
parent 50ecfbc614
commit 0b7d8b9a0d
4 changed files with 194 additions and 1 deletions

View file

@ -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",

View file

@ -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,

View file

@ -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

View file

@ -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)