From 46119852f00375d1f1e948464472ace352c25334 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 16:23:54 -0700 Subject: [PATCH] feat(agents): scim source storage and contracts --- .../20260925220100_agent_scim/migration.sql | 105 ++++++++ .../litellm_proxy_extras/schema.prisma | 42 +++ litellm/proxy/_lazy_openapi_snapshot.json | 25 ++ .../scim/source_endpoints.py | 128 ++++++++++ litellm/proxy/schema.prisma | 42 +++ litellm/repositories/table_repositories.py | 8 + .../scim_agent_provisioning.py | 43 ++++ .../proxy/management_endpoints/scim_v2.py | 12 + schema.prisma | 42 +++ .../scim/test_source_endpoints.py | 241 ++++++++++++++++++ .../test_scim_agent_provisioning.py | 24 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 9 + 12 files changed, 721 insertions(+) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260925220100_agent_scim/migration.sql create mode 100644 litellm/proxy/management_endpoints/scim/source_endpoints.py create mode 100644 litellm/types/proxy/management_endpoints/scim_agent_provisioning.py create mode 100644 tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py create mode 100644 tests/unit/types/proxy/management_endpoints/test_scim_agent_provisioning.py 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 fc7325ddac3..d4ba05e93e7 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -88,6 +88,7 @@ model LiteLLM_AgentsTable { spend_window DateTime? lifetime_budget_spend Float @default(0.0) 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? @@ -107,6 +108,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()) @@ -139,13 +141,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 { @@ -301,6 +342,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 89bdd6647ce..a50c0ec31ae 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -49206,6 +49206,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": { @@ -49872,6 +49887,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/source_endpoints.py b/litellm/proxy/management_endpoints/scim/source_endpoints.py new file mode 100644 index 00000000000..fb7136f0a12 --- /dev/null +++ b/litellm/proxy/management_endpoints/scim/source_endpoints.py @@ -0,0 +1,128 @@ +from itertools import chain +from types import SimpleNamespace +from typing import Annotated, Final +from uuid import uuid4 + +from fastapi import APIRouter, Depends, HTTPException +from prisma import Json +from prisma.types import ( + LiteLLM_SCIMSourceCreateInput, + LiteLLM_SCIMSourceOrderByInput, + LiteLLM_SCIMSourceUpdateInput, + LiteLLM_SCIMSourceWhereUniqueInput, + LiteLLM_VerificationTokenWhereUniqueInput, +) +from pydantic import BaseModel + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, hash_token +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.utils import PrismaClient +from litellm.repositories.chunked_in import find_many_in +from litellm.repositories.table_repositories import AccessGroupRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import ( + SCIMSourceConfig, + SCIMSourceCreate, + SCIMSourceResponse, +) + +router: Final = APIRouter(prefix="/sources") + + +class _AdminPolicy(BaseModel): + user_role: str | None = None + allowed_routes: tuple[str, ...] | None = None + + +def _require_source_admin(auth: UserAPIKeyAuth) -> None: + policy: Final = _AdminPolicy.model_validate(auth, from_attributes=True) + if policy.user_role != LitellmUserRoles.PROXY_ADMIN or policy.allowed_routes: + raise HTTPException(403, "Provisioning configuration requires an unrestricted proxy administrator") + + +def source_response(source: object) -> SCIMSourceResponse: + return SCIMSourceResponse.model_validate(source, from_attributes=True) + + +async def _client() -> PrismaClient: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "Provisioning configuration requires a database") + return prisma_client + + +@router.get("", response_model=tuple[SCIMSourceResponse, ...]) +async def list_sources(auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)]) -> tuple[SCIMSourceResponse, ...]: + _require_source_admin(auth) + client: Final = await _client() + order: Final[LiteLLM_SCIMSourceOrderByInput] = {"display_name": "asc"} + async with client.tx() as tx: + sources: Final = await tx.litellm_scimsource.find_many(order=order) + return tuple(source_response(source) for source in sources) + + +@router.post("", response_model=SCIMSourceResponse, status_code=201) +async def create_source( + data: SCIMSourceCreate, auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] +) -> SCIMSourceResponse: + _require_source_admin(auth) + client: Final = await _client() + token_hash: Final = hash_token(data.provisioning_token.get_secret_value()) + source_filter: Final[LiteLLM_SCIMSourceWhereUniqueInput] = {"key_hash": token_hash} + async with client.tx() as tx: + key: Final = await VerificationTokenRepository(SimpleNamespace(db=tx)).table.find_unique( + where=LiteLLM_VerificationTokenWhereUniqueInput(token=token_hash) + ) + if key is None or tuple(key.allowed_routes or ()) != ("/scim/*",): + raise HTTPException(400, "Select a dedicated token restricted to /scim/*") + if await tx.litellm_scimsource.find_unique(where=source_filter) is not None: + raise HTTPException(409, "This token already belongs to a provisioning source") + group_ids: Final = tuple( + frozenset(chain.from_iterable(mapping.access_group_ids for mapping in data.group_mappings)) + ) + groups: Final = await find_many_in( + AccessGroupRepository(SimpleNamespace(db=tx)).table, "access_group_id", group_ids + ) + if frozenset(group.access_group_id for group in groups) != frozenset(group_ids): + raise HTTPException(400, "A mapped access group does not exist") + create_data: Final[LiteLLM_SCIMSourceCreateInput] = LiteLLM_SCIMSourceCreateInput( + source_id=str(uuid4()), + display_name=data.display_name, + tenant_id=str(data.tenant_id), + key_hash=token_hash, + enabled=data.enabled, + group_mappings=Json(data.model_dump(mode="json")["group_mappings"]), + ) + source: Final = await tx.litellm_scimsource.create(data=create_data) + return source_response(source) + + +@router.put("/{source_id}", response_model=SCIMSourceResponse) +async def update_source( + source_id: str, data: SCIMSourceConfig, auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] +) -> SCIMSourceResponse: + _require_source_admin(auth) + client: Final = await _client() + source_filter: Final[LiteLLM_SCIMSourceWhereUniqueInput] = {"source_id": source_id} + async with client.tx() as tx: + source: Final = await tx.litellm_scimsource.find_unique(where=source_filter) + if source is None: + raise HTTPException(404, "Provisioning source not found") + if str(data.tenant_id) != source.tenant_id: + raise HTTPException(409, "A provisioning source's tenant is immutable") + group_ids: Final = tuple( + frozenset(chain.from_iterable(mapping.access_group_ids for mapping in data.group_mappings)) + ) + groups: Final = await find_many_in( + AccessGroupRepository(SimpleNamespace(db=tx)).table, "access_group_id", group_ids + ) + if frozenset(group.access_group_id for group in groups) != frozenset(group_ids): + raise HTTPException(400, "A mapped access group does not exist") + update_data: Final[LiteLLM_SCIMSourceUpdateInput] = { + "display_name": data.display_name, + "enabled": data.enabled, + "group_mappings": Json(data.model_dump(mode="json")["group_mappings"]), + } + updated: Final = await tx.litellm_scimsource.update(where=source_filter, data=update_data) + return source_response(updated) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index fc7325ddac3..d4ba05e93e7 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -88,6 +88,7 @@ model LiteLLM_AgentsTable { spend_window DateTime? lifetime_budget_spend Float @default(0.0) 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? @@ -107,6 +108,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()) @@ -139,13 +141,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 { @@ -301,6 +342,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 4e511a2ec93..fc78bb63eff 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -267,5 +267,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 fc7325ddac3..d4ba05e93e7 100644 --- a/schema.prisma +++ b/schema.prisma @@ -88,6 +88,7 @@ model LiteLLM_AgentsTable { spend_window DateTime? lifetime_budget_spend Float @default(0.0) 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? @@ -107,6 +108,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()) @@ -139,13 +141,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 { @@ -301,6 +342,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_source_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py new file mode 100644 index 00000000000..9de7f0c1110 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py @@ -0,0 +1,241 @@ +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, hash_token +from litellm.proxy.management_endpoints.scim.source_endpoints import create_source, list_sources, update_source +from litellm.proxy.utils import PrismaClient +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIMSourceConfig, SCIMSourceCreate + +TENANT: Final = "11111111-1111-4111-8111-111111111111" +ADMIN: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth", + [ + UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, allowed_routes=["/scim/*"]), + ], +) +async def test_provisioning_credentials_cannot_manage_sources(auth: UserAPIKeyAuth) -> None: + with pytest.raises(HTTPException) as failure: + await list_sources(auth) + assert failure.value.status_code == 403 + + +def source_database(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + client: Final = MagicMock(spec=PrismaClient) + tx: Final = client.tx.return_value.__aenter__.return_value + monkeypatch.setattr(proxy_server, "prisma_client", client) + tx.litellm_scimsource.find_unique = AsyncMock(return_value=None) + tx.litellm_verificationtoken.find_unique = AsyncMock(return_value=SimpleNamespace(allowed_routes=["/scim/*"])) + tx.litellm_accessgrouptable.find_many = AsyncMock(return_value=[]) + tx.litellm_scimsource.create = AsyncMock() + return tx + + +@pytest.mark.asyncio +@pytest.mark.parametrize("routes", [[], ["/scim/*", "openai_routes"], ["/user/*"]]) +async def test_source_token_must_be_restricted_to_scim_only(routes: list[str], monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_verificationtoken.find_unique.return_value = SimpleNamespace(allowed_routes=routes) + request: Final = SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token") + with pytest.raises(HTTPException) as failure: + await create_source(request, ADMIN) + assert failure.value.status_code == 400 + tx.litellm_scimsource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_source_token_cannot_be_reused_for_another_tenant(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="existing-source") + request: Final = SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token") + with pytest.raises(HTTPException) as failure: + await create_source(request, ADMIN) + assert failure.value.status_code == 409 + tx.litellm_scimsource.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_source_creation_stores_a_hash_and_never_returns_the_secret(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.create.return_value = SimpleNamespace( + source_id="source", + display_name="Source", + tenant_id=TENANT, + enabled=True, + group_mappings=[], + ) + request: Final = SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token") + result: Final = await create_source(request, ADMIN) + stored: Final = tx.litellm_scimsource.create.call_args.kwargs["data"] + assert stored["key_hash"] == hash_token("test-token") + assert "test-token" not in str(stored) + assert result.source_id == "source" + assert "key_hash" not in result.model_dump() + assert "provisioning_token" not in result.model_dump() + + +@pytest.mark.asyncio +async def test_source_update_rejects_tenant_rebinding(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source", tenant_id=TENANT) + request: Final = SCIMSourceConfig(display_name="Renamed", tenant_id="22222222-2222-4222-8222-222222222222") + with pytest.raises(HTTPException) as failure: + await update_source("source", request, ADMIN) + assert failure.value.status_code == 409 + tx.litellm_scimsource.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_source_creation_rejects_a_mapping_to_missing_access_groups(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + request: Final = SCIMSourceCreate( + display_name="Source", + tenant_id=TENANT, + provisioning_token="test-token", + group_mappings=[{"external_group_id": TENANT, "access_group_ids": ["missing"]}], + ) + with pytest.raises(HTTPException) as failure: + await create_source(request, ADMIN) + assert failure.value.status_code == 400 + tx.litellm_scimsource.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", [True, False]) +async def test_source_update_preserves_tenant_and_maps_existing_access_groups( + enabled: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source", tenant_id=TENANT) + tx.litellm_accessgrouptable.find_many.return_value = [SimpleNamespace(access_group_id="group")] + request: Final = SCIMSourceConfig( + display_name="Renamed", + tenant_id=TENANT, + enabled=enabled, + group_mappings=[{"external_group_id": TENANT, "access_group_ids": ["group"]}], + ) + tx.litellm_scimsource.update = AsyncMock( + return_value=SimpleNamespace(source_id="source", **request.model_dump(mode="json")) + ) + result: Final = await update_source("source", request, ADMIN) + assert result.enabled is enabled + assert result.group_mappings[0].access_group_ids == ("group",) + stored: Final = tx.litellm_scimsource.update.call_args.kwargs["data"] + assert stored["enabled"] is enabled + assert "tenant_id" not in stored and "key_hash" not in stored + + +@pytest.mark.asyncio +async def test_missing_source_update_is_not_an_upsert(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + request: Final = SCIMSourceConfig(display_name="Source", tenant_id=TENANT) + with pytest.raises(HTTPException) as failure: + await update_source("missing", request, ADMIN) + assert failure.value.status_code == 404 + tx.litellm_scimsource.update.assert_not_called() + + +@pytest.mark.asyncio +async def test_source_list_returns_public_configuration_without_credentials(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.find_many = AsyncMock( + return_value=[ + SimpleNamespace( + source_id="source", + display_name="Source", + tenant_id=TENANT, + enabled=True, + group_mappings=[], + key_hash="private-hash", + ) + ] + ) + result: Final = await list_sources(ADMIN) + assert len(result) == 1 and result[0].source_id == "source" + assert "private-hash" not in str(result) + assert "key_hash" not in result[0].model_dump() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["list", "create", "update"]) +@pytest.mark.parametrize("routes", [["/user/*"], ["openai_routes"], ["/scim/v2/sources"]]) +async def test_any_scoped_administrator_is_denied_source_configuration(operation: str, routes: list[str]) -> None: + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, allowed_routes=routes) + request: Final = ( + list_sources(auth) + if operation == "list" + else create_source( + SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token"), auth + ) + if operation == "create" + else update_source("source", SCIMSourceConfig(display_name="Source", tenant_id=TENANT), auth) + ) + with pytest.raises(HTTPException) as failure: + await request + assert failure.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_source_configuration_requires_available_database(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + with pytest.raises(HTTPException) as failure: + await list_sources(ADMIN) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_source_update_rejects_missing_mapped_group_before_writing(monkeypatch: pytest.MonkeyPatch) -> None: + tx: Final = source_database(monkeypatch) + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source", tenant_id=TENANT) + request: Final = SCIMSourceConfig( + display_name="Renamed", + tenant_id=TENANT, + group_mappings=[{"external_group_id": TENANT, "access_group_ids": ["missing"]}], + ) + with pytest.raises(HTTPException) as failure: + await update_source("source", request, ADMIN) + assert failure.value.status_code == 400 + tx.litellm_scimsource.update.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "update"]) +async def test_source_mapping_accepts_access_groups_across_query_batches( + operation: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE + + tx: Final = source_database(monkeypatch) + group_ids: Final = tuple(f"group-{index}" for index in range(IN_LIST_CHUNK_SIZE + 1)) + groups: Final = [SimpleNamespace(access_group_id=group_id) for group_id in group_ids] + tx.litellm_accessgrouptable.find_many.side_effect = [groups[:-1], groups[-1:]] + request: Final = SCIMSourceCreate( + display_name="Source", + tenant_id=TENANT, + provisioning_token="test-token", + group_mappings=[{"external_group_id": TENANT, "access_group_ids": group_ids}], + ) + stored: Final = SimpleNamespace( + source_id="source", **request.model_dump(mode="json", exclude={"provisioning_token"}) + ) + tx.litellm_scimsource.create.return_value = stored + tx.litellm_scimsource.update = AsyncMock(return_value=stored) + if operation == "update": + tx.litellm_scimsource.find_unique.return_value = stored + result: Final = ( + await create_source(request, ADMIN) if operation == "create" else await update_source("source", request, ADMIN) + ) + assert result.group_mappings[0].access_group_ids == group_ids + assert tx.litellm_accessgrouptable.find_many.await_count == 2 diff --git a/tests/unit/types/proxy/management_endpoints/test_scim_agent_provisioning.py b/tests/unit/types/proxy/management_endpoints/test_scim_agent_provisioning.py new file mode 100644 index 00000000000..d12771f4a3a --- /dev/null +++ b/tests/unit/types/proxy/management_endpoints/test_scim_agent_provisioning.py @@ -0,0 +1,24 @@ +import pytest +from pydantic import ValidationError + +from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIM_AGENT_USER_SCHEMA +from litellm.types.proxy.management_endpoints.scim_v2 import SCIMUser + + +@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 + + +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"}) + diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f199f132812..4f1616a3c28 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -43091,6 +43091,14 @@ export interface components { /** Lookback Hours */ lookback_hours?: number | null; }; + /** SCIMAgentUser */ + SCIMAgentUser: { + /** + * Identityparentid + * Format: uuid + */ + identityParentId: string; + }; /** SCIMEnterpriseUser */ SCIMEnterpriseUser: { /** Costcenter */ @@ -43304,6 +43312,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; };