From ee9d32ccaf9a5eb0abf94af4fc2de6aca03d5b42 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 16:23:55 -0700 Subject: [PATCH 1/5] feat(agents): scim activation and authorization --- litellm/proxy/_lazy_openapi_snapshot.json | 391 ++++++++ .../auth/agent_access_groups.py | 14 +- .../auth/agent_permission_handler.py | 2 + .../auth/managed_authorization.py | 8 +- litellm/proxy/agent_endpoints/endpoints.py | 2 + .../proxy/agent_endpoints/identity_store.py | 117 ++- .../proxy/agent_endpoints/managed_identity.py | 51 +- .../management_endpoints/scim/scim_v2.py | 459 +++++++-- litellm/types/agents.py | 3 + litellm/types/proxy/agent_identity.py | 17 +- .../auth/test_managed_agent_access.py | 36 + .../auth/test_agent_permission_handler.py | 2 + .../auth/test_managed_authorization.py | 10 + .../agent_endpoints/test_identity_store.py | 235 ++++- .../agent_endpoints/test_managed_identity.py | 111 +++ .../scim/test_agent_provisioning.py | 38 + .../scim/test_human_provisioning.py | 16 + .../scim/test_scim_v2_discovery.py | 54 +- .../scim/test_scim_v2_endpoints.py | 869 +++++++++++++++++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 233 +++++ 20 files changed, 2534 insertions(+), 134 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index a50c0ec31ae..3fedd5a8ff9 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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" + ] + } } } }, diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py index 67547e82f24..02e0feb866b 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py +++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py @@ -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) diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 4e8880d37c5..d591bb7193d 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index d945c6e1ad2..a01370e59c4 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -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") diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index a76ec57b630..e1b317e3104 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -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, diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py index ac0ebac7865..6ece3925a63 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -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") diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py index e2289e3bf47..8e64ea3639c 100644 --- a/litellm/proxy/agent_endpoints/managed_identity.py +++ b/litellm/proxy/agent_endpoints/managed_identity.py @@ -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") diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index c07aab47bd0..ff4d0c27c54 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -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( diff --git a/litellm/types/agents.py b/litellm/types/agents.py index d59b6fa6ce9..38027d7ed79 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -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 diff --git a/litellm/types/proxy/agent_identity.py b/litellm/types/proxy/agent_identity.py index 2dc43976e0f..b1ceba79577 100644 --- a/litellm/types/proxy/agent_identity.py +++ b/litellm/types/proxy/agent_identity.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py index 90403e5553f..7344a75125d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py @@ -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( diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a1a022fdd35..9c64d82d0ce 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -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( diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index bc84e9ddd03..4ae6d4fa10d 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -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() diff --git a/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py index 005f0b4c074..c567c45c9aa 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py @@ -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",) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py index 55b1a0d366e..ecd65ca026b 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py index bd23720d9ca..7774e5555c8 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_agent_provisioning.py @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py b/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py index 886cd1a177d..fe54c42c326 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_human_provisioning.py @@ -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( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py index 6ced5264267..022f8e9292e 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py @@ -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" diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index c50ae9ce1cf..6b82f9f77ba 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -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) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 4f1616a3c28..6afa9f418b8 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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?: { From d8efd6a151c8c5d36f2d84da469e9bf2c60f31a5 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:37:23 -0700 Subject: [PATCH 2/5] fix(scim): guard source token regeneration and route changes --- .../key_management_endpoints.py | 22 +++++++ .../test_key_management_endpoints.py | 57 +++++++++++++++++++ 2 files changed, 79 insertions(+) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2de9ddc2577..a8124ac6a12 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -142,6 +142,7 @@ from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, + SCIMSourceRepository, ) from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository @@ -2744,6 +2745,13 @@ async def _process_single_key_update( prisma_client=prisma_client, ) + if ( + "allowed_routes" in update_key_request.model_fields_set + and tuple(update_key_request.allowed_routes or ()) != ("/scim/*",) + and prisma_client is not None + ): + await _reject_source_bound_key_change(prisma_client, existing_key_row) + _check_disable_global_guardrails_caller_permission( update_key_request.disable_global_guardrails, update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict @@ -5532,6 +5540,18 @@ async def _insert_deprecated_key( ) +async def _reject_source_bound_key_change(prisma_client: PrismaClient, key: LiteLLM_VerificationToken) -> None: + if tuple(key.allowed_routes or ()) != ("/scim/*",): + return + source: Final = await SCIMSourceRepository(prisma_client, use_writer=True).table.find_unique( + where={"key_hash": key.token} + ) + if source is not None: + raise HTTPException( + 409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed" + ) + + async def _execute_virtual_key_regeneration( *, prisma_client: PrismaClient, @@ -5549,6 +5569,8 @@ async def _execute_virtual_key_regeneration( from litellm.proxy import proxy_server from litellm.proxy.proxy_server import hash_token + await _reject_source_bound_key_change(prisma_client, key_in_db) + # Mirror the /key/update ownership rebind guard. See helper docstring. _validate_caller_can_change_key_ownership( data=data, diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 5ea38ce23d5..265d781a7e4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21097,3 +21097,60 @@ class TestTeamAdminMemberKeyBudgetUpdate: ) assert exc.value.status_code == 403 assert "member_key_budgets" not in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_source_bound_key_cannot_regenerate_before_any_write(): + from types import SimpleNamespace + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source")) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + with _patch_regenerate_side_effects(): + with pytest.raises(HTTPException) as denied: + await _execute_virtual_key_regeneration( + prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), + ) + assert denied.value.status_code == 409 + client.db.litellm_verificationtoken.update.assert_not_awaited() + client.db.litellm_verificationtoken.create.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("routes", [None, [], ["llm_api_routes"], ["/scim/*", "/key/info"]]) +async def test_source_bound_key_cannot_remove_its_route_restriction(routes): + from types import SimpleNamespace + from litellm.proxy.management_endpoints.key_management_endpoints import _process_single_key_update + + client = _make_regenerate_mock_prisma() + client.update_data = AsyncMock(return_value={"data": {"key_alias": "source"}}) + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source")) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + with pytest.raises(HTTPException) as denied: + await _process_single_key_update( + update_key_request=UpdateKeyRequest(key=key.token, allowed_routes=routes), existing_key_row=key, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + prisma_client=client, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), llm_router=None, + ) + assert denied.value.status_code == 409 + client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_unbound_scim_key_can_still_regenerate(): + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + with _patch_regenerate_side_effects(): + result = await _execute_virtual_key_regeneration( + prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), + ) + assert result.key is not None + client.db.litellm_verificationtoken.update.assert_awaited_once() From d362a91d2b31c0a6f3769fdd451b0bc4cb1c6b76 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:52:40 -0700 Subject: [PATCH 3/5] fix(scim): serialize source binding with token mutations --- .../key_management_endpoints.py | 91 +++++++++++++------ .../scim/source_endpoints.py | 7 +- .../scim/test_source_endpoints.py | 16 ++++ .../test_key_management_endpoints.py | 85 +++++++++++++++++ 4 files changed, 168 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a8124ac6a12..aded71edd78 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2434,14 +2434,19 @@ async def _update_key_row_with_soft_budget( ) -> _KeyUpdateResult: hashed_token: Final = _hash_token_if_needed(key) key_where: Final[_KeyRowWhere] = {"token": hashed_token} - tx: _KeyUpdateTx async with prisma_client.tx() as tx: - update_values: Final = await _apply_soft_budget_update( - data=data, - non_default_values=non_default_values, - db=tx, - existing_key_row=existing_key_row, - changed_by=changed_by, + if "allowed_routes" in data.model_fields_set: + await _lock_and_validate_source_key_change(tx, hashed_token, data.allowed_routes) + update_values: Final = ( + await _apply_soft_budget_update( + data=data, + non_default_values=non_default_values, + db=tx, + existing_key_row=existing_key_row, + changed_by=changed_by, + ) + if "soft_budget" in data.model_fields_set + else non_default_values ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( @@ -3498,6 +3503,9 @@ async def update_key_fn( await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data) + if "allowed_routes" in data.model_fields_set and tuple(data.allowed_routes or ()) != ("/scim/*",): + await _reject_source_bound_key_change(prisma_client, existing_key_row) + # Enforce upperbound key params on update (don't fill defaults) _enforce_upperbound_key_params(data, fill_defaults=False) non_default_values: Final = await prepare_key_update_data( @@ -3551,7 +3559,7 @@ async def update_key_fn( existing_key_row=existing_key_row, changed_by=changed_by, ) - if "soft_budget" in data.model_fields_set + if {"soft_budget", "allowed_routes"}.intersection(data.model_fields_set) else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) ) @@ -5484,6 +5492,7 @@ async def _insert_deprecated_key( old_token_hash: str, new_token_hash: str, grace_period: str | None, + tx: "Prisma | None" = None, ) -> None: """ Insert old key into deprecated table so it remains valid during grace period. @@ -5514,7 +5523,12 @@ async def _insert_deprecated_key( try: revoke_at: Final = datetime.now(timezone.utc) + timedelta(seconds=grace_seconds) - await _deprecated_verification_token_table(prisma_client).upsert( + table: Final = ( + tx.litellm_deprecatedverificationtoken + if tx is not None + else _deprecated_verification_token_table(prisma_client) + ) + await table.upsert( where={"token": old_token_hash}, data={ "create": { @@ -5540,6 +5554,24 @@ async def _insert_deprecated_key( ) +async def _lock_and_validate_source_key_change( + tx: "Prisma", token: str, allowed_routes: Sequence[str] | None = None +) -> None: + from litellm.proxy.management_endpoints.scim.source_endpoints import lock_provisioning_token + + await lock_provisioning_token(tx, token) + key: Final = await tx.litellm_verificationtoken.find_unique(where={"token": token}) + if key is None: + raise HTTPException(409, "The key changed during the request; retry with the current key") + if tuple(allowed_routes or ()) == ("/scim/*",): + return + source: Final = await tx.litellm_scimsource.find_unique(where={"key_hash": token}) + if source is not None: + raise HTTPException( + 409, "A provisioning source token cannot be regenerated or have its SCIM restriction removed" + ) + + async def _reject_source_bound_key_change(prisma_client: PrismaClient, key: LiteLLM_VerificationToken) -> None: if tuple(key.allowed_routes or ()) != ("/scim/*",): return @@ -5661,27 +5693,26 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) - await _persist_deleted_verification_tokens( - keys=[key_in_db], - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=litellm_changed_by, - ) - - # If grace period set, insert deprecated key so old key remains valid - await _insert_deprecated_key( - prisma_client=prisma_client, - old_token_hash=hashed_api_key, - new_token_hash=new_token_hash, - grace_period=data.grace_period if data else None, - ) - - updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table( - VerificationTokenRepository(prisma_client) - ).update( - where={"token": hashed_api_key}, - data=with_settings_updated_at(jsonified_update_data), - ) + async with prisma_client.tx() as tx: + await _lock_and_validate_source_key_change(tx, hashed_api_key) + await _persist_deleted_verification_tokens( + keys=[key_in_db], + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=litellm_changed_by, + tx=tx, + ) + await _insert_deprecated_key( + prisma_client=prisma_client, + old_token_hash=hashed_api_key, + new_token_hash=new_token_hash, + grace_period=data.grace_period if data else None, + tx=tx, + ) + updated_token: Final = await tx.litellm_verificationtoken.update( + where={"token": hashed_api_key}, + data=with_settings_updated_at(jsonified_update_data), + ) updated_token_dict: Final[dict[str, object]] = dict(updated_token) if updated_token is not None else {} updated_token_dict["key"] = new_token updated_token_dict["token_id"] = updated_token_dict.pop("token") diff --git a/litellm/proxy/management_endpoints/scim/source_endpoints.py b/litellm/proxy/management_endpoints/scim/source_endpoints.py index fb7136f0a12..0bd78ccb195 100644 --- a/litellm/proxy/management_endpoints/scim/source_endpoints.py +++ b/litellm/proxy/management_endpoints/scim/source_endpoints.py @@ -4,7 +4,7 @@ from typing import Annotated, Final from uuid import uuid4 from fastapi import APIRouter, Depends, HTTPException -from prisma import Json +from prisma import Json, Prisma from prisma.types import ( LiteLLM_SCIMSourceCreateInput, LiteLLM_SCIMSourceOrderByInput, @@ -62,6 +62,10 @@ async def list_sources(auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth return tuple(source_response(source) for source in sources) +async def lock_provisioning_token(tx: Prisma, token_hash: str) -> None: + await tx.execute_raw('SELECT 1 FROM "LiteLLM_VerificationToken" WHERE token = $1 FOR UPDATE', token_hash) + + @router.post("", response_model=SCIMSourceResponse, status_code=201) async def create_source( data: SCIMSourceCreate, auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)] @@ -71,6 +75,7 @@ async def create_source( token_hash: Final = hash_token(data.provisioning_token.get_secret_value()) source_filter: Final[LiteLLM_SCIMSourceWhereUniqueInput] = {"key_hash": token_hash} async with client.tx() as tx: + await lock_provisioning_token(tx, token_hash) key: Final = await VerificationTokenRepository(SimpleNamespace(db=tx)).table.find_unique( where=LiteLLM_VerificationTokenWhereUniqueInput(token=token_hash) ) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py index 9de7f0c1110..c60cc2d9118 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_source_endpoints.py @@ -34,6 +34,7 @@ def source_database(monkeypatch: pytest.MonkeyPatch): client: Final = MagicMock(spec=PrismaClient) tx: Final = client.tx.return_value.__aenter__.return_value monkeypatch.setattr(proxy_server, "prisma_client", client) + tx.execute_raw = AsyncMock() tx.litellm_scimsource.find_unique = AsyncMock(return_value=None) tx.litellm_verificationtoken.find_unique = AsyncMock(return_value=SimpleNamespace(allowed_routes=["/scim/*"])) tx.litellm_accessgrouptable.find_many = AsyncMock(return_value=[]) @@ -239,3 +240,18 @@ async def test_source_mapping_accepts_access_groups_across_query_batches( ) assert result.group_mappings[0].access_group_ids == group_ids assert tx.litellm_accessgrouptable.find_many.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("current_key", [None, SimpleNamespace(allowed_routes=["/*"])]) +async def test_source_creation_revalidates_key_after_waiting_for_mutation(monkeypatch, current_key): + tx = source_database(monkeypatch) + + async def complete_concurrent_mutation(*args): + tx.litellm_verificationtoken.find_unique.return_value = current_key + + tx.execute_raw.side_effect = complete_concurrent_mutation + with pytest.raises(HTTPException) as denied: + await create_source(SCIMSourceCreate(display_name="Source", tenant_id=TENANT, provisioning_token="test-token"), ADMIN) + assert denied.value.status_code == 400 + tx.litellm_scimsource.create.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 265d781a7e4..fad0f290462 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -12755,6 +12755,10 @@ def _make_regenerate_mock_prisma(): return_value=None ) mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data) + mock_prisma_client.tx = MagicMock() + mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db + mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=_make_regenerate_existing_key()) return mock_prisma_client @@ -14823,6 +14827,11 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha return_value=None ) mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data) + mock_prisma_client.tx = MagicMock() + mock_prisma_client.tx.return_value.__aenter__.return_value = mock_prisma_client.db + mock_prisma_client.db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key) + mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -21154,3 +21163,79 @@ async def test_unbound_scim_key_can_still_regenerate(): ) assert result.key is not None client.db.litellm_verificationtoken.update.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("routes", [None, [], ["/*"], ["/scim/*", "/key/info"]]) +async def test_key_update_endpoint_preserves_source_token_restriction(monkeypatch, routes): + from types import SimpleNamespace + + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn + + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + _wire_update_key_fn(monkeypatch, key) + client = proxy_server.prisma_client + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=SimpleNamespace(source_id="source")) + request = MagicMock() + request.query_params = {} + with pytest.raises(ProxyException) as denied: + await update_key_fn( + request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=routes), + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + ) + assert str(denied.value.code) == "409" + client.update_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_regeneration_rechecks_source_binding_before_transaction_writes(): + from types import SimpleNamespace + + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + tx = AsyncMock() + client.tx = MagicMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source") + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + tx.litellm_verificationtoken.find_unique.return_value = key + with _patch_regenerate_side_effects(): + with pytest.raises(HTTPException) as denied: + await _execute_virtual_key_regeneration( + prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=None, + user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, + user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), + ) + assert denied.value.status_code == 409 + tx.litellm_verificationtoken.update.assert_not_awaited() + tx.litellm_deletedverificationtoken.create_many.assert_not_awaited() + client.db.litellm_verificationtoken.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bound,routes,missing,expected", [(True, ["/*"], False, 409), (True, ["/scim/*"], False, None), (False, [], False, None), (False, [], True, 409)]) +async def test_route_write_rechecks_binding_and_preserves_supported_updates(bound, routes, missing, expected): + from types import SimpleNamespace + + from litellm.proxy.management_endpoints.key_management_endpoints import _update_key_row_with_soft_budget + + client = _make_regenerate_mock_prisma() + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"]}) + tx = client.db + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="source") if bound else None + tx.litellm_verificationtoken.find_unique.return_value = None if missing else key + tx.litellm_verificationtoken.update.return_value = key.model_copy(update={"allowed_routes": routes}) + request = UpdateKeyRequest(key=key.token, allowed_routes=routes) + write = _update_key_row_with_soft_budget(client, key.token, request, {"allowed_routes": routes}, key, "admin") + if expected: + with pytest.raises(HTTPException) as denied: + await write + assert denied.value.status_code == expected + tx.litellm_verificationtoken.update.assert_not_awaited() + else: + response = await write + assert response["data"]["allowed_routes"] == routes + tx.litellm_budgettable.update.assert_not_awaited() From f4d6091b6c49d2081ed56b1f44457d2a9be42879 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:56:46 -0700 Subject: [PATCH 4/5] test(scim): preserve rotation grace period transaction writes --- .../test_key_management_endpoints.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index fad0f290462..d011cd8cbee 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21239,3 +21239,22 @@ async def test_route_write_rechecks_binding_and_preserves_supported_updates(boun response = await write assert response["data"]["allowed_routes"] == routes tx.litellm_budgettable.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("transactional", [False, True]) +async def test_rotation_grace_period_is_written_to_the_selected_database(transactional): + from litellm.proxy.management_endpoints.key_management_endpoints import _insert_deprecated_key + + client = _make_regenerate_mock_prisma() + transaction = AsyncMock() + before = datetime.now(timezone.utc) + await _insert_deprecated_key(client, "old-token", "new-token", "1h", tx=transaction if transactional else None) + selected = transaction if transactional else client.db + unused = client.db if transactional else transaction + saved = selected.litellm_deprecatedverificationtoken.upsert.call_args.kwargs + assert saved["where"] == {"token": "old-token"} + assert saved["data"]["create"]["active_token_id"] == "new-token" + assert saved["data"]["update"]["active_token_id"] == "new-token" + assert before + timedelta(hours=1) <= saved["data"]["create"]["revoke_at"] <= datetime.now(timezone.utc) + timedelta(hours=1) + unused.litellm_deprecatedverificationtoken.upsert.assert_not_awaited() From 88f3d22eca5d0f9fa259e6eea664917b54f95716 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:29:11 -0700 Subject: [PATCH 5/5] fix(scim): include key permissions in source lifecycle transactions --- .../key_management_endpoints.py | 45 ++++++++---- .../object_permission_utils.py | 16 ++++- .../test_key_management_endpoints.py | 69 +++++++++++++++++++ 3 files changed, 113 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index aded71edd78..bc4ee68e011 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2437,16 +2437,22 @@ async def _update_key_row_with_soft_budget( async with prisma_client.tx() as tx: if "allowed_routes" in data.model_fields_set: await _lock_and_validate_source_key_change(tx, hashed_token, data.allowed_routes) + permission_values: Final = await _handle_update_object_permission( + data_json=dict(non_default_values), + existing_key_row=existing_key_row, + prisma_client=prisma_client, + tx=tx, + ) update_values: Final = ( await _apply_soft_budget_update( data=data, - non_default_values=non_default_values, + non_default_values=permission_values, db=tx, existing_key_row=existing_key_row, changed_by=changed_by, ) if "soft_budget" in data.model_fields_set - else non_default_values + else permission_values ) include_object_permission: Final[prisma.types.LiteLLM_VerificationTokenInclude] = {"object_permission": True} updated_row: Final = await tx.litellm_verificationtoken.update( @@ -2573,6 +2579,8 @@ async def _handle_update_object_permission( data_json: dict, existing_key_row: LiteLLM_VerificationToken, prisma_client: PrismaClient, + *, + tx: "Prisma | None" = None, ) -> dict: """Persist the requested object permission row and swap it for its id, only after the key policy allowed the write.""" if "object_permission" not in data_json: @@ -2582,6 +2590,7 @@ async def _handle_update_object_permission( data_json=data_json, existing_object_permission_id=existing_key_row.object_permission_id, prisma_client=prisma_client, + tx=tx, ) # Add the object_permission_id to data_json if one was created/updated @@ -3544,10 +3553,15 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, - existing_key_row=existing_key_row, - prisma_client=prisma_client, + uses_transaction: Final = bool(data.model_fields_set.intersection(("soft_budget", "allowed_routes"))) + update_values: Final = ( + non_default_values + if uses_transaction + else await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + ) ) changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name response: Final = ( @@ -3559,7 +3573,7 @@ async def update_key_fn( existing_key_row=existing_key_row, changed_by=changed_by, ) - if {"soft_budget", "allowed_routes"}.intersection(data.model_fields_set) + if uses_transaction else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key})) ) @@ -5678,14 +5692,6 @@ async def _execute_virtual_key_regeneration( request=data if data is not None else RegenerateKeyRequest(), ), ) - update_values: Final = await _handle_update_object_permission( - data_json=non_default_values, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) - update_data.update(update_values) - jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) - # Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash, # but their cached jwt_key_mapping entries still point at the old token (LIT-5379). jwt_mapping_cache_keys: Final = await get_jwt_key_mapping_cache_keys_for_token( @@ -5695,6 +5701,15 @@ async def _execute_virtual_key_regeneration( async with prisma_client.tx() as tx: await _lock_and_validate_source_key_change(tx, hashed_api_key) + update_values: Final = await _handle_update_object_permission( + data_json=non_default_values, + existing_key_row=key_in_db, + prisma_client=prisma_client, + tx=tx, + ) + jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object( + data={**update_data, **update_values} + ) await _persist_deleted_verification_tokens( keys=[key_in_db], prisma_client=prisma_client, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 569f2ebd278..68ac96a609d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -25,6 +25,7 @@ from litellm.repositories.prisma_protocols import DatabaseClient, TableActions from litellm.repositories.table_repositories import MCPServerRepository if TYPE_CHECKING: + from prisma import Prisma from prisma import models as prisma_models from litellm.proxy._types import ( @@ -84,6 +85,8 @@ async def prepare_object_permission_upsert( new_object_permission: Mapping[str, object], existing_object_permission_id: str | None, prisma_client: PrismaClient, + *, + tx: "Prisma | None" = None, ) -> ObjectPermissionUpsert: """ Read-and-merge half of an object permission upsert; performs no writes. @@ -101,7 +104,10 @@ async def prepare_object_permission_upsert( update cannot leave permission changes live. """ object_permission_id: Final = existing_object_permission_id or str(uuid.uuid4()) - existing_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique( + permission_table: Final = ( + tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table + ) + existing_object_permission: Final = await permission_table.find_unique( where={"object_permission_id": object_permission_id}, ) existing_fields: Final[dict[str, object]] = ( @@ -134,6 +140,8 @@ async def handle_update_object_permission_common( data_json: dict, existing_object_permission_id: str | None, prisma_client: PrismaClient | None, + *, + tx: "Prisma | None" = None, ) -> str | None: """ Common logic for handling object permission updates across organizations, teams, and keys. @@ -170,8 +178,12 @@ async def handle_update_object_permission_common( new_object_permission=new_object_permission if isinstance(new_object_permission, dict) else {}, existing_object_permission_id=existing_object_permission_id, prisma_client=prisma_client, + tx=tx, ) - created_object_permission_row: Final = await ObjectPermissionRepository(prisma_client).table.upsert( + permission_table: Final = ( + tx.litellm_objectpermissiontable if tx is not None else ObjectPermissionRepository(prisma_client).table + ) + created_object_permission_row: Final = await permission_table.upsert( where={"object_permission_id": upsert.object_permission_id}, data={ "create": upsert.record, diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index d011cd8cbee..2e0ce748bcd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21258,3 +21258,72 @@ async def test_rotation_grace_period_is_written_to_the_selected_database(transac assert saved["data"]["update"]["active_token_id"] == "new-token" assert before + timedelta(hours=1) <= saved["data"]["create"]["revoke_at"] <= datetime.now(timezone.utc) + timedelta(hours=1) unused.litellm_deprecatedverificationtoken.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["update", "regenerate"]) +async def test_source_binding_race_denial_preserves_object_permissions(monkeypatch, operation): + from types import SimpleNamespace + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn, _execute_virtual_key_regeneration + + key = _make_regenerate_existing_key().model_copy(update={"allowed_routes": ["/scim/*"], "object_permission_id": "permission-before"}) + if operation == "update": + _wire_update_key_fn(monkeypatch, key) + client = proxy_server.prisma_client + else: + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + events = [] + _record_object_permission_writes(client, events) + tx = AsyncMock() + client.tx = MagicMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_verificationtoken.find_unique.return_value = key + tx.litellm_scimsource.find_unique.return_value = SimpleNamespace(source_id="concurrent-source") + permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"]) + if operation == "update": + request = MagicMock() + request.query_params = {} + call = update_key_fn(request=request, data=UpdateKeyRequest(key=key.token, allowed_routes=["/*"], object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None) + else: + call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock()) + with _patch_regenerate_side_effects(): + with pytest.raises((HTTPException, ProxyException)) as denied: + await call + assert str(getattr(denied.value, "code", getattr(denied.value, "status_code", None))) == "409" + assert events == [] + client.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + tx.litellm_objectpermissiontable.upsert.assert_not_awaited() + tx.litellm_verificationtoken.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["update", "regenerate"]) +async def test_permission_update_joins_key_write_transaction(operation): + from litellm.proxy.management_endpoints.key_management_endpoints import _execute_virtual_key_regeneration, _update_key_row_with_soft_budget + + client = _make_regenerate_mock_prisma() + client.writer_db.litellm_scimsource.find_unique = AsyncMock(return_value=None) + key = _make_regenerate_existing_key().model_copy(update={"object_permission_id": "permission-before"}) + tx = AsyncMock() + client.tx.return_value.__aenter__.return_value = tx + tx.litellm_verificationtoken.find_unique.return_value = key + tx.litellm_scimsource.find_unique.return_value = None + tx.litellm_objectpermissiontable.find_unique.return_value = None + tx.litellm_objectpermissiontable.upsert.return_value = MagicMock(object_permission_id="permission-before") + tx.litellm_verificationtoken.update.side_effect = RuntimeError("token write failed") + permission = LiteLLM_ObjectPermissionBase(vector_stores=["vs-after"]) + if operation == "update": + call = _update_key_row_with_soft_budget(client, key.token, UpdateKeyRequest(key=key.token, allowed_routes=[], object_permission=permission), {"allowed_routes": [], "object_permission": permission.model_dump()}, key, "admin") + else: + call = _execute_virtual_key_regeneration(prisma_client=client, key_in_db=key, hashed_api_key=key.token, key="sk-original", data=RegenerateKeyRequest(object_permission=permission), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock()) + with _patch_regenerate_side_effects(): + with pytest.raises(RuntimeError, match="token write failed"): + await call + client.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + saved = tx.litellm_objectpermissiontable.upsert.await_args.kwargs + assert saved["where"] == {"object_permission_id": "permission-before"} + assert saved["data"]["update"]["vector_stores"] == ["vs-after"] + assert tx.litellm_verificationtoken.update.await_args.kwargs["data"]["object_permission_id"] == "permission-before" + assert client.tx.return_value.__aexit__.await_args.args[0] is RuntimeError