mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(scim): fail group sync when a member add or user creation fails (LIT-5105) (#37688)
* fix(scim): fail group sync when a member add or user creation fails Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(scim): apply ruff format Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
387a948263
commit
a030b33188
2 changed files with 144 additions and 10 deletions
|
|
@ -20,7 +20,7 @@ from fastapi import (
|
|||
Response,
|
||||
)
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
from typing_extensions import TypedDict, assert_never
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -627,6 +627,40 @@ def _admitted_member_ids(classified: Iterable[_ClassifiedGroupMember], created_i
|
|||
)
|
||||
|
||||
|
||||
class _UserIdWhere(TypedDict):
|
||||
user_id: ReadOnly[str]
|
||||
|
||||
|
||||
class _ScimErrorDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
|
||||
|
||||
async def _ensure_group_member_user(
|
||||
user_id: str,
|
||||
created_via: str,
|
||||
prisma_client: PrismaClient,
|
||||
) -> NewUserResponse | None:
|
||||
"""The created user, or None when the id already resolves to a user row (a
|
||||
concurrent provisioning request won the creation race after our lookup missed).
|
||||
|
||||
Raises:
|
||||
HTTPException: 500 when the user can neither be created nor found. The
|
||||
request has to fail so the identity provider retries, instead of recording
|
||||
success for a member the roster silently dropped.
|
||||
"""
|
||||
created: Final = await _create_user_if_not_exists(user_id=user_id, created_via=created_via)
|
||||
if created is not None:
|
||||
return created
|
||||
where: Final[_UserIdWhere] = {"user_id": user_id}
|
||||
existing: Final = await _table(UserRepository(prisma_client)).find_unique(where=where)
|
||||
if existing is not None:
|
||||
return None
|
||||
detail: Final[_ScimErrorDetail] = {
|
||||
"error": f"Failed to create user '{user_id}' while provisioning group membership."
|
||||
}
|
||||
raise HTTPException(status_code=500, detail=detail)
|
||||
|
||||
|
||||
async def _resolve_group_member_ids(
|
||||
members: Sequence[SCIMMember],
|
||||
created_via: str,
|
||||
|
|
@ -644,7 +678,8 @@ async def _resolve_group_member_ids(
|
|||
Raises:
|
||||
HTTPException: 400 when a member id is empty, or when scim_upsert_user is
|
||||
False and a member id is neither an existing user, an existing team, nor a
|
||||
member declared to be something other than a user.
|
||||
member declared to be something other than a user. 500 when a member's
|
||||
user row can neither be created nor found.
|
||||
"""
|
||||
classified: Final = tuple([await _classify_group_member(member, prisma_client) for member in members])
|
||||
partition: Final = _partition_classified_members(classified)
|
||||
|
|
@ -665,10 +700,14 @@ async def _resolve_group_member_ids(
|
|||
},
|
||||
)
|
||||
|
||||
unique_unknown_ids: Final = tuple(dict.fromkeys(partition.unknown_ids))
|
||||
creations: Final = tuple(
|
||||
[
|
||||
(user_id, await _create_user_if_not_exists(user_id=user_id, created_via=created_via))
|
||||
for user_id in partition.unknown_ids
|
||||
(
|
||||
user_id,
|
||||
await _ensure_group_member_user(user_id=user_id, created_via=created_via, prisma_client=prisma_client),
|
||||
)
|
||||
for user_id in unique_unknown_ids
|
||||
]
|
||||
)
|
||||
created_users: Final = tuple(created for _, created in creations if created is not None)
|
||||
|
|
@ -676,10 +715,7 @@ async def _resolve_group_member_ids(
|
|||
return GroupMemberExtractionResult(
|
||||
existing_member_ids=partition.resolved_ids,
|
||||
created_users=created_users,
|
||||
all_member_ids=_admitted_member_ids(
|
||||
classified,
|
||||
frozenset(user_id for user_id, created in creations if created is not None),
|
||||
),
|
||||
all_member_ids=_admitted_member_ids(classified, frozenset(unique_unknown_ids)),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2379,19 +2415,25 @@ async def _apply_group_patch_updates(group_id: str, update_data: dict[str, objec
|
|||
|
||||
|
||||
async def _handle_group_membership_changes(group_id: str, current_members: set[str], final_members: set[str]):
|
||||
"""Handle adding/removing members from the group."""
|
||||
"""Handle adding/removing members from the group.
|
||||
|
||||
Runs strict: a genuine add or remove failure propagates so the group request
|
||||
fails and the identity provider retries, instead of reporting success for a
|
||||
member the roster never received. Idempotent no-ops (already in / already out
|
||||
of the team) are still swallowed by patch_team_membership.
|
||||
"""
|
||||
members_to_add: Final = final_members - current_members
|
||||
members_to_remove: Final = current_members - final_members
|
||||
|
||||
verbose_proxy_logger.debug("members_to_add: %s", members_to_add)
|
||||
verbose_proxy_logger.debug("members_to_remove: %s", members_to_remove)
|
||||
|
||||
# Use existing helper functions for team membership changes
|
||||
for member_id in members_to_add:
|
||||
await patch_team_membership(
|
||||
user_id=member_id,
|
||||
teams_ids_to_add_user_to=[group_id],
|
||||
teams_ids_to_remove_user_from=[],
|
||||
raise_on_error=True,
|
||||
)
|
||||
|
||||
for member_id in members_to_remove:
|
||||
|
|
@ -2399,6 +2441,7 @@ async def _handle_group_membership_changes(group_id: str, current_members: set[s
|
|||
user_id=member_id,
|
||||
teams_ids_to_add_user_to=[],
|
||||
teams_ids_to_remove_user_from=[group_id],
|
||||
raise_on_error=True,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
|
|||
_apply_group_patch_updates,
|
||||
_extract_group_member_ids,
|
||||
_extract_ids_from_path_filter,
|
||||
_handle_group_membership_changes,
|
||||
_handle_team_membership_changes,
|
||||
_parse_member_entries,
|
||||
_process_group_patch_operations,
|
||||
|
|
@ -1297,6 +1298,10 @@ async def test_update_group_metadata_serialization_issue(mocker):
|
|||
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
|
||||
AsyncMock(return_value=mock_scim_group_response),
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
|
||||
AsyncMock(),
|
||||
)
|
||||
|
||||
# Call the function that had the bug
|
||||
await update_group(group_id=group_id, group=scim_group)
|
||||
|
|
@ -4453,3 +4458,89 @@ async def test_get_groups_members_are_typed_as_users(mocker):
|
|||
response = await get_groups(startIndex=1, count=10, filter=None)
|
||||
|
||||
assert [m.type for m in response.Resources[0].members] == ["User"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_raises_when_creation_fails(mocker, scim_upsert_user_enabled):
|
||||
"""A member whose user row can neither be found nor created must fail the
|
||||
request. Regression: the resolver silently dropped that member and the group
|
||||
write reported success, so the IdP recorded the user as provisioned while the
|
||||
team roster was missing them."""
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _resolve_group_member_ids(
|
||||
members=[SCIMMember(value="member-1")],
|
||||
created_via="scim_group_membership",
|
||||
prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "member-1" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_admits_member_created_concurrently(mocker, scim_upsert_user_enabled):
|
||||
"""When creation fails because a concurrent request already created the user,
|
||||
the member is still admitted: the id resolves to a real user row, so failing
|
||||
or dropping it would be wrong either way."""
|
||||
prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set())
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
side_effect=[None, LiteLLM_UserTable(user_id="raced-user")]
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
|
||||
result = await _resolve_group_member_ids(
|
||||
members=[SCIMMember(value="raced-user")],
|
||||
created_via="scim_group_membership",
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert result.all_member_ids == ["raced-user"]
|
||||
assert len(result.created_users) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_group_membership_changes_propagates_add_failure(mocker):
|
||||
"""A genuine roster add failure must fail the group request so the IdP retries.
|
||||
Regression: patch_team_membership ran with raise_on_error=False here, so a
|
||||
failed team_member_add was logged and swallowed and the SCIM group sync
|
||||
reported success with members missing from the team."""
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
|
||||
AsyncMock(side_effect=HTTPException(status_code=500, detail={"error": "db write failed"})),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
await _handle_group_membership_changes(
|
||||
group_id="group-1", current_members=set(), final_members={"user-1"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_group_membership_changes_already_in_team_is_noop(mocker):
|
||||
"""The strict path must keep treating an already-enrolled member as a no-op
|
||||
and continue with the remaining members instead of failing the sync."""
|
||||
mock_team_member_add = mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
|
||||
AsyncMock(
|
||||
side_effect=ProxyException(
|
||||
message="already in team",
|
||||
type=ProxyErrorTypes.team_member_already_in_team.value,
|
||||
param=None,
|
||||
code=400,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
await _handle_group_membership_changes(
|
||||
group_id="group-1", current_members=set(), final_members={"user-1", "user-2"}
|
||||
)
|
||||
|
||||
assert mock_team_member_add.await_count == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue