From 4782d435edd6aeddca7d92d8d17d8c75c81be9c5 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 18 Jun 2025 16:24:55 -0700 Subject: [PATCH] [Fix] SCIM - Add SCIM PATCH and PUT Ops for Users (#11863) * fix SCIM memberships Patch * fixes for SCIM updates * fixes for SCIM * working provisioning for teams on SCIM * working user patch / PUT ops SCIM * fixes SCIM * test_scim_v2_endpoints.py * handle_existing_user_by_email * fixes for provisioning SCIMUser * fixes SCIM provisioning * test scim v2 * fixes for linting * fix _apply_patch_ops * fixes code QA check for team membership checks --- .../management_endpoints/scim/scim_errors.py | 13 - .../scim/scim_transformations.py | 2 + .../management_endpoints/scim/scim_v2.py | 828 +++++++++++++----- .../proxy/management_endpoints/scim/utils.py | 20 - .../proxy/management_endpoints/scim_v2.py | 8 +- .../scim/test_scim_patch_user.py | 105 +++ .../scim/test_scim_v2_endpoints.py | 436 ++++++++- 7 files changed, 1163 insertions(+), 249 deletions(-) delete mode 100644 litellm/proxy/management_endpoints/scim/scim_errors.py delete mode 100644 litellm/proxy/management_endpoints/scim/utils.py create mode 100644 tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py diff --git a/litellm/proxy/management_endpoints/scim/scim_errors.py b/litellm/proxy/management_endpoints/scim/scim_errors.py deleted file mode 100644 index 9e09f6e095..0000000000 --- a/litellm/proxy/management_endpoints/scim/scim_errors.py +++ /dev/null @@ -1,13 +0,0 @@ -from fastapi import HTTPException - - -class ScimUserAlreadyExists(HTTPException): - """ - Exception raised when a user already exists in the database. - """ - def __init__(self, message: str, scim_type: str = "uniqueness"): - super().__init__(status_code=409, detail=message) - self.message = message - self.scim_type = scim_type - self.schemas = ["urn:ietf:params:scim:api:messages:2.0:Error"] - \ No newline at end of file diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index f68f728e2a..bb07cdbd77 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -124,6 +124,8 @@ class ScimTransformations: # Get team members scim_members: List[SCIMMember] = [] for member in team.members_with_roles or []: + if isinstance(member, dict): + member = Member(**member) scim_members.append( SCIMMember( value=ScimTransformations._get_scim_member_value(member), diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index e1dce38fb8..b88413f180 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -5,7 +5,7 @@ This is an enterprise feature and requires a premium license. """ import uuid -from typing import List, Optional +from typing import Any, Dict, List, Optional, Set, Tuple, TypedDict from fastapi import ( APIRouter, @@ -19,28 +19,85 @@ from fastapi import ( ) from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ( LiteLLM_UserTable, LitellmUserRoles, Member, NewTeamRequest, NewUserRequest, + TeamMemberAddRequest, + TeamMemberDeleteRequest, UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.internal_user_endpoints import new_user -from litellm.proxy.management_endpoints.scim.scim_errors import ScimUserAlreadyExists from litellm.proxy.management_endpoints.scim.scim_transformations import ( ScimTransformations, ) -from litellm.proxy.management_endpoints.scim.utils import ( - _check_user_exists, - _extract_error_message, +from litellm.proxy.management_endpoints.team_endpoints import ( + new_team, + team_member_add, + team_member_delete, ) -from litellm.proxy.management_endpoints.team_endpoints import new_team from litellm.proxy.utils import _premium_user_check, handle_exception_on_proxy from litellm.types.proxy.management_endpoints.scim_v2 import * + +class UserProvisionerHelpers: + """Helper methods for user provisioning operations.""" + + @staticmethod + async def handle_existing_user_by_email( + prisma_client, + new_user_request: NewUserRequest + ) -> Optional[SCIMUser]: + """ + Check if a user with the given email already exists and update them if found. + + Args: + prisma_client: Database client + new_user_request: New user request data + + Returns: + SCIMUser if user was updated, None if no existing user found + """ + if not new_user_request.user_email: + return None + + existing_user = await prisma_client.db.litellm_usertable.find_first( + where={"user_email": new_user_request.user_email} + ) + + if not existing_user: + return None + + # Update the user + updated_user = await prisma_client.db.litellm_usertable.update( + where={"user_id": existing_user.user_id}, + data={ + "user_id": new_user_request.user_id, + "user_email": new_user_request.user_email, + "user_alias": new_user_request.user_alias, + "teams": new_user_request.teams, + "metadata": safe_dumps(new_user_request.metadata), + }, + ) + + return await ScimTransformations.transform_litellm_user_to_scim_user(updated_user) + + +class ScimUserData(TypedDict): + """Typed structure for extracted SCIM user data.""" + user_email: Optional[str] + user_alias: Optional[str] + sso_user_id: Optional[str] + teams: List[str] + given_name: Optional[str] + family_name: Optional[str] + active: Optional[bool] + + scim_router = APIRouter( prefix="/scim/v2", tags=["✨ SCIM v2 (Enterprise Only)"], @@ -48,6 +105,137 @@ scim_router = APIRouter( ) +# Helper functions for common operations +async def _get_prisma_client_or_raise_exception(): + """Check if database is connected and raise HTTPException if not.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail={"error": "No database connected"}) + return prisma_client + + +async def _check_user_exists(user_id: str): + """Check if user exists and return user, raise 404 if not found.""" + prisma_client = await _get_prisma_client_or_raise_exception() + + user = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_id} + ) + + if not user: + raise HTTPException( + status_code=404, detail={"error": f"User not found with ID: {user_id}"} + ) + + return user + + +async def _check_team_exists(team_id: str): + """Check if team exists and return team, raise 404 if not found.""" + prisma_client = await _get_prisma_client_or_raise_exception() + + team = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + + if not team: + raise HTTPException( + status_code=404, detail={"error": f"Group not found with ID: {team_id}"} + ) + + return team + + +def _extract_scim_user_data(user: SCIMUser) -> ScimUserData: + """Extract common data from SCIMUser object.""" + user_email = None + if user.emails and len(user.emails) > 0: + user_email = user.emails[0].value + + user_alias = None + if user.name and user.name.givenName: + user_alias = user.name.givenName + + teams = [] + if user.groups: + teams = [group.value for group in user.groups] + + return { + "user_email": user_email, + "user_alias": user_alias, + "sso_user_id": user.externalId, + "teams": teams, + "given_name": user.name.givenName if user.name else None, + "family_name": user.name.familyName if user.name else None, + "active": user.active, + } + + +def _build_scim_metadata(given_name: Optional[str], family_name: Optional[str], active: Optional[bool] = None) -> Dict[str, Any]: + """Build metadata dictionary with SCIM data.""" + metadata: Dict[str, Any] = { + "scim_metadata": LiteLLM_UserScimMetadata( + givenName=given_name, + familyName=family_name, + ).model_dump() + } + + if active is not None: + metadata["scim_active"] = active + + return metadata + + +async def _extract_group_member_ids(group: SCIMGroup) -> List[str]: + """Extract valid member IDs from SCIMGroup, verifying users exist.""" + prisma_client = await _get_prisma_client_or_raise_exception() + member_ids = [] + + if group.members: + for member in group.members: + # Check if user exists + user = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": member.value} + ) + if user: + member_ids.append(member.value) + + return member_ids + + +async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: + """Get SCIMMember objects with display names for a list of member IDs.""" + prisma_client = await _get_prisma_client_or_raise_exception() + members: List[SCIMMember] = [] + + for member_id in member_ids: + user = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": member_id} + ) + if user: + display_name = user.user_email or user.user_id + members.append(SCIMMember(value=user.user_id, display=display_name)) + + return members + + +async def _handle_team_membership_changes(user_id: str, existing_teams: List[str], new_teams: List[str]) -> None: + """Handle adding/removing user from teams based on changes.""" + existing_teams_set = set(existing_teams) + new_teams_set = set(new_teams) + + teams_to_add = new_teams_set - existing_teams_set + teams_to_remove = existing_teams_set - new_teams_set + + if teams_to_add or teams_to_remove: + await patch_team_membership( + user_id=user_id, + teams_ids_to_add_user_to=list(teams_to_add), + teams_ids_to_remove_user_from=list(teams_to_remove), + ) + + # Dependency to set the correct SCIM Content-Type async def set_scim_content_type(response: Response): """Sets the Content-Type header to application/scim+json""" @@ -71,12 +259,8 @@ async def get_users( """ Get a list of users according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: + prisma_client = await _get_prisma_client_or_raise_exception() # Parse filter if provided (basic support) where_conditions = {} if filter: @@ -134,21 +318,9 @@ async def get_user( """ Get a single user by ID according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: - user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_id} - ) - - if not user: - raise HTTPException( - status_code=404, detail={"error": f"User not found with ID: {user_id}"} - ) - + user = await _check_user_exists(user_id) + # Convert to SCIM format scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(user) return scim_user @@ -168,54 +340,55 @@ async def create_user( """ Create a user according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: verbose_proxy_logger.debug("SCIM CREATE USER request: %s", user) + prisma_client = await _get_prisma_client_or_raise_exception() - # Extract user data - user_email = user.emails[0].value if user.emails else None - user_id = user.userName or str(uuid.uuid4()) + # Extract data from SCIM user + user_data = _extract_scim_user_data(user) - # Check for duplicate username - if await _check_user_exists(prisma_client, user.userName): - raise ScimUserAlreadyExists( - message=f"User already exists with username: {user.userName}" + # Check if user already exists + if user.userName: + existing_user = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user.userName} ) - - # Attempt to create user - try: - created_user = await new_user( - data=NewUserRequest( - user_id=user_id, - user_email=user_email, - user_alias=user.name.givenName, - teams=[group.value for group in user.groups] if user.groups else None, - metadata={ - "scim_metadata": LiteLLM_UserScimMetadata( - givenName=user.name.givenName, - familyName=user.name.familyName, - ).model_dump() - }, - auto_create_key=False, - ), - ) - except HTTPException as e: - # Convert duplicate email errors to SCIM 409 - if e.status_code == 400 and "already exists" in str(e.detail): - raise ScimUserAlreadyExists( - message=_extract_error_message(e) + if existing_user: + raise HTTPException( + status_code=409, + detail={"error": f"User already exists with username: {user.userName}"}, ) - raise e - # Transform and return SCIM user - return await ScimTransformations.transform_litellm_user_to_scim_user(created_user) + # Create user in database + user_id = user.userName or str(uuid.uuid4()) + metadata = _build_scim_metadata(user_data["given_name"], user_data["family_name"]) + new_user_request = NewUserRequest( + user_id=user_id, + user_email=user_data["user_email"], + user_alias=user_data["user_alias"], + teams=user_data["teams"], + metadata=metadata, + auto_create_key=False, + ) + + # Check if user with email already exists and update if found + existing_user_scim = await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=prisma_client, + new_user_request=new_user_request + ) - except HTTPException: - raise # Let HTTPExceptions (including ScimUserAlreadyExists) propagate directly + if existing_user_scim: + return existing_user_scim + + created_user = await new_user( + data=new_user_request, + ) + + scim_user = await ScimTransformations.transform_litellm_user_to_scim_user( + user=created_user + ) + return scim_user + except HTTPException as e: # allow exceptions like SCIMUserAlreadyExists to be raised + raise e except Exception as e: raise handle_exception_on_proxy(e) @@ -231,14 +404,55 @@ async def update_user( user: SCIMUser = Body(...), ): """ - Update a user according to SCIM v2 protocol + Update a user according to SCIM v2 protocol (full replacement) """ - from litellm.proxy.proxy_server import prisma_client + verbose_proxy_logger.debug("SCIM PUT USER request: %s", user) - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) try: - return None + prisma_client = await _get_prisma_client_or_raise_exception() + existing_user = await _check_user_exists(user_id) + + # Extract data from SCIM user + user_data = _extract_scim_user_data(user) + + # Build metadata with SCIM data + metadata = _build_scim_metadata( + user_data["given_name"], + user_data["family_name"], + user_data["active"] + ) + + # Handle team membership changes + await _handle_team_membership_changes( + user_id=user_id, + existing_teams=existing_user.teams or [], + new_teams=user_data["teams"] + ) + + # Update user with all new data (full replacement) + update_data = { + "user_email": user_data["user_email"], + "user_alias": user_data["user_alias"], + "sso_user_id": user_data["sso_user_id"], + "teams": user_data["teams"], + "metadata": metadata, + } + + # Serialize metadata to JSON string for Prisma to avoid GraphQL parsing issues + if "metadata" in update_data and isinstance(update_data["metadata"], dict): + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + update_data["metadata"] = safe_dumps(update_data["metadata"]) + + updated_user = await prisma_client.db.litellm_usertable.update( + where={"user_id": user_id}, + data=update_data, + ) + + # Convert back to SCIM format + scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(updated_user) + + return scim_user + except Exception as e: raise handle_exception_on_proxy(e) @@ -254,21 +468,9 @@ async def delete_user( """ Delete a user according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: - # Check if user exists - existing_user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_id} - ) - - if not existing_user: - raise HTTPException( - status_code=404, detail={"error": f"User not found with ID: {user_id}"} - ) + prisma_client = await _get_prisma_client_or_raise_exception() + existing_user = await _check_user_exists(user_id) # Get teams user belongs to teams = [] @@ -297,6 +499,156 @@ async def delete_user( raise handle_exception_on_proxy(e) +def _extract_group_values(value: Any) -> List[str]: + """Return group ids from a SCIM patch value.""" + group_values: List[str] = [] + if isinstance(value, list): + for v in value: + if isinstance(v, dict) and v.get("value"): + group_values.append(str(v.get("value"))) + elif isinstance(v, str): + group_values.append(v) + elif isinstance(value, dict): + if value.get("value"): + group_values.append(str(value.get("value"))) + elif isinstance(value, str): + group_values.append(value) + return group_values + + +def _handle_displayname_update(op_type: str, value: Any, update_data: Dict[str, Any]) -> None: + """Handle displayname updates.""" + if op_type == "remove": + update_data["user_alias"] = None + else: + update_data["user_alias"] = str(value) + + +def _handle_externalid_update(op_type: str, value: Any, update_data: Dict[str, Any]) -> None: + """Handle externalid updates.""" + if op_type == "remove": + update_data["sso_user_id"] = None + else: + update_data["sso_user_id"] = str(value) + + +def _handle_active_update(op_type: str, value: Any, metadata: Dict[str, Any]) -> None: + """Handle active status updates.""" + if op_type == "remove": + metadata.pop("scim_active", None) + else: + bool_val = value + if isinstance(value, str): + bool_val = value.lower() == "true" + else: + bool_val = bool(value) + metadata["scim_active"] = bool_val + + +def _handle_name_update(path: str, op_type: str, value: Any, scim_metadata: Dict[str, Any]) -> None: + """Handle name field updates (givenName, familyName).""" + if path == "name.givenname": + if op_type == "remove": + scim_metadata.pop("givenName", None) + else: + scim_metadata["givenName"] = str(value) + elif path == "name.familyname": + if op_type == "remove": + scim_metadata.pop("familyName", None) + else: + scim_metadata["familyName"] = str(value) + + +def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str]) -> Optional[Set[str]]: + """Handle group/team membership operations.""" + group_values = _extract_group_values(value) + if op_type == "replace": + return set(group_values) + elif op_type == "add": + teams_set.update(group_values) + elif op_type == "remove": + for gid in group_values: + teams_set.discard(gid) + return None + + +def _handle_generic_metadata(path: str, op_type: str, value: Any, metadata: Dict[str, Any]) -> None: + """Handle generic metadata operations for unknown paths.""" + if op_type == "remove": + metadata.pop(path, None) + else: + metadata[path] = value + + +def _apply_patch_ops( + existing_user: LiteLLM_UserTable, + patch_ops: SCIMPatchOp, +) -> Tuple[Dict[str, Any], Set[str]]: + """Apply patch operations and return update data and final team set.""" + update_data: Dict[str, Any] = {} + metadata = existing_user.metadata or {} + scim_metadata = metadata.get("scim_metadata", {}) + + teams_set: Set[str] = set(existing_user.teams or []) + replace_team_set: Optional[Set[str]] = None + + for op in patch_ops.Operations: + path = (op.path or "").lower() + value = op.value + op_type = op.op + + if path == "displayname": + _handle_displayname_update(op_type, value, update_data) + elif path == "externalid": + _handle_externalid_update(op_type, value, update_data) + elif path == "active": + _handle_active_update(op_type, value, metadata) + elif path in ("name.givenname", "name.familyname"): + _handle_name_update(path, op_type, value, scim_metadata) + elif path.startswith("groups"): + new_replace_set = _handle_group_operations(op_type, value, teams_set) + if new_replace_set is not None: + replace_team_set = new_replace_set + else: + _handle_generic_metadata(path, op_type, value, metadata) + + final_team_set = replace_team_set if replace_team_set is not None else teams_set + metadata["scim_metadata"] = scim_metadata + update_data["metadata"] = metadata + return update_data, final_team_set + +async def patch_team_membership( + user_id: str, + teams_ids_to_add_user_to: List[str], + teams_ids_to_remove_user_from: List[str], +) -> bool: + """ + Add or remove user from teams + """ + for _team_id in teams_ids_to_add_user_to: + try: + await team_member_add( + data=TeamMemberAddRequest( + team_id=_team_id, + member=Member(user_id=user_id, role="user"), + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}") + + for _team_id in teams_ids_to_remove_user_from: + try: + await team_member_delete( + data=TeamMemberDeleteRequest(team_id=_team_id, user_id=user_id), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}") + + + return True + @scim_router.patch( "/Users/{user_id}", response_model=SCIMUser, @@ -310,25 +662,39 @@ async def patch_user( """ Patch a user according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - verbose_proxy_logger.debug("SCIM PATCH USER request: %s", patch_ops) try: - # Check if user exists - existing_user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_id} + prisma_client = await _get_prisma_client_or_raise_exception() + existing_user = await _check_user_exists(user_id) + + update_data, final_team_set = _apply_patch_ops( + existing_user=existing_user, + patch_ops=patch_ops, ) - if not existing_user: - raise HTTPException( - status_code=404, detail={"error": f"User not found with ID: {user_id}"} - ) + # Handle team membership changes + await _handle_team_membership_changes( + user_id=user_id, + existing_teams=existing_user.teams or [], + new_teams=list(final_team_set) + ) - return None + update_data["teams"] = list(final_team_set) + + # Serialize metadata to JSON string for Prisma to avoid GraphQL parsing issues + if "metadata" in update_data and isinstance(update_data["metadata"], dict): + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + update_data["metadata"] = safe_dumps(update_data["metadata"]) + + updated_user = await prisma_client.db.litellm_usertable.update( + where={"user_id": user_id}, + data=update_data, + ) + + scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(updated_user) + + return scim_user except Exception as e: raise handle_exception_on_proxy(e) @@ -349,12 +715,8 @@ async def get_groups( """ Get a list of groups according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: + prisma_client = await _get_prisma_client_or_raise_exception() # Parse filter if provided (basic support) where_conditions = {} if filter: @@ -379,18 +741,9 @@ async def get_groups( # Convert to SCIM format scim_groups = [] for team in teams: - # Get team members - members = [] - for member_id in team.members or []: - member = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": member_id} - ) - if member: - display_name = member.user_email or member.user_id - members.append( - SCIMMember(value=member.user_id, display=display_name) - ) - + # Get team members with display names + members = await _get_team_members_display(team.members or []) + verbose_proxy_logger.debug(f"SCIM GET GROUPS members: {members}") team_alias = getattr(team, "team_alias", team.team_id) team_created_at = team.created_at.isoformat() if team.created_at else None team_updated_at = team.updated_at.isoformat() if team.updated_at else None @@ -408,6 +761,7 @@ async def get_groups( ) scim_groups.append(scim_group) + verbose_proxy_logger.debug(f"SCIM GET GROUPS response: {scim_groups}") return SCIMListResponse( totalResults=total_count, startIndex=startIndex, @@ -431,25 +785,13 @@ async def get_group( """ Get a single group by ID according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: - team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": group_id} - ) - - if not team: - raise HTTPException( - status_code=404, - detail={"error": f"Group not found with ID: {group_id}"}, - ) + team = await _check_team_exists(group_id) scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( team ) + verbose_proxy_logger.debug(f"SCIM GET GROUP response: {scim_group}") return scim_group except Exception as e: @@ -468,12 +810,9 @@ async def create_group( """ Create a group according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: + prisma_client = await _get_prisma_client_or_raise_exception() + # Generate ID if not provided team_id = group.id or str(uuid.uuid4()) @@ -488,16 +827,9 @@ async def create_group( detail={"error": f"Group already exists with ID: {team_id}"}, ) - # Extract members - members_with_roles: List[Member] = [] - if group.members: - for member in group.members: - # Check if user exists - user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": member.value} - ) - if user: - members_with_roles.append(Member(user_id=member.value, role="user")) + # Extract valid member IDs + member_ids = await _extract_group_member_ids(group) + members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_ids] # Create team in database created_team = await new_team( @@ -531,33 +863,12 @@ async def update_group( """ Update a group according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: - # Check if team exists - existing_team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": group_id} - ) + prisma_client = await _get_prisma_client_or_raise_exception() + existing_team = await _check_team_exists(group_id) - if not existing_team: - raise HTTPException( - status_code=404, - detail={"error": f"Group not found with ID: {group_id}"}, - ) - - # Extract members - member_ids = [] - if group.members: - for member in group.members: - # Check if user exists - user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": member.value} - ) - if user: - member_ids.append(member.value) + # Extract valid member IDs + member_ids = await _extract_group_member_ids(group) # Update team in database existing_metadata = existing_team.metadata if existing_team.metadata else {} @@ -602,14 +913,7 @@ async def update_group( ) # Get updated members for response - members = [] - for member_id in member_ids: - user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": member_id} - ) - if user: - display_name = user.user_email or user.user_id - members.append(SCIMMember(value=user.user_id, display=display_name)) + members = await _get_team_members_display(member_ids) team_created_at = ( updated_team.created_at.isoformat() if updated_team.created_at else None @@ -645,22 +949,9 @@ async def delete_group( """ Delete a group according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - try: - # Check if team exists - existing_team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": group_id} - ) - - if not existing_team: - raise HTTPException( - status_code=404, - detail={"error": f"Group not found with ID: {group_id}"}, - ) + prisma_client = await _get_prisma_client_or_raise_exception() + existing_team = await _check_team_exists(group_id) # For each member, remove this team from their teams list for member_id in existing_team.members or []: @@ -684,6 +975,122 @@ async def delete_group( raise handle_exception_on_proxy(e) +async def _process_group_patch_operations( + patch_ops: SCIMPatchOp, + existing_team, + prisma_client +) -> Tuple[Dict[str, Any], Set[str]]: + """Process patch operations for a group and return update data and final members.""" + update_data: Dict[str, Any] = {} + + # Create a fresh copy of existing metadata to avoid Prisma issues + existing_metadata = existing_team.metadata or {} + metadata = dict(existing_metadata) if existing_metadata else {} + + # Track member changes + current_members = set(existing_team.members or []) + final_members = current_members.copy() + + # Process each patch operation + for op in patch_ops.Operations: + path = (op.path or "").lower() + value = op.value + op_type = op.op + + if path == "displayname": + if op_type == "remove": + update_data["team_alias"] = None + else: + update_data["team_alias"] = str(value) + elif path == "externalid": + if op_type == "remove": + metadata.pop("externalId", None) + else: + metadata["externalId"] = str(value) + elif path.startswith("members"): + # Handle member operations + member_values = _extract_group_values(value) + # Validate that users exist + valid_members = [] + for member_id in member_values: + user = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": member_id} + ) + if user: + valid_members.append(member_id) + + if op_type == "replace": + final_members = set(valid_members) + elif op_type == "add": + final_members.update(valid_members) + elif op_type == "remove": + for member_id in valid_members: + final_members.discard(member_id) + else: + # Handle other generic metadata + if op_type == "remove": + metadata.pop(path, None) + else: + metadata[path] = value + + # Include metadata in update data if it exists + if metadata: + update_data["metadata"] = metadata + + return update_data, final_members + + +async def _apply_group_patch_updates( + group_id: str, + update_data: Dict[str, Any], + final_members: Set[str], + prisma_client +): + """Apply patch updates to the group in the database.""" + # Serialize metadata if present + if "metadata" in update_data and isinstance(update_data["metadata"], dict): + update_data["metadata"] = safe_dumps(update_data["metadata"]) + + # Update members list + update_data["members"] = list(final_members) + + # Update team in database + updated_team = await prisma_client.db.litellm_teamtable.update( + where={"team_id": group_id}, + data=update_data, + ) + + return updated_team + + +async def _handle_group_membership_changes( + group_id: str, + current_members: Set[str], + final_members: Set[str] +): + """Handle adding/removing members from the group.""" + members_to_add = final_members - current_members + members_to_remove = current_members - final_members + + verbose_proxy_logger.debug(f"members_to_add: {members_to_add}") + verbose_proxy_logger.debug(f"members_to_remove: {members_to_remove}") + + # Use existing helper functions for team membership changes + for member_id in members_to_add: + await patch_team_membership( + user_id=member_id, + teams_ids_to_add_user_to=[group_id], + teams_ids_to_remove_user_from=[], + ) + + for member_id in members_to_remove: + await patch_team_membership( + user_id=member_id, + teams_ids_to_add_user_to=[], + teams_ids_to_remove_user_from=[group_id], + ) + + @scim_router.patch( "/Groups/{group_id}", response_model=SCIMGroup, @@ -697,24 +1104,35 @@ async def patch_group( """ Patch a group according to SCIM v2 protocol """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException(status_code=500, detail={"error": "No database connected"}) - verbose_proxy_logger.debug("SCIM PATCH GROUP request: %s", patch_ops) try: - # Check if group exists - existing_team = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": group_id} + prisma_client = await _get_prisma_client_or_raise_exception() + existing_team = await _check_team_exists(group_id) + + # Process patch operations + update_data, final_members = await _process_group_patch_operations( + patch_ops, existing_team, prisma_client + ) + + # Track current members for comparison + current_members = set(existing_team.members or []) + + # Apply updates to the database + updated_team = await _apply_group_patch_updates( + group_id, update_data, final_members, prisma_client ) - if not existing_team: - raise HTTPException( - status_code=404, - detail={"error": f"Group not found with ID: {group_id}"}, - ) - return None + # Handle user-team relationship changes + await _handle_group_membership_changes( + group_id, current_members, final_members + ) + + # Convert to SCIM format and return + scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( + updated_team + ) + return scim_group + except Exception as e: raise handle_exception_on_proxy(e) diff --git a/litellm/proxy/management_endpoints/scim/utils.py b/litellm/proxy/management_endpoints/scim/utils.py deleted file mode 100644 index e899781b20..0000000000 --- a/litellm/proxy/management_endpoints/scim/utils.py +++ /dev/null @@ -1,20 +0,0 @@ -from fastapi import HTTPException - - -async def _check_user_exists(prisma_client, user_name: str) -> bool: - """Check if user already exists by username""" - if not user_name: - return False - - existing_user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_name} - ) - return existing_user is not None - - -def _extract_error_message(http_exception: HTTPException) -> str: - """Extract error message from HTTPException detail""" - if isinstance(http_exception.detail, dict): - return http_exception.detail.get("error", "User already exists") - return str(http_exception.detail) - diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index fa78b8b650..f97553f371 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -22,8 +22,8 @@ class SCIMResource(BaseModel): class SCIMUserName(BaseModel): - familyName: str - givenName: str + familyName: Optional[str] = None + givenName: Optional[str] = None formatted: Optional[str] = None middleName: Optional[str] = None honorificPrefix: Optional[str] = None @@ -43,8 +43,8 @@ class SCIMUserGroup(BaseModel): class SCIMUser(SCIMResource): - userName: str - name: SCIMUserName + userName: Optional[str] = None + name: Optional[SCIMUserName] = None displayName: Optional[str] = None active: bool = True emails: Optional[List[SCIMUserEmail]] = None diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py new file mode 100644 index 0000000000..db96333b3d --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py @@ -0,0 +1,105 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.management_endpoints.scim.scim_v2 import patch_user +from litellm.types.proxy.management_endpoints.scim_v2 import ( + SCIMPatchOp, + SCIMPatchOperation, +) + + +@pytest.mark.asyncio +async def test_patch_user_updates_fields(): + mock_user = LiteLLM_UserTable( + user_id="user-1", + user_email="test@example.com", + user_alias="Old", + teams=[], + metadata={}, + ) + + async def mock_update(*, where, data): + if "user_alias" in data: + mock_user.user_alias = data["user_alias"] + if "metadata" in data: + mock_user.metadata = data["metadata"] + if "teams" in data: + mock_user.teams = data["teams"] + if "sso_user_id" in data: + mock_user.sso_user_id = data["sso_user_id"] + return mock_user + + mock_client = MagicMock() + mock_db = MagicMock() + mock_client.db = mock_db + mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) + mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) + mock_db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation(op="replace", path="displayName", value="New Name"), + SCIMPatchOperation(op="replace", path="active", value="False"), + ] + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client): + result = await patch_user(user_id="user-1", patch_ops=patch_ops) + + mock_db.litellm_usertable.update.assert_called_once() + assert result.displayName == "New Name" + assert mock_user.metadata.get("scim_active") is False + + +@pytest.mark.asyncio +async def test_patch_user_manages_group_memberships(): + mock_user = LiteLLM_UserTable( + user_id="user-2", + user_email="test@example.com", + user_alias="Old", + teams=["old-team"], + metadata={}, + ) + + async def mock_update(*, where, data): + if "teams" in data: + mock_user.teams = data["teams"] + if "metadata" in data: + mock_user.metadata = data["metadata"] + return mock_user + + mock_client = MagicMock() + mock_db = MagicMock() + mock_client.db = mock_db + mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) + mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) + + async def mock_add(data, user_api_key_dict): + mock_user.teams.append(data.team_id) + + async def mock_delete(data, user_api_key_dict): + if data.team_id in mock_user.teams: + mock_user.teams.remove(data.team_id) + + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation(op="add", path="groups", value=[{"value": "new-team"}]), + SCIMPatchOperation(op="remove", path="groups", value=[{"value": "old-team"}]), + ] + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client), patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock(side_effect=mock_add), + ) as mock_add_fn, patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(side_effect=mock_delete), + ) as mock_del_fn: + await patch_user(user_id="user-2", patch_ops=patch_ops) + + assert mock_add_fn.called + assert mock_del_fn.called + assert mock_user.teams == ["new-team"] + diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 8f7ec1f9b4..0cc3436d81 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -1,12 +1,23 @@ from unittest.mock import AsyncMock import pytest +from fastapi import HTTPException +from litellm.proxy._types import NewUserRequest, ProxyException from litellm.proxy.management_endpoints.scim.scim_errors import ScimUserAlreadyExists -from litellm.proxy.management_endpoints.scim.scim_v2 import create_user +from litellm.proxy.management_endpoints.scim.scim_v2 import ( + UserProvisionerHelpers, + _handle_team_membership_changes, + create_user, + patch_user, + update_user, +) from litellm.types.proxy.management_endpoints.scim_v2 import ( + SCIMPatchOp, + SCIMPatchOperation, SCIMUser, SCIMUserEmail, + SCIMUserGroup, SCIMUserName, ) @@ -22,21 +33,432 @@ async def test_create_user_existing_user_conflict(mocker): emails=[SCIMUserEmail(value="existing@example.com")], ) - mock_prisma = mocker.MagicMock() + # Create a properly structured mock for the prisma client + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value={"user_id": "existing-user"}) - mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + # Mock the _get_prisma_client_or_raise_exception to return our mock mocker.patch( - "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", - AsyncMock(return_value=True), + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), ) + mocked_new_user = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.new_user", AsyncMock(), ) - with pytest.raises(ScimUserAlreadyExists) as exc_info: + with pytest.raises(HTTPException) as exc_info: await create_user(user=scim_user) + # Check that it's an HTTPException with status 409 assert exc_info.value.status_code == 409 - assert "existing-user" in exc_info.value.message + assert "existing-user" in str(exc_info.value.detail) mocked_new_user.assert_not_called() + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_no_email(mocker): + """Should return None when new_user_request has no email""" + mock_prisma_client = mocker.MagicMock() + + new_user_request = NewUserRequest( + user_id="test-user", + user_email=None, # No email provided + user_alias="Test User", + teams=[], + metadata={}, + auto_create_key=False, + ) + + result = await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, + new_user_request=new_user_request + ) + + assert result is None + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_no_existing_user(mocker): + """Should return None when no existing user is found with the email""" + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + + new_user_request = NewUserRequest( + user_id="test-user", + user_email="test@example.com", + user_alias="Test User", + teams=["team1"], + metadata={"key": "value"}, + auto_create_key=False, + ) + + result = await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, + new_user_request=new_user_request + ) + + assert result is None + mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with( + where={"user_email": "test@example.com"} + ) + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_existing_user_updated(mocker): + """Should update existing user and return SCIMUser when user with email exists""" + # Mock existing user - create a proper mock object with attributes + existing_user = mocker.MagicMock() + existing_user.user_id = "old-user-id" + existing_user.user_email = "test@example.com" + existing_user.user_alias = "Old Name" + existing_user.teams = ["old-team"] + existing_user.metadata = {"old": "data"} + + # Mock updated user + updated_user = { + "user_id": "new-user-id", + "user_email": "test@example.com", + "user_alias": "New Name", + "teams": ["new-team"], + "metadata": '{"new": "data"}' + } + + # Mock SCIM user to be returned + mock_scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="new-user-id", + userName="new-user-id", + name=SCIMUserName(familyName="Name", givenName="New"), + emails=[SCIMUserEmail(value="test@example.com")], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) + + # Mock the transformation function + mock_transform = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=mock_scim_user) + ) + + new_user_request = NewUserRequest( + user_id="new-user-id", + user_email="test@example.com", + user_alias="New Name", + teams=["new-team"], + metadata={"new": "data"}, + auto_create_key=False, + ) + + result = await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, + new_user_request=new_user_request + ) + + # Verify the result + assert result == mock_scim_user + + # Verify database operations + mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with( + where={"user_email": "test@example.com"} + ) + + mock_prisma_client.db.litellm_usertable.update.assert_called_once_with( + where={"user_id": "old-user-id"}, + data={ + "user_id": "new-user-id", + "user_email": "test@example.com", + "user_alias": "New Name", + "teams": ["new-team"], + "metadata": '{"new": "data"}', + }, + ) + + # Verify transformation was called + mock_transform.assert_called_once_with(updated_user) + + +@pytest.mark.asyncio +async def test_handle_team_membership_changes_no_changes(mocker): + """Should not call patch_team_membership when existing teams equal new teams""" + mock_patch_team_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock() + ) + + # Same teams - no changes + await _handle_team_membership_changes( + user_id="test-user", + existing_teams=["team1", "team2"], + new_teams=["team1", "team2"] + ) + + # Should not be called since no changes + mock_patch_team_membership.assert_not_called() + + +@pytest.mark.asyncio +async def test_handle_team_membership_changes_add_teams(mocker): + """Should call patch_team_membership with teams to add""" + mock_patch_team_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock() + ) + + # Adding teams + await _handle_team_membership_changes( + user_id="test-user", + existing_teams=["team1"], + new_teams=["team1", "team2", "team3"] + ) + + mock_patch_team_membership.assert_called_once_with( + user_id="test-user", + teams_ids_to_add_user_to=["team2", "team3"], # Order might vary due to set operations + teams_ids_to_remove_user_from=[] + ) + + +@pytest.mark.asyncio +async def test_handle_team_membership_changes_remove_teams(mocker): + """Should call patch_team_membership with teams to remove""" + mock_patch_team_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock() + ) + + # Removing teams + await _handle_team_membership_changes( + user_id="test-user", + existing_teams=["team1", "team2", "team3"], + new_teams=["team1"] + ) + + mock_patch_team_membership.assert_called_once_with( + user_id="test-user", + teams_ids_to_add_user_to=[], + teams_ids_to_remove_user_from=["team2", "team3"] # Order might vary due to set operations + ) + + +@pytest.mark.asyncio +async def test_handle_team_membership_changes_add_and_remove(mocker): + """Should call patch_team_membership with both teams to add and remove""" + mock_patch_team_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock() + ) + + # Both adding and removing teams + await _handle_team_membership_changes( + user_id="test-user", + existing_teams=["team1", "team2"], + new_teams=["team2", "team3"] + ) + + # team1 should be removed, team3 should be added, team2 stays + mock_patch_team_membership.assert_called_once_with( + user_id="test-user", + teams_ids_to_add_user_to=["team3"], + teams_ids_to_remove_user_from=["team1"] + ) + + +@pytest.mark.asyncio +async def test_update_user_success(mocker): + """Should successfully update user with PUT request""" + # Mock existing user + existing_user = mocker.MagicMock() + existing_user.teams = ["old-team"] + + # Mock updated user + updated_user = { + "user_id": "test-user", + "user_email": "updated@example.com", + "user_alias": "Updated User", + "teams": ["new-team"], + "metadata": '{"scim_metadata": {"givenName": "Updated", "familyName": "User"}}' + } + + # Mock SCIM user for request + scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + userName="test-user", + name=SCIMUserName(familyName="User", givenName="Updated"), + emails=[SCIMUserEmail(value="updated@example.com")], + groups=[SCIMUserGroup(value="new-team")] + ) + + # Mock SCIM user for response + response_scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="test-user", + userName="test-user", + name=SCIMUserName(familyName="User", givenName="Updated"), + emails=[SCIMUserEmail(value="updated@example.com")], + ) + + # Mock prisma client + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) + + # Mock dependencies + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(return_value=existing_user) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock() + ) + mock_transform = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=response_scim_user) + ) + + # Call update_user + result = await update_user(user_id="test-user", user=scim_user) + + # Verify result + assert result == response_scim_user + + # Verify database update was called with correct data + mock_prisma_client.db.litellm_usertable.update.assert_called_once() + call_args = mock_prisma_client.db.litellm_usertable.update.call_args + assert call_args[1]["where"] == {"user_id": "test-user"} + assert call_args[1]["data"]["user_email"] == "updated@example.com" + assert call_args[1]["data"]["teams"] == ["new-team"] + + +@pytest.mark.asyncio +async def test_update_user_not_found(mocker): + """Should raise 404 when user doesn't exist""" + scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + userName="nonexistent-user", + name=SCIMUserName(familyName="User", givenName="Test"), + emails=[SCIMUserEmail(value="test@example.com")], + ) + + # Mock dependencies to raise HTTPException for user not found + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mocker.MagicMock()) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "User not found"})) + ) + + # Should raise ProxyException (which wraps the HTTPException) + with pytest.raises(ProxyException): + await update_user(user_id="nonexistent-user", user=scim_user) + + +@pytest.mark.asyncio +async def test_patch_user_success(mocker): + """Should successfully patch user with PATCH request""" + # Mock existing user + existing_user = mocker.MagicMock() + existing_user.teams = ["team1"] + existing_user.metadata = {} + + # Mock updated user + updated_user = { + "user_id": "test-user", + "user_alias": "Patched User", + "teams": ["team1", "team2"], + "metadata": '{"scim_metadata": {}}' + } + + # Mock patch operations + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="replace", path="displayName", value="Patched User"), + SCIMPatchOperation(op="add", path="groups", value=[{"value": "team2"}]) + ] + ) + + # Mock response SCIM user + response_scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="test-user", + userName="test-user", + name=SCIMUserName(familyName="User", givenName="Patched"), + ) + + # Mock prisma client + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) + + # Mock dependencies + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(return_value=existing_user) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock() + ) + mock_transform = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=response_scim_user) + ) + + # Call patch_user + result = await patch_user(user_id="test-user", patch_ops=patch_ops) + + # Verify result + assert result == response_scim_user + + # Verify database update was called + mock_prisma_client.db.litellm_usertable.update.assert_called_once() + call_args = mock_prisma_client.db.litellm_usertable.update.call_args + assert call_args[1]["where"] == {"user_id": "test-user"} + + +@pytest.mark.asyncio +async def test_patch_user_not_found(mocker): + """Should raise 404 when user doesn't exist for patch""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="replace", path="displayName", value="New Name") + ] + ) + + # Mock dependencies to raise HTTPException for user not found + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mocker.MagicMock()) + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "User not found"})) + ) + + # Should raise ProxyException (which wraps the HTTPException) + with pytest.raises(ProxyException): + await patch_user(user_id="nonexistent-user", patch_ops=patch_ops)