mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(sso): paginate through all pages when fetching service principal group assignments (#33149)
get_group_ids_from_service_principal only read the first page of the Graph API appRoleAssignedTo response, so tenants with more than 100 groups assigned to the enterprise application silently lost group memberships during SSO login. Loop over @odata.nextLink with the same MAX_GRAPH_API_PAGES cap that get_user_groups_from_graph_api already uses, and warn when the cap is hit. Ported from #32792 by @saisurya237 so CI can run. Fixes #32790 Co-authored-by: saisurya237 <saisurya.abhishek237@gmail.com>
This commit is contained in:
parent
948a43cd64
commit
bf501c38a5
2 changed files with 72 additions and 14 deletions
|
|
@ -4034,30 +4034,41 @@ class MicrosoftSSOHandler:
|
|||
base_url = MicrosoftSSOHandler.get_graph_api_base_url()
|
||||
# Endpoint to get app role assignments for the given service principal
|
||||
endpoint = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo"
|
||||
url = base_url + endpoint
|
||||
next_link: str | None = base_url + endpoint
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"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] = []
|
||||
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", []):
|
||||
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"),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
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 group assignments may not be included."
|
||||
)
|
||||
|
||||
return group_ids, service_principal_teams
|
||||
|
||||
|
|
|
|||
|
|
@ -461,6 +461,53 @@ 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_paginates_through_all_pages():
|
||||
# Arrange
|
||||
page_one = {
|
||||
"@odata.nextLink": "https://graph.microsoft.com/v1.0/servicePrincipals/sp-123/appRoleAssignedTo?$skiptoken=page2",
|
||||
"value": [
|
||||
{
|
||||
"principalType": "Group",
|
||||
"principalId": "group-on-page-1",
|
||||
"principalDisplayName": "Group On Page 1",
|
||||
}
|
||||
],
|
||||
}
|
||||
page_two = {
|
||||
"value": [
|
||||
{
|
||||
"principalType": "Group",
|
||||
"principalId": "group-on-page-2",
|
||||
"principalDisplayName": "Group On Page 2",
|
||||
}
|
||||
],
|
||||
}
|
||||
responses = [page_one, page_two]
|
||||
|
||||
async def mock_get(url, *args, **kwargs):
|
||||
mock = MagicMock()
|
||||
mock.json.return_value = responses.pop(0)
|
||||
return mock
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.get = mock_get
|
||||
|
||||
# Act
|
||||
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
|
||||
assert group_ids == ["group-on-page-1", "group-on-page-2"]
|
||||
assert [team["principalId"] for team in teams] == [
|
||||
"group-on-page-1",
|
||||
"group-on-page-2",
|
||||
]
|
||||
|
||||
|
||||
def test_get_group_ids_from_graph_api_response():
|
||||
# Arrange
|
||||
mock_response = MicrosoftGraphAPIUserGroupResponse(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue