Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 23 additions & 2 deletions litellm/proxy/management_endpoints/internal_user_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,9 @@
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_METADATA_KEY,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
Expand Down Expand Up @@ -719,6 +722,17 @@ async def _get_user_info_teams(
return team_list, teams_1


def _redact_scim_enterprise_metadata(
Comment thread
veria-ai[bot] marked this conversation as resolved.
metadata: Optional[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""SCIM enterprise attributes are persisted in user metadata so reporting can
group on them, but they are directory-only fields that generic user-info
endpoints must not surface; SCIM clients read them through the SCIM endpoints."""
if not isinstance(metadata, dict) or SCIM_ENTERPRISE_METADATA_KEY not in metadata:
return metadata
return {k: v for k, v in metadata.items() if k != SCIM_ENTERPRISE_METADATA_KEY}


def _build_user_info_response(
user_id: Optional[str],
user_info: Optional[Any],
Expand All @@ -739,6 +753,9 @@ def _build_user_info_response(
)
if isinstance(_user_info, dict):
_user_info.pop("password", None)
_user_info["metadata"] = _redact_scim_enterprise_metadata(
_user_info.get("metadata")
)

return UserInfoResponse(
user_id=user_id,
Expand Down Expand Up @@ -983,7 +1000,7 @@ async def user_info_v2(
models=user_data.get("models") or [],
budget_duration=user_data.get("budget_duration"),
budget_reset_at=user_data.get("budget_reset_at"),
metadata=user_data.get("metadata"),
metadata=_redact_scim_enterprise_metadata(user_data.get("metadata")),
created_at=user_data.get("created_at"),
updated_at=user_data.get("updated_at"),
sso_user_id=user_data.get("sso_user_id"),
Expand Down Expand Up @@ -2098,9 +2115,13 @@ async def get_users(
user_list: List[LiteLLM_UserTableWithKeyCount] = []
if users is not None:
for user in users:
user_dump = user.model_dump()
user_dump["metadata"] = _redact_scim_enterprise_metadata(
user_dump.get("metadata")
)
user_list.append(
LiteLLM_UserTableWithKeyCount(
**user.model_dump(), key_count=user_key_counts.get(user.user_id, 0)
**user_dump, key_count=user_key_counts.get(user.user_id, 0)
)
)
else:
Expand Down
11 changes: 10 additions & 1 deletion litellm/proxy/management_endpoints/scim/scim_transformations.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,16 @@ async def transform_litellm_user_to_scim_user(
scim_active = metadata.get("scim_active")
active = True if scim_active is None else bool(scim_active)

schemas = ["urn:ietf:params:scim:schemas:core:2.0:User"]
enterprise_user = None
if metadata.get(SCIM_ENTERPRISE_METADATA_KEY):
enterprise_user = SCIMEnterpriseUser.model_validate(
metadata[SCIM_ENTERPRISE_METADATA_KEY]
)
schemas.append(SCIM_ENTERPRISE_USER_SCHEMA)

return SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
schemas=schemas,
id=user.user_id,
userName=ScimTransformations._get_scim_user_name(user),
displayName=ScimTransformations._get_scim_user_name(user),
Expand All @@ -62,6 +70,7 @@ async def transform_litellm_user_to_scim_user(
emails=emails,
groups=groups,
active=active,
enterprise_user=enterprise_user,
meta={
"resourceType": "User",
"created": user_created_at,
Expand Down
17 changes: 15 additions & 2 deletions litellm/proxy/management_endpoints/scim/scim_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,7 @@ class ScimUserData(TypedDict):
given_name: Optional[str]
family_name: Optional[str]
active: Optional[bool]
enterprise: Optional[SCIMEnterpriseUser]


class GroupMemberExtractionResult(BaseModel):
Expand Down Expand Up @@ -199,11 +200,15 @@ def _extract_scim_user_data(user: SCIMUser) -> ScimUserData:
"given_name": user.name.givenName if user.name else None,
"family_name": user.name.familyName if user.name else None,
"active": user.active,
"enterprise": user.enterprise_user,
}


def _build_scim_metadata(
given_name: Optional[str], family_name: Optional[str], active: Optional[bool] = None
given_name: Optional[str],
family_name: Optional[str],
active: Optional[bool] = None,
enterprise: Optional[SCIMEnterpriseUser] = None,
) -> Dict[str, Any]:
"""Build metadata dictionary with SCIM data."""
metadata: Dict[str, Any] = {
Expand All @@ -216,6 +221,11 @@ def _build_scim_metadata(
if active is not None:
metadata["scim_active"] = active

if enterprise is not None:
metadata[SCIM_ENTERPRISE_METADATA_KEY] = enterprise.model_dump(
by_alias=True, exclude_none=True
)

return metadata


Expand Down Expand Up @@ -999,7 +1009,9 @@ async def create_user(
# Create user in database
user_id = user.userName or str(uuid.uuid4())
metadata = _build_scim_metadata(
user_data["given_name"], user_data["family_name"]
user_data["given_name"],
user_data["family_name"],
enterprise=user_data["enterprise"],
)

default_role: Optional[
Expand Down Expand Up @@ -1088,6 +1100,7 @@ async def update_user(
user_data["given_name"],
user_data["family_name"],
scim_active_for_metadata,
enterprise=user_data["enterprise"],
)

await _handle_team_membership_changes(
Expand Down
51 changes: 50 additions & 1 deletion litellm/types/proxy/management_endpoints/scim_v2.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,20 @@
from typing import Any, Dict, List, Literal, Optional, Union

from fastapi import HTTPException
from pydantic import BaseModel, ConfigDict, EmailStr, field_validator
from pydantic import (
BaseModel,
ConfigDict,
EmailStr,
Field,
field_validator,
model_serializer,
)
from pydantic_core.core_schema import SerializerFunctionWrapHandler

SCIM_ENTERPRISE_USER_SCHEMA = (
"urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"
)
SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise"


class LiteLLM_UserScimMetadata(BaseModel):
Expand Down Expand Up @@ -42,13 +55,49 @@ class SCIMUserGroup(BaseModel):
type: Optional[str] = "direct" # direct or indirect


class SCIMUserManager(BaseModel):
model_config = ConfigDict(populate_by_name=True)

value: Optional[str] = None
displayName: Optional[str] = None
ref: Optional[str] = Field(default=None, alias="$ref")


class SCIMEnterpriseUser(BaseModel):
model_config = ConfigDict(populate_by_name=True)

employeeNumber: Optional[str] = None
costCenter: Optional[str] = None
organization: Optional[str] = None
division: Optional[str] = None
department: Optional[str] = None
manager: Optional[SCIMUserManager] = None


class SCIMUser(SCIMResource):
model_config = ConfigDict(populate_by_name=True)

userName: Optional[str] = None
name: Optional[SCIMUserName] = None
displayName: Optional[str] = None
active: bool = True
emails: Optional[List[SCIMUserEmail]] = None
groups: Optional[List[SCIMUserGroup]] = None
enterprise_user: Optional[SCIMEnterpriseUser] = Field(
default=None,
alias=SCIM_ENTERPRISE_USER_SCHEMA,
serialization_alias=SCIM_ENTERPRISE_USER_SCHEMA,
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.

@model_serializer(mode="wrap")
def _omit_absent_enterprise(
self, handler: SerializerFunctionWrapHandler
) -> Dict[str, Any]:
dumped = handler(self)
if self.enterprise_user is None:
dumped.pop(SCIM_ENTERPRISE_USER_SCHEMA, None)
dumped.pop("enterprise_user", None)
return dumped


class SCIMMember(BaseModel):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,10 @@
ScimTransformations,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIMEnterpriseUser,
SCIMPatchOperation,
SCIMUser,
)


Expand Down Expand Up @@ -149,6 +152,77 @@ async def test_transform_user_with_scim_metadata(
assert scim_user.name.givenName == "Test"
assert scim_user.name.familyName == "User"

@pytest.mark.asyncio
async def test_transform_user_with_enterprise_metadata(self, mock_prisma_client):
mock_client, mock_find_unique = mock_prisma_client
mock_find_unique.return_value = None

user = LiteLLM_UserTable(
user_id="user-ent",
user_email="ent@example.com",
user_alias=None,
teams=[],
created_at=None,
updated_at=None,
metadata={
"scim_enterprise": {"costCenter": "CC-42", "department": "Platform"}
},
)

with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
user
)

assert scim_user.enterprise_user is not None
assert scim_user.enterprise_user.costCenter == "CC-42"
assert scim_user.enterprise_user.department == "Platform"
assert SCIM_ENTERPRISE_USER_SCHEMA in scim_user.schemas

@pytest.mark.asyncio
async def test_transform_user_without_enterprise_metadata_omits_schema(
self, mock_user, mock_prisma_client
):
mock_client, mock_find_unique = mock_prisma_client
team1 = LiteLLM_TeamTable(
team_id="team-1", team_alias="Team One", members_with_roles=[]
)
team2 = LiteLLM_TeamTable(
team_id="team-2", team_alias="Team Two", members_with_roles=[]
)
mock_find_unique.side_effect = [team1, team2]

with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
mock_user
)

assert scim_user.enterprise_user is None
assert SCIM_ENTERPRISE_USER_SCHEMA not in scim_user.schemas

def test_scim_user_serialization_omits_absent_enterprise_urn(self):
without_enterprise = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="user-1",
userName="user@example.com",
)
dumped = without_enterprise.model_dump(by_alias=True)
assert SCIM_ENTERPRISE_USER_SCHEMA not in dumped
assert "enterprise_user" not in dumped
assert SCIM_ENTERPRISE_USER_SCHEMA not in dumped["schemas"]

with_enterprise = SCIMUser(
schemas=[
"urn:ietf:params:scim:schemas:core:2.0:User",
SCIM_ENTERPRISE_USER_SCHEMA,
],
id="user-2",
userName="ent@example.com",
enterprise_user=SCIMEnterpriseUser(costCenter="CC-42"),
)
dumped_ent = with_enterprise.model_dump(by_alias=True)
assert dumped_ent[SCIM_ENTERPRISE_USER_SCHEMA]["costCenter"] == "CC-42"

@pytest.mark.asyncio
async def test_transform_litellm_team_to_scim_group(
self, mock_team, mock_prisma_client
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
update_user,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIMGroup,
SCIMMember,
SCIMPatchOp,
Expand Down Expand Up @@ -115,6 +116,59 @@ async def test_create_user_defaults_to_viewer(mocker, monkeypatch):
assert called_args.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY


@pytest.mark.asyncio
async def test_create_user_ingests_enterprise_extension(mocker, monkeypatch):
"""A SCIM create payload carrying the enterprise extension block should land
in the created user's metadata under scim_enterprise"""

scim_user = SCIMUser.model_validate(
{
"schemas": [
"urn:ietf:params:scim:schemas:core:2.0:User",
SCIM_ENTERPRISE_USER_SCHEMA,
],
"userName": "ent-user",
"name": {"familyName": "User", "givenName": "Ent"},
"emails": [{"value": "ent@example.com"}],
SCIM_ENTERPRISE_USER_SCHEMA: {
"costCenter": "CC-42",
"department": "Platform",
},
}
)

mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)

monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False)

mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)

new_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_user",
AsyncMock(return_value=NewUserRequest(user_id="ent-user")),
)

mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=scim_user),
)

await create_user(user=scim_user)

created_metadata = new_user_mock.call_args.kwargs["data"].metadata
assert created_metadata["scim_enterprise"] == {
"costCenter": "CC-42",
"department": "Platform",
}


@pytest.mark.asyncio
async def test_create_user_uses_default_internal_user_params_role(mocker, monkeypatch):
"""If role is set in default_internal_user_params, new user should use that role"""
Expand Down
Loading
Loading