diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index efe18cb68ad..a65500c80dc 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -203,7 +203,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): start_time=start_time, end_time=end_time, ) - + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): await self._async_log_event_base( kwargs=kwargs, @@ -212,7 +212,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): end_time=end_time, ) pass - async def _async_log_event_base(self, kwargs, response_obj, start_time, end_time): try: @@ -242,7 +241,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): verbose_logger.exception(f"s3 Layer Error - {str(e)}") pass - async def async_upload_data_to_s3( self, batch_logging_element: s3BatchLoggingElement ): @@ -277,8 +275,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): # Prepare the URL url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}" - if self.s3_endpoint_url: - url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key + if self.s3_endpoint_url and self.s3_bucket_name: + url = ( + self.s3_endpoint_url + + "/" + + self.s3_bucket_name + + "/" + + batch_logging_element.s3_object_key + ) # Convert JSON to string json_string = safe_dumps(batch_logging_element.payload) @@ -420,8 +424,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): # Prepare the URL url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}" - if self.s3_endpoint_url: - url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key + if self.s3_endpoint_url and self.s3_bucket_name: + url = ( + self.s3_endpoint_url + + "/" + + self.s3_bucket_name + + "/" + + batch_logging_element.s3_object_key + ) # Convert JSON to string json_string = safe_dumps(batch_logging_element.payload) @@ -462,14 +472,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): except Exception as e: verbose_logger.exception(f"Error uploading to s3: {str(e)}") - async def _download_object_from_s3(self, s3_object_key: str) -> Optional[dict]: """ Download and parse JSON object from S3. - + Args: s3_object_key: The S3 object key to download - + Returns: Optional[dict]: The parsed JSON object or None if not found/error """ @@ -481,7 +490,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): from botocore.awsrequest import AWSRequest except ImportError: raise ImportError("Missing boto3 to call S3. Run 'pip install boto3'.") - + try: from litellm.litellm_core_utils.asyncify import asyncify @@ -506,8 +515,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): # Prepare the URL url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{s3_object_key}" - if self.s3_endpoint_url: - url = self.s3_endpoint_url + "/" + s3_object_key + if self.s3_endpoint_url and self.s3_bucket_name: + url = ( + self.s3_endpoint_url + + "/" + + self.s3_bucket_name + + "/" + + s3_object_key + ) # Prepare the request for GET operation # For GET requests, we need x-amz-content-sha256 with hash of empty string @@ -533,12 +548,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): response = await self.async_httpx_client.get(url, headers=signed_headers) if response.status_code != 200: - verbose_logger.exception("S3 object not found, saw response=", response.text) + verbose_logger.exception( + "S3 object not found, saw response=", response.text + ) return None - + # Parse JSON response return response.json() - + except Exception as e: verbose_logger.exception(f"Error downloading from S3: {str(e)}") return None @@ -551,11 +568,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): Get the proxy server request from cold storage Allows fetching a dict of the proxy server request from s3 or GCS bucket. - + Args: request_id: The unique request ID to search for start_time: The start time of the request (datetime or ISO string) - + Returns: Optional[dict]: The request data dictionary or None if not found """ @@ -564,5 +581,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): downloaded_object = await self._download_object_from_s3(object_key) return downloaded_object except Exception as e: - verbose_logger.exception(f"Error retrieving object {object_key} from cold storage: {str(e)}") - return None \ No newline at end of file + verbose_logger.exception( + f"Error retrieving object {object_key} from cold storage: {str(e)}" + ) + return None diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 327b269d1d4..c59e3bb24e8 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -28,6 +28,7 @@ from litellm.types.files import ( get_file_type_from_extension, is_gemini_1_5_accepted_file_type, ) +from litellm.types.utils import LlmProviders from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionAssistantMessage, @@ -492,7 +493,8 @@ def _transform_request_body( data["generationConfig"] = generation_config if cached_content is not None: data["cachedContent"] = cached_content - if labels is not None: + # Only add labels for Vertex AI endpoints (not Google GenAI/AI Studio) and only if non-empty + if labels and custom_llm_provider != LlmProviders.GEMINI: data["labels"] = labels except Exception as e: raise e @@ -647,3 +649,5 @@ def _transform_system_message( return SystemInstructions(parts=system_content_blocks), messages return None, messages + + diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index f84b4df42d0..6720f0c3b71 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -18,6 +18,7 @@ from fastapi import ( Response, ) from typing_extensions import TypedDict +from pydantic import BaseModel import litellm from litellm._logging import verbose_proxy_logger @@ -29,6 +30,7 @@ from litellm.proxy._types import ( Member, NewTeamRequest, NewUserRequest, + NewUserResponse, TeamMemberAddRequest, TeamMemberDeleteRequest, UserAPIKeyAuth, @@ -101,6 +103,13 @@ class ScimUserData(TypedDict): active: Optional[bool] +class GroupMemberExtractionResult(BaseModel): + """Result of extracting and processing group members.""" + existing_member_ids: List[str] + created_users: List[NewUserResponse] + all_member_ids: List[str] # existing + newly created + + scim_router = APIRouter( prefix="/scim/v2", tags=["✨ SCIM v2 (Enterprise Only)"], @@ -190,21 +199,47 @@ def _build_scim_metadata(given_name: Optional[str], family_name: Optional[str], return metadata -async def _extract_group_member_ids(group: SCIMGroup) -> List[str]: - """Extract valid member IDs from SCIMGroup, verifying users exist.""" +async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult: + """ + Extract member IDs from SCIMGroup, creating users that don't exist. + + Returns: + GroupMemberExtractionResult with existing members, created users, and all member IDs + """ prisma_client = await _get_prisma_client_or_raise_exception() - member_ids = [] + existing_member_ids = [] + created_users = [] + all_member_ids = [] if group.members: for member in group.members: + user_id = member.value + # Check if user exists user = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": member.value} + where={"user_id": user_id} ) + if user: - member_ids.append(member.value) + existing_member_ids.append(user_id) + all_member_ids.append(user_id) + else: + # Create the user if they don't exist using our helper + created_user = await _create_user_if_not_exists( + user_id=user_id, + created_via="scim_group_membership" + ) + + if created_user: + created_users.append(created_user) + all_member_ids.append(user_id) + # If creation failed, user is skipped (logged in helper) - return member_ids + return GroupMemberExtractionResult( + existing_member_ids=existing_member_ids, + created_users=created_users, + all_member_ids=all_member_ids + ) async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: @@ -239,6 +274,51 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str ) +async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_group") -> Optional[NewUserResponse]: + """ + Helper function to create a user if they don't exist. + + Args: + user_id: The user ID to create + created_via: Context for where the user was created from + + Returns: + LiteLLM_UserTable if user was created, None if creation failed + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import new_user + + try: + # Get default role for new internal users + default_role: Optional[ + Literal[ + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + LitellmUserRoles.INTERNAL_USER, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + ] + ] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + if litellm.default_internal_user_params: + default_role = litellm.default_internal_user_params.get("user_role") + + new_user_request = NewUserRequest( + user_id=user_id, + user_email=user_id, # We don't have email from group membership + user_alias=None, + teams=[], # Teams will be added separately + metadata={"created_via": created_via}, + auto_create_key=False, + user_role=default_role, + ) + + created_user = await new_user(data=new_user_request) + verbose_proxy_logger.info(f"Created user {user_id} via {created_via}") + return created_user + + except Exception as e: + verbose_proxy_logger.exception(f"Failed to create user {user_id}: {e}") + return None + + async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[str]: """ Get the IDs of the members from a team. @@ -256,6 +336,8 @@ async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[s member_user_ids.append(user_id) return member_user_ids + + # 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""" @@ -914,9 +996,9 @@ async def create_group( detail={"error": f"Group already exists with ID: {team_id}"}, ) - # 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] + # Extract and process group members (creating users that don't exist) + member_result = await _extract_group_member_ids(group) + members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_result.all_member_ids] # Create team in database created_team = await new_team( @@ -959,9 +1041,10 @@ async def update_group( prisma_client = await _get_prisma_client_or_raise_exception() existing_team = await _check_team_exists(group_id) - # Extract valid member IDs - member_ids = await _extract_group_member_ids(group) - verbose_proxy_logger.debug(f"SCIM PUT GROUP member_ids: {member_ids}") + # Extract and process group members (creating users that don't exist) + member_result = await _extract_group_member_ids(group) + verbose_proxy_logger.debug(f"SCIM PUT GROUP all_member_ids: {member_result.all_member_ids}") + verbose_proxy_logger.debug(f"SCIM PUT GROUP created_users: {len(member_result.created_users)}") # Prepare update data existing_metadata = existing_team.metadata if existing_team.metadata else {} @@ -978,10 +1061,10 @@ async def update_group( data=update_data, ) - # Handle user-team relationship changes using the same approach as patch_group + # Handle user-team relationship changes current_members = set(await _get_team_member_user_ids_from_team(existing_team)) verbose_proxy_logger.debug(f"SCIM PUT GROUP current_members: {current_members}") - final_members = set(member_ids) + final_members = set(member_result.all_member_ids) verbose_proxy_logger.debug(f"SCIM PUT GROUP final_members: {final_members}") await _handle_group_membership_changes( @@ -1075,7 +1158,7 @@ async def _process_group_patch_operations( elif path.startswith("members"): # Handle member operations member_values = _extract_group_values(value) - # Validate that users exist + # Create users that don't exist and get all valid member IDs valid_members = [] for member_id in member_values: user = await prisma_client.db.litellm_usertable.find_unique( @@ -1083,6 +1166,16 @@ async def _process_group_patch_operations( ) if user: valid_members.append(member_id) + else: + # Create the user if they don't exist using our helper + created_user = await _create_user_if_not_exists( + user_id=member_id, + created_via="scim_group_patch" + ) + + if created_user: + valid_members.append(member_id) + # If creation failed, user is skipped (logged in helper) if op_type == "replace": final_members = set(valid_members) diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/test_litellm/integrations/test_s3_v2.py index ea14a4c8017..d31d783ba0a 100644 --- a/tests/test_litellm/integrations/test_s3_v2.py +++ b/tests/test_litellm/integrations/test_s3_v2.py @@ -10,6 +10,7 @@ from litellm.types.utils import StandardLoggingPayload class TestS3V2UnitTests: """Test that S3 v2 integration only uses safe_dumps and not json.dumps""" + def test_s3_v2_source_code_analysis(self): """Test that S3 v2 source code only imports and uses safe_dumps""" import inspect @@ -18,7 +19,139 @@ class TestS3V2UnitTests: # Get the source code of the s3_v2 module source_code = inspect.getsource(s3_v2) - + # Verify that json.dumps is not used directly in the code - assert "json.dumps(" not in source_code, \ - "S3 v2 should not use json.dumps directly" \ No newline at end of file + assert ( + "json.dumps(" not in source_code + ), "S3 v2 should not use json.dumps directly" + + @patch('asyncio.create_task') + @patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush') + def test_s3_v2_endpoint_url(self, mock_periodic_flush, mock_create_task): + """testing s3 endpoint url""" + from unittest.mock import AsyncMock, MagicMock + from litellm.types.integrations.s3_v2 import s3BatchLoggingElement + + # Mock periodic_flush and create_task to prevent async task creation during init + mock_periodic_flush.return_value = None + mock_create_task.return_value = None + + # Mock response for all tests + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.raise_for_status = MagicMock() + + # Create a test batch logging element + test_element = s3BatchLoggingElement( + s3_object_key="2025-09-14/test-key.json", + payload={"test": "data"}, + s3_object_download_filename="test-file.json" + ) + + # Test 1: Custom endpoint URL with bucket name + s3_logger = S3Logger( + s3_bucket_name="test-bucket", + s3_endpoint_url="https://s3.amazonaws.com", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1" + ) + + s3_logger.async_httpx_client = AsyncMock() + s3_logger.async_httpx_client.put.return_value = mock_response + + asyncio.run(s3_logger.async_upload_data_to_s3(test_element)) + + call_args = s3_logger.async_httpx_client.put.call_args + assert call_args is not None + url = call_args[0][0] + expected_url = "https://s3.amazonaws.com/test-bucket/2025-09-14/test-key.json" + assert url == expected_url, f"Expected URL {expected_url}, got {url}" + + # Test 2: MinIO-compatible endpoint + s3_logger_minio = S3Logger( + s3_bucket_name="litellm-logs", + s3_endpoint_url="https://minio.example.com:9000", + s3_aws_access_key_id="minio-key", + s3_aws_secret_access_key="minio-secret", + s3_region_name="us-east-1" + ) + + s3_logger_minio.async_httpx_client = AsyncMock() + s3_logger_minio.async_httpx_client.put.return_value = mock_response + + asyncio.run(s3_logger_minio.async_upload_data_to_s3(test_element)) + + call_args_minio = s3_logger_minio.async_httpx_client.put.call_args + assert call_args_minio is not None + url_minio = call_args_minio[0][0] + expected_minio_url = "https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json" + assert url_minio == expected_minio_url, f"Expected MinIO URL {expected_minio_url}, got {url_minio}" + + # Test 3: Custom endpoint without bucket name (should fall back to default) + s3_logger_no_bucket = S3Logger( + s3_endpoint_url="https://s3.amazonaws.com", + s3_aws_access_key_id="test-key", + s3_aws_secret_access_key="test-secret", + s3_region_name="us-east-1" + ) + + s3_logger_no_bucket.async_httpx_client = AsyncMock() + s3_logger_no_bucket.async_httpx_client.put.return_value = mock_response + + asyncio.run(s3_logger_no_bucket.async_upload_data_to_s3(test_element)) + + call_args_no_bucket = s3_logger_no_bucket.async_httpx_client.put.call_args + assert call_args_no_bucket is not None + url_no_bucket = call_args_no_bucket[0][0] + # Should use default S3 URL format when bucket is missing (bucket becomes None in URL) + assert "s3.us-east-1.amazonaws.com" in url_no_bucket + assert "https://" in url_no_bucket + # Should not include the custom endpoint since bucket is missing + assert "https://s3.amazonaws.com/" not in url_no_bucket + + # Test 4: Sync upload method with custom endpoint + s3_logger_sync = S3Logger( + s3_bucket_name="sync-bucket", + s3_endpoint_url="https://custom.s3.endpoint.com", + s3_aws_access_key_id="sync-key", + s3_aws_secret_access_key="sync-secret", + s3_region_name="us-east-1" + ) + + mock_sync_client = MagicMock() + mock_sync_client.put.return_value = mock_response + + with patch('litellm.integrations.s3_v2._get_httpx_client', return_value=mock_sync_client): + s3_logger_sync.upload_data_to_s3(test_element) + + call_args_sync = mock_sync_client.put.call_args + assert call_args_sync is not None + url_sync = call_args_sync[0][0] + expected_sync_url = "https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json" + assert url_sync == expected_sync_url, f"Expected sync URL {expected_sync_url}, got {url_sync}" + + # Test 5: Download method with custom endpoint + s3_logger_download = S3Logger( + s3_bucket_name="download-bucket", + s3_endpoint_url="https://download.s3.endpoint.com", + s3_aws_access_key_id="download-key", + s3_aws_secret_access_key="download-secret", + s3_region_name="us-east-1" + ) + + mock_download_response = MagicMock() + mock_download_response.status_code = 200 + mock_download_response.json = MagicMock(return_value={"downloaded": "data"}) + s3_logger_download.async_httpx_client = AsyncMock() + s3_logger_download.async_httpx_client.get.return_value = mock_download_response + + result = asyncio.run(s3_logger_download._download_object_from_s3("2025-09-14/download-test-key.json")) + + call_args_download = s3_logger_download.async_httpx_client.get.call_args + assert call_args_download is not None + url_download = call_args_download[0][0] + expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json" + assert url_download == expected_download_url, f"Expected download URL {expected_download_url}, got {url_download}" + + assert result == {"downloaded": "data"} \ No newline at end of file diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index d6d33258576..4da2976e1f9 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,4 +1,7 @@ -from litellm.llms.vertex_ai.gemini.transformation import check_if_part_exists_in_parts +from litellm.llms.vertex_ai.gemini.transformation import ( + check_if_part_exists_in_parts, + _transform_request_body, +) def test_check_if_part_exists_in_parts(): @@ -73,3 +76,82 @@ def test_check_if_part_exists_in_parts_camel_case_snake_case(): } assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing) + + +# Tests for issue #14556: Labels field provider-aware filtering +def test_google_genai_excludes_labels(): + """Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'""" + messages = [{"role": "user", "content": "test"}] + optional_params = {"labels": {"project": "test", "team": "ai"}} + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="gemini", + litellm_params=litellm_params, + cached_content=None, + ) + + # Google GenAI/AI Studio should NOT include labels + assert "labels" not in result + assert "contents" in result + + +def test_vertex_ai_includes_labels(): + """Test that Vertex AI endpoints include labels when custom_llm_provider='vertex_ai'""" + messages = [{"role": "user", "content": "test"}] + optional_params = {"labels": {"project": "test", "team": "ai"}} + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params=litellm_params, + cached_content=None, + ) + + # Vertex AI SHOULD include labels + assert "labels" in result + assert result["labels"] == {"project": "test", "team": "ai"} + + + +def test_metadata_to_labels_vertex_only(): + """Test that metadata->labels conversion only happens for Vertex AI""" + messages = [{"role": "user", "content": "test"}] + optional_params = {} + litellm_params = { + "metadata": { + "requester_metadata": { + "user": "john_doe", + "project": "test-project" + } + } + } + + # Google GenAI/AI Studio should not include labels from metadata + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params.copy(), + custom_llm_provider="gemini", + litellm_params=litellm_params.copy(), + cached_content=None, + ) + assert "labels" not in result + + # Vertex AI should include labels from metadata + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params.copy(), + custom_llm_provider="vertex_ai", + litellm_params=litellm_params.copy(), + cached_content=None, + ) + assert "labels" in result + assert result["labels"] == {"user": "john_doe", "project": "test-project"} 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 959275787c8..5cbd602268d 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 @@ -7,6 +7,7 @@ from litellm.proxy._types import LitellmUserRoles, NewUserRequest, ProxyExceptio from litellm.proxy.management_endpoints.scim.scim_v2 import ( UserProvisionerHelpers, _handle_team_membership_changes, + create_group, create_user, get_service_provider_config, patch_user, @@ -910,4 +911,280 @@ async def test_update_group_e2e(mocker): assert len(result.members) == 3 # Verify SCIM transformation was called with updated team - ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team) \ No newline at end of file + ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team) + + +@pytest.mark.asyncio +async def test_create_group_with_nonexistent_users_creates_users(mocker): + """ + Test that creating a group with non-existent users creates those users. + This tests the scenario: Group Push ['new user', existing users...] + """ + # Test data + group_id = "test-group-123" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Test Group", + members=[ + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist + SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist + ] + ) + + ######################################################### + # We expect new-user-1 and new-user-2 to be created + ######################################################### + + # Mock prisma client + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + + # Mock team operations - team doesn't exist yet + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + + # Mock user lookup - only existing-user exists + def mock_user_lookup(where): + user_id = where["user_id"] + if user_id == "existing-user": + mock_user = mocker.MagicMock() + mock_user.user_id = user_id + return mock_user + return None # new-user-1 and new-user-2 don't exist + + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) + + # Mock dependencies + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client) + ) + + # Mock new_user function to track user creation + mock_new_user = mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user", + AsyncMock() + ) + + # Mock created users return values + def mock_new_user_side_effect(data): + from litellm.proxy._types import LiteLLM_UserTable + return LiteLLM_UserTable( + user_id=data.user_id, + user_email=data.user_email, + metadata=data.metadata, + teams=data.teams, + user_role=data.user_role + ) + + mock_new_user.side_effect = mock_new_user_side_effect + + # Mock new_team function + mock_created_team = mocker.MagicMock() + mock_created_team.team_id = group_id + mock_created_team.team_alias = "Test Group" + + mock_new_team = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mock_created_team) + ) + + # Mock SCIM transformation + expected_scim_response = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Test Group", + members=[ + SCIMMember(value="existing-user", display="existing-user"), + SCIMMember(value="new-user-1", display="new-user-1"), + SCIMMember(value="new-user-2", display="new-user-2") + ] + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=expected_scim_response) + ) + + # Execute the create_group function + result = await create_group(group=scim_group) + + ######################################################### + # Assert that new-user-1 and new-user-2 were created + ######################################################### + + # Verify that new_user was called exactly twice (for new-user-1 and new-user-2) + assert mock_new_user.call_count == 2 + + # Check the user creation calls + created_user_ids = set() + for call in mock_new_user.call_args_list: + user_request = call.kwargs["data"] + created_user_ids.add(user_request.user_id) + assert user_request.metadata["created_via"] == "scim_group_membership" + assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + assert user_request.auto_create_key is False + assert user_request.teams == [] # Teams added separately + + assert created_user_ids == {"new-user-1", "new-user-2"} + + # Verify team creation was called with all members (existing + created) + mock_new_team.assert_called_once() + team_request = mock_new_team.call_args.kwargs["data"] + assert team_request.team_id == group_id + assert team_request.team_alias == "Test Group" + + # Verify all members are in the team (existing + newly created) + member_user_ids = {member.user_id for member in team_request.members_with_roles} + assert member_user_ids == {"existing-user", "new-user-1", "new-user-2"} + + # Verify response + assert result.id == group_id + assert result.displayName == "Test Group" + assert len(result.members) == 3 + + +@pytest.mark.asyncio +async def test_update_group_with_nonexistent_users_creates_users(mocker): + """ + Test that updating a group with non-existent users creates those users. + This tests the scenario where a group is updated with members that don't exist in user table. + """ + # Test data + group_id = "existing-group-456" + + # Mock existing team + mock_existing_team = mocker.MagicMock() + mock_existing_team.team_id = group_id + mock_existing_team.team_alias = "Old Group Name" + mock_existing_team.members = ["old-user"] + mock_existing_team.members_with_roles = [{"user_id": "old-user", "role": "user"}] + mock_existing_team.metadata = {"existing": "data"} + + # SCIM group update request + scim_group_update = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Updated Group Name", + members=[ + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-3", display="New User 3"), # This user doesn't exist + SCIMMember(value="new-user-4", display="New User 4"), # This user doesn't exist + ] + ) + + # Mock prisma client + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + + # Mock team operations + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) + + # Mock updated team response + mock_updated_team = mocker.MagicMock() + mock_updated_team.team_id = group_id + mock_updated_team.team_alias = "Updated Group Name" + mock_updated_team.members = ["existing-user", "new-user-3", "new-user-4"] + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) + + # Mock user lookup - only existing-user exists + def mock_user_lookup(where): + user_id = where["user_id"] + if user_id == "existing-user": + mock_user = mocker.MagicMock() + mock_user.user_id = user_id + return mock_user + return None # new-user-3 and new-user-4 don't exist + + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) + + # 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_team_exists", + AsyncMock(return_value=mock_existing_team) + ) + + # Mock new_user function to track user creation + mock_new_user = mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user", + AsyncMock() + ) + + # Mock created users return values + def mock_new_user_side_effect(data): + from litellm.proxy._types import LiteLLM_UserTable + return LiteLLM_UserTable( + user_id=data.user_id, + user_email=data.user_email, + metadata=data.metadata, + teams=data.teams, + user_role=data.user_role + ) + + mock_new_user.side_effect = mock_new_user_side_effect + + # Mock group membership changes + mock_handle_group_membership_changes = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_group_membership_changes", + AsyncMock() + ) + + # Mock SCIM transformation + expected_scim_response = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Updated Group Name", + members=[ + SCIMMember(value="existing-user", display="existing-user"), + SCIMMember(value="new-user-3", display="new-user-3"), + SCIMMember(value="new-user-4", display="new-user-4") + ] + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=expected_scim_response) + ) + + # Execute the update_group function + result = await update_group(group_id=group_id, group=scim_group_update) + + # Verify that new_user was called exactly twice (for new-user-3 and new-user-4) + assert mock_new_user.call_count == 2 + + # Check the user creation calls + created_user_ids = set() + for call in mock_new_user.call_args_list: + user_request = call.kwargs["data"] + created_user_ids.add(user_request.user_id) + assert user_request.metadata["created_via"] == "scim_group_membership" + assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + assert user_request.auto_create_key is False + assert user_request.teams == [] # Teams added separately + + assert created_user_ids == {"new-user-3", "new-user-4"} + + # Verify team update was called + mock_prisma_client.db.litellm_teamtable.update.assert_called_once() + update_call = mock_prisma_client.db.litellm_teamtable.update.call_args + assert update_call[1]["where"]["team_id"] == group_id + assert update_call[1]["data"]["team_alias"] == "Updated Group Name" + + # Verify group membership changes were handled with all members (existing + created) + mock_handle_group_membership_changes.assert_called_once() + membership_call = mock_handle_group_membership_changes.call_args + assert membership_call[1]["group_id"] == group_id + assert membership_call[1]["final_members"] == {"existing-user", "new-user-3", "new-user-4"} + + # Verify response + assert result.id == group_id + assert result.displayName == "Updated Group Name" + assert len(result.members) == 3 \ No newline at end of file