diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220100_agent_scim/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220100_agent_scim/migration.sql new file mode 100644 index 00000000000..4623f34f94e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220100_agent_scim/migration.sql @@ -0,0 +1,105 @@ +ALTER TABLE "LiteLLM_AgentIdentity" ADD COLUMN IF NOT EXISTS "provisioning_source_id" TEXT; + +ALTER TABLE "LiteLLM_VerifiedSubject" ADD COLUMN IF NOT EXISTS "agent_id" TEXT, +ADD COLUMN IF NOT EXISTS "parent_client_id" TEXT, +ADD COLUMN IF NOT EXISTS "scim_resource_id" TEXT; + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_SCIMSource" ( + "source_id" TEXT NOT NULL, + "display_name" TEXT NOT NULL, + "tenant_id" TEXT NOT NULL, + "key_hash" TEXT NOT NULL, + "enabled" BOOLEAN NOT NULL DEFAULT true, + "group_mappings" JSONB NOT NULL DEFAULT '[]', + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_SCIMSource_pkey" PRIMARY KEY ("source_id") +); + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_SCIMResource" ( + "id" TEXT NOT NULL, + "source_id" TEXT NOT NULL, + "kind" TEXT NOT NULL, + "external_id" TEXT NOT NULL, + "user_name" TEXT, + "display_name" TEXT NOT NULL, + "document" JSONB NOT NULL, + "active" BOOLEAN NOT NULL DEFAULT true, + "deleted" BOOLEAN NOT NULL DEFAULT false, + "local_id" TEXT, + "human_email" TEXT, + "human_subject_key" TEXT, + "member_ids" TEXT[] DEFAULT ARRAY[]::TEXT[], + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_SCIMResource_pkey" PRIMARY KEY ("id") +); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_scim_resource_id_key" ON "LiteLLM_VerifiedSubject"("scim_resource_id"); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_agent_id_idx" ON "LiteLLM_VerifiedSubject"("agent_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMSource_key_hash_key" ON "LiteLLM_SCIMSource"("key_hash"); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_source_id_kind_idx" ON "LiteLLM_SCIMResource"("source_id", "kind"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_source_id_kind_external_id_key" ON "LiteLLM_SCIMResource"("source_id", "kind", "external_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_source_id_kind_user_name_key" ON "LiteLLM_SCIMResource"("source_id", "kind", "user_name"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_local_id_key" ON "LiteLLM_SCIMResource"("local_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_human_email_key" ON "LiteLLM_SCIMResource"("human_email"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_SCIMResource_human_subject_key_key" ON "LiteLLM_SCIMResource"("human_subject_key"); + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_agent_id_fkey') THEN + ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE; + END IF; +END $$; + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_scim_resource_id_fkey') THEN + ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_scim_resource_id_fkey" FOREIGN KEY ("scim_resource_id") REFERENCES "LiteLLM_SCIMResource"("id") ON DELETE RESTRICT ON UPDATE CASCADE; + END IF; +END $$; + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_SCIMResource_source_id_fkey') THEN + ALTER TABLE "LiteLLM_SCIMResource" ADD CONSTRAINT "LiteLLM_SCIMResource_source_id_fkey" FOREIGN KEY ("source_id") REFERENCES "LiteLLM_SCIMSource"("source_id") ON DELETE RESTRICT ON UPDATE CASCADE; + END IF; +END $$; + +-- AddCheckConstraint +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_kind_shape') THEN + ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_kind_shape" CHECK ( + ("kind" = 'human' AND "user_id" IS NOT NULL AND "agent_id" IS NULL + AND "parent_client_id" IS NULL AND "scim_resource_id" IS NULL AND "verified_via" = 'sso_interactive') + OR + ("kind" = 'agent_user' AND "user_id" IS NULL AND "parent_client_id" IS NOT NULL + AND "scim_resource_id" IS NOT NULL AND "verified_via" = 'scim') + ); + END IF; +END $$; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 92dfe53e29e..98ce879eedf 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -85,8 +85,9 @@ model LiteLLM_AgentsTable { identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] budget_id String? @unique - spend_window DateTime? + spend_window DateTime? litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) + verified_subjects LiteLLM_VerifiedSubject[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -106,6 +107,7 @@ model LiteLLM_AgentIdentity { tenant_id String client_id String service_principal_id String? + provisioning_source_id String? required_roles String[] @default([]) required_scopes String[] @default(["user_impersonation"]) revision String @default(uuid()) @@ -138,13 +140,52 @@ model LiteLLM_VerifiedSubject { kind String @default("human") user_id String? user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + parent_client_id String? + scim_resource_id String? @unique + scim_resource LiteLLM_SCIMResource? @relation(fields: [scim_resource_id], references: [id], onDelete: Restrict) verified_via String @default("sso_interactive") verified_at DateTime @default(now()) @@unique([issuer, tenant_id, oid]) @@index([user_id]) + @@index([agent_id]) } +model LiteLLM_SCIMSource { + source_id String @id @default(uuid()) + display_name String + tenant_id String + key_hash String @unique + enabled Boolean @default(true) + group_mappings Json @default("[]") + resources LiteLLM_SCIMResource[] + created_at DateTime @default(now()) + updated_at DateTime @updatedAt +} +model LiteLLM_SCIMResource { + id String @id @default(uuid()) + source_id String + source LiteLLM_SCIMSource @relation(fields: [source_id], references: [source_id], onDelete: Restrict) + kind String + external_id String + user_name String? + display_name String + document Json + active Boolean @default(true) + deleted Boolean @default(false) + local_id String? @unique + human_email String? @unique + human_subject_key String? @unique + member_ids String[] @default([]) + subject LiteLLM_VerifiedSubject? + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + @@unique([source_id, kind, external_id]) + @@unique([source_id, kind, user_name]) + @@index([source_id, kind]) +} model LiteLLM_OrganizationTable { @@ -300,6 +341,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index ed7460ac815..5719c7a69b7 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -47134,6 +47134,21 @@ "title": "HTTPValidationError", "type": "object" }, + "SCIMAgentUser": { + "additionalProperties": false, + "properties": { + "identityParentId": { + "format": "uuid", + "title": "Identityparentid", + "type": "string" + } + }, + "required": [ + "identityParentId" + ], + "title": "SCIMAgentUser", + "type": "object" + }, "SCIMEnterpriseUser": { "properties": { "costCenter": { @@ -47800,6 +47815,16 @@ } ] }, + "urn:ietf:params:scim:schemas:extension:litellmAgent:2.0:User": { + "anyOf": [ + { + "$ref": "#/components/schemas/SCIMAgentUser" + }, + { + "type": "null" + } + ] + }, "userName": { "anyOf": [ { diff --git a/litellm/proxy/management_endpoints/scim/agent_provisioning.py b/litellm/proxy/management_endpoints/scim/agent_provisioning.py new file mode 100644 index 00000000000..1ba1b3a95a4 --- /dev/null +++ b/litellm/proxy/management_endpoints/scim/agent_provisioning.py @@ -0,0 +1,200 @@ +import re +from collections import deque +from dataclasses import dataclass +from functools import reduce +from itertools import chain +from typing import Final, Literal + +from fastapi import HTTPException +from prisma import Prisma +from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource +from prisma.types import ( + LiteLLM_SCIMResourceUpdateInput, + LiteLLM_SCIMResourceWhereUniqueInput, +) +from pydantic import TypeAdapter, ValidationError + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import PrismaClient +from litellm.repositories.table_repositories import SCIMSourceRepository +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import ( + SCIM_AGENT_USER_SCHEMA, +) +from litellm.types.proxy.management_endpoints.scim_v2 import ( + SCIMGroup, + SCIMMember, + SCIMPatchOp, + SCIMPatchOperation, + SCIMUser, + SCIMUserName, +) + + +@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": [{"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 = { + "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 = {"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 () + return tuple({member.value: member for member in (*tuple(current.members or ()), *incoming)}.values()) + + +def _patch_group_operation( + current: SCIMGroup | SCIMProvisioningFailure, item: SCIMPatchOperation +) -> SCIMGroup | SCIMProvisioningFailure: + if isinstance(current, SCIMProvisioningFailure): + return current + if (item.path or "").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 group_members_after_patch(group: SCIMGroup, patch: SCIMPatchOp) -> SCIMGroup | SCIMProvisioningFailure: + return reduce(_patch_group_operation, patch.Operations, group) + + +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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 92dfe53e29e..98ce879eedf 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -85,8 +85,9 @@ model LiteLLM_AgentsTable { identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] budget_id String? @unique - spend_window DateTime? + spend_window DateTime? litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) + verified_subjects LiteLLM_VerifiedSubject[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -106,6 +107,7 @@ model LiteLLM_AgentIdentity { tenant_id String client_id String service_principal_id String? + provisioning_source_id String? required_roles String[] @default([]) required_scopes String[] @default(["user_impersonation"]) revision String @default(uuid()) @@ -138,13 +140,52 @@ model LiteLLM_VerifiedSubject { kind String @default("human") user_id String? user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + parent_client_id String? + scim_resource_id String? @unique + scim_resource LiteLLM_SCIMResource? @relation(fields: [scim_resource_id], references: [id], onDelete: Restrict) verified_via String @default("sso_interactive") verified_at DateTime @default(now()) @@unique([issuer, tenant_id, oid]) @@index([user_id]) + @@index([agent_id]) } +model LiteLLM_SCIMSource { + source_id String @id @default(uuid()) + display_name String + tenant_id String + key_hash String @unique + enabled Boolean @default(true) + group_mappings Json @default("[]") + resources LiteLLM_SCIMResource[] + created_at DateTime @default(now()) + updated_at DateTime @updatedAt +} +model LiteLLM_SCIMResource { + id String @id @default(uuid()) + source_id String + source LiteLLM_SCIMSource @relation(fields: [source_id], references: [source_id], onDelete: Restrict) + kind String + external_id String + user_name String? + display_name String + document Json + active Boolean @default(true) + deleted Boolean @default(false) + local_id String? @unique + human_email String? @unique + human_subject_key String? @unique + member_ids String[] @default([]) + subject LiteLLM_VerifiedSubject? + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + @@unique([source_id, kind, external_id]) + @@unique([source_id, kind, user_name]) + @@index([source_id, kind]) +} model LiteLLM_OrganizationTable { @@ -300,6 +341,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 292694747f3..11fd91dc933 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -263,5 +263,13 @@ class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteL table_name = "litellm_adaptiveroutersession" +class SCIMSourceRepository(PrismaTableRepository["prisma_models.LiteLLM_SCIMSource"]): + table_name = "litellm_scimsource" + + +class SCIMResourceRepository(PrismaTableRepository["prisma_models.LiteLLM_SCIMResource"]): + table_name = "litellm_scimresource" + + class RetiredAgentRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgent"]): table_name = "litellm_retiredagent" diff --git a/litellm/types/proxy/management_endpoints/scim_agent_provisioning.py b/litellm/types/proxy/management_endpoints/scim_agent_provisioning.py new file mode 100644 index 00000000000..fe82d7b3702 --- /dev/null +++ b/litellm/types/proxy/management_endpoints/scim_agent_provisioning.py @@ -0,0 +1,43 @@ +from typing import Final +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, SecretStr + +SCIM_AGENT_USER_SCHEMA: Final = "urn:ietf:params:scim:schemas:extension:litellmAgent:2.0:User" + + +def canonical_directory_id(value: str) -> str: + try: + return str(UUID(value)) + except ValueError: + return value + + +class SCIMAgentUser(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + identityParentId: UUID + + +class SCIMGroupMapping(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + external_group_id: UUID + access_group_ids: tuple[str, ...] = Field(min_length=1) + + +class SCIMSourceConfig(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + display_name: str = Field(min_length=1) + tenant_id: UUID + enabled: bool = True + group_mappings: tuple[SCIMGroupMapping, ...] = () + + +class SCIMSourceCreate(SCIMSourceConfig): + provisioning_token: SecretStr + + +class SCIMSourceResponse(SCIMSourceConfig): + source_id: str diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 6f2c48ab283..ef4be90d9a8 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -13,6 +13,8 @@ from pydantic import ( ) from pydantic_core.core_schema import SerializerFunctionWrapHandler +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIM_AGENT_USER_SCHEMA, SCIMAgentUser + SCIM_ENTERPRISE_USER_SCHEMA: Final = "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User" SCIM_ENTERPRISE_METADATA_KEY: Final = "scim_enterprise" SCIM_ENTITLEMENTS_METADATA_KEY: Final = "scim_entitlements" @@ -106,6 +108,7 @@ class SCIMEnterpriseUser(BaseModel): class SCIMUser(SCIMResource): model_config = ConfigDict(populate_by_name=True) + agent_user: SCIMAgentUser | None = Field(default=None, alias=SCIM_AGENT_USER_SCHEMA) userName: str | None = None name: SCIMUserName | None = None displayName: str | None = None @@ -120,9 +123,18 @@ class SCIMUser(SCIMResource): serialization_alias=SCIM_ENTERPRISE_USER_SCHEMA, ) + @model_validator(mode="after") + def validate_agent_extension(self) -> "SCIMUser": + if SCIM_AGENT_USER_SCHEMA in self.schemas and self.agent_user is None: + raise ValueError("The agent-user schema requires identityParentId") + return self + @model_serializer(mode="wrap") def _omit_absent_optional_blocks(self, handler: SerializerFunctionWrapHandler) -> dict[str, object]: dumped: Final = handler(self) + if self.agent_user is None: + dumped.pop(SCIM_AGENT_USER_SCHEMA, None) + dumped.pop("agent_user", None) if self.enterprise_user is None: dumped.pop(SCIM_ENTERPRISE_USER_SCHEMA, None) dumped.pop("enterprise_user", None) diff --git a/schema.prisma b/schema.prisma index 92dfe53e29e..98ce879eedf 100644 --- a/schema.prisma +++ b/schema.prisma @@ -85,8 +85,9 @@ model LiteLLM_AgentsTable { identity LiteLLM_AgentIdentity? retired_identities LiteLLM_RetiredAgentIdentity[] budget_id String? @unique - spend_window DateTime? + spend_window DateTime? litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id]) + verified_subjects LiteLLM_VerifiedSubject[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -106,6 +107,7 @@ model LiteLLM_AgentIdentity { tenant_id String client_id String service_principal_id String? + provisioning_source_id String? required_roles String[] @default([]) required_scopes String[] @default(["user_impersonation"]) revision String @default(uuid()) @@ -138,13 +140,52 @@ model LiteLLM_VerifiedSubject { kind String @default("human") user_id String? user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + parent_client_id String? + scim_resource_id String? @unique + scim_resource LiteLLM_SCIMResource? @relation(fields: [scim_resource_id], references: [id], onDelete: Restrict) verified_via String @default("sso_interactive") verified_at DateTime @default(now()) @@unique([issuer, tenant_id, oid]) @@index([user_id]) + @@index([agent_id]) } +model LiteLLM_SCIMSource { + source_id String @id @default(uuid()) + display_name String + tenant_id String + key_hash String @unique + enabled Boolean @default(true) + group_mappings Json @default("[]") + resources LiteLLM_SCIMResource[] + created_at DateTime @default(now()) + updated_at DateTime @updatedAt +} +model LiteLLM_SCIMResource { + id String @id @default(uuid()) + source_id String + source LiteLLM_SCIMSource @relation(fields: [source_id], references: [source_id], onDelete: Restrict) + kind String + external_id String + user_name String? + display_name String + document Json + active Boolean @default(true) + deleted Boolean @default(false) + local_id String? @unique + human_email String? @unique + human_subject_key String? @unique + member_ids String[] @default([]) + subject LiteLLM_VerifiedSubject? + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + @@unique([source_id, kind, external_id]) + @@unique([source_id, kind, user_name]) + @@index([source_id, kind]) +} model LiteLLM_OrganizationTable { @@ -300,6 +341,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py new file mode 100644 index 00000000000..fb0bea69f70 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py @@ -0,0 +1,283 @@ +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException +from pydantic import ValidationError + +from litellm.proxy.management_endpoints.scim.agent_provisioning import ( + SCIMProvisioningFailure, + apply_user_patch, + group_members_after_patch, +) +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIM_AGENT_USER_SCHEMA +from litellm.types.proxy.management_endpoints.scim_v2 import SCIMGroup, SCIMMember, SCIMPatchOp, SCIMUser + +PARENT: Final = "11111111-1111-4111-8111-111111111111" + +SUBJECT: Final = "22222222-2222-4222-8222-222222222222" + + +@pytest.mark.parametrize( + "value,expected", + [ + ("ABCDEF00-1234-4234-9234-123456789ABC", "abcdef00-1234-4234-9234-123456789abc"), + ("opaque-directory-id", "opaque-directory-id"), + ], +) +def test_directory_ids_normalize_uuids_without_rewriting_opaque_ids(value: str, expected: str) -> None: + from litellm.types.proxy.management_endpoints.scim_agent_provisioning import canonical_directory_id + + assert canonical_directory_id(value) == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", [True, False, None]) +async def test_source_token_lookup_uses_current_writer_state(enabled: bool | None) -> None: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.management_endpoints.scim.agent_provisioning import source_for_auth + + source: Final = SimpleNamespace(source_id="source", enabled=enabled) if enabled is not None else None + client: Final = MagicMock() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=source) + auth: Final = UserAPIKeyAuth(token="source-token-hash") + if enabled is False: + with pytest.raises(HTTPException) as failure: + await source_for_auth(auth, client) + assert failure.value.status_code == 403 + else: + assert await source_for_auth(auth, client) is source + client.writer_db.litellm_scimsource.find_unique.assert_awaited_once_with(where={"key_hash": "source-token-hash"}) + client.db.litellm_scimsource.find_unique.assert_not_called() + assert await source_for_auth(None, client) is None + assert await source_for_auth(UserAPIKeyAuth(), client) is None + + +@pytest.mark.asyncio +async def test_directory_documents_use_authoritative_activity_and_membership() -> None: + from litellm.proxy.management_endpoints.scim.agent_provisioning import ( + group_document, + remove_group_member, + user_document, + ) + + row: Final = SimpleNamespace(id="row", active=False, document={"userName": "subject", "active": True}) + assert user_document(row).active is False + assert user_document(row).id == "row" + group: Final = SimpleNamespace( + id="group", + member_ids=["keep", "remove"], + document={"displayName": "Directory", "members": [{"value": "stale"}]}, + ) + assert [member.value for member in group_document(group).members] == ["keep", "remove"] + client: Final = MagicMock() + client.litellm_scimresource.update = AsyncMock() + await remove_group_member(client, group, "remove") + client.litellm_scimresource.update.assert_awaited_once_with(where={"id": "group"}, data={"member_ids": ["keep"]}) + assert group.member_ids == ["keep", "remove"] + + +def agent_user() -> SCIMUser: + return SCIMUser.model_validate( + { + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User", SCIM_AGENT_USER_SCHEMA], + "id": "stable-scim-id", + "externalId": SUBJECT, + "userName": "agent@example.com", + "active": True, + SCIM_AGENT_USER_SCHEMA: {"identityParentId": PARENT}, + } + ) + + +@pytest.mark.parametrize("wire_value, expected", [("False", False), ("True", True), (False, False), (True, True)]) +def test_entra_boolean_patch_preserves_identity_and_returns_json_boolean(wire_value: object, expected: bool) -> None: + original: Final = agent_user() + result: Final = apply_user_patch( + original, SCIMPatchOp(Operations=[{"op": "Replace", "path": "active", "value": wire_value}]) + ) + assert isinstance(result, SCIMUser) + assert result.active is expected + assert result.id == original.id + assert result.externalId == SUBJECT + assert result.agent_user == original.agent_user + assert result.model_dump(by_alias=True)["active"] is expected + assert original.active is True + + +@pytest.mark.parametrize( + "path, value", + [ + ("externalId", PARENT), + (SCIM_AGENT_USER_SCHEMA + ":identityParentId", SUBJECT), + (None, {SCIM_AGENT_USER_SCHEMA: {"identityParentId": SUBJECT}}), + ("active", "garbage"), + ], +) +def test_identity_rebinding_and_invalid_active_patch_are_rejected(path: str | None, value: object) -> None: + result: Final = apply_user_patch( + agent_user(), SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}]) + ) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +def test_rename_does_not_reenable_a_disabled_subject() -> None: + original: Final = agent_user().model_copy(update={"active": False}) + result: Final = apply_user_patch( + original, SCIMPatchOp(Operations=[{"op": "replace", "value": {"displayName": "Renamed"}}]) + ) + assert isinstance(result, SCIMUser) + assert result.displayName == "Renamed" + assert result.active is False + assert result.id == original.id + + +def test_agent_schema_without_parent_is_not_a_human() -> None: + with pytest.raises(ValidationError, match="identityParentId"): + SCIMUser.model_validate({"schemas": [SCIM_AGENT_USER_SCHEMA], "userName": "agent@example.com"}) + + +def test_group_removal_preserves_other_members_and_repeat_removal_is_idempotent() -> None: + group: Final = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + displayName="Mixed", + members=[SCIMMember(value="agent"), SCIMMember(value="human")], + ) + patch: Final = SCIMPatchOp(Operations=[{"op": "Remove", "path": 'members[value eq "agent"]'}]) + result: Final = group_members_after_patch(group, patch) + assert isinstance(result, SCIMGroup) + assert result.members == [SCIMMember(value="human")] + assert group_members_after_patch(result, patch) == result + assert group.members == [SCIMMember(value="agent"), SCIMMember(value="human")] + + +def test_group_replace_removes_omitted_members_and_empty_replace_removes_all() -> None: + group: Final = SCIMGroup( + schemas=[], displayName="Mixed", members=[SCIMMember(value="agent"), SCIMMember(value="human")] + ) + result: Final = group_members_after_patch( + group, SCIMPatchOp(Operations=[{"op": "replace", "path": "members", "value": [{"value": "human"}]}]) + ) + assert isinstance(result, SCIMGroup) + assert result.members == [SCIMMember(value="human")] + empty: Final = group_members_after_patch( + result, SCIMPatchOp(Operations=[{"op": "replace", "path": "members", "value": []}]) + ) + assert isinstance(empty, SCIMGroup) + assert empty.members == [] + + +@pytest.mark.parametrize( + "path,value", + [ + (SCIM_AGENT_USER_SCHEMA + ":identityParentId", PARENT), + (None, {SCIM_AGENT_USER_SCHEMA: {"identityParentId": PARENT}}), + ("schemas", [SCIM_AGENT_USER_SCHEMA]), + ], +) +def test_human_patch_cannot_smuggle_an_agent_identity(path: str | None, value: object) -> None: + from litellm.proxy.management_endpoints.scim.agent_provisioning import patch_changes_identity + + patch: Final = SCIMPatchOp(Operations=[{"op": "add", "path": path, "value": value}]) + assert patch_changes_identity(patch) + + +@pytest.mark.parametrize("identity_marker", [True, False]) +def test_deep_patch_checks_identity_without_exhausting_the_call_stack(identity_marker: bool) -> None: + from functools import reduce + + from litellm.proxy.management_endpoints.scim.agent_provisioning import patch_changes_identity + + leaf: Final = {"identityParentId": PARENT} if identity_marker else {"displayName": "Renamed"} + nested: Final = reduce(lambda value, _: {"nested": [value]}, range(1200), leaf) + patch: Final = SCIMPatchOp(Operations=[{"op": "replace", "value": nested}]) + assert patch_changes_identity(patch) is identity_marker + + +def test_patch_error_does_not_apply_later_operations() -> None: + patch: Final = SCIMPatchOp( + Operations=[ + {"op": "replace", "path": "externalId", "value": "foreign-subject"}, + {"op": "replace", "path": "displayName", "value": "renamed"}, + ] + ) + result: Final = apply_user_patch(agent_user(), patch) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +@pytest.mark.parametrize( + "operation,expected", + [ + ({"op": "add", "path": "members", "value": [{"value": "human"}, {"value": "second"}]}, ["human", "second"]), + ({"op": "remove", "path": "members", "value": [{"value": "human"}]}, []), + ({"op": "remove", "path": "members"}, []), + ({"op": "replace", "path": "displayName", "value": "Renamed"}, ["human"]), + ], +) +def test_group_patch_add_remove_and_rename_keep_membership_consistent( + operation: dict[str, object], expected: list[str] +) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch(group, SCIMPatchOp.model_validate({"Operations": [operation]})) + assert isinstance(result, SCIMGroup) + assert [member.value for member in result.members or []] == expected + assert result.displayName == ("Renamed" if operation["path"] == "displayName" else "Original") + + +@pytest.mark.parametrize( + "operation", + [ + {"op": "replace", "path": "externalId", "value": "foreign"}, + {"op": "add", "path": "members", "value": [{"display": "missing-id"}]}, + ], +) +def test_group_patch_failure_cannot_apply_subsequent_membership_changes(operation: dict[str, object]) -> None: + group: Final = SCIMGroup(schemas=[], displayName="Original", members=[SCIMMember(value="human")]) + result: Final = group_members_after_patch( + group, SCIMPatchOp.model_validate({"Operations": [operation, {"op": "remove", "path": "members"}]}) + ) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + assert group.members == [SCIMMember(value="human")] + + +@pytest.mark.parametrize( + "operation", + [ + {"op": "remove"}, + {"op": "replace", "value": "invalid"}, + {"op": "remove", "path": "userName"}, + ], +) +def test_invalid_agent_profile_patch_is_rejected(operation: dict[str, object]) -> None: + result: Final = apply_user_patch(agent_user(), SCIMPatchOp.model_validate({"Operations": [operation]})) + assert isinstance(result, SCIMProvisioningFailure) + assert result.status == 400 + + +@pytest.mark.parametrize( + "path,value", + [("name.givenName", "New"), ("name.familyName", "Family"), ('emails[type eq "work"].value', "new@example.com")], +) +def test_native_profile_accepts_standard_entra_subattribute_updates(path: str, value: str) -> None: + user: Final = SCIMUser.model_validate( + { + **agent_user().model_dump(by_alias=True), + "name": {"givenName": "Old", "familyName": "Original"}, + "emails": [{"type": "work", "value": "old@example.com"}, {"type": "home", "value": "home@example.com"}], + } + ) + result: Final = apply_user_patch(user, SCIMPatchOp(Operations=[{"op": "replace", "path": path, "value": value}])) + assert isinstance(result, SCIMUser) + assert result.externalId == user.externalId + assert result.agent_user == user.agent_user + if path.startswith("name."): + assert getattr(result.name, path.split(".")[1]) == value + assert result.emails == user.emails + else: + assert result.emails[0].value == value + assert result.emails[1].value == "home@example.com" + assert result.name == user.name diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c3a356a6098..4a3fe152d0c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -41734,6 +41734,14 @@ export interface components { /** Run Id */ run_id: string; }; + /** SCIMAgentUser */ + SCIMAgentUser: { + /** + * Identityparentid + * Format: uuid + */ + identityParentId: string; + }; /** SCIMEnterpriseUser */ SCIMEnterpriseUser: { /** Costcenter */ @@ -41947,6 +41955,7 @@ export interface components { /** Schemas */ schemas: string[]; "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"?: components["schemas"]["SCIMEnterpriseUser"] | null; + "urn:ietf:params:scim:schemas:extension:litellmAgent:2.0:User"?: components["schemas"]["SCIMAgentUser"] | null; /** Username */ userName?: string | null; };