mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): allow /key/update to identify the key by key_alias (#34851)
* fix(proxy): allow /key/update to identify the key by key_alias * fix(ui): drop machine-dependent union-order churn from generated schema.d.ts
This commit is contained in:
parent
fe1670fc06
commit
40878a1ed5
5 changed files with 277 additions and 16 deletions
|
|
@ -1169,7 +1169,6 @@ class GenerateKeyResponse(KeyRequestBase):
|
|||
class UpdateKeyRequest(KeyRequestBase):
|
||||
# Note: the defaults of all Params here MUST BE NONE
|
||||
# else they will get overwritten
|
||||
key: str # type: ignore
|
||||
duration: Optional[str] = None
|
||||
spend: Optional[float] = None
|
||||
metadata: Optional[dict] = None
|
||||
|
|
@ -1186,6 +1185,12 @@ class UpdateKeyRequest(KeyRequestBase):
|
|||
raise ValueError("temp_budget_increase and temp_budget_expiry must be set together")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_key_identifier(self) -> "UpdateKeyRequest":
|
||||
if self.key is None and self.key_alias is None:
|
||||
raise ValueError("either key or key_alias must be provided")
|
||||
return self
|
||||
|
||||
|
||||
class RegenerateKeyRequest(GenerateKeyRequest):
|
||||
# This needs to be different from UpdateKeyRequest, because "key" is optional for this
|
||||
|
|
|
|||
|
|
@ -2023,7 +2023,7 @@ def _validate_max_budget(max_budget: Optional[float]) -> None:
|
|||
|
||||
|
||||
async def _get_and_validate_existing_key(
|
||||
token: str, prisma_client: Optional[PrismaClient]
|
||||
token: str | None, prisma_client: Optional[PrismaClient], key_alias: str | None = None
|
||||
) -> LiteLLM_VerificationToken:
|
||||
"""
|
||||
Get existing key from database and validate it exists.
|
||||
|
|
@ -2031,12 +2031,13 @@ async def _get_and_validate_existing_key(
|
|||
Args:
|
||||
token: The key token to look up
|
||||
prisma_client: Prisma client instance
|
||||
key_alias: Alias to look the key up by when token is not provided
|
||||
|
||||
Returns:
|
||||
LiteLLM_VerificationToken: The existing key row
|
||||
|
||||
Raises:
|
||||
ProxyException: 404 if key is not found
|
||||
ProxyException: 404 if key is not found, 400 if the alias matches multiple keys
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -2044,19 +2045,65 @@ async def _get_and_validate_existing_key(
|
|||
detail={"error": "Database not connected"},
|
||||
)
|
||||
|
||||
hashed_token = _hash_token_if_needed(token=token)
|
||||
if token is not None:
|
||||
hashed_token = _hash_token_if_needed(token=token)
|
||||
|
||||
existing_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token})
|
||||
existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository(
|
||||
prisma_client
|
||||
).table.find_unique(where={"token": hashed_token})
|
||||
|
||||
if existing_key_row is None:
|
||||
if existing_key_row is None:
|
||||
raise ProxyException(
|
||||
message="Key not found.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
return existing_key_row
|
||||
|
||||
if key_alias is None:
|
||||
raise ProxyException(
|
||||
message="either key or key_alias must be provided",
|
||||
type=ProxyErrorTypes.bad_request_error,
|
||||
param="key",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
where={"key_alias": key_alias}, take=2
|
||||
)
|
||||
|
||||
if len(rows) == 0:
|
||||
raise ProxyException(
|
||||
message=f"Key not found. No key with key_alias='{key_alias}'.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key_alias",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
if len(rows) > 1:
|
||||
raise ProxyException(
|
||||
message=f"Multiple keys share key_alias='{key_alias}', so it cannot be used as an identifier.",
|
||||
type=ProxyErrorTypes.bad_request_error,
|
||||
param="key_alias",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> str:
|
||||
if data.key is not None:
|
||||
return data.key
|
||||
if existing_key_row.token is None:
|
||||
raise ProxyException(
|
||||
message="Key not found.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
return existing_key_row
|
||||
return existing_key_row.token
|
||||
|
||||
|
||||
async def _process_single_key_update(
|
||||
|
|
@ -2508,8 +2555,8 @@ async def update_key_fn(
|
|||
Update an existing API key's parameters.
|
||||
|
||||
Parameters:
|
||||
- key: str - The key to update
|
||||
- key_alias: Optional[str] - User-friendly key alias
|
||||
- key: Optional[str] - The key to update. Either key or key_alias must be provided.
|
||||
- key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
|
||||
- user_id: Optional[str] - User ID associated with key
|
||||
- team_id: Optional[str] - Team ID associated with key
|
||||
- agent_id: Optional[str] - The agent id associated with the key.
|
||||
|
|
@ -2592,14 +2639,14 @@ async def update_key_fn(
|
|||
detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"},
|
||||
)
|
||||
|
||||
data_json: dict = data.model_dump(exclude_unset=True)
|
||||
key = data_json.pop("key")
|
||||
|
||||
# get the row from db
|
||||
existing_key_row = await _get_and_validate_existing_key(
|
||||
token=data.key,
|
||||
prisma_client=prisma_client,
|
||||
key_alias=data.key_alias,
|
||||
)
|
||||
key = _resolve_token_to_update(data=data, existing_key_row=existing_key_row)
|
||||
data.key = key
|
||||
|
||||
await _validate_update_key_data(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -2507,6 +2507,195 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch):
|
|||
assert "Authentication Error" not in str(exc_info.value.message)
|
||||
|
||||
|
||||
def _setup_update_key_mocks(monkeypatch, mock_prisma_client):
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
monkeypatch.setattr("litellm.store_audit_logs", False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_by_alias_only(monkeypatch):
|
||||
"""
|
||||
/key/update identified by key_alias alone resolves the key row via
|
||||
find_many on the alias and updates using the resolved token.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
update_key_fn,
|
||||
)
|
||||
|
||||
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
|
||||
key_in_db = LiteLLM_VerificationToken(
|
||||
token=hashed_token,
|
||||
key_alias="prod-alias",
|
||||
user_id="test-user",
|
||||
max_budget=200.0,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[key_in_db]
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
return_value={"data": {"max_budget": 50.0}}
|
||||
)
|
||||
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
|
||||
)
|
||||
|
||||
request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
|
||||
) as mock_delete_cache:
|
||||
mock_delete_cache.return_value = None
|
||||
result = await update_key_fn(
|
||||
request=MagicMock(),
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with(
|
||||
where={"key_alias": "prod-alias"}, take=2
|
||||
)
|
||||
assert request_data.key == hashed_token
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_not_called()
|
||||
mock_prisma_client.update_data.assert_awaited_once()
|
||||
assert mock_prisma_client.update_data.call_args.kwargs["token"] == hashed_token
|
||||
assert (
|
||||
mock_prisma_client.update_data.call_args.kwargs["data"]["token"] == hashed_token
|
||||
)
|
||||
assert result["key"] == hashed_token
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_by_alias_not_found_returns_404(monkeypatch):
|
||||
"""
|
||||
/key/update with a key_alias matching no key returns 404.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
update_key_fn,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_key_fn(
|
||||
request=MagicMock(),
|
||||
data=UpdateKeyRequest(key_alias="no-such-alias", max_budget=50.0),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "404"
|
||||
assert "not found" in str(exc_info.value.message).lower()
|
||||
mock_prisma_client.update_data.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_by_duplicate_alias_returns_400(monkeypatch):
|
||||
"""
|
||||
/key/update with a key_alias shared by multiple keys returns 400
|
||||
instead of silently updating one of them.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
update_key_fn,
|
||||
)
|
||||
|
||||
rows = [
|
||||
LiteLLM_VerificationToken(token="hashed-token-1", key_alias="dup-alias"),
|
||||
LiteLLM_VerificationToken(token="hashed-token-2", key_alias="dup-alias"),
|
||||
]
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
|
||||
return_value=rows
|
||||
)
|
||||
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_key_fn(
|
||||
request=MagicMock(),
|
||||
data=UpdateKeyRequest(key_alias="dup-alias", max_budget=50.0),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "400"
|
||||
assert "multiple keys" in str(exc_info.value.message).lower()
|
||||
mock_prisma_client.update_data.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch):
|
||||
"""
|
||||
Regression: passing both key and key_alias keeps today's behavior. The key
|
||||
identifies the row (find_unique, never find_many) and key_alias is the new
|
||||
alias to set; the response echoes the caller-passed key.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
update_key_fn,
|
||||
)
|
||||
|
||||
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
|
||||
key_in_db = LiteLLM_VerificationToken(
|
||||
token=hashed_token,
|
||||
key_alias="old-name",
|
||||
user_id="test-user",
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=key_in_db
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
return_value={"data": {"key_alias": "new-name"}}
|
||||
)
|
||||
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
|
||||
) as mock_delete_cache:
|
||||
mock_delete_cache.return_value = None
|
||||
result = await update_key_fn(
|
||||
request=MagicMock(),
|
||||
data=UpdateKeyRequest(key="sk-test-key", key_alias="new-name"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many.assert_not_called()
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once()
|
||||
assert mock_prisma_client.update_data.call_args.kwargs["token"] == "sk-test-key"
|
||||
assert result["key"] == "sk-test-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_key_existing_key_succeeds(monkeypatch):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -139,3 +139,23 @@ def test_key_request_router_settings_keeps_enable_tag_filtering():
|
|||
dumped = req.router_settings.model_dump(exclude_none=True)
|
||||
assert dumped["enable_tag_filtering"] is True
|
||||
assert dumped["num_retries"] == 2
|
||||
|
||||
|
||||
def test_update_key_request_requires_key_or_key_alias():
|
||||
"""``/key/update`` can be addressed by ``key`` or by ``key_alias``;
|
||||
a request with neither has no way to identify the target key and must
|
||||
fail validation before hitting the endpoint."""
|
||||
import pydantic
|
||||
|
||||
from litellm.proxy._types import UpdateKeyRequest
|
||||
|
||||
with pytest.raises(pydantic.ValidationError, match="either key or key_alias must be provided"):
|
||||
UpdateKeyRequest(max_budget=10.0)
|
||||
|
||||
by_key = UpdateKeyRequest(key="sk-1234")
|
||||
assert by_key.key == "sk-1234"
|
||||
assert by_key.key_alias is None
|
||||
|
||||
by_alias = UpdateKeyRequest(key_alias="my-alias")
|
||||
assert by_alias.key is None
|
||||
assert by_alias.key_alias == "my-alias"
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -6935,8 +6935,8 @@ export interface paths {
|
|||
* @description Update an existing API key's parameters.
|
||||
*
|
||||
* Parameters:
|
||||
* - key: str - The key to update
|
||||
* - key_alias: Optional[str] - User-friendly key alias
|
||||
* - key: Optional[str] - The key to update. Either key or key_alias must be provided.
|
||||
* - key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases)
|
||||
* - user_id: Optional[str] - User ID associated with key
|
||||
* - team_id: Optional[str] - Team ID associated with key
|
||||
* - agent_id: Optional[str] - The agent id associated with the key.
|
||||
|
|
@ -32437,7 +32437,7 @@ export interface components {
|
|||
/** Guardrails */
|
||||
guardrails?: string[] | null;
|
||||
/** Key */
|
||||
key: string;
|
||||
key?: string | null;
|
||||
/** Key Alias */
|
||||
key_alias?: string | null;
|
||||
/** Max Budget */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue