From fb3a5bc1b7a1ccfeea7439d343e02c6cff1d15d5 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 17 Jul 2025 22:28:29 -0700 Subject: [PATCH] feat(internal_user_endpoints.py): new `/user/bulk_update` endpoint (#12720) * feat(internal_user_endpoints.py): new `/user/bulk_update` endpoint enable bulk updating users on the UI * refactor: cleanup unused import --- .../internal_user_endpoints.py | 378 +++++++++++++----- .../internal_user_endpoints.py | 27 +- 2 files changed, 296 insertions(+), 109 deletions(-) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 83de641fbf..bd94d46ec4 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -43,7 +43,10 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendMetrics, ) from litellm.types.proxy.management_endpoints.internal_user_endpoints import ( + BulkUpdateUserRequest, + BulkUpdateUserResponse, UserListResponse, + UserUpdateResult, ) if TYPE_CHECKING: @@ -786,6 +789,143 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di return non_default_values +async def _update_single_user_helper( + user_request: UpdateUserRequest, + user_api_key_dict: UserAPIKeyAuth, + litellm_changed_by: Optional[str] = None, +) -> Dict[str, Any]: + """ + Helper function to update a single user. + Used by both user_update and bulk_user_update endpoints. + + Returns the updated user data or raises an exception on failure. + """ + from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client + + if prisma_client is None: + raise Exception("Not connected to DB!") + + # Validate user identifier + if not user_request.user_id and not user_request.user_email: + raise ValueError("Either user_id or user_email must be provided") + + # Convert to data format expected by update logic + data_json: dict = user_request.model_dump(exclude_unset=True) + + # Apply update transformations (reuse existing logic) + non_default_values = _update_internal_user_params( + data_json=data_json, data=user_request + ) + + # Get existing user data for audit logging and metadata preparation + existing_user_row: Optional[BaseModel] = None + if user_request.user_id: + existing_user_row = await prisma_client.db.litellm_usertable.find_first( + where={"user_id": user_request.user_id} + ) + elif user_request.user_email: + existing_user_row = await prisma_client.db.litellm_usertable.find_first( + where={"user_email": user_request.user_email} + ) + + if existing_user_row is not None: + existing_user_row = LiteLLM_UserTable( + **existing_user_row.model_dump(exclude_none=True) + ) + + existing_metadata = ( + cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) + if existing_user_row is not None + else {} + ) + + non_default_values = prepare_metadata_fields( + data=user_request, + non_default_values=non_default_values, + existing_metadata=existing_metadata or {}, + ) + + # Perform the update + response: Optional[Dict[str, Any]] = None + + if user_request.user_id and len(user_request.user_id) > 0: + non_default_values["user_id"] = user_request.user_id + response = await prisma_client.update_data( + user_id=user_request.user_id, + data=non_default_values, + table_name="user", + ) + elif user_request.user_email: + # Handle email-based updates + existing_user_rows = await prisma_client.get_data( + key_val={"user_email": user_request.user_email}, + table_name="user", + query_type="find_all", + ) + + if ( + existing_user_rows + and isinstance(existing_user_rows, list) + and len(existing_user_rows) > 0 + ): + for existing_user in existing_user_rows: + non_default_values["user_id"] = existing_user.user_id + response = await prisma_client.update_data( + user_id=existing_user.user_id, + data=non_default_values, + table_name="user", + ) + break # Update first matching user + else: + # Create new user if not found + non_default_values["user_id"] = str(uuid.uuid4()) + non_default_values["user_email"] = user_request.user_email + response = await prisma_client.insert_data( + data=non_default_values, table_name="user" + ) + + # Create audit log for successful update + if response is not None: + try: + updated_user_row = await prisma_client.db.litellm_usertable.find_first( + where={"user_id": response["user_id"]} + ) + + if updated_user_row: + user_row_typed = LiteLLM_UserTable( + **updated_user_row.model_dump(exclude_none=True) + ) + + # Create audit log asynchronously + asyncio.create_task( + UserManagementEventHooks.create_internal_user_audit_log( + user_id=user_row_typed.user_id, + action="updated", + litellm_changed_by=litellm_changed_by + or user_api_key_dict.user_id, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + before_value=( + existing_user_row.model_dump_json(exclude_none=True) + if existing_user_row + else None + ), + after_value=user_row_typed.model_dump_json(exclude_none=True), + ) + ) + except Exception as audit_error: + verbose_proxy_logger.warning( + f"Failed to create audit log for user {response.get('user_id')}: {audit_error}" + ) + + if response is None: + raise HTTPException( + status_code=400, + detail={"error": "Failed to update user"}, + ) + return response + + @router.post( "/user/update", tags=["Internal User management"], @@ -842,116 +982,13 @@ async def user_update( - object_permission: Optional[LiteLLM_ObjectPermissionBase] - internal user-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission. """ - from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client - try: - data_json: dict = data.model_dump(exclude_unset=True) - # get the row from db - if prisma_client is None: - raise Exception("Not connected to DB!") - - # get non default values for key - non_default_values = _update_internal_user_params( - data_json=data_json, data=data - ) - - existing_user_row: Optional[BaseModel] = None - if data.user_id is not None: - existing_user_row = await prisma_client.db.litellm_usertable.find_first( - where={"user_id": data.user_id} - ) - if existing_user_row is not None: - existing_user_row = LiteLLM_UserTable( - **existing_user_row.model_dump(exclude_none=True) - ) - - existing_metadata = ( - cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) - if existing_user_row is not None - else {} - ) - - non_default_values = prepare_metadata_fields( - data=data, - non_default_values=non_default_values, - existing_metadata=existing_metadata or {}, - ) - - ## ADD USER, IF NEW ## verbose_proxy_logger.debug("/user/update: Received data = %s", data) - response: Optional[Any] = None - if data.user_id is not None and len(data.user_id) > 0: - non_default_values["user_id"] = data.user_id # type: ignore - verbose_proxy_logger.debug("In update user, user_id condition block.") - response = await prisma_client.update_data( - user_id=data.user_id, - data=non_default_values, - table_name="user", - ) - verbose_proxy_logger.debug( - f"received response from updating prisma client. response={response}" - ) - elif data.user_email is not None: - non_default_values["user_id"] = str(uuid.uuid4()) - non_default_values["user_email"] = data.user_email - ## user email is not unique acc. to prisma schema -> future improvement - ### for now: check if it exists in db, if not - insert it - existing_user_rows = await prisma_client.get_data( - key_val={"user_email": data.user_email}, - table_name="user", - query_type="find_all", - ) - if existing_user_rows is None or ( - isinstance(existing_user_rows, list) and len(existing_user_rows) == 0 - ): - response = await prisma_client.insert_data( - data=non_default_values, table_name="user" - ) - elif isinstance(existing_user_rows, list) and len(existing_user_rows) > 0: - for existing_user in existing_user_rows: - response = await prisma_client.update_data( - user_id=existing_user.user_id, - data=non_default_values, - table_name="user", - ) - - if response is not None: # emit audit log - try: - user_row: BaseModel = ( - await prisma_client.db.litellm_usertable.find_first( - where={"user_id": response["user_id"]} - ) - ) - - user_row_litellm_typed = LiteLLM_UserTable( - **user_row.model_dump(exclude_none=True) - ) - - asyncio.create_task( - UserManagementEventHooks.create_internal_user_audit_log( - user_id=user_row_litellm_typed.user_id, - action="updated", - litellm_changed_by=user_api_key_dict.user_id, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - before_value=( - existing_user_row.model_dump_json(exclude_none=True) - if existing_user_row - else None - ), - after_value=user_row_litellm_typed.model_dump_json( - exclude_none=True - ), - ) - ) - except Exception as e: - verbose_proxy_logger.warning( - "Unable to create audit log for user on `/user/update` - {}".format( - str(e) - ) - ) - return response # type: ignore - # update based on remaining passed in values + response = await _update_single_user_helper( + user_request=data, + user_api_key_dict=user_api_key_dict, + ) + return response except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.user_update(): Exception occured - {}".format( @@ -976,6 +1013,131 @@ async def user_update( ) +@router.post( + "/user/bulk_update", + tags=["Internal User management"], + dependencies=[Depends(user_api_key_auth)], + response_model=BulkUpdateUserResponse, +) +@management_endpoint_wrapper +async def bulk_user_update( + data: BulkUpdateUserRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + litellm_changed_by: Optional[str] = Header( + None, + description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability", + ), +): + """ + Bulk update multiple users at once. + + This endpoint allows updating multiple users in a single request. Each user update + is processed independently - if some updates fail, others will still succeed. + + Parameters: + - users: List[UpdateUserRequest] - List of user update requests + + Returns: + - results: List of individual update results + - total_requested: Total number of users requested for update + - successful_updates: Number of successful updates + - failed_updates: Number of failed updates + + Example request: + ```bash + curl --location 'http://0.0.0.0:4000/user/bulk_update' \ + --header 'Authorization: Bearer sk-1234' \ + --header 'Content-Type: application/json' \ + --data '{ + "users": [ + { + "user_id": "user1", + "user_role": "internal_user", + "max_budget": 100.0 + }, + { + "user_email": "user2@example.com", + "user_role": "internal_user_viewer", + "max_budget": 50.0 + } + ] + }' + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": "Database not connected"}, + ) + + if not data.users: + raise HTTPException( + status_code=400, + detail={"error": "At least one user update request is required"}, + ) + + # Limit batch size to prevent overwhelming the system + MAX_BATCH_SIZE = 100 + if len(data.users) > MAX_BATCH_SIZE: + raise HTTPException( + status_code=400, + detail={"error": f"Maximum {MAX_BATCH_SIZE} users can be updated at once"}, + ) + + results: List[UserUpdateResult] = [] + successful_updates = 0 + failed_updates = 0 + + # Process each user update independently + for user_request in data.users: + try: + response = await _update_single_user_helper( + user_request=user_request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=litellm_changed_by, + ) + # Record success + results.append( + UserUpdateResult( + user_id=( + response.get("user_id") if response else user_request.user_id + ), + user_email=user_request.user_email, + success=True, + updated_user=response, + ) + ) + successful_updates += 1 + except Exception as e: + verbose_proxy_logger.exception( + f"Failed to update user {user_request.user_id or user_request.user_email}: {e}" + ) + # Record failure + error_message = str(e) + verbose_proxy_logger.error( + f"Failed to update user {user_request.user_id or user_request.user_email}: {error_message}" + ) + + results.append( + UserUpdateResult( + user_id=user_request.user_id, + user_email=user_request.user_email, + success=False, + error=error_message, + ) + ) + failed_updates += 1 + + return BulkUpdateUserResponse( + results=results, + total_requested=len(data.users), + successful_updates=successful_updates, + failed_updates=failed_updates, + ) + + async def get_user_key_counts( prisma_client, user_ids: Optional[List[str]] = None, diff --git a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py index 5c2c5bf371..a3a68a04c9 100644 --- a/litellm/types/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/types/proxy/management_endpoints/internal_user_endpoints.py @@ -3,7 +3,7 @@ from typing import Any, Dict, List, Literal, Optional, Union from fastapi import HTTPException from pydantic import BaseModel, EmailStr -from litellm.proxy._types import LiteLLM_UserTableWithKeyCount +from litellm.proxy._types import LiteLLM_UserTableWithKeyCount, UpdateUserRequest class UserListResponse(BaseModel): @@ -16,3 +16,28 @@ class UserListResponse(BaseModel): page: int page_size: int total_pages: int + + +class BulkUpdateUserRequest(BaseModel): + """Request for bulk user updates""" + + users: List[UpdateUserRequest] # List of user update requests + + +class UserUpdateResult(BaseModel): + """Result of a single user update operation""" + + user_id: Optional[str] = None + user_email: Optional[str] = None + success: bool + error: Optional[str] = None + updated_user: Optional[Dict[str, Any]] = None + + +class BulkUpdateUserResponse(BaseModel): + """Response for bulk user update operations""" + + results: List[UserUpdateResult] + total_requested: int + successful_updates: int + failed_updates: int