mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(agents): scim provisioning operations
This commit is contained in:
parent
46119852f0
commit
d5d2c5bea1
13 changed files with 3561 additions and 338 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
667
litellm/proxy/management_endpoints/scim/agent_provisioning.py
Normal file
667
litellm/proxy/management_endpoints/scim/agent_provisioning.py
Normal 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
|
||||
255
litellm/proxy/management_endpoints/scim/human_provisioning.py
Normal file
255
litellm/proxy/management_endpoints/scim/human_provisioning.py
Normal 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
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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()
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue