mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: edge case when key alias empty
This commit is contained in:
parent
50ecfbc614
commit
0b7d8b9a0d
4 changed files with 194 additions and 1 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue