mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): paginate Microsoft SSO app role assignments
This commit is contained in:
parent
bf02a4a47f
commit
e9858c29b8
2 changed files with 80 additions and 13 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue