diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6ffd6816ca6..c6d4ee1120a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e7ee5ffa849..a94a75fdfa3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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, 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 51f72f91dc3..867ef759fb3 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 @@ -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): """ diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index c8e0b3a730a..5354de182a0 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -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" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 541d7a17ae1..79f03978b12 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */