This commit is contained in:
joshua-berri 2026-10-01 01:59:47 +00:00 • committed by GitHub
commit 4b6f6634f4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 2909 additions and 178 deletions

View file

@ -2662,6 +2662,17 @@
"title": "Provider",
"type": "string"
},
"provisioning_source_id": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Provisioning Source Id"
},
"required_roles": {
"default": [],
"items": {
@ -3173,6 +3184,25 @@
],
"title": "Created By"
},
"directory_access_group_ids": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Directory Access Group Ids"
},
"directory_active": {
"default": true,
"title": "Directory Active",
"type": "boolean"
},
"enabled": {
"default": true,
"title": "Enabled",
@ -3680,6 +3710,17 @@
"title": "Provider",
"type": "string"
},
"provisioning_source_id": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Provisioning Source Id"
},
"required_roles": {
"default": [],
"items": {
@ -3874,6 +3915,25 @@
},
"ManagedAgentIdentityStatus": {
"properties": {
"directory_access_group_ids": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Directory Access Group Ids"
},
"directory_active": {
"default": true,
"title": "Directory Active",
"type": "boolean"
},
"enabled": {
"default": true,
"title": "Enabled",
@ -49407,6 +49467,30 @@
"title": "SCIMGroup",
"type": "object"
},
"SCIMGroupMapping": {
"additionalProperties": false,
"properties": {
"access_group_ids": {
"items": {
"type": "string"
},
"minItems": 1,
"title": "Access Group Ids",
"type": "array"
},
"external_group_id": {
"format": "uuid",
"title": "External Group Id",
"type": "string"
}
},
"required": [
"external_group_id",
"access_group_ids"
],
"title": "SCIMGroupMapping",
"type": "object"
},
"SCIMListResponse": {
"properties": {
"Resources": {
@ -49752,6 +49836,120 @@
"title": "SCIMServiceProviderConfig",
"type": "object"
},
"SCIMSourceConfig": {
"additionalProperties": false,
"properties": {
"display_name": {
"minLength": 1,
"title": "Display Name",
"type": "string"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
},
"group_mappings": {
"default": [],
"items": {
"$ref": "#/components/schemas/SCIMGroupMapping"
},
"title": "Group Mappings",
"type": "array"
},
"tenant_id": {
"format": "uuid",
"title": "Tenant Id",
"type": "string"
}
},
"required": [
"display_name",
"tenant_id"
],
"title": "SCIMSourceConfig",
"type": "object"
},
"SCIMSourceCreate": {
"additionalProperties": false,
"properties": {
"display_name": {
"minLength": 1,
"title": "Display Name",
"type": "string"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
},
"group_mappings": {
"default": [],
"items": {
"$ref": "#/components/schemas/SCIMGroupMapping"
},
"title": "Group Mappings",
"type": "array"
},
"provisioning_token": {
"format": "password",
"title": "Provisioning Token",
"type": "string",
"writeOnly": true
},
"tenant_id": {
"format": "uuid",
"title": "Tenant Id",
"type": "string"
}
},
"required": [
"display_name",
"tenant_id",
"provisioning_token"
],
"title": "SCIMSourceCreate",
"type": "object"
},
"SCIMSourceResponse": {
"additionalProperties": false,
"properties": {
"display_name": {
"minLength": 1,
"title": "Display Name",
"type": "string"
},
"enabled": {
"default": true,
"title": "Enabled",
"type": "boolean"
},
"group_mappings": {
"default": [],
"items": {
"$ref": "#/components/schemas/SCIMGroupMapping"
},
"title": "Group Mappings",
"type": "array"
},
"source_id": {
"title": "Source Id",
"type": "string"
},
"tenant_id": {
"format": "uuid",
"title": "Tenant Id",
"type": "string"
}
},
"required": [
"display_name",
"tenant_id",
"source_id"
],
"title": "SCIMSourceResponse",
"type": "object"
},
"SCIMUser-Input": {
"properties": {
"active": {
@ -51443,6 +51641,199 @@
"scim"
]
}
},
"/scim/v2/sources": {
"get": {
"operationId": "list_sources_scim_v2_sources_get",
"parameters": [
{
"in": "query",
"name": "feature",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Feature"
}
}
],
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"items": {
"$ref": "#/components/schemas/SCIMSourceResponse"
},
"title": "Response List Sources Scim V2 Sources Get",
"type": "array"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "List Sources",
"tags": [
"scim"
]
},
"post": {
"operationId": "create_source_scim_v2_sources_post",
"parameters": [
{
"in": "query",
"name": "feature",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Feature"
}
}
],
"requestBody": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/SCIMSourceCreate"
}
}
},
"required": true
},
"responses": {
"201": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/SCIMSourceResponse"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Create Source",
"tags": [
"scim"
]
}
},
"/scim/v2/sources/{source_id}": {
"put": {
"operationId": "update_source_scim_v2_sources__source_id__put",
"parameters": [
{
"in": "path",
"name": "source_id",
"required": true,
"schema": {
"title": "Source Id",
"type": "string"
}
},
{
"in": "query",
"name": "feature",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Feature"
}
}
],
"requestBody": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/SCIMSourceConfig"
}
}
},
"required": true
},
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/SCIMSourceResponse"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Update Source",
"tags": [
"scim"
]
}
}
}
},

View file

@ -95,7 +95,19 @@ async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentA
async def manual_ids(_agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ())
async def directory_ids(_agent_id: str) -> AccessGroupIds:
return agent.directory_access_group_ids or ()
manual: Final = await resolve_agent_access_group_ceiling(
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
)
return (manual,) if manual is not None else ()
directory: Final = (
await resolve_agent_access_group_ceiling(
agent.agent_id, load_access_group_ids=directory_ids, load_access_group=authoritative_group
)
if agent.directory_access_group_ids
else AgentAccessGroupCeiling((), frozenset(), frozenset(), frozenset())
if agent.directory_access_group_ids is not None
else None
)
return tuple(ceiling for ceiling in (manual, directory) if ceiling is not None)

View file

@ -186,6 +186,8 @@ class AgentRequestHandler:
elif isinstance(target, AgentResponse) and target.identity_managed:
if (
not target.enabled
or not target.directory_active
or target.directory_access_group_ids == ()
or target.identity is None
or not target.identity.active
or user_api_key_auth is None

View file

@ -221,7 +221,13 @@ def actor_admission_failure(
agent: AgentResponse,
context: ManagedAgentContext | None,
) -> AgentIdentityFailure | None:
if not agent.enabled or agent.identity is None or not agent.identity.active:
if (
not agent.enabled
or not agent.directory_active
or agent.directory_access_group_ids == ()
or agent.identity is None
or not agent.identity.active
):
return AgentIdentityFailure(message="Agent execution is disabled")
if context is None:
return AgentIdentityFailure(message="This agent requires its bound identity provider token")

View file

@ -522,6 +522,8 @@ async def get_agent_identity_status(
return ManagedAgentIdentityStatus(
identity=agent.identity,
identity_managed=agent.identity_managed,
directory_active=agent.directory_active,
directory_access_group_ids=agent.directory_access_group_ids,
enabled=agent.enabled,
execution_mode=agent.execution_mode,
last_authenticated_at=agent.identity.last_authenticated_at if agent.identity else None,

View file

@ -1,6 +1,8 @@
import json
from collections.abc import Mapping
from datetime import datetime, timezone
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject
@ -10,6 +12,8 @@ from litellm.repositories.table_repositories import (
AgentsRepository,
RetiredAgentIdentityRepository,
RetiredAgentRepository,
SCIMResourceRepository,
SCIMSourceRepository,
VerifiedSubjectRepository,
)
from litellm.types.agents import AgentResponse
@ -17,6 +21,7 @@ from litellm.types.proxy.agent_identity import (
AgentIdentityFailure,
ManagedAgentContext,
MicrosoftInteractiveSubject,
VerifiedAgentSubject,
VerifiedHumanSubject,
)
@ -43,6 +48,8 @@ class AgentIdentityStore:
AgentIdentityRepository(client, use_writer=True),
VerifiedSubjectRepository(client, use_writer=True),
RetiredAgentIdentityRepository(client, use_writer=True),
SCIMSourceRepository(client, use_writer=True),
SCIMResourceRepository(client, use_writer=True),
RetiredAgentRepository(client, use_writer=True),
cache=cache,
)
@ -53,6 +60,8 @@ class AgentIdentityStore:
identities: AgentIdentityRepository,
humans: VerifiedSubjectRepository,
retired: RetiredAgentIdentityRepository | None = None,
sources: SCIMSourceRepository | None = None,
resources: SCIMResourceRepository | None = None,
retired_agents: RetiredAgentRepository | None = None,
*,
cache: UserApiKeyCache | None = None,
@ -61,6 +70,8 @@ class AgentIdentityStore:
self.identities = identities
self.humans = humans
self.retired = retired
self.sources = sources
self.resources = resources
self.retired_agents = retired_agents
self.cache = cache
@ -75,7 +86,10 @@ class AgentIdentityStore:
row: Final = await self.agents.table.find_unique(where=where, include=include)
if row is None:
return None
return AgentResponse.model_validate(row.model_dump())
agent: Final = AgentResponse.model_validate(row.model_dump())
if agent.identity is None or agent.identity.provisioning_source_id is None:
return agent
return await self.directory_policy(agent)
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent policy could not be loaded")
@ -124,11 +138,28 @@ class AgentIdentityStore:
if not isinstance(issuer, str) or not isinstance(tenant, str) or not isinstance(client, str):
return None
agent_id: Final = await self._bound_agent_id(tenant, client)
if agent_id is None or isinstance(agent_id, AgentIdentityFailure):
if isinstance(agent_id, AgentIdentityFailure):
return agent_id
agent: Final = await self.agent(agent_id)
agent: Final = await self.agent(agent_id) if agent_id is not None else None
if isinstance(agent, AgentIdentityFailure):
return agent
proven: Final = (
await self.subject(issuer, tenant, claims.get("oid"))
if agent is None
or (agent.identity is not None and agent.identity.provisioning_source_id is not None)
or isinstance(claims.get("scp"), str)
else None
)
if isinstance(proven, AgentIdentityFailure):
return proven
if (
proven is not None
and proven.kind == "agent_user"
and (agent_id is None or proven.agent_id != agent_id or proven.parent_client_id != client)
):
return AgentIdentityFailure(message="Provisioned subject does not match an active agent binding")
if agent_id is None:
return None
if (
agent is None
or not agent.identity_managed
@ -137,19 +168,27 @@ class AgentIdentityStore:
or not agent.identity.active
):
return AgentIdentityFailure(message="Agent is disabled or no longer bound to an identity")
subject: Final = classify_agent_subject(agent.identity, claims, agent.execution_mode)
native: Final = (
VerifiedAgentSubject.model_validate(proven.model_dump())
if proven is not None and proven.kind == "agent_user" and proven.agent_id is not None
else None
)
if native is not None and (
agent.identity.provisioning_source_id is None
or not agent.directory_active
or agent.directory_access_group_ids == ()
):
return AgentIdentityFailure(message="Provisioned agent is inactive or has no mapped directory entitlement")
subject: Final = classify_agent_subject(agent.identity, claims, agent.execution_mode, native_subject=native)
if isinstance(subject, AgentIdentityFailure):
return subject
if subject.kind == "application":
if subject.kind in ("application", "agent_user"):
return ManagedAgentContext(
agent_id=agent.agent_id,
binding_revision=agent.identity.revision,
mode=subject.mode,
subject_oid=subject.oid,
)
proven: Final = await self.subject(issuer, tenant, claims.get("oid"))
if isinstance(proven, AgentIdentityFailure):
return proven
human: Final = (
VerifiedHumanSubject.model_validate(proven.model_dump())
if proven is not None
@ -181,6 +220,68 @@ class AgentIdentityStore:
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Subject classification is unavailable")
async def directory_policy(self, agent: AgentResponse) -> AgentResponse | AgentIdentityFailure:
from pydantic import TypeAdapter
from litellm.types.proxy.management_endpoints.scim_agent_provisioning import (
SCIMGroupMapping,
canonical_directory_id,
)
if self.sources is None or self.resources is None or agent.identity is None:
return AgentIdentityFailure(code="policy_unavailable", message="Provisioning state is unavailable")
from prisma.types import LiteLLM_SCIMResourceWhereInput, LiteLLM_SCIMSourceWhereUniqueInput
source_where: Final[LiteLLM_SCIMSourceWhereUniqueInput] = {"source_id": agent.identity.provisioning_source_id}
resource_where: Final[LiteLLM_SCIMResourceWhereInput] = {
"source_id": agent.identity.provisioning_source_id,
"kind": "Users",
"local_id": agent.agent_id,
}
source: Final = await self.sources.table.find_unique(where=source_where)
subjects: Final = await self.resources.table.find_many(where=resource_where)
live: Final = tuple(subject for subject in subjects if subject.active and not subject.deleted)
if source is None or not source.enabled or len(live) != 1:
return agent.model_copy(
update=MappingProxyType({"directory_active": False, "directory_access_group_ids": ()})
)
from litellm.proxy.management_endpoints.scim.agent_provisioning import user_document
directory_user: Final = user_document(live[0])
if (
source.tenant_id != agent.identity.tenant_id
or directory_user.agent_user is None
or str(directory_user.agent_user.identityParentId) != agent.identity.client_id
):
return AgentIdentityFailure(message="Directory identity does not match the registered binding")
proven: Final = await self.subject(agent.identity.issuer, source.tenant_id, live[0].external_id)
if isinstance(proven, AgentIdentityFailure):
return proven
if (
proven is None
or proven.kind != "agent_user"
or proven.verified_via != "scim"
or proven.agent_id != agent.agent_id
or proven.parent_client_id != agent.identity.client_id
or proven.scim_resource_id != live[0].id
):
return AgentIdentityFailure(message="Directory subject does not match the provisioned resource")
groups_where: Final[LiteLLM_SCIMResourceWhereInput] = {
"source_id": source.source_id,
"kind": "Groups",
"deleted": False,
"active": True,
"member_ids": {"has": live[0].id},
}
groups: Final = await self.resources.table.find_many(where=groups_where)
mappings: Final = TypeAdapter(tuple[SCIMGroupMapping, ...]).validate_python(source.group_mappings)
external_ids: Final = frozenset(canonical_directory_id(group.external_id) for group in groups)
mapped: Final = tuple(
mapping.access_group_ids for mapping in mappings if str(mapping.external_group_id) in external_ids
)
ids: Final = tuple(sorted(frozenset(chain.from_iterable(mapped))))
return agent.model_copy(update=MappingProxyType({"directory_active": True, "directory_access_group_ids": ids}))
async def retired_agent(self, agent_id: str) -> bool | AgentIdentityFailure:
if self.retired_agents is None:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")

View file

@ -16,6 +16,7 @@ from litellm.types.proxy.agent_identity import (
AgentIdentityFailure,
AgentSubject,
EntraIdentityConfig,
VerifiedAgentSubject,
)
_MODE: Final = TypeAdapter(AgentExecutionMode)
@ -27,6 +28,7 @@ class IdentityFields(TypedDict, total=False):
client_id: ReadOnly[str]
issuer: ReadOnly[str]
service_principal_id: ReadOnly[str | None]
provisioning_source_id: ReadOnly[str | None]
required_roles: ReadOnly[tuple[str, ...]]
required_scopes: ReadOnly[tuple[str, ...]]
active: ReadOnly[bool]
@ -94,7 +96,15 @@ def _configuration_failure(
mode: AgentExecutionMode,
enabling_without_binding: bool,
) -> AgentIdentityFailure | None:
if identity is not None and mode != "delegated" and not identity.service_principal_id:
if identity is not None and identity.provisioning_source_id is not None:
if mode != "autonomous" or not identity.required_scopes:
return AgentIdentityFailure(message="Provisioned native agent-users require autonomous mode and scopes")
if (
identity is not None
and mode != "delegated"
and not identity.service_principal_id
and not identity.provisioning_source_id
):
return AgentIdentityFailure(
message="Autonomous mode requires the Enterprise application service-principal object ID"
)
@ -129,6 +139,8 @@ def managed_write_fields(
return failure
empty: Final[ManagedWriteFields] = {}
identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
if isinstance(identity_fields, AgentIdentityFailure):
return identity_fields
budget_fields: Final = (
_budget_write(incoming["budget"], existing, updated_by) if "budget" in incoming else empty
)
@ -143,7 +155,19 @@ def managed_write_fields(
return AgentIdentityFailure(message=f"Invalid agent identity or budget configuration: {exc}")
def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
def _identity_write(
identity: EntraIdentityConfig | None, existing: AgentResponse | None
) -> ManagedWriteFields | AgentIdentityFailure:
source_id: Final = existing.identity.provisioning_source_id if existing and existing.identity else None
if existing is not None and existing.identity is not None and source_id is not None:
if identity is None or (identity.tenant_id, identity.client_id, identity.provisioning_source_id) != (
existing.identity.tenant_id,
existing.identity.client_id,
source_id,
):
return AgentIdentityFailure(message="A directory-owned identity cannot be rebound through agent settings")
elif identity is not None and identity.provisioning_source_id is not None:
return AgentIdentityFailure(message="Only SCIM provisioning can create a directory-owned identity")
if identity is None:
unbind: Final[ManagedWriteFields] = {
**(
@ -167,6 +191,7 @@ def _identity_write(identity: EntraIdentityConfig | None, existing: AgentRespons
"tenant_id": identity.tenant_id,
"client_id": identity.client_id,
"service_principal_id": identity.service_principal_id,
"provisioning_source_id": identity.provisioning_source_id,
"required_roles": identity.required_roles,
"required_scopes": identity.required_scopes,
"issuer": identity.issuer,
@ -248,6 +273,8 @@ def classify_agent_subject(
binding: AgentIdentityBinding,
claims: Mapping[str, object],
allowed_mode: AgentExecutionMode,
*,
native_subject: VerifiedAgentSubject | None = None,
) -> AgentSubject | AgentIdentityFailure:
if (claims.get("iss"), claims.get("tid"), claims.get("azp")) != (
binding.issuer,
@ -259,6 +286,26 @@ def classify_agent_subject(
if not isinstance(oid, str) or not oid:
return AgentIdentityFailure(message="Entra token must identify its object subject")
scope: Final = claims.get("scp")
if native_subject is not None:
if (
(
native_subject.issuer,
native_subject.tenant_id,
native_subject.parent_client_id,
native_subject.agent_id,
native_subject.oid,
)
!= (binding.issuer, binding.tenant_id, binding.client_id, binding.agent_id, oid)
or allowed_mode == "delegated"
or claims.get("idtyp") == "app"
or not isinstance(scope, str)
or not binding.required_scopes
or not frozenset(binding.required_scopes).issubset(scope.split())
):
return AgentIdentityFailure(message="Token does not match the provisioned native agent-user")
return AgentSubject(kind="agent_user", oid=oid, mode="autonomous")
if binding.provisioning_source_id is not None:
return AgentIdentityFailure(message="A verified provisioned agent-user token is required")
facets: Final = claims.get("xms_sub_fct")
if facets is not None and (not isinstance(facets, str) or "13" in facets.split()):
return AgentIdentityFailure(message="Native agent-user authentication is not supported by this binding")

View file

@ -142,6 +142,7 @@ from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
DeletedVerificationTokenRepository,
DeprecatedVerificationTokenRepository,
SCIMSourceRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
@ -2433,14 +2434,25 @@ async def _update_key_row_with_soft_budget(
) -> _KeyUpdateResult:
hashed_token: Final = _hash_token_if_needed(key)
key_where: Final[_KeyRowWhere] = {"token": hashed_token}
tx: _KeyUpdateTx
async with prisma_client.tx() as tx:
update_values: Final = await _apply_soft_budget_update(
data=data,
non_default_values=non_default_values,
db=tx,
if "allowed_routes" in data.model_fields_set:
await _lock_and_validate_source_key_change(tx, hashed_token, data.allowed_routes)
permission_values: Final = await _handle_update_object_permission(
data_json=dict(non_default_values),
existing_key_row=existing_key_row,
changed_by=changed_by,
prisma_client=prisma_client,
tx=tx,
)
update_values: Final = (
await _apply_soft_budget_update(
data=data,
non_default_values=permission_values,
db=tx,
existing_key_row=existing_key_row,
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
else permission_values
)
include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True}
updated_row: Final = await tx.litellm_verificationtoken.update(
@ -2567,6 +2579,8 @@ async def _handle_update_object_permission(
data_json: dict,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient,
*,
tx: "Prisma | None" = None,
) -> dict:
"""Persist the requested object permission row and swap it for its id, only after the key policy allowed the write."""
if "object_permission" not in data_json:
@ -2576,6 +2590,7 @@ async def _handle_update_object_permission(
data_json=data_json,
existing_object_permission_id=existing_key_row.object_permission_id,
prisma_client=prisma_client,
tx=tx,
)
# Add the object_permission_id to data_json if one was created/updated
@ -2744,6 +2759,13 @@ async def _process_single_key_update(
prisma_client=prisma_client,
)
if (
"allowed_routes" in update_key_request.model_fields_set
and tuple(update_key_request.allowed_routes or ()) != ("/scim/*",)
and prisma_client is not None
):
await _reject_source_bound_key_change(prisma_client, existing_key_row)
_check_disable_global_guardrails_caller_permission(
update_key_request.disable_global_guardrails,
update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict
@ -3490,6 +3512,9 @@ async def update_key_fn(
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data)
if "allowed_routes" in data.model_fields_set and tuple(data.allowed_routes or ()) != ("/scim/*",):
await _reject_source_bound_key_change(prisma_client, existing_key_row)
# Enforce upperbound key params on update (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
non_default_values: Final = await prepare_key_update_data(
@ -3528,10 +3553,15 @@ async def update_key_fn(
if prisma_client is None:
raise Exception("Not connected to DB!")
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
uses_transaction: Final = bool(data.model_fields_set.intersection(("soft_budget", "allowed_routes")))
update_values: Final = (
non_default_values
if uses_transaction
else await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
)
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
response: Final = (
@ -3543,7 +3573,7 @@ async def update_key_fn(
existing_key_row=existing_key_row,
changed_by=changed_by,
)
if "soft_budget" in data.model_fields_set
if uses_transaction
else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key}))
)
@ -5476,6 +5506,7 @@ async def _insert_deprecated_key(
old_token_hash: str,
new_token_hash: str,
grace_period: str | None,
tx: "Prisma | None" = None,
) -> None:
"""
Insert old key into deprecated table so it remains valid during grace period.
@ -5506,7 +5537,12 @@ async def _insert_deprecated_key(
try:
revoke_at: Final = datetime.now(timezone.utc) + timedelta(seconds=grace_seconds)
await _deprecated_verification_token_table(prisma_client).upsert(
table: Final = (
tx.litellm_deprecatedverificationtoken
if tx is not None
else _deprecated_verification_token_table(prisma_client)
)
await table.upsert(
where={"token": old_token_hash},
data={
"create": {
@ -5532,6 +5568,36 @@ async def _insert_deprecated_key(
)
async def _lock_and_validate_source_key_change(
tx: "Prisma", token: str, allowed_routes: Sequence[str] | None = None
) -> None:
from litellm.proxy.management_endpoints.scim.source_endpoints import lock_provisioning_token
await lock_provisioning_token(tx, token)
key: Final = await tx.litellm_verificationtoken.find_unique(where={"token": token})
if key is None:
raise HTTPException(409, "The key changed during the request; retry with the current key")
if tuple(allowed_routes or ()) == ("/scim/*",):
return
source: Final = await tx.litellm_scimsource.find_unique(where={"key_hash": token})
if source is not None:
raise HTTPException(
409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed"
)
async def _reject_source_bound_key_change(prisma_client: PrismaClient, key: LiteLLM_VerificationToken) -> None:
if tuple(key.allowed_routes or ()) != ("/scim/*",):
return
source: Final = await SCIMSourceRepository(prisma_client, use_writer=True).table.find_unique(
where={"key_hash": key.token}
)
if source is not None:
raise HTTPException(
409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed"
)
async def _execute_virtual_key_regeneration(
*,
prisma_client: PrismaClient,
@ -5549,6 +5615,8 @@ async def _execute_virtual_key_regeneration(
from litellm.proxy import proxy_server
from litellm.proxy.proxy_server import hash_token
await _reject_source_bound_key_change(prisma_client, key_in_db)
# Mirror the /key/update ownership rebind guard. See helper docstring.
_validate_caller_can_change_key_ownership(
data=data,
@ -5624,14 +5692,6 @@ async def _execute_virtual_key_regeneration(
request=data if data is not None else RegenerateKeyRequest(),
),
)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=key_in_db,
prisma_client=prisma_client,
)
update_data.update(update_values)
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
# but their cached jwt_key_mapping entries still point at the old token (LIT-5379).
jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_token(
@ -5639,27 +5699,35 @@ async def _execute_virtual_key_regeneration(
prisma_client=prisma_client,
)
await _persist_deleted_verification_tokens(
keys=[key_in_db],
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
# If grace period set, insert deprecated key so old key remains valid
await _insert_deprecated_key(
prisma_client=prisma_client,
old_token_hash=hashed_api_key,
new_token_hash=new_token_hash,
grace_period=data.grace_period if data else None,
)
updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table(
VerificationTokenRepository(prisma_client)
).update(
where={"token": hashed_api_key},
data=with_settings_updated_at(jsonified_update_data),
)
async with prisma_client.tx() as tx:
await _lock_and_validate_source_key_change(tx, hashed_api_key)
update_values: Final = await _handle_update_object_permission(
data_json=non_default_values,
existing_key_row=key_in_db,
prisma_client=prisma_client,
tx=tx,
)
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(
data={**update_data, **update_values}
)
await _persist_deleted_verification_tokens(
keys=[key_in_db],
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
tx=tx,
)
await _insert_deprecated_key(
prisma_client=prisma_client,
old_token_hash=hashed_api_key,
new_token_hash=new_token_hash,
grace_period=data.grace_period if data else None,
tx=tx,
)
updated_token: Final = await tx.litellm_verificationtoken.update(
where={"token": hashed_api_key},
data=with_settings_updated_at(jsonified_update_data),
)
updated_token_dict: Final[dict[str, object]] = dict(updated_token) if updated_token is not None else {}
updated_token_dict["key"] = new_token
updated_token_dict["token_id"] = updated_token_dict.pop("token")

View file

@ -9,8 +9,9 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from functools import partial
from itertools import chain
from typing import TYPE_CHECKING, Final, NamedTuple, Protocol, overload
from itertools import chain, groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, NamedTuple, Protocol, TypeVar, overload
from fastapi import (
APIRouter,
@ -22,6 +23,7 @@ from fastapi import (
Request,
Response,
)
from prisma.types import LiteLLM_SCIMResourceWhereInput
from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
@ -72,9 +74,11 @@ from litellm.proxy.utils import (
_premium_user_check,
handle_exception_on_proxy,
)
from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE, find_many_in
from litellm.repositories.table_repositories import (
InvitationLinkRepository,
OrganizationMembershipRepository,
SCIMResourceRepository,
TeamMembershipRepository,
)
from litellm.repositories.team_repository import TeamRepository
@ -88,6 +92,8 @@ if TYPE_CHECKING:
from prisma import Prisma
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationToken
from litellm.proxy.management_endpoints.scim.agent_provisioning import AgentProvisioningService
class _UserTableClient(Protocol):
async def find_first(self, where: Mapping[str, object]) -> LiteLLM_UserTable | None: ...
@ -176,6 +182,7 @@ class UserProvisionerHelpers:
prisma_client: PrismaClient,
new_user_request: NewUserRequest,
admin_group: str | None = None,
auth: UserAPIKeyAuth | None = None,
) -> SCIMUser | None:
"""
Check if a user with the given email already exists and update them if found.
@ -212,8 +219,12 @@ class UserProvisionerHelpers:
if not existing_user:
return None
requested_teams: Final = list(dict.fromkeys(new_user_request.teams or []))
await _assert_legacy_source_access(auth, "Users", existing_user.user_id)
requested_teams: Final = list(
dict.fromkeys(team if isinstance(team, str) else team.team_id for team in new_user_request.teams or [])
)
new_teams: Final = requested_teams if requested_teams else list(existing_user.teams or [])
await assert_legacy_team_changes_unowned(auth, existing_user.teams or (), new_teams)
if new_user_request.user_id != existing_user.user_id:
verbose_proxy_logger.info(
@ -276,6 +287,21 @@ scim_router: Final = APIRouter(
dependencies=[Depends(_premium_user_check)],
)
from litellm.proxy.management_endpoints.scim.source_endpoints import router as scim_source_router
scim_router.include_router(scim_source_router)
async def _agent_provisioning_service(auth: UserAPIKeyAuth | None) -> "AgentProvisioningService | None":
from litellm.proxy.management_endpoints.scim.agent_provisioning import AgentProvisioningService, source_for_auth
if not isinstance(auth, UserAPIKeyAuth):
return None
client: Final = await _get_prisma_client_or_raise_exception()
source: Final = await source_for_auth(auth, client)
return AgentProvisioningService(client, source) if source is not None else None
SCIM_MAX_PAGE_SIZE: Final = 100
@ -450,6 +476,219 @@ def _resolve_scim_user_role(
return default_role
async def _source_owned_ids(
prisma_client: "PrismaClient | _GroupWriteDatabase", kind: Literal["Users", "Groups"], local_ids: tuple[str, ...]
) -> frozenset[str]:
if not local_ids:
return frozenset()
table: Final = SCIMResourceRepository(prisma_client, use_writer=True).table
source_kind: Final = LiteLLM_SCIMResourceWhereInput(kind=kind)
by_scim_id: Final = await find_many_in(table, "id", local_ids, where=source_kind)
by_local_id: Final = await find_many_in(table, "local_id", local_ids, where=source_kind)
resources: Final = chain(by_scim_id, by_local_id)
return frozenset(filter(None, chain.from_iterable((resource.id, resource.local_id) for resource in resources)))
_UNOWNED_USER_PREDICATES: Final[Mapping[str, str]] = MappingProxyType(
{
"username": "($1::text IS NULL OR u.user_email = $1 OR u.user_id = $1)",
"emails.value": "($1::text IS NULL OR u.user_email = $1)",
}
)
_UNOWNED_USERS_FROM: Final = """
FROM "LiteLLM_UserTable" u
WHERE {predicate}
AND NOT EXISTS (SELECT 1 FROM "LiteLLM_SCIMResource" r WHERE r.kind = 'Users' AND r.local_id = u.user_id)
"""
_UNOWNED_TEAMS_FROM: Final = """
FROM "LiteLLM_TeamTable" t
WHERE ($1::text IS NULL OR t.team_alias = $1)
AND NOT EXISTS (SELECT 1 FROM "LiteLLM_SCIMResource" r WHERE r.kind = 'Groups' AND r.local_id = t.team_id)
"""
class _LocalIdRow(BaseModel):
local_id: str
class _CountRow(BaseModel):
total: int
_LOCAL_ID_ROWS: Final = TypeAdapter(tuple[_LocalIdRow, ...])
_COUNT_ROWS: Final = TypeAdapter(tuple[_CountRow, ...])
@dataclass(frozen=True)
class _UnownedPage:
local_ids: tuple[str, ...]
total: int
async def _unowned_page(
prisma_client: PrismaClient,
kind: Literal["Users", "Groups"],
filter_value: str | None,
start_index: int,
page_size: int,
*,
filter_attribute: str = "username",
) -> _UnownedPage:
body: Final = (
_UNOWNED_USERS_FROM.format(predicate=_UNOWNED_USER_PREDICATES[filter_attribute])
if kind == "Users"
else _UNOWNED_TEAMS_FROM
)
column: Final = "u.user_id" if kind == "Users" else "t.team_id"
created_at: Final = "u.created_at" if kind == "Users" else "t.created_at"
page_sql: Final = f"SELECT {column} AS local_id {body} ORDER BY {created_at} DESC LIMIT $2 OFFSET $3"
count_sql: Final = f"SELECT COUNT(*)::int AS total {body}"
async with prisma_client.tx() as tx:
page_rows: Final = await tx.query_raw(page_sql, filter_value, page_size, start_index - 1)
count_rows: Final = await tx.query_raw(count_sql, filter_value)
local_ids: Final = tuple(row.local_id for row in _LOCAL_ID_ROWS.validate_python(page_rows))
totals: Final = _COUNT_ROWS.validate_python(count_rows)
return _UnownedPage(local_ids=local_ids, total=totals[0].total if totals else 0)
_RowT = TypeVar("_RowT")
def _in_page_order(local_ids: Sequence[str], rows: Iterable[_RowT], key: Callable[[_RowT], str]) -> tuple[_RowT, ...]:
by_id: Final = {key(row): row for row in rows}
return tuple(by_id[local_id] for local_id in local_ids if local_id in by_id)
async def _assert_legacy_source_access(
auth: UserAPIKeyAuth | None, kind: Literal["Users", "Groups"], local_id: str
) -> None:
if auth is None:
return
client: Final = await _get_prisma_client_or_raise_exception()
if await _source_owned_ids(client, kind, (local_id,)):
raise HTTPException(403, "This record is owned by a different provisioning source")
async def assert_legacy_team_changes_unowned(
auth: UserAPIKeyAuth | None, current: Sequence[str], proposed: Sequence[str]
) -> None:
if auth is None:
return
changed: Final = tuple(frozenset(current) ^ frozenset(proposed))
if not changed:
return
client: Final = await _get_prisma_client_or_raise_exception()
if await _source_owned_ids(client, "Groups", changed):
raise HTTPException(403, "Team membership is owned by a different provisioning source")
_FOLDED_EMAIL_USERS_SQL: Final = """
SELECT user_id, LOWER(user_email) AS folded_email
FROM "LiteLLM_UserTable"
WHERE LOWER(user_email) = ANY($1::text[])
"""
class _FoldedEmailRow(BaseModel):
user_id: str
folded_email: str
_FOLDED_EMAIL_ROWS: Final = TypeAdapter(tuple[_FoldedEmailRow, ...])
def _pair_key(pair: tuple[str, str]) -> str:
return pair[0]
def _pair_user_id(pair: tuple[str, str]) -> str:
return pair[1]
def _ids_by_key(pairs: Iterable[tuple[str | None, str]]) -> Mapping[str, frozenset[str]]:
keyed: Final = sorted(((key, user_id) for key, user_id in pairs if key), key=_pair_key)
grouped: Final = groupby(keyed, _pair_key)
return MappingProxyType({key: frozenset(map(_pair_user_id, group)) for key, group in grouped})
async def _users_by_folded_email(
prisma_client: PrismaClient, subjects: tuple[str, ...]
) -> Mapping[str, frozenset[str]]:
"""Case-insensitive email match in chunked writer reads; `find_many_in` cannot fold the chunked field."""
folded: Final = tuple(dict.fromkeys(subject.lower() for subject in subjects))
starts: Final = range(0, len(folded), IN_LIST_CHUNK_SIZE)
async def _page(start: int) -> tuple[_FoldedEmailRow, ...]:
async with prisma_client.tx() as tx:
rows: Final = await tx.query_raw(_FOLDED_EMAIL_USERS_SQL, folded[start : start + IN_LIST_CHUNK_SIZE])
return _FOLDED_EMAIL_ROWS.validate_python(rows)
pages: Final = tuple([await _page(start) for start in starts])
return _ids_by_key((row.folded_email, row.user_id) for row in chain.from_iterable(pages))
async def _accounts_named_by_member_values(
values: tuple[str, ...], prisma_client: PrismaClient
) -> Mapping[str, frozenset[str]]:
"""``_accounts_named_by_member_value`` for many values in O(chunks) writer reads: exact user id, exact
stripped SSO id, case-insensitive stripped email."""
users: Final = _table(UserRepository(prisma_client, use_writer=True))
subjects: Final = tuple(dict.fromkeys(value.strip() for value in values))
by_id: Final = _ids_by_key((row.user_id, row.user_id) for row in await find_many_in(users, "user_id", values))
by_sso: Final = _ids_by_key(
(row.sso_user_id, row.user_id) for row in await find_many_in(users, "sso_user_id", subjects)
)
by_email: Final = await _users_by_folded_email(prisma_client, subjects)
empty: Final = frozenset[str]()
return MappingProxyType(
{
value: by_id.get(value, empty)
| by_sso.get(value.strip(), empty)
| by_email.get(value.strip().lower(), empty)
for value in values
}
)
async def _assert_legacy_members_unowned(auth: UserAPIKeyAuth | None, members: Sequence[SCIMMember]) -> None:
"""403 for a legacy (non-source) write naming a source-owned subject by SCIM id, local id, SSO id or email."""
if auth is None or not members:
return
client: Final = await _get_prisma_client_or_raise_exception()
values: Final = tuple(dict.fromkeys(_member_value(member) for member in members))
owned_directly: Final = await _source_owned_ids(client, "Users", values)
direct: Final = next((value for value in values if value in owned_directly), None)
if direct is not None:
raise _owned_member_error(direct)
named: Final = await _accounts_named_by_member_values(values, client)
alias_ids: Final = tuple(frozenset(chain.from_iterable(named.values())))
owned_by_alias: Final = await _source_owned_ids(client, "Users", alias_ids)
aliased: Final = next((value for value, ids in named.items() if not owned_by_alias.isdisjoint(ids)), None)
if aliased is not None:
raise _owned_member_error(aliased)
def _owned_member_error(value: str) -> HTTPException:
return HTTPException(
403, f"Group member '{value}' is owned by a different provisioning source and cannot be added here"
)
def _patched_members(op: SCIMPatchOperation) -> tuple[SCIMMember, ...]:
"""The members named by a ``members`` patch operation, from its value or its path filter."""
if op.value is not None:
return _parse_member_entries(op.value)
return tuple(SCIMMember(value=member_id) for member_id in _extract_ids_from_path_filter(op.path, "members"))
def _members_a_patch_admits(patch_ops: SCIMPatchOp) -> tuple[SCIMMember, ...]:
"""The members an ``add`` or ``replace`` operation on ``members`` would put on the roster."""
roster_ops: Final = tuple(
op for op in patch_ops.Operations if op.op != "remove" and (op.path or "").lower().startswith("members")
)
return tuple(chain.from_iterable(_patched_members(op) for op in roster_ops))
async def _scim_groups_from_team_ids(
prisma_client: "PrismaClient | _GroupWriteDatabase", team_ids: list[str]
) -> list[SCIMUserGroup]:
@ -458,6 +697,7 @@ async def _scim_groups_from_team_ids(
team's alias so admin-group matching by display name works the same way it
does on PUT (where SCIM groups carry display names natively).
"""
source_owned: Final = await _source_owned_ids(prisma_client, "Groups", tuple(team_ids))
teams: Final = [
await _table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id}) for team_id in team_ids
]
@ -467,6 +707,7 @@ async def _scim_groups_from_team_ids(
display=team.team_alias if team is not None else None,
)
for team_id, team in zip(team_ids, teams)
if team_id not in source_owned
]
@ -490,7 +731,11 @@ async def write_scim_member_roles(
if admin_group is None:
return
default_role: Final = _default_scim_user_role()
for user_id in user_ids:
candidates: Final = tuple(user_ids)
source_owned: Final = await _source_owned_ids(prisma_client, "Users", candidates)
for user_id in candidates:
if user_id in source_owned:
continue
user = await _table(UserRepository(prisma_client)).find_unique(where={"user_id": user_id})
if user is None:
continue
@ -1200,6 +1445,9 @@ def _get_resource_types(base_url: str = "/scim/v2") -> Sequence[SCIMResourceType
name="User",
description="User Account",
endpoint="/Users",
schemaExtensions=[ # mutable-ok: SCIMResourceType list contract
SCIMSchemaExtension(schema_=SCIM_AGENT_USER_SCHEMA, required=False)
],
schema_="urn:ietf:params:scim:schemas:core:2.0:User",
meta={
"location": f"{base_url}/ResourceTypes/User",
@ -1425,6 +1673,14 @@ def _get_schemas() -> Sequence[SCIMSchema]:
"resourceType": "Schema",
},
),
SCIMSchema(
id=SCIM_AGENT_USER_SCHEMA,
name="LiteLLMAgentUser",
description="Entra agent-user identity, enabled through a trusted provisioning source",
attributes=[ # mutable-ok: SCIMSchema list contract
SCIMSchemaAttribute(name="identityParentId", type="string", required=True, mutability="immutable")
],
),
]
@ -1600,10 +1856,14 @@ async def get_users(
startIndex: int = Query(1, ge=1),
count: int = Query(10, ge=0),
filter: str | None = Query(None),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Get a list of users according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.list("Users", startIndex, count, filter)
page_size: Final = min(count, SCIM_MAX_PAGE_SIZE)
verbose_proxy_logger.debug(
"SCIM GET USERS request: startIndex=%s count=%s filter=%s",
@ -1613,33 +1873,7 @@ async def get_users(
)
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
# Parse filter if provided (basic support)
where_conditions: Final[dict[str, object]] = {}
if filter:
# 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)
if parsed_filter:
filter_attribute, filter_value = parsed_filter
if filter_attribute == "username":
where_conditions["OR"] = [
{"user_email": filter_value},
{"user_id": filter_value},
]
elif filter_attribute == "emails.value":
where_conditions["user_email"] = filter_value
# Get users from database
users: Final[Sequence[LiteLLM_UserTable]] = await _table(UserRepository(prisma_client)).find_many(
where=where_conditions,
skip=(startIndex - 1),
take=page_size,
order={"created_at": "desc"},
)
# Get total count for pagination
total_count: Final = await _table(UserRepository(prisma_client)).count(where=where_conditions)
users, total_count = await _legacy_users_page(prisma_client, auth, filter, startIndex, page_size)
# Convert to SCIM format
scim_users: Final[list[SCIMUser]] = []
@ -1666,10 +1900,15 @@ async def get_users(
)
async def get_user(
user_id: str = Path(..., title="User ID"),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Get a single user by ID according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.get("Users", user_id)
await _assert_legacy_source_access(auth, "Users", user_id)
verbose_proxy_logger.debug("SCIM GET USER request for user_id=%s", user_id)
try:
user: Final = await _check_user_exists(user_id)
@ -1690,16 +1929,23 @@ async def get_user(
)
async def create_user(
user: SCIMUser = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Create a user according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.create_user(user)
if user.agent_user is not None:
raise HTTPException(400, "Configure an Entra provisioning source before provisioning agent-users")
try:
verbose_proxy_logger.debug("SCIM CREATE USER request: %s", user.model_dump())
prisma_client: Final = await _get_prisma_client_or_raise_exception()
# Extract data from SCIM user
user_data: Final = _extract_scim_user_data(user)
await assert_legacy_team_changes_unowned(auth, (), user_data["teams"] or ())
# Check if user already exists
if user.userName:
@ -1739,6 +1985,7 @@ async def create_user(
prisma_client=prisma_client,
new_user_request=new_user_request,
admin_group=admin_group,
auth=auth,
)
if existing_user_scim:
@ -1830,10 +2077,17 @@ async def finish_provisioned_user_update(user_id: str, tokens: tuple[str, ...])
async def update_user(
user_id: str = Path(..., title="User ID"),
user: SCIMUser = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Update a user according to SCIM v2 protocol (full replacement)
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.update_user(user_id, user)
if user.agent_user is not None:
raise HTTPException(400, "Configure an Entra provisioning source before provisioning agent-users")
await _assert_legacy_source_access(auth, "Users", user_id)
verbose_proxy_logger.debug(
"SCIM PUT USER request for user_id=%s: %s",
user_id,
@ -1851,6 +2105,7 @@ async def update_user(
client_set_active: Final = "active" in user.model_fields_set
metadata: Final = replacement.metadata
target_teams: Final = replacement.teams
await assert_legacy_team_changes_unowned(auth, existing_user.teams or (), target_teams or ())
await _handle_team_membership_changes(
user_id=user_id,
existing_teams=existing_user.teams,
@ -1964,10 +2219,16 @@ async def write_scim_group_deletion(tx: "Prisma", group_id: str, admin_group: st
)
async def delete_user(
user_id: str = Path(..., title="User ID"),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Delete a user according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
await service.delete("Users", user_id)
return Response(status_code=204)
await _assert_legacy_source_access(auth, "Users", user_id)
verbose_proxy_logger.debug("SCIM DELETE USER request for user_id=%s", user_id)
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
@ -1986,7 +2247,9 @@ async def delete_user(
response_model=tuple[SCIMPlaceholder, ...],
dependencies=(Depends(user_api_key_auth),),
)
async def list_placeholders() -> tuple[SCIMPlaceholder, ...]:
async def list_placeholders(
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
) -> tuple[SCIMPlaceholder, ...]:
"""
List user rows whose id is another account's SSO identity or email.
@ -1998,6 +2261,8 @@ async def list_placeholders() -> tuple[SCIMPlaceholder, ...]:
its own or owns virtual keys is left out: someone uses that account.
"""
try:
if await _agent_provisioning_service(auth) is not None:
raise HTTPException(403, "Provisioning source tokens cannot access global placeholder maintenance")
prisma_client: Final = await _get_prisma_client_or_raise_exception()
async with prisma_client.tx() as tx:
return await UserRepository(prisma_client).find_shadowing_placeholders(tx)
@ -2026,6 +2291,7 @@ def _placeholder_rejection(placeholder: LiteLLM_UserTable, resolved: tuple[str,
)
async def merge_placeholder(
user_id: str = Path(..., title="User ID"),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
) -> SCIMPlaceholderMergeResult:
"""
Fold a placeholder user into the one account its id names by SSO identity or email.
@ -2036,6 +2302,8 @@ async def merge_placeholder(
has an SSO identity of its own, owns virtual keys, or names no account or several.
"""
try:
if await _agent_provisioning_service(auth) is not None:
raise HTTPException(403, "Provisioning source tokens cannot access global placeholder maintenance")
prisma_client: Final = await _get_prisma_client_or_raise_exception()
placeholder: Final = await _check_user_exists(user_id)
resolved: Final = tuple(
@ -2433,10 +2701,19 @@ async def patch_team_membership(
async def patch_user(
user_id: str = Path(..., title="User ID"),
patch_ops: SCIMPatchOp = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Patch a user according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.update_user(user_id, patch_ops)
from litellm.proxy.management_endpoints.scim.agent_provisioning import patch_changes_identity
if patch_changes_identity(patch_ops):
raise HTTPException(400, "SCIM PATCH cannot change a subject's identity classification")
await _assert_legacy_source_access(auth, "Users", user_id)
verbose_proxy_logger.debug(
"SCIM PATCH USER request for user_id=%s: %s",
user_id,
@ -2454,6 +2731,8 @@ async def patch_user(
patch_ops=patch_ops,
)
await assert_legacy_team_changes_unowned(auth, existing_user.teams or (), tuple(final_team_set))
patched_metadata: Final = update_data.get("metadata")
new_active: Final = _scim_active_value(patched_metadata if isinstance(patched_metadata, Mapping) else None)
@ -2467,7 +2746,7 @@ async def patch_user(
update_data["teams"] = list(final_team_set)
admin_group: Final = await _get_scim_admin_group()
if admin_group is not None:
if admin_group is not None and not await _source_owned_ids(prisma_client, "Users", (user_id,)):
update_data["user_role"] = _resolve_scim_user_role(
await _scim_groups_from_team_ids(prisma_client, list(final_team_set)),
admin_group,
@ -2499,10 +2778,63 @@ async def patch_user(
raise handle_exception_on_proxy(e)
def _legacy_user_filter(filter: str | None) -> tuple[str, str | None]:
# 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) if filter else None
if parsed_filter and parsed_filter[0] in _UNOWNED_USER_PREDICATES:
return parsed_filter
return ("username", None)
async def _legacy_users_page(
prisma_client: PrismaClient, auth: UserAPIKeyAuth | None, filter: str | None, start_index: int, page_size: int
) -> tuple[Sequence[LiteLLM_UserTable], int]:
filter_attribute, filter_value = _legacy_user_filter(filter)
table: Final = _table(UserRepository(prisma_client, use_writer=auth is not None))
if auth is not None:
page: Final = await _unowned_page(
prisma_client, "Users", filter_value, start_index, page_size, filter_attribute=filter_attribute
)
rows: Final = await find_many_in(table, "user_id", page.local_ids)
return _in_page_order(page.local_ids, rows, lambda user: user.user_id), page.total
where_conditions: Final[dict[str, object]] = {}
if filter_value is not None and filter_attribute == "username":
where_conditions["OR"] = [{"user_email": filter_value}, {"user_id": filter_value}]
elif filter_value is not None:
where_conditions["user_email"] = filter_value
users: Final = await table.find_many(
where=where_conditions, skip=start_index - 1, take=page_size, order={"created_at": "desc"}
)
return users, await table.count(where=where_conditions)
class _TeamWhereConditions(TypedDict, total=False):
"""The team columns SCIM GET /Groups can filter on, as Prisma where-conditions."""
team_alias: str
team_alias: ReadOnly[str]
def _legacy_team_alias(filter: str | None) -> str | None:
# Very basic filter support - only handling displayName eq
return filter.split("displayName eq ")[1].strip("\"'") if filter and "displayName eq" in filter else None
async def _legacy_teams_page(
prisma_client: PrismaClient, auth: UserAPIKeyAuth | None, filter: str | None, start_index: int, page_size: int
) -> tuple[Sequence[LiteLLM_TeamTable], int]:
team_alias: Final = _legacy_team_alias(filter)
table: Final = _table(TeamRepository(prisma_client, use_writer=auth is not None))
if auth is not None:
page: Final = await _unowned_page(prisma_client, "Groups", team_alias, start_index, page_size)
rows: Final = await find_many_in(table, "team_id", page.local_ids)
return _in_page_order(page.local_ids, rows, lambda team: team.team_id), page.total
where_conditions: Final[_TeamWhereConditions] = {} if team_alias is None else {"team_alias": team_alias}
teams: Final = await table.find_many(
where=where_conditions, skip=start_index - 1, take=page_size, order={"created_at": "desc"}
)
return teams, await table.count(where=where_conditions)
# Group Endpoints
@ -2516,10 +2848,14 @@ async def get_groups(
startIndex: int = Query(1, ge=1),
count: int = Query(10, ge=0),
filter: str | None = Query(None),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Get a list of groups according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.list("Groups", startIndex, count, filter)
page_size: Final = min(count, SCIM_MAX_PAGE_SIZE)
verbose_proxy_logger.debug(
"SCIM GET GROUPS request: startIndex=%s count=%s filter=%s",
@ -2529,24 +2865,7 @@ async def get_groups(
)
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
# Parse filter if provided (basic support)
where_conditions: Final[_TeamWhereConditions] = {}
if filter:
# Very basic filter support - only handling displayName eq
if "displayName eq" in filter:
team_alias = filter.split("displayName eq ")[1].strip("\"'")
where_conditions["team_alias"] = team_alias
# Get teams from database
teams: Final = await _table(TeamRepository(prisma_client)).find_many(
where=where_conditions,
skip=(startIndex - 1),
take=page_size,
order={"created_at": "desc"},
)
# Get total count for pagination
total_count: Final = await _table(TeamRepository(prisma_client)).count(where=where_conditions)
teams, total_count = await _legacy_teams_page(prisma_client, auth, filter, startIndex, page_size)
# Convert to SCIM format
scim_groups: Final[list[SCIMGroup]] = []
@ -2594,10 +2913,15 @@ async def get_groups(
)
async def get_group(
group_id: str = Path(..., title="Group ID"),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Get a single group by ID according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.get("Groups", group_id)
await _assert_legacy_source_access(auth, "Groups", group_id)
verbose_proxy_logger.debug("SCIM GET GROUP request for group_id=%s", group_id)
try:
team: Final = await _check_team_exists(group_id)
@ -2787,10 +3111,14 @@ def _group_replacement_data(existing: LiteLLM_TeamTable, group: SCIMGroup) -> Ma
)
async def create_group(
group: SCIMGroup = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Create a group according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.create_group(group)
verbose_proxy_logger.debug(
"SCIM CREATE GROUP request: %s",
group.model_dump(),
@ -2810,6 +3138,7 @@ async def create_group(
detail={"error": f"Group already exists with ID: {team_id}"},
)
await _assert_legacy_members_unowned(auth, group.members or ())
# Extract and validate group members (all users must exist)
member_result: Final = await _extract_group_member_ids(group)
members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_result.all_member_ids]
@ -2842,10 +3171,15 @@ async def create_group(
async def update_group(
group_id: str = Path(..., title="Group ID"),
group: SCIMGroup = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Update a group according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.update_group(group_id, group)
await _assert_legacy_source_access(auth, "Groups", group_id)
verbose_proxy_logger.debug(
"SCIM PUT GROUP request for group_id=%s: %s",
group_id,
@ -2854,6 +3188,7 @@ async def update_group(
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
existing_team: Final = await _check_team_exists(group_id)
await _assert_legacy_members_unowned(auth, group.members or ())
# Extract and validate group members (all users must exist)
member_result: Final = await _extract_group_member_ids(group)
@ -2904,10 +3239,16 @@ async def update_group(
)
async def delete_group(
group_id: str = Path(..., title="Group ID"),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Delete a group according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
await service.delete("Groups", group_id)
return Response(status_code=204)
await _assert_legacy_source_access(auth, "Groups", group_id)
verbose_proxy_logger.debug("SCIM DELETE GROUP request for group_id=%s", group_id)
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
@ -2971,13 +3312,7 @@ async def _process_group_patch_operations(
metadata["externalId"] = str(value)
elif path.startswith("members"):
# Handle member operations
patched_members = (
_parse_member_entries(value)
if value is not None
else tuple(
SCIMMember(value=member_id) for member_id in _extract_ids_from_path_filter(op.path, "members")
)
)
patched_members = _patched_members(op)
if op_type == "remove":
final_members = final_members - await _member_ids_to_drop(
@ -3087,10 +3422,15 @@ async def _handle_group_membership_changes(group_id: str, current_members: set[s
async def patch_group(
group_id: str = Path(..., title="Group ID"),
patch_ops: SCIMPatchOp = Body(...),
auth: Annotated[UserAPIKeyAuth | None, Depends(user_api_key_auth)] = None,
):
"""
Patch a group according to SCIM v2 protocol
"""
service: Final = await _agent_provisioning_service(auth)
if service is not None:
return await service.update_group(group_id, patch_ops)
await _assert_legacy_source_access(auth, "Groups", group_id)
verbose_proxy_logger.debug(
"SCIM PATCH GROUP request for group_id=%s: %s",
group_id,
@ -3100,6 +3440,7 @@ async def patch_group(
try:
prisma_client: Final = await _get_prisma_client_or_raise_exception()
existing_team: Final = await _check_team_exists(group_id)
await _assert_legacy_members_unowned(auth, _members_a_patch_admits(patch_ops))
# Process patch operations
update_data, final_members, replace_target = await _process_group_patch_operations(

View file

@ -4,7 +4,7 @@ from typing import Annotated, Final
from uuid import uuid4
from fastapi import APIRouter, Depends, HTTPException
from prisma import Json
from prisma import Json, Prisma
from prisma.types import (
LiteLLM_SCIMSourceCreateInput,
LiteLLM_SCIMSourceOrderByInput,
@ -62,6 +62,10 @@ async def list_sources(auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth
return tuple(source_response(source) for source in sources)
async def lock_provisioning_token(tx: Prisma, token_hash: str) -> None:
await tx.execute_raw('SELECT 1 FROM "LiteLLM_VerificationToken" WHERE token = $1 FOR UPDATE', token_hash)
@router.post("", response_model=SCIMSourceResponse, status_code=201)
async def create_source(
data: SCIMSourceCreate, auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)]
@ -71,6 +75,7 @@ async def create_source(
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:
await lock_provisioning_token(tx, token_hash)
key: Final = await VerificationTokenRepository(SimpleNamespace(db=tx)).table.find_unique(
where=LiteLLM_VerificationTokenWhereUniqueInput(token=token_hash)
)

View file

@ -25,6 +25,7 @@ from litellm.repositories.prisma_protocols import DatabaseClient, TableActions
from litellm.repositories.table_repositories import MCPServerRepository
if TYPE_CHECKING:
from prisma import Prisma
from prisma import models as prisma_models
from litellm.proxy._types import (
@ -84,6 +85,8 @@ async def prepare_object_permission_upsert(
new_object_permission: Mapping[str, object],
existing_object_permission_id: str | None,
prisma_client: PrismaClient,
*,
tx: "Prisma | None" = None,
) -> ObjectPermissionUpsert:
"""
Read-and-merge half of an object permission upsert; performs no writes.
@ -101,7 +104,10 @@ async def prepare_object_permission_upsert(
update cannot leave permission changes live.
"""
object_permission_id: Final = existing_object_permission_id or str(uuid.uuid4())
existing_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
permission_table: Final = (
tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table
)
existing_object_permission: Final = await permission_table.find_unique(
where={"object_permission_id": object_permission_id},
)
existing_fields: Final[dict[str, object]] = (
@ -134,6 +140,8 @@ async def handle_update_object_permission_common(
data_json: dict,
existing_object_permission_id: str | None,
prisma_client: PrismaClient | None,
*,
tx: "Prisma | None" = None,
) -> str | None:
"""
Common logic for handling object permission updates across organizations, teams, and keys.
@ -170,8 +178,12 @@ async def handle_update_object_permission_common(
new_object_permission=new_object_permission if isinstance(new_object_permission, dict) else {},
existing_object_permission_id=existing_object_permission_id,
prisma_client=prisma_client,
tx=tx,
)
created_object_permission_row: Final = await ObjectPermissionRepository(prisma_client).table.upsert(
permission_table: Final = (
tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table
)
created_object_permission_row: Final = await permission_table.upsert(
where={"object_permission_id": upsert.object_permission_id},
data={
"create": upsert.record,

View file

@ -351,6 +351,9 @@ class AgentResponse(BaseModel):
budget_id: str | None = None
lifetime_budget_spend: float = 0.0
litellm_budget_table: AgentBudgetState | None = None
directory_active: bool = True
directory_access_group_ids: tuple[str, ...] | None = None
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
enabled: bool = True

View file

@ -13,6 +13,7 @@ class EntraIdentityConfig(BaseModel):
provider: Literal["microsoft_entra"]
tenant_id: str
client_id: str
provisioning_source_id: str | None = None
service_principal_id: str | None = None
required_roles: tuple[str, ...] = ()
required_scopes: tuple[str, ...] = Field(
@ -38,6 +39,7 @@ class AgentIdentityBinding(BaseModel):
provider: Literal["microsoft_entra"]
tenant_id: str
client_id: str
provisioning_source_id: str | None = None
service_principal_id: str | None = None
issuer: str
required_roles: tuple[str, ...] = ()
@ -65,7 +67,7 @@ class AgentBudgetState(BaseModel):
class AgentSubject(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["application", "delegated_subject"]
kind: Literal["application", "delegated_subject", "agent_user"]
oid: str
mode: Literal["autonomous", "delegated"]
@ -105,8 +107,21 @@ class MicrosoftInteractiveSubject(BaseModel):
class ManagedAgentIdentityStatus(BaseModel):
directory_active: bool = True
directory_access_group_ids: tuple[str, ...] | None = None
identity: AgentIdentityBinding | None = None
identity_managed: bool = False
enabled: bool = True
execution_mode: AgentExecutionMode = "autonomous"
last_authenticated_at: datetime | None = None
class VerifiedAgentSubject(BaseModel):
model_config = ConfigDict(frozen=True)
issuer: str
tenant_id: str
oid: str
agent_id: str
parent_client_id: str
scim_resource_id: str

View file

@ -114,6 +114,42 @@ async def test_access_groups_cap_agent_servers_without_granting_new_ones(
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
"directory_groups,directory_servers,expected",
[
(None, (), ["slack"]),
((), (), []),
(("directory",), ("slack", "linear"), ["slack"]),
(("directory",), ("linear",), []),
],
)
async def test_directory_and_manual_groups_independently_cap_mcp_access(
monkeypatch: pytest.MonkeyPatch,
directory_groups: tuple[str, ...] | None,
directory_servers: tuple[str, ...],
expected: list[str],
) -> None:
from litellm.models.access_group import LiteLLM_AccessGroupTable
manual: Final = LiteLLM_AccessGroupTable(
access_group_id="manual", access_group_name="Manual", access_mcp_server_ids=["slack"]
)
directory: Final = LiteLLM_AccessGroupTable(
access_group_id="directory", access_group_name="Directory", access_mcp_server_ids=list(directory_servers)
)
lookup: Final = AsyncMock(side_effect=[manual, directory] if directory_groups else [manual])
monkeypatch.setattr(auth_checks, "get_access_object", lookup)
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
update={"access_group_ids": ["manual"], "directory_access_group_ids": directory_groups}
)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == expected
assert lookup.await_count == (2 if directory_groups else 1)
assert all(call.kwargs["check_db_only"] for call in lookup.await_args_list)
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(

View file

@ -662,6 +662,8 @@ class TestAgentRequestHandler:
[
({}, True),
({"enabled": False}, False),
({"directory_active": False}, False),
({"directory_access_group_ids": ()}, False),
],
)
async def test_managed_invocation_requires_local_and_directory_admission(

View file

@ -334,6 +334,16 @@ async def test_invocation_cannot_bypass_missing_policy_permission_or_invalid_pri
assert auth.agent_invocation_cost is None
def test_native_directory_agent_cannot_authenticate_with_a_virtual_key() -> None:
policy: Final = agent(
identity=BINDING.model_copy(update={"provisioning_source_id": "source"}), execution_mode="autonomous"
)
failure: Final = actor_admission_failure(policy, None)
assert isinstance(failure, AgentIdentityFailure)
assert failure.code == "identity_denied"
assert "bound identity provider token" in failure.message
@pytest.mark.asyncio
async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None:
database: Final = MagicMock()

View file

@ -1,3 +1,4 @@
import json
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Final
@ -280,6 +281,209 @@ async def test_retired_client_cannot_fall_back_to_ordinary_user_authentication()
assert "retired" in result.message
def native_store():
from prisma.models import LiteLLM_SCIMResource, LiteLLM_SCIMSource
from litellm.repositories.table_repositories import SCIMResourceRepository, SCIMSourceRepository
from litellm.types.proxy.management_endpoints.scim_agent_provisioning import SCIM_AGENT_USER_SCHEMA
now: Final = datetime.now(timezone.utc)
native: Final = LiteLLM_VerifiedSubject(
subject_id="native-subject",
kind="agent_user",
issuer=ISSUER,
tenant_id=TENANT,
oid=HUMAN,
agent_id="agent-one",
parent_client_id=CLIENT,
scim_resource_id="directory-user",
verified_via="scim",
verified_at=now,
)
store, agents, identities, humans = setup_store(human=native)
binding: Final = BINDING.model_copy(update={"provisioning_source_id": "source", "service_principal_id": None})
agents.find_unique.return_value = stored_agent(identity=binding, execution_mode="autonomous")
identities.find_unique.return_value = binding
source: Final = LiteLLM_SCIMSource(
source_id="source",
display_name="Directory",
tenant_id=TENANT,
key_hash="hash",
enabled=True,
group_mappings=json.dumps([{"external_group_id": PRINCIPAL, "access_group_ids": ["read"]}]),
created_at=now,
updated_at=now,
)
resource: Final = LiteLLM_SCIMResource(
id="directory-user",
source_id="source",
kind="Users",
external_id=HUMAN,
local_id="agent-one",
display_name="Native",
member_ids=[],
document=json.dumps(
{
"schemas": [],
"userName": "native@example.com",
"externalId": HUMAN,
SCIM_AGENT_USER_SCHEMA: {"identityParentId": CLIENT},
}
),
active=True,
deleted=False,
created_at=now,
updated_at=now,
)
group: Final = resource.model_copy(update={"id": "directory-group", "kind": "Groups", "external_id": PRINCIPAL})
sources: Final = AsyncMock()
resources: Final = AsyncMock()
sources.find_unique.return_value = source
resources.find_many.side_effect = lambda **query: [resource] if query["where"]["kind"] == "Users" else [group]
client: Final = SimpleNamespace(db=SimpleNamespace(litellm_scimsource=sources, litellm_scimresource=resources))
return (
AgentIdentityStore(
store.agents,
store.identities,
store.humans,
sources=SCIMSourceRepository(client),
resources=SCIMResourceRepository(client),
),
sources,
resources,
humans,
native,
resource,
)
@pytest.mark.asyncio
async def test_native_subject_is_autonomous_and_directory_grants_are_separate() -> None:
store, _, _, _, _, _ = native_store()
context: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
assert isinstance(context, ManagedAgentContext)
assert context.mode == "autonomous"
assert context.user_id is None
assert context.subject_oid == HUMAN
agent: Final = await store.agent("agent-one")
assert isinstance(agent, AgentResponse)
assert agent.directory_active is True
assert agent.directory_access_group_ids == ("read",)
assert agent.access_group_ids is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"field,value",
[
("kind", "human"),
("verified_via", "sso_interactive"),
("agent_id", "foreign-agent"),
("parent_client_id", PRINCIPAL),
("scim_resource_id", "foreign-resource"),
],
)
async def test_directory_policy_rejects_mismatched_subject_ownership(field: str, value: str) -> None:
store, _, _, humans, native, _ = native_store()
humans.find_unique.return_value = native.model_copy(update={field: value})
result: Final = await store.agent("agent-one")
assert isinstance(result, AgentIdentityFailure)
assert "subject" in result.message
@pytest.mark.asyncio
@pytest.mark.parametrize("state", ["source-disabled", "resource-disabled", "resource-deleted", "no-groups"])
async def test_native_revocation_cannot_fall_back_to_human_or_unrestricted_agent(state: str) -> None:
store, sources, resources, _, _, resource = native_store()
if state == "source-disabled":
sources.find_unique.return_value = sources.find_unique.return_value.model_copy(update={"enabled": False})
elif state == "no-groups":
resources.find_many.side_effect = lambda **query: [resource] if query["where"]["kind"] == "Users" else []
else:
resources.find_many.side_effect = None
resources.find_many.return_value = [
resource.model_copy(
update={
"active": state != "resource-disabled",
"deleted": state == "resource-deleted",
}
)
]
result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
assert isinstance(result, AgentIdentityFailure)
assert result.code == "identity_denied"
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", ["both", "second-only"])
async def test_multiple_directory_groups_union_and_removal_preserves_remaining_grants(groups: str) -> None:
store, sources, resources, _, _, resource = native_store()
sources.find_unique.return_value = sources.find_unique.return_value.model_copy(
update={
"group_mappings": [
{"external_group_id": PRINCIPAL, "access_group_ids": ["read"]},
{"external_group_id": TENANT, "access_group_ids": ["write"]},
]
}
)
first: Final = resource.model_copy(update={"id": "group-one", "kind": "Groups", "external_id": PRINCIPAL})
second: Final = resource.model_copy(update={"id": "group-two", "kind": "Groups", "external_id": TENANT})
rows: Final = [first, second] if groups == "both" else [second]
resources.find_many.side_effect = lambda **query: [resource] if query["where"]["kind"] == "Users" else rows
result: Final = await store.agent("agent-one")
assert isinstance(result, AgentResponse)
assert result.directory_access_group_ids == (("read", "write") if groups == "both" else ("write",))
assert result.access_group_ids is None
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["source-unavailable", "wrong-tenant", "missing-subject", "subject-unavailable"])
async def test_native_directory_policy_fails_closed_when_correspondence_cannot_be_proven(failure: str) -> None:
store, sources, _, humans, _, _ = native_store()
if failure == "source-unavailable":
sources.find_unique.side_effect = RuntimeError("writer unavailable")
elif failure == "wrong-tenant":
sources.find_unique.return_value = sources.find_unique.return_value.model_copy(update={"tenant_id": HUMAN})
elif failure == "missing-subject":
humans.find_unique.return_value = None
else:
humans.find_unique.side_effect = RuntimeError("writer unavailable")
result: Final = await store.agent("agent-one")
assert isinstance(result, AgentIdentityFailure)
@pytest.mark.asyncio
async def test_verified_subject_cannot_switch_to_an_unrelated_registered_parent() -> None:
store, _, _, humans, native, _ = native_store()
humans.find_unique.return_value = native.model_copy(update={"parent_client_id": PRINCIPAL})
result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
assert isinstance(result, AgentIdentityFailure)
assert "does not match" in result.message
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["human", "agent_user", "untrusted", "missing"])
async def test_human_lookup_requires_trusted_interactive_enrollment(kind: str) -> None:
store, _, _, humans = setup_store()
native: Final = native_store()[4]
humans.find_unique.return_value = (
None
if kind == "missing"
else native.model_copy(update={"kind": "human", "verified_via": "sso_interactive", "user_id": "local-human"})
if kind == "human"
else native.model_copy(update={"kind": "human", "verified_via": "untrusted"})
if kind == "untrusted"
else native
)
result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
if kind == "human":
assert isinstance(result, ManagedAgentContext)
assert result.mode == "delegated"
assert result.user_id == "local-human"
else:
assert isinstance(result, AgentIdentityFailure)
@pytest.mark.asyncio
async def test_missing_revision_cannot_create_entra_authentication_evidence() -> None:
store, _, identities, _ = setup_store()
@ -318,6 +522,15 @@ async def test_retired_binding_denies_and_history_outage_cannot_become_legacy_fa
)
@pytest.mark.asyncio
async def test_missing_directory_repositories_cannot_admit_a_provisioned_agent() -> None:
store, _, _, _ = setup_store()
native: Final = stored_agent(identity=BINDING.model_copy(update={"provisioning_source_id": "source"}))
result: Final = await store.directory_policy(native)
assert isinstance(result, AgentIdentityFailure)
assert result.code == "policy_unavailable"
@pytest.mark.asyncio
async def test_new_binding_is_enforced_after_another_worker_commits_it() -> None:
_, agents, identities, humans = setup_store()
@ -436,7 +649,7 @@ async def test_resolver_preserves_unconfigured_and_unrelated_authentication() ->
@pytest.mark.asyncio
@pytest.mark.parametrize("registered", [True, False])
async def test_application_and_unregistered_clients_do_not_depend_on_human_subject_storage(registered: bool) -> None:
async def test_subject_outage_only_blocks_clients_that_need_directory_classification(registered: bool) -> None:
store, _, identities, humans = setup_store()
identities.find_unique.return_value = BINDING if registered else None
humans.find_unique.side_effect = RuntimeError("subject database unavailable")
@ -445,6 +658,22 @@ async def test_application_and_unregistered_clients_do_not_depend_on_human_subje
assert isinstance(result, ManagedAgentContext)
assert result.mode == "autonomous"
assert result.user_id is None
humans.find_unique.assert_not_awaited()
else:
assert result is None
humans.find_unique.assert_not_awaited()
assert isinstance(result, AgentIdentityFailure)
assert result.code == "policy_unavailable"
humans.find_unique.assert_awaited_once()
@pytest.mark.asyncio
async def test_directory_group_guid_case_does_not_remove_agent_grants() -> None:
store, sources, resources, _, _, resource = native_store()
guid: Final = "abcdefab-abcd-4abc-8abc-abcdefabcdef"
sources.find_unique.return_value = sources.find_unique.return_value.model_copy(
update={"group_mappings": [{"external_group_id": guid, "access_group_ids": ["read"]}]}
)
group: Final = resource.model_copy(update={"id": "group", "kind": "Groups", "external_id": guid.upper()})
resources.find_many.side_effect = lambda **query: [resource] if query["where"]["kind"] == "Users" else [group]
result: Final = await store.agent("agent-one")
assert isinstance(result, AgentResponse)
assert result.directory_access_group_ids == ("read",)

View file

@ -170,6 +170,101 @@ def test_each_application_binding_records_its_history_atomically() -> None:
assert replacement["retired_identities"]["create"]["client_id"] == HUMAN
def test_native_agent_user_uses_proven_subject_without_optional_facets() -> None:
from litellm.types.proxy.agent_identity import VerifiedAgentSubject
subject: Final = VerifiedAgentSubject(
issuer=ISSUER,
tenant_id=TENANT,
oid=HUMAN,
agent_id=BINDING.agent_id,
parent_client_id=CLIENT,
scim_resource_id="scim-subject",
)
result: Final = classify_agent_subject(
BINDING,
claims(oid=HUMAN, scp="user_impersonation"),
"autonomous",
native_subject=subject,
)
assert result == AgentSubject(kind="agent_user", oid=HUMAN, mode="autonomous")
@pytest.mark.parametrize(
"overrides",
[
{"oid": PRINCIPAL},
{"azp": HUMAN},
{"tid": HUMAN},
{"scp": "unrelated"},
{"scp": None},
{"idtyp": "app"},
],
)
def test_native_subject_binding_rejects_other_subjects_parents_and_scopes(overrides: dict[str, object]) -> None:
from litellm.types.proxy.agent_identity import VerifiedAgentSubject
subject: Final = VerifiedAgentSubject(
issuer=ISSUER,
tenant_id=TENANT,
oid=HUMAN,
agent_id=BINDING.agent_id,
parent_client_id=CLIENT,
scim_resource_id="scim-subject",
)
result: Final = classify_agent_subject(
BINDING,
claims(**{"oid": HUMAN, "scp": "user_impersonation", **overrides}),
"both",
native_subject=subject,
)
assert isinstance(result, AgentIdentityFailure)
@pytest.mark.parametrize("change", [None, {"client_id": HUMAN}, {"provisioning_source_id": "another-source"}])
def test_directory_owned_identity_cannot_be_unbound_or_reassigned(change: dict[str, str] | None) -> None:
native_binding: Final = BINDING.model_copy(update={"provisioning_source_id": "source"})
existing: Final = AgentResponse(
agent_id="agent-one",
agent_name="Native",
agent_card_params={},
identity=native_binding,
identity_managed=True,
execution_mode="autonomous",
)
identity: Final = (
None
if change is None
else {
"provider": "microsoft_entra",
"tenant_id": TENANT,
"client_id": CLIENT,
"provisioning_source_id": "source",
**change,
}
)
result: Final = managed_write_fields({"identity": identity}, existing, "admin")
assert isinstance(result, AgentIdentityFailure)
assert "directory-owned" in result.message
def test_manual_registration_cannot_claim_directory_ownership() -> None:
result: Final = managed_write_fields(
{
"identity": {
"provider": "microsoft_entra",
"tenant_id": TENANT,
"client_id": CLIENT,
"provisioning_source_id": "source",
}
},
None,
"admin",
)
assert isinstance(result, AgentIdentityFailure)
assert "Only SCIM" in result.message
def test_unchanged_binding_preserves_revision_and_authentication_evidence() -> None:
configuration: Final = BINDING.model_dump(
exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
@ -206,6 +301,15 @@ def test_explicit_empty_scope_requirements_can_be_registered_and_preserved(mode:
assert not isinstance(updated, AgentIdentityFailure)
assert updated["execution_mode"] == mode
@pytest.mark.parametrize("mode", ["delegated", "both"])
def test_native_directory_identity_cannot_switch_to_delegated_execution(mode: str) -> None:
agent: Final = managed_agent().model_copy(
update={"identity": BINDING.model_copy(update={"provisioning_source_id": "source"})}
)
result: Final = managed_write_fields({"execution_mode": mode}, agent, "admin")
assert isinstance(result, AgentIdentityFailure)
assert "autonomous mode" in result.message
@pytest.mark.parametrize(
"incoming",
@ -354,3 +458,10 @@ def test_recreated_or_converted_lifetime_budget_gets_a_fresh_allowance(previous_
assert "spend" not in result
assert "create" in result["litellm_budget_table"]
assert "update" not in result["litellm_budget_table"]
def test_directory_binding_cannot_fall_back_to_an_application_token() -> None:
binding: Final = BINDING.model_copy(update={"provisioning_source_id": "source"})
result: Final = classify_agent_subject(binding, claims(), "autonomous")
assert isinstance(result, AgentIdentityFailure)
assert "verified provisioned agent-user" in result.message

View file

@ -951,6 +951,44 @@ async def test_group_patch_failure_does_not_sync_members(case: str, monkeypatch:
tx.litellm_scimresource.update_many.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"kind,operation", [(kind, op) for kind in ("Users", "Groups") for op in ("get", "update", "patch", "delete")]
)
async def test_legacy_scim_token_cannot_access_source_owned_local_record(
kind: str, operation: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from unittest.mock import MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.management_endpoints.scim import scim_v2
service, _, native = provisioning_fixture()
client = MagicMock()
client.writer_db.litellm_scimresource.find_many = AsyncMock(
return_value=[native.model_copy(update={"kind": kind, "local_id": "local-owned"})]
)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=client))
lookup = AsyncMock(side_effect=AssertionError("legacy operation reached the source-owned local record"))
monkeypatch.setattr(scim_v2, "_check_user_exists" if kind == "Users" else "_check_team_exists", lookup)
arguments = {
"user_id" if kind == "Users" else "group_id": "local-owned",
"auth": UserAPIKeyAuth(api_key="legacy-scim"),
}
if operation == "update":
arguments["user" if kind == "Users" else "group"] = (
SCIMUser(schemas=[], userName="changed@example.com")
if kind == "Users"
else SCIMGroup(schemas=[], displayName="Changed")
)
if operation == "patch":
arguments["patch_ops"] = SCIMPatchOp(Operations=[{"op": "replace", "path": "displayName", "value": "Changed"}])
endpoint = getattr(scim_v2, operation + ("_user" if kind == "Users" else "_group"))
with pytest.raises(HTTPException) as denied:
await endpoint(**arguments)
assert denied.value.status_code == 403
lookup.assert_not_awaited()
@pytest.mark.parametrize(

View file

@ -20,6 +20,22 @@ TENANT: Final = "11111111-1111-4111-8111-111111111111"
SUBJECT: Final = "22222222-2222-4222-8222-222222222222"
@pytest.mark.asyncio
async def test_source_ownership_collects_scim_and_local_ids_across_query_batches() -> None:
from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE
client: Final = MagicMock(spec=PrismaClient)
client.writer_db = MagicMock()
client.db = MagicMock()
local_ids: Final = tuple(f"subject-{index}" for index in range(IN_LIST_CHUNK_SIZE + 1))
by_scim_id: Final = SimpleNamespace(id=local_ids[-1], local_id="local-human")
by_local_id: Final = SimpleNamespace(id="scim-human", local_id=local_ids[0])
client.writer_db.litellm_scimresource.find_many = AsyncMock(side_effect=[[], [by_scim_id], [], [by_local_id]])
owned: Final = await scim_v2._source_owned_ids(client, "Users", local_ids)
assert owned == frozenset((local_ids[-1], "local-human", "scim-human", local_ids[0]))
client.db.litellm_scimresource.find_many.assert_not_called()
def human_fixture():
now: Final = datetime.now(timezone.utc)
source: Final = LiteLLM_SCIMSource(

View file

@ -27,9 +27,7 @@ from litellm.types.proxy.management_endpoints.scim_v2 import (
)
def _make_mock_request(
base_url="http://localhost:4000/", url="http://localhost:4000/scim/v2"
):
def _make_mock_request(base_url="http://localhost:4000/", url="http://localhost:4000/scim/v2"):
"""Create a mock FastAPI Request object."""
request = MagicMock()
request.method = "GET"
@ -67,9 +65,7 @@ class TestGetResourceTypes:
def test_custom_base_url(self):
resource_types = _get_resource_types("https://example.com/scim/v2")
user_rt = next(rt for rt in resource_types if rt.id == "User")
assert (
user_rt.meta["location"] == "https://example.com/scim/v2/ResourceTypes/User"
)
assert user_rt.meta["location"] == "https://example.com/scim/v2/ResourceTypes/User"
def test_model_dump_uses_schema_key(self):
"""Ensure model_dump() outputs 'schema' not 'schema_'."""
@ -82,16 +78,15 @@ class TestGetResourceTypes:
class TestGetSchemas:
def test_returns_user_and_group_schemas(self):
schemas = _get_schemas()
assert len(schemas) == 2
assert len(schemas) == 3
ids = [s.id for s in schemas]
assert "urn:ietf:params:scim:schemas:core:2.0:User" in ids
assert "urn:ietf:params:scim:schemas:core:2.0:Group" in ids
assert "urn:ietf:params:scim:schemas:extension:litellmAgent:2.0:User" in ids
def test_user_schema_has_required_attributes(self):
schemas = _get_schemas()
user_schema = next(
s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:User"
)
user_schema = next(s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:User")
attr_names = [a.name for a in user_schema.attributes]
assert "userName" in attr_names
assert "name" in attr_names
@ -101,9 +96,7 @@ class TestGetSchemas:
def test_group_schema_has_required_attributes(self):
schemas = _get_schemas()
group_schema = next(
s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group"
)
group_schema = next(s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group")
attr_names = [a.name for a in group_schema.attributes]
assert "displayName" in attr_names
assert "members" in attr_names
@ -112,9 +105,7 @@ class TestGetSchemas:
"""IdPs read the schema to learn we understand ``members.type``, which is how
a nested group announces itself."""
schemas = _get_schemas()
group_schema = next(
s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group"
)
group_schema = next(s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group")
members = next(a for a in group_schema.attributes if a.name == "members")
member_type = next(a for a in members.subAttributes or [] if a.name == "type")
assert member_type.type == "string"
@ -123,9 +114,7 @@ class TestGetSchemas:
def test_schema_meta_fields(self):
schemas = _get_schemas()
user_schema = next(
s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:User"
)
user_schema = next(s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:User")
assert user_schema.meta is not None
assert user_schema.meta["resourceType"] == "Schema"
@ -139,9 +128,7 @@ class TestGetScimBase:
request = _make_mock_request()
result = await get_scim_base(request)
assert result["schemas"] == [
"urn:ietf:params:scim:api:messages:2.0:ListResponse"
]
assert result["schemas"] == ["urn:ietf:params:scim:api:messages:2.0:ListResponse"]
assert result["totalResults"] == 2
assert len(result["Resources"]) == 2
@ -170,10 +157,7 @@ class TestGetScimBase:
result = await get_scim_base(request)
user_resource = next(r for r in result["Resources"] if r["id"] == "User")
assert (
user_resource["meta"]["location"]
== "https://proxy.example.com/scim/v2/ResourceTypes/User"
)
assert user_resource["meta"]["location"] == "https://proxy.example.com/scim/v2/ResourceTypes/User"
class TestGetResourceTypesEndpoint:
@ -182,9 +166,7 @@ class TestGetResourceTypesEndpoint:
request = _make_mock_request()
result = await get_resource_types(request)
assert result["schemas"] == [
"urn:ietf:params:scim:api:messages:2.0:ListResponse"
]
assert result["schemas"] == ["urn:ietf:params:scim:api:messages:2.0:ListResponse"]
assert result["totalResults"] == 2
@pytest.mark.asyncio
@ -232,10 +214,8 @@ class TestGetSchemasEndpoint:
request = _make_mock_request()
result = await get_schemas(request)
assert result["schemas"] == [
"urn:ietf:params:scim:api:messages:2.0:ListResponse"
]
assert result["totalResults"] == 2
assert result["schemas"] == ["urn:ietf:params:scim:api:messages:2.0:ListResponse"]
assert result["totalResults"] == 3
@pytest.mark.asyncio
async def test_resources_have_correct_ids(self):
@ -251,9 +231,7 @@ class TestGetSchemaById:
@pytest.mark.asyncio
async def test_get_user_schema(self):
request = _make_mock_request()
result = await get_schema(
request, schema_id="urn:ietf:params:scim:schemas:core:2.0:User"
)
result = await get_schema(request, schema_id="urn:ietf:params:scim:schemas:core:2.0:User")
assert result["id"] == "urn:ietf:params:scim:schemas:core:2.0:User"
assert result["name"] == "User"
@ -262,9 +240,7 @@ class TestGetSchemaById:
@pytest.mark.asyncio
async def test_get_group_schema(self):
request = _make_mock_request()
result = await get_schema(
request, schema_id="urn:ietf:params:scim:schemas:core:2.0:Group"
)
result = await get_schema(request, schema_id="urn:ietf:params:scim:schemas:core:2.0:Group")
assert result["id"] == "urn:ietf:params:scim:schemas:core:2.0:Group"
assert result["name"] == "Group"

View file

@ -2,8 +2,11 @@ import json
import logging
import time
from collections.abc import Callable, Mapping, Sequence
from itertools import chain
from types import MappingProxyType
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from itertools import chain, groupby
from operator import itemgetter
from types import MappingProxyType, ModuleType
from typing import Final
from unittest.mock import AsyncMock, MagicMock, call
@ -18,6 +21,7 @@ from litellm.proxy._types import (
LitellmUserRoles,
Member,
NewUserRequest,
NewUserRequestTeam,
NewUserResponse,
ProxyErrorTypes,
ProxyException,
@ -53,11 +57,13 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
update_user,
user_api_key_auth,
)
from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIM_MANAGED_TEAM_METADATA_KEY,
SCIM_TEAM_DATA_METADATA_KEY,
SCIMGroup,
SCIMListResponse,
SCIMMember,
SCIMPatchOp,
SCIMPatchOperation,
@ -507,12 +513,7 @@ async def test_scim_collection_endpoints_clamp_requested_page_size(
):
"""SCIM list endpoints accept zero and cap larger client page requests."""
mock_prisma_client = MagicMock()
mock_prisma_client.db = MagicMock()
table = MagicMock()
table.find_many = AsyncMock(return_value=[])
table.count = AsyncMock(return_value=0)
mock_prisma_client.db.litellm_usertable = table
mock_prisma_client.db.litellm_teamtable = table
query_raw = _writer_query_raw(mock_prisma_client, AsyncMock(side_effect=([], [{"total": 0}])))
mocker.patch( # test-quality-ok: HTTP validation requires an in-memory database boundary.
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
@ -522,12 +523,9 @@ async def test_scim_collection_endpoints_clamp_requested_page_size(
response = await client.get(f"/scim/v2/{endpoint}?startIndex=1&count={requested_count}")
assert response.status_code == 200
table.find_many.assert_awaited_once_with(
where={},
skip=0,
take=effective_count,
order={"created_at": "desc"},
)
page_sql, filter_value, limit, offset = query_raw.await_args_list[0].args
assert "LIMIT $2 OFFSET $3" in page_sql
assert (filter_value, limit, offset) == (None, effective_count, 0)
assert response.json()["itemsPerPage"] == 0
@ -827,7 +825,10 @@ async def test_handle_existing_user_by_email_roster_changes_use_existing_user_id
@pytest.mark.asyncio
async def test_handle_existing_user_by_email_syncs_roster_and_dedups_teams(mocker):
@pytest.mark.parametrize("structured_teams", [False, True])
async def test_handle_existing_user_by_email_syncs_roster_and_dedups_teams(
mocker: MockerFixture, structured_teams: bool
) -> None:
"""Existing-email upsert must add the user to the team roster via the shared
team_member_add path and dedup the teams built from repeated SCIM groups.
@ -862,7 +863,11 @@ async def test_handle_existing_user_by_email_syncs_roster_and_dedups_teams(mocke
user_id="same-id",
user_email="member@example.com",
user_alias="Member",
teams=["team-a", "team-a", "team-b"],
teams=(
[NewUserRequestTeam(team_id=team_id) for team_id in ("team-a", "team-a", "team-b")]
if structured_teams
else ["team-a", "team-a", "team-b"]
),
metadata={},
auto_create_key=False,
)
@ -1326,6 +1331,7 @@ async def test_update_user_put_with_valueless_entitlements_deactivates_user(scim
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.writer_db.litellm_scimresource.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user)
@ -2689,7 +2695,8 @@ async def test_update_user_demotes_when_default_params_lack_user_role(mocker, mo
@pytest.mark.asyncio
async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group(mocker, monkeypatch):
@pytest.mark.parametrize("source_owned", [False, True])
async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group(mocker, monkeypatch, source_owned: bool):
"""PATCH that drops the admin team from the resulting team set must write the
non-admin default, mirroring the PUT demotion path."""
from litellm.proxy.proxy_server import proxy_config
@ -2721,6 +2728,10 @@ async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group(mocke
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.writer_db = mock_prisma_client.db
mock_prisma_client.db.litellm_scimresource.find_many = AsyncMock(
return_value=[mocker.MagicMock(id="source-user", local_id="demote-me")] if source_owned else []
)
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user)
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
@ -2751,11 +2762,15 @@ async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group(mocke
await patch_user(user_id="demote-me", patch_ops=patch_ops)
call_args = mock_prisma_client.db.litellm_usertable.update.call_args
assert call_args[1]["data"]["user_role"] == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
if source_owned:
assert "user_role" not in call_args[1]["data"]
else:
assert call_args[1]["data"]["user_role"] == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
@pytest.mark.asyncio
async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch):
@pytest.mark.parametrize("source_owned", [False, True])
async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch, source_owned: bool):
"""PATCH carries groups as team ids, so admin-group matching must fall back to
each team's display name; an admin group configured as a human-readable alias
grants PROXY_ADMIN even when the team id differs."""
@ -2788,6 +2803,10 @@ async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.writer_db = mock_prisma_client.db
mock_prisma_client.db.litellm_scimresource.find_many = AsyncMock(
return_value=[mocker.MagicMock(id="source-user", local_id="promote-me")] if source_owned else []
)
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user)
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
@ -2818,7 +2837,10 @@ async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch)
await patch_user(user_id="promote-me", patch_ops=patch_ops)
call_args = mock_prisma_client.db.litellm_usertable.update.call_args
assert call_args[1]["data"]["user_role"] == LitellmUserRoles.PROXY_ADMIN
if source_owned:
assert "user_role" not in call_args[1]["data"]
else:
assert call_args[1]["data"]["user_role"] == LitellmUserRoles.PROXY_ADMIN
def _scim_admin_prisma(mocker, *, user_teams):
@ -2841,6 +2863,8 @@ def _scim_admin_prisma(mocker, *, user_teams):
prisma.db.litellm_usertable.update = AsyncMock(return_value=user)
prisma.db.litellm_teamtable = mocker.MagicMock()
prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=_team_find_unique)
prisma.writer_db = prisma.db
prisma.db.litellm_scimresource.find_many = AsyncMock(return_value=[])
return prisma
@ -6176,3 +6200,808 @@ async def test_merge_placeholder_refuses_rows_that_are_not_a_lone_placeholder(
assert reason in str(exc_info.value.message)
team_member_add_mock.assert_not_awaited()
prisma_client.db.litellm_usertable.delete.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route,arguments,dispatch",
[
("get_users", {"startIndex": 2, "count": 5, "filter": None}, "list"),
("get_groups", {"startIndex": 1, "count": 5, "filter": None}, "list"),
("get_user", {"user_id": "scoped-id"}, "get"),
("get_group", {"group_id": "scoped-id"}, "get"),
("create_user", {"user": SCIMUser(schemas=[], userName="human@example.com")}, "create_user"),
("create_group", {"group": SCIMGroup(schemas=[], displayName="Directory")}, "create_group"),
(
"update_user",
{"user_id": "scoped-id", "user": SCIMUser(schemas=[], userName="human@example.com")},
"update_user",
),
(
"update_group",
{"group_id": "scoped-id", "group": SCIMGroup(schemas=[], displayName="Directory")},
"update_group",
),
("patch_user", {"user_id": "scoped-id", "patch_ops": SCIMPatchOp(Operations=[])}, "update_user"),
("patch_group", {"group_id": "scoped-id", "patch_ops": SCIMPatchOp(Operations=[])}, "update_group"),
("delete_user", {"user_id": "scoped-id"}, "delete"),
("delete_group", {"group_id": "scoped-id"}, "delete"),
],
)
async def test_scoped_routes_propagate_source_denial_without_falling_back_to_legacy_humans(
route: str,
arguments: dict[str, object],
dispatch: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
from litellm.proxy.management_endpoints.scim.agent_provisioning import AgentProvisioningService
service: Final = AsyncMock(spec=AgentProvisioningService)
handler: Final = getattr(service, dispatch)
handler.side_effect = HTTPException(403, "Provisioning source disabled")
resolver: Final = AsyncMock(return_value=service)
legacy_database: Final = AsyncMock()
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", resolver)
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", legacy_database)
auth: Final = UserAPIKeyAuth(token="scoped-hash")
with pytest.raises(HTTPException) as failure:
await getattr(scim_v2, route)(**arguments, auth=auth)
assert failure.value.status_code == 403
resolver.assert_awaited_once_with(auth)
handler.assert_awaited_once()
legacy_database.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("ownership", ["new", "linked", "legacy"])
async def test_source_owned_groups_cannot_grant_global_admin(mocker, ownership: str) -> None:
from types import SimpleNamespace
from litellm.proxy.management_endpoints.scim.scim_v2 import _resolve_scim_user_role, _scim_groups_from_team_ids
prisma = _scim_admin_prisma(mocker, user_teams=["litellm-admins"])
resource = SimpleNamespace(
id="litellm-admins" if ownership == "new" else "source-group",
local_id="litellm-admins" if ownership == "linked" else None,
)
prisma.writer_db = SimpleNamespace(
litellm_scimresource=SimpleNamespace(
find_many=AsyncMock(return_value=[] if ownership == "legacy" else [resource])
)
)
groups = await _scim_groups_from_team_ids(prisma, ["litellm-admins"])
role = _resolve_scim_user_role(groups, "litellm-admins", LitellmUserRoles.INTERNAL_USER_VIEW_ONLY)
assert role == (LitellmUserRoles.PROXY_ADMIN if ownership == "legacy" else LitellmUserRoles.INTERNAL_USER_VIEW_ONLY)
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
async def test_source_membership_sync_preserves_manually_assigned_global_role(mocker, role) -> None:
from types import SimpleNamespace
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_scim_admin_group",
new=AsyncMock(return_value="litellm-admins"),
)
prisma = _scim_admin_prisma(mocker, user_teams=["litellm-admins"])
prisma.db.litellm_usertable.find_unique.return_value.user_role = role
prisma.db.litellm_scimresource.find_many.return_value = [SimpleNamespace(id="source-user", local_id="member-1")]
await _recompute_scim_member_roles(prisma, ["member-1"])
prisma.db.litellm_usertable.update.assert_not_awaited()
assert prisma.db.litellm_usertable.find_unique.return_value.user_role == role
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["create", "replace", "patch"])
async def test_legacy_token_cannot_classify_a_user_as_native_agent(
method: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
from litellm.types.proxy.management_endpoints.scim_v2 import SCIM_AGENT_USER_SCHEMA
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
database: Final = AsyncMock()
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", database)
extension: Final = {"identityParentId": "11111111-1111-4111-8111-111111111111"}
user: Final = SCIMUser.model_validate(
{"schemas": [], "userName": "agent@example.com", SCIM_AGENT_USER_SCHEMA: extension}
)
request: Final = (
scim_v2.create_user(user=user)
if method == "create"
else scim_v2.update_user(user_id="human", user=user)
if method == "replace"
else scim_v2.patch_user(
user_id="human",
patch_ops=SCIMPatchOp(Operations=[{"op": "add", "path": SCIM_AGENT_USER_SCHEMA, "value": extension}]),
)
)
with pytest.raises(HTTPException) as failure:
await request
assert failure.value.status_code == 400
database.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["Users", "Groups"])
async def test_source_delete_returns_no_content_without_legacy_fallback(
kind: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
from litellm.proxy.management_endpoints.scim.agent_provisioning import AgentProvisioningService
service: Final = AsyncMock(spec=AgentProvisioningService)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=service))
database: Final = AsyncMock()
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", database)
result: Final = (
await scim_v2.delete_user(user_id="owned") if kind == "Users" else await scim_v2.delete_group(group_id="owned")
)
assert result.status_code == 204
service.delete.assert_awaited_once_with(kind, "owned")
database.assert_not_awaited()
_OWNED_HUMAN: Final = LiteLLM_UserTable(
user_id="human-local", user_email="human@example.com", sso_user_id="00u1human", teams=["directory-team"]
)
_ORDINARY_USER: Final = LiteLLM_UserTable(user_id="ordinary-user", user_email="ordinary@example.com")
@dataclass(frozen=True, slots=True)
class _OwnedResource:
id: str
local_id: str
_LEGACY_TEAM: Final = LiteLLM_TeamTable(team_id="legacy-team", team_alias="legacy")
def _filter_subjects(filter_: object) -> tuple[str, ...]:
"""The literal(s) a Prisma string filter compares against: ``"x"``, ``{"equals": "x"}`` or ``{"in": [..]}``."""
if isinstance(filter_, str):
return (filter_,)
assert isinstance(filter_, dict), filter_
return tuple(filter_["in"]) if "in" in filter_ else (str(filter_["equals"]),)
def _legacy_membership_prisma(
mocker: MockerFixture, users: tuple[LiteLLM_UserTable, ...] = (_OWNED_HUMAN, _ORDINARY_USER)
) -> MagicMock:
"""One source-owned human (SCIM id ``scim-human``, local id ``human-local``), one source-owned
native agent-user (``nat-000001`` / ``agent-1``) and the given legacy users (one ordinary by default)."""
resources: Final = (_OwnedResource("scim-human", _OWNED_HUMAN.user_id), _OwnedResource("nat-000001", "agent-1"))
def _subjects(where: Mapping[str, object]) -> frozenset[str]:
clauses: Final = tuple(where.get("OR", (where,)))
filters: Final = chain.from_iterable(clause.values() for clause in clauses)
return frozenset(chain.from_iterable(map(_filter_subjects, filters)))
def _row_keys(row: LiteLLM_UserTable) -> tuple[tuple[str, LiteLLM_UserTable], ...]:
return tuple((key, row) for key in (row.user_id, row.sso_user_id, row.user_email) if key)
keyed: Final = sorted(chain.from_iterable(map(_row_keys, users)), key=itemgetter(0))
by_key: Final = MappingProxyType(
{key: tuple(row for _, row in group) for key, group in groupby(keyed, itemgetter(0))}
)
async def _users_find_many(where: Mapping[str, object], take: int | None = None) -> tuple[LiteLLM_UserTable, ...]:
hits: Final = chain.from_iterable(by_key.get(subject, ()) for subject in _subjects(where))
return tuple(MappingProxyType({row.user_id: row for row in hits}).values())
async def _users_find_unique(where: Mapping[str, str]) -> LiteLLM_UserTable | None:
return next((row for row in users if row.user_id == where["user_id"]), None)
async def _resources_find_many(where: Mapping[str, object]) -> tuple[_OwnedResource, ...]:
membership: Final = next(clause for clause in where["AND"] if "kind" not in clause)
field, filter_ = next(iter(membership.items()))
return tuple(row for row in resources if (row.id if field == "id" else row.local_id) in filter_["in"])
async def _team_find_unique(where: Mapping[str, str]) -> LiteLLM_TeamTable | None:
return _LEGACY_TEAM if where["team_id"] == _LEGACY_TEAM.team_id else None
by_folded_email: Final = MappingProxyType(
{row.user_email.lower(): row for row in users if row.user_email is not None}
)
async def _folded_email_rows(sql: str, folded: Sequence[str]) -> tuple[Mapping[str, str], ...]:
assert "LOWER(user_email) = ANY($1::text[])" in sql, sql
hits: Final = tuple(filter(None, map(by_folded_email.get, folded)))
return tuple(MappingProxyType({"user_id": row.user_id, "folded_email": row.user_email.lower()}) for row in hits)
prisma: Final = mocker.MagicMock()
prisma.db.litellm_usertable.find_many = AsyncMock(side_effect=_users_find_many)
prisma.writer_db.litellm_usertable.find_many = AsyncMock(side_effect=_users_find_many)
writer_tx: Final = mocker.MagicMock()
writer_tx.query_raw = AsyncMock(side_effect=_folded_email_rows)
prisma.tx.return_value.__aenter__ = AsyncMock(return_value=writer_tx)
prisma.tx.return_value.__aexit__ = AsyncMock(return_value=False)
prisma.db.litellm_usertable.find_unique = AsyncMock(side_effect=_users_find_unique)
prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=_team_find_unique)
prisma.db.litellm_teamtable.update = AsyncMock(return_value=_LEGACY_TEAM)
prisma.writer_db.litellm_scimresource.find_many = AsyncMock(side_effect=_resources_find_many)
return prisma
def _roster_written(roster: AsyncMock) -> set[str]:
"""The ``final_members`` the route handed to ``_handle_group_membership_changes``, however it was passed."""
args, kwargs = roster.call_args
return kwargs["final_members"] if "final_members" in kwargs else args[2]
def _legacy_route_call(route: str, members: tuple[str, ...], auth: UserAPIKeyAuth | None):
from litellm.proxy.management_endpoints.scim import scim_v2
group: Final = SCIMGroup(schemas=[], displayName="legacy", members=[SCIMMember(value=value) for value in members])
if route == "create":
return scim_v2.create_group(group=group, auth=auth)
if route == "replace":
return scim_v2.update_group(group_id="legacy-team", group=group, auth=auth)
entries: Final = [{"value": value} for value in members]
operations: Final = [SCIMPatchOperation(op=route.removeprefix("patch-"), path="members", value=entries)]
return scim_v2.patch_group(group_id="legacy-team", patch_ops=SCIMPatchOp(Operations=operations), auth=auth)
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["create", "replace", "patch-add", "patch-replace"])
@pytest.mark.parametrize(
"member",
[
pytest.param("scim-human", id="human-scim-id"),
pytest.param("human-local", id="human-local-id"),
pytest.param("human@example.com", id="human-email"),
pytest.param("00u1human", id="human-sso"),
pytest.param("nat-000001", id="native-scim-id"),
pytest.param("agent-1", id="native-local-id"),
],
)
async def test_legacy_key_cannot_put_a_source_owned_subject_on_its_roster(
route: str, member: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""A /scim/* key that is not a provisioning source gets 403 before any user is provisioned
or any team roster is written, whichever way it names the directory-owned subject."""
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
monkeypatch.setattr(scim_v2, "_check_team_exists", AsyncMock(return_value=_LEGACY_TEAM))
side_effects: Final = tuple(
mocker.patch.object(scim_v2, name, AsyncMock())
for name in ("_ensure_group_member_user", "new_team", "_handle_group_membership_changes", "team_member_add")
)
with pytest.raises(ProxyException) as failure:
await _legacy_route_call(route, ("ordinary-user", member), UserAPIKeyAuth(token="legacy-hash"))
assert str(failure.value.code) == "403", failure.value.message
assert member in failure.value.message
for side_effect in side_effects:
side_effect.assert_not_awaited()
_FORTY_THOUSAND: Final = 40_000
_OWNERSHIP_READS_PER_PASS: Final = 2 * -(-_FORTY_THOUSAND // IN_LIST_CHUNK_SIZE)
@pytest.mark.asyncio
async def test_legacy_key_naming_40k_native_ids_is_refused_after_the_batched_ownership_read_alone(
mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Native ids are owned by SCIM id, so the chunked resource read settles it: no per-member user reads."""
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
members: Final = tuple(f"nat-{index:06d}" for index in range(1, _FORTY_THOUSAND + 1))
with pytest.raises(HTTPException) as failure:
await scim_v2._assert_legacy_members_unowned(
UserAPIKeyAuth(token="legacy-hash"), tuple(SCIMMember(value=value) for value in members)
)
assert failure.value.status_code == 403
assert prisma.writer_db.litellm_scimresource.find_many.await_count == _OWNERSHIP_READS_PER_PASS
assert prisma.writer_db.litellm_usertable.find_many.await_count == 0
assert prisma.db.litellm_usertable.find_many.await_count == 0
@pytest.mark.asyncio
async def test_legacy_key_naming_40k_ordinary_users_resolves_aliases_in_chunked_reads(
mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Values found by id, SSO id or email exactly never fall back to a single-value read."""
from litellm.proxy.management_endpoints.scim import scim_v2
ordinary: Final = tuple(
LiteLLM_UserTable(user_id=f"u-{index}", user_email=f"u{index}@example.com", sso_user_id=f"sso-{index}")
for index in range(_FORTY_THOUSAND)
)
prisma: Final = _legacy_membership_prisma(mocker, users=(_OWNED_HUMAN, *ordinary))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
spellings: Final = (
lambda row: row.user_id,
lambda row: f" {row.sso_user_id} ",
lambda row: row.user_email,
)
members: Final = tuple(spellings[index % 3](row) for index, row in enumerate(ordinary))
await scim_v2._assert_legacy_members_unowned(
UserAPIKeyAuth(token="legacy-hash"), tuple(SCIMMember(value=value) for value in members)
)
assert prisma.writer_db.litellm_scimresource.find_many.await_count == 2 * _OWNERSHIP_READS_PER_PASS
assert prisma.writer_db.litellm_usertable.find_many.await_count == _OWNERSHIP_READS_PER_PASS
assert prisma.tx.call_count == _OWNERSHIP_READS_PER_PASS // 2
assert prisma.db.litellm_usertable.find_many.await_count == 0
assert prisma.db.litellm_usertable.find_unique.await_count == 0
@pytest.mark.asyncio
async def test_legacy_key_naming_an_owned_human_by_differently_cased_email_is_refused_on_the_writer(
mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""The alias is found by folded email on the writer while the read replica, lagging, knows no such user."""
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
replica: Final = AsyncMock(return_value=())
prisma.db.litellm_usertable.find_many = replica
prisma.db.query_raw = replica
with pytest.raises(HTTPException) as failure:
await scim_v2._assert_legacy_members_unowned(
UserAPIKeyAuth(token="legacy-hash"),
(SCIMMember(value="ordinary@example.com"), SCIMMember(value=" HUMAN@Example.com ")),
)
assert failure.value.status_code == 403
assert "HUMAN@Example.com" in str(failure.value.detail)
replica.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["create", "replace", "patch-add"])
async def test_ordinary_legacy_group_members_still_reach_the_roster(
route: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
monkeypatch.setattr(scim_v2, "_check_team_exists", AsyncMock(return_value=_LEGACY_TEAM))
monkeypatch.setattr(scim_v2, "_recompute_scim_member_roles", AsyncMock())
monkeypatch.setattr(scim_v2, "_get_team_member_user_ids_from_team", AsyncMock(return_value=[]))
monkeypatch.setattr(scim_v2, "_apply_group_patch_updates", AsyncMock(return_value=_LEGACY_TEAM))
roster: Final = mocker.patch.object(scim_v2, "_handle_group_membership_changes", AsyncMock())
created: Final = mocker.patch.object(
scim_v2, "new_team", AsyncMock(return_value=LiteLLM_TeamTable(team_id="legacy-team"))
)
monkeypatch.setattr(
scim_v2.ScimTransformations,
"transform_litellm_team_to_scim_group",
AsyncMock(return_value=SCIMGroup(schemas=[], id="legacy-team", displayName="legacy")),
)
result: Final = await _legacy_route_call(
route, ("ordinary-user", "ordinary@example.com"), UserAPIKeyAuth(token="legacy-hash")
)
assert result.id == "legacy-team"
if route == "create":
assert [member.user_id for member in created.call_args.kwargs["data"].members_with_roles] == ["ordinary-user"]
else:
assert _roster_written(roster) == {"ordinary-user"}
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["create", "replace"])
async def test_trusted_source_sync_still_places_its_own_humans_on_the_roster(
route: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""The in-process directory sync calls the legacy routes without a key and names its own
humans by local id; ownership is not a reason to refuse it."""
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
monkeypatch.setattr(scim_v2, "_check_team_exists", AsyncMock(return_value=_LEGACY_TEAM))
monkeypatch.setattr(scim_v2, "_recompute_scim_member_roles", AsyncMock())
monkeypatch.setattr(scim_v2, "_get_team_member_user_ids_from_team", AsyncMock(return_value=[]))
roster: Final = mocker.patch.object(scim_v2, "_handle_group_membership_changes", AsyncMock())
created: Final = mocker.patch.object(
scim_v2, "new_team", AsyncMock(return_value=LiteLLM_TeamTable(team_id="legacy-team"))
)
monkeypatch.setattr(
scim_v2.ScimTransformations,
"transform_litellm_team_to_scim_group",
AsyncMock(return_value=SCIMGroup(schemas=[], id="legacy-team", displayName="legacy")),
)
result: Final = await _legacy_route_call(route, ("human-local",), None)
assert result.id == "legacy-team"
if route == "create":
assert [member.user_id for member in created.call_args.kwargs["data"].members_with_roles] == ["human-local"]
else:
assert _roster_written(roster) == {"human-local"}
@pytest.mark.asyncio
async def test_empty_source_ownership_lookup_never_queries_the_writer(mocker: MockerFixture) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = mocker.MagicMock()
assert await scim_v2._source_owned_ids(prisma, "Users", ()) == frozenset()
prisma.writer_db.litellm_scimresource.find_many.assert_not_called()
@dataclass(frozen=True, slots=True)
class _Directory:
"""Legacy rows plus the set a provisioning source owns; ``query_raw`` evaluates the route's
``NOT EXISTS`` page/count statements against them, ``find_many`` only hydrates the ids it is handed."""
users: tuple[LiteLLM_UserTable, ...]
teams: tuple[LiteLLM_TeamTable, ...]
owned_local_ids: frozenset[str]
def _unowned(self, sql: str, filter_value: str | None) -> tuple[str, ...]:
rows: Final[Sequence[LiteLLM_UserTable | LiteLLM_TeamTable]] = (
self.users if 'FROM "LiteLLM_UserTable"' in sql else self.teams
)
assert 'NOT EXISTS (SELECT 1 FROM "LiteLLM_SCIMResource"' in sql, sql
ordered: Final = sorted(rows, key=lambda row: row.created_at or datetime.min, reverse=True)
ids: Final = tuple(
row.user_id if isinstance(row, LiteLLM_UserTable) else row.team_id
for row in ordered
if self._matches(row, filter_value)
)
return tuple(local_id for local_id in ids if local_id not in self.owned_local_ids)
@staticmethod
def _matches(row: LiteLLM_UserTable | LiteLLM_TeamTable, filter_value: str | None) -> bool:
if filter_value is None:
return True
if isinstance(row, LiteLLM_TeamTable):
return row.team_alias == filter_value
return filter_value in (row.user_id, row.user_email)
async def query_raw(self, sql: str, *params: object) -> tuple[Mapping[str, object], ...]:
filter_value: Final = params[0]
assert filter_value is None or isinstance(filter_value, str)
unowned: Final = self._unowned(sql, filter_value)
if sql.startswith("SELECT COUNT(*)"):
return ({"total": len(unowned)},)
assert "LIMIT $2 OFFSET $3" in sql and len(params) == 3, (sql, params)
limit, offset = params[1], params[2]
assert isinstance(limit, int) and isinstance(offset, int)
return tuple({"local_id": local_id} for local_id in unowned[offset : offset + limit])
def hydrate(self, key: str) -> AsyncMock:
rows: Final[Sequence[LiteLLM_UserTable | LiteLLM_TeamTable]] = self.users if key == "user_id" else self.teams
async def _find_many(where: Mapping[str, object]) -> tuple[LiteLLM_UserTable | LiteLLM_TeamTable, ...]:
wanted: Final = frozenset(_filter_subjects(where[key]))
assert len(wanted) <= scim_v2_module().SCIM_MAX_PAGE_SIZE, len(wanted)
return tuple(row for row in rows if _local_id(row) in wanted)
return AsyncMock(side_effect=_find_many)
def _local_id(row: LiteLLM_UserTable | LiteLLM_TeamTable) -> str:
return row.user_id if isinstance(row, LiteLLM_UserTable) else row.team_id
def _writer_query_raw(prisma: MagicMock, query_raw: AsyncMock) -> AsyncMock:
"""Route ``async with prisma.tx() as tx: tx.query_raw(...)`` to ``query_raw``."""
prisma.tx.return_value.__aenter__.return_value.query_raw = query_raw
return query_raw
def scim_v2_module() -> ModuleType:
from litellm.proxy.management_endpoints.scim import scim_v2
return scim_v2
def _directory_prisma(mocker: MockerFixture, directory: _Directory) -> MagicMock:
prisma: Final = mocker.MagicMock()
_writer_query_raw(prisma, AsyncMock(side_effect=directory.query_raw))
prisma.db.litellm_usertable.find_many = directory.hydrate("user_id")
prisma.db.litellm_teamtable.find_many = directory.hydrate("team_id")
prisma.writer_db.litellm_usertable.find_many = directory.hydrate("user_id")
prisma.writer_db.litellm_teamtable.find_many = directory.hydrate("team_id")
prisma.writer_db.litellm_scimresource.find_many = AsyncMock(
side_effect=AssertionError("legacy listing must not materialise the owned resource set")
)
return prisma
def _listing_route_stubs(mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch, prisma: MagicMock) -> None:
scim_v2: Final = scim_v2_module()
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
monkeypatch.setattr(scim_v2, "_get_team_member_user_ids_from_team", AsyncMock(return_value=[]))
monkeypatch.setattr(scim_v2, "_get_team_members_display", AsyncMock(return_value=[]))
mocker.patch.object(
scim_v2.ScimTransformations,
"transform_litellm_user_to_scim_user",
AsyncMock(side_effect=lambda user: SCIMUser(schemas=[], id=user.user_id, userName=user.user_id)),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["Users", "Groups"])
async def test_legacy_key_listing_does_not_see_source_owned_records(
kind: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""GET /Users and /Groups from a non-source key hide what GET /{id} already refuses with 403,
and totalResults counts only the visible records."""
scim_v2: Final = scim_v2_module()
directory: Final = _Directory(
users=(_OWNED_HUMAN, _ORDINARY_USER),
teams=(LiteLLM_TeamTable(team_id="directory-team", team_alias="directory"), _LEGACY_TEAM),
owned_local_ids=frozenset({_OWNED_HUMAN.user_id, "directory-team"}),
)
_listing_route_stubs(mocker, monkeypatch, _directory_prisma(mocker, directory))
auth: Final = UserAPIKeyAuth(token="legacy-hash")
listed: Final = (
await scim_v2.get_users(startIndex=1, count=10, filter=None, auth=auth)
if kind == "Users"
else await scim_v2.get_groups(startIndex=1, count=10, filter=None, auth=auth)
)
assert [resource.id for resource in listed.Resources] == ["ordinary-user" if kind == "Users" else "legacy-team"]
assert listed.totalResults == 1
def _large_directory(owned: int, legacy: int) -> _Directory:
"""``owned`` source-owned users and teams interleaved by creation time with ``legacy`` ordinary ones."""
epoch: Final = datetime(2026, 1, 1, tzinfo=timezone.utc)
users: Final = tuple(
LiteLLM_UserTable(
user_id=f"u{index:06d}", user_email=f"u{index}@example.com", created_at=epoch + timedelta(seconds=index)
)
for index in range(owned + legacy)
)
teams: Final = tuple(
LiteLLM_TeamTable(
team_id=f"t{index:06d}", team_alias=f"team {index}", created_at=epoch + timedelta(seconds=index)
)
for index in range(owned + legacy)
)
owned_ids: Final = frozenset(
chain(
(f"u{index:06d}" for index in range(0, owned + legacy, 2)),
(f"t{index:06d}" for index in range(0, owned + legacy, 2)),
)
)
return _Directory(users=users, teams=teams, owned_local_ids=owned_ids)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["Users", "Groups"])
async def test_legacy_key_listing_pages_a_large_directory_without_materialising_owned_ids(
kind: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Twenty thousand rows, half owned by a source: every page is the newest unowned rows in order,
totalResults is the unowned count, hydration reads one page of ids and the owned set is never loaded."""
scim_v2: Final = scim_v2_module()
directory: Final = _large_directory(owned=10_000, legacy=10_000)
prisma: Final = _directory_prisma(mocker, directory)
_listing_route_stubs(mocker, monkeypatch, prisma)
auth: Final = UserAPIKeyAuth(token="legacy-hash")
prefix: Final = "u" if kind == "Users" else "t"
async def _page(start_index: int) -> SCIMListResponse:
if kind == "Users":
return await scim_v2.get_users(startIndex=start_index, count=100, filter=None, auth=auth)
return await scim_v2.get_groups(startIndex=start_index, count=100, filter=None, auth=auth)
first, second, last = await _page(1), await _page(101), await _page(9_901)
newest_unowned: Final = tuple(f"{prefix}{index:06d}" for index in range(19_999, -1, -1) if index % 2)
assert [resource.id for resource in first.Resources] == list(newest_unowned[:100])
assert [resource.id for resource in second.Resources] == list(newest_unowned[100:200])
assert [resource.id for resource in last.Resources] == list(newest_unowned[9_900:])
assert (first.totalResults, second.totalResults, last.totalResults) == (10_000, 10_000, 10_000)
assert (first.itemsPerPage, second.itemsPerPage, last.itemsPerPage) == (100, 100, 100)
prisma.writer_db.litellm_scimresource.find_many.assert_not_called()
assert prisma.tx.return_value.__aenter__.return_value.query_raw.await_count == 6
@pytest.mark.asyncio
async def test_legacy_key_username_filter_ignores_source_owned_match(
mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""The Okta deprovisioning lookup ``userName eq`` finds an ordinary user but reports total 0 for an owned one."""
scim_v2: Final = scim_v2_module()
directory: Final = _Directory(
users=(_OWNED_HUMAN, _ORDINARY_USER), teams=(), owned_local_ids=frozenset({_OWNED_HUMAN.user_id})
)
_listing_route_stubs(mocker, monkeypatch, _directory_prisma(mocker, directory))
auth: Final = UserAPIKeyAuth(token="legacy-hash")
owned: Final = await scim_v2.get_users(
startIndex=1, count=10, filter=f'userName eq "{_OWNED_HUMAN.user_email}"', auth=auth
)
ordinary: Final = await scim_v2.get_users(
startIndex=1, count=10, filter=f'userName eq "{_ORDINARY_USER.user_email}"', auth=auth
)
assert (owned.totalResults, owned.Resources) == (0, [])
assert (ordinary.totalResults, [resource.id for resource in ordinary.Resources]) == (1, ["ordinary-user"])
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["Users", "Groups"])
async def test_trusted_listing_keeps_the_plain_legacy_query(
kind: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""With no key (trusted flow) the route still lists every row through the legacy Prisma query."""
scim_v2: Final = scim_v2_module()
prisma: Final = mocker.MagicMock()
_writer_query_raw(prisma, AsyncMock(side_effect=AssertionError("trusted listing must not change query shape")))
prisma.db.litellm_usertable.find_many = AsyncMock(return_value=(_OWNED_HUMAN, _ORDINARY_USER))
prisma.db.litellm_usertable.count = AsyncMock(return_value=2)
prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=(_LEGACY_TEAM,))
prisma.db.litellm_teamtable.count = AsyncMock(return_value=1)
_listing_route_stubs(mocker, monkeypatch, prisma)
listed: Final = (
await scim_v2.get_users(startIndex=1, count=10, filter=None, auth=None)
if kind == "Users"
else await scim_v2.get_groups(startIndex=1, count=10, filter='displayName eq "legacy"', auth=None)
)
assert listed.totalResults == (2 if kind == "Users" else 1)
table: Final = prisma.db.litellm_usertable if kind == "Users" else prisma.db.litellm_teamtable
table.find_many.assert_awaited_once_with(
where={} if kind == "Users" else {"team_alias": "legacy"}, skip=0, take=10, order={"created_at": "desc"}
)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["Users", "Groups"])
async def test_legacy_listing_hydrates_writer_page_despite_replica_lag(
kind: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
scim_v2: Final = scim_v2_module()
directory: Final = _Directory(
users=(_ORDINARY_USER,), teams=(_LEGACY_TEAM,), owned_local_ids=frozenset()
)
prisma: Final = _directory_prisma(mocker, directory)
prisma.db.litellm_usertable.find_many = AsyncMock(return_value=())
prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=())
_listing_route_stubs(mocker, monkeypatch, prisma)
auth: Final = UserAPIKeyAuth(token="legacy-hash")
listed: Final = (
await scim_v2.get_users(startIndex=1, count=10, filter=None, auth=auth)
if kind == "Users"
else await scim_v2.get_groups(startIndex=1, count=10, filter=None, auth=auth)
)
assert [resource.id for resource in listed.Resources] == ["ordinary-user" if kind == "Users" else "legacy-team"]
assert listed.totalResults == listed.itemsPerPage == 1
prisma.db.litellm_usertable.find_many.assert_not_awaited()
prisma.db.litellm_teamtable.find_many.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"method,path", [("GET", "/scim/v2/placeholders"), ("POST", "/scim/v2/placeholders/shadow/merge")]
)
@pytest.mark.parametrize("enabled", [True, False])
async def test_source_token_cannot_read_or_merge_global_placeholders(
monkeypatch: pytest.MonkeyPatch, method: str, path: str, enabled: bool
) -> None:
from types import SimpleNamespace
from litellm.proxy import proxy_server
database: Final = MagicMock()
database.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(enabled=enabled))
monkeypatch.setattr(proxy_server, "prisma_client", database)
app: Final = FastAPI()
app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler)
app.include_router(scim_router)
app.dependency_overrides[_premium_user_check] = lambda: None
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
token="source-token-hash", allowed_routes=["/scim/*"], user_role=LitellmUserRoles.PROXY_ADMIN
)
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
response: Final = await client.request(method, path)
assert response.status_code == 403, response.text
database.writer_db.litellm_scimsource.find_unique.assert_awaited_once_with(where={"key_hash": "source-token-hash"})
database.tx.assert_not_called()
database.db.litellm_usertable.find_unique.assert_not_called()
database.db.litellm_usertable.delete.assert_not_called()
database.db.litellm_teammembership.create.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["create", "replace", "patch-add", "patch-replace"])
async def test_legacy_user_routes_cannot_join_a_source_owned_team(
route: str, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
prisma: Final = _legacy_membership_prisma(mocker)
async def owned(_client: object, kind: str, ids: tuple[str, ...]) -> frozenset[str]:
return frozenset(ids) & {"directory-team"} if kind == "Groups" else frozenset()
monkeypatch.setattr(scim_v2, "_source_owned_ids", owned)
monkeypatch.setattr(scim_v2, "_agent_provisioning_service", AsyncMock(return_value=None))
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=prisma))
monkeypatch.setattr(scim_v2, "_check_user_exists", AsyncMock(return_value=_ORDINARY_USER))
monkeypatch.setattr(scim_v2, "_get_scim_admin_group", AsyncMock(return_value=None))
create: Final = AsyncMock(side_effect=AssertionError("user creation must follow ownership validation"))
roster: Final = AsyncMock(side_effect=AssertionError("roster writes must follow ownership validation"))
monkeypatch.setattr(scim_v2, "new_user", create)
monkeypatch.setattr(scim_v2, "_handle_team_membership_changes", roster)
user: Final = SCIMUser(schemas=[], userName="new-user", groups=[SCIMUserGroup(value="directory-team")])
auth: Final = UserAPIKeyAuth(token="legacy-hash")
pending: Final = (
scim_v2.create_user(user=user, auth=auth)
if route == "create"
else scim_v2.update_user(user_id="ordinary-user", user=user, auth=auth)
if route == "replace"
else scim_v2.patch_user(
user_id="ordinary-user",
patch_ops=SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op=route.removeprefix("patch-"), path="groups", value=[{"value": "directory-team"}]
)
]
),
auth=auth,
)
)
with pytest.raises((HTTPException, ProxyException)) as failure:
await pending
status: Final = failure.value.status_code if isinstance(failure.value, HTTPException) else failure.value.code
assert str(status) == "403"
create.assert_not_awaited()
roster.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"current,proposed,trusted,denied",
[
((), ("ordinary-team",), False, False),
((), ("directory-team",), True, False),
(("directory-team",), ("directory-team",), False, False),
(("directory-team",), (), False, True),
],
)
async def test_legacy_team_changes_preserve_directory_ownership(
current: tuple[str, ...],
proposed: tuple[str, ...],
trusted: bool,
denied: bool,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy.management_endpoints.scim import scim_v2
async def owned(_client: object, kind: str, ids: tuple[str, ...]) -> frozenset[str]:
assert kind == "Groups"
return frozenset(ids) & {"directory-team"}
monkeypatch.setattr(scim_v2, "_get_prisma_client_or_raise_exception", AsyncMock(return_value=object()))
monkeypatch.setattr(scim_v2, "_source_owned_ids", owned)
if denied:
with pytest.raises(HTTPException) as failure:
await scim_v2.assert_legacy_team_changes_unowned(UserAPIKeyAuth(), current, proposed)
assert failure.value.status_code == 403
else:
await scim_v2.assert_legacy_team_changes_unowned(None if trusted else UserAPIKeyAuth(), current, proposed)

View file

@ -34,6 +34,7 @@ def source_database(monkeypatch: pytest.MonkeyPatch):
client: Final = MagicMock(spec=PrismaClient)
tx: Final = client.tx.return_value.__aenter__.return_value
monkeypatch.setattr(proxy_server, "prisma_client", client)
tx.execute_raw = AsyncMock()
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=[])
@ -239,3 +240,18 @@ async def test_source_mapping_accepts_access_groups_across_query_batches(
)
assert result.group_mappings[0].access_group_ids == group_ids
assert tx.litellm_accessgrouptable.find_many.await_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize("current_key", [None, SimpleNamespace(allowed_routes=["/*"])])
async def test_source_creation_revalidates_key_after_waiting_for_mutation(monkeypatch, current_key):
tx = source_database(monkeypatch)
async def complete_concurrent_mutation(*args):
tx.litellm_verificationtoken.find_unique.return_value = current_key
tx.execute_raw.side_effect = complete_concurrent_mutation
with pytest.raises(HTTPException) as denied:
await create_source(SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token"), ADMIN)
assert denied.value.status_code == 400
tx.litellm_scimsource.create.assert_not_awaited()

View file

@ -12755,6 +12755,10 @@ def _make_regenerate_mock_prisma():
return_value=None
)
mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data)
mock_prisma_client.tx = MagicMock()
mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db
mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=_make_regenerate_existing_key())
return mock_prisma_client
@ -14823,6 +14827,11 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha
return_value=None
)
mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data)
mock_prisma_client.tx = MagicMock()
mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db
mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key)
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
@ -21097,3 +21106,224 @@ class TestTeamAdminMemberKeyBudgetUpdate:
)
assert exc.value.status_code == 403
assert "member_key_budgets" not in str(exc.value.detail)
@pytest.mark.asyncio
async def test_source_bound_key_cannot_regenerate_before_any_write():
from types import SimpleNamespace
from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration
client = _make_regenerate_mock_prisma()
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source"))
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
with _patch_regenerate_side_effects():
with pytest.raises(HTTPException) as denied:
await _execute_virtual_key_regeneration(
prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None,
user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None,
user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(),
)
assert denied.value.status_code == 409
client.db.litellm_verificationtoken.update.assert_not_awaited()
client.db.litellm_verificationtoken.create.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("routes", [None, [], ["llm_api_routes"], ["/scim/*", "/key/info"]])
async def test_source_bound_key_cannot_remove_its_route_restriction(routes):
from types import SimpleNamespace
from litellm.proxy.management_endpoints.key_management_endpoints import _process_single_key_update
client = _make_regenerate_mock_prisma()
client.update_data = AsyncMock(return_value={"data": {"key_alias": "source"}})
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source"))
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
with pytest.raises(HTTPException) as denied:
await _process_single_key_update(
update_key_request=UpdateKeyRequest(key=key.token, allowed_routes=routes), existing_key_row=key,
user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None,
prisma_client=client, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), llm_router=None,
)
assert denied.value.status_code == 409
client.update_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_unbound_scim_key_can_still_regenerate():
from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration
client = _make_regenerate_mock_prisma()
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
with _patch_regenerate_side_effects():
result = await _execute_virtual_key_regeneration(
prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None,
user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None,
user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(),
)
assert result.key is not None
client.db.litellm_verificationtoken.update.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("routes", [None, [], ["/*"], ["/scim/*", "/key/info"]])
async def test_key_update_endpoint_preserves_source_token_restriction(monkeypatch, routes):
from types import SimpleNamespace
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
_wire_update_key_fn(monkeypatch, key)
client = proxy_server.prisma_client
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source"))
request = MagicMock()
request.query_params = {}
with pytest.raises(ProxyException) as denied:
await update_key_fn(
request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=routes),
user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None,
)
assert str(denied.value.code) == "409"
client.update_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_regeneration_rechecks_source_binding_before_transaction_writes():
from types import SimpleNamespace
from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration
client = _make_regenerate_mock_prisma()
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
tx = AsyncMock()
client.tx = MagicMock()
client.tx.return_value.__aenter__.return_value = tx
tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source")
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
tx.litellm_verificationtoken.find_unique.return_value = key
with _patch_regenerate_side_effects():
with pytest.raises(HTTPException) as denied:
await _execute_virtual_key_regeneration(
prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None,
user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None,
user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(),
)
assert denied.value.status_code == 409
tx.litellm_verificationtoken.update.assert_not_awaited()
tx.litellm_deletedverificationtoken.create_many.assert_not_awaited()
client.db.litellm_verificationtoken.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("bound,routes,missing,expected", [(True, ["/*"], False, 409), (True, ["/scim/*"], False, None), (False, [], False, None), (False, [], True, 409)])
async def test_route_write_rechecks_binding_and_preserves_supported_updates(bound, routes, missing, expected):
from types import SimpleNamespace
from litellm.proxy.management_endpoints.key_management_endpoints import _update_key_row_with_soft_budget
client = _make_regenerate_mock_prisma()
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]})
tx = client.db
tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source") if bound else None
tx.litellm_verificationtoken.find_unique.return_value = None if missing else key
tx.litellm_verificationtoken.update.return_value = key.model_copy(update={"allowed_routes": routes})
request = UpdateKeyRequest(key=key.token, allowed_routes=routes)
write = _update_key_row_with_soft_budget(client, key.token, request, {"allowed_routes": routes}, key, "admin")
if expected:
with pytest.raises(HTTPException) as denied:
await write
assert denied.value.status_code == expected
tx.litellm_verificationtoken.update.assert_not_awaited()
else:
response = await write
assert response["data"]["allowed_routes"] == routes
tx.litellm_budgettable.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("transactional", [False, True])
async def test_rotation_grace_period_is_written_to_the_selected_database(transactional):
from litellm.proxy.management_endpoints.key_management_endpoints import _insert_deprecated_key
client = _make_regenerate_mock_prisma()
transaction = AsyncMock()
before = datetime.now(timezone.utc)
await _insert_deprecated_key(client, "old-token", "new-token", "1h", tx=transaction if transactional else None)
selected = transaction if transactional else client.db
unused = client.db if transactional else transaction
saved = selected.litellm_deprecatedverificationtoken.upsert.call_args.kwargs
assert saved["where"] == {"token": "old-token"}
assert saved["data"]["create"]["active_token_id"] == "new-token"
assert saved["data"]["update"]["active_token_id"] == "new-token"
assert before + timedelta(hours=1) <= saved["data"]["create"]["revoke_at"] <= datetime.now(timezone.utc) + timedelta(hours=1)
unused.litellm_deprecatedverificationtoken.upsert.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ["update", "regenerate"])
async def test_source_binding_race_denial_preserves_object_permissions(monkeypatch, operation):
from types import SimpleNamespace
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn, _execute_virtual_key_regeneration
key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"], "object_permission_id": "permission-before"})
if operation == "update":
_wire_update_key_fn(monkeypatch, key)
client = proxy_server.prisma_client
else:
client = _make_regenerate_mock_prisma()
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
events = []
_record_object_permission_writes(client, events)
tx = AsyncMock()
client.tx = MagicMock()
client.tx.return_value.__aenter__.return_value = tx
tx.litellm_verificationtoken.find_unique.return_value = key
tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source")
permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"])
if operation == "update":
request = MagicMock()
request.query_params = {}
call = update_key_fn(request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=["/*"], object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None)
else:
call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock())
with _patch_regenerate_side_effects():
with pytest.raises((HTTPException, ProxyException)) as denied:
await call
assert str(getattr(denied.value, "code", getattr(denied.value, "status_code", None))) == "409"
assert events == []
client.db.litellm_objectpermissiontable.upsert.assert_not_awaited()
tx.litellm_objectpermissiontable.upsert.assert_not_awaited()
tx.litellm_verificationtoken.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ["update", "regenerate"])
async def test_permission_update_joins_key_write_transaction(operation):
from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration, _update_key_row_with_soft_budget
client = _make_regenerate_mock_prisma()
client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None)
key = _make_regenerate_existing_key().model_copy(update={"object_permission_id": "permission-before"})
tx = AsyncMock()
client.tx.return_value.__aenter__.return_value = tx
tx.litellm_verificationtoken.find_unique.return_value = key
tx.litellm_scimsource.find_unique.return_value = None
tx.litellm_objectpermissiontable.find_unique.return_value = None
tx.litellm_objectpermissiontable.upsert.return_value = MagicMock(object_permission_id="permission-before")
tx.litellm_verificationtoken.update.side_effect = RuntimeError("token write failed")
permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"])
if operation == "update":
call = _update_key_row_with_soft_budget(client, key.token, UpdateKeyRequest(key=key.token, allowed_routes=[], object_permission=permission), {"allowed_routes": [], "object_permission": permission.model_dump()}, key, "admin")
else:
call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock())
with _patch_regenerate_side_effects():
with pytest.raises(RuntimeError, match="token write failed"):
await call
client.db.litellm_objectpermissiontable.upsert.assert_not_awaited()
saved = tx.litellm_objectpermissiontable.upsert.await_args.kwargs
assert saved["where"] == {"object_permission_id": "permission-before"}
assert saved["data"]["update"]["vector_stores"] == ["vs-after"]
assert tx.litellm_verificationtoken.update.await_args.kwargs["data"]["object_permission_id"] == "permission-before"
assert client.tx.return_value.__aexit__.await_args.args[0] is RuntimeError

View file

@ -14403,6 +14403,41 @@ export interface paths {
patch?: never;
trace?: never;
};
"/scim/v2/sources": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/** List Sources */
get: operations["list_sources_scim_v2_sources_get"];
put?: never;
/** Create Source */
post: operations["create_source_scim_v2_sources_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/scim/v2/sources/{source_id}": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
/** Update Source */
put: operations["update_source_scim_v2_sources__source_id__put"];
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/search": {
parameters: {
query?: never;
@ -24813,6 +24848,8 @@ export interface components {
* @constant
*/
provider: "microsoft_entra";
/** Provisioning Source Id */
provisioning_source_id?: string | null;
/**
* Required Roles
* @default []
@ -24989,6 +25026,13 @@ export interface components {
created_at?: string | null;
/** Created By */
created_by?: string | null;
/** Directory Access Group Ids */
directory_access_group_ids?: string[] | null;
/**
* Directory Active
* @default true
*/
directory_active: boolean;
/**
* Enabled
* @default true
@ -30889,6 +30933,8 @@ export interface components {
* @constant
*/
provider: "microsoft_entra";
/** Provisioning Source Id */
provisioning_source_id?: string | null;
/**
* Required Roles
* @default []
@ -36570,6 +36616,13 @@ export interface components {
};
/** ManagedAgentIdentityStatus */
ManagedAgentIdentityStatus: {
/** Directory Access Group Ids */
directory_access_group_ids?: string[] | null;
/**
* Directory Active
* @default true
*/
directory_active: boolean;
/**
* Enabled
* @default true
@ -43141,6 +43194,16 @@ export interface components {
/** Schemas */
schemas: string[];
};
/** SCIMGroupMapping */
SCIMGroupMapping: {
/** Access Group Ids */
access_group_ids: string[];
/**
* External Group Id
* Format: uuid
*/
external_group_id: string;
};
/** SCIMListResponse */
SCIMListResponse: {
/** Resources */
@ -43283,6 +43346,73 @@ export interface components {
*/
sort: components["schemas"]["SCIMFeature"];
};
/** SCIMSourceConfig */
SCIMSourceConfig: {
/** Display Name */
display_name: string;
/**
* Enabled
* @default true
*/
enabled: boolean;
/**
* Group Mappings
* @default []
*/
group_mappings: components["schemas"]["SCIMGroupMapping"][];
/**
* Tenant Id
* Format: uuid
*/
tenant_id: string;
};
/** SCIMSourceCreate */
SCIMSourceCreate: {
/** Display Name */
display_name: string;
/**
* Enabled
* @default true
*/
enabled: boolean;
/**
* Group Mappings
* @default []
*/
group_mappings: components["schemas"]["SCIMGroupMapping"][];
/**
* Provisioning Token
* Format: password
*/
provisioning_token: string;
/**
* Tenant Id
* Format: uuid
*/
tenant_id: string;
};
/** SCIMSourceResponse */
SCIMSourceResponse: {
/** Display Name */
display_name: string;
/**
* Enabled
* @default true
*/
enabled: boolean;
/**
* Group Mappings
* @default []
*/
group_mappings: components["schemas"]["SCIMGroupMapping"][];
/** Source Id */
source_id: string;
/**
* Tenant Id
* Format: uuid
*/
tenant_id: string;
};
/** SCIMUser */
"SCIMUser-Input": {
/**
@ -67826,6 +67956,109 @@ export interface operations {
};
};
};
list_sources_scim_v2_sources_get: {
parameters: {
query?: {
feature?: string | null;
};
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["SCIMSourceResponse"][];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
create_source_scim_v2_sources_post: {
parameters: {
query?: {
feature?: string | null;
};
header?: never;
path?: never;
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["SCIMSourceCreate"];
};
};
responses: {
/** @description Successful Response */
201: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["SCIMSourceResponse"];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
update_source_scim_v2_sources__source_id__put: {
parameters: {
query?: {
feature?: string | null;
};
header?: never;
path: {
source_id: string;
};
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["SCIMSourceConfig"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["SCIMSourceResponse"];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
search_search_post: {
parameters: {
query?: {