From 748ea4ece9164c0277875a4bb437c54ffd4e9aa8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 14 Mar 2026 15:35:22 -0700 Subject: [PATCH] fix: use safe orjson imports and fix Black formatting for CI Replace hard orjson removal with try/except safe imports that preserve perf when orjson is available. Revert transformation.py to response.json() matching main. Fix test assertions to use safe_dumps() for orjson-agnostic JSON comparison. Apply Black formatting to sidecar files. Co-Authored-By: Claude Opus 4.6 --- .../llms/custom_httpx/sidecar_transport.py | 6 +- .../llms/openai_like/chat/transformation.py | 3 +- litellm/proxy/proxy_server.py | 12 +- litellm/proxy/sidecar_client.py | 8 +- .../integrations/arize/test_arize_utils.py | 81 +- .../agent_endpoints/test_agent_registry.py | 3 +- .../scim/test_scim_v2_endpoints.py | 786 ++++++++++-------- 7 files changed, 494 insertions(+), 405 deletions(-) diff --git a/litellm/llms/custom_httpx/sidecar_transport.py b/litellm/llms/custom_httpx/sidecar_transport.py index 501ff825da7..bc0250a39db 100644 --- a/litellm/llms/custom_httpx/sidecar_transport.py +++ b/litellm/llms/custom_httpx/sidecar_transport.py @@ -68,7 +68,11 @@ class LiteLLMSidecarTransport(httpx.AsyncBaseTransport): provider_base = f"{parsed.scheme}://{parsed.host}" if parsed.port and parsed.port not in (80, 443): provider_base += f":{parsed.port}" - path = parsed.raw_path.decode("ascii") if isinstance(parsed.raw_path, bytes) else str(parsed.raw_path) + path = ( + parsed.raw_path.decode("ascii") + if isinstance(parsed.raw_path, bytes) + else str(parsed.raw_path) + ) # Extract auth header if present api_key = "" diff --git a/litellm/llms/openai_like/chat/transformation.py b/litellm/llms/openai_like/chat/transformation.py index 86d7b97809d..1c8cd574c01 100644 --- a/litellm/llms/openai_like/chat/transformation.py +++ b/litellm/llms/openai_like/chat/transformation.py @@ -2,7 +2,6 @@ OpenAI-like chat completion transformation """ -import json from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union import httpx @@ -95,7 +94,7 @@ class OpenAILikeChatConfig(OpenAIGPTConfig): custom_llm_provider: Optional[str], base_model: Optional[str], ) -> ModelResponse: - response_json = json.loads(response.content) + response_json = response.json() logging_obj.post_call( input=messages, api_key="", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 87fad21e4ea..e7ac1ecfb33 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -946,16 +946,22 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 shared_aiohttp_session = await _initialize_shared_aiohttp_session() ## Initialize Rust sidecar client (optional, for high-perf forwarding) - _use_sidecar = os.environ.get("USE_SIDECAR", "").lower() == "true" or general_settings.get("use_sidecar", False) + _use_sidecar = os.environ.get( + "USE_SIDECAR", "" + ).lower() == "true" or general_settings.get("use_sidecar", False) if _use_sidecar: # Ensure env var is set so AsyncHTTPHandler._should_use_sidecar_transport() picks it up os.environ["USE_SIDECAR"] = "true" from litellm.proxy.sidecar_client import init_sidecar_client - _sidecar_port = int(os.environ.get("SIDECAR_PORT", general_settings.get("sidecar_port", 8787))) + _sidecar_port = int( + os.environ.get("SIDECAR_PORT", general_settings.get("sidecar_port", 8787)) + ) os.environ.setdefault("SIDECAR_PORT", str(_sidecar_port)) - _sidecar_binary = os.environ.get("SIDECAR_BINARY", general_settings.get("sidecar_binary", "")) + _sidecar_binary = os.environ.get( + "SIDECAR_BINARY", general_settings.get("sidecar_binary", "") + ) await init_sidecar_client( port=_sidecar_port, binary=_sidecar_binary or None, diff --git a/litellm/proxy/sidecar_client.py b/litellm/proxy/sidecar_client.py index 15503cf9dce..14666133928 100644 --- a/litellm/proxy/sidecar_client.py +++ b/litellm/proxy/sidecar_client.py @@ -58,9 +58,7 @@ class SidecarClient: self._healthy = await self._check_health() if self._healthy: - verbose_proxy_logger.info( - f"Sidecar client connected to {self.sidecar_url}" - ) + verbose_proxy_logger.info(f"Sidecar client connected to {self.sidecar_url}") else: verbose_proxy_logger.warning( f"Sidecar not available at {self.sidecar_url}, will use fallback" @@ -144,9 +142,7 @@ class SidecarClient: self._process.terminate() try: await asyncio.wait_for( - asyncio.get_event_loop().run_in_executor( - None, self._process.wait - ), + asyncio.get_event_loop().run_in_executor(None, self._process.wait), timeout=5, ) except asyncio.TimeoutError: diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index 9a9f3d5afc7..726fe9e9a26 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -18,6 +18,7 @@ from litellm.integrations._types.open_inference import ( ) from litellm.integrations.arize.arize import ArizeLogger from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.types.utils import Choices, StandardCallbackDynamicParams @@ -88,7 +89,7 @@ def test_arize_set_attributes(): # Metadata attached to the span span.set_attribute.assert_any_call( - SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None}) + SpanAttributes.METADATA, safe_dumps({"key_1": "value_1", "key_2": None}) ) # Basic LLM information @@ -147,7 +148,7 @@ def test_arize_set_attributes(): # Invocation parameters span.set_attribute.assert_any_call( - SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}' + SpanAttributes.LLM_INVOCATION_PARAMETERS, safe_dumps({"user": "test_user"}) ) # User ID @@ -178,10 +179,20 @@ def test_arize_set_attributes_responses_api(): Verifies that multiple output types are correctly handled. """ from unittest.mock import MagicMock - from litellm.types.llms.openai import ResponsesAPIResponse, ResponseAPIUsage, OutputTokensDetails - from openai.types.responses import ResponseReasoningItem, ResponseOutputMessage, ResponseOutputText + + from openai.types.responses import ( + ResponseOutputMessage, + ResponseOutputText, + ResponseReasoningItem, + ) from openai.types.responses.response_reasoning_item import Summary + from litellm.types.llms.openai import ( + OutputTokensDetails, + ResponseAPIUsage, + ResponsesAPIResponse, + ) + span = MagicMock() # Mocked tracing span to test attribute setting # Construct kwargs to simulate a real LLM request scenario @@ -212,11 +223,8 @@ def test_arize_set_attributes_responses_api(): id="reasoning-001", type="reasoning", summary=[ - Summary( - text="First, I need to analyze...", - type="summary_text" - ) - ] + Summary(text="First, I need to analyze...", type="summary_text") + ], ), ResponseOutputMessage( id="msg-001", @@ -229,17 +237,15 @@ def test_arize_set_attributes_responses_api(): text="The answer is 42", type="output_text", ) - ] - ) + ], + ), ], usage=ResponseAPIUsage( input_tokens=120, output_tokens=250, total_tokens=370, - output_tokens_details=OutputTokensDetails( - reasoning_tokens=180 - ) - ) + output_tokens_details=OutputTokensDetails(reasoning_tokens=180), + ), ) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) @@ -247,21 +253,18 @@ def test_arize_set_attributes_responses_api(): # Verify reasoning summary was set (index 0) span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_REASONING_SUMMARY}", - "First, I need to analyze..." + "First, I need to analyze...", ) # Verify message content was set (index 1) - span.set_attribute.assert_any_call( - SpanAttributes.OUTPUT_VALUE, - "The answer is 42" - ) + span.set_attribute.assert_any_call(SpanAttributes.OUTPUT_VALUE, "The answer is 42") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.1.{MessageAttributes.MESSAGE_CONTENT}", - "The answer is 42" + "The answer is 42", ) span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.1.{MessageAttributes.MESSAGE_ROLE}", - "assistant" + "assistant", ) # Verify token counts including reasoning tokens @@ -335,42 +338,34 @@ def test_construct_dynamic_arize_headers(): # Test with all parameters present dynamic_params_full = StandardCallbackDynamicParams( - arize_api_key="test_api_key", - arize_space_id="test_space_id" + arize_api_key="test_api_key", arize_space_id="test_space_id" ) arize_logger = ArizeLogger() - + headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_full) - expected_headers = { - "api_key": "test_api_key", - "arize-space-id": "test_space_id" - } + expected_headers = {"api_key": "test_api_key", "arize-space-id": "test_space_id"} assert headers == expected_headers - + # Test with only space_id dynamic_params_space_id_only = StandardCallbackDynamicParams( arize_space_id="test_space_id" ) - + headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_id_only) - expected_headers = { - "arize-space-id": "test_space_id" - } + expected_headers = {"arize-space-id": "test_space_id"} assert headers == expected_headers - + # Test with empty parameters dict dynamic_params_empty = StandardCallbackDynamicParams() - + headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_empty) assert headers == {} # test with space key and api key dynamic_params_space_key_and_api_key = StandardCallbackDynamicParams( - arize_space_key="test_space_key", - arize_api_key="test_api_key" + arize_space_key="test_space_key", arize_api_key="test_api_key" ) - headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_key_and_api_key) - expected_headers = { - "arize-space-id": "test_space_key", - "api_key": "test_api_key" - } + headers = arize_logger.construct_dynamic_otel_headers( + dynamic_params_space_key_and_api_key + ) + expected_headers = {"arize-space-id": "test_space_key", "api_key": "test_api_key"} diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index ddd9cc09c8e..66ab0145ccf 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -4,6 +4,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry @@ -110,5 +111,5 @@ async def test_update_agent_in_db_preserves_explicit_static_headers_and_extra_he call_kwargs = mock_update.call_args.kwargs update_data = call_kwargs["data"] - assert update_data["static_headers"] == '{"Authorization": "Bearer xyz"}' + assert update_data["static_headers"] == safe_dumps({"Authorization": "Bearer xyz"}) assert update_data["extra_headers"] == ["X-Custom-Header"] 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 1355ca0abbe..4e02097acfe 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 @@ -3,7 +3,13 @@ from unittest.mock import AsyncMock import pytest from fastapi import HTTPException -from litellm.proxy._types import LitellmUserRoles, NewUserRequest, NewUserResponse, ProxyException +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy._types import ( + LitellmUserRoles, + NewUserRequest, + NewUserResponse, + ProxyException, +) from litellm.proxy.management_endpoints.scim.scim_v2 import ( UserProvisionerHelpers, _extract_group_member_ids, @@ -45,14 +51,16 @@ async def test_create_user_existing_user_conflict(mocker): 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"}) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value={"user_id": "existing-user"} + ) # Mock the _get_prisma_client_or_raise_exception to return our mock mocker.patch( "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(), @@ -61,7 +69,7 @@ async def test_create_user_existing_user_conflict(mocker): with pytest.raises(HTTPException) as exc_info: await create_user(user=scim_user) - # Check that it's an HTTPException with status 409 + # Check that it's an HTTPException with status 409 assert exc_info.value.status_code == 409 assert "existing-user" in str(exc_info.value.detail) mocked_new_user.assert_not_called() @@ -84,9 +92,7 @@ async def test_create_user_defaults_to_viewer(mocker, monkeypatch): mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) - monkeypatch.setattr( - "litellm.default_internal_user_params", None, raising=False - ) + monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -211,7 +217,10 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp "BUG: _update_litellm_setting did not update litellm.default_internal_user_params in memory. " "The local variable reassignment (in_memory_var = ...) doesn't propagate back." ) - assert litellm.default_internal_user_params.get("user_role") == LitellmUserRoles.INTERNAL_USER + assert ( + litellm.default_internal_user_params.get("user_role") + == LitellmUserRoles.INTERNAL_USER + ) # Step 3: Create a user via SCIM scim_user = SCIMUser( @@ -256,7 +265,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp 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 @@ -265,12 +274,11 @@ async def test_handle_existing_user_by_email_no_email(mocker): metadata={}, auto_create_key=False, ) - + result = await UserProvisionerHelpers.handle_existing_user_by_email( - prisma_client=mock_prisma_client, - new_user_request=new_user_request + prisma_client=mock_prisma_client, new_user_request=new_user_request ) - + assert result is None @@ -281,21 +289,20 @@ async def test_handle_existing_user_by_email_no_existing_user(mocker): 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", + 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 + 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"} @@ -312,16 +319,16 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): 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_email": "test@example.com", "user_alias": "New Name", "teams": ["new-team"], - "metadata": '{"new": "data"}' + "metadata": '{"new": "data"}', } - + # Mock SCIM user to be returned mock_scim_user = SCIMUser( schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], @@ -330,52 +337,55 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): 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_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) + 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"], + 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 + 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_email": "test@example.com", "user_alias": "New Name", "teams": ["new-team"], - "metadata": '{"new": "data"}', + "metadata": safe_dumps({"new": "data"}), }, ) - + # Verify transformation was called mock_transform.assert_called_once_with(updated_user) @@ -385,16 +395,16 @@ 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() + AsyncMock(), ) - + # Same teams - no changes await _handle_team_membership_changes( user_id="test-user", existing_teams=["team1", "team2"], - new_teams=["team1", "team2"] + new_teams=["team1", "team2"], ) - + # Should not be called since no changes mock_patch_team_membership.assert_not_called() @@ -404,19 +414,19 @@ 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() + AsyncMock(), ) - + # Adding teams await _handle_team_membership_changes( user_id="test-user", existing_teams=["team1"], - new_teams=["team1", "team2", "team3"] + new_teams=["team1", "team2", "team3"], ) - + # Verify the call was made once mock_patch_team_membership.assert_called_once() - + # Check the arguments more flexibly to handle order variations call_args = mock_patch_team_membership.call_args assert call_args[1]["user_id"] == "test-user" @@ -429,19 +439,19 @@ 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() + AsyncMock(), ) - + # Removing teams await _handle_team_membership_changes( user_id="test-user", existing_teams=["team1", "team2", "team3"], - new_teams=["team1"] + new_teams=["team1"], ) - + # Verify the call was made once mock_patch_team_membership.assert_called_once() - + # Check the arguments more flexibly to handle order variations call_args = mock_patch_team_membership.call_args assert call_args[1]["user_id"] == "test-user" @@ -454,19 +464,19 @@ 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() + AsyncMock(), ) - + # Both adding and removing teams await _handle_team_membership_changes( user_id="test-user", existing_teams=["team1", "team2"], - new_teams=["team2", "team3"] + new_teams=["team2", "team3"], ) - + # Verify the call was made once mock_patch_team_membership.assert_called_once() - + # Check the arguments - team1 should be removed, team3 should be added, team2 stays call_args = mock_patch_team_membership.call_args assert call_args[1]["user_id"] == "test-user" @@ -480,25 +490,25 @@ async def test_update_user_success(mocker): # 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"}}' + "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")] + groups=[SCIMUserGroup(value="new-team")], ) - + # Mock SCIM user for response response_scim_user = SCIMUser( schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], @@ -507,37 +517,39 @@ async def test_update_user_success(mocker): 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_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) + AsyncMock(return_value=mock_prisma_client), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", - AsyncMock(return_value=existing_user) + AsyncMock(return_value=existing_user), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", - AsyncMock() + AsyncMock(), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", - AsyncMock(return_value=response_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 @@ -555,17 +567,21 @@ async def test_update_user_not_found(mocker): 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()) + 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"})) + 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) @@ -578,24 +594,24 @@ async def test_patch_user_success(mocker): 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": {}}' + "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"}]) - ] + 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"], @@ -603,37 +619,39 @@ async def test_patch_user_success(mocker): 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_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) + AsyncMock(return_value=mock_prisma_client), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", - AsyncMock(return_value=existing_user) + AsyncMock(return_value=existing_user), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", - AsyncMock() + AsyncMock(), ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", - AsyncMock(return_value=response_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 @@ -647,19 +665,23 @@ async def test_patch_user_not_found(mocker): 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()) + 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"})) + 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) @@ -671,13 +693,15 @@ async def test_get_service_provider_config(mocker): # Mock the Request object mock_request = mocker.MagicMock() mock_request.url = "https://example.com/scim/v2/ServiceProviderConfig" - + # Call the endpoint result = await get_service_provider_config(mock_request) - + # Verify it returns the correct response assert isinstance(result, SCIMServiceProviderConfig) - assert result.schemas == ["urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig"] + assert result.schemas == [ + "urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig" + ] assert result.patch.supported is True assert result.bulk.supported is False assert result.meta is not None @@ -688,9 +712,9 @@ async def test_get_service_provider_config(mocker): async def test_update_group_metadata_serialization_issue(mocker): """ Test that update_group properly serializes metadata to avoid Prisma DataError. - + This test reproduces the issue where metadata was passed as a dict instead of - a JSON string, causing: "Invalid argument type. `metadata` should be of any + a JSON string, causing: "Invalid argument type. `metadata` should be of any of the following types: `JsonNullValueInput`, `Json`" """ from litellm.proxy.management_endpoints.scim.scim_v2 import update_group @@ -702,9 +726,9 @@ async def test_update_group_metadata_serialization_issue(mocker): schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], id=group_id, displayName="Test Group", - members=[SCIMMember(value="user1", display="User One")] + members=[SCIMMember(value="user1", display="User One")], ) - + # Mock existing team with metadata mock_existing_team = mocker.MagicMock() mock_existing_team.team_id = group_id @@ -713,7 +737,7 @@ async def test_update_group_metadata_serialization_issue(mocker): mock_existing_team.metadata = {"existing_key": "existing_value"} mock_existing_team.created_at = None mock_existing_team.updated_at = None - + # Mock updated team response mock_updated_team = mocker.MagicMock() mock_updated_team.team_id = group_id @@ -721,63 +745,72 @@ async def test_update_group_metadata_serialization_issue(mocker): mock_updated_team.members = ["user1"] mock_updated_team.created_at = None mock_updated_team.updated_at = None - + # 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_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_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) - + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_existing_team + ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock( + return_value=mock_updated_team + ) + # Mock user operations mock_user = mocker.MagicMock() mock_user.user_id = "user1" mock_user.user_email = "user1@example.com" # Add proper string value for user_email mock_user.teams = [group_id] - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mock_user + ) mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=mock_user) - + # Mock the _get_prisma_client_or_raise_exception to return our mock mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", AsyncMock(return_value=mock_prisma_client), ) - + # Mock the transformation function mock_scim_group_response = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], id=group_id, displayName="Test Group", - members=[SCIMMember(value="user1", display="User One")] + members=[SCIMMember(value="user1", display="User One")], ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", AsyncMock(return_value=mock_scim_group_response), ) - + # Call the function that had the bug await update_group(group_id=group_id, group=scim_group) - + # Verify the team update was called mock_prisma_client.db.litellm_teamtable.update.assert_called_once() - + # Get the call arguments to verify metadata serialization call_args = mock_prisma_client.db.litellm_teamtable.update.call_args update_data = call_args[1]["data"] - + # Verify that metadata is properly serialized as a string, not a dict # This is the critical check that would have caught the original bug assert "metadata" in update_data metadata = update_data["metadata"] - + # The fix should ensure metadata is serialized as a JSON string - assert isinstance(metadata, str), f"metadata should be a JSON string, but got {type(metadata)}" - + assert isinstance( + metadata, str + ), f"metadata should be a JSON string, but got {type(metadata)}" + # Verify we can parse it back to verify it contains the expected data import json + parsed_metadata = json.loads(metadata) assert "existing_key" in parsed_metadata assert "scim_data" in parsed_metadata @@ -788,7 +821,7 @@ async def test_team_membership_management(mocker): """ Test that team membership changes work correctly: - Adding members to team - - Removing members from team + - Removing members from team - members_with_roles is used as source of truth """ from litellm.proxy._types import Member @@ -801,51 +834,55 @@ async def test_team_membership_management(mocker): mock_team = mocker.MagicMock() mock_team.members_with_roles = [ Member(user_id="user1", role="user"), - Member(user_id="user2", role="user") + Member(user_id="user2", role="user"), ] mock_team.members = ["user1", "user2", "user3"] # This should be ignored - + # Test that members_with_roles is source of truth member_ids = await _get_team_member_user_ids_from_team(mock_team) assert set(member_ids) == {"user1", "user2"} assert "user3" not in member_ids # Should not be included even though in members - + # Mock patch_team_membership function mock_patch_team_membership = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", - AsyncMock() + AsyncMock(), ) - + # Test adding and removing members group_id = "test-group-id" current_members = {"user1", "user2"} final_members = {"user2", "user3", "user4"} # Remove user1, add user3 and user4 - + await _handle_group_membership_changes( - group_id=group_id, - current_members=current_members, - final_members=final_members + group_id=group_id, current_members=current_members, final_members=final_members ) - + # Verify patch_team_membership was called correctly assert mock_patch_team_membership.call_count == 3 - + # Check calls for adding members - add_calls = [call for call in mock_patch_team_membership.call_args_list - if call[1]["teams_ids_to_add_user_to"] == [group_id]] + add_calls = [ + call + for call in mock_patch_team_membership.call_args_list + if call[1]["teams_ids_to_add_user_to"] == [group_id] + ] assert len(add_calls) == 2 # user3 and user4 - + add_user_ids = {call[1]["user_id"] for call in add_calls} assert add_user_ids == {"user3", "user4"} - - # Check calls for removing members - remove_calls = [call for call in mock_patch_team_membership.call_args_list - if call[1]["teams_ids_to_remove_user_from"] == [group_id]] + + # Check calls for removing members + remove_calls = [ + call + for call in mock_patch_team_membership.call_args_list + if call[1]["teams_ids_to_remove_user_from"] == [group_id] + ] assert len(remove_calls) == 1 # user1 - + remove_user_ids = {call[1]["user_id"] for call in remove_calls} assert remove_user_ids == {"user1"} - + # Verify all calls have correct structure for call in mock_patch_team_membership.call_args_list: assert "user_id" in call[1] @@ -854,7 +891,9 @@ async def test_team_membership_management(mocker): # Each call should either add OR remove, not both add_teams = call[1]["teams_ids_to_add_user_to"] remove_teams = call[1]["teams_ids_to_remove_user_from"] - assert (len(add_teams) > 0) != (len(remove_teams) > 0) # XOR - one should be empty + assert (len(add_teams) > 0) != ( + len(remove_teams) > 0 + ) # XOR - one should be empty @pytest.mark.asyncio @@ -873,7 +912,7 @@ async def test_update_group_e2e(mocker): # Setup test data group_id = "test-team-123" - + # Mock existing team in database existing_team = LiteLLM_TeamTable( team_id=group_id, @@ -881,11 +920,11 @@ async def test_update_group_e2e(mocker): members=["user1", "user2"], # This should be ignored members_with_roles=[ Member(user_id="user1", role="user"), - Member(user_id="user2", role="user") + Member(user_id="user2", role="user"), ], - metadata={"existing_key": "existing_value"} + metadata={"existing_key": "existing_value"}, ) - + # Mock updated SCIM group request scim_group_update = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], @@ -894,19 +933,21 @@ async def test_update_group_e2e(mocker): members=[ SCIMMember(value="user2", display="User Two"), # Keep user2 SCIMMember(value="user3", display="User Three"), # Add user3 - SCIMMember(value="user4", display="User Four") # Add user4 - ] + SCIMMember(value="user4", display="User Four"), # Add user4 + ], ) - + # 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 database operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) - + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=existing_team + ) + # Mock the updated team that gets returned from database updated_team = LiteLLM_TeamTable( team_id=group_id, @@ -915,32 +956,36 @@ async def test_update_group_e2e(mocker): members_with_roles=[ Member(user_id="user2", role="user"), Member(user_id="user3", role="user"), - Member(user_id="user4", role="user") + Member(user_id="user4", role="user"), ], metadata={ "existing_key": "existing_value", - "scim_data": scim_group_update.model_dump() - } + "scim_data": scim_group_update.model_dump(), + }, ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) - + mock_prisma_client.db.litellm_teamtable.update = AsyncMock( + return_value=updated_team + ) + # Mock user validation (all users exist) mock_user = mocker.MagicMock() mock_user.user_id = "test-user" - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) - + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mock_user + ) + # Mock dependencies mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", - AsyncMock(return_value=mock_prisma_client) + AsyncMock(return_value=mock_prisma_client), ) - + # Mock patch_team_membership to track membership changes mock_patch_team_membership = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", - AsyncMock() + AsyncMock(), ) - + # Mock SCIM transformation expected_scim_response = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], @@ -948,67 +993,78 @@ async def test_update_group_e2e(mocker): displayName="Updated Team Name", members=[ SCIMMember(value="user2", display="user2"), - SCIMMember(value="user3", display="user3"), - SCIMMember(value="user4", display="user4") - ] + SCIMMember(value="user3", display="user3"), + SCIMMember(value="user4", display="user4"), + ], ) mocker.patch.object( ScimTransformations, "transform_litellm_team_to_scim_group", - AsyncMock(return_value=expected_scim_response) + AsyncMock(return_value=expected_scim_response), ) - + # Execute the update_group function result = await update_group(group_id=group_id, group=scim_group_update) - + # Verify database update was called with correct data mock_prisma_client.db.litellm_teamtable.update.assert_called_once() update_call_args = mock_prisma_client.db.litellm_teamtable.update.call_args - + # Check the update parameters assert update_call_args[1]["where"]["team_id"] == group_id update_data = update_call_args[1]["data"] assert update_data["team_alias"] == "Updated Team Name" - + # Verify metadata includes both existing data and SCIM data metadata_str = update_data["metadata"] import json + metadata = json.loads(metadata_str) assert metadata["existing_key"] == "existing_value" assert "scim_data" in metadata assert metadata["scim_data"]["displayName"] == "Updated Team Name" - + # Verify team membership changes were handled correctly - assert mock_patch_team_membership.call_count == 3 # Remove user1, add user3, add user4 - + assert ( + mock_patch_team_membership.call_count == 3 + ) # Remove user1, add user3, add user4 + # Check membership changes call_args_list = mock_patch_team_membership.call_args_list - + # Find remove operation (user1) - remove_calls = [call for call in call_args_list - if call[1]["teams_ids_to_remove_user_from"] == [group_id]] + remove_calls = [ + call + for call in call_args_list + if call[1]["teams_ids_to_remove_user_from"] == [group_id] + ] assert len(remove_calls) == 1 assert remove_calls[0][1]["user_id"] == "user1" assert remove_calls[0][1]["teams_ids_to_add_user_to"] == [] - + # Find add operations (user3, user4) - add_calls = [call for call in call_args_list - if call[1]["teams_ids_to_add_user_to"] == [group_id]] + add_calls = [ + call + for call in call_args_list + if call[1]["teams_ids_to_add_user_to"] == [group_id] + ] assert len(add_calls) == 2 add_user_ids = {call[1]["user_id"] for call in add_calls} assert add_user_ids == {"user3", "user4"} - + # Verify all add calls have empty remove lists for call in add_calls: assert call[1]["teams_ids_to_remove_user_from"] == [] - + # Verify the response assert result.id == group_id assert result.displayName == "Updated Team Name" 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) + ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with( + updated_team + ) @pytest.mark.asyncio @@ -1018,17 +1074,15 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): Per SCIM 2.0 protocol, users must exist before being added to groups. This prevents security issues where users not assigned to app get provisioned via group membership. """ + # Mock the feature flag to False (SCIM 2.0 strict mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": False - } - } - + return {"litellm_settings": {"scim_upsert_user": False}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data group_id = "test-group-123" scim_group = SCIMGroup( @@ -1036,25 +1090,31 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): 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 - ] + 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 the request to be rejected with 400 error ######################################################### - + # 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"] @@ -1063,23 +1123,27 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): 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_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) + AsyncMock(return_value=mock_prisma_client), ) - + # Execute the create_group function - should raise ProxyException with pytest.raises(ProxyException) as exc_info: await create_group(group=scim_group) - + # Verify it's a 400 Bad Request assert int(exc_info.value.code) == 400 assert "does not exist" in str(exc_info.value.message) - assert "new-user-1" in str(exc_info.value.message) or "new-user-2" in str(exc_info.value.message) + assert "new-user-1" in str(exc_info.value.message) or "new-user-2" in str( + exc_info.value.message + ) @pytest.mark.asyncio @@ -1088,20 +1152,18 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): Test that updating a group with non-existent users is rejected when scim_upsert_user is False. Per SCIM 2.0 protocol, users must exist before being added to groups. """ + # Mock the feature flag to False (SCIM 2.0 strict mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": False - } - } - + return {"litellm_settings": {"scim_upsert_user": False}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data group_id = "existing-group-456" - + # Mock existing team mock_existing_team = mocker.MagicMock() mock_existing_team.team_id = group_id @@ -1109,35 +1171,45 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): 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 - ] + 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_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_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"] @@ -1146,47 +1218,51 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): 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_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) + AsyncMock(return_value=mock_prisma_client), ) - + mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists", - AsyncMock(return_value=mock_existing_team) + AsyncMock(return_value=mock_existing_team), ) - + # Execute the update_group function - should raise ProxyException with pytest.raises(ProxyException) as exc_info: await update_group(group_id=group_id, group=scim_group_update) - + # Verify it's a 400 Bad Request assert int(exc_info.value.code) == 400 assert "does not exist" in str(exc_info.value.message) - assert "new-user-3" in str(exc_info.value.message) or "new-user-4" in str(exc_info.value.message) + assert "new-user-3" in str(exc_info.value.message) or "new-user-4" in str( + exc_info.value.message + ) @pytest.mark.asyncio -async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker, monkeypatch): +async def test_create_group_with_nonexistent_users_creates_when_flag_true( + mocker, monkeypatch +): """ Test that creating a group with non-existent users creates them when scim_upsert_user is True. This preserves backward compatible behavior. """ + # Mock the feature flag to True (backward compatible mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": True - } - } - + return {"litellm_settings": {"scim_upsert_user": True}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data group_id = "test-group-123" scim_group = SCIMGroup( @@ -1194,21 +1270,27 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker 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 - should be created - SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist - should be created - ] + SCIMMember( + value="existing-user", display="Existing User" + ), # This user exists + SCIMMember( + value="new-user-1", display="New User 1" + ), # This user doesn't exist - should be created + SCIMMember( + value="new-user-2", display="New User 2" + ), # This user doesn't exist - should 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 initially def mock_user_lookup(where): user_id = where["user_id"] @@ -1217,87 +1299,93 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker 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_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=mock_user_lookup + ) + # Mock user creation created_user_1 = NewUserResponse(user_id="new-user-1", key="test-key-1") created_user_2 = NewUserResponse(user_id="new-user-2", key="test-key-2") mock_create_user = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", - AsyncMock(side_effect=[created_user_1, created_user_2]) + AsyncMock(side_effect=[created_user_1, created_user_2]), ) - + # Mock new_team mock_team = mocker.MagicMock() mock_team.team_id = group_id mock_new_team = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.new_team", - AsyncMock(return_value=mock_team) + AsyncMock(return_value=mock_team), ) - + # Mock transformation mock_scim_group = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], id=group_id, displayName="Test Group", - members=[] + members=[], ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", - AsyncMock(return_value=mock_scim_group) + AsyncMock(return_value=mock_scim_group), ) - + # Mock dependencies mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", - AsyncMock(return_value=mock_prisma_client) + AsyncMock(return_value=mock_prisma_client), ) - + # Execute the create_group function - should succeed result = await create_group(group=scim_group) - + # Verify users were created assert mock_create_user.call_count == 2 - assert mock_create_user.call_args_list[0].kwargs['user_id'] == "new-user-1" - assert mock_create_user.call_args_list[1].kwargs['user_id'] == "new-user-2" - + assert mock_create_user.call_args_list[0].kwargs["user_id"] == "new-user-1" + assert mock_create_user.call_args_list[1].kwargs["user_id"] == "new-user-2" + # Verify team was created mock_new_team.assert_called_once() @pytest.mark.asyncio -async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, monkeypatch): +async def test_extract_group_member_ids_with_flag_true_creates_users( + mocker, monkeypatch +): """ Test that _extract_group_member_ids creates users when scim_upsert_user is True. """ + # Mock the feature flag to True (backward compatible mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": True - } - } - + return {"litellm_settings": {"scim_upsert_user": True}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data scim_group = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], id="test-group", 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 - should be created - ] + SCIMMember( + value="existing-user", display="Existing User" + ), # This user exists + SCIMMember( + value="new-user-1", display="New User 1" + ), # This user doesn't exist - should be created + ], ) - + # Mock prisma client mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - + # Mock user lookup - only existing-user exists initially def mock_user_lookup(where): user_id = where["user_id"] @@ -1306,35 +1394,36 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon mock_user.user_id = user_id return mock_user return None # new-user-1 doesn't exist - - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) - + + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=mock_user_lookup + ) + # Mock user creation created_user = NewUserResponse(user_id="new-user-1", key="test-key-1") mock_create_user = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", - AsyncMock(return_value=created_user) + AsyncMock(return_value=created_user), ) - + # Mock dependencies mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", - AsyncMock(return_value=mock_prisma_client) + AsyncMock(return_value=mock_prisma_client), ) - + # Execute the function result = await _extract_group_member_ids(scim_group) - + # Verify result assert "existing-user" in result.existing_member_ids assert "existing-user" in result.all_member_ids assert "new-user-1" in result.all_member_ids assert len(result.created_users) == 1 - + # Verify user was created mock_create_user.assert_called_once_with( - user_id="new-user-1", - created_via="scim_group_membership" + user_id="new-user-1", created_via="scim_group_membership" ) @@ -1343,33 +1432,35 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa """ Test that _extract_group_member_ids rejects non-existent users when scim_upsert_user is False. """ + # Mock the feature flag to False (SCIM 2.0 strict mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": False - } - } - + return {"litellm_settings": {"scim_upsert_user": False}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data scim_group = SCIMGroup( schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], id="test-group", 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 - should be rejected - ] + SCIMMember( + value="existing-user", display="Existing User" + ), # This user exists + SCIMMember( + value="new-user-1", display="New User 1" + ), # This user doesn't exist - should be rejected + ], ) - + # Mock prisma client mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - + # Mock user lookup - only existing-user exists def mock_user_lookup(where): user_id = where["user_id"] @@ -1378,19 +1469,21 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa mock_user.user_id = user_id return mock_user return None # new-user-1 doesn't exist - - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) - + + 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) + AsyncMock(return_value=mock_prisma_client), ) - + # Execute the function - should raise HTTPException with pytest.raises(HTTPException) as exc_info: await _extract_group_member_ids(scim_group) - + # Verify it's a 400 Bad Request assert exc_info.value.status_code == 400 assert "does not exist" in str(exc_info.value.detail) @@ -1398,119 +1491,114 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa @pytest.mark.asyncio -async def test_process_group_patch_operations_with_flag_true_creates_users(mocker, monkeypatch): +async def test_process_group_patch_operations_with_flag_true_creates_users( + mocker, monkeypatch +): """ Test that _process_group_patch_operations creates users when scim_upsert_user is True. """ + # Mock the feature flag to True (backward compatible mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": True - } - } - + return {"litellm_settings": {"scim_upsert_user": True}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations=[ SCIMPatchOperation( - op="add", - path="members", - value=[{"value": "new-user-1"}] + op="add", path="members", value=[{"value": "new-user-1"}] ) - ] + ], ) - + # Mock existing team mock_existing_team = mocker.MagicMock() mock_existing_team.members = [] mock_existing_team.metadata = {} - + # Mock prisma client mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - + # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) - + # Mock user creation created_user = NewUserResponse(user_id="new-user-1", key="test-key-1") mock_create_user = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", - AsyncMock(return_value=created_user) + AsyncMock(return_value=created_user), ) - + # Execute the function update_data, final_members = await _process_group_patch_operations( patch_ops=patch_ops, existing_team=mock_existing_team, - prisma_client=mock_prisma_client + prisma_client=mock_prisma_client, ) - + # Verify result assert "new-user-1" in final_members - + # Verify user was created mock_create_user.assert_called_once_with( - user_id="new-user-1", - created_via="scim_group_patch" + user_id="new-user-1", created_via="scim_group_patch" ) @pytest.mark.asyncio -async def test_process_group_patch_operations_with_flag_false_rejects(mocker, monkeypatch): +async def test_process_group_patch_operations_with_flag_false_rejects( + mocker, monkeypatch +): """ Test that _process_group_patch_operations rejects non-existent users when scim_upsert_user is False. """ + # Mock the feature flag to False (SCIM 2.0 strict mode) async def mock_get_config(): - return { - "litellm_settings": { - "scim_upsert_user": False - } - } - + return {"litellm_settings": {"scim_upsert_user": False}} + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - + # Test data patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations=[ SCIMPatchOperation( - op="add", - path="members", - value=[{"value": "new-user-1"}] + op="add", path="members", value=[{"value": "new-user-1"}] ) - ] + ], ) - + # Mock existing team mock_existing_team = mocker.MagicMock() mock_existing_team.members = [] mock_existing_team.metadata = {} - + # Mock prisma client mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - + # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) - + # Execute the function - should raise HTTPException with pytest.raises(HTTPException) as exc_info: await _process_group_patch_operations( patch_ops=patch_ops, existing_team=mock_existing_team, - prisma_client=mock_prisma_client + prisma_client=mock_prisma_client, ) - + # Verify it's a 400 Bad Request assert exc_info.value.status_code == 400 assert "does not exist" in str(exc_info.value.detail)