feat(agents): scim provisioning operations

This commit is contained in:
Joshua Valluru 2026-09-28 16:23:55 -07:00
parent 46119852f0
commit d5d2c5bea1
13 changed files with 3561 additions and 338 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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))

View file

@ -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(

View file

@ -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(

View file

@ -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``:

File diff suppressed because it is too large Load diff

View file

@ -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()

View file

@ -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()

View file

@ -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