mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix SCIM memberships Patch
This commit is contained in:
parent
c39b8f2178
commit
f42480fb52
2 changed files with 254 additions and 3 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
Loading…
Add table
Reference in a new issue