From 61c50cafebf4d5a982127b083127764635cde629 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 9 Apr 2025 15:27:14 -0700 Subject: [PATCH] testing for msft group assignment --- .../proxy/management_endpoints/test_ui_sso.py | 193 ++++++++++++++++++ 1 file changed, 193 insertions(+) diff --git a/tests/litellm/proxy/management_endpoints/test_ui_sso.py b/tests/litellm/proxy/management_endpoints/test_ui_sso.py index 14b66883618..606f3833beb 100644 --- a/tests/litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/litellm/proxy/management_endpoints/test_ui_sso.py @@ -19,6 +19,10 @@ from litellm.proxy.management_endpoints.ui_sso import ( GoogleSSOHandler, MicrosoftSSOHandler, ) +from litellm.types.proxy.management_endpoints.ui_sso import ( + MicrosoftGraphAPIUserGroupDirectoryObject, + MicrosoftGraphAPIUserGroupResponse, +) def test_microsoft_sso_handler_openid_from_response(): @@ -186,3 +190,192 @@ def test_get_google_callback_response(): assert result.get("sub") == "google123" assert result.get("given_name") == "Google" assert result.get("family_name") == "User" + + +@pytest.mark.asyncio +async def test_get_user_groups_from_graph_api(): + # Arrange + mock_response = { + "@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects", + "value": [ + { + "@odata.type": "#microsoft.graph.group", + "id": "group1", + "displayName": "Group 1", + }, + { + "@odata.type": "#microsoft.graph.group", + "id": "group2", + "displayName": "Group 2", + }, + ], + } + + async def mock_get(*args, **kwargs): + mock = MagicMock() + mock.json.return_value = mock_response + return mock + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + # Act + result = await MicrosoftSSOHandler.get_user_groups_from_graph_api( + access_token="mock_token" + ) + + # Assert + assert isinstance(result, list) + assert len(result) == 2 + assert "group1" in result + assert "group2" in result + + +@pytest.mark.asyncio +async def test_get_user_groups_pagination(): + # Arrange + first_response = { + "@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects", + "@odata.nextLink": "https://graph.microsoft.com/v1.0/me/memberOf?$skiptoken=page2", + "value": [ + { + "@odata.type": "#microsoft.graph.group", + "id": "group1", + "displayName": "Group 1", + }, + ], + } + second_response = { + "@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects", + "value": [ + { + "@odata.type": "#microsoft.graph.group", + "id": "group2", + "displayName": "Group 2", + }, + ], + } + + responses = [first_response, second_response] + current_response = {"index": 0} + + async def mock_get(*args, **kwargs): + mock = MagicMock() + mock.json.return_value = responses[current_response["index"]] + current_response["index"] += 1 + return mock + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + # Act + result = await MicrosoftSSOHandler.get_user_groups_from_graph_api( + access_token="mock_token" + ) + + # Assert + assert isinstance(result, list) + assert len(result) == 2 + assert "group1" in result + assert "group2" in result + assert current_response["index"] == 2 # Verify both pages were fetched + + +@pytest.mark.asyncio +async def test_get_user_groups_empty_response(): + # Arrange + mock_response = { + "@odata.context": "https://graph.microsoft.com/v1.0/$metadata#directoryObjects", + "value": [], + } + + async def mock_get(*args, **kwargs): + mock = MagicMock() + mock.json.return_value = mock_response + return mock + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + # Act + result = await MicrosoftSSOHandler.get_user_groups_from_graph_api( + access_token="mock_token" + ) + + # Assert + assert isinstance(result, list) + assert len(result) == 0 + + +@pytest.mark.asyncio +async def test_get_user_groups_error_handling(): + # Arrange + async def mock_get(*args, **kwargs): + raise Exception("API Error") + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + # Act + result = await MicrosoftSSOHandler.get_user_groups_from_graph_api( + access_token="mock_token" + ) + + # Assert + assert isinstance(result, list) + assert len(result) == 0 + + +def test_get_group_ids_from_graph_api_response(): + # Arrange + mock_response = MicrosoftGraphAPIUserGroupResponse( + odata_context="https://graph.microsoft.com/v1.0/$metadata#directoryObjects", + odata_nextLink=None, + value=[ + MicrosoftGraphAPIUserGroupDirectoryObject( + odata_type="#microsoft.graph.group", + id="group1", + displayName="Group 1", + description=None, + deletedDateTime=None, + roleTemplateId=None, + ), + MicrosoftGraphAPIUserGroupDirectoryObject( + odata_type="#microsoft.graph.group", + id="group2", + displayName="Group 2", + description=None, + deletedDateTime=None, + roleTemplateId=None, + ), + MicrosoftGraphAPIUserGroupDirectoryObject( + odata_type="#microsoft.graph.group", + id=None, # Test handling of None id + displayName="Invalid Group", + description=None, + deletedDateTime=None, + roleTemplateId=None, + ), + ], + ) + + # Act + result = MicrosoftSSOHandler._get_group_ids_from_graph_api_response(mock_response) + + # Assert + assert isinstance(result, list) + assert len(result) == 2 + assert "group1" in result + assert "group2" in result