diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 24b4a8047ec..99b8335a6ea 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3846,7 +3846,7 @@ class _TokenInFilter(TypedDict): async def get_jwt_key_mapping_cache_keys_for_tokens( hashed_tokens: Sequence[str], - prisma_client: PrismaClient, + prisma_client: DatabaseClient, ) -> tuple[str, ...]: """Cache keys of every JWT claim mapped to any of the given virtual keys.""" if not hashed_tokens: diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 866a747690e..4659c92b76e 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -15,7 +15,7 @@ All /budget management endpoints import math from collections.abc import Mapping from types import MappingProxyType -from typing import Final +from typing import TYPE_CHECKING, Final from fastapi import APIRouter, Depends, HTTPException @@ -28,8 +28,12 @@ from litellm.proxy.management_endpoints.common_utils import ( ) from litellm.proxy.utils import jsonify_object from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import AgentsRepository +if TYPE_CHECKING: + from prisma.models import LiteLLM_BudgetTable as PrismaBudget + router: Final = APIRouter() @@ -69,8 +73,6 @@ async def new_budget( - model_max_budget: Optional[dict] - Specify max budget for a given model. Example: {"openai/gpt-4o-mini": {"max_budget": 100.0, "budget_duration": "1d", "tpm_limit": 100000, "rpm_limit": 100000}} - budget_reset_at: Optional[datetime] - Datetime when the initial budget is reset. Default is now. """ - from prisma.errors import UniqueViolationError - from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client if prisma_client is None: @@ -79,6 +81,18 @@ async def new_budget( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) + return await create_budget( + budget_obj=budget_obj, + table=BudgetRepository(prisma_client).table, + created_by=user_api_key_dict.user_id or litellm_proxy_admin_name, + ) + + +async def create_budget( + *, budget_obj: BudgetNewRequest, table: TableActions["PrismaBudget"], created_by: str +) -> "PrismaBudget": + from prisma.errors import UniqueViolationError + # Validate budget values are not negative if budget_obj.max_budget is not None and (not math.isfinite(budget_obj.max_budget) or budget_obj.max_budget < 0): raise HTTPException( @@ -111,11 +125,11 @@ async def new_budget( budget_obj_json: Final = budget_obj.model_dump(exclude_none=True) budget_obj_jsonified: Final[dict[str, object]] = jsonify_object(budget_obj_json) # mutable-ok: prisma create input try: - response: Final = await BudgetRepository(prisma_client).table.create( + response: Final = await table.create( data={ **budget_obj_jsonified, - "created_by": user_api_key_dict.user_id or litellm_proxy_admin_name, - "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, + "created_by": created_by, + "updated_by": created_by, } ) except Exception as e: diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index ee1ebcb5ce9..143480d235e 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -468,6 +468,18 @@ async def _fetch_user_team_ids(user_id: str, prisma_client: "PrismaClient") -> t return tuple(user_row.teams) if user_row is not None else () +async def check_user_license_capacity(prisma_client: "PrismaClient", *, users: UserRepository | None = None) -> None: + from litellm.proxy.proxy_server import _license_check + + repository: Final = UserRepository(prisma_client) if users is None else users + billable_users: Final = await repository.count_billable_users() + if billable_users and _license_check.is_over_limit(total_users=billable_users): + raise HTTPException( + status_code=403, + detail="License is over limit. Please contact support@berri.ai to upgrade your license.", + ) + + @router.post( "/user/new", tags=["Internal User management"], @@ -544,7 +556,7 @@ async def new_user( ``` """ try: - from litellm.proxy.proxy_server import _license_check, prisma_client + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: raise HTTPException(status_code=400, detail=CommonProxyErrors.db_not_connected_error.value) @@ -560,13 +572,7 @@ async def new_user( await _check_duplicate_user_id(data.user_id, prisma_client) await _check_duplicate_user_email(data.user_email, prisma_client) - # Check if license is over limit - billable_users: Final = await UserRepository(prisma_client).count_billable_users() - if billable_users and _license_check.is_over_limit(total_users=billable_users): - raise HTTPException( - status_code=403, - detail="License is over limit. Please contact support@berri.ai to upgrade your license.", - ) + await check_user_license_capacity(prisma_client) # Only proxy admins can create administrative users # Check if user_api_key_dict is actually a UserAPIKeyAuth instance (not a Depends object) diff --git a/litellm/proxy/management_endpoints/scim/agent_provisioning.py b/litellm/proxy/management_endpoints/scim/agent_provisioning.py new file mode 100644 index 00000000000..95ab3c11bd1 --- /dev/null +++ b/litellm/proxy/management_endpoints/scim/agent_provisioning.py @@ -0,0 +1,667 @@ +import re +from collections import deque +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping +from contextlib import asynccontextmanager +from dataclasses import dataclass +from datetime import timedelta +from functools import reduce, wraps +from itertools import chain +from types import MappingProxyType +from typing import TYPE_CHECKING, Concatenate, Final, Literal, ParamSpec, TypeVar +from uuid import UUID, uuid4 + +from fastapi import HTTPException +from prisma import Json, Prisma +from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource +from prisma.types import ( + LiteLLM_AgentsTableCreateInput, + LiteLLM_SCIMResourceCreateInput, + LiteLLM_SCIMResourceUpdateInput, + LiteLLM_SCIMResourceWhereInput, + LiteLLM_SCIMResourceWhereUniqueInput, + LiteLLM_VerifiedSubjectCreateInput, +) +from pydantic import TypeAdapter, ValidationError + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import PrismaClient +from litellm.repositories.base_repository import is_unique_violation +from litellm.repositories.chunked_in import count_in, find_many_in +from litellm.repositories.table_repositories import SCIMResourceRepository, SCIMSourceRepository +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import ( + SCIM_AGENT_USER_SCHEMA, + canonical_directory_id, +) +from litellm.types.proxy.management_endpoints.scim_v2 import ( + SCIMGroup, + SCIMListResponse, + SCIMMember, + SCIMPatchOp, + SCIMPatchOperation, + SCIMUser, + SCIMUserName, +) + +if TYPE_CHECKING: + from litellm.proxy.management_endpoints.scim.scim_v2 import ProvisionedGroupWrite + + +@dataclass(frozen=True, slots=True) +class _GroupDatabase: + db: Prisma + + +@dataclass(frozen=True, slots=True) +class SCIMProvisioningFailure: + status: int + message: str + + +def reject(failure: SCIMProvisioningFailure) -> None: + raise HTTPException(failure.status, failure.message) + + +async def source_for_auth(auth: object, client: PrismaClient) -> LiteLLM_SCIMSource | None: + if not isinstance(auth, UserAPIKeyAuth) or not auth.token: + return None + source: Final = await SCIMSourceRepository(client, use_writer=True).table.find_unique( + where={"key_hash": auth.token} + ) + if source is not None and not source.enabled: + reject(SCIMProvisioningFailure(403, "This provisioning source is disabled")) + return source + + +def user_document(row: LiteLLM_SCIMResource) -> SCIMUser: + return SCIMUser.model_validate( + {**TypeAdapter(dict[str, object]).validate_python(row.document), "id": row.id, "active": row.active} + ) + + +def group_document(row: LiteLLM_SCIMResource) -> SCIMGroup: + return SCIMGroup.model_validate( + { + **TypeAdapter(dict[str, object]).validate_python(row.document), + "id": row.id, + "members": tuple(SCIMMember(value=value) for value in row.member_ids), + } + ) + + +async def remove_group_member(tx: Prisma, group: LiteLLM_SCIMResource, member_id: str) -> None: + where: Final[LiteLLM_SCIMResourceWhereUniqueInput] = {"id": group.id} + data: Final[LiteLLM_SCIMResourceUpdateInput] = { + "member_ids": [member for member in group.member_ids if member != member_id] + } + await tx.litellm_scimresource.update(where=where, data=data) + + +def _user_changes(item: SCIMPatchOperation, current: SCIMUser) -> dict[str, object] | SCIMProvisioningFailure: + allowed: Final = MappingProxyType( + { + "active": "active", + "displayname": "displayName", + "username": "userName", + "name": "name", + "emails": "emails", + } + ) + if item.path is None: + value: Final = item.value + if item.op == "remove" or not isinstance(value, dict): + return SCIMProvisioningFailure(400, "An object value or attribute path is required") + changes: Final = TypeAdapter(dict[str, object]).validate_python(value) + if any(key not in allowed.values() for key in changes): + return SCIMProvisioningFailure(400, "Agent subject and parent identity are immutable") + return changes + name_fields: Final = MappingProxyType({"name." + name.lower(): name for name in SCIMUserName.model_fields}) + name_field: Final = name_fields.get(item.path.lower()) + if name_field is not None: + return { + "name": { + **(current.name.model_dump() if current.name else {}), + name_field: None if item.op == "remove" else item.value, + } + } + email_type: Final = re.fullmatch(r'emails\[type eq "([^"\r\n]+)"\]\.value', item.path, re.IGNORECASE) + if email_type is not None: + others: Final = tuple(email.model_dump() for email in current.emails or () if email.type != email_type[1]) + selected: Final = next( + (email.model_dump() for email in current.emails or () if email.type == email_type[1]), + {"type": email_type[1]}, + ) + return {"emails": others if item.op == "remove" else ({**selected, "value": item.value}, *others)} + key: Final = allowed.get(item.path.lower()) + if key is None: + return SCIMProvisioningFailure(400, "This attribute is immutable or unsupported for an agent-user") + return {key: None if item.op == "remove" else item.value} + + +def _patch_user_operation( + current: SCIMUser | SCIMProvisioningFailure, item: SCIMPatchOperation +) -> SCIMUser | SCIMProvisioningFailure: + if isinstance(current, SCIMProvisioningFailure): + return current + changes: Final = _user_changes(item, current) + if isinstance(changes, SCIMProvisioningFailure): + return changes + try: + updated: Final = SCIMUser.model_validate({**current.model_dump(by_alias=True), **changes}) + except ValidationError: + return SCIMProvisioningFailure(400, "Invalid agent-user attribute value") + if not updated.userName: + return SCIMProvisioningFailure(400, "userName is required") + return updated + + +def apply_user_patch(user: SCIMUser, patch: SCIMPatchOp) -> SCIMUser | SCIMProvisioningFailure: + return reduce(_patch_user_operation, patch.Operations, user) + + +def _patched_members(current: SCIMGroup, item: SCIMPatchOperation) -> tuple[SCIMMember, ...] | SCIMProvisioningFailure: + path: Final = item.path or "" + selected: Final = re.fullmatch(r'members\[value eq "([^"\r\n]+)"\]', path, re.IGNORECASE) + if selected is not None and item.op == "remove": + return tuple(member for member in current.members or () if member.value != selected[1]) + if path.lower() != "members": + return SCIMProvisioningFailure(400, "Unsupported group PATCH attribute") + try: + incoming: Final = TypeAdapter(tuple[SCIMMember, ...]).validate_python(() if item.value is None else item.value) + except ValidationError: + return SCIMProvisioningFailure(400, "Invalid group members") + if item.op == "replace": + return incoming + if item.op == "remove": + removed: Final = frozenset(member.value for member in incoming) + return tuple(member for member in current.members or () if member.value not in removed) if incoming else () + merged: Final = MappingProxyType({member.value: member for member in (*(current.members or ()), *incoming)}) + return tuple(merged.values()) + + +def _patch_group_operation( + current: SCIMGroup | SCIMProvisioningFailure, item: SCIMPatchOperation +) -> SCIMGroup | SCIMProvisioningFailure: + if isinstance(current, SCIMProvisioningFailure): + return current + if item.path is None: + return _replace_group_attributes(current, item) + if item.path.lower() == "displayname" and item.op != "remove" and isinstance(item.value, str): + return current.model_copy(update={"displayName": item.value}) + members: Final = _patched_members(current, item) + if isinstance(members, SCIMProvisioningFailure): + return members + return current.model_copy(update={"members": list(members)}) + + +def _replace_group_attributes(current: SCIMGroup, item: SCIMPatchOperation) -> SCIMGroup | SCIMProvisioningFailure: + raw: Final[object] = item.value + if item.op == "remove" or not isinstance(raw, dict): + return SCIMProvisioningFailure(400, "An object value or attribute path is required") + fields: Final = TypeAdapter(Mapping[str, object]).validate_python(raw) + if any(key.lower() not in ("displayname", "members") for key in fields): + return SCIMProvisioningFailure(400, "Unsupported group PATCH attribute") + display_name: Final = next((value for key, value in fields.items() if key.lower() == "displayname"), None) + if display_name is not None and (not isinstance(display_name, str) or not display_name): + return SCIMProvisioningFailure(400, "Invalid group displayName") + renamed: Final = current if display_name is None else current.model_copy(update={"displayName": display_name}) + member_values: Final = next((value for key, value in fields.items() if key.lower() == "members"), None) + if member_values is None: + return renamed + return _patch_group_operation(renamed, SCIMPatchOperation(op=item.op, path="members", value=member_values)) + + +def group_members_after_patch(group: SCIMGroup, patch: SCIMPatchOp) -> SCIMGroup | SCIMProvisioningFailure: + return reduce(_patch_group_operation, patch.Operations, group) + + +Parameters = ParamSpec("Parameters") +Result = TypeVar("Result") + + +def serialized_source( + operation: Callable[Concatenate["AgentProvisioningService", Parameters], Awaitable[Result]], +) -> Callable[Concatenate["AgentProvisioningService", Parameters], Awaitable[Result]]: + @wraps(operation) + async def execute( + service: "AgentProvisioningService", *args: Parameters.args, **kwargs: Parameters.kwargs + ) -> Result: + async with service.source_transaction(): + return await operation(service, *args, **kwargs) + + return execute + + +def _identity_patch_children(value: object) -> tuple[object, ...] | Literal[True]: + if isinstance(value, dict): + fields: Final = TypeAdapter(dict[str, object]).validate_python(value) + if any( + key.lower().startswith(SCIM_AGENT_USER_SCHEMA.lower()) or key.lower() in ("agent_user", "identityparentid") + for key in fields + ): + return True + return tuple(fields.values()) + if isinstance(value, list): + return TypeAdapter(tuple[object, ...]).validate_python(value) + if isinstance(value, str) and value.lower().startswith(SCIM_AGENT_USER_SCHEMA.lower()): + return True + return () + + +def patch_changes_identity(patch: SCIMPatchOp) -> bool: + pending: Final = deque( # mutable-ok: work queue avoids recursion on arbitrarily nested untrusted PATCH values + chain.from_iterable((item.path, item.value) for item in patch.Operations) + ) + while pending: + match _identity_patch_children(pending.popleft()): + case True: + return True + case children: + pending.extend(children) + return False + + +class AgentProvisioningService: + def __init__(self, client: PrismaClient, source: LiteLLM_SCIMSource) -> None: + self.client = client + self.source = source + + @asynccontextmanager + async def source_transaction(self) -> AsyncGenerator[Prisma]: + async with self.client.tx(timeout=timedelta(seconds=30)) as tx: + await tx.execute_raw( + "SELECT pg_advisory_xact_lock(hashtextextended($1, 0))", "scim-source:" + self.source.source_id + ) + current: Final = await tx.litellm_scimsource.find_unique(where={"source_id": self.source.source_id}) + if current is None or not current.enabled or current.key_hash != self.source.key_hash: + raise HTTPException(403, "This provisioning source is disabled or its token has changed") + yield tx + + async def list( + self, kind: Literal["Users", "Groups"], start: int, count: int, filter_value: str | None + ) -> SCIMListResponse: + from litellm.proxy.management_endpoints.scim.scim_v2 import parse_scim_eq_filter + + parsed: Final = parse_scim_eq_filter(filter_value) if filter_value else None + fields: Final = MappingProxyType( + { + "username": "user_name", + "externalid": "external_id", + "displayname": "display_name", + "id": "id", + } + ) + if filter_value and (parsed is None or parsed[0] not in fields): + reject(SCIMProvisioningFailure(400, "Unsupported SCIM filter")) + filter_clause: Final[LiteLLM_SCIMResourceWhereInput] = ( + {"user_name": parsed[1]} + if parsed and parsed[0] == "username" + else {"external_id": canonical_directory_id(parsed[1])} + if parsed and parsed[0] == "externalid" + else {"display_name": parsed[1]} + if parsed and parsed[0] == "displayname" + else {"id": parsed[1]} + if parsed + else {} + ) + where: Final[LiteLLM_SCIMResourceWhereInput] = { + "source_id": self.source.source_id, + "kind": kind, + "deleted": False, + **filter_clause, + } + async with self.client.tx() as tx: + rows: Final = await tx.litellm_scimresource.find_many( + where=where, skip=start - 1, take=min(count, 100), order={"id": "asc"} + ) + total: Final = await tx.litellm_scimresource.count(where=where) + return SCIMListResponse( + totalResults=total, + startIndex=start, + itemsPerPage=len(rows), + Resources=[user_document(row) for row in rows] + if kind == "Users" + else [group_document(row) for row in rows], + ) + + async def get(self, kind: Literal["Users", "Groups"], resource_id: str) -> SCIMUser | SCIMGroup: + async with self.client.tx() as tx: + row: Final = await self._resource(tx, kind, resource_id) + return user_document(row) if kind == "Users" else group_document(row) + + async def _resource(self, tx: Prisma, kind: Literal["Users", "Groups"], resource_id: str) -> LiteLLM_SCIMResource: + where: Final[LiteLLM_SCIMResourceWhereUniqueInput] = {"id": resource_id} + row: Final = await tx.litellm_scimresource.find_unique(where=where) + if row is None or row.source_id != self.source.source_id or row.kind != kind or row.deleted: + raise HTTPException(404, "SCIM resource not found in this provisioning source") + return row + + async def create_user(self, user: SCIMUser) -> SCIMUser: + from litellm.proxy.management_endpoints.scim.human_provisioning import SourceHumanProvisioner + + async with self.source_transaction() as tx: + if not user.externalId or not user.userName: + raise HTTPException(400, "externalId and userName are required") + if user.agent_user is not None: + return await self._create_native( + tx, + user, + external_id=user.externalId, + user_name=user.userName, + parent=str(user.agent_user.identityParentId), + ) + result: Final = await SourceHumanProvisioner(self.client, self.source).create_in_transaction(tx, user) + await SourceHumanProvisioner.finish_update(result) + return result.document + + async def _create_native( + self, tx: Prisma, user: SCIMUser, *, external_id: str, user_name: str, parent: str + ) -> SCIMUser: + try: + oid: Final = str(UUID(external_id)) + except ValueError: + raise HTTPException(400, "Agent externalId must be the Entra object ID") + tenant: Final = self.source.tenant_id + issuer: Final = f"https://login.microsoftonline.com/{tenant}/v2.0" + try: + previous: Final = await tx.litellm_scimresource.find_unique( + where={ + "source_id_kind_external_id": { + "source_id": self.source.source_id, + "kind": "Users", + "external_id": oid, + } + } + ) + if previous is not None: + if previous.deleted: + raise HTTPException(409, "This subject was deleted; automatic recreation is not permitted") + return await self._update_native(tx, previous, user) + registered: Final = await tx.litellm_agentidentity.find_unique( + where={ + "provider_tenant_id_client_id": { + "provider": "microsoft_entra", + "tenant_id": tenant, + "client_id": parent, + } + } + ) + if registered is not None: + raise HTTPException( + 409, "This parent identity is already registered; automatic adoption is not permitted" + ) + agent_id: Final = str(uuid4()) + scim_id: Final = str(uuid4()) + document: Final = user.model_copy(update={"id": scim_id}) + resource_data: Final[LiteLLM_SCIMResourceCreateInput] = LiteLLM_SCIMResourceCreateInput( + id=scim_id, + source_id=self.source.source_id, + kind="Users", + external_id=oid, + user_name=user_name, + display_name=user.displayName or user_name, + document=Json(document.model_dump(by_alias=True, mode="json", exclude_none=True)), + active=user.active, + local_id=agent_id, + ) + await tx.litellm_scimresource.create(data=resource_data) + agent_data: Final[LiteLLM_AgentsTableCreateInput] = LiteLLM_AgentsTableCreateInput( + agent_id=agent_id, + agent_name=user.displayName or user_name, + agent_card_params=Json({}), + identity_managed=True, + enabled=False, + execution_mode="autonomous", + created_by="scim:" + self.source.source_id, + updated_by="scim:" + self.source.source_id, + identity={ + "create": { + "provider": "microsoft_entra", + "tenant_id": tenant, + "issuer": issuer, + "client_id": parent, + "provisioning_source_id": self.source.source_id, + "revision": str(uuid4()), + } + }, + retired_identities={ + "create": { + "provider": "microsoft_entra", + "tenant_id": tenant, + "issuer": issuer, + "client_id": parent, + } + }, + ) + await tx.litellm_agentstable.create(data=agent_data) + subject_data: Final[LiteLLM_VerifiedSubjectCreateInput] = LiteLLM_VerifiedSubjectCreateInput( + issuer=issuer, + tenant_id=tenant, + oid=oid, + kind="agent_user", + agent_id=agent_id, + parent_client_id=parent, + scim_resource_id=scim_id, + verified_via="scim", + ) + await tx.litellm_verifiedsubject.create(data=subject_data) + return document + except Exception as exc: + if is_unique_violation(exc): + raise HTTPException(409, "The subject, parent identity or agent name is already registered") from exc + raise + + async def _update_native(self, tx: Prisma, row: LiteLLM_SCIMResource, user: SCIMUser) -> SCIMUser: + old: Final = user_document(row) + try: + subject_id: Final = str(UUID(user.externalId or "")) + except ValueError: + raise HTTPException(409, "Agent subject and parent identity are immutable") from None + if user.agent_user != old.agent_user or subject_id != row.external_id: + raise HTTPException(409, "Agent subject and parent identity are immutable") + if row.local_id is None or await tx.litellm_agentstable.find_unique(where={"agent_id": row.local_id}) is None: + raise HTTPException(409, "The registered agent was deleted; automatic recreation is not permitted") + if not user.userName: + raise HTTPException(400, "userName is required") + document: Final = user.model_copy( + update={"id": row.id, "active": user.active if "active" in user.model_fields_set else row.active} + ) + updated: Final = await tx.litellm_scimresource.update_many( + where={"id": row.id, "updated_at": row.updated_at}, + data={ + "user_name": user.userName, + "display_name": user.displayName or user.userName or row.display_name, + "document": Json(document.model_dump(by_alias=True, mode="json", exclude_none=True)), + "active": document.active, + }, + ) + if updated != 1: + raise HTTPException(409, "The provisioned subject changed concurrently; retry") + return document + + async def update_user(self, resource_id: str, user: SCIMUser | SCIMPatchOp) -> SCIMUser: + from litellm.proxy.management_endpoints.scim.human_provisioning import SourceHumanProvisioner + + async with self.source_transaction() as tx: + row: Final = await self._resource(tx, "Users", resource_id) + current: Final = user_document(row) + if current.agent_user is not None: + updated: Final = apply_user_patch(current, user) if isinstance(user, SCIMPatchOp) else user + if isinstance(updated, SCIMProvisioningFailure): + raise HTTPException(updated.status, updated.message) + return await self._update_native(tx, row, updated) + if isinstance(user, SCIMUser) and ( + user.agent_user is not None + or canonical_directory_id(user.externalId or "") != canonical_directory_id(row.external_id) + ): + raise HTTPException(409, "A human subject cannot be rebound or converted into an agent-user") + if isinstance(user, SCIMPatchOp) and patch_changes_identity(user): + raise HTTPException(409, "A human cannot be converted into an agent-user") + result: Final = await SourceHumanProvisioner(self.client, self.source).update_in_transaction(tx, row, user) + await SourceHumanProvisioner.finish_update(result) + return result.document + + @serialized_source + async def delete(self, kind: Literal["Users", "Groups"], resource_id: str) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + + where: Final[LiteLLM_SCIMResourceWhereUniqueInput] = {"id": resource_id} + async with self.client.tx() as tx: + row: Final = await tx.litellm_scimresource.find_unique(where=where) + if row is None or row.source_id != self.source.source_id or row.kind != kind: + raise HTTPException(404, "SCIM resource not found in this provisioning source") + if row.deleted: + return + if row.local_id is not None: + try: + if kind == "Groups": + await scim_v2.delete_group(group_id=row.local_id) + elif user_document(row).agent_user is None: + await scim_v2.delete_user(user_id=row.local_id) + except HTTPException as exc: + if exc.status_code != 404: + raise + async with self.client.tx() as tx: + if kind == "Users": + memberships: Final[LiteLLM_SCIMResourceWhereInput] = { + "source_id": self.source.source_id, + "kind": "Groups", + "member_ids": {"has": row.id}, + } + groups: Final = await tx.litellm_scimresource.find_many(where=memberships) + for group in groups: + await remove_group_member(tx, group, row.id) + retired: Final[LiteLLM_SCIMResourceUpdateInput] = {"active": False, "deleted": True, "member_ids": []} + await tx.litellm_scimresource.update(where=where, data=retired) + + async def create_group(self, group: SCIMGroup) -> SCIMGroup: + from litellm.proxy.management_endpoints.scim import scim_v2 + + admin_group: Final = await scim_v2.provisioning_group_admin_role() + async with self.source_transaction() as tx: + document, result = await self._create_group(tx, group, admin_group) + await scim_v2.finish_provisioned_group(result) + return document + + async def _create_group( + self, tx: Prisma, group: SCIMGroup, admin_group: str | None + ) -> tuple[SCIMGroup, "ProvisionedGroupWrite | None"]: + if not group.externalId: + raise HTTPException(400, "externalId is required for a directory group") + external_id: Final = canonical_directory_id(group.externalId) + old: Final = await tx.litellm_scimresource.find_unique( + where={ + "source_id_kind_external_id": { + "source_id": self.source.source_id, + "kind": "Groups", + "external_id": external_id, + } + } + ) + if old is not None and old.deleted: + raise HTTPException(409, "This directory group was deleted") + if old is not None: + return await self._update_group(tx, old.id, group, admin_group) + members: Final = tuple(dict.fromkeys(member.value for member in group.members or ())) + scim_id: Final = str(uuid4()) + document: Final = group.model_copy(update={"id": scim_id, "externalId": external_id}) + await self._validate_members(members, tx=tx) + resource_data: Final[LiteLLM_SCIMResourceCreateInput] = LiteLLM_SCIMResourceCreateInput( + id=scim_id, + source_id=self.source.source_id, + kind="Groups", + external_id=external_id, + display_name=group.displayName, + document=Json(document.model_dump(by_alias=True, mode="json", exclude_none=True)), + member_ids=list(members), + ) + row: Final = await tx.litellm_scimresource.create(data=resource_data) + result: Final = await self._sync_human_members(tx, row, admin_group) + return group_document(row), result + + async def _validate_members(self, members: tuple[str, ...], *, tx: Prisma | None = None) -> None: + if not members: + return + count: Final = await count_in( + SCIMResourceRepository(self.client, use_writer=True).table + if tx is None + else SCIMResourceRepository(_GroupDatabase(tx)).table, + "id", + members, + where=LiteLLM_SCIMResourceWhereInput(source_id=self.source.source_id, kind="Users", deleted=False), + ) + if count != len(frozenset(members)): + raise HTTPException(400, "Group members must exist in this provisioning source") + + async def update_group(self, resource_id: str, change: SCIMGroup | SCIMPatchOp) -> SCIMGroup: + from litellm.proxy.management_endpoints.scim import scim_v2 + + admin_group: Final = await scim_v2.provisioning_group_admin_role() + async with self.source_transaction() as tx: + document, result = await self._update_group(tx, resource_id, change, admin_group) + await scim_v2.finish_provisioned_group(result) + return document + + async def _update_group( + self, tx: Prisma, resource_id: str, change: SCIMGroup | SCIMPatchOp, admin_group: str | None + ) -> tuple[SCIMGroup, "ProvisionedGroupWrite | None"]: + old: Final = await self._resource(tx, "Groups", resource_id) + updated: Final = ( + group_members_after_patch(group_document(old), change) if isinstance(change, SCIMPatchOp) else change + ) + if isinstance(updated, SCIMProvisioningFailure): + raise HTTPException(updated.status, updated.message) + if canonical_directory_id(updated.externalId or "") != canonical_directory_id(old.external_id): + raise HTTPException(409, "Directory group externalId is immutable") + members: Final = tuple(dict.fromkeys(member.value for member in updated.members or ())) + await self._validate_members(members, tx=tx) + data: Final = LiteLLM_SCIMResourceUpdateInput( + display_name=updated.displayName, + document=Json(updated.model_dump(by_alias=True, mode="json", exclude_none=True)), + member_ids=list(members), + ) + count: Final = await tx.litellm_scimresource.update_many( + where=LiteLLM_SCIMResourceWhereInput(id=old.id, updated_at=old.updated_at), data=data + ) + if count != 1: + raise HTTPException(409, "The group changed concurrently; retry") + row: Final = await self._resource(tx, "Groups", resource_id) + result: Final = await self._sync_human_members(tx, row, admin_group) + return group_document(row), result + + async def _sync_human_members( + self, tx: Prisma, group: LiteLLM_SCIMResource, admin_group: str | None + ) -> "ProvisionedGroupWrite | None": + from litellm.proxy.management_endpoints.scim import scim_v2 + + users: Final = await find_many_in( + SCIMResourceRepository(_GroupDatabase(tx)).table, + "id", + group.member_ids, + where={ + "source_id": self.source.source_id, + "kind": "Users", + "deleted": False, + }, + ) + if any(row.local_id is None and user_document(row).agent_user is None for row in users): + raise HTTPException(409, "The human provisioned record is incomplete") + humans: Final = [ + SCIMMember(value=row.local_id) + for row in users + if row.local_id is not None and user_document(row).agent_user is None + ] + if not humans and group.local_id is None: + return None + document: Final = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group.local_id or group.id, + externalId=group.external_id, + displayName=group.display_name, + members=humans, + ) + result: Final = await scim_v2.write_provisioned_group(tx, self.client, document, admin_group) + if group.local_id is None: + where: Final[LiteLLM_SCIMResourceWhereUniqueInput] = {"id": group.id} + linked: Final[LiteLLM_SCIMResourceUpdateInput] = {"local_id": result.team_id} + await tx.litellm_scimresource.update(where=where, data=linked) + return result diff --git a/litellm/proxy/management_endpoints/scim/human_provisioning.py b/litellm/proxy/management_endpoints/scim/human_provisioning.py new file mode 100644 index 00000000000..fcd08a1e401 --- /dev/null +++ b/litellm/proxy/management_endpoints/scim/human_provisioning.py @@ -0,0 +1,255 @@ +from collections.abc import Mapping +from dataclasses import dataclass +from functools import reduce +from typing import Final +from uuid import uuid4 + +from fastapi import HTTPException +from prisma import Json, Prisma +from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource +from prisma.types import ( + LiteLLM_SCIMResourceCreateInput, + LiteLLM_SCIMResourceUpdateInput, + LiteLLM_SCIMResourceWhereUniqueInput, + LiteLLM_UserTableCreateInput, + LiteLLM_UserTableWhereInput, +) +from pydantic import TypeAdapter + +from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles +from litellm.proxy.management_endpoints.internal_user_endpoints import check_user_license_capacity +from litellm.proxy.utils import PrismaClient +from litellm.repositories.base_repository import is_unique_violation +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import canonical_directory_id +from litellm.types.proxy.management_endpoints.scim_v2 import SCIMPatchOp, SCIMPatchOperation, SCIMUser, SCIMUserEmail + + +def human_email(user: SCIMUser) -> str: + preferred: Final = next((email.value for email in user.emails or () if email.primary), None) + first: Final = user.emails[0].value if user.emails else None + email: Final = preferred or first or user.userName + if not email: + raise HTTPException(400, "userName or email is required") + return email.casefold() + + +def changes_readonly_attribute(operation: SCIMPatchOperation) -> bool: + path: Final = (operation.path or "").lower() + value: Final = operation.value + fields: Final = TypeAdapter(dict[str, object]).validate_python(value) if isinstance(value, dict) else {} + attributes: Final = tuple(key.lower() for key in fields) if not path else (path,) + return any(attribute.startswith(("groups", "externalid")) for attribute in attributes) + + +def changes_email_attribute(operation: SCIMPatchOperation) -> bool: + value: Final = operation.value + fields: Final = TypeAdapter(dict[str, object]).validate_python(value) if isinstance(value, dict) else {} + attributes: Final = (operation.path,) if operation.path else tuple(fields) + return any(attribute.casefold().split("[", 1)[0].split(".", 1)[0] == "emails" for attribute in attributes) + + +def validate_human_patch(patch: SCIMPatchOp) -> None: + if any(changes_email_attribute(operation) for operation in patch.Operations): + raise HTTPException(400, "Use PUT to replace a provisioned human's email") + if any(changes_readonly_attribute(operation) for operation in patch.Operations): + raise HTTPException(400, "externalId is immutable; update group membership through this source's Groups") + + +def patched_username(current: str | None, operation: SCIMPatchOperation) -> str | None: + direct: Final = bool(operation.path and operation.path.casefold() == "username") + fields: Final = ( + TypeAdapter(Mapping[str, object]).validate_python(operation.value) + if operation.path is None and isinstance(operation.value, dict) + else None + ) + if not direct and (fields is None or "userName" not in fields): + return current + candidate: Final = ( + None + if operation.op == "remove" + else operation.value + if direct + else fields["userName"] + if fields is not None + else None + ) + if not isinstance(candidate, str) or not candidate: + raise HTTPException(400, "userName is required") + return candidate + + +@dataclass(frozen=True, slots=True) +class _HumanWriteDatabase: + db: Prisma + + +@dataclass(frozen=True, slots=True) +class _HumanUpdate: + document: SCIMUser + local_id: str + tokens: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class SourceHumanProvisioner: + client: PrismaClient + source: LiteLLM_SCIMSource + + async def create(self, user: SCIMUser) -> SCIMUser: + async with self.client.tx() as tx: + result: Final = await self.create_in_transaction(tx, user) + await self.finish_update(result) + return result.document + + async def create_in_transaction(self, tx: Prisma, user: SCIMUser) -> _HumanUpdate: + from litellm.proxy.management_endpoints.scim.agent_provisioning import user_document + + if not user.externalId or not user.userName: + raise HTTPException(400, "externalId and userName are required") + try: + row: Final = await self.reserve_in_transaction(tx, user) + except Exception as exc: + if is_unique_violation(exc): + raise HTTPException(409, "This human identity belongs to another provisioning record") from exc + raise + if row.deleted: + raise HTTPException(409, "This subject was deleted; automatic recreation is not permitted") + if user_document(row).agent_user is not None: + raise HTTPException(409, "A provisioned agent-user cannot become a human") + return await self.update_in_transaction(tx, row, user) + + async def reserve(self, user: SCIMUser) -> LiteLLM_SCIMResource: + async with self.client.tx() as tx: + return await self.reserve_in_transaction(tx, user) + + async def reserve_in_transaction(self, tx: Prisma, user: SCIMUser) -> LiteLLM_SCIMResource: + if user.externalId is None or user.userName is None: + raise HTTPException(400, "externalId and userName are required") + external_id: Final = canonical_directory_id(user.externalId) + resource_filter: Final[LiteLLM_SCIMResourceWhereUniqueInput] = { + "source_id_kind_external_id": { + "source_id": self.source.source_id, + "kind": "Users", + "external_id": external_id, + } + } + existing: Final = await tx.litellm_scimresource.find_unique(where=resource_filter) + if existing is not None: + return existing + email: Final = human_email(user) + user_filter: Final[LiteLLM_UserTableWhereInput] = { + "OR": [ + {"user_id": user.userName}, + {"user_email": {"equals": email, "mode": "insensitive"}}, + ] + } + matches: Final = await tx.litellm_usertable.find_many(where=user_filter) + if matches: + raise HTTPException(409, "This local user already exists; automatic directory adoption is not permitted") + await check_user_license_capacity(self.client, users=UserRepository(_HumanWriteDatabase(tx))) + local_id: Final = user.userName + local_data: Final[LiteLLM_UserTableCreateInput] = { + "user_id": local_id, + "user_email": email, + "user_role": LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + "teams": [], + } + await tx.litellm_usertable.create(data=local_data) + scim_id: Final = str(uuid4()) + document: Final = user.model_copy(update={"id": scim_id, "externalId": external_id}) + data: Final = LiteLLM_SCIMResourceCreateInput( + id=scim_id, + source_id=self.source.source_id, + kind="Users", + external_id=external_id, + user_name=user.userName, + display_name=user.displayName or user.userName, + document=Json(document.model_dump(by_alias=True, mode="json", exclude_none=True)), + active=user.active, + local_id=local_id, + human_email=email, + human_subject_key=f"{self.source.tenant_id}:{external_id}", + ) + return await tx.litellm_scimresource.create(data=data) + + async def update(self, row: LiteLLM_SCIMResource, change: SCIMUser | SCIMPatchOp) -> SCIMUser: + async with self.client.tx() as tx: + result: Final = await self.update_in_transaction(tx, row, change) + await self.finish_update(result) + return result.document + + async def update_in_transaction( + self, tx: Prisma, row: LiteLLM_SCIMResource, change: SCIMUser | SCIMPatchOp + ) -> _HumanUpdate: + from litellm.proxy.management_endpoints.scim import scim_v2 + + if row.local_id is None: + raise HTTPException(409, "The human provisioned record is incomplete") + username: Final = ( + change.userName + if isinstance(change, SCIMUser) + else reduce(patched_username, change.Operations, row.user_name) + ) + if username is None: + raise HTTPException(400, "userName is required") + if isinstance(change, SCIMPatchOp): + validate_human_patch(change) + writer: Final = _HumanWriteDatabase(tx) + users: Final = UserRepository(writer) + existing: Final = await users.find_by_id(row.local_id, id_field="user_id") + if existing is None: + raise HTTPException(409, "The human local identity was removed; automatic recreation is not permitted") + if isinstance(change, SCIMUser): + await self.claim_email(tx, row, human_email(change)) + normalized: Final = ( + change.model_copy( + update={"groups": None, "emails": [SCIMUserEmail(value=human_email(change), primary=True)]} + ) + if isinstance(change, SCIMUser) + else change + ) + data, active, previous_active = scim_v2.prepare_provisioned_user_update(existing, normalized) + updated_row: Final = await users.table.update(where={"user_id": row.local_id}, data=data) + if updated_row is None: + raise HTTPException(409, "The human local identity was removed; automatic recreation is not permitted") + updated: Final = LiteLLM_UserTable.model_validate(updated_row.model_dump()) + tokens: Final = ( + await scim_v2.write_user_keys_blocked( + VerificationTokenRepository(writer).table, user_id=row.local_id, blocked=not active + ) + if active is not None and active != (True if previous_active is None else previous_active) + else () + ) + result: Final = await scim_v2.ScimTransformations.transform_litellm_user_to_scim_user(updated, tx=tx) + document: Final = result.model_copy(update={"id": row.id, "externalId": row.external_id, "userName": username}) + update_data: Final[LiteLLM_SCIMResourceUpdateInput] = { + "document": Json(document.model_dump(by_alias=True, mode="json", exclude_none=True)), + "active": document.active, + "user_name": document.userName, + } + await tx.litellm_scimresource.update(where={"id": row.id}, data=update_data) + return _HumanUpdate(document=document, local_id=row.local_id, tokens=tokens) + + @staticmethod + async def finish_update(result: _HumanUpdate) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + + await scim_v2.finish_provisioned_user_update(result.local_id, result.tokens) + + async def claim_email(self, tx: Prisma, row: LiteLLM_SCIMResource, email: str) -> None: + resource_filter: Final[LiteLLM_SCIMResourceWhereUniqueInput] = {"id": row.id} + update_data: Final[LiteLLM_SCIMResourceUpdateInput] = {"human_email": email} + try: + local_filter: Final[LiteLLM_UserTableWhereInput] = { + "user_email": {"equals": email, "mode": "insensitive"}, + "NOT": {"user_id": row.local_id}, + } + if await tx.litellm_usertable.find_many(where=local_filter): + raise HTTPException(409, "This email belongs to another local user") + await tx.litellm_scimresource.update(where=resource_filter, data=update_data) + except Exception as exc: + if is_unique_violation(exc): + raise HTTPException(409, "This email belongs to another provisioning record") from exc + raise diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 6a67093fde8..d8294eb9af3 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import Final, TypeVar +from typing import TYPE_CHECKING, Final, TypeVar from pydantic import ValidationError @@ -13,6 +13,9 @@ from litellm.proxy._types import ( from litellm.repositories.team_repository import TeamRepository from litellm.types.proxy.management_endpoints.scim_v2 import * +if TYPE_CHECKING: + from prisma import Prisma + T = TypeVar("T") @@ -25,17 +28,23 @@ class ScimTransformations: @staticmethod async def transform_litellm_user_to_scim_user( user: LiteLLM_UserTable | NewUserResponse, + *, + tx: "Prisma | None" = None, ) -> SCIMUser: from litellm.proxy.proxy_server import prisma_client - if prisma_client is None: + if prisma_client is None and tx is None: raise HTTPException(status_code=500, detail={"error": "No database connected"}) # Get user's teams/groups groups: Final = [] team_ids: Final[list[str]] = user.teams or [] # mutable-ok: scim reads the user row's team ids for team_id in team_ids: - team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) + team = ( + await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) + if tx is None + else await tx.litellm_teamtable.find_unique(where={"team_id": team_id}) + ) if team: team_alias = getattr(team, "team_alias", team.team_id) groups.append(SCIMUserGroup(value=team.team_id, display=team_alias)) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 0cf201b3a00..e914e0c8d79 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -53,10 +53,20 @@ from litellm.proxy.management_endpoints.scim.scim_transformations import ( ScimTransformations, ) from litellm.proxy.management_endpoints.team_endpoints import ( + CreatedTeam, + TeamMemberAddition, + TeamMemberRemoval, + add_team_members_in_transaction, + create_team, + delete_team_member_in_transaction, + finish_team_creation, + finish_team_member_addition, + finish_team_member_removal, new_team, team_member_add, team_member_delete, ) +from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL from litellm.proxy.utils import ( PrismaClient, _premium_user_check, @@ -75,6 +85,7 @@ from litellm.repositories.verification_token_repository import ( from litellm.types.proxy.management_endpoints.scim_v2 import * if TYPE_CHECKING: + from prisma import Prisma from prisma.models import LiteLLM_VerificationToken as PrismaVerificationToken @@ -439,7 +450,9 @@ def _resolve_scim_user_role( return default_role -async def _scim_groups_from_team_ids(prisma_client: PrismaClient, team_ids: list[str]) -> list[SCIMUserGroup]: +async def _scim_groups_from_team_ids( + prisma_client: "PrismaClient | _GroupWriteDatabase", team_ids: list[str] +) -> list[SCIMUserGroup]: """ Build SCIMUserGroup objects from team ids, populating display from each team's alias so admin-group matching by display name works the same way it @@ -468,6 +481,14 @@ async def _recompute_scim_member_roles(prisma_client: PrismaClient, user_ids: It if admin_group is None: return + await write_scim_member_roles(prisma_client, user_ids, admin_group) + + +async def write_scim_member_roles( + prisma_client: "PrismaClient | _GroupWriteDatabase", user_ids: Iterable[str], admin_group: str | None +) -> None: + if admin_group is None: + return default_role: Final = _default_scim_user_role() for user_id in user_ids: user = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": user_id}) @@ -588,7 +609,9 @@ async def _users_named_by_member_value( return tuple(dict.fromkeys(row.user_id for row in rows)) -async def _accounts_named_by_member_value(value: str, prisma_client: PrismaClient) -> tuple[str, ...]: +async def _accounts_named_by_member_value( + value: str, prisma_client: "PrismaClient | _GroupWriteDatabase" +) -> tuple[str, ...]: """Every user id this member value names, by user id, SSO identity or email. Classification needs to know whether the value is one account's ``user_id`` and @@ -1011,51 +1034,52 @@ async def _set_user_keys_blocked(user_id: str, blocked: bool) -> int: prisma_client: Final = await _get_prisma_client_or_raise_exception() - if blocked: - # `blocked` is a nullable column with no default, so existing rows - # typically hold NULL; treat NULL as "not blocked" since SQL equality - # on NULL would otherwise silently skip them. - candidates = await _table(VerificationTokenRepository(prisma_client)).find_many( - where={ - "user_id": user_id, - "OR": [{"blocked": False}, {"blocked": None}], - }, - ) - affected_keys = candidates - else: - candidates = await _table(VerificationTokenRepository(prisma_client)).find_many( - where={"user_id": user_id, "blocked": True}, - ) - affected_keys = [k for k in candidates if _key_was_scim_blocked(k.metadata)] - - if not affected_keys: - return 0 - - for key_row in affected_keys: - current_metadata: dict[str, object] = dict(key_row.metadata) if isinstance(key_row.metadata, dict) else {} - if blocked: - new_metadata = {**current_metadata, SCIM_BLOCKED_METADATA_KEY: True} - else: - new_metadata = {k: v for k, v in current_metadata.items() if k != SCIM_BLOCKED_METADATA_KEY} - await _table(VerificationTokenRepository(prisma_client)).update( - where={"token": key_row.token}, - data={"blocked": blocked, "metadata": safe_dumps(new_metadata)}, - ) - - for key_row in affected_keys: + tokens: Final = await write_user_keys_blocked( + _table(VerificationTokenRepository(prisma_client)), user_id=user_id, blocked=blocked + ) + for token in tokens: await _delete_cache_key_object( - hashed_token=key_row.token, + hashed_token=token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) + if tokens: + verbose_proxy_logger.info( + "SCIM: %s %d virtual key(s) for user_id=%s", + "blocked" if blocked else "unblocked", + len(tokens), + user_id, + ) + return len(tokens) - verbose_proxy_logger.info( - "SCIM: %s %d virtual key(s) for user_id=%s", - "blocked" if blocked else "unblocked", - len(affected_keys), - user_id, + +async def write_user_keys_blocked( + table: _VerificationTokenTableClient, *, user_id: str, blocked: bool +) -> tuple[str, ...]: + candidates: Final = await table.find_many( + where=( + {"user_id": user_id, "OR": [{"blocked": False}, {"blocked": None}]} + if blocked + else {"user_id": user_id, "blocked": True} + ) ) - return len(affected_keys) + affected_keys: Final = tuple(key for key in candidates if blocked or _key_was_scim_blocked(key.metadata)) + for key_row, current_metadata in ( + (key, TypeAdapter(dict[str, object]).validate_python(key.metadata) if isinstance(key.metadata, dict) else {}) + for key in affected_keys + ): + await table.update( + where={"token": key_row.token}, + data={ + "blocked": blocked, + "metadata": safe_dumps( + {**current_metadata, SCIM_BLOCKED_METADATA_KEY: True} + if blocked + else {key: value for key, value in current_metadata.items() if key != SCIM_BLOCKED_METADATA_KEY} + ), + }, + ) + return tuple(key.token for key in affected_keys) async def _delete_rows_referencing_user(prisma_client: PrismaClient, *, user_id: str) -> None: @@ -1553,7 +1577,7 @@ async def get_service_provider_config(request: Request): return SCIMServiceProviderConfig(meta=meta) -def _parse_scim_eq_filter(scim_filter: str) -> tuple[str, str] | None: +def parse_scim_eq_filter(scim_filter: str) -> tuple[str, str] | None: """Parse the SCIM equality filters Okta uses before user lifecycle changes.""" match: Final = re.match( r"""\s*([\w.]+)\s+eq\s+(['"]?)(.*?)\2\s*$""", @@ -1595,7 +1619,7 @@ async def get_users( # Okta locates users by userName before deprovisioning. LiteLLM # exposes SCIM userName from user_email, while older SCIM-created # users may still have user_id == userName, so support both. - parsed_filter: Final = _parse_scim_eq_filter(filter) + parsed_filter: Final = parse_scim_eq_filter(filter) if parsed_filter: filter_attribute, filter_value = parsed_filter if filter_attribute == "username": @@ -1733,6 +1757,70 @@ async def create_user( raise handle_exception_on_proxy(e) +@dataclass(frozen=True, slots=True) +class _UserReplacement: + data: Mapping[str, object] + metadata: Mapping[str, object] + teams: Sequence[str] + + +def _prepare_user_replacement(existing_user: LiteLLM_UserTable, user: SCIMUser) -> _UserReplacement: + user_data: Final = _extract_scim_user_data(user) + metadata: Final = _build_scim_metadata( + user_data["given_name"], + user_data["family_name"], + user_data["active"] if "active" in user.model_fields_set else _user_scim_active(existing_user), + enterprise=user_data["enterprise"], + entitlements=user_data["entitlements"], + roles=user_data["roles"], + ) + teams: Final = user_data["teams"] or existing_user.teams + return _UserReplacement( + data={ + "user_email": user_data["user_email"], + "user_alias": user_data["user_alias"], + "sso_user_id": user_data["sso_user_id"], + "teams": teams, + "metadata": safe_dumps(metadata), + }, + metadata=metadata, + teams=teams, + ) + + +def prepare_provisioned_user_update( + existing: LiteLLM_UserTable, change: SCIMUser | SCIMPatchOp +) -> tuple[Mapping[str, object], bool | None, bool | None]: + previous_active: Final = _user_scim_active(existing) + if isinstance(change, SCIMUser): + replacement: Final = _prepare_user_replacement(existing, change) + return replacement.data, _scim_active_value(replacement.metadata), previous_active + updates, teams = _apply_patch_ops(existing, change) + metadata: Final = updates.get("metadata") + typed_metadata: Final = ( + TypeAdapter(Mapping[str, object]).validate_python(metadata) if isinstance(metadata, Mapping) else None + ) + return ( + { + **updates, + "teams": list(teams), + "metadata": safe_dumps(metadata), + }, + _scim_active_value(typed_metadata), + previous_active, + ) + + +async def finish_provisioned_user_update(user_id: str, tokens: tuple[str, ...]) -> None: + from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache + + await evict_and_broadcast(cache_keys=(user_id,), user_api_key_cache=user_api_key_cache) + for token in tokens: + await _delete_cache_key_object( + hashed_token=token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj + ) + + @scim_router.put( "/Users/{user_id}", response_model=SCIMUser, @@ -1759,40 +1847,17 @@ async def update_user( prev_active: Final = _user_scim_active(existing_user) user_data: Final = _extract_scim_user_data(user) - - # SCIM PUT may legally omit `active` (full-replace with the field absent). - # Pydantic fills the model default, so distinguish "client sent active" - # from "client omitted it" via model_fields_set, and preserve the prior - # SCIM active state when omitted — otherwise a vanilla PUT to a - # deactivated user would silently re-enable them and unblock their keys. + replacement: Final = _prepare_user_replacement(existing_user, user) client_set_active: Final = "active" in user.model_fields_set - scim_active_for_metadata: Final = user_data["active"] if client_set_active else prev_active - - metadata: Final = _build_scim_metadata( - user_data["given_name"], - user_data["family_name"], - scim_active_for_metadata, - enterprise=user_data["enterprise"], - entitlements=user_data["entitlements"], - roles=user_data["roles"], - ) - - # SCIM User.groups is readOnly (RFC 7643 4.1.2): IdPs sync membership via /Groups and send - # no groups or `groups: []` on profile PUTs, so empty means unspecified, not "remove from every team" - target_teams: Final = user_data["teams"] or existing_user.teams + metadata: Final = replacement.metadata + target_teams: Final = replacement.teams await _handle_team_membership_changes( user_id=user_id, existing_teams=existing_user.teams, new_teams=target_teams, ) - update_data: Final = { - "user_email": user_data["user_email"], - "user_alias": user_data["user_alias"], - "sso_user_id": user_data["sso_user_id"], - "teams": target_teams, - "metadata": safe_dumps(metadata), - } + update_data: Final = dict(replacement.data) admin_group: Final = await _get_scim_admin_group() if admin_group is not None and user_data["teams"]: @@ -2539,6 +2604,144 @@ def _new_team_request_with_defaults( ) +@dataclass(frozen=True, slots=True) +class _GroupWriteDatabase: + db: "Prisma" + + @property + def writer_db(self) -> "Prisma": + return self.db + + +@dataclass(frozen=True, slots=True) +class ProvisionedGroupWrite: + team_id: str + created: CreatedTeam | None + removals: tuple[TeamMemberRemoval, ...] + additions: tuple[TeamMemberAddition, ...] + + +async def provisioning_group_admin_role() -> str | None: + return await _get_scim_admin_group() + + +async def _validate_provisioned_group_member(tx: "Prisma", database: _GroupWriteDatabase, user_id: str) -> None: + local_rows: Final = await tx.query_raw( + 'SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = $1 FOR KEY SHARE', user_id + ) + if not local_rows: + raise HTTPException(409, "The human local identity was removed; automatic recreation is not permitted") + named: Final = await _accounts_named_by_member_value(user_id, database) + if named != (user_id,): + raise HTTPException(400, "Group member identity is ambiguous") + + +async def write_provisioned_group( + tx: "Prisma", client: PrismaClient, group: SCIMGroup, admin_group: str | None +) -> ProvisionedGroupWrite: + if group.id is None: + raise HTTPException(409, "The provisioned group identity is incomplete") + database: Final = _GroupWriteDatabase(tx) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, group.id) + member_ids: Final = tuple(sorted(frozenset(member.value for member in group.members or ()))) + for user_id in member_ids: + await _validate_provisioned_group_member(tx, database, user_id) + existing: Final = await TeamRepository(database).find_by_id(group.id, id_field="team_id") + if existing is None: + created: Final = await create_team( + _new_team_request_with_defaults( + group.id, group.displayName, tuple(Member(user_id=value, role="user") for value in member_ids) + ), + auth, + transaction=tx, + ) + await write_scim_member_roles(database, member_ids, admin_group) + return ProvisionedGroupWrite(team_id=group.id, created=created, removals=(), additions=()) + current: Final = frozenset(await _get_team_member_user_ids_from_team(existing)) + final: Final = frozenset(member_ids) + await _table(TeamRepository(database)).update( + where={"team_id": group.id}, data=_group_replacement_data(existing, group) + ) + addition_result: Final = await _add_provisioned_group_members(tx, client, existing, final - current, auth) + removals: Final = tuple( + [ + await delete_team_member_in_transaction( + tx=tx, + data=TeamMemberDeleteRequest(team_id=group.id, user_id=user_id), + existing_team_row=existing, + prisma_client=client, + user_api_key_dict=auth, + ) + for user_id in sorted(current - final) + ] + ) + await write_scim_member_roles( + database, current | final if existing.team_alias != group.displayName else current ^ final, admin_group + ) + return ProvisionedGroupWrite( + team_id=group.id, + created=None, + removals=removals, + additions=(addition_result,) if addition_result is not None else (), + ) + + +async def _add_provisioned_group_members( + tx: "Prisma", + client: PrismaClient, + team: LiteLLM_TeamTable, + member_ids: frozenset[str], + auth: UserAPIKeyAuth, +) -> TeamMemberAddition | None: + from litellm.proxy.proxy_server import litellm_proxy_admin_name + + if not member_ids: + return None + before: Final = tuple(team.members_with_roles) + members: Final = [Member(user_id=user_id, role="user") for user_id in sorted(member_ids)] + _, users, memberships = await add_team_members_in_transaction( + tx=tx, + data=TeamMemberAddRequest(team_id=team.team_id, member=members), + complete_team_data=team, + prisma_client=client, + user_api_key_dict=auth, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + return TeamMemberAddition( + team=team.model_copy(deep=True), + before=before, + users=tuple(users), + memberships=tuple(memberships), + existing_user_ids=member_ids, + ) + + +async def finish_provisioned_group(result: ProvisionedGroupWrite | None) -> None: + if result is None: + return + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + if result.created is not None: + await finish_team_creation(result.created, auth) + for addition in result.additions: + await finish_team_member_addition(addition, auth) + for removal in result.removals: + await finish_team_member_removal(removal, auth) + + +def _group_replacement_data(existing: LiteLLM_TeamTable, group: SCIMGroup) -> Mapping[str, object]: + return { + "team_alias": group.displayName, + "metadata": safe_dumps( + { + **(existing.metadata or {}), + SCIM_TEAM_DATA_METADATA_KEY: group.model_dump(), + SCIM_MANAGED_TEAM_METADATA_KEY: True, + } + ), + } + + @scim_router.post( "/Groups", response_model=SCIMGroup, @@ -2620,18 +2823,7 @@ async def update_group( verbose_proxy_logger.debug("SCIM PUT GROUP all_member_ids: %s", member_result.all_member_ids) verbose_proxy_logger.debug("SCIM PUT GROUP created_users: %s", len(member_result.created_users)) - # Prepare update data - existing_metadata: Final = existing_team.metadata if existing_team.metadata else {} - updated_metadata: Final = { - **existing_metadata, - SCIM_TEAM_DATA_METADATA_KEY: group.model_dump(), - SCIM_MANAGED_TEAM_METADATA_KEY: True, - } - - update_data: Final = { - "team_alias": group.displayName, - "metadata": safe_dumps(updated_metadata), - } + update_data: Final = _group_replacement_data(existing_team, group) # Update team in database updated_team: Final = await _table(TeamRepository(prisma_client)).update( diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index a66d781dd61..78901142d43 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -179,11 +179,12 @@ from litellm.proxy.management_helpers.utils import ( from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.organization_repository import OrganizationRepository -from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.prisma_protocols import DatabaseClient, TableActions from litellm.repositories.table_repositories import ( AccessGroupRepository, DeletedTeamRepository, ModelTableRepository, + ObjectPermissionRepository, OrganizationMembershipRepository, TeamMembershipRepository, ) @@ -393,7 +394,7 @@ UPDATE "LiteLLM_UserTable" SET teams = array_remove(teams, $1) WHERE $1 = ANY(te _INCLUDE_MODEL_TABLE: Final = MappingProxyType({"litellm_model_table": True}) -def _team_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_TeamTable]": +def _team_db(prisma_client: DatabaseClient | None) -> "TableActions[prisma_models.LiteLLM_TeamTable]": return TeamRepository(prisma_client).table @@ -520,7 +521,7 @@ class TeamMemberBudgetHandler: metadata.pop(key, None) @staticmethod - def should_create_budget( + def shouldcreate_budget( team_member_budget: float | None = None, team_member_rpm_limit: int | None = None, team_member_tpm_limit: int | None = None, @@ -546,6 +547,7 @@ class TeamMemberBudgetHandler: team_member_tpm_limit: int | None = None, team_member_budget_duration: str | None = None, explicitly_set_fields: AbstractSet[str] = frozenset(), + table: "TableActions[prisma_models.LiteLLM_BudgetTable] | None" = None, ) -> dict: """Create team member budget table with provided limits. @@ -554,8 +556,10 @@ class TeamMemberBudgetHandler: """ from litellm.proxy._types import BudgetNewRequest from litellm.proxy.management_endpoints.budget_management_endpoints import ( + create_budget, new_budget, ) + from litellm.proxy.proxy_server import litellm_proxy_admin_name if data.team_alias is not None: budget_id = f"team-{data.team_alias.replace(' ', '-')}-budget-{uuid.uuid4().hex}" @@ -581,9 +585,14 @@ class TeamMemberBudgetHandler: if team_member_budget_duration is not None: budget_request.budget_duration = team_member_budget_duration - team_member_budget_table: Final = await _as_budget_write(new_budget)( - budget_obj=budget_request, - user_api_key_dict=user_api_key_dict, + team_member_budget_table: Final = ( + await _as_budget_write(new_budget)(budget_obj=budget_request, user_api_key_dict=user_api_key_dict) + if table is None + else await create_budget( + budget_obj=budget_request, + table=table, + created_by=user_api_key_dict.user_id or litellm_proxy_admin_name, + ) ) # Add team_member_budget_id as metadata field to team table @@ -1005,7 +1014,7 @@ def check_org_team_rpm_tpm_limits( async def _check_org_team_limits( org_table: LiteLLM_OrganizationTable, data: NewTeamRequest | UpdateTeamRequest, - prisma_client: PrismaClient, + prisma_client: DatabaseClient, ) -> None: """ Check organization team limits including: @@ -1425,14 +1434,43 @@ async def new_team( }' ``` """ + result: Final = await create_team(data, user_api_key_dict) + await finish_team_creation(result, user_api_key_dict, litellm_changed_by) + try: + return result.team.model_dump() + except Exception: + return result.team.dict() + + +@dataclass(frozen=True, slots=True) +class _TeamWriteDatabase: + db: "Prisma" + + +@dataclass(frozen=True, slots=True) +class CreatedTeam: + team: "prisma_models.LiteLLM_TeamTable" + snapshot: LiteLLM_TeamTable + access_groups: tuple[str, ...] + + +async def write_team_creation( + tx: _TeamCreateTx, data: Mapping[str, object] +) -> tuple["prisma_models.LiteLLM_TeamTable", tuple[str, ...]]: + team: Final = await tx.litellm_teamtable.create(data=data, include=_INCLUDE_MODEL_TABLE) + groups: Final = await reconcile_team_access_group_membership(tx, team.team_id) + return team, tuple(groups) + + +async def create_team( + data: NewTeamRequest, + user_api_key_dict: UserAPIKeyAuth, + *, + transaction: "Prisma | None" = None, +) -> CreatedTeam: try: - from litellm.proxy.management_helpers.audit_logs import ( - get_audit_log_changed_by, - is_audit_logging_enabled, - ) from litellm.proxy.proxy_server import ( _license_check, - create_audit_log_for_update, general_settings, litellm_proxy_admin_name, llm_router, @@ -1444,6 +1482,8 @@ async def new_team( if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) + writer: Final = prisma_client if transaction is None else _TeamWriteDatabase(transaction) + # Validate budget values are not negative if data.max_budget is not None and (not math.isfinite(data.max_budget) or data.max_budget < 0): raise HTTPException( @@ -1494,7 +1534,9 @@ async def new_team( ) # Check if license is over limit - total_teams: Final = await _team_db(prisma_client).count() + total_teams: Final = await ( + _team_db(prisma_client) if transaction is None else _team_tx_db(transaction) + ).count() if total_teams and _license_check.is_team_count_over_limit(team_count=total_teams): raise HTTPException( status_code=403, @@ -1512,8 +1554,10 @@ async def new_team( }, ) # Check if team_id exists already - _existing_team_id: Final = await prisma_client.get_data( - team_id=data.team_id, table_name="team", query_type="find_unique" + _existing_team_id: Final = ( + await prisma_client.get_data(team_id=data.team_id, table_name="team", query_type="find_unique") + if transaction is None + else await _team_tx_db(transaction).find_unique(where={"team_id": data.team_id}) ) if _existing_team_id is not None: raise HTTPException( @@ -1558,12 +1602,20 @@ async def new_team( # check org key limits - done here to handle inheriting org id from team if data.organization_id is not None and prisma_client is not None: try: - org_table = await get_org_object( - org_id=data.organization_id, - user_api_key_cache=user_api_key_cache, - prisma_client=prisma_client, - include_budget_table=True, - ) + if transaction is None: + org_table = await get_org_object( + org_id=data.organization_id, + user_api_key_cache=user_api_key_cache, + prisma_client=prisma_client, + include_budget_table=True, + ) + else: + org_row: Final = await OrganizationRepository(writer).table.find_unique( + where={"organization_id": data.organization_id}, include={"litellm_budget_table": True} + ) + org_table = ( + LiteLLM_OrganizationTable.model_validate(org_row.model_dump()) if org_row is not None else None + ) except OrganizationNotFoundError: org_table = None if org_table is None: @@ -1575,7 +1627,7 @@ async def new_team( await _check_org_team_limits( org_table=org_table, data=data, - prisma_client=prisma_client, + prisma_client=writer, ) if ( @@ -1621,7 +1673,7 @@ async def new_team( await validate_router_settings_weights( data.router_settings, team_id=data.team_id, - prisma_client=prisma_client, + prisma_client=writer, llm_router=llm_router, ) @@ -1633,7 +1685,10 @@ async def new_team( created_by=user_api_key_dict.user_id or litellm_proxy_admin_name, updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name, ) - model_dict: Final = await _model_db(prisma_client).create({**litellm_modeltable.json(exclude_none=True)}) + model_table: Final = ModelTableRepository( + prisma_client if transaction is None else _TeamWriteDatabase(transaction) + ).table + model_dict: Final = await model_table.create({**litellm_modeltable.json(exclude_none=True)}) _model_id = model_dict.id @@ -1648,10 +1703,13 @@ async def new_team( ) data_json = await _set_object_permission( data_json=data_json, - prisma_client=prisma_client, + prisma_client=writer, + table=ObjectPermissionRepository(_TeamWriteDatabase(transaction)).table + if transaction is not None + else None, ) - if TeamMemberBudgetHandler.should_create_budget( + if TeamMemberBudgetHandler.shouldcreate_budget( team_member_budget=data.team_member_budget, team_member_rpm_limit=data.team_member_rpm_limit, team_member_tpm_limit=data.team_member_tpm_limit, @@ -1666,6 +1724,7 @@ async def new_team( team_member_tpm_limit=data.team_member_tpm_limit, team_member_budget_duration=data.team_member_budget_duration, explicitly_set_fields=data.model_fields_set, + table=BudgetRepository(_TeamWriteDatabase(transaction)).table if transaction is not None else None, ) ## ADD TO TEAM TABLE @@ -1731,65 +1790,83 @@ async def new_team( complete_team_data_dict = prisma_client.jsonify_team_object(db_data=complete_team_data_dict) team_creation_data: Final[Mapping[str, object]] = complete_team_data_dict - tx: _TeamCreateTx - async with prisma_client.db.tx() as tx: - team_row: Final[prisma_models.LiteLLM_TeamTable] = await tx.litellm_teamtable.create( - data=team_creation_data, - include=_INCLUDE_MODEL_TABLE, - ) - affected_access_groups: Final = await reconcile_team_access_group_membership(tx, team_row.team_id) - - await invalidate_access_group_caches(affected_access_groups) + if transaction is None: + async with prisma_client.db.tx() as tx: + team_row, affected_access_groups = await write_team_creation(tx, team_creation_data) + await invalidate_access_group_caches(affected_access_groups) + else: + team_row, affected_access_groups = await write_team_creation(transaction, team_creation_data) ## ADD TEAM ID TO USER TABLE ## team_member_add_request: Final = TeamMemberAddRequest( team_id=data.team_id, member=members_with_roles, ) - await _add_team_members_to_team( - data=team_member_add_request, - complete_team_data=team_row, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, + if transaction is None: + await _add_team_members_to_team( + data=team_member_add_request, + complete_team_data=team_row, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + else: + await add_team_members_in_transaction( + tx=transaction, + data=team_member_add_request, + complete_team_data=LiteLLM_TeamTable.model_validate(team_row.model_dump()), + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + + return CreatedTeam( + team=team_row, + snapshot=complete_team_data, + access_groups=tuple(affected_access_groups) if transaction is not None else (), ) - - if is_audit_logging_enabled(): - created_team_snapshot: Final = complete_team_data.model_copy( - update={"members_with_roles": list(team_row.members_with_roles)} - ) - _updated_values = created_team_snapshot.json(exclude_none=True) - - _updated_values = json.dumps(_updated_values, default=str) - - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=get_audit_log_changed_by( - litellm_changed_by=litellm_changed_by, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ), - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.TEAM_TABLE_NAME, - object_id=data.team_id, - action="created", - updated_values=_updated_values, - before_value=None, - ) - ) - ) - - try: - return team_row.model_dump() - except Exception: - return team_row.dict() except Exception as e: raise handle_exception_on_proxy(e) +async def finish_team_creation( + result: CreatedTeam, + user_api_key_dict: UserAPIKeyAuth, + litellm_changed_by: str | None = None, +) -> None: + from litellm.proxy.management_helpers.audit_logs import get_audit_log_changed_by, is_audit_logging_enabled + from litellm.proxy.proxy_server import create_audit_log_for_update, litellm_proxy_admin_name + + await invalidate_access_group_caches(result.access_groups) + if is_audit_logging_enabled(): + created_team_snapshot: Final = result.snapshot.model_copy( + update={"members_with_roles": list(result.team.members_with_roles)} + ) + _updated_values = created_team_snapshot.json(exclude_none=True) + + _updated_values = json.dumps(_updated_values, default=str) + + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=get_audit_log_changed_by( + litellm_changed_by=litellm_changed_by, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ), + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.TEAM_TABLE_NAME, + object_id=result.team.team_id, + action="created", + updated_values=_updated_values, + before_value=None, + ) + ) + ) + + async def _create_team_update_audit_log( existing_team_row: _AuditableTeamRow, updated_kv: dict, @@ -2455,7 +2532,7 @@ async def update_team( }, } - if _team_member_fields_in_request and TeamMemberBudgetHandler.should_create_budget( + if _team_member_fields_in_request and TeamMemberBudgetHandler.shouldcreate_budget( team_member_budget=data.team_member_budget, team_member_rpm_limit=data.team_member_rpm_limit, team_member_tpm_limit=data.team_member_tpm_limit, @@ -2992,37 +3069,56 @@ async def _add_team_members_to_team( a waiter that can deadlock the pool, since enough concurrent adds for one team would hold every connection waiting on the lock while the holder waits for a free one. """ - gone_detail: Final[_ErrorDetail] = {"error": f"Team={data.team_id} was deleted while this member add was running"} async with prisma_client.tx() as tx: - await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) - - locked_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) - if locked_members is None: - raise HTTPException(status_code=404, detail=gone_detail) - complete_team_data.members_with_roles = locked_members - - updated_users, updated_team_memberships = await _process_team_members( + return await add_team_members_in_transaction( + tx=tx, data=data, complete_team_data=complete_team_data, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, litellm_proxy_admin_name=litellm_proxy_admin_name, - tx=tx, ) - await _update_team_members_list( - data=data, - complete_team_data=complete_team_data, - updated_users=updated_users, - ) - _db_team_members: Final = [m.model_dump() for m in complete_team_data.members_with_roles] - updated_team: Final = await _team_tx_db(tx).update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members)}, - ) - if updated_team is None: - raise HTTPException(status_code=404, detail=gone_detail) +async def add_team_members_in_transaction( + *, + tx: "Prisma", + data: TeamMemberAddRequest, + complete_team_data: LiteLLM_TeamTable, + prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, + litellm_proxy_admin_name: str, +) -> tuple["prisma_models.LiteLLM_TeamTable", list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: + gone_detail: Final[_ErrorDetail] = {"error": f"Team={data.team_id} was deleted while this member add was running"} + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) + + locked_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) + if locked_members is None: + raise HTTPException(status_code=404, detail=gone_detail) + complete_team_data.members_with_roles = locked_members + + updated_users, updated_team_memberships = await _process_team_members( + data=data, + complete_team_data=complete_team_data, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + tx=tx, + ) + + await _update_team_members_list( + data=data, + complete_team_data=complete_team_data, + updated_users=updated_users, + ) + + _db_team_members: Final = [m.model_dump() for m in complete_team_data.members_with_roles] + updated_team: Final = await _team_tx_db(tx).update( + where={"team_id": data.team_id}, + data={"members_with_roles": json.dumps(_db_team_members)}, + ) + if updated_team is None: + raise HTTPException(status_code=404, detail=gone_detail) return updated_team, updated_users, updated_team_memberships @@ -3375,7 +3471,6 @@ async def team_member_add( ``` """ - from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast from litellm.proxy.proxy_server import ( litellm_proxy_admin_name, premium_user, @@ -3470,27 +3565,15 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) - await evict_and_broadcast( - cache_keys=tuple(sorted(user.user_id for user in updated_users)), - user_api_key_cache=user_api_key_cache, - ) - await _evict_created_membership_caches( - user_ids=(tm.user_id for tm in updated_team_memberships), - team_id=data.team_id, - user_api_key_cache=user_api_key_cache, - ) - - _emit_team_members_metric(complete_team_data) - - _schedule_team_member_add_audit_logs( - team_id=data.team_id, - team_alias=complete_team_data.team_alias, - updated_users=updated_users, - existing_user_ids=pre_existing_user_ids, - before_members=members_before_add, - after_members=tuple(complete_team_data.members_with_roles), - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, + await finish_team_member_addition( + TeamMemberAddition( + team=complete_team_data, + before=members_before_add, + users=tuple(updated_users), + memberships=tuple(updated_team_memberships), + existing_user_ids=pre_existing_user_ids, + ), + user_api_key_dict, ) return TeamAddMemberResponse.model_validate( @@ -3502,6 +3585,42 @@ async def team_member_add( ) +@dataclass(frozen=True, slots=True) +class TeamMemberAddition: + team: LiteLLM_TeamTable + before: tuple[Member, ...] + users: tuple[LiteLLM_UserTable, ...] + memberships: tuple[LiteLLM_TeamMembership, ...] + existing_user_ids: frozenset[str] + + +async def finish_team_member_addition(result: TeamMemberAddition, user_api_key_dict: UserAPIKeyAuth) -> None: + from litellm.proxy.proxy_server import litellm_proxy_admin_name, user_api_key_cache + + await evict_and_broadcast( + cache_keys=tuple(sorted(user.user_id for user in result.users)), + user_api_key_cache=user_api_key_cache, + ) + await _evict_created_membership_caches( + user_ids=(tm.user_id for tm in result.memberships), + team_id=result.team.team_id, + user_api_key_cache=user_api_key_cache, + ) + + _emit_team_members_metric(result.team) + + _schedule_team_member_add_audit_logs( + team_id=result.team.team_id, + team_alias=result.team.team_alias, + updated_users=result.users, + existing_user_ids=result.existing_user_ids, + before_members=result.before, + after_members=tuple(result.team.members_with_roles), + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + + def _is_member_addressed_by(member: Member, data: TeamMemberDeleteRequest) -> bool: return (data.user_id is not None and member.user_id is not None and data.user_id == member.user_id) or ( data.user_email is not None and member.user_email is not None and data.user_email == member.user_email @@ -3576,11 +3695,7 @@ async def _team_member_delete( data: TeamMemberDeleteRequest, user_api_key_dict: UserAPIKeyAuth, ) -> tuple[LiteLLM_TeamTable, tuple[Member, ...], tuple[Member, ...]]: - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) @@ -3615,139 +3730,170 @@ async def _team_member_delete( }, ) - ## DELETE MEMBER FROM TEAM - # Everything from here on runs under the team's advisory lock, the same one - # /team/member_add and /team/delete take: without it, this endpoint's own row-level - # update lock used to be the only thing serializing it against a concurrent member_add, - # and only by accident (their SELECT ... FOR UPDATE contended for the same row lock this - # UPDATE takes). Now that member_add reads under the advisory lock instead, this has to - # take it too, and re-read the roster under it rather than off the snapshot validated - # above, or a member_add that commits in between can have its addition silently - # overwritten by this delete computing from stale data. async with prisma_client.tx() as tx: - await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) + removal: Final = await delete_team_member_in_transaction( + tx=tx, + data=data, + existing_team_row=existing_team_row, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + await finish_team_member_removal(removal, user_api_key_dict) + return removal.team, removal.before, removal.after - fresh_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) - if fresh_members is None: - raise HTTPException( - status_code=400, - detail={"error": f"Team id={data.team_id} does not exist in db"}, + +@dataclass(frozen=True, slots=True) +class TeamMemberRemoval: + team: LiteLLM_TeamTable + before: tuple[Member, ...] + after: tuple[Member, ...] + keys: tuple["prisma_models.LiteLLM_VerificationToken", ...] + jwt_mapping_cache_keys: tuple[str, ...] + user_ids: frozenset[str] + + +async def delete_team_member_in_transaction( + *, + tx: "Prisma", + data: TeamMemberDeleteRequest, + existing_team_row: LiteLLM_TeamTable, + prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, +) -> TeamMemberRemoval: + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) + + fresh_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) + if fresh_members is None: + raise HTTPException( + status_code=400, + detail={"error": f"Team id={data.team_id} does not exist in db"}, + ) + + removed_team_members, new_team_members = _cleanup_members_with_roles( + existing_team_row=LiteLLM_TeamTable(team_id=data.team_id, members_with_roles=fresh_members), + data=data, + ) + + existing_team_row.members_with_roles = new_team_members + + _db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members] + + ## DELETE TEAM ID from USER ROW, IF EXISTS ## + # get user row + removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None) + addressed_user_ids: Final = ( + removed_user_ids if removed_team_members else frozenset((data.user_id,) if data.user_id is not None else ()) + ) + key_val: Final[Mapping[str, object]] = ( + {"user_id": {"in": sorted(addressed_user_ids)}} if addressed_user_ids else {"user_email": data.user_email} + ) + member_tx: Final[_MemberDeleteTx] = tx + existing_user_rows: Final = await member_tx.litellm_usertable.find_many(where=key_val) + + # A user row can outlive its roster entry, and until the team is off user.teams the user + # still sees it and still fails key creation against it, so removal has to clear it too + stale_user_rows: Final = tuple(user for user in existing_user_rows if data.team_id in user.teams) + + # Also clean up any existing team membership rows for this user and team. An email can + # match several user rows, so with no roster entry to name the member, only the rows + # actually carrying the team are the ones this request is allowed to touch + cleanup_user_rows: Final = existing_user_rows if removed_team_members else stale_user_rows + user_ids_to_delete: Final = addressed_user_ids.union(user.user_id for user in cleanup_user_rows if user.user_id) + + if not removed_team_members and not stale_user_rows: + raise HTTPException(status_code=400, detail={"error": "User not found in team"}) + + ## DELETE KEYS CREATED BY USER FOR THIS TEAM + # Fetch keys before deletion so their audit records can be persisted alongside the delete. + # An empty user_ids_to_delete still resolves cleanly: prisma's "in": [] matches no rows. + keys_to_delete: Final = await member_tx.litellm_verificationtoken.find_many( + where={ + "user_id": {"in": sorted(user_ids_to_delete)}, + "team_id": data.team_id, + } + ) + jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens( + hashed_tokens=tuple(key.token for key in keys_to_delete), + prisma_client=_TeamWriteDatabase(tx), + ) + + if removed_team_members: + await _team_tx_db(tx).update( + where={"team_id": data.team_id}, + data={"members_with_roles": json.dumps(_db_new_team_members)}, + ) + + for existing_user in stale_user_rows: + await tx.litellm_usertable.update( + where={"user_id": existing_user.user_id}, + data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}}, + ) + + for _uid in sorted(user_ids_to_delete): + await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid}) + + if user_ids_to_delete: + if keys_to_delete: + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _persist_deleted_verification_tokens, ) - removed_team_members, new_team_members = _cleanup_members_with_roles( - existing_team_row=LiteLLM_TeamTable(team_id=data.team_id, members_with_roles=fresh_members), - data=data, - ) + await _persist_deleted_verification_tokens( + keys=keys_to_delete, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + tx=tx, + ) - existing_team_row.members_with_roles = new_team_members - - _db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members] - - ## DELETE TEAM ID from USER ROW, IF EXISTS ## - # get user row - removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None) - addressed_user_ids: Final = ( - removed_user_ids if removed_team_members else frozenset((data.user_id,) if data.user_id is not None else ()) - ) - key_val: Final[Mapping[str, object]] = ( - {"user_id": {"in": sorted(addressed_user_ids)}} if addressed_user_ids else {"user_email": data.user_email} - ) - member_tx: Final[_MemberDeleteTx] = tx - existing_user_rows: Final = await member_tx.litellm_usertable.find_many(where=key_val) - - # A user row can outlive its roster entry, and until the team is off user.teams the user - # still sees it and still fails key creation against it, so removal has to clear it too - stale_user_rows: Final = tuple(user for user in existing_user_rows if data.team_id in user.teams) - - # Also clean up any existing team membership rows for this user and team. An email can - # match several user rows, so with no roster entry to name the member, only the rows - # actually carrying the team are the ones this request is allowed to touch - cleanup_user_rows: Final = existing_user_rows if removed_team_members else stale_user_rows - user_ids_to_delete: Final = addressed_user_ids.union(user.user_id for user in cleanup_user_rows if user.user_id) - - if not removed_team_members and not stale_user_rows: - raise HTTPException(status_code=400, detail={"error": "User not found in team"}) - - ## DELETE KEYS CREATED BY USER FOR THIS TEAM - # Fetch keys before deletion so their audit records can be persisted alongside the delete. - # An empty user_ids_to_delete still resolves cleanly: prisma's "in": [] matches no rows. - keys_to_delete: Final = await member_tx.litellm_verificationtoken.find_many( + await tx.litellm_verificationtoken.delete_many( where={ "user_id": {"in": sorted(user_ids_to_delete)}, "team_id": data.team_id, } ) - jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_tokens( - hashed_tokens=tuple(key.token for key in keys_to_delete), - prisma_client=prisma_client, - ) - if removed_team_members: - await _team_tx_db(tx).update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_new_team_members)}, - ) + return TeamMemberRemoval( + team=existing_team_row, + before=tuple(fresh_members), + after=tuple(new_team_members), + keys=tuple(keys_to_delete), + jwt_mapping_cache_keys=jwt_mapping_cache_keys, + user_ids=frozenset(user_ids_to_delete), + ) - for existing_user in stale_user_rows: - await tx.litellm_usertable.update( - where={"user_id": existing_user.user_id}, - data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}}, - ) - for _uid in sorted(user_ids_to_delete): - await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid}) +async def finish_team_member_removal(removal: TeamMemberRemoval, user_api_key_dict: UserAPIKeyAuth) -> None: + from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - if user_ids_to_delete: - if keys_to_delete: - from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, - ) - - await _persist_deleted_verification_tokens( - keys=keys_to_delete, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=None, - tx=tx, - ) - - await tx.litellm_verificationtoken.delete_many( - where={ - "user_id": {"in": sorted(user_ids_to_delete)}, - "team_id": data.team_id, - } - ) - - if keys_to_delete: + if removal.keys: KeyManagementEventHooks.create_key_deleted_audit_logs( - keys_being_deleted=keys_to_delete, + keys_being_deleted=removal.keys, user_api_key_dict=user_api_key_dict, litellm_changed_by=None, ) await delete_cache_team_object( - team_id=data.team_id, - team_alias=existing_team_row.team_alias, + team_id=removal.team.team_id, + team_alias=removal.team.team_alias, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) await delete_cache_key_objects( - hashed_tokens=tuple(key.token for key in keys_to_delete), + hashed_tokens=tuple(key.token for key in removal.keys), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) - await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache) - await evict_and_broadcast(cache_keys=tuple(sorted(user_ids_to_delete)), user_api_key_cache=user_api_key_cache) - for user_id in sorted(user_ids_to_delete): + await evict_and_broadcast(cache_keys=removal.jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache) + await evict_and_broadcast(cache_keys=tuple(sorted(removal.user_ids)), user_api_key_cache=user_api_key_cache) + for user_id in sorted(removal.user_ids): await invalidate_team_member_spend_state( user_id=user_id, - team_id=data.team_id, + team_id=removal.team.team_id, user_api_key_cache=user_api_key_cache, ) - _emit_team_members_metric(existing_team_row) - - return existing_team_row, tuple(fresh_members), tuple(new_team_members) + _emit_team_members_metric(removal.team) @router.post( diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index b12a689429d..569f2ebd278 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -21,6 +21,7 @@ from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key from litellm.proxy.utils import PrismaClient from litellm.repositories.object_permission_repository import ObjectPermissionRepository +from litellm.repositories.prisma_protocols import DatabaseClient, TableActions from litellm.repositories.table_repositories import MCPServerRepository if TYPE_CHECKING: @@ -201,7 +202,9 @@ async def invalidate_cached_object_permissions( async def _set_object_permission( data_json: dict, - prisma_client: PrismaClient | None, + prisma_client: DatabaseClient | None, + *, + table: "TableActions[prisma_models.LiteLLM_ObjectPermissionTable] | None" = None, ): """ Creates the LiteLLM_ObjectPermissionTable record for the key/team. @@ -230,7 +233,8 @@ async def _set_object_permission( if "mcp_tool_permissions" in clean_data: clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"]) - created_permission: Final = await ObjectPermissionRepository(prisma_client).table.create(data=clean_data) + permission_table: Final = ObjectPermissionRepository(prisma_client).table if table is None else table + created_permission: Final = await permission_table.create(data=clean_data) data_json["object_permission_id"] = created_permission.object_permission_id data_json.pop("object_permission") @@ -259,7 +263,7 @@ def _mcp_server_identifier_matches(server: object, identifier: str) -> bool: async def _get_db_mcp_servers_by_identifiers( identifiers: AbstractSet[str], - prisma_client: PrismaClient | None, + prisma_client: DatabaseClient | None, ) -> "Sequence[prisma_models.LiteLLM_MCPServerTable]": if prisma_client is None or not identifiers: return [] @@ -278,7 +282,7 @@ async def _get_db_mcp_servers_by_identifiers( async def _resolve_mcp_server_identifiers_to_ids( identifiers: AbstractSet[str], - prisma_client: PrismaClient | None, + prisma_client: DatabaseClient | None, ) -> dict[str, set[str]]: """ Resolve MCP permission entries written as server_id, alias, or server_name @@ -335,7 +339,7 @@ def _mcp_tool_permission_entries(raw: object) -> Mapping[str, frozenset[str]]: async def reject_ambiguous_mcp_tool_permission_keys( new_mcp_tool_permissions: object, existing_mcp_tool_permissions: object, - prisma_client: PrismaClient | None, + prisma_client: DatabaseClient | None, ) -> None: """ A name or alias shared by several MCP servers cannot key ``mcp_tool_permissions``: diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py new file mode 100644 index 00000000000..40eea1d7618 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py @@ -0,0 +1,1246 @@ +import asyncio +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.management_endpoints.scim.agent_provisioning import ( + SCIMProvisioningFailure, + apply_user_patch, + group_members_after_patch, +) +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIM_AGENT_USER_SCHEMA +from litellm.types.proxy.management_endpoints.scim_v2 import SCIMGroup, SCIMMember, SCIMPatchOp, SCIMUser + +PARENT: Final = "11111111-1111-4111-8111-111111111111" + +SUBJECT: Final = "22222222-2222-4222-8222-222222222222" + + + + +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", [True, False, None]) +async def test_source_token_lookup_uses_current_writer_state(enabled: bool | None) -> None: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.management_endpoints.scim.agent_provisioning import source_for_auth + + source: Final = SimpleNamespace(source_id="source", enabled=enabled) if enabled is not None else None + client: Final = MagicMock() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=source) + auth: Final = UserAPIKeyAuth(token="source-token-hash") + if enabled is False: + with pytest.raises(HTTPException) as failure: + await source_for_auth(auth, client) + assert failure.value.status_code == 403 + else: + assert await source_for_auth(auth, client) is source + client.writer_db.litellm_scimsource.find_unique.assert_awaited_once_with(where={"key_hash": "source-token-hash"}) + client.db.litellm_scimsource.find_unique.assert_not_called() + assert await source_for_auth(None, client) is None + assert await source_for_auth(UserAPIKeyAuth(), client) is None + + +@pytest.mark.asyncio +async def test_directory_documents_use_authoritative_activity_and_membership() -> None: + from litellm.proxy.management_endpoints.scim.agent_provisioning import ( + group_document, + remove_group_member, + user_document, + ) + + row: Final = SimpleNamespace( + id="row", active=False, document={"schemas": [], "userName": "subject", "active": True} + ) + assert user_document(row).active is False + assert user_document(row).id == "row" + group: Final = SimpleNamespace( + id="group", + member_ids=["keep", "remove"], + document={"schemas": [], "displayName": "Directory", "members": [{"value": "stale"}]}, + ) + assert [member.value for member in group_document(group).members] == ["keep", "remove"] + client: Final = MagicMock() + client.litellm_scimresource.update = AsyncMock() + await remove_group_member(client, group, "remove") + client.litellm_scimresource.update.assert_awaited_once_with(where={"id": "group"}, data={"member_ids": ["keep"]}) + assert group.member_ids == ["keep", "remove"] + + +def agent_user() -> SCIMUser: + return SCIMUser.model_validate( + { + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User", SCIM_AGENT_USER_SCHEMA], + "id": "stable-scim-id", + "externalId": SUBJECT, + "userName": "agent@example.com", + "active": True, + SCIM_AGENT_USER_SCHEMA: {"identityParentId": PARENT}, + } + ) + + +@pytest.mark.parametrize("wire_value, expected", [("False", False), ("True", True), (False, False), (True, True)]) +def test_entra_boolean_patch_preserves_identity_and_returns_json_boolean(wire_value: object, expected: bool) -> None: + original: Final = agent_user() + result: Final = apply_user_patch( + original, SCIMPatchOp(Operations=[{"op": "Replace", "path": "active", "value": wire_value}]) + ) + assert isinstance(result, SCIMUser) + assert result.active is expected + assert result.id == original.id + assert result.externalId == SUBJECT + assert result.agent_user == original.agent_user + assert result.model_dump(by_alias=True)["active"] is expected + assert original.active is True + + +@pytest.mark.parametrize( + "path, value", + [ + ("externalId", PARENT), + (SCIM_AGENT_USER_SCHEMA + ":identityParentId", SUBJECT), + (None, {SCIM_AGENT_USER_SCHEMA: {"identityParentId": SUBJECT}}), + ("active", "garbage"), + ], +) +def test_identity_rebinding_and_invalid_active_patch_are_rejected(path: str | None, value: object) -> None: + result: Final = apply_user_patch( + agent_user(), SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}]) + ) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +def test_rename_does_not_reenable_a_disabled_subject() -> None: + original: Final = agent_user().model_copy(update={"active": False}) + result: Final = apply_user_patch( + original, SCIMPatchOp(Operations=[{"op": "replace", "value": {"displayName": "Renamed"}}]) + ) + assert isinstance(result, SCIMUser) + assert result.displayName == "Renamed" + assert result.active is False + assert result.id == original.id + + + + +def test_group_removal_preserves_other_members_and_repeat_removal_is_idempotent() -> None: + group: Final = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + displayName="Mixed", + members=[SCIMMember(value="agent"), SCIMMember(value="human")], + ) + patch: Final = SCIMPatchOp(Operations=[{"op": "Remove", "path": 'members[value eq "agent"]'}]) + result: Final = group_members_after_patch(group, patch) + assert isinstance(result, SCIMGroup) + assert result.members == [SCIMMember(value="human")] + assert group_members_after_patch(result, patch) == result + assert group.members == [SCIMMember(value="agent"), SCIMMember(value="human")] + + +def test_group_replace_removes_omitted_members_and_empty_replace_removes_all() -> None: + group: Final = SCIMGroup( + schemas=[], displayName="Mixed", members=[SCIMMember(value="agent"), SCIMMember(value="human")] + ) + result: Final = group_members_after_patch( + group, SCIMPatchOp(Operations=[{"op": "replace", "path": "members", "value": [{"value": "human"}]}]) + ) + assert isinstance(result, SCIMGroup) + assert result.members == [SCIMMember(value="human")] + empty: Final = group_members_after_patch( + result, SCIMPatchOp(Operations=[{"op": "replace", "path": "members", "value": []}]) + ) + assert isinstance(empty, SCIMGroup) + assert empty.members == [] + + +@pytest.mark.parametrize( + "path,value", + [ + (SCIM_AGENT_USER_SCHEMA + ":identityParentId", PARENT), + (None, {SCIM_AGENT_USER_SCHEMA: {"identityParentId": PARENT}}), + ("schemas", [SCIM_AGENT_USER_SCHEMA]), + ], +) +def test_human_patch_cannot_smuggle_an_agent_identity(path: str | None, value: object) -> None: + from litellm.proxy.management_endpoints.scim.agent_provisioning import patch_changes_identity + + patch: Final = SCIMPatchOp(Operations=[{"op": "add", "path": path, "value": value}]) + assert patch_changes_identity(patch) + + +@pytest.mark.parametrize("identity_marker", [True, False]) +def test_deep_patch_checks_identity_without_exhausting_the_call_stack(identity_marker: bool) -> None: + from functools import reduce + + from litellm.proxy.management_endpoints.scim.agent_provisioning import patch_changes_identity + + leaf: Final = {"identityParentId": PARENT} if identity_marker else {"displayName": "Renamed"} + nested: Final = reduce(lambda value, _: {"nested": [value]}, range(1200), leaf) + patch: Final = SCIMPatchOp(Operations=[{"op": "replace", "value": nested}]) + assert patch_changes_identity(patch) is identity_marker + + +def test_patch_error_does_not_apply_later_operations() -> None: + patch: Final = SCIMPatchOp( + Operations=[ + {"op": "replace", "path": "externalId", "value": "foreign-subject"}, + {"op": "replace", "path": "displayName", "value": "renamed"}, + ] + ) + result: Final = apply_user_patch(agent_user(), patch) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +@pytest.mark.parametrize( + "operation,expected", + [ + ({"op": "add", "path": "members", "value": [{"value": "human"}, {"value": "second"}]}, ["human", "second"]), + ({"op": "remove", "path": "members", "value": [{"value": "human"}]}, []), + ({"op": "remove", "path": "members"}, []), + ({"op": "replace", "path": "displayName", "value": "Renamed"}, ["human"]), + ], +) +def test_group_patch_add_remove_and_rename_keep_membership_consistent( + operation: dict[str, object], expected: list[str] +) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch(group, SCIMPatchOp.model_validate({"Operations": [operation]})) + assert isinstance(result, SCIMGroup) + assert [member.value for member in result.members or []] == expected + assert result.displayName == ("Renamed" if operation["path"] == "displayName" else "Original") + + +@pytest.mark.parametrize( + "operation", + [ + {"op": "replace", "path": "externalId", "value": "foreign"}, + {"op": "add", "path": "members", "value": [{"display": "missing-id"}]}, + ], +) +def test_group_patch_failure_cannot_apply_subsequent_membership_changes(operation: dict[str, object]) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch( + group, SCIMPatchOp.model_validate({"Operations": [operation, {"op": "remove", "path": "members"}]}) + ) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + assert group.members == [SCIMMember(value="human")] + + +@pytest.mark.parametrize( + "operation", + [ + {"op": "remove"}, + {"op": "replace", "value": "invalid"}, + {"op": "remove", "path": "userName"}, + ], +) +def test_invalid_agent_profile_patch_is_rejected(operation: dict[str, object]) -> None: + result: Final = apply_user_patch(agent_user(), SCIMPatchOp.model_validate({"Operations": [operation]})) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +@pytest.mark.parametrize( + "path,value", + [("name.givenName", "New"), ("name.familyName", "Family"), ('emails[type eq "work"].value', "new@example.com")], +) +def test_native_profile_accepts_standard_entra_subattribute_updates(path: str, value: str) -> None: + user: Final = SCIMUser.model_validate( + { + **agent_user().model_dump(by_alias=True), + "name": {"givenName": "Old", "familyName": "Original"}, + "emails": [{"type": "work", "value": "old@example.com"}, {"type": "home", "value": "home@example.com"}], + } + ) + result: Final = apply_user_patch(user, SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}])) + assert isinstance(result, SCIMUser) + assert result.externalId == user.externalId + assert result.agent_user == user.agent_user + if path.startswith("name."): + assert getattr(result.name, path.split(".")[1]) == value + assert result.emails == user.emails + else: + assert result.emails[0].value == value + assert result.emails[1].value == "home@example.com" + assert result.name == user.name + + +def provisioning_fixture(): + from datetime import datetime, timezone + from unittest.mock import AsyncMock, MagicMock + + from prisma import Prisma + from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource + + from litellm.proxy.management_endpoints.scim.agent_provisioning import AgentProvisioningService + from litellm.proxy.utils import PrismaClient + + now: Final = datetime.now(timezone.utc) + source: Final = LiteLLM_SCIMSource( + source_id="source", + display_name="Source", + tenant_id=PARENT, + key_hash="hash", + enabled=True, + group_mappings="[]", + created_at=now, + updated_at=now, + ) + row: Final = LiteLLM_SCIMResource( + id="stable-scim-id", + source_id="source", + kind="Users", + external_id=SUBJECT, + user_name="agent@example.com", + display_name="Agent", + document=agent_user().model_dump_json(by_alias=True), + active=True, + deleted=False, + local_id="registered-agent", + member_ids=[], + created_at=now, + updated_at=now, + ) + client: Final = MagicMock(spec=PrismaClient) + tx: Final = MagicMock(spec=Prisma) + client.tx.return_value.__aenter__.return_value = tx + client.writer_db = tx + tx.execute_raw = AsyncMock(return_value=1) + tx.litellm_scimsource.find_unique = AsyncMock(return_value=source) + tx.litellm_scimresource.find_unique = AsyncMock(return_value=row) + tx.litellm_scimresource.update_many = AsyncMock(return_value=1) + tx.litellm_scimresource.update = AsyncMock(return_value=row) + tx.litellm_scimresource.find_many = AsyncMock(return_value=[]) + tx.litellm_scimresource.count = AsyncMock(return_value=0) + tx.litellm_agentstable.find_unique = AsyncMock(return_value={"agent_id": row.local_id}) + return AgentProvisioningService(client, source), tx, row + + +@pytest.mark.asyncio +async def test_profile_put_cannot_reenable_native_user_when_active_is_omitted() -> None: + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"active": False}) + profile: Final = agent_user().model_dump(by_alias=True) + incoming: Final = SCIMUser.model_validate({key: value for key, value in profile.items() if key != "active"}) + result: Final = await service.update_user(row.id, incoming) + assert result.active is False + assert result.id == row.id + assert tx.litellm_scimresource.update_many.call_args.kwargs["data"]["active"] is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("username", [None, ""]) +async def test_native_profile_update_rejects_blank_username_without_writing(username: str | None) -> None: + service, tx, row = provisioning_fixture() + request: Final = SCIMUser.model_validate({**agent_user().model_dump(by_alias=True), "userName": username}) + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, request) + assert failure.value.status_code == 400 + assert failure.value.detail == "userName is required" + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_deleted_registration_is_not_recreated_by_directory_replay() -> None: + from fastapi import HTTPException + + service, tx, row = provisioning_fixture() + tx.litellm_agentstable.find_unique.return_value = None + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, agent_user()) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_disabled_source_is_rechecked_after_waiting_for_its_lock() -> None: + from fastapi import HTTPException + + service, tx, row = provisioning_fixture() + tx.litellm_scimsource.find_unique.return_value = service.source.model_copy(update={"enabled": False}) + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, agent_user()) + assert failure.value.status_code == 403 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_concurrent_native_profile_update_returns_conflict() -> None: + from fastapi import HTTPException + + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.update_many.return_value = 0 + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, agent_user()) + assert failure.value.status_code == 409 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["Users", "Groups"]) +async def test_scoped_delete_propagates_human_and_team_deprovisioning( + kind: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + document: Final = SCIMUser(schemas=[], userName="human@example.com").model_dump(mode="json") + row: Final = native.model_copy(update={"kind": kind, "document": document}) + tx.litellm_scimresource.find_unique.return_value = row + delete_user: Final = AsyncMock() + delete_group: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "delete_user", delete_user) + monkeypatch.setattr(scim_v2, "delete_group", delete_group) + await service.delete(kind, row.id) + if kind == "Users": + delete_user.assert_awaited_once_with(user_id=row.local_id) + delete_group.assert_not_awaited() + else: + delete_group.assert_awaited_once_with(group_id=row.local_id) + delete_user.assert_not_awaited() + tx.litellm_scimresource.update.assert_awaited_once_with( + where={"id": row.id}, + data={"active": False, "deleted": True, "member_ids": []}, + ) + + +@pytest.mark.asyncio +async def test_failed_human_deletion_remains_retryable(monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import AsyncMock + + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + row: Final = native.model_copy(update={"document": SCIMUser(schemas=[], userName="human@example.com").model_dump()}) + tx.litellm_scimresource.find_unique.return_value = row + monkeypatch.setattr(scim_v2, "delete_user", AsyncMock(side_effect=HTTPException(503, "unavailable"))) + with pytest.raises(HTTPException) as failure: + await service.delete("Users", row.id) + assert failure.value.status_code == 503 + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_native_delete_is_idempotent_and_does_not_delete_a_human(monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, row = provisioning_fixture() + delete_user: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "delete_user", delete_user) + await service.delete("Users", row.id) + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"deleted": True, "active": False}) + await service.delete("Users", row.id) + delete_user.assert_not_awaited() + tx.litellm_scimresource.update.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_native_delete_removes_group_links_before_tombstoning_the_subject() -> None: + from unittest.mock import call + + service, tx, row = provisioning_fixture() + group: Final = row.model_copy(update={"id": "mixed-group", "kind": "Groups", "member_ids": [row.id, "human"]}) + tx.litellm_scimresource.find_many.return_value = [group] + await service.delete("Users", row.id) + tx.litellm_scimresource.update.assert_has_awaits( + [ + call(where={"id": group.id}, data={"member_ids": ["human"]}), + call(where={"id": row.id}, data={"active": False, "deleted": True, "member_ids": []}), + ] + ) + tx.litellm_scimresource.find_many.assert_awaited_once_with( + where={"source_id": "source", "kind": "Groups", "member_ids": {"has": row.id}} + ) + + +@pytest.mark.asyncio +async def test_new_native_user_is_created_disabled_with_exact_subject_parent_and_source() -> None: + service, tx, _ = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = None + tx.litellm_scimresource.create = AsyncMock() + tx.litellm_agentidentity.find_unique = AsyncMock(return_value=None) + tx.litellm_agentstable.create = AsyncMock() + tx.litellm_verifiedsubject.create = AsyncMock() + result: Final = await service.create_user(agent_user()) + resource: Final = tx.litellm_scimresource.create.call_args.kwargs["data"] + agent: Final = tx.litellm_agentstable.create.call_args.kwargs["data"] + subject: Final = tx.litellm_verifiedsubject.create.call_args.kwargs["data"] + assert result.id == resource["id"] == subject["scim_resource_id"] + assert agent["agent_id"] == resource["local_id"] == subject["agent_id"] + assert agent["enabled"] is False + assert agent["execution_mode"] == "autonomous" + assert agent["identity"]["create"]["provisioning_source_id"] == service.source.source_id + assert agent["identity"]["create"]["client_id"] == subject["parent_client_id"] == PARENT + assert subject["oid"] == resource["external_id"] == SUBJECT + assert subject["kind"] == "agent_user" + assert "user_id" not in subject + + +@pytest.mark.asyncio +async def test_native_create_replay_uses_existing_registration() -> None: + service, tx, row = provisioning_fixture() + tx.litellm_agentstable.create = AsyncMock() + result: Final = await service.create_user(agent_user()) + assert result.id == row.id + tx.litellm_agentstable.create.assert_not_awaited() + assert tx.litellm_scimresource.update_many.call_args.kwargs["where"]["id"] == row.id + + +@pytest.mark.asyncio +async def test_native_create_cannot_adopt_a_preexisting_parent() -> None: + service, tx, _ = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = None + tx.litellm_agentidentity.find_unique = AsyncMock(return_value={"agent_id": "other"}) + tx.litellm_scimresource.create = AsyncMock() + with pytest.raises(HTTPException) as failure: + await service.create_user(agent_user()) + assert failure.value.status_code == 409 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("changed", [{"externalId": "invalid"}, {"externalId": None}, {"userName": None}]) +async def test_native_create_rejects_missing_or_malformed_identity(changed: dict[str, object]) -> None: + service, tx, _ = provisioning_fixture() + tx.litellm_scimresource.create = AsyncMock() + with pytest.raises(HTTPException) as failure: + await service.create_user(agent_user().model_copy(update=changed)) + assert failure.value.status_code == 400 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", ["missing", "foreign", "deleted", "wrong-kind"]) +async def test_resource_reads_are_source_and_kind_scoped(state: str) -> None: + service, tx, row = provisioning_fixture() + changes: Final = { + "foreign": {"source_id": "foreign"}, + "deleted": {"deleted": True}, + "wrong-kind": {"kind": "Groups"}, + } + tx.litellm_scimresource.find_unique.return_value = ( + None if state == "missing" else row.model_copy(update=changes[state]) + ) + with pytest.raises(HTTPException) as failure: + await service.get("Users", row.id) + assert failure.value.status_code == 404 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "attribute,field", + [("userName", "user_name"), ("externalId", "external_id"), ("displayName", "display_name"), ("id", "id")], +) +async def test_scim_filter_keeps_source_boundary_and_total_count(attribute: str, field: str) -> None: + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.find_many.return_value = [row] + tx.litellm_scimresource.count = AsyncMock(return_value=3) + result: Final = await service.list("Users", 2, 500, f'{attribute} eq "value"') + query: Final = tx.litellm_scimresource.find_many.call_args.kwargs + assert query["where"] == {"source_id": "source", "kind": "Users", "deleted": False, field: "value"} + assert query["skip"] == 1 and query["take"] == 100 + assert result.totalResults == 3 + assert result.itemsPerPage == 1 + assert result.Resources[0].id == row.id + + +@pytest.mark.asyncio +@pytest.mark.parametrize("filter_value", ['active eq "true"', 'userName sw "a"']) +async def test_unsupported_filter_is_rejected_before_database_query(filter_value: str) -> None: + service, tx, _ = provisioning_fixture() + with pytest.raises(HTTPException) as failure: + await service.list("Users", 1, 10, filter_value) + assert failure.value.status_code == 400 + tx.litellm_scimresource.find_many.assert_not_awaited() + + +def group_rows(native): + from litellm.types.proxy.management_endpoints.scim_v2 import SCIMGroup, SCIMUser + + human: Final = native.model_copy( + update={ + "id": "human-scim", + "local_id": "local-human", + "external_id": "external-human", + "document": SCIMUser(schemas=[], userName="human@example.com").model_dump(), + } + ) + document: Final = SCIMGroup(schemas=[], externalId="directory-group", displayName="Mixed") + group: Final = native.model_copy( + update={ + "id": "group-scim", + "kind": "Groups", + "local_id": None, + "external_id": "directory-group", + "document": document.model_dump(), + "member_ids": [native.id, human.id], + } + ) + return human, group, document + + +@pytest.mark.asyncio +async def test_group_validation_reads_all_member_batches() -> None: + from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE + + service, tx, _ = provisioning_fixture() + members: Final = tuple(f"member-{index}" for index in range(IN_LIST_CHUNK_SIZE + 1)) + tx.litellm_scimresource.count.side_effect = [IN_LIST_CHUNK_SIZE, 1] + await service._validate_members(members) + assert tx.litellm_scimresource.count.await_count == 2 + tx.litellm_scimresource.count.side_effect = [IN_LIST_CHUNK_SIZE, 0] + with pytest.raises(HTTPException) as failure: + await service._validate_members(members) + assert failure.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_group_sync_includes_humans_from_later_batches(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + members: Final = [f"member-{index}" for index in range(IN_LIST_CHUNK_SIZE)] + [human.id] + native_rows: Final = [native.model_copy(update={"id": member}) for member in members[:-1]] + tx.litellm_scimresource.find_many.side_effect = [native_rows, [human]] + tx.litellm_teamtable.find_unique = AsyncMock(return_value=None) + create_team: Final = AsyncMock(return_value=scim_v2.ProvisionedGroupWrite(team_id="local-team", created=None, removals=(), additions=())) + monkeypatch.setattr(scim_v2, "write_provisioned_group", create_team) + await service._sync_human_members(tx, group.model_copy(update={"member_ids": members}), None) + create_team.assert_awaited_once() + assert create_team.call_args.args[2].members == [SCIMMember(value=human.local_id)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "update"]) +async def test_group_member_validation_uses_the_source_transaction(operation: str) -> None: + service, tx, native = provisioning_fixture() + _, group, document = group_rows(native) + single_member: Final = group.model_copy(update={"member_ids": [native.id]}) + request: Final = document.model_copy(update={"members": [SCIMMember(value=native.id)]}) + tx.litellm_scimresource.find_unique.return_value = None if operation == "create" else single_member + tx.litellm_scimresource.create = AsyncMock(return_value=single_member) + tx.litellm_scimresource.count.return_value = 1 + tx.litellm_scimresource.find_many.return_value = [native] + writer: Final = MagicMock() + writer.litellm_scimresource.count = AsyncMock(return_value=1) + writer.litellm_scimresource.find_many = AsyncMock(return_value=[native]) + service.client.writer_db = writer + result: Final = ( + await service.create_group(request) if operation == "create" else await service.update_group(group.id, request) + ) + assert result.id == group.id + assert result.members == [SCIMMember(value=native.id)] + tx.litellm_scimresource.count.assert_awaited_once() + writer.litellm_scimresource.count.assert_not_awaited() + service.client.tx.assert_called_once() + assert service.client.tx.call_args.kwargs["timeout"].total_seconds() == 30 + if operation == "create": + tx.litellm_scimresource.create.assert_awaited_once() + else: + assert tx.litellm_scimresource.update_many.call_args.kwargs["where"] == { + "id": group.id, + "updated_at": group.updated_at, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("include_human", [True, False]) +async def test_group_create_keeps_agents_out_of_human_teams( + include_human: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + rows: Final = [native, human] if include_human else [native] + group: Final = group.model_copy(update={"member_ids": [row.id for row in rows]}) + tx.litellm_scimresource.find_unique.return_value = None + tx.litellm_scimresource.find_many.return_value = rows + tx.litellm_scimresource.count.return_value = len(rows) + tx.litellm_scimresource.create = AsyncMock(return_value=group) + tx.litellm_teamtable.find_unique = AsyncMock(return_value=None) + create_team: Final = AsyncMock(return_value=scim_v2.ProvisionedGroupWrite(team_id="local-team", created=None, removals=(), additions=())) + monkeypatch.setattr(scim_v2, "write_provisioned_group", create_team) + result: Final = await service.create_group( + document.model_copy(update={"members": [SCIMMember(value=row.id) for row in rows]}) + ) + assert result.id == group.id + assert {member.value for member in result.members} == set(group.member_ids) + if include_human: + create_team.assert_awaited_once() + assert create_team.call_args.args[2].members == [SCIMMember(value=human.local_id)] + assert tx.litellm_scimresource.update.call_args.kwargs == { + "where": {"id": group.id}, + "data": {"local_id": "local-team"}, + } + else: + create_team.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_group_create_rejects_members_missing_from_its_source(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + _, _, document = group_rows(native) + tx.litellm_scimresource.find_unique.return_value = None + tx.litellm_scimresource.create = AsyncMock() + create_team: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "write_provisioned_group", create_team) + with pytest.raises(HTTPException) as failure: + await service.create_group(document.model_copy(update={"members": [SCIMMember(value="foreign-member")]})) + assert failure.value.status_code == 400 + tx.litellm_scimresource.create.assert_not_awaited() + create_team.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_group_replay_reconciles_removed_humans_but_preserves_agent_membership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from types import SimpleNamespace + + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + _, group, document = group_rows(native) + old: Final = group.model_copy(update={"local_id": "local-team"}) + updated: Final = old.model_copy(update={"member_ids": [native.id]}) + tx.litellm_scimresource.find_unique.side_effect = [old, old, updated] + tx.litellm_scimresource.find_many.return_value = [native] + tx.litellm_scimresource.count.return_value = 1 + tx.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(team_id="local-team")) + update_team: Final = AsyncMock(return_value=scim_v2.ProvisionedGroupWrite(team_id="local-team", created=None, removals=(), additions=())) + monkeypatch.setattr(scim_v2, "write_provisioned_group", update_team) + result: Final = await service.create_group(document.model_copy(update={"members": [SCIMMember(value=native.id)]})) + assert result.id == old.id + assert result.members == [SCIMMember(value=native.id)] + assert tx.litellm_scimresource.update_many.call_args.kwargs["data"]["member_ids"] == [native.id] + update_team.assert_awaited_once() + assert update_team.call_args.args[2].id == "local-team" + assert update_team.call_args.args[2].members == [] + + +@pytest.mark.asyncio +async def test_group_external_id_is_immutable() -> None: + service, tx, native = provisioning_fixture() + _, group, document = group_rows(native) + tx.litellm_scimresource.find_unique.return_value = group + with pytest.raises(HTTPException) as failure: + await service.update_group(group.id, document.model_copy(update={"externalId": "other-directory-group"})) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ["create", "update"]) +async def test_native_replay_accepts_equivalent_uppercase_object_id(method: str) -> None: + service, tx, row = provisioning_fixture() + subject: Final = "abcdef01-abcd-4abc-8abc-abcdef012345" + user: Final = agent_user().model_copy(update={"externalId": subject.upper()}) + stored: Final = row.model_copy( + update={"external_id": subject, "document": user.model_dump(by_alias=True, mode="json")} + ) + tx.litellm_scimresource.find_unique.return_value = stored + result: Final = await service.create_user(user) if method == "create" else await service.update_user(row.id, user) + assert result.id == row.id + assert result.externalId == subject.upper() + tx.litellm_scimresource.update_many.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("external_id", [None, "invalid", PARENT]) +async def test_native_put_cannot_remove_or_change_subject(external_id: str | None) -> None: + service, tx, row = provisioning_fixture() + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, agent_user().model_copy(update={"externalId": external_id})) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_native_patch_rejects_identity_mutation_without_writing() -> None: + service, tx, row = provisioning_fixture() + with pytest.raises(HTTPException) as failure: + await service.update_user( + row.id, SCIMPatchOp(Operations=[{"op": "replace", "path": "externalId", "value": PARENT}]) + ) + assert failure.value.status_code == 400 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", [True, False]) +async def test_source_lookup_uses_writer_and_rejects_disabled_source(enabled: bool) -> None: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.management_endpoints.scim.agent_provisioning import source_for_auth + + service, _, _ = provisioning_fixture() + service.client.writer_db.litellm_scimsource.find_unique = AsyncMock( + return_value=service.source.model_copy(update={"enabled": enabled}) + ) + if enabled: + result: Final = await source_for_auth(UserAPIKeyAuth(token="source-token-hash"), service.client) + assert result.source_id == service.source.source_id + else: + with pytest.raises(HTTPException) as failure: + await source_for_auth(UserAPIKeyAuth(token="source-token-hash"), service.client) + assert failure.value.status_code == 403 + service.client.writer_db.litellm_scimsource.find_unique.assert_awaited_once_with( + where={"key_hash": "source-token-hash"} + ) + assert await source_for_auth(None, service.client) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["put", "patch"]) +async def test_source_owned_human_cannot_be_reclassified_as_agent(operation: str) -> None: + service, tx, row = provisioning_fixture() + human: Final = SCIMUser(schemas=[], externalId=SUBJECT, userName="human@example.com") + tx.litellm_scimresource.find_unique.return_value = row.model_copy( + update={"document": human.model_dump(mode="json")} + ) + incoming: Final = ( + agent_user() + if operation == "put" + else SCIMPatchOp( + Operations=[{"op": "add", "path": SCIM_AGENT_USER_SCHEMA + ":identityParentId", "value": PARENT}] + ) + ) + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, incoming) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_deleted_native_subject_cannot_be_recreated_by_post_replay() -> None: + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"deleted": True}) + with pytest.raises(HTTPException) as failure: + await service.create_user(agent_user()) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", ["foreign", "missing", "wrong-kind"]) +async def test_scoped_delete_rejects_unknown_or_foreign_resources(state: str) -> None: + service, tx, row = provisioning_fixture() + changed: Final = {"source_id": "other"} if state == "foreign" else {"kind": "Groups"} + tx.litellm_scimresource.find_unique.return_value = None if state == "missing" else row.model_copy(update=changed) + with pytest.raises(HTTPException) as failure: + await service.delete("Users", row.id) + assert failure.value.status_code == 404 + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create-replay", "put"]) +async def test_scoped_human_routes_preserve_reserved_identity_and_use_human_provisioner( + monkeypatch: pytest.MonkeyPatch, operation: str +) -> None: + from litellm.proxy._types import LiteLLM_UserTable + + service, tx, row = provisioning_fixture() + human: Final = SCIMUser(schemas=[], userName="human@example.com", externalId=SUBJECT) + reserved: Final = row.model_copy(update={"document": human.model_dump(mode="json"), "local_id": "human"}) + tx.litellm_scimresource.find_unique.return_value = reserved + tx.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="human")) + tx.litellm_usertable.find_many = AsyncMock(return_value=[]) + tx.litellm_agentstable.create = AsyncMock() + tx.litellm_usertable.update = AsyncMock( + return_value=LiteLLM_UserTable(user_id="human", user_email=human.userName, metadata={"scim_active": True}) + ) + result: Final = ( + await service.create_user(human) if operation == "create-replay" else await service.update_user(row.id, human) + ) + assert result.id == row.id + assert result.externalId == SUBJECT + assert result.agent_user is None + assert tx.litellm_usertable.update.await_args.kwargs["where"] == {"user_id": "human"} + assert tx.litellm_usertable.update.await_args.kwargs["data"]["user_email"] == human.userName + service.client.tx.assert_called_once() + tx.litellm_agentstable.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", ["missing-external-id", "deleted"]) +async def test_group_create_cannot_recreate_deleted_group_or_omit_directory_id(state: str) -> None: + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"kind": "Groups", "deleted": True}) + with pytest.raises(HTTPException) as failure: + await service.create_group( + SCIMGroup(schemas=[], displayName="Group", externalId=None if state == "missing-external-id" else SUBJECT) + ) + assert failure.value.status_code == (400 if state == "missing-external-id" else 409) + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["Users", "Groups"]) +async def test_get_returns_the_owned_directory_document(kind: str) -> None: + service, tx, row = provisioning_fixture() + if kind == "Groups": + _, group, _ = group_rows(row) + tx.litellm_scimresource.find_unique.return_value = group + result: Final = await service.get(kind, tx.litellm_scimresource.find_unique.return_value.id) + assert result.id == tx.litellm_scimresource.find_unique.return_value.id + assert isinstance(result, SCIMUser if kind == "Users" else SCIMGroup) + + +@pytest.mark.asyncio +async def test_native_insert_unique_collision_is_a_conflict() -> None: + from prisma.errors import UniqueViolationError + + service, tx, _ = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = None + tx.litellm_agentidentity.find_unique = AsyncMock(return_value=None) + tx.litellm_agentstable.create = AsyncMock( + side_effect=UniqueViolationError({"user_facing_error": {"error_code": "P2002", "message": "Duplicate"}}) + ) + tx.litellm_scimresource.create = AsyncMock() + tx.litellm_verifiedsubject.create = AsyncMock() + with pytest.raises(HTTPException) as failure: + await service.create_user(agent_user()) + assert failure.value.status_code == 409 + tx.litellm_verifiedsubject.create.assert_not_awaited() + assert isinstance(failure.value.__cause__, UniqueViolationError) + assert service.client.tx.return_value.__aexit__.call_args.args[1] is failure.value + + +@pytest.mark.asyncio +async def test_missing_native_local_id_rejects_profile_update() -> None: + service, tx, row = provisioning_fixture() + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"local_id": None}) + with pytest.raises(HTTPException) as failure: + await service.update_user(row.id, agent_user()) + assert failure.value.status_code == 409 + tx.litellm_scimresource.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("case", ["invalid-patch", "concurrent"]) +async def test_group_patch_failure_does_not_sync_members(case: str, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + _, group, _ = group_rows(native) + tx.litellm_scimresource.find_unique.return_value = group.model_copy(update={"member_ids": []}) + tx.litellm_scimresource.update_many.return_value = 0 + change: Final = SCIMPatchOp( + Operations=[ + {"op": "replace", "path": "externalId" if case == "invalid-patch" else "displayName", "value": "changed"} + ] + ) + sync: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "update_group", sync) + with pytest.raises(HTTPException) as failure: + await service.update_group(group.id, change) + assert failure.value.status_code == (400 if case == "invalid-patch" else 409) + sync.assert_not_awaited() + if case == "invalid-patch": + tx.litellm_scimresource.update_many.assert_not_awaited() + + + + +@pytest.mark.parametrize( + "value, expected_name, expected_members", + [ + ({"displayName": "Renamed"}, "Renamed", ["human"]), + ({"displayName": "Renamed", "members": [{"value": "agent"}]}, "Renamed", ["agent"]), + ({"members": [{"value": "agent"}]}, "Original", ["agent"]), + ], +) +def test_pathless_group_replace_applies_display_name_and_members( + value: dict[str, object], expected_name: str, expected_members: list[str] +) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch(group, SCIMPatchOp(Operations=[{"op": "replace", "value": value}])) + assert isinstance(result, SCIMGroup), result + assert result.displayName == expected_name + assert [member.value for member in result.members or []] == expected_members + + +@pytest.mark.parametrize( + "value", + [ + {"externalId": "foreign"}, + {"displayName": ""}, + {"displayName": "Renamed", "members": [{"display": "no-id"}]}, + "x", + ], +) +def test_pathless_group_replace_rejects_unsupported_or_invalid_attributes(value: object) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch(group, SCIMPatchOp(Operations=[{"op": "replace", "value": value}])) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "update"]) +@pytest.mark.parametrize("identity", ["missing", "ambiguous", "incomplete"]) +async def test_source_group_rejects_unusable_human_identity_without_placeholder_creation( + operation: str, identity: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + member = human.model_copy(update={"local_id": None}) if identity == "incomplete" else human + tx.litellm_scimresource.find_unique.return_value = None if operation == "create" else group + tx.litellm_scimresource.create = AsyncMock(return_value=group) + tx.litellm_scimresource.find_many.return_value = [member] + tx.litellm_scimresource.count.return_value = 1 + tx.query_raw = AsyncMock(return_value=[] if identity == "missing" else [{"user_id": human.local_id}]) + tx.litellm_usertable.find_many = AsyncMock(return_value=[ + LiteLLM_UserTable(user_id=human.local_id), LiteLLM_UserTable(user_id="other") + ]) + placeholder = AsyncMock() + effects = AsyncMock() + monkeypatch.setattr(scim_v2, "_create_user_if_not_exists", placeholder) + monkeypatch.setattr(scim_v2, "finish_provisioned_group", effects) + request = document.model_copy(update={"members": [SCIMMember(value=human.id)]}) + with pytest.raises(HTTPException) as failure: + await service.create_group(request) if operation == "create" else await service.update_group(group.id, request) + assert failure.value.status_code == (400 if identity == "ambiguous" else 409) + assert service.client.tx.return_value.__aexit__.await_args.args[1] is failure.value + placeholder.assert_not_awaited() + effects.assert_not_awaited() + tx.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure_stage", ["directory", "team", "keys"]) +@pytest.mark.parametrize("failure_type", [RuntimeError, asyncio.CancelledError]) +async def test_group_replacement_rolls_back_local_and_directory_writes_then_retries( + failure_stage: str, failure_type, monkeypatch: pytest.MonkeyPatch +) -> None: + import copy + from contextlib import asynccontextmanager + + from litellm.proxy._types import LiteLLM_TeamTable, LiteLLM_UserTable, Member + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + group = group.model_copy(update={"local_id": group.id, "member_ids": [human.id]}) + state = {"directory": group, "members": [Member(user_id=human.local_id, role="user")], + "alias": group.display_name, "teams": [group.id], "membership": True, "keys_deleted": False} + initial = copy.deepcopy(state) + fail = True + transaction_open = False + + @asynccontextmanager + async def transaction(**kwargs): + nonlocal transaction_open + snapshot = copy.deepcopy(state) + transaction_open = True + try: + yield tx + except BaseException: + state.clear() + state.update(snapshot) + raise + finally: + transaction_open = False + + async def directory_update(*, where, data): + state["directory"] = state["directory"].model_copy(update={"member_ids": data["member_ids"], "display_name": data["display_name"], "document": data["document"].data}) + if fail and failure_stage == "directory": + raise failure_type("directory interrupted") + return 1 + + async def team_read(**kwargs): + return LiteLLM_TeamTable(team_id=group.id, team_alias=state["alias"], members_with_roles=state["members"], metadata={}) + + async def team_write(*, where, data): + state["alias"] = data.get("team_alias", state["alias"]) + if "members_with_roles" in data: + import json + state["members"] = [Member.model_validate(value) for value in json.loads(data["members_with_roles"])] + if fail and failure_stage == "team": + raise failure_type("team interrupted") + return await team_read() + + async def query(sql, *args): + return [{"members_with_roles": [member.model_dump() for member in state["members"]]}] if "SELECT members_with_roles" in sql else [] + + async def user_update(*, where, data): + state["teams"] = data["teams"]["set"] + + async def membership_delete(**kwargs): + state["membership"] = False + return 1 + + async def key_delete(**kwargs): + state["keys_deleted"] = True + if fail and failure_stage == "keys": + raise failure_type("key deletion interrupted") + return 0 + + async def effects(result): + assert not transaction_open + assert state["directory"].member_ids == [] + assert state["members"] == state["teams"] == [] + assert state["membership"] is False and state["keys_deleted"] is True + assert len(result.removals) == 1 + assert result.removals[0].user_ids == frozenset([human.local_id]) + + service.client.tx.side_effect = transaction + tx.query_raw = AsyncMock(side_effect=query) + tx.litellm_scimresource.find_unique.side_effect = lambda **kwargs: state["directory"] + tx.litellm_scimresource.update_many.side_effect = directory_update + tx.litellm_scimresource.find_many.return_value = [] + tx.litellm_teamtable.find_unique = AsyncMock(side_effect=team_read) + tx.litellm_teamtable.update = AsyncMock(side_effect=team_write) + tx.litellm_usertable.find_many = AsyncMock(side_effect=lambda **kwargs: [LiteLLM_UserTable(user_id=human.local_id, teams=state["teams"])]) + tx.litellm_usertable.update = AsyncMock(side_effect=user_update) + tx.litellm_teammembership.delete_many = AsyncMock(side_effect=membership_delete) + tx.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + tx.litellm_verificationtoken.delete_many = AsyncMock(side_effect=key_delete) + finish = AsyncMock(side_effect=effects) + monkeypatch.setattr(scim_v2, "finish_provisioned_group", finish) + monkeypatch.setattr(scim_v2, "provisioning_group_admin_role", AsyncMock(return_value=None)) + change = document.model_copy(update={"displayName": "Renamed", "members": []}) + with pytest.raises(failure_type): + await service.update_group(group.id, change) + assert state == initial + finish.assert_not_awaited() + fail = False + result = await service.update_group(group.id, change) + assert result.id == group.id and result.members == [] and result.displayName == "Renamed" + finish.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure_type", [RuntimeError, asyncio.CancelledError]) +async def test_group_creation_link_failure_rolls_back_and_replays_one_identity( + failure_type, monkeypatch: pytest.MonkeyPatch +) -> None: + from contextlib import asynccontextmanager + + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + state = {"directory": None, "team": None} + fail = True + transaction_open = False + + @asynccontextmanager + async def transaction(**kwargs): + nonlocal transaction_open + before = dict(state) + transaction_open = True + try: + yield tx + except BaseException: + state.clear() + state.update(before) + raise + finally: + transaction_open = False + + async def create_resource(*, data): + row = group.model_copy(update={"id": data["id"], "member_ids": [human.id], "document": data["document"].data}) + state["directory"] = row + return row + + async def link_resource(*, where, data): + state["directory"] = state["directory"].model_copy(update=data) + if fail: + raise failure_type("link interrupted") + return state["directory"] + + async def write_team(writer, client, incoming, admin_group): + assert writer is tx and transaction_open + state["team"] = incoming.id + return scim_v2.ProvisionedGroupWrite(team_id=incoming.id, created=None, removals=(), additions=()) + + async def finish_group(result): + assert not transaction_open + assert state["directory"].local_id == state["team"] == result.team_id + + service.client.tx.side_effect = transaction + tx.litellm_scimresource.find_unique.side_effect = lambda **kwargs: state["directory"] + tx.litellm_scimresource.create = AsyncMock(side_effect=create_resource) + tx.litellm_scimresource.update.side_effect = link_resource + tx.litellm_scimresource.find_many.return_value = [human] + tx.litellm_scimresource.count.return_value = 1 + effects = AsyncMock(side_effect=finish_group) + monkeypatch.setattr(scim_v2, "write_provisioned_group", AsyncMock(side_effect=write_team)) + monkeypatch.setattr(scim_v2, "finish_provisioned_group", effects) + monkeypatch.setattr(scim_v2, "provisioning_group_admin_role", AsyncMock(return_value=None)) + request = document.model_copy(update={"members": [SCIMMember(value=human.id)]}) + with pytest.raises(failure_type): + await service.create_group(request) + assert state == {"directory": None, "team": None} + effects.assert_not_awaited() + fail = False + result = await service.create_group(request) + replay = await service.create_group(request) + assert result.id == replay.id == state["team"] + assert effects.await_count == 2 + assert tx.litellm_scimresource.create.await_count == 2 + + +@pytest.mark.asyncio +async def test_group_addition_reuses_member_writes_and_defers_existing_audit_and_cache_effects( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable + from litellm.proxy.management_endpoints import team_endpoints + from litellm.proxy.management_endpoints.scim import scim_v2 + + service, tx, native = provisioning_fixture() + human, group, document = group_rows(native) + local_user = LiteLLM_UserTable(user_id=human.local_id, teams=[]) + local_team = LiteLLM_TeamTable(team_id=group.id, members_with_roles=[], metadata={"team_member_budget_id": "inherited"}) + tx.query_raw = AsyncMock(side_effect=lambda sql, *args: [{"members_with_roles": []}] if "SELECT members_with_roles" in sql else [{"user_id": human.local_id}]) + tx.litellm_budgettable.find_unique = AsyncMock(return_value=SimpleNamespace(budget_id="inherited")) + tx.litellm_usertable.find_many = AsyncMock(return_value=[local_user]) + tx.litellm_usertable.upsert = AsyncMock(return_value=local_user) + tx.litellm_usertable.update_many = AsyncMock(return_value=1) + tx.litellm_teamtable.find_unique = AsyncMock(return_value=local_team) + tx.litellm_teamtable.update = AsyncMock(return_value=local_team) + tx.litellm_teammembership.upsert = AsyncMock(return_value=LiteLLM_TeamMembership( + team_id=group.id, user_id=human.local_id, budget_id="inherited" + )) + invalidate = AsyncMock() + audit = MagicMock() + monkeypatch.setattr(team_endpoints, "evict_and_broadcast", invalidate) + monkeypatch.setattr(team_endpoints, "_schedule_team_member_add_audit_logs", audit) + request = document.model_copy(update={"id": group.id, "members": [SCIMMember(value=human.local_id)]}) + result = await scim_v2.write_provisioned_group(tx, service.client, request, None) + assert result.created is None and result.removals == () + assert len(result.additions) == 1 + assert result.additions[0].users[0].user_id == human.local_id + membership = tx.litellm_teammembership.upsert.await_args.kwargs["data"] + assert membership["create"]["budget_id"] == "inherited" + assert membership["create"]["user_id"] == human.local_id + assert tx.litellm_usertable.update_many.await_count == 1 + service.client.tx.assert_not_called() + invalidate.assert_not_awaited() + audit.assert_not_called() + await scim_v2.finish_provisioned_group(result) + assert any(call.kwargs.get("cache_keys") == (human.local_id,) for call in invalidate.await_args_list) + audit.assert_called_once() + assert audit.call_args.kwargs["existing_user_ids"] == frozenset([human.local_id]) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py b/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py new file mode 100644 index 00000000000..886cd1a177d --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py @@ -0,0 +1,559 @@ +import asyncio +import json +from datetime import datetime, timezone +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException +from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource + +from litellm.proxy import proxy_server +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.management_endpoints.scim import scim_v2 +from litellm.proxy.management_endpoints.scim.human_provisioning import SourceHumanProvisioner, human_email +from litellm.proxy.utils import PrismaClient +from litellm.types.proxy.management_endpoints.scim_v2 import SCIMPatchOp, SCIMUser, SCIMUserEmail + +TENANT: Final = "11111111-1111-4111-8111-111111111111" +SUBJECT: Final = "22222222-2222-4222-8222-222222222222" + + +def human_fixture(): + now: Final = datetime.now(timezone.utc) + source: Final = LiteLLM_SCIMSource( + source_id="source", + display_name="Directory", + tenant_id=TENANT, + key_hash="hash", + enabled=True, + group_mappings="[]", + created_at=now, + updated_at=now, + ) + user: Final = SCIMUser(schemas=[], userName="human@example.com", externalId=SUBJECT, displayName="Human") + row: Final = LiteLLM_SCIMResource( + id="stable-scim-id", + source_id=source.source_id, + kind="Users", + external_id=SUBJECT, + user_name=user.userName, + display_name="Human", + document=user.model_dump_json(), + active=True, + deleted=False, + local_id="local-human", + human_email=user.userName, + member_ids=[], + created_at=now, + updated_at=now, + ) + client: Final = MagicMock(spec=PrismaClient) + tx: Final = client.tx.return_value.__aenter__.return_value + tx.litellm_scimresource.find_unique = AsyncMock(return_value=None) + tx.litellm_scimresource.create = AsyncMock(return_value=row) + tx.litellm_scimresource.update = AsyncMock(return_value=row) + tx.litellm_usertable.find_many = AsyncMock(return_value=[]) + tx.litellm_usertable.find_unique = AsyncMock(return_value=None) + tx.litellm_usertable.create = AsyncMock() + + async def save_local(*, where, data): + before = tx.litellm_usertable.find_unique.return_value.model_dump() + values = {**before, **data} + if isinstance(values.get("metadata"), str): + values["metadata"] = json.loads(values["metadata"]) + return LiteLLM_UserTable.model_validate(values) + + tx.litellm_usertable.update = AsyncMock(side_effect=save_local) + tx.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + tx.litellm_verificationtoken.update = AsyncMock() + + client.db = MagicMock() + client.db.litellm_usertable.count = AsyncMock(return_value=0) + tx.litellm_usertable.count = client.db.litellm_usertable.count + return SourceHumanProvisioner(client, source), tx, row, user + + +@pytest.mark.parametrize( + "emails,expected", + [ + (None, "human@example.com"), + ([{"value": "FIRST@example.com"}], "first@example.com"), + ([{"value": "FIRST@example.com"}, {"value": "PRIMARY@example.com", "primary": True}], "primary@example.com"), + ], +) +def test_human_ownership_email_uses_primary_and_normalizes_case(emails: object, expected: str) -> None: + user: Final = SCIMUser.model_validate({"schemas": [], "userName": "human@example.com", "emails": emails}) + assert human_email(user) == expected + + +def test_missing_human_ownership_identifier_is_rejected() -> None: + with pytest.raises(HTTPException) as failure: + human_email(SCIMUser(schemas=[])) + assert failure.value.status_code == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing", ["externalId", "userName"]) +async def test_direct_reservation_requires_complete_directory_identity(missing: str) -> None: + service, tx, _, user = human_fixture() + with pytest.raises(HTTPException) as failure: + await service.reserve(user.model_copy(update={missing: None})) + assert failure.value.status_code == 400 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_human_create_does_not_hide_storage_failure() -> None: + service, tx, _, user = human_fixture() + tx.litellm_scimresource.find_unique.side_effect = ConnectionError("unavailable") + with pytest.raises(ConnectionError, match="unavailable"): + await service.create(user) + tx.litellm_usertable.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_incomplete_human_record_cannot_be_updated() -> None: + service, tx, row, user = human_fixture() + with pytest.raises(HTTPException) as failure: + await service.update(row.model_copy(update={"local_id": None}), user) + assert failure.value.status_code == 409 + tx.litellm_usertable.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path,value", + [ + ("emails", [{"value": "NEW@example.com"}]), + ('emails[type eq "work"].value', "NEW@example.com"), + (None, {"emails": [{"value": "NEW@example.com"}]}), + ], +) +async def test_scoped_email_patch_requires_put_before_mutation(path: str | None, value: object) -> None: + service, tx, row, _ = human_fixture() + change: Final = SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}]) + with pytest.raises(HTTPException, match="PUT") as failure: + await service.update(row, change) + assert failure.value.status_code == 400 + tx.litellm_usertable.find_unique.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "emails", [None, [{"value": "FIRST@example.com"}, {"value": "PRIMARY@example.com", "primary": True}]] +) +async def test_human_put_claims_the_same_email_it_writes(emails: object, monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable(user_id=row.local_id, user_email=row.human_email, teams=[]) + change: Final = SCIMUser.model_validate({**user.model_dump(), "emails": emails}) + expected: Final = human_email(change) + + result: Final = await service.update(row, change) + assert result.id == row.id + assert result.emails and result.emails[0].value == expected + assert tx.litellm_usertable.update.await_args.kwargs["data"]["user_email"] == expected + assert tx.litellm_scimresource.update.await_args_list[0].kwargs["data"] == {"human_email": expected} + + +@pytest.mark.asyncio +async def test_reservation_replay_preserves_identity_before_creating_a_local_user() -> None: + service, tx, row, user = human_fixture() + tx.litellm_scimresource.find_unique.return_value = row + assert await service.reserve(user) == row + tx.litellm_scimresource.create.assert_not_awaited() + tx.litellm_usertable.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_reservation_claims_email_subject_and_local_identity() -> None: + service, tx, _, user = human_fixture() + await service.reserve(user) + data: Final = tx.litellm_scimresource.create.call_args.kwargs["data"] + assert data["local_id"] == user.userName + assert data["human_email"] == user.userName + assert data["human_subject_key"] == f"{TENANT}:{SUBJECT}" + assert data["id"] == data["document"].data["id"] + tx.litellm_usertable.create.assert_awaited_once_with( + data={ + "user_id": user.userName, + "user_email": user.userName, + "user_role": "internal_user_viewer", + "teams": [], + } + ) + + +@pytest.mark.asyncio +async def test_ambiguous_local_human_match_is_rejected_before_reserving() -> None: + service, tx, _, user = human_fixture() + tx.litellm_usertable.find_many.return_value = [SimpleNamespace(user_id="one"), SimpleNamespace(user_id="two")] + with pytest.raises(HTTPException) as failure: + await service.reserve(user) + assert failure.value.status_code == 409 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing", ["externalId", "userName"]) +async def test_incomplete_identity_cannot_be_reserved(missing: str) -> None: + service, tx, _, user = human_fixture() + with pytest.raises(HTTPException) as failure: + await service.create(user.model_copy(update={missing: None})) + assert failure.value.status_code == 400 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_missing_local_human_is_not_recreated_or_email_adopted(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + tx.litellm_scimresource.find_unique.return_value = row + create: Final = AsyncMock(return_value=user.model_copy(update={"id": "unrelated-admin"})) + update: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "create_user", create) + monkeypatch.setattr(scim_v2, "update_user", update) + with pytest.raises(HTTPException) as failure: + await service.create(user) + assert failure.value.status_code == 409 + create.assert_not_awaited() + update.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_disappearing_local_human_aborts_update_before_key_or_document_writes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + service, tx, row, user = human_fixture() + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable( + user_id=row.local_id, user_email=row.human_email, metadata={"scim_active": True} + ) + tx.litellm_usertable.update = AsyncMock(return_value=None) + evict_user: Final = AsyncMock() + evict_key: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "evict_and_broadcast", evict_user) + monkeypatch.setattr(scim_v2, "_delete_cache_key_object", evict_key) + + with pytest.raises(HTTPException) as failure: + await service.update(row, user.model_copy(update={"active": False})) + + assert failure.value.status_code == 409 + assert "automatic recreation is not permitted" in failure.value.detail + assert service.client.tx.return_value.__aexit__.await_args.args[1] is failure.value + tx.litellm_scimresource.update.assert_awaited_once() + assert tx.litellm_scimresource.update.await_args.kwargs["data"] == {"human_email": row.human_email} + tx.litellm_usertable.create.assert_not_awaited() + tx.litellm_verificationtoken.find_many.assert_not_awaited() + tx.litellm_verificationtoken.update.assert_not_awaited() + evict_user.assert_not_awaited() + evict_key.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_deleted_human_cannot_be_recreated_by_replay(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + tx.litellm_scimresource.find_unique.return_value = row.model_copy(update={"deleted": True}) + create: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "create_user", create) + with pytest.raises(HTTPException) as failure: + await service.create(user) + assert failure.value.status_code == 409 + create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path,value", [("groups", []), ("externalId", "foreign"), (None, {"externalId": "foreign"})]) +async def test_scoped_human_patch_cannot_modify_directory_owned_correspondence(path: str | None, value: object) -> None: + service, tx, row, _ = human_fixture() + with pytest.raises(HTTPException) as failure: + await service.update(row, SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}])) + assert failure.value.status_code == 400 + tx.litellm_usertable.find_unique.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_local_insert_failure_aborts_resource_reservation() -> None: + service, tx, _, user = human_fixture() + tx.litellm_usertable.create.side_effect = RuntimeError("interrupted") + with pytest.raises(RuntimeError, match="interrupted"): + await service.reserve(user) + tx.litellm_scimresource.create.assert_not_awaited() + assert service.client.tx.return_value.__aexit__.call_args.args[0] is RuntimeError + + +@pytest.mark.asyncio +@pytest.mark.parametrize("phase", ["reservation", "email-update"]) +async def test_ownership_collision_is_a_conflict_before_legacy_user_mutation( + phase: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from prisma.errors import UniqueViolationError + + service, tx, row, user = human_fixture() + collision: Final = UniqueViolationError( + {"user_facing_error": {"error_code": "P2002", "message": "Unique identity"}} + ) + create: Final = AsyncMock() + update: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "create_user", create) + monkeypatch.setattr(scim_v2, "update_user", update) + if phase == "reservation": + tx.litellm_scimresource.create.side_effect = collision + else: + tx.litellm_scimresource.find_unique.return_value = row + tx.litellm_scimresource.update.side_effect = collision + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable(user_id=row.local_id, user_email=row.human_email, teams=[]) + with pytest.raises(HTTPException) as failure: + await service.create(user) + assert failure.value.status_code == 409 + create.assert_not_awaited() + update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_human_patch_preserves_scim_id_and_updates_activity(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._types import LiteLLM_UserTable + + service, tx, row, user = human_fixture() + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable( + user_id=row.local_id, user_email="human@example.com" + ) + operations: Final = SCIMPatchOp(Operations=[{"op": "replace", "path": "active", "value": False}]) + result: Final = await service.update(row, operations) + assert result.id == row.id and result.externalId == row.external_id + assert result.active is False + assert json.loads(tx.litellm_usertable.update.await_args.kwargs["data"]["metadata"])["scim_active"] is False + assert tx.litellm_scimresource.update.call_args.kwargs["data"]["active"] is False + + +@pytest.mark.asyncio +async def test_unavailable_ownership_database_is_not_reported_as_a_conflict() -> None: + service, tx, _, user = human_fixture() + tx.litellm_scimresource.find_unique.side_effect = RuntimeError("database unavailable") + with pytest.raises(RuntimeError, match="database unavailable"): + await service.create(user) + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_source_cannot_adopt_an_existing_local_or_sso_user() -> None: + service, tx, _, user = human_fixture() + tx.litellm_usertable.find_many.return_value = [SimpleNamespace(user_id="existing-admin")] + with pytest.raises(HTTPException) as failure: + await service.reserve(user) + assert failure.value.status_code == 409 + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["put", "display", "username", "object"]) +async def test_source_username_is_independent_of_local_display_name( + operation: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._types import LiteLLM_UserTable + + service, tx, row, user = human_fixture() + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable( + user_id=row.local_id, user_email="human@example.com" + ) + changes: Final = { + "put": user.model_copy(update={"userName": "renamed@example.com", "emails": [SCIMUserEmail(value="human@example.com", primary=True)]}), + "display": SCIMPatchOp(Operations=[{"op": "replace", "path": "displayName", "value": "Display Name"}]), + "username": SCIMPatchOp(Operations=[{"op": "replace", "path": "userName", "value": "renamed@example.com"}]), + "object": SCIMPatchOp(Operations=[{"op": "replace", "value": {"userName": "renamed@example.com"}}]), + } + result: Final = await service.update(row, changes[operation]) + expected: Final = "renamed@example.com" if operation in ("put", "username", "object") else "human@example.com" + assert result.userName == expected + assert result.displayName == "human@example.com" + if operation == "display": + assert tx.litellm_usertable.update.await_args.kwargs["data"]["user_alias"] == "Display Name" + assert result.id == row.id + assert tx.litellm_scimresource.update.call_args.kwargs["data"]["user_name"] == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation,value", [("remove", None), ("replace", ""), ("replace", 123)]) +async def test_invalid_username_is_rejected_before_local_mutation(operation: str, value: object) -> None: + service, tx, row, _ = human_fixture() + change: Final = SCIMPatchOp(Operations=[{"op": operation, "path": "userName", "value": value}]) + with pytest.raises(HTTPException, match="userName is required") as failure: + await service.update(row, change) + assert failure.value.status_code == 400 + tx.litellm_usertable.find_unique.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_human_guid_case_replays_the_same_reserved_identity() -> None: + service, tx, row, user = human_fixture() + guid: Final = "abcdefab-abcd-4abc-8abc-abcdefabcdef" + tx.litellm_scimresource.find_unique.side_effect = lambda **query: ( + row if query["where"]["source_id_kind_external_id"]["external_id"] == guid else None + ) + result: Final = await service.reserve(user.model_copy(update={"externalId": guid.upper()})) + assert result.id == row.id + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_email_update_cannot_claim_an_unrelated_local_human(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable(user_id=row.local_id, user_email=row.human_email, teams=[]) + tx.litellm_usertable.find_many.return_value = [SimpleNamespace(user_id="unrelated-admin")] + update: Final = AsyncMock() + monkeypatch.setattr(scim_v2, "update_user", update) + with pytest.raises(HTTPException) as failure: + await service.update( + row, SCIMUser.model_validate({**user.model_dump(), "emails": [{"value": "ADMIN@example.com"}]}) + ) + assert failure.value.status_code == 409 + update.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_source_human_creation_preserves_license_limit(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, _, user = human_fixture() + service.client.db.litellm_usertable.count.side_effect = [3, 0] + license_check: Final = MagicMock() + license_check.is_over_limit.return_value = True + monkeypatch.setattr(proxy_server, "_license_check", license_check) + with pytest.raises(HTTPException) as failure: + await service.reserve(user) + assert failure.value.status_code == 403 + license_check.is_over_limit.assert_called_once_with(total_users=3) + tx.litellm_usertable.create.assert_not_awaited() + tx.litellm_scimresource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_replayed_human_does_not_consume_another_license_seat(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + tx.litellm_scimresource.find_unique.return_value = row + license_check: Final = MagicMock() + license_check.is_over_limit.return_value = True + monkeypatch.setattr(proxy_server, "_license_check", license_check) + assert await service.reserve(user) == row + license_check.is_over_limit.assert_not_called() + tx.litellm_usertable.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_existing_native_agent_cannot_be_reclassified_as_human() -> None: + from litellm.types.proxy.management_endpoints.scim_v2 import SCIM_AGENT_USER_SCHEMA + + service, tx, row, user = human_fixture() + native: Final = SCIMUser.model_validate( + { + **user.model_dump(), + SCIM_AGENT_USER_SCHEMA: {"identityParentId": TENANT}, + } + ) + tx.litellm_scimresource.find_unique.return_value = row.model_copy( + update={"document": native.model_dump(by_alias=True, mode="json")} + ) + with pytest.raises(HTTPException) as failure: + await service.create(user) + assert failure.value.status_code == 409 + tx.litellm_usertable.find_unique.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_put_without_username_is_rejected_before_anything_is_written(monkeypatch: pytest.MonkeyPatch) -> None: + service, tx, row, user = human_fixture() + legacy_put: Final = AsyncMock(return_value=user) + monkeypatch.setattr(scim_v2, "update_user", legacy_put) + with pytest.raises(HTTPException, match="userName is required") as failure: + await service.update( + row, user.model_copy(update={"userName": None, "emails": [SCIMUserEmail(value="human@example.com")]}) + ) + assert failure.value.status_code == 400 + legacy_put.assert_not_awaited() + tx.litellm_usertable.find_unique.assert_not_awaited() + tx.litellm_scimresource.update.assert_not_called() + + +@pytest.mark.parametrize("failure_type", [RuntimeError, asyncio.CancelledError]) +@pytest.mark.parametrize("failure_stage", ["local", "keys", "document"]) +@pytest.mark.asyncio +async def test_human_update_rolls_back_and_retry_preserves_identity( + failure_type, failure_stage: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from contextlib import asynccontextmanager + + service, tx, row, user = human_fixture() + state = {"human_email": row.human_email, "user_email": row.human_email, "blocked": False} + failure_enabled = True + transaction_open = False + + @asynccontextmanager + async def transaction(): + nonlocal transaction_open + before = dict(state) + transaction_open = True + try: + yield tx + except BaseException: + state.clear() + state.update(before) + raise + finally: + transaction_open = False + + async def save_resource(*, where, data): + state.update(data) + if "document" in data and failure_stage == "document" and failure_enabled: + raise failure_type("directory write interrupted") + return row + + async def save_user(*, where, data): + state["user_email"] = data["user_email"] + if failure_stage == "local" and failure_enabled: + raise failure_type("local write interrupted") + return LiteLLM_UserTable( + user_id=row.local_id, user_email=data["user_email"], metadata=json.loads(data["metadata"]), teams=[] + ) + + async def save_key(*, where, data): + state["blocked"] = data["blocked"] + if failure_stage == "keys" and failure_enabled: + raise failure_type("key write interrupted") + + async def invalidate_key(**kwargs): + assert not transaction_open + assert state["blocked"] is True + + service.client.tx.side_effect = transaction + tx.litellm_scimresource.update = AsyncMock(side_effect=save_resource) + tx.litellm_usertable.find_unique.return_value = LiteLLM_UserTable( + user_id=row.local_id, user_email=row.human_email, teams=[], metadata={"scim_active": True} + ) + tx.litellm_usertable.update = AsyncMock(side_effect=save_user) + tx.litellm_verificationtoken.find_many.return_value = [SimpleNamespace(token="human-key", metadata={})] + tx.litellm_verificationtoken.update = AsyncMock(side_effect=save_key) + evict_user = AsyncMock() + evict_key = AsyncMock(side_effect=invalidate_key) + monkeypatch.setattr(scim_v2, "evict_and_broadcast", evict_user) + monkeypatch.setattr(scim_v2, "_delete_cache_key_object", evict_key) + change: Final = user.model_copy( + update={"emails": [SCIMUserEmail(value="replacement@example.com", primary=True)], "active": False} + ) + with pytest.raises(failure_type): + await service.update(row, change) + assert state == {"human_email": row.human_email, "user_email": row.human_email, "blocked": False} + evict_user.assert_not_awaited() + evict_key.assert_not_awaited() + + failure_enabled = False + result: Final = await service.update(row, change) + assert result.id == row.id and result.externalId == row.external_id + assert result.active is False + assert state["human_email"] == state["user_email"] == "replacement@example.com" + assert state["blocked"] is True + assert state["active"] is False + evict_user.assert_awaited_once() + evict_key.assert_awaited_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5f0488b0c76..bef82f88293 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -5749,9 +5749,10 @@ class _InjectedMemberDeleteFailure(Exception): pass +@pytest.mark.parametrize("failure_type", [_InjectedMemberDeleteFailure, asyncio.CancelledError]) @pytest.mark.asyncio async def test_team_member_delete_is_atomic_across_its_four_writes( - mock_db_client, mock_admin_auth + mock_db_client, mock_admin_auth, failure_type, monkeypatch ): """ /team/member_delete's four cleanups (team roster, user.teams, team @@ -5767,6 +5768,11 @@ async def test_team_member_delete_is_atomic_across_its_four_writes( from litellm.proxy._types import TeamMemberDeleteRequest from litellm.proxy.management_endpoints.team_endpoints import team_member_delete + from litellm.proxy.management_endpoints import team_endpoints + + invalidate_team: Final = AsyncMock() + monkeypatch.setattr(team_endpoints, "delete_cache_team_object", invalidate_team) + test_team_id = "team-del-atomic-123" test_user_id = "user-atomic@example.com" @@ -5793,7 +5799,7 @@ async def test_team_member_delete_is_atomic_across_its_four_writes( return_value=[mock_user_row] ) mock_db_client.db.litellm_usertable.update = AsyncMock( - side_effect=_InjectedMemberDeleteFailure("boom between writes 1 and 2") + side_effect=failure_type("boom between writes 1 and 2") ) mock_db_client.db.litellm_teammembership = MagicMock() @@ -5805,7 +5811,7 @@ async def test_team_member_delete_is_atomic_across_its_four_writes( _wire_member_delete_tx(mock_db_client) - with pytest.raises(_InjectedMemberDeleteFailure): + with pytest.raises(failure_type): await team_member_delete( data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id), user_api_key_dict=mock_admin_auth, @@ -5815,7 +5821,8 @@ async def test_team_member_delete_is_atomic_across_its_four_writes( mock_db_client.db.litellm_teamtable.update.assert_awaited_once() mock_db_client.tx.assert_called_once() aexit_args = mock_db_client.tx.return_value.__aexit__.await_args.args - assert aexit_args[0] is _InjectedMemberDeleteFailure + assert aexit_args[0] is failure_type + invalidate_team.assert_not_awaited() # Writes queued behind the failure inside that same transaction never ran. mock_db_client.db.litellm_teammembership.delete_many.assert_not_awaited() @@ -16884,3 +16891,115 @@ def test_list_team_v2_answers_503_no_db_connection_when_the_callers_user_read_hi assert response.status_code == 503, response.text assert response.json() == _DB_OUTAGE_503_BODY + + +@pytest.mark.asyncio +async def test_member_removal_writes_nothing_when_team_disappeared_under_lock() -> None: + from litellm.proxy._types import TeamMemberDeleteRequest + from litellm.proxy.management_endpoints.team_endpoints import delete_team_member_in_transaction + + tx: Final = MagicMock() + tx.query_raw = AsyncMock(return_value=[]) + tx.litellm_teamtable.update = AsyncMock() + tx.litellm_usertable.update = AsyncMock() + tx.litellm_verificationtoken.delete_many = AsyncMock() + with pytest.raises(HTTPException) as error: + await delete_team_member_in_transaction( + tx=tx, + data=TeamMemberDeleteRequest(team_id="removed-team", user_id="member"), + existing_team_row=LiteLLM_TeamTable(team_id="removed-team", members_with_roles=[]), + prisma_client=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + assert error.value.status_code == 400 + assert error.value.detail == {"error": "Team id=removed-team does not exist in db"} + tx.litellm_teamtable.update.assert_not_awaited() + tx.litellm_usertable.update.assert_not_awaited() + tx.litellm_verificationtoken.delete_many.assert_not_awaited() + + +@pytest.mark.parametrize( + ("member_duration", "explicit_fields", "expected_duration"), + [(None, frozenset(), "30d"), (None, frozenset({"team_member_budget_duration"}), None), ("7d", frozenset(), "7d")], +) +@pytest.mark.asyncio +async def test_team_member_budget_uses_supplied_table_and_preserves_defaults( + member_duration: str | None, explicit_fields: frozenset[str], expected_duration: str | None +) -> None: + from types import SimpleNamespace + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler + + table: Final = MagicMock() + table.create = AsyncMock(return_value=SimpleNamespace(budget_id="member-budget")) + result: Final = await TeamMemberBudgetHandler.create_team_member_budget_table( + data=NewTeamRequest(team_alias="Directory team", budget_duration="30d"), + new_team_data_json={"team_member_budget": 12.0}, + user_api_key_dict=UserAPIKeyAuth(user_id="provisioner"), + team_member_budget=12.0, + team_member_budget_duration=member_duration, + explicitly_set_fields=explicit_fields, + table=table, + ) + table.create.assert_awaited_once() + written: Final = table.create.await_args.kwargs["data"] + assert written["max_budget"] == 12.0 + assert written.get("budget_duration") == expected_duration + assert written["created_by"] == written["updated_by"] == "provisioner" + assert result["metadata"]["team_member_budget_id"] == "member-budget" + assert "team_member_budget" not in result + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_membership", [False, True]) +async def test_team_creation_uses_supplied_transaction_for_inherited_records( + mock_db_client, mock_admin_auth, monkeypatch: pytest.MonkeyPatch, fail_membership: bool +) -> None: + import litellm + + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewTeamRequest, ProxyException + from litellm.proxy.management_endpoints import team_endpoints + + tx = MagicMock() + row = LiteLLM_TeamTable(team_id="scim-team", team_alias="Directory team", members_with_roles=[]) + tx.litellm_teamtable.count = AsyncMock(return_value=0) + tx.litellm_teamtable.find_unique = AsyncMock(return_value=None) + tx.litellm_teamtable.create = AsyncMock(return_value=row) + tx.litellm_teamtable.update = AsyncMock( + return_value=row, side_effect=RuntimeError("membership failed") if fail_membership else None + ) + tx.litellm_modeltable.create = AsyncMock(return_value=SimpleNamespace(id=42)) + tx.litellm_objectpermissiontable.create = AsyncMock(return_value=SimpleNamespace(object_permission_id="permissions")) + tx.litellm_budgettable.create = AsyncMock(return_value=SimpleNamespace(budget_id="member-budget")) + + async def query(sql, *args): + return [{"members_with_roles": []}] if "SELECT members_with_roles" in sql else [] + + tx.query_raw = AsyncMock(side_effect=query) + invalidate = AsyncMock() + monkeypatch.setattr(team_endpoints, "invalidate_access_group_caches", invalidate) + monkeypatch.setattr(litellm, "default_team_params", None) + mock_db_client.jsonify_team_object = lambda db_data: db_data + request = NewTeamRequest( + team_id="scim-team", team_alias="Directory team", model_aliases={"friendly": "model"}, + object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["store"]), + team_member_budget=5, team_member_budget_duration="30d", models=["model"], + ) + if fail_membership: + with pytest.raises(ProxyException, match="membership failed"): + await team_endpoints.create_team(request, mock_admin_auth, transaction=tx) + else: + result = await team_endpoints.create_team(request, mock_admin_auth, transaction=tx) + assert result.team.team_id == "scim-team" + assert result.snapshot.models == ["model"] + + assert tx.litellm_modeltable.create.await_args.args[0]["model_aliases"] == '{"friendly": "model"}' + assert tx.litellm_objectpermissiontable.create.await_args.kwargs["data"]["vector_stores"] == ["store"] + assert tx.litellm_budgettable.create.await_args.kwargs["data"]["max_budget"] == 5 + payload = tx.litellm_teamtable.create.await_args.kwargs["data"] + assert payload["model_id"] == 42 + assert payload["object_permission_id"] == "permissions" + metadata = json.loads(payload["metadata"]) if isinstance(payload["metadata"], str) else payload["metadata"] + assert metadata["team_member_budget_id"] == "member-budget" + invalidate.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 2fba39b30f6..ae1bfe021f0 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -28,8 +28,9 @@ from litellm.proxy.management_helpers.object_permission_utils import ( ) +@pytest.mark.parametrize("use_transaction", [False, True]) @pytest.mark.asyncio -async def test_set_object_permission(): +async def test_set_object_permission(use_transaction: bool): """ Test that _set_object_permission correctly: 1. Creates an object permission record in the database @@ -48,6 +49,9 @@ async def test_set_object_permission(): ) mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) + write_table = MagicMock() if use_transaction else mock_prisma_client.db.litellm_objectpermissiontable + write_table.create = AsyncMock(return_value=mock_created_permission) + # Test data with object_permission data_json = { "user_id": "test_user", @@ -63,7 +67,7 @@ async def test_set_object_permission(): # Call the function result = await _set_object_permission( - data_json=data_json, prisma_client=mock_prisma_client + data_json=data_json, prisma_client=mock_prisma_client, table=write_table if use_transaction else None ) # Verify object_permission_id was added to result @@ -73,10 +77,12 @@ async def test_set_object_permission(): assert "object_permission" not in result # Verify create was called - mock_prisma_client.db.litellm_objectpermissiontable.create.assert_called_once() + write_table.create.assert_awaited_once() + if use_transaction: + mock_prisma_client.db.litellm_objectpermissiontable.create.assert_not_awaited() # Verify the data passed to create excludes None values and object_permission_id - call_args = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args + call_args = write_table.create.call_args created_data = call_args.kwargs["data"] assert "object_permission_id" not in created_data