diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 989fad7cd0b..abdbb317891 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -40,6 +40,7 @@ ) from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_permissions_caller_permission, + _enforce_email_prefix_on_key_alias, generate_key_helper_fn, prepare_metadata_fields, ) @@ -98,25 +99,41 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d auto_create_key = data_json.pop("auto_create_key", True) if auto_create_key is False: - data_json["table_name"] = "user" # only create a user, don't create key if 'auto_create_key' set to False + data_json["table_name"] = ( + "user" # only create a user, don't create key if 'auto_create_key' set to False + ) if litellm.default_internal_user_params and ( - data.user_role != LitellmUserRoles.PROXY_ADMIN.value and data.user_role != LitellmUserRoles.PROXY_ADMIN + data.user_role != LitellmUserRoles.PROXY_ADMIN.value + and data.user_role != LitellmUserRoles.PROXY_ADMIN ): for key, value in litellm.default_internal_user_params.items(): if key == "available_teams": continue elif key not in data_json or data_json[key] is None: data_json[key] = value - elif key == "models" and isinstance(data_json[key], list) and len(data_json[key]) == 0: + elif ( + key == "models" + and isinstance(data_json[key], list) + and len(data_json[key]) == 0 + ): data_json[key] = value ## INTERNAL USER ROLE ONLY DEFAULT PARAMS ## - if data.user_role is not None and data.user_role == LitellmUserRoles.INTERNAL_USER.value: - if litellm.max_internal_user_budget is not None and data_json.get("max_budget") is None: + if ( + data.user_role is not None + and data.user_role == LitellmUserRoles.INTERNAL_USER.value + ): + if ( + litellm.max_internal_user_budget is not None + and data_json.get("max_budget") is None + ): data_json["max_budget"] = litellm.max_internal_user_budget - if litellm.internal_user_budget_duration is not None and data_json.get("budget_duration") is None: + if ( + litellm.internal_user_budget_duration is not None + and data_json.get("budget_duration") is None + ): data_json["budget_duration"] = litellm.internal_user_budget_duration data_json.pop("teams", None) # handled separately @@ -154,18 +171,24 @@ async def _check_duplicate_user_field( if case_insensitive: where_clause[field_name]["mode"] = "insensitive" - existing_user = await UserRepository(prisma_client).table.find_first(where=where_clause) + existing_user = await UserRepository(prisma_client).table.find_first( + where=where_clause + ) if existing_user is not None: existing_value = getattr(existing_user, field_name, value) error_label = label or field_name raise HTTPException( status_code=409, - detail={"error": f"User with {error_label} {existing_value} already exists"}, + detail={ + "error": f"User with {error_label} {existing_value} already exists" + }, ) -async def _check_duplicate_user_email(user_email: Optional[str], prisma_client: Any) -> None: +async def _check_duplicate_user_email( + user_email: Optional[str], prisma_client: Any +) -> None: """ Helper function to check if a user email already exists in the database. """ @@ -249,7 +272,9 @@ async def _add_user_to_team( user_api_key_dict=user_api_key_dict, ) except HTTPException as e: - if e.status_code == 400 and ("already exists" in str(e) or "doesn't exist" in str(e)): + if e.status_code == 400 and ( + "already exists" in str(e) or "doesn't exist" in str(e) + ): verbose_proxy_logger.debug( "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( str(e) @@ -257,7 +282,7 @@ async def _add_user_to_team( ) else: verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): Exception occured - {}".format( + "litellm.proxy.management_endpoints..new_user(): Exception occured - {}".format( str(e) ) ) @@ -268,7 +293,10 @@ async def _add_user_to_team( str(e) ) ) - elif isinstance(e, ProxyException) and ProxyErrorTypes.team_member_already_in_team in e.type: + elif ( + isinstance(e, ProxyException) + and ProxyErrorTypes.team_member_already_in_team in e.type + ): verbose_proxy_logger.debug( "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( str(e) @@ -410,7 +438,9 @@ async def new_user( from litellm.proxy.proxy_server import _license_check, prisma_client if prisma_client is None: - raise HTTPException(status_code=400, detail=CommonProxyErrors.db_not_connected_error.value) + raise HTTPException( + status_code=400, detail=CommonProxyErrors.db_not_connected_error.value + ) if prisma_client is None: raise HTTPException( @@ -433,7 +463,8 @@ async def new_user( # Check if user_api_key_dict is actually a UserAPIKeyAuth instance (not a Depends object) # This can happen when the function is called directly in tests if ( - data.user_role in [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] + data.user_role + in [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] and isinstance(user_api_key_dict, UserAPIKeyAuth) and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN ): @@ -449,11 +480,19 @@ async def new_user( data_json = data.json() # type: ignore data_json = _update_internal_new_user_params(data_json, data) + data_json["key_alias"] = await _enforce_email_prefix_on_key_alias( + prisma_client=prisma_client, + key_alias=cast(Optional[str], data_json.get("key_alias")), + owner_email=data.user_email, + fallback_email=user_api_key_dict.user_email, + ) _hash_password_in_dict(data_json) teams = data.teams if teams is None: teams = check_if_default_team_set() - organization_ids = cast(Optional[List[str]], data_json.pop("organizations", None)) + organization_ids = cast( + Optional[List[str]], data_json.pop("organizations", None) + ) response = await generate_key_helper_fn(request_type="user", **data_json) # Admin UI Logic @@ -514,7 +553,9 @@ async def new_user( return new_user_response except Exception as e: - verbose_proxy_logger.exception("/user/new: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception( + "/user/new: Exception occured - {}".format(str(e)) + ) raise handle_exception_on_proxy(e) @@ -599,14 +640,18 @@ def get_user_id_from_request(request: Request) -> Optional[str]: return user_id -def _normalize_user_info_user_id(request: Request, user_id: Optional[str]) -> Optional[str]: +def _normalize_user_info_user_id( + request: Request, user_id: Optional[str] +) -> Optional[str]: """Normalize URL-decoded user_id while preserving '+' characters.""" if user_id is not None and " " in user_id: return get_user_id_from_request(request=request) return user_id -def _enforce_user_info_access(user_id: Optional[str], user_api_key_dict: UserAPIKeyAuth) -> None: +def _enforce_user_info_access( + user_id: Optional[str], user_api_key_dict: UserAPIKeyAuth +) -> None: """Re-validate that the caller may read the resolved ``user_id`` after URL-decoding has been finalized. @@ -671,7 +716,9 @@ async def _get_user_info_teams( query_type="find_all", ) elif user_api_key_dict.user_id is not None and user_id is None: - caller_user_info = await prisma_client.get_data(user_id=user_api_key_dict.user_id) + caller_user_info = await prisma_client.get_data( + user_id=user_api_key_dict.user_id + ) caller_team_ids = getattr(caller_user_info, "teams", None) if caller_team_ids: teams_2 = await prisma_client.get_data( @@ -715,10 +762,14 @@ def _build_user_info_response( returned_keys = _process_keys_for_user_info(keys=keys, all_teams=teams_1) team_list.sort(key=lambda x: getattr(x, "team_alias", "") or "") - _user_info = user_info.model_dump() if isinstance(user_info, BaseModel) else user_info + _user_info = ( + user_info.model_dump() if isinstance(user_info, BaseModel) else user_info + ) if isinstance(_user_info, dict): _user_info.pop("password", None) - _user_info["metadata"] = _redact_scim_enterprise_metadata(_user_info.get("metadata")) + _user_info["metadata"] = _redact_scim_enterprise_metadata( + _user_info.get("metadata") + ) return UserInfoResponse( user_id=user_id, @@ -737,7 +788,9 @@ def _build_user_info_response( @management_endpoint_wrapper async def user_info( request: Request, - user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"), + user_id: Optional[str] = fastapi.Query( + default=None, description="User ID in the request parameters" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -763,8 +816,13 @@ async def user_info( raise Exception( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if user_id is None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: - return await _get_user_info_for_proxy_admin(user_api_key_dict=user_api_key_dict) + if ( + user_id is None + and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + ): + return await _get_user_info_for_proxy_admin( + user_api_key_dict=user_api_key_dict + ) elif user_id is None: user_id = user_api_key_dict.user_id ## GET USER ROW ## @@ -803,7 +861,11 @@ async def user_info( return response_data except Exception as e: - verbose_proxy_logger.exception("litellm.proxy.proxy_server.user_info(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.user_info(): Exception occured - {}".format( + str(e) + ) + ) raise handle_exception_on_proxy(e) @@ -831,7 +893,9 @@ async def _check_user_info_v2_access( # Helper: fetch the target user row (reused across branches) async def _fetch_target_user(): - return await UserRepository(prisma_client).table.find_unique(where={"user_id": target_user_id}) + return await UserRepository(prisma_client).table.find_unique( + where={"user_id": target_user_id} + ) # Rule 1: Proxy admins — fetch and return the target row directly if _user_has_admin_view(user_api_key_dict): @@ -854,10 +918,14 @@ async def _fetch_target_user(): return None # Get all teams the caller belongs to - teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.teams}}) + teams = await TeamRepository(prisma_client).table.find_many( + where={"team_id": {"in": caller_user.teams}} + ) for team in teams: team_obj = LiteLLM_TeamTable(**team.model_dump()) - if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): + if _is_user_team_admin( + user_api_key_dict=user_api_key_dict, team_obj=team_obj + ): # Check if target user is in this team if team.team_id in (target_user.teams or []): return target_user @@ -874,7 +942,9 @@ async def _fetch_target_user(): @management_endpoint_wrapper async def user_info_v2( request: Request, - user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"), + user_id: Optional[str] = fastapi.Query( + default=None, description="User ID in the request parameters" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -952,7 +1022,9 @@ async def user_info_v2( ) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.user_info_v2(): Exception occured - {}".format(str(e)) + "litellm.proxy.proxy_server.user_info_v2(): Exception occured - {}".format( + str(e) + ) ) raise handle_exception_on_proxy(e) @@ -1006,7 +1078,9 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth): admin_user_info = await prisma_client.get_data(user_id=admin_user_id) if admin_user_info is not None: admin_user_info = ( - admin_user_info.model_dump() if isinstance(admin_user_info, BaseModel) else admin_user_info + admin_user_info.model_dump() + if isinstance(admin_user_info, BaseModel) + else admin_user_info ) if isinstance(admin_user_info, dict): admin_user_info.pop("password", None) @@ -1048,8 +1122,14 @@ def _process_keys_for_user_info( if _key.get("team_id") == UI_SESSION_TOKEN_TEAM_ID: continue - if "team_id" in _key and _key["team_id"] is not None and _key["team_id"] != "litellm-dashboard": - team_info = get_team_from_list(team_list=all_teams, team_id=_key["team_id"]) + if ( + "team_id" in _key + and _key["team_id"] is not None + and _key["team_id"] != "litellm-dashboard" + ): + team_info = get_team_from_list( + team_list=all_teams, team_id=_key["team_id"] + ) if team_info is not None: team_alias = getattr(team_info, "team_alias", None) _key["team_alias"] = team_alias @@ -1099,9 +1179,13 @@ def _update_internal_user_params( ): # applies internal user limits, if user role updated non_default_values["max_budget"] = litellm.max_internal_user_budget - if "budget_duration" not in non_default_values: # applies internal user limits, if user role updated + if ( + "budget_duration" not in non_default_values + ): # applies internal user limits, if user role updated if is_internal_user and litellm.internal_user_budget_duration is not None: - non_default_values["budget_duration"] = litellm.internal_user_budget_duration + non_default_values["budget_duration"] = ( + litellm.internal_user_budget_duration + ) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time non_default_values["budget_reset_at"] = get_budget_reset_time( @@ -1123,9 +1207,13 @@ async def _schedule_user_update_audit_log( if prisma_client is None: return try: - updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]}) + updated_user_row = await UserRepository(prisma_client).table.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)) + user_row_typed = LiteLLM_UserTable( + **updated_user_row.model_dump(exclude_none=True) + ) asyncio.create_task( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_row_typed.user_id, @@ -1133,12 +1221,18 @@ async def _schedule_user_update_audit_log( 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), + 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}") + verbose_proxy_logger.warning( + f"Failed to create audit log for user {response.get('user_id')}: {audit_error}" + ) def _check_user_update_authz( @@ -1147,12 +1241,19 @@ def _check_user_update_authz( existing_user_row: Optional[BaseModel], ) -> None: """Authorization checks for /user/update — raises HTTPException on failure.""" - if user_request.user_role is not None and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: - raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.") + if ( + user_request.user_role is not None + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + ): + raise HTTPException( + status_code=403, detail="Only proxy admins can modify user roles." + ) if existing_user_row is not None: typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) - if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row): + if not can_user_call_user_update( + user_api_key_dict=user_api_key_dict, user_info=typed_row + ): raise HTTPException( status_code=403, detail={ @@ -1183,7 +1284,9 @@ async def _invalidate_user_spend_counter_if_changed( if non_default_values.get("spend") is not None: from litellm.proxy.proxy_server import _invalidate_spend_counter - await _invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") + await _invalidate_spend_counter( + counter_key=f"spend:user:{non_default_values['user_id']}" + ) async def _update_single_user_helper( @@ -1211,7 +1314,9 @@ async def _update_single_user_helper( ) data_json: dict = user_request.model_dump(exclude_unset=True) - non_default_values = _update_internal_user_params(data_json=data_json, data=user_request) + non_default_values = _update_internal_user_params( + data_json=data_json, data=user_request + ) _hash_password_in_dict(non_default_values) existing_user_row: Optional[BaseModel] = None @@ -1227,17 +1332,26 @@ async def _update_single_user_helper( _check_user_update_authz(user_request, user_api_key_dict, existing_user_row) if existing_user_row is not None: - existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + existing_user_row = LiteLLM_UserTable( + **existing_user_row.model_dump(exclude_none=True) + ) # Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers # must not be able to raise their own budget/spend fields. # can_user_call_user_update() already restricts non-admins to self-updates, # so this guard only fires for self-escalation attempts. _target_user_id = user_request.user_id or ( - getattr(existing_user_row, "user_id", None) if existing_user_row is not None else None + getattr(existing_user_row, "user_id", None) + if existing_user_row is not None + else None ) - _is_self_update = _target_user_id is not None and user_api_key_dict.user_id == _target_user_id - if _is_self_update and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + _is_self_update = ( + _target_user_id is not None and user_api_key_dict.user_id == _target_user_id + ) + if ( + _is_self_update + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + ): _protected_fields = ("max_budget", "soft_budget", "spend") for _field in _protected_fields: if _field in non_default_values: @@ -1249,7 +1363,9 @@ async def _update_single_user_helper( ) existing_metadata = ( - cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) if existing_user_row is not None else {} + cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) + if existing_user_row is not None + else {} ) non_default_values = prepare_metadata_fields( @@ -1279,7 +1395,11 @@ async def _update_single_user_helper( query_type="find_all", ) - if existing_user_rows and isinstance(existing_user_rows, list) and len(existing_user_rows) > 0: + 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( @@ -1292,7 +1412,9 @@ async def _update_single_user_helper( # 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") + response = await prisma_client.insert_data( + data=non_default_values, table_name="user" + ) if response is not None: await _schedule_user_update_audit_log( @@ -1400,7 +1522,9 @@ async def user_update( return response except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.user_update(): Exception occured - {}".format(str(e)) + "litellm.proxy.proxy_server.user_update(): Exception occured - {}".format( + str(e) + ) ) verbose_proxy_logger.debug(traceback.format_exc()) if isinstance(e, HTTPException): @@ -1441,7 +1565,11 @@ async def bulk_update_processed_users( # Record success results.append( UserUpdateResult( - user_id=(response.get("user_id") if response else user_request.user_id), + user_id=( + response.get("user_id") + if response + else user_request.user_id + ), user_email=user_request.user_email, success=True, updated_user=response, @@ -1555,17 +1683,26 @@ async def bulk_user_update( ) # Only proxy admins can modify user_role in bulk updates - _bulk_role = getattr(data.user_updates, "user_role", None) if data.user_updates else None + _bulk_role = ( + getattr(data.user_updates, "user_role", None) if data.user_updates else None + ) if _bulk_role is None and data.users: - _bulk_role = next((u.user_role for u in data.users if u.user_role is not None), None) - if _bulk_role is not None and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + _bulk_role = next( + (u.user_role for u in data.users if u.user_role is not None), None + ) + if ( + _bulk_role is not None + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + ): raise HTTPException( status_code=403, detail="Only proxy admins can modify user roles.", ) # Determine the list of users to update - users_to_update: Union[List[UpdateUserRequest], List[UpdateUserRequestNoUserIDorEmail]] = [] + users_to_update: Union[ + List[UpdateUserRequest], List[UpdateUserRequestNoUserIDorEmail] + ] = [] if data.all_users and data.user_updates: # Only proxy admins can update all users at once @@ -1575,7 +1712,9 @@ async def bulk_user_update( detail="Only proxy admins can update all users at once.", ) # Optimized path for updating all users directly in database - all_users_in_db = await UserRepository(prisma_client).table.find_many(order={"created_at": "desc"}) + all_users_in_db = await UserRepository(prisma_client).table.find_many( + order={"created_at": "desc"} + ) if not all_users_in_db: raise HTTPException( @@ -1595,7 +1734,9 @@ async def bulk_user_update( # Apply update transformations (reuse existing logic) data_json: dict = data.user_updates.model_dump(exclude_unset=True) - non_default_values = _update_internal_user_params(data_json=data_json, data=data.user_updates) + non_default_values = _update_internal_user_params( + data_json=data_json, data=data.user_updates + ) # Remove user identification fields since we're updating by user_id non_default_values.pop("user_id", None) @@ -1630,7 +1771,8 @@ async def bulk_user_update( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_api_key_dict.user_id or "", action="updated", - litellm_changed_by=litellm_changed_by or user_api_key_dict.user_id, + 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=f"Updated {len(all_users_in_db)} users", @@ -1638,7 +1780,9 @@ async def bulk_user_update( ) ) except Exception as audit_error: - verbose_proxy_logger.warning(f"Failed to create bulk audit log: {audit_error}") + verbose_proxy_logger.warning( + f"Failed to create bulk audit log: {audit_error}" + ) except Exception as e: verbose_proxy_logger.exception(f"Failed to perform bulk update: {e}") @@ -1726,7 +1870,9 @@ async def get_user_key_counts( return result -def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[Dict[str, str]]: +def _validate_sort_params( + sort_by: Optional[str], sort_order: str +) -> Optional[Dict[str, str]]: order_by: Dict[str, str] = {} if sort_by is None: @@ -1743,7 +1889,9 @@ def _validate_sort_params(sort_by: Optional[str], sort_order: str) -> Optional[D if sort_by not in valid_columns: raise HTTPException( status_code=400, - detail={"error": f"Invalid sort column. Must be one of: {', '.join(valid_columns)}"}, + detail={ + "error": f"Invalid sort column. Must be one of: {', '.join(valid_columns)}" + }, ) # Validate sort_order @@ -1778,7 +1926,9 @@ async def _authorize_user_list_request( if user_api_key_dict.user_id is None: raise HTTPException( status_code=403, - detail={"error": "Only proxy admins and organization admins can list users."}, + detail={ + "error": "Only proxy admins and organization admins can list users." + }, ) try: caller_user = await get_user_object( @@ -1791,12 +1941,16 @@ async def _authorize_user_list_request( except ValueError: raise HTTPException( status_code=403, - detail={"error": "Only proxy admins and organization admins can list users."}, + detail={ + "error": "Only proxy admins and organization admins can list users." + }, ) if caller_user is None: raise HTTPException( status_code=403, - detail={"error": "Only proxy admins and organization admins can list users."}, + detail={ + "error": "Only proxy admins and organization admins can list users." + }, ) allowed_org_ids = [ @@ -1807,17 +1961,23 @@ async def _authorize_user_list_request( if not allowed_org_ids: raise HTTPException( status_code=403, - detail={"error": "Only proxy admins and organization admins can list users."}, + detail={ + "error": "Only proxy admins and organization admins can list users." + }, ) # If client also sent organization_ids, intersect with allowed orgs if organization_ids: - requested = set(oid.strip() for oid in organization_ids.split(",") if oid.strip()) + requested = set( + oid.strip() for oid in organization_ids.split(",") if oid.strip() + ) intersection = list(requested & set(allowed_org_ids)) if not intersection: raise HTTPException( status_code=403, - detail={"error": "You do not have org_admin access to the requested organization(s)."}, + detail={ + "error": "You do not have org_admin access to the requested organization(s)." + }, ) allowed_org_ids = intersection @@ -1831,18 +1991,32 @@ async def _authorize_user_list_request( response_model=UserListResponse, ) async def get_users( - role: Optional[str] = fastapi.Query(default=None, description="Filter users by role"), - user_ids: Optional[str] = fastapi.Query(default=None, description="Get list of users by user_ids"), - sso_user_ids: Optional[str] = fastapi.Query(default=None, description="Get list of users by sso_user_id"), - user_email: Optional[str] = fastapi.Query(default=None, description="Filter users by partial email match"), - team: Optional[str] = fastapi.Query(default=None, description="Filter users by team id"), + role: Optional[str] = fastapi.Query( + default=None, description="Filter users by role" + ), + user_ids: Optional[str] = fastapi.Query( + default=None, description="Get list of users by user_ids" + ), + sso_user_ids: Optional[str] = fastapi.Query( + default=None, description="Get list of users by sso_user_id" + ), + user_email: Optional[str] = fastapi.Query( + default=None, description="Filter users by partial email match" + ), + team: Optional[str] = fastapi.Query( + default=None, description="Filter users by team id" + ), page: int = fastapi.Query(default=1, ge=1, description="Page number"), - page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"), + page_size: int = fastapi.Query( + default=25, ge=1, le=100, description="Number of items per page" + ), sort_by: Optional[str] = fastapi.Query( default=None, description="Column to sort by (e.g. 'user_id', 'user_email', 'created_at', 'spend')", ), - sort_order: str = fastapi.Query(default="asc", description="Sort order ('asc' or 'desc')"), + sort_order: str = fastapi.Query( + default="asc", description="Sort order ('asc' or 'desc')" + ), organization_ids: Optional[str] = fastapi.Query( default=None, description="Filter users by organization membership. Comma-separated list of org IDs.", @@ -1936,9 +2110,13 @@ async def get_users( } if organization_ids: - org_id_list = [oid.strip() for oid in organization_ids.split(",") if oid.strip()] + org_id_list = [ + oid.strip() for oid in organization_ids.split(",") if oid.strip() + ] if org_id_list: - where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_id_list}}} + where_conditions["organization_memberships"] = { + "some": {"organization_id": {"in": org_id_list}} + } ## Filter any none fastapi.Query params - e.g. where_conditions: {'user_email': {'contains': Query(None), 'mode': 'insensitive'}, 'teams': {'has': Query(None)}} where_conditions = {k: v for k, v in where_conditions.items() if v is not None} @@ -1946,22 +2124,30 @@ async def get_users( # Build order_by conditions order_by: Optional[Dict[str, str]] = ( - _validate_sort_params(sort_by, sort_order) if sort_by is not None and isinstance(sort_by, str) else None + _validate_sort_params(sort_by, sort_order) + if sort_by is not None and isinstance(sort_by, str) + else None ) users = await UserRepository(prisma_client).table.find_many( where=where_conditions, skip=skip, take=page_size, - order=(order_by if order_by else {"created_at": "desc"}), # Default to created_at desc if no sort specified + order=( + order_by if order_by else {"created_at": "desc"} + ), # Default to created_at desc if no sort specified ) # Get total count of user rows - total_count = await UserRepository(prisma_client).table.count(where=where_conditions) + total_count = await UserRepository(prisma_client).table.count( + where=where_conditions + ) # Get key count for each user if users is not None: - user_key_counts = await get_user_key_counts(prisma_client, [user.user_id for user in users]) + user_key_counts = await get_user_key_counts( + prisma_client, [user.user_id for user in users] + ) else: user_key_counts = {} @@ -1975,8 +2161,14 @@ async def get_users( if users is not None: for user in users: user_dump = user.model_dump() - user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) - user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0))) + user_dump["metadata"] = _redact_scim_enterprise_metadata( + user_dump.get("metadata") + ) + user_list.append( + LiteLLM_UserTableWithKeyCount( + **user_dump, key_count=user_key_counts.get(user.user_id, 0) + ) + ) else: user_list = [] @@ -2045,7 +2237,9 @@ async def delete_user( # cross-check data.user_ids against the caller's scope, so without this # loop an org-admin of org-A could delete users in org-B by supplying # {"user_ids": [victim_in_org_B], "organization_id": "org-A"}. - caller_is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + caller_is_proxy_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) caller_admin_org_ids: set = set() if not caller_is_proxy_admin: caller_memberships = ( @@ -2058,20 +2252,24 @@ async def delete_user( if user_api_key_dict.user_id else [] ) - caller_admin_org_ids = {m.organization_id for m in caller_memberships if m.organization_id} + caller_admin_org_ids = { + m.organization_id for m in caller_memberships if m.organization_id + } if not caller_admin_org_ids: raise HTTPException( status_code=403, - detail={"error": "Only PROXY_ADMIN or ORG_ADMIN users may delete users."}, + detail={ + "error": "Only PROXY_ADMIN or ORG_ADMIN users may delete users." + }, ) # Batch-fetch target memberships once before the per-user loop. Avoids # an N+1 DB call when delete_user is called with a large user_ids list. target_org_ids_by_user: Dict[str, set] = {} if not caller_is_proxy_admin: - all_target_memberships = await OrganizationMembershipRepository(prisma_client).table.find_many( - where={"user_id": {"in": data.user_ids}} - ) + all_target_memberships = await OrganizationMembershipRepository( + prisma_client + ).table.find_many(where={"user_id": {"in": data.user_ids}}) for m in all_target_memberships: if not m.organization_id: continue @@ -2079,7 +2277,9 @@ async def delete_user( # check that all teams passed exist for user_id in data.user_ids: - user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) + user_row = await UserRepository(prisma_client).table.find_unique( + where={"user_id": user_id} + ) if user_row is None: raise HTTPException( @@ -2131,7 +2331,9 @@ async def delete_user( ) ## CLEANUP MEMBERS_WITH_ROLES - fetch_all_teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_row.teams}}) + fetch_all_teams = await TeamRepository(prisma_client).table.find_many( + where={"team_id": {"in": user_row.teams}} + ) teams_to_update = [] for team in fetch_all_teams: is_member_in_team, new_team_members = _cleanup_members_with_roles( @@ -2143,7 +2345,9 @@ async def delete_user( ), ) if is_member_in_team: - _db_new_team_members: List[dict] = [m.model_dump() for m in new_team_members] + _db_new_team_members: List[dict] = [ + m.model_dump() for m in new_team_members + ] team.members_with_roles = json.dumps(_db_new_team_members) teams_to_update.append(team) @@ -2157,7 +2361,9 @@ async def delete_user( # End of Audit logging ## DELETE ASSOCIATED KEYS - await VerificationTokenRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}}) + await VerificationTokenRepository(prisma_client).table.delete_many( + where={"user_id": {"in": data.user_ids}} + ) ## DELETE ASSOCIATED INVITATION LINKS await InvitationLinkRepository(prisma_client).table.delete_many( @@ -2171,13 +2377,19 @@ async def delete_user( ) ## DELETE ASSOCIATED ORGANIZATION MEMBERSHIPS - await OrganizationMembershipRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}}) + await OrganizationMembershipRepository(prisma_client).table.delete_many( + where={"user_id": {"in": data.user_ids}} + ) ## DELETE ASSOCIATED TEAM MEMBERSHIPS - await TeamMembershipRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}}) + await TeamMembershipRepository(prisma_client).table.delete_many( + where={"user_id": {"in": data.user_ids}} + ) ## DELETE USERS - deleted_users = await UserRepository(prisma_client).table.delete_many(where={"user_id": {"in": data.user_ids}}) + deleted_users = await UserRepository(prisma_client).table.delete_many( + where={"user_id": {"in": data.user_ids}} + ) return deleted_users @@ -2205,14 +2417,18 @@ async def add_internal_user_to_organization( try: # Check if organization_id exists - organization_row = await OrganizationRepository(prisma_client).table.find_unique( - where={"organization_id": organization_id} - ) + organization_row = await OrganizationRepository( + prisma_client + ).table.find_unique(where={"organization_id": organization_id}) if organization_row is None: - raise Exception(f"Organization not found, passed organization_id={organization_id}") + raise Exception( + f"Organization not found, passed organization_id={organization_id}" + ) # Create a new organization membership entry - new_membership = await OrganizationMembershipRepository(prisma_client).table.create( + new_membership = await OrganizationMembershipRepository( + prisma_client + ).table.create( data={ "user_id": user_id, "organization_id": organization_id, @@ -2268,7 +2484,9 @@ async def _resolve_org_filter_for_user_search( # This allows team admins who are org members to search users in their org. member_org_ids: List[str] = [] if caller_user is not None: - member_org_ids = [m.organization_id for m in (caller_user.organization_memberships or [])] + member_org_ids = [ + m.organization_id for m in (caller_user.organization_memberships or []) + ] if member_org_ids: return member_org_ids @@ -2312,13 +2530,17 @@ async def _resolve_team_org_filter( except HTTPException: raise HTTPException( status_code=403, - detail={"error": f"scope_user_search_to_org is enabled but team '{team_id}' was not found."}, + detail={ + "error": f"scope_user_search_to_org is enabled but team '{team_id}' was not found." + }, ) if not _is_user_team_admin(user_api_key_dict, team_obj): raise HTTPException( status_code=403, - detail={"error": "scope_user_search_to_org is enabled. You must be an admin of this team to search users."}, + detail={ + "error": "scope_user_search_to_org is enabled. You must be an admin of this team to search users." + }, ) if team_obj.organization_id: @@ -2342,14 +2564,22 @@ async def _resolve_team_org_filter( }, ) async def ui_view_users( - user_id: Optional[str] = fastapi.Query(default=None, description="User ID in the request parameters"), - user_email: Optional[str] = fastapi.Query(default=None, description="User email in the request parameters"), + user_id: Optional[str] = fastapi.Query( + default=None, description="User ID in the request parameters" + ), + user_email: Optional[str] = fastapi.Query( + default=None, description="User email in the request parameters" + ), team_id: Optional[str] = fastapi.Query( default=None, description="Team ID — used when a team admin searches for users to add to their team", ), - page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1), - page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=100), + page: int = fastapi.Query( + default=1, description="Page number for pagination", ge=1 + ), + page_size: int = fastapi.Query( + default=50, description="Number of items per page", ge=1, le=100 + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -2403,10 +2633,14 @@ async def ui_view_users( # Apply org filter when scope_user_search_to_org is ON and caller is not proxy admin if org_filter_ids is not None: - where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_filter_ids}}} + where_conditions["organization_memberships"] = { + "some": {"organization_id": {"in": org_filter_ids}} + } # Query users with pagination and filters - users: Optional[List[BaseModel]] = await UserRepository(prisma_client).table.find_many( + users: Optional[List[BaseModel]] = await UserRepository( + prisma_client + ).table.find_many( where=where_conditions, skip=skip, take=page_size, @@ -2428,14 +2662,23 @@ async def ui_view_users( # Using shared metric helper implementations from common_daily_activity -async def _resolve_user_email_metadata(prisma_client: "PrismaClient", records: list[Any]) -> dict[str, dict]: +async def _resolve_user_email_metadata( + prisma_client: "PrismaClient", records: list[Any] +) -> dict[str, dict]: """Map each user_id on the page to its email/alias so the Usage dashboard can label the 'Spend Per User' chart with the email instead of the raw UUID.""" - user_ids = {record.user_id for record in records if getattr(record, "user_id", None)} + user_ids = { + record.user_id for record in records if getattr(record, "user_id", None) + } if not user_ids: return {} - users = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": list(user_ids)}}) - return {user.user_id: {"user_email": user.user_email, "user_alias": user.user_alias} for user in users} + users = await UserRepository(prisma_client).table.find_many( + where={"user_id": {"in": list(user_ids)}} + ) + return { + user.user_id: {"user_email": user.user_email, "user_alias": user.user_alias} + for user in users + } @router.get( @@ -2466,8 +2709,12 @@ async def get_user_daily_activity( default=None, description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.", ), - page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1), - page_size: int = fastapi.Query(default=50, description="Items per page", ge=1, le=1000), + page: int = fastapi.Query( + default=1, description="Page number for pagination", ge=1 + ), + page_size: int = fastapi.Query( + default=50, description="Items per page", ge=1, le=1000 + ), timezone: Optional[int] = fastapi.Query( default=None, description="Timezone offset in minutes from UTC (e.g., 480 for PST). " @@ -2517,7 +2764,9 @@ async def get_user_daily_activity( if user_id != caller_user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Non-admin users can only view their own spend data."}, + detail={ + "error": "Non-admin users can only view their own spend data." + }, ) entity_id = user_id @@ -2534,13 +2783,17 @@ async def get_user_daily_activity( page=page, page_size=page_size, timezone_offset_minutes=timezone, - resolve_entity_metadata=lambda records: _resolve_user_email_metadata(prisma_client, records), + resolve_entity_metadata=lambda records: _resolve_user_email_metadata( + prisma_client, records + ), ) except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception("/spend/daily/analytics: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception( + "/spend/daily/analytics: Exception occured - {}".format(str(e)) + ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Failed to fetch analytics: {str(e)}"}, @@ -2612,7 +2865,9 @@ async def get_user_daily_activity_aggregated( if user_id != caller_user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Non-admin users can only view their own spend data."}, + detail={ + "error": "Non-admin users can only view their own spend data." + }, ) entity_id = user_id @@ -2632,7 +2887,9 @@ async def get_user_daily_activity_aggregated( except HTTPException: raise except Exception as e: - verbose_proxy_logger.exception("/user/daily/activity/aggregated: Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception( + "/user/daily/activity/aggregated: Exception occured - {}".format(str(e)) + ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Failed to fetch analytics: {str(e)}"}, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d92aea57063..cef6ef54d81 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -145,6 +145,74 @@ def _is_team_key(data: Union[GenerateKeyRequest, LiteLLM_VerificationToken]): return data.team_id is not None +def _get_email_local_part(email: str) -> str: + return email.split("@", maxsplit=1)[0].strip().lower() + + +def _normalize_key_alias_segment(value: str) -> str: + camel_spaced = re.sub(r"([a-z0-9])([A-Z])", r"\1-\2", value) + normalized = re.sub(r"[^a-z0-9]+", "-", camel_spaced.lower()) + return normalized.strip("-") + + +async def _get_user_email_by_id(prisma_client: Optional[PrismaClient], user_id: Optional[str]) -> Optional[str]: + if prisma_client is None or user_id is None: + return None + user = await UserRepository(prisma_client).find_by_id(user_id) + if user is None: + return None + user_email = user.user_email + if user_email is None: + return None + stripped_email = user_email.strip() + return stripped_email or None + + +async def _enforce_email_prefix_on_key_alias( + *, + prisma_client: Optional[PrismaClient], + key_alias: Optional[str], + owner_email: Optional[str], + fallback_email: Optional[str], +) -> Optional[str]: + if key_alias is None: + return None + ui_settings = await get_ui_settings_cached() + if ui_settings.get("enforce_email_prefix_on_key_alias") is not True: + return key_alias + resolved_email = next( + (email.strip() for email in (owner_email, fallback_email) if email is not None and email.strip()), + None, + ) + if resolved_email is None: + raise HTTPException( + status_code=400, + detail={ + "error": "key_alias requires a resolvable owner or caller email when enforce_email_prefix_on_key_alias is enabled." + }, + ) + prefix = _get_email_local_part(resolved_email) + if not prefix: + raise HTTPException( + status_code=400, + detail={ + "error": "key_alias requires a resolvable owner or caller email when enforce_email_prefix_on_key_alias is enabled." + }, + ) + normalized_alias = key_alias.strip() + prefix_with_dash = f"{prefix}-" + alias_suffix = ( + normalized_alias[len(prefix_with_dash) :] if normalized_alias.startswith(prefix_with_dash) else normalized_alias + ) + normalized_suffix = _normalize_key_alias_segment(alias_suffix) + if not normalized_suffix: + raise HTTPException( + status_code=400, + detail={"error": "key_alias must contain at least one alphanumeric character after normalization."}, + ) + return f"{prefix}-{normalized_suffix}" + + def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: Optional[str]) -> Optional[Member]: if user_id is None: return None @@ -994,6 +1062,17 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) + owner_email = await _get_user_email_by_id( + prisma_client=prisma_client, + user_id=data_json.get("user_id"), + ) + data_json["key_alias"] = await _enforce_email_prefix_on_key_alias( + prisma_client=prisma_client, + key_alias=data_json.get("key_alias", None), + owner_email=data_json.get("user_email") or owner_email, + fallback_email=user_api_key_dict.user_email, + ) + _validate_key_alias_format(key_alias=data_json.get("key_alias", None)) await _enforce_unique_key_alias( @@ -1859,6 +1938,8 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_ async def prepare_key_update_data( data: Union[UpdateKeyRequest, RegenerateKeyRequest], existing_key_row: LiteLLM_VerificationToken, + user_api_key_dict: Optional[UserAPIKeyAuth] = None, + prisma_client: Optional[PrismaClient] = None, ): data_json: dict = data.model_dump(exclude_unset=True) data_json.pop("key", None) @@ -1946,6 +2027,20 @@ async def prepare_key_update_data( data=data, non_default_values=non_default_values, existing_metadata=_metadata ) + if "key_alias" in non_default_values: + non_default_values["key_alias"] = await _enforce_email_prefix_on_key_alias( + prisma_client=prisma_client, + key_alias=cast(Optional[str], non_default_values.get("key_alias")), + owner_email=( + await _get_user_email_by_id( + prisma_client=prisma_client, + user_id=cast(Optional[str], non_default_values.get("user_id", existing_key_row.user_id)), + ) + ) + or getattr(existing_key_row, "user_email", None), + fallback_email=user_api_key_dict.user_email if user_api_key_dict is not None else None, + ) + return non_default_values @@ -2146,7 +2241,12 @@ async def _process_single_key_update( ) # Prepare update data - non_default_values = await prepare_key_update_data(data=update_key_request, existing_key_row=existing_key_row) + non_default_values = await prepare_key_update_data( + data=update_key_request, + existing_key_row=existing_key_row, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) # Update key in database if prisma_client is None: @@ -2588,7 +2688,12 @@ async def update_key_fn( # Enforce upperbound key params on update (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) - non_default_values = await prepare_key_update_data(data=data, existing_key_row=existing_key_row) + non_default_values = await prepare_key_update_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) # Only validate key_alias format if it's actually being changed new_key_alias = non_default_values.get("key_alias", None) @@ -4444,7 +4549,12 @@ async def _execute_virtual_key_regeneration( if data is not None: # Enforce upperbound key params on regenerate (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) - non_default_values = await prepare_key_update_data(data=data, existing_key_row=key_in_db) + non_default_values = await prepare_key_update_data( + data=data, + existing_key_row=key_in_db, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) # Only validate key_alias format if it's actually being changed new_key_alias = non_default_values.get("key_alias") if new_key_alias != key_in_db.key_alias: diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index a8926d26047..4337b2e5753 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -182,6 +182,11 @@ class UISettings(BaseModel): description="If true, shows the Chat page in the UI sidebar, letting users chat with an LLM and connect their own MCP server credentials via OAuth.", ) + enforce_email_prefix_on_key_alias: bool = Field( + default=False, + description="If true, key_alias values are normalized to kebab case and prefixed with the owner's email local-part.", + ) + class UISettingsResponse(SettingsResponse): """Response model for UI settings""" @@ -206,6 +211,7 @@ class UISettingsResponse(SettingsResponse): "disable_custom_api_keys", "disable_key_generate_for_org_admin", "enable_chat_ui", + "enforce_email_prefix_on_key_alias", } # Flags that must be synced from the persisted UISettings into @@ -219,6 +225,7 @@ class UISettingsResponse(SettingsResponse): "disable_vector_stores_for_internal_users", "allow_vector_stores_for_team_admins", "disable_key_generate_for_org_admin", + "enforce_email_prefix_on_key_alias", ] # Extension point: packages outside OSS (e.g. litellm_enterprise) can @@ -354,7 +361,9 @@ async def add_allowed_ip( if store_model_in_db is not True: raise HTTPException( status_code=500, - detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, ) # Load existing config @@ -562,7 +571,9 @@ async def get_default_team_settings(): ) -async def update_default_team_member_budget(teams: List[NewUserRequestTeam], user_api_key_dict: UserAPIKeyAuth): +async def update_default_team_member_budget( + teams: List[NewUserRequestTeam], user_api_key_dict: UserAPIKeyAuth +): """ 1. Update the max member budget for the team """ @@ -590,7 +601,9 @@ async def update_default_team_member_budget(teams: List[NewUserRequestTeam], use async def _update_litellm_setting( - settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings], + settings: Union[ + DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings + ], settings_key: str, success_message: str, user_api_key_dict: UserAPIKeyAuth, @@ -613,7 +626,9 @@ async def _update_litellm_setting( if store_model_in_db is not True: raise HTTPException( status_code=500, - detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, ) in_memory_var = settings.model_dump(exclude_none=True) @@ -670,7 +685,9 @@ async def update_internal_user_settings( Update the default internal user parameters for SSO users. These settings will be applied to new users who sign in via SSO. """ - if settings.teams is not None and all(isinstance(team, NewUserRequestTeam) for team in settings.teams): + if settings.teams is not None and all( + isinstance(team, NewUserRequestTeam) for team in settings.teams + ): await update_default_team_member_budget( settings.teams, user_api_key_dict=user_api_key_dict, # type: ignore @@ -726,7 +743,9 @@ async def get_sso_settings(): ) # Get SSO config from dedicated table - sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) + sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique( + where={"id": "sso_config"} + ) # Initialize with defaults sso_settings_dict = {} @@ -763,15 +782,29 @@ async def get_sso_settings(): sso_config = SSOConfig( google_client_id=decrypted_sso_settings_dict.get("google_client_id", None), - google_client_secret=decrypted_sso_settings_dict.get("google_client_secret", None), - microsoft_client_id=decrypted_sso_settings_dict.get("microsoft_client_id", None), - microsoft_client_secret=decrypted_sso_settings_dict.get("microsoft_client_secret", None), + google_client_secret=decrypted_sso_settings_dict.get( + "google_client_secret", None + ), + microsoft_client_id=decrypted_sso_settings_dict.get( + "microsoft_client_id", None + ), + microsoft_client_secret=decrypted_sso_settings_dict.get( + "microsoft_client_secret", None + ), microsoft_tenant=decrypted_sso_settings_dict.get("microsoft_tenant", None), generic_client_id=decrypted_sso_settings_dict.get("generic_client_id", None), - generic_client_secret=decrypted_sso_settings_dict.get("generic_client_secret", None), - generic_authorization_endpoint=decrypted_sso_settings_dict.get("generic_authorization_endpoint", None), - generic_token_endpoint=decrypted_sso_settings_dict.get("generic_token_endpoint", None), - generic_userinfo_endpoint=decrypted_sso_settings_dict.get("generic_userinfo_endpoint", None), + generic_client_secret=decrypted_sso_settings_dict.get( + "generic_client_secret", None + ), + generic_authorization_endpoint=decrypted_sso_settings_dict.get( + "generic_authorization_endpoint", None + ), + generic_token_endpoint=decrypted_sso_settings_dict.get( + "generic_token_endpoint", None + ), + generic_userinfo_endpoint=decrypted_sso_settings_dict.get( + "generic_userinfo_endpoint", None + ), proxy_base_url=decrypted_sso_settings_dict.get("proxy_base_url", None), user_email=decrypted_sso_settings_dict.get("user_email"), ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"), @@ -838,7 +871,9 @@ async def update_sso_settings( if store_model_in_db is not True: raise HTTPException( status_code=500, - detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, ) # Update environment variables @@ -861,7 +896,9 @@ async def update_sso_settings( # before-snapshot has the same shape as after_value, and rely on # create_config_audit_log's secret-name redaction to mask the # *_client_secret fields before the audit row is written. - existing_sso_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) + existing_sso_record = await SSOConfigRepository(prisma_client).table.find_unique( + where={"id": "sso_config"} + ) before_sso_data: Optional[Dict[str, Any]] = None if existing_sso_record and existing_sso_record.sso_settings: stored = existing_sso_record.sso_settings @@ -892,7 +929,9 @@ async def update_sso_settings( # Clear environment variable if value is null/empty os.environ.pop(env_var_name, None) - encrypted_sso_data = proxy_config._encrypt_env_variables(environment_variables=sso_data) + encrypted_sso_data = proxy_config._encrypt_env_variables( + environment_variables=sso_data + ) # Save to dedicated SSO table await SSOConfigRepository(prisma_client).table.upsert( @@ -937,7 +976,9 @@ async def update_sso_settings( env_vars_to_remove = set(env_var_mapping.values()) filtered_env_vars = { - key: value for key, value in environment_variables.items() if key not in env_vars_to_remove + key: value + for key, value in environment_variables.items() + if key not in env_vars_to_remove } await ConfigRepository(prisma_client).table.update( @@ -1034,7 +1075,9 @@ async def update_ui_theme_settings( if store_model_in_db is not True: raise HTTPException( status_code=500, - detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, ) # Load existing config @@ -1097,7 +1140,10 @@ async def update_ui_theme_settings( # Handle environment variable encryption if needed stored_config = config.copy() - if "environment_variables" in stored_config and len(stored_config["environment_variables"]) > 0: + if ( + "environment_variables" in stored_config + and len(stored_config["environment_variables"]) > 0 + ): # Only encrypt if there are environment variables to encrypt stored_config["environment_variables"] = proxy_config._encrypt_env_variables( environment_variables=stored_config["environment_variables"] @@ -1182,9 +1228,13 @@ async def update_mcp_semantic_filter_settings( from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is not None: - await proxy_config._init_semantic_filter_settings_in_db(prisma_client=prisma_client) + await proxy_config._init_semantic_filter_settings_in_db( + prisma_client=prisma_client + ) except Exception as e: - verbose_proxy_logger.warning(f"Failed to reinitialize MCP semantic filter settings immediately: {e}") + verbose_proxy_logger.warning( + f"Failed to reinitialize MCP semantic filter settings immediately: {e}" + ) return result @@ -1211,17 +1261,23 @@ async def get_ui_settings_cached() -> Dict[str, Any]: if prisma_client is None: return {} - db_record = await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}) + db_record = await UISettingsRepository(prisma_client).table.find_unique( + where={"id": "ui_settings"} + ) ui_settings: Dict[str, Any] = {} if db_record and db_record.ui_settings: raw = db_record.ui_settings ui_settings = json.loads(raw) if isinstance(raw, str) else dict(raw) # Sanitize - ui_settings = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS} + ui_settings = { + k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS + } # 3. Populate cache with TTL - await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) + await user_api_key_cache.async_set_cache( + key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL + ) return ui_settings @@ -1246,7 +1302,9 @@ async def get_ui_settings(): ui_settings: Dict[str, Any] = {} - db_record = await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}) + db_record = await UISettingsRepository(prisma_client).table.find_unique( + where={"id": "ui_settings"} + ) if db_record and db_record.ui_settings: ui_settings_json = db_record.ui_settings @@ -1256,11 +1314,15 @@ async def get_ui_settings(): ui_settings = dict(ui_settings_json) # Sanitize any unexpected keys from persisted config before returning - ui_settings = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS} + ui_settings = { + k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS + } # Sync runtime flags into general_settings so the proxy picks them up # at runtime (covers server restart scenarios). - _flags_to_sync = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings} + _flags_to_sync = { + k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings + } if _flags_to_sync: from litellm.proxy.proxy_server import general_settings @@ -1269,7 +1331,9 @@ async def get_ui_settings(): # Refresh DualCache so other code paths (e.g. /user/filter/ui) see fresh values from litellm.proxy.proxy_server import user_api_key_cache - await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) + await user_api_key_cache.async_set_cache( + key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL + ) # Build config-like object for schema helper config: Dict[str, Any] = {"litellm_settings": {"ui_settings": ui_settings}} @@ -1301,7 +1365,9 @@ async def update_ui_settings( ) if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can update UI settings.") + raise HTTPException( + status_code=403, detail="Only proxy admins can update UI settings." + ) if prisma_client is None: raise HTTPException( @@ -1312,7 +1378,9 @@ async def update_ui_settings( if store_model_in_db is not True: raise HTTPException( status_code=500, - detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, ) # Validate against the same effective class GET advertises, so @@ -1328,7 +1396,10 @@ async def update_ui_settings( # Reject enterprise-only settings up front so the caller gets a clear # signal instead of a silent drop. - blocked_enterprise_keys = sorted((settings_dict.keys() & _ENTERPRISE_ONLY_UI_SETTINGS) - ALLOWED_UI_SETTINGS_FIELDS) + blocked_enterprise_keys = sorted( + (settings_dict.keys() & _ENTERPRISE_ONLY_UI_SETTINGS) + - ALLOWED_UI_SETTINGS_FIELDS + ) if blocked_enterprise_keys: raise HTTPException( status_code=403, @@ -1341,12 +1412,16 @@ async def update_ui_settings( ) # Enforce allowlist and drop anything unexpected - incoming = {k: v for k, v in settings_dict.items() if k in ALLOWED_UI_SETTINGS_FIELDS} + incoming = { + k: v for k, v in settings_dict.items() if k in ALLOWED_UI_SETTINGS_FIELDS + } # Merge with existing persisted settings so a partial PATCH doesn't # overwrite fields the caller didn't send. existing: dict = {} - db_existing = await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}) + db_existing = await UISettingsRepository(prisma_client).table.find_unique( + where={"id": "ui_settings"} + ) if db_existing and db_existing.ui_settings: raw = db_existing.ui_settings existing = json.loads(raw) if isinstance(raw, str) else dict(raw) @@ -1368,7 +1443,9 @@ async def update_ui_settings( # Sync runtime flags to general_settings so the proxy picks them up # at runtime (general_settings is checked in pre-call utils). - _flags_to_sync = {k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings} + _flags_to_sync = { + k: ui_settings[k] for k in _RUNTIME_GENERAL_SETTINGS_FLAGS if k in ui_settings + } if _flags_to_sync: from litellm.proxy.proxy_server import general_settings @@ -1377,8 +1454,12 @@ async def update_ui_settings( # Invalidate + set DualCache so subsequent reads see the new values immediately from litellm.proxy.proxy_server import user_api_key_cache - sanitized = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS} - await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=sanitized, ttl=UI_SETTINGS_CACHE_TTL) + sanitized = { + k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS + } + await user_api_key_cache.async_set_cache( + key=UI_SETTINGS_CACHE_KEY, value=sanitized, ttl=UI_SETTINGS_CACHE_TTL + ) asyncio.create_task( create_config_audit_log( @@ -1424,7 +1505,9 @@ async def upload_logo(file: UploadFile = File(...)): # Validate file size (max 5MB) file_content = await file.read() if len(file_content) > 5 * 1024 * 1024: # 5MB - raise HTTPException(status_code=400, detail="File size too large. Maximum size is 5MB.") + raise HTTPException( + status_code=400, detail="File size too large. Maximum size is 5MB." + ) # Create uploads directory if it doesn't exist current_dir = os.path.dirname(os.path.abspath(__file__)) diff --git a/terraform/litellm/aws/config/prod/terraform.tfvars b/terraform/litellm/aws/config/prod/terraform.tfvars index 27438b9bec6..7e7ac8226a7 100644 --- a/terraform/litellm/aws/config/prod/terraform.tfvars +++ b/terraform/litellm/aws/config/prod/terraform.tfvars @@ -71,7 +71,6 @@ proxy_config = { litellm_settings = { callbacks = ["smtp_email"] mcp_semantic_tool_filter = { - enabled = true embedding_model = "text-embedding-3-small" top_k = 5 diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index ce2d04f0d26..81e7fdbbf50 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -5,6 +5,7 @@ from types import SimpleNamespace import pytest +from fastapi import HTTPException from fastapi.testclient import TestClient sys.path.insert( @@ -1511,6 +1512,158 @@ async def mock_check_duplicate_user_id(*args, **kwargs): litellm.default_internal_user_params = original_default_params +@pytest.mark.asyncio +async def test_new_user_applies_email_prefix_to_auto_created_key(mocker): + mock_prisma_client = mocker.MagicMock() + + async def mock_count(*args, **kwargs): + return 5 + + mock_prisma_client.db.litellm_usertable.count = mock_count + + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + mocker.AsyncMock(return_value=None), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + mocker.AsyncMock(return_value=None), + ) + mock_license_check = mocker.MagicMock() + mock_license_check.is_over_limit.return_value = False + + mock_generate_key_helper_fn = mocker.AsyncMock( + return_value={ + "user_id": "test-user-123", + "token": "sk-test-token-123", + "expires": None, + "max_budget": 100, + } + ) + mock_user_created_hook = mocker.AsyncMock() + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn", + mock_generate_key_helper_fn, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.UserManagementEventHooks.async_user_created_hook", + mock_user_created_hook, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._enforce_email_prefix_on_key_alias", + mocker.AsyncMock(return_value="owner.user-test-key-name"), + ) + + response = await new_user( + data=NewUserRequest( + user_email="owner.user@example.com", + user_role="internal_user", + key_alias="TestKey Name", + ), + user_api_key_dict=UserAPIKeyAuth(user_id="test_admin", user_email="caller@example.com"), + ) + + call_kwargs = mock_generate_key_helper_fn.call_args.kwargs + assert call_kwargs["key_alias"] == "owner.user-test-key-name" + assert response.user_id == "test-user-123" + + +@pytest.mark.asyncio +async def test_new_user_rejects_alias_without_resolvable_email(mocker): + mock_prisma_client = mocker.MagicMock() + + async def mock_count(*args, **kwargs): + return 5 + + mock_prisma_client.db.litellm_usertable.count = mock_count + + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + mocker.AsyncMock(return_value=None), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + mocker.AsyncMock(return_value=None), + ) + mock_license_check = mocker.MagicMock() + mock_license_check.is_over_limit.return_value = False + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._enforce_email_prefix_on_key_alias", + mocker.AsyncMock(side_effect=HTTPException(status_code=400, detail={"error": "missing email"})), + ) + + with pytest.raises(ProxyException) as exc_info: + await new_user( + data=NewUserRequest(user_email=None, user_role="internal_user", key_alias="Needs Alias"), + user_api_key_dict=UserAPIKeyAuth(user_id="test_admin", user_email=None), + ) + + assert str(exc_info.value.code) == "400" + assert "missing email" in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_new_user_without_alias_skips_email_prefix_enforcement(mocker): + mock_prisma_client = mocker.MagicMock() + + async def mock_count(*args, **kwargs): + return 5 + + mock_prisma_client.db.litellm_usertable.count = mock_count + + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + mocker.AsyncMock(return_value=None), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id", + mocker.AsyncMock(return_value=None), + ) + mock_license_check = mocker.MagicMock() + mock_license_check.is_over_limit.return_value = False + + mock_generate_key_helper_fn = mocker.AsyncMock( + return_value={ + "user_id": "test-user-123", + "token": "sk-test-token-123", + "expires": None, + "max_budget": 100, + } + ) + enforce_mock = mocker.AsyncMock(return_value=None) + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn", + mock_generate_key_helper_fn, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.UserManagementEventHooks.async_user_created_hook", + mocker.AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._enforce_email_prefix_on_key_alias", + enforce_mock, + ) + + await new_user( + data=NewUserRequest(user_email="owner.user@example.com", user_role="internal_user"), + user_api_key_dict=UserAPIKeyAuth(user_id="test_admin", user_email="caller@example.com"), + ) + + call_kwargs = mock_generate_key_helper_fn.call_args.kwargs + assert call_kwargs.get("key_alias") is None + enforce_mock.assert_awaited_once() + assert enforce_mock.await_args.kwargs["key_alias"] is None + + def test_update_internal_new_user_params_proxy_admin_role(): """ Test that default_internal_user_params are NOT applied when user_role is PROXY_ADMIN 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 7b97ae60443..3febaa38545 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 @@ -33,6 +33,7 @@ _check_org_key_limits, _check_team_key_limits, _common_key_generation_helper, + _enforce_email_prefix_on_key_alias, _enforce_upperbound_key_params, _get_and_validate_existing_key, _list_key_helper, @@ -3480,7 +3481,7 @@ async def test_generate_key_team_member_inherits_org_skips_membership_check(): ) with ( - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), patch("litellm.proxy.proxy_server.llm_router", None), patch("litellm.proxy.proxy_server.premium_user", True), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), @@ -3492,6 +3493,10 @@ async def test_generate_key_team_member_inherits_org_skips_membership_check(): "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team", new_callable=AsyncMock, ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + new=AsyncMock(return_value=None), + ), patch( "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", mock_validate_org, @@ -3562,10 +3567,14 @@ async def test_generate_key_foreign_org_without_team_still_enforces_membership() ) with ( - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), patch("litellm.proxy.proxy_server.llm_router", None), patch("litellm.proxy.proxy_server.premium_user", True), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + new=AsyncMock(return_value=None), + ), patch( "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", mock_validate_org, @@ -3632,7 +3641,7 @@ async def test_generate_key_foreign_org_with_mismatched_team_still_enforces_memb ) with ( - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), patch("litellm.proxy.proxy_server.llm_router", None), patch("litellm.proxy.proxy_server.premium_user", True), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), @@ -3648,6 +3657,10 @@ async def test_generate_key_foreign_org_with_mismatched_team_still_enforces_memb "litellm_enterprise.proxy.management_endpoints.key_management_endpoints.apply_enterprise_key_management_params", side_effect=lambda data, team_table: data, ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + new=AsyncMock(return_value=None), + ), patch( "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", mock_validate_org, @@ -7891,7 +7904,11 @@ async def test_default_key_generate_params_object_permission_applied_when_absent team_table=None, ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + created_data = ( + mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ + "data" + ] + ) assert created_data["vector_stores"] == ["default-vs"] finally: litellm.default_key_generate_params = original_value @@ -7956,7 +7973,11 @@ async def test_default_key_generate_params_object_permission_merges_partial( team_table=None, ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + created_data = ( + mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ + "data" + ] + ) assert created_data["agents"] == ["agent-1"] assert created_data["vector_stores"] == ["default-vs"] finally: @@ -8023,7 +8044,11 @@ async def test_default_key_generate_params_object_permission_does_not_override_e team_table=None, ) - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + created_data = ( + mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ + "data" + ] + ) assert created_data["vector_stores"] == ["explicit-vs"] finally: litellm.default_key_generate_params = original_value @@ -8068,6 +8093,10 @@ async def test_default_key_generate_params_object_permission_not_rejected_for_no ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + AsyncMock(return_value=None), + ) original_value = litellm.default_key_generate_params litellm.default_key_generate_params = { @@ -8088,7 +8117,11 @@ async def test_default_key_generate_params_object_permission_not_rejected_for_no ) assert response is not None - created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"] + created_data = ( + mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ + "data" + ] + ) assert created_data["vector_stores"] == ["default-vs"] finally: litellm.default_key_generate_params = original_value @@ -9520,7 +9553,9 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc mock_updated_key.key_alias = "my-alias" mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) - mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"key_alias": "my-alias"}} + ) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=mock_existing_key ) @@ -9535,11 +9570,20 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc 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) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={}), + ) def mock_hash_token(token): return test_hashed_token monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + AsyncMock(return_value=None), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None) async def mock_delete_cache_key_object(**kwargs): pass @@ -9758,7 +9802,9 @@ async def test_update_key_team_member_with_permission_can_update_non_budget( mock_prisma_client = AsyncMock() mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) - mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"key_alias": "renamed-by-member"}} + ) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=mock_existing_key ) @@ -9794,9 +9840,17 @@ async def mock_delete_cache_key_object(**kwargs): 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) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={}), + ) monkeypatch.setattr( "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + AsyncMock(return_value=None), + ) mock_request = MagicMock() mock_request.query_params = {} @@ -11276,7 +11330,9 @@ async def test_non_admin_update_key_explicit_null_allowed_routes_rejected(self): assert "allowed_routes" in str(exc_info.value.message) @pytest.mark.asyncio - async def test_non_admin_regenerate_key_explicit_empty_allowed_routes_rejected(self): + async def test_non_admin_regenerate_key_explicit_empty_allowed_routes_rejected( + self, + ): """`regenerate_key_fn` rejects a non-admin when `allowed_routes` is present as `[]` in the request body.""" from litellm.proxy._types import RegenerateKeyRequest @@ -11303,7 +11359,9 @@ async def test_non_admin_regenerate_key_explicit_empty_allowed_routes_rejected(s assert "allowed_routes" in str(exc_info.value.message) @pytest.mark.asyncio - async def test_non_admin_regenerate_key_allowed_routes_rejected_before_enterprise_gate(self): + async def test_non_admin_regenerate_key_allowed_routes_rejected_before_enterprise_gate( + self, + ): """`regenerate_key_fn` runs `_check_allowed_routes_caller_permission` before the `premium_user` check, so a non-premium proxy still returns the allowed_routes rejection (403) rather than the enterprise-license @@ -13743,7 +13801,9 @@ async def test_budget_limits_admin_unrestricted(monkeypatch): @pytest.mark.asyncio @pytest.mark.parametrize("non_finite", [float("nan"), float("inf"), float("-inf")]) -async def test_budget_limits_window_non_finite_rejected_for_non_admin(monkeypatch, non_finite): +async def test_budget_limits_window_non_finite_rejected_for_non_admin( + monkeypatch, non_finite +): """A non-admin caller submitting a non-finite `budget_limits` window gets 400. The finite-number invariant applies before role / ceiling checks.""" @@ -13773,7 +13833,9 @@ async def test_budget_limits_window_non_finite_rejected_for_non_admin(monkeypatc @pytest.mark.asyncio @pytest.mark.parametrize("non_finite", [float("nan"), float("inf"), float("-inf")]) -async def test_budget_limits_window_non_finite_rejected_for_admin(monkeypatch, non_finite): +async def test_budget_limits_window_non_finite_rejected_for_admin( + monkeypatch, non_finite +): """The finite-number invariant applies to every caller including proxy admin.""" monkeypatch.setattr( @@ -13862,7 +13924,9 @@ async def test_budget_limits_session_token_team_key_uses_team_ceiling(monkeypatc @pytest.mark.asyncio -async def test_budget_limits_session_token_team_key_over_team_budget_rejected(monkeypatch): +async def test_budget_limits_session_token_team_key_over_team_budget_rejected( + monkeypatch, +): """Same shape, but window exceeds the team's `max_budget`.""" monkeypatch.setattr( "litellm.proxy.management_endpoints.key_management_endpoints.litellm.default_key_generate_params", @@ -14006,7 +14070,9 @@ async def test_permissions_admin_can_set_any(monkeypatch): @pytest.mark.asyncio -async def test_permissions_explicit_empty_rejected_for_non_admin_on_generate(monkeypatch): +async def test_permissions_explicit_empty_rejected_for_non_admin_on_generate( + monkeypatch, +): """`_common_key_generation_helper` rejects a non-admin when `permissions` is present in the request body, even as `{}`. Omit-default stays allowed; that carve-out lives in @@ -14218,7 +14284,9 @@ async def test_regenerate_key_non_admin_permissions_rejected(monkeypatch): @pytest.mark.asyncio -async def test_regenerate_key_non_admin_permissions_explicit_empty_rejected(monkeypatch): +async def test_regenerate_key_non_admin_permissions_explicit_empty_rejected( + monkeypatch, +): """`regenerate_key_fn` rejects a non-admin when `permissions` is present as `{}` in the request body.""" from litellm.proxy._types import RegenerateKeyRequest @@ -14243,7 +14311,9 @@ async def test_regenerate_key_non_admin_permissions_explicit_empty_rejected(monk @pytest.mark.asyncio -async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_gate(monkeypatch): +async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_gate( + monkeypatch, +): """`regenerate_key_fn` runs `_check_permissions_caller_permission` before the `premium_user` check, so a non-premium proxy still returns the permissions rejection (403) rather than the enterprise-license @@ -14270,3 +14340,446 @@ async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_g assert int(exc.value.code) == 403 assert "permissions" in str(exc.value.message) assert "Enterprise" not in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_enforce_email_prefix_on_key_alias_prefers_owner_email(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await _enforce_email_prefix_on_key_alias( + prisma_client=None, + key_alias="TestKey Name", + owner_email="owner.user@example.com", + fallback_email="caller@example.com", + ) + + assert result == "owner.user-test-key-name" + + +@pytest.mark.asyncio +async def test_enforce_email_prefix_on_key_alias_leaves_alias_unchanged_when_flag_off( + monkeypatch, +): + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={}), + ) + + result = await _enforce_email_prefix_on_key_alias( + prisma_client=None, + key_alias="TestKey Name", + owner_email="owner.user@example.com", + fallback_email="caller@example.com", + ) + + assert result == "TestKey Name" + + +@pytest.mark.asyncio +async def test_enforce_email_prefix_on_key_alias_uses_caller_fallback(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await _enforce_email_prefix_on_key_alias( + prisma_client=None, + key_alias="testKey Name", + owner_email=None, + fallback_email="caller.user@example.com", + ) + + assert result == "caller.user-test-key-name" + + +@pytest.mark.asyncio +async def test_enforce_email_prefix_on_key_alias_avoids_double_prefix(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await _enforce_email_prefix_on_key_alias( + prisma_client=None, + key_alias="owner.user-TestKey Name", + owner_email="owner.user@example.com", + fallback_email="caller@example.com", + ) + + assert result == "owner.user-test-key-name" + + +@pytest.mark.asyncio +async def test_enforce_email_prefix_on_key_alias_rejects_without_email(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + with pytest.raises(HTTPException) as exc_info: + await _enforce_email_prefix_on_key_alias( + prisma_client=None, + key_alias="test-key", + owner_email=None, + fallback_email=None, + ) + + assert exc_info.value.status_code == 400 + assert "resolvable owner or caller email" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_common_key_generation_helper_applies_owner_email_prefix(monkeypatch): + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin-1", + user_email="caller@example.com", + ) + captured_kwargs = {} + + async def mock_generate_key_helper_fn(**kwargs): + captured_kwargs.update(kwargs) + return {"key": "sk-test", "expires": None, "user_id": kwargs.get("user_id")} + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + mock_generate_key_helper_fn, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + AsyncMock(return_value="owner.user@example.com"), + ) + + result = await _common_key_generation_helper( + data=GenerateKeyRequest( + key_alias="TestKey Name", + user_id="user-1", + ), + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + + assert result is not None + assert captured_kwargs["key_alias"] == "owner.user-test-key-name" + + +@pytest.mark.asyncio +async def test_common_key_generation_helper_uses_caller_email_for_service_account_alias( + monkeypatch, +): + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin-1", + user_email="caller.user@example.com", + ) + captured_kwargs = {} + + async def mock_generate_key_helper_fn(**kwargs): + captured_kwargs.update(kwargs) + return {"key": "sk-test", "expires": None, "user_id": kwargs.get("user_id")} + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + mock_generate_key_helper_fn, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await _common_key_generation_helper( + data=GenerateKeyRequest(key_alias="Service Account Key"), + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + + assert result is not None + assert captured_kwargs["key_alias"] == "caller.user-service-account-key" + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_applies_email_prefix_on_changed_user_id( + monkeypatch, +): + existing_key = LiteLLM_VerificationToken( + token="hashed-token", + key_alias="old-alias", + user_id="old-user", + user_email="old.user@example.com", + metadata={}, + aliases={}, + config={}, + permissions={}, + model_max_budget={}, + budget_duration=None, + models=[], + spend=0.0, + max_budget=None, + team_id=None, + ) + prisma_client = MagicMock() + mock_user_repository = MagicMock() + mock_user_repository.find_by_id = AsyncMock( + return_value=LiteLLM_UserTable( + user_id="new-user", user_email="new.owner@example.com" + ) + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.UserRepository", + lambda _: mock_user_repository, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await prepare_key_update_data( + data=UpdateKeyRequest( + key="hashed-token", user_id="new-user", key_alias="TestKey Name" + ), + existing_key_row=existing_key, + user_api_key_dict=UserAPIKeyAuth( + user_id="caller", user_email="caller@example.com" + ), + prisma_client=prisma_client, + ) + + assert result["key_alias"] == "new.owner-test-key-name" + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_falls_back_to_caller_email(monkeypatch): + existing_key = LiteLLM_VerificationToken( + token="hashed-token", + key_alias="old-alias", + user_id=None, + user_email=None, + metadata={}, + aliases={}, + config={}, + permissions={}, + model_max_budget={}, + budget_duration=None, + models=[], + spend=0.0, + max_budget=None, + team_id=None, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await prepare_key_update_data( + data=UpdateKeyRequest(key="hashed-token", key_alias="Renamed Alias"), + existing_key_row=existing_key, + user_api_key_dict=UserAPIKeyAuth( + user_id="caller", user_email="caller.user@example.com" + ), + prisma_client=None, + ) + + assert result["key_alias"] == "caller.user-renamed-alias" + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_without_alias_leaves_key_alias_unchanged( + monkeypatch, +): + existing_key = LiteLLM_VerificationToken( + token="hashed-token", + key_alias="old-alias", + user_id="old-user", + user_email="old.user@example.com", + metadata={}, + aliases={}, + config={}, + permissions={}, + model_max_budget={}, + budget_duration=None, + models=[], + spend=0.0, + max_budget=None, + team_id=None, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + result = await prepare_key_update_data( + data=UpdateKeyRequest(key="hashed-token", models=["gpt-4"]), + existing_key_row=existing_key, + user_api_key_dict=UserAPIKeyAuth( + user_id="caller", user_email="caller.user@example.com" + ), + prisma_client=None, + ) + + assert "key_alias" not in result + + +@pytest.mark.asyncio +async def test_update_key_non_budget_fields_apply_email_prefix(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = AsyncMock() + mock_proxy_logging_obj = MagicMock() + + test_hashed_token = ( + "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + ) + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = "internal_user" + mock_existing_key.user_email = "owner.user@example.com" + mock_existing_key.created_by = "internal_user" + mock_existing_key.team_id = None + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = None + mock_existing_key.models = [] + mock_existing_key.metadata = {} + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": "internal_user", + "team_id": None, + "max_budget": 10.0, + } + + captured_update = {} + + async def mock_update_data(*args, **kwargs): + captured_update.update(kwargs["data"]) + return {"data": kwargs["data"]} + + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.update_data = AsyncMock(side_effect=mock_update_data) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj + ) + 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) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + AsyncMock(), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", + AsyncMock(), + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._get_user_email_by_id", + AsyncMock(return_value=None), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_update", None) + + mock_request = MagicMock() + mock_request.query_params = {} + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + user_email="caller@example.com", + ) + + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, key_alias="My TestKey"), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert result is not None + assert captured_update["key_alias"] == "owner.user-my-test-key" + + +@pytest.mark.asyncio +async def test_update_key_non_budget_fields_reject_empty_normalized_alias(monkeypatch): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + test_hashed_token = ( + "c1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + ) + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = "internal_user" + mock_existing_key.user_email = "owner.user@example.com" + mock_existing_key.created_by = "internal_user" + mock_existing_key.team_id = None + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = None + mock_existing_key.models = [] + mock_existing_key.metadata = {} + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": "internal_user", + "team_id": None, + "max_budget": 10.0, + } + + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + 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) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_ui_settings_cached", + AsyncMock(return_value={"enforce_email_prefix_on_key_alias": True}), + ) + + mock_request = MagicMock() + mock_request.query_params = {} + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + user_email="caller@example.com", + ) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, key_alias="---"), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "400" diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 69845ec59c2..5ceaac2ac87 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -1336,6 +1336,73 @@ def test_update_ui_settings_persists_and_syncs_disable_key_generate_for_org_admi # Synced into general_settings so the enforcement helper sees it assert general_settings.get(flag_name) is True + def test_update_ui_settings_persists_email_prefix_alias_flag( + self, mock_auth, monkeypatch + ): + """Email prefix alias enforcement flag must be allowlisted, persisted, and synced to general_settings.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + mock_user_auth = UserAPIKeyAuth( + user_id="test-user-123", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth + + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + + general_settings: dict = {} + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", general_settings + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + payload = {"enforce_email_prefix_on_key_alias": True} + + try: + response = client.patch("/update/ui_settings", json=payload) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 200 + data = response.json() + assert data["settings"]["enforce_email_prefix_on_key_alias"] is True + + stored_settings = json.loads( + mock_prisma.db.litellm_uisettings.upsert.call_args.kwargs["data"]["create"]["ui_settings"] + ) + assert stored_settings["enforce_email_prefix_on_key_alias"] is True + assert general_settings.get("enforce_email_prefix_on_key_alias") is True + + def test_get_ui_settings_returns_email_prefix_alias_flag( + self, mock_auth, monkeypatch + ): + """Email prefix alias enforcement flag should round-trip through the UI settings GET response.""" + from unittest.mock import AsyncMock, MagicMock + + mock_prisma = MagicMock() + mock_db_record = MagicMock() + mock_db_record.ui_settings = {"enforce_email_prefix_on_key_alias": True} + mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( + return_value=mock_db_record + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + data = response.json() + assert data["values"]["enforce_email_prefix_on_key_alias"] is True + assert ( + "enforce_email_prefix_on_key_alias" in data["field_schema"]["properties"] + ) + def test_get_sso_settings_from_database( self, mock_proxy_config, mock_auth, monkeypatch ):