mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(agents): scim activation and authorization
This commit is contained in:
parent
585f671bad
commit
ee9d32ccaf
20 changed files with 2534 additions and 134 deletions
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
233
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
233
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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?: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue