diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 065464aa565..1f2da7ad8c0 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -4041,23 +4041,35 @@ class MicrosoftSSOHandler: "Content-Type": "application/json", } - response = await async_client.get(url, headers=headers) - response_json = response.json() - verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}") group_ids: List[str] = [] service_principal_teams: List[MicrosoftServicePrincipalTeam] = [] + next_link: Optional[str] = url + page_count = 0 - for _object in response_json.get("value", []): - if _object.get("principalType") == "Group": - # Append the group ID to the list - group_ids.append(_object.get("principalId")) - # Append the service principal team to the list - service_principal_teams.append( - MicrosoftServicePrincipalTeam( - principalDisplayName=_object.get("principalDisplayName"), - principalId=_object.get("principalId"), + while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES: + response = await async_client.get(next_link, headers=headers) + response_json = response.json() + verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}") + + for _object in response_json.get("value", []): + principal_id = _object.get("principalId") + if _object.get("principalType") == "Group" and principal_id is not None: + group_ids.append(principal_id) + service_principal_teams.append( + MicrosoftServicePrincipalTeam( + principalDisplayName=_object.get("principalDisplayName"), + principalId=principal_id, + ) ) - ) + + next_link = response_json.get("@odata.nextLink") + page_count += 1 + + if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES: + verbose_proxy_logger.warning( + f"Reached maximum page limit of {MicrosoftSSOHandler.MAX_GRAPH_API_PAGES}. " + "Some service principal groups may not be included." + ) return group_ids, service_principal_teams diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 045e15f8b8b..92470fb5e08 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -461,6 +461,61 @@ async def test_get_group_ids_from_service_principal_uses_configured_graph_endpoi ] +@pytest.mark.asyncio +async def test_get_group_ids_from_service_principal_follows_next_link(): + next_link = "https://graph.microsoft.com/v1.0/servicePrincipals/sp-123/appRoleAssignedTo?$skiptoken=page-2" + first_response = MagicMock() + first_response.json.return_value = { + "value": [ + { + "principalType": "Group", + "principalId": "group-1", + "principalDisplayName": "Group 1", + }, + { + "principalType": "User", + "principalId": "user-1", + "principalDisplayName": "User 1", + }, + ], + "@odata.nextLink": next_link, + } + second_response = MagicMock() + second_response.json.return_value = { + "value": [ + { + "principalType": "Group", + "principalId": "group-2", + "principalDisplayName": "Group 2", + } + ] + } + async_client = MagicMock() + async_client.get = AsyncMock(side_effect=[first_response, second_response]) + + group_ids, teams = await MicrosoftSSOHandler.get_group_ids_from_service_principal( + service_principal_id="sp-123", + async_client=async_client, + access_token="mock_token", + ) + + assert group_ids == ["group-1", "group-2"] + assert teams == [ + MicrosoftServicePrincipalTeam( + principalDisplayName="Group 1", + principalId="group-1", + ), + MicrosoftServicePrincipalTeam( + principalDisplayName="Group 2", + principalId="group-2", + ), + ] + assert [call.args[0] for call in async_client.get.await_args_list] == [ + "https://graph.microsoft.com/v1.0/servicePrincipals/sp-123/appRoleAssignedTo", + next_link, + ] + + def test_get_group_ids_from_graph_api_response(): # Arrange mock_response = MicrosoftGraphAPIUserGroupResponse(