fix SCIM memberships Patch

This commit is contained in:
Ishaan Jaff 2025-06-18 08:43:52 -07:00
parent c39b8f2178
commit f42480fb52
2 changed files with 254 additions and 3 deletions

View file

@ -5,7 +5,7 @@ This is an enterprise feature and requires a premium license.
"""
import uuid
from typing import List, Optional
from typing import Any, Dict, List, Optional, Set
from fastapi import (
APIRouter,
@ -25,6 +25,8 @@ from litellm.proxy._types import (
Member,
NewTeamRequest,
NewUserRequest,
TeamMemberAddRequest,
TeamMemberDeleteRequest,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -32,7 +34,11 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
from litellm.proxy.management_endpoints.scim.scim_transformations import (
ScimTransformations,
)
from litellm.proxy.management_endpoints.team_endpoints import new_team
from litellm.proxy.management_endpoints.team_endpoints import (
new_team,
team_member_add,
team_member_delete,
)
from litellm.proxy.utils import _premium_user_check, handle_exception_on_proxy
from litellm.types.proxy.management_endpoints.scim_v2 import *
@ -291,6 +297,122 @@ async def delete_user(
raise handle_exception_on_proxy(e)
def _extract_group_values(value: Any) -> List[str]:
"""Return group ids from a SCIM patch value."""
group_values: List[str] = []
if isinstance(value, list):
for v in value:
if isinstance(v, dict) and v.get("value"):
group_values.append(str(v.get("value")))
elif isinstance(v, str):
group_values.append(v)
elif isinstance(value, dict):
if value.get("value"):
group_values.append(str(value.get("value")))
elif isinstance(value, str):
group_values.append(value)
return group_values
def _apply_patch_ops( # noqa: PLR0915 -- allow complex logic
existing_user: LiteLLM_UserTable,
patch_ops: SCIMPatchOp,
) -> tuple[Dict[str, Any], Set[str]]:
"""Apply patch operations and return update data and final team set."""
update_data: Dict[str, Any] = {}
metadata = existing_user.metadata or {}
scim_metadata = metadata.get("scim_metadata", {})
teams_set: Set[str] = set(existing_user.teams or [])
replace_team_set: Optional[Set[str]] = None
for op in patch_ops.Operations:
path = (op.path or "").lower()
value = op.value
op_type = op.op
if path == "displayname":
if op_type == "remove":
update_data["user_alias"] = None
else:
update_data["user_alias"] = str(value)
elif path == "active":
if op_type == "remove":
metadata.pop("scim_active", None)
else:
bool_val = value
if isinstance(value, str):
bool_val = value.lower() == "true"
else:
bool_val = bool(value)
metadata["scim_active"] = bool_val
elif path == "externalid":
if op_type == "remove":
update_data["sso_user_id"] = None
else:
update_data["sso_user_id"] = str(value)
elif path == "name.givenname":
if op_type == "remove":
scim_metadata.pop("givenName", None)
else:
scim_metadata["givenName"] = str(value)
elif path == "name.familyname":
if op_type == "remove":
scim_metadata.pop("familyName", None)
else:
scim_metadata["familyName"] = str(value)
elif path.startswith("groups"):
group_values = _extract_group_values(value)
if op_type == "replace":
replace_team_set = set(group_values)
elif op_type == "add":
teams_set.update(group_values)
elif op_type == "remove":
for gid in group_values:
teams_set.discard(gid)
else:
if op_type == "remove":
metadata.pop(path, None)
else:
metadata[path] = value
final_team_set = replace_team_set if replace_team_set is not None else teams_set
metadata["scim_metadata"] = scim_metadata
update_data["metadata"] = metadata
return update_data, final_team_set
async def patch_team_membership(
user_id: str,
teams_ids_to_add_user_to: List[str],
teams_ids_to_remove_user_from: List[str],
) -> bool:
"""
Add or remove user from teams
"""
for _team_id in teams_ids_to_add_user_to:
try:
await team_member_add(
data=TeamMemberAddRequest(
team_id=_team_id,
member=Member(user_id=user_id, role="user"),
),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
except Exception as e:
verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}")
for _team_id in teams_ids_to_remove_user_from:
try:
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=_team_id, user_id=user_id),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
except Exception as e:
verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}")
return True
@scim_router.patch(
"/Users/{user_id}",
response_model=SCIMUser,
@ -322,7 +444,31 @@ async def patch_user(
status_code=404, detail={"error": f"User not found with ID: {user_id}"}
)
return None
update_data, final_team_set = _apply_patch_ops(
existing_user=existing_user,
patch_ops=patch_ops,
)
existing_teams = set(existing_user.teams or [])
added_groups = final_team_set - existing_teams
removed_groups = existing_teams - final_team_set
await patch_team_membership(
user_id=user_id,
teams_ids_to_add_user_to=list(added_groups),
teams_ids_to_remove_user_from=list(removed_groups),
)
update_data["teams"] = list(final_team_set)
updated_user = await prisma_client.db.litellm_usertable.update(
where={"user_id": user_id},
data=update_data,
)
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(updated_user)
return scim_user
except Exception as e:
raise handle_exception_on_proxy(e)

View file

@ -0,0 +1,105 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.scim.scim_v2 import patch_user
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIMPatchOp,
SCIMPatchOperation,
)
@pytest.mark.asyncio
async def test_patch_user_updates_fields():
mock_user = LiteLLM_UserTable(
user_id="user-1",
user_email="test@example.com",
user_alias="Old",
teams=[],
metadata={},
)
async def mock_update(*, where, data):
if "user_alias" in data:
mock_user.user_alias = data["user_alias"]
if "metadata" in data:
mock_user.metadata = data["metadata"]
if "teams" in data:
mock_user.teams = data["teams"]
if "sso_user_id" in data:
mock_user.sso_user_id = data["sso_user_id"]
return mock_user
mock_client = MagicMock()
mock_db = MagicMock()
mock_client.db = mock_db
mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user)
mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update)
mock_db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(op="replace", path="displayName", value="New Name"),
SCIMPatchOperation(op="replace", path="active", value="False"),
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
result = await patch_user(user_id="user-1", patch_ops=patch_ops)
mock_db.litellm_usertable.update.assert_called_once()
assert result.displayName == "New Name"
assert mock_user.metadata.get("scim_active") is False
@pytest.mark.asyncio
async def test_patch_user_manages_group_memberships():
mock_user = LiteLLM_UserTable(
user_id="user-2",
user_email="test@example.com",
user_alias="Old",
teams=["old-team"],
metadata={},
)
async def mock_update(*, where, data):
if "teams" in data:
mock_user.teams = data["teams"]
if "metadata" in data:
mock_user.metadata = data["metadata"]
return mock_user
mock_client = MagicMock()
mock_db = MagicMock()
mock_client.db = mock_db
mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user)
mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update)
async def mock_add(data, user_api_key_dict):
mock_user.teams.append(data.team_id)
async def mock_delete(data, user_api_key_dict):
if data.team_id in mock_user.teams:
mock_user.teams.remove(data.team_id)
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(op="add", path="groups", value=[{"value": "new-team"}]),
SCIMPatchOperation(op="remove", path="groups", value=[{"value": "old-team"}]),
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_client), patch(
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
AsyncMock(side_effect=mock_add),
) as mock_add_fn, patch(
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete",
AsyncMock(side_effect=mock_delete),
) as mock_del_fn:
await patch_user(user_id="user-2", patch_ops=patch_ops)
assert mock_add_fn.called
assert mock_del_fn.called
assert mock_user.teams == ["new-team"]