From 320a688bae823b4c334256d4eb37d619a484092f Mon Sep 17 00:00:00 2001
From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Date: Sat, 26 Sep 2026 17:12:00 -0700
Subject: [PATCH] feat(agents): integrate Entra identity registration and
authorization
---
.../migration.sql | 97 ++++
.../litellm_proxy_extras/schema.prisma | 57 +++
.../mcp_server/auth/managed_agent_access.py | 67 +++
.../mcp_server/auth/user_api_key_auth_mcp.py | 95 +++-
.../mcp_server/idp_token_exchange.py | 5 +
.../mcp_server/mcp_server_manager.py | 13 +-
.../mcp_server/rest_endpoints.py | 33 +-
.../mcp_server/ui_session_utils.py | 5 +-
litellm/proxy/_lazy_openapi_snapshot.json | 349 ++++++++++++-
litellm/proxy/_types.py | 25 +-
.../proxy/agent_endpoints/a2a_endpoints.py | 5 +
litellm/proxy/agent_endpoints/a2a_routing.py | 2 +-
.../proxy/agent_endpoints/agent_registry.py | 178 ++++---
.../auth/agent_access_groups.py | 21 +-
.../auth/agent_permission_handler.py | 164 +++++-
.../auth/managed_authorization.py | 232 +++++++++
litellm/proxy/agent_endpoints/endpoints.py | 109 +++-
litellm/proxy/agent_endpoints/identity.py | 17 +
.../proxy/agent_endpoints/identity_store.py | 224 ++++++++
.../proxy/agent_endpoints/managed_identity.py | 220 ++++++++
litellm/proxy/auth/auth_checks.py | 123 +++--
litellm/proxy/auth/handle_jwt.py | 114 ++++-
litellm/proxy/auth/user_api_key_auth.py | 54 +-
litellm/proxy/common_request_processing.py | 12 +-
.../proxy/common_utils/http_parsing_utils.py | 35 +-
.../common_utils/registry_read_through.py | 5 +-
.../proxy/hooks/proxy_track_cost_callback.py | 10 +-
litellm/proxy/image_endpoints/endpoints.py | 18 +-
litellm/proxy/litellm_pre_call_utils.py | 14 +-
.../sso/agent_subject_enrollment.py | 49 ++
litellm/proxy/management_endpoints/ui_sso.py | 22 +
litellm/proxy/proxy_server.py | 17 +-
litellm/proxy/schema.prisma | 57 +++
.../spend_management_endpoints.py | 32 +-
.../spend_tracking/spend_tracking_utils.py | 1 +
.../object_permission_repository.py | 7 +-
litellm/repositories/table_repositories.py | 23 +-
litellm/repositories/team_repository.py | 9 +-
litellm/repositories/user_repository.py | 7 +-
litellm/types/agents.py | 18 +-
litellm/types/proxy/agent_identity.py | 96 ++++
schema.prisma | 57 +++
.../auth/test_managed_agent_access.py | 301 +++++++++++
.../auth/test_user_api_key_auth_mcp.py | 94 +++-
.../mcp_server/test_discoverable_endpoints.py | 2 +-
.../mcp_server/test_idp_token_exchange.py | 11 +
.../mcp_server/test_mcp_server_manager.py | 44 +-
.../mcp_server/test_proxy_api_credentials.py | 4 +-
.../mcp_server/test_rest_endpoints.py | 28 +-
.../mcp_server/test_ui_session_utils.py | 4 +-
.../auth/test_agent_permission_handler.py | 287 ++++++++++-
.../auth/test_managed_authorization.py | 484 ++++++++++++++++++
.../agent_endpoints/test_a2a_endpoints.py | 7 +-
.../agent_endpoints/test_agent_registry.py | 426 ++++++++++++---
.../proxy/agent_endpoints/test_endpoints.py | 260 +++++++++-
.../proxy/agent_endpoints/test_identity.py | 27 +
.../agent_endpoints/test_identity_store.py | 402 +++++++++++++++
.../agent_endpoints/test_managed_identity.py | 257 ++++++++++
.../proxy/auth/test_auth_checks.py | 162 +++++-
.../proxy/auth/test_handle_jwt.py | 411 +++++++++++----
.../proxy/auth/test_user_api_key_auth.py | 200 ++++++++
.../common_utils/test_http_parsing_utils.py | 42 +-
.../test_registry_read_through.py | 34 +-
.../hooks/test_proxy_track_cost_callback.py | 31 ++
.../sso/test_agent_subject_enrollment.py | 114 +++++
.../test_mcp_management_endpoints.py | 4 +-
.../test_team_endpoints.py | 12 +-
.../proxy/management_endpoints/test_ui_sso.py | 3 +
.../proxy/proxy_server/test_proxy_config.py | 40 +-
.../test_spend_management_endpoints.py | 46 +-
.../test_spend_query_optimization.py | 10 +-
.../test_spend_tracking_utils.py | 36 ++
.../proxy/test_route_a2a_models.py | 3 +
.../_components/AgentIdentityDetails.test.tsx | 43 ++
.../_components/AgentIdentityDetails.tsx | 81 +++
.../_components/AgentIdentityFields.tsx | 259 ++++++++++
.../agents/_components/AgentsPanel.tsx | 6 +-
.../agents/_components/AgentsTable.test.tsx | 6 +
.../agents/_components/AgentsTableColumns.tsx | 1 +
.../add_agent_form.integration.test.tsx | 48 +-
.../_components/add_agent_form.test.tsx | 10 +-
.../agents/_components/add_agent_form.tsx | 27 +-
.../agents/_components/agent_config.ts | 29 +-
.../agents/_components/agent_identity.test.ts | 83 +++
.../agents/_components/agent_identity.ts | 101 ++++
.../agent_info.integration.test.tsx | 45 ++
.../agents/_components/agent_info.test.tsx | 27 +-
.../agents/_components/agent_info.tsx | 13 +-
.../agent_management/AgentSelector.test.tsx | 29 ++
.../agent_management/AgentSelector.tsx | 11 +-
.../src/components/agents/types.ts | 5 +
.../MCPServerSelector.test.tsx | 38 ++
.../MCPServerSelector.tsx | 14 +-
.../permissions/AgentPermissions.tsx | 2 +-
.../src/components/team/TeamInfo.test.tsx | 4 +-
.../src/components/team/TeamInfo.tsx | 10 +-
.../RequestLogsTableColumns.test.tsx | 14 +
.../view_logs/RequestLogsTableColumns.tsx | 2 +-
ui/litellm-dashboard/src/lib/http/schema.d.ts | 213 +++++++-
99 files changed, 6995 insertions(+), 610 deletions(-)
create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql
create mode 100644 litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py
create mode 100644 litellm/proxy/agent_endpoints/auth/managed_authorization.py
create mode 100644 litellm/proxy/agent_endpoints/identity.py
create mode 100644 litellm/proxy/agent_endpoints/identity_store.py
create mode 100644 litellm/proxy/agent_endpoints/managed_identity.py
create mode 100644 litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py
create mode 100644 litellm/types/proxy/agent_identity.py
create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py
create mode 100644 tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_identity.py
create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_identity_store.py
create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
create mode 100644 tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py
create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx
create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx
create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql
new file mode 100644
index 00000000000..06cf03b26b5
--- /dev/null
+++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql
@@ -0,0 +1,97 @@
+-- AlterTable
+ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "enabled" BOOLEAN NOT NULL DEFAULT true,
+ADD COLUMN IF NOT EXISTS "execution_mode" TEXT NOT NULL DEFAULT 'autonomous',
+ADD COLUMN IF NOT EXISTS "identity_managed" BOOLEAN NOT NULL DEFAULT false;
+
+-- AlterTable
+ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "billing_agent_id" TEXT;
+
+-- CreateTable
+CREATE TABLE IF NOT EXISTS "LiteLLM_AgentIdentity" (
+ "agent_id" TEXT NOT NULL,
+ "active" BOOLEAN NOT NULL DEFAULT true,
+ "provider" TEXT NOT NULL,
+ "issuer" TEXT NOT NULL,
+ "tenant_id" TEXT NOT NULL,
+ "client_id" TEXT NOT NULL,
+ "service_principal_id" TEXT,
+ "required_roles" TEXT[] DEFAULT ARRAY[]::TEXT[],
+ "required_scopes" TEXT[] DEFAULT ARRAY['user_impersonation']::TEXT[],
+ "revision" TEXT NOT NULL,
+ "last_authenticated_at" TIMESTAMP(3),
+
+ CONSTRAINT "LiteLLM_AgentIdentity_pkey" PRIMARY KEY ("agent_id")
+);
+
+-- CreateTable
+CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgentIdentity" (
+ "binding_id" TEXT NOT NULL,
+ "agent_id" TEXT,
+ "provider" TEXT NOT NULL,
+ "issuer" TEXT NOT NULL,
+ "tenant_id" TEXT NOT NULL,
+ "client_id" TEXT NOT NULL,
+
+ CONSTRAINT "LiteLLM_RetiredAgentIdentity_pkey" PRIMARY KEY ("binding_id")
+);
+
+-- CreateTable
+CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgent" (
+ "original_agent_id" TEXT NOT NULL,
+ "retired_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
+
+ CONSTRAINT "LiteLLM_RetiredAgent_pkey" PRIMARY KEY ("original_agent_id")
+);
+
+-- CreateTable
+CREATE TABLE IF NOT EXISTS "LiteLLM_VerifiedSubject" (
+ "subject_id" TEXT NOT NULL,
+ "issuer" TEXT NOT NULL,
+ "tenant_id" TEXT NOT NULL,
+ "oid" TEXT NOT NULL,
+ "kind" TEXT NOT NULL DEFAULT 'human',
+ "user_id" TEXT,
+ "verified_via" TEXT NOT NULL DEFAULT 'sso_interactive',
+ "verified_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
+
+ CONSTRAINT "LiteLLM_VerifiedSubject_pkey" PRIMARY KEY ("subject_id")
+);
+
+-- CreateIndex
+CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_AgentIdentity"("provider", "tenant_id", "client_id");
+
+-- CreateIndex
+CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_issuer_service_principal_id_key" ON "LiteLLM_AgentIdentity"("issuer", "service_principal_id");
+
+-- CreateIndex
+CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_RetiredAgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_RetiredAgentIdentity"("provider", "tenant_id", "client_id");
+
+-- CreateIndex
+CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_user_id_idx" ON "LiteLLM_VerifiedSubject"("user_id");
+
+-- CreateIndex
+CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_issuer_tenant_id_oid_key" ON "LiteLLM_VerifiedSubject"("issuer", "tenant_id", "oid");
+
+-- AddForeignKey
+DO $$
+BEGIN
+ IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentIdentity_agent_id_fkey') THEN
+ ALTER TABLE "LiteLLM_AgentIdentity" ADD CONSTRAINT "LiteLLM_AgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE;
+ END IF;
+END $$;
+
+-- AddForeignKey
+DO $$
+BEGIN
+ IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_RetiredAgentIdentity_agent_id_fkey') THEN
+ ALTER TABLE "LiteLLM_RetiredAgentIdentity" ADD CONSTRAINT "LiteLLM_RetiredAgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE;
+ END IF;
+END $$;
+
+-- AddForeignKey
+DO $$
+BEGIN
+ IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_user_id_fkey') THEN
+ ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE CASCADE ON UPDATE CASCADE;
+ END IF;
+END $$;
diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
index 69c63d9ecd6..8535d1004f7 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
@@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
+ identity_managed Boolean @default(false)
+ enabled Boolean @default(true)
+ execution_mode String @default("autonomous")
+ identity LiteLLM_AgentIdentity?
+ retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
+model LiteLLM_AgentIdentity {
+ agent_id String @id
+ active Boolean @default(true)
+ agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ service_principal_id String?
+ required_roles String[] @default([])
+ required_scopes String[] @default(["user_impersonation"])
+ revision String @default(uuid())
+ last_authenticated_at DateTime?
+ @@unique([provider, tenant_id, client_id])
+ @@unique([issuer, service_principal_id])
+}
+
+model LiteLLM_RetiredAgentIdentity {
+ binding_id String @id @default(uuid())
+ agent_id String?
+ agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ @@unique([provider, tenant_id, client_id])
+}
+
+model LiteLLM_RetiredAgent {
+ original_agent_id String @id
+ retired_at DateTime @default(now())
+}
+
+model LiteLLM_VerifiedSubject {
+ subject_id String @id @default(uuid())
+ issuer String
+ tenant_id String
+ oid String
+ kind String @default("human")
+ user_id String?
+ user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
+ verified_via String @default("sso_interactive")
+ verified_at DateTime @default(now())
+ @@unique([issuer, tenant_id, oid])
+ @@index([user_id])
+}
+
+
+
+
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
+ verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@@ -674,6 +730,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
+ billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?
diff --git a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py
new file mode 100644
index 00000000000..dd7390a784d
--- /dev/null
+++ b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py
@@ -0,0 +1,67 @@
+from types import MappingProxyType
+from typing import Final
+
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
+from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
+from litellm.types.proxy.agent_identity import AgentIdentityFailure
+
+
+async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
+ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
+
+ human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
+ return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
+
+
+async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
+ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
+
+ agent: Final = auth.managed_agent_policy
+ if agent is None:
+ return ()
+
+ try:
+ base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
+ ceilings: Final = await resolve_managed_agent_ceilings(agent)
+ expanded: Final = tuple(
+ frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
+ for ceiling in ceilings
+ )
+ own: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
+ context: Final = auth.managed_agent_context
+ if context is None or context.mode == "autonomous":
+ return tuple(sorted(own))
+ if context.user_id is None:
+ return ()
+ human: Final = await _delegated_resource_subject(context.user_id)
+ allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human)
+ return tuple(sorted(own.intersection(allowed)))
+ except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
+ )
+
+
+async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
+ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
+
+ if server_id not in await managed_agent_servers(auth):
+ return []
+ try:
+ own: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
+ context: Final = auth.managed_agent_context
+ if context is None or context.mode == "autonomous":
+ return own
+ if context.user_id is None:
+ return []
+ human: Final = await _delegated_resource_subject(context.user_id)
+ human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human)
+ if own is None:
+ return human_tools
+ return own if human_tools is None else sorted(frozenset(own).intersection(human_tools))
+ except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
+ )
diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
index a93ffaeac9f..ec92bbba677 100644
--- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
+++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
@@ -67,7 +67,6 @@ from litellm.repositories.table_repositories import (
AgentsRepository,
MCPServerRepository,
)
-from litellm.repositories.user_repository import UserRepository
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
@@ -1086,7 +1085,7 @@ class MCPRequestHandler:
assert_never(identity.subject_type)
@staticmethod
- async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
+ async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
@@ -1111,6 +1110,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
+ check_db_only=requires_fresh_policy,
)
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
@@ -1119,6 +1119,7 @@ class MCPRequestHandler:
if user_object is not None and object_permission is None and user_object.object_permission_id:
object_permission = await get_object_permission(
object_permission_id=user_object.object_permission_id,
+ check_db_only=requires_fresh_policy,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
@@ -1147,6 +1148,7 @@ class MCPRequestHandler:
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
admitted.mcp_admitted_user_subject = True
+ admitted.requires_fresh_policy = requires_fresh_policy
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
# several teams under its own identity, so without this a cross-team user outruns every team's
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
@@ -1597,6 +1599,11 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
+ if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
+ from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
+
+ return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
+
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
try:
@@ -1606,7 +1613,7 @@ class MCPRequestHandler:
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return MCPServerAccess(
- server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
+ server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
)
# Get allowed servers from key and team
@@ -1703,7 +1710,7 @@ class MCPRequestHandler:
if user_api_key_auth and user_api_key_auth.agent_id:
agent_capped: Final = _agent_capped_servers(
allowed_mcp_servers,
- await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
+ await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
)
if agent_capped is not None:
@@ -1829,10 +1836,12 @@ class MCPRequestHandler:
scoped.object_permission = auth.object_permission
scoped.object_permission_id = auth.object_permission_id
scoped.access_group_ids = auth.access_group_ids
+ scoped.requires_fresh_policy = auth.requires_fresh_policy
+ scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
return scoped
@staticmethod
- async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
+ async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
direct grants, plus every team they are a live roster member of.
@@ -1886,6 +1895,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(auth and auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
@@ -1941,11 +1951,11 @@ class MCPRequestHandler:
roster instead of by grant charged unrelated teams' buckets)."""
return [
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
- for source in await MCPRequestHandler._admitted_subject_sources(auth)
+ for source in await MCPRequestHandler.admitted_subject_sources(auth)
]
@staticmethod
- async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
+ async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
"""Union of what each of the admitted subject's sources reaches, each answered by the
canonical resolver so no rule is reimplemented for this caller shape."""
reachable: Final[set[str]] = set()
@@ -2007,7 +2017,7 @@ class MCPRequestHandler:
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
@staticmethod
- async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
+ async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
sources that actually grant that server.
@@ -2088,6 +2098,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not team_obj:
@@ -2171,6 +2182,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
@staticmethod
@@ -2219,12 +2231,17 @@ class MCPRequestHandler:
if not user_api_key_auth:
return None
+ if user_api_key_auth.managed_agent_policy is not None:
+ from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools
+
+ return await managed_agent_tools(server_id, user_api_key_auth)
+
try:
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
if _is_mcp_admitted_user_subject(user_api_key_auth):
- return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
+ return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
# Get key and team object permissions (already loaded in main auth flow)
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
@@ -2334,7 +2351,7 @@ class MCPRequestHandler:
if user_api_key_auth.agent_id:
# Pre-fetch agent object_permission once to avoid a duplicate DB query.
agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
- agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server(
+ agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(
server_id=server_id,
user_api_key_auth=user_api_key_auth,
agent_object_permission=agent_obj_perm,
@@ -2456,6 +2473,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not raw_server_ids:
return []
@@ -2502,6 +2520,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@@ -2550,7 +2569,7 @@ class MCPRequestHandler:
"""Get allowed MCP servers a caller inherits from the team it is pinned to.
Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not
- fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``,
+ fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``,
and each of those sources pins a single ``team_id`` before reaching this point. Keeping the
fan-out here as well would be a second multi-team path to drift from that one.
"""
@@ -2568,7 +2587,7 @@ class MCPRequestHandler:
which must NOT silently gain the union across every team the user belongs to), and it covers
each single-source auth an admitted subject fans out into — those pin a team_id, so they land
on the first branch. The admitted subject itself never reaches here: it resolves per source
- in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
+ in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
resolves to no teams exactly as before."""
if user_api_key_auth is None or not user_api_key_auth.team_id:
return []
@@ -2596,6 +2615,7 @@ class MCPRequestHandler:
user_id_upsert=False,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises
verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e)
@@ -2667,6 +2687,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
return []
@@ -2680,6 +2701,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers)
@@ -2716,6 +2738,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
raise unloadable from e
@@ -2961,7 +2984,9 @@ class MCPRequestHandler:
return None
user_id: Final = user_api_key_auth.user_id
- object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
+ object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(
+ user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy
+ )
if object_permission_id is None:
return None
@@ -2971,6 +2996,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if object_permission is None:
raise ValueError(
@@ -2979,7 +3005,9 @@ class MCPRequestHandler:
return object_permission
@staticmethod
- async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
+ async def _user_object_permission_id(
+ user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
+ ) -> str | None:
"""The permission row this human's user row links to, or None when they link none.
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
@@ -2988,16 +3016,23 @@ class MCPRequestHandler:
whether someone is entitled is the state that existed before this level, so it places no
ceiling. Only a link we DID resolve can make the caller deny.
"""
+ from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import user_api_key_cache
cache_key: Final = user_object_permission_id_cache_key(user_id)
try:
- cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
+ cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key)
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
return None
if isinstance(cached, str) and cached:
return cached
- user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
+ user_row: Final = await get_user_object(
+ user_id=user_id,
+ prisma_client=prisma_client,
+ user_api_key_cache=user_api_key_cache,
+ user_id_upsert=False,
+ check_db_only=check_db_only,
+ )
linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None
object_permission_id: Final = linked if isinstance(linked, str) and linked else None
await user_api_key_cache.async_set_cache(
@@ -3006,7 +3041,9 @@ class MCPRequestHandler:
ttl=get_management_object_ttl(user_api_key_cache),
)
return object_permission_id
- except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
+ except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
+ if check_db_only:
+ raise HTTPException(503, "User policy is unavailable") from e
verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e)
return None
@@ -3119,9 +3156,13 @@ class MCPRequestHandler:
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
- channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
+ channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot
disagree."""
- if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
+ if (
+ user_api_key_auth is None
+ or user_api_key_auth.mcp_explicit_grants_only
+ or not user_api_key_has_admin_view(user_api_key_auth)
+ ):
return False
object_permission: Final = user_api_key_auth.object_permission
credential_scoped: Final = (
@@ -3302,6 +3343,10 @@ class MCPRequestHandler:
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
+ if user_api_key_auth.managed_agent_policy is not None:
+ permission: Final = user_api_key_auth.managed_agent_policy.object_permission
+ return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None
+
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
@@ -3319,7 +3364,7 @@ class MCPRequestHandler:
)
@staticmethod
- async def _get_allowed_mcp_servers_for_agent(
+ async def get_allowed_mcp_servers_for_agent(
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> list[str]:
@@ -3363,7 +3408,7 @@ class MCPRequestHandler:
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
except Exception as e:
- if isinstance(e, UnloadableEntitlementError):
+ if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
return []
@@ -3390,7 +3435,7 @@ class MCPRequestHandler:
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
@staticmethod
- async def _get_agent_tool_permissions_for_server(
+ async def get_agent_tool_permissions_for_server(
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
@@ -3432,9 +3477,9 @@ class MCPRequestHandler:
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
- return list(agent_tools) if agent_tools else None
+ return list(agent_tools) if agent_tools is not None else None
except Exception as e:
- if isinstance(e, UnloadableEntitlementError):
+ if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
return None
@@ -3548,6 +3593,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@@ -3591,6 +3637,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
verbose_logger.debug("team_obj is None")
diff --git a/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py b/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py
index a437df17e6a..15632eb4783 100644
--- a/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py
+++ b/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py
@@ -170,6 +170,11 @@ async def identity_from_subject_token(
return _refusal_for(denied, denied.message)
except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures
return _refusal_for(denied, denied)
+ if result.get("agent_id") is not None:
+ return SubjectTokenRefusal(
+ error="invalid_request",
+ description="Agent tokens require direct JWT authentication; this exchange supports users only",
+ )
user_id: Final = result["user_id"]
if user_id is None:
return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows")
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 31896d9ddc5..cdd1dcb64c5 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -3418,7 +3418,9 @@ class MCPServerManager:
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
which precomputes both for its fallback path, does not compute them twice."""
- if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None:
+ if user_api_key_auth is not None and (
+ user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only
+ ):
return set()
if allow_all_server_ids is None:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
@@ -3467,9 +3469,14 @@ class MCPServerManager:
2. If admin and no object_permission, return all servers
3. Otherwise, use standard permission checks
"""
+ if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
+ managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
+ return managed if access is None else [server for server in managed if server in access.server_ids]
+
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
+ explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only)
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
@@ -3501,7 +3508,7 @@ class MCPServerManager:
# only keys without their own mcp_servers list get submitted servers unioned in.
submitted_server_ids: Final = (
[]
- if has_explicit_object_permission
+ if has_explicit_object_permission or explicit_grants_only
else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
)
@@ -3570,7 +3577,7 @@ class MCPServerManager:
return [
server_id
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
- if scope is None or server_id == scope
+ if not explicit_grants_only and (scope is None or server_id == scope)
]
async def resolve_toolset_tool_permissions(
diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
index 02694f110b1..bae1043691b 100644
--- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py
@@ -5,7 +5,7 @@ from dataclasses import dataclass
from datetime import datetime
from traceback import walk_tb
from types import MappingProxyType
-from typing import TYPE_CHECKING, Any, Final, Literal
+from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
from uuid import uuid4
import anyio
@@ -14,6 +14,7 @@ import httpx2
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from pydantic import ValidationError
from starlette.datastructures import Headers
+from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
@@ -62,7 +63,27 @@ if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.types.mcp import MCPAuth
-from litellm.types.utils import CallTypes
+from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
+
+
+class _MCPModelMetadata(TypedDict):
+ model_group: ReadOnly[str]
+
+
+def _stamp_mcp_tool_metadata(logging_obj: "LiteLLMLoggingObj | None", server_id: str, tool_name: str) -> None:
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
+
+ if logging_obj is None:
+ return
+ server: Final = global_mcp_server_manager.get_mcp_server_by_id(
+ server_id
+ ) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
+ metadata: Final[StandardLoggingMCPToolCall] = {
+ "name": tool_name,
+ "mcp_server_name": server.name if server is not None else server_id,
+ }
+ logging_obj.model_call_details["mcp_tool_call_metadata"] = metadata
+
MCP_AVAILABLE: bool = True
try:
@@ -1143,6 +1164,12 @@ if MCP_AVAILABLE:
},
)
+ data["model"] = f"MCP: {tool_name}"
+ model_metadata: Final[_MCPModelMetadata] = {
+ **(data.get("metadata") or MappingProxyType({})),
+ "model_group": f"MCP: {tool_name}",
+ }
+ data["metadata"] = model_metadata
proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
_request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below
try:
@@ -1176,6 +1203,8 @@ if MCP_AVAILABLE:
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
+ _stamp_mcp_tool_metadata(logging_obj, server_id, tool_name)
+
# Resolve allowed MCP servers with IP filtering
(
allowed_mcp_servers,
diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py
index 107a4818de1..901259c18ad 100644
--- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py
+++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py
@@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
+ check_db_only=user_api_key_auth.requires_fresh_policy,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
@@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey
)
try:
- admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id)
+ admitted: Final = await MCPRequestHandler.reload_admitted_user(
+ user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
+ )
except HTTPException as e:
verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail)
return None
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index ded4db6d2aa..f374ba09d07 100644
--- a/litellm/proxy/_lazy_openapi_snapshot.json
+++ b/litellm/proxy/_lazy_openapi_snapshot.json
@@ -2378,6 +2378,19 @@
"title": "Agent Name",
"type": "string"
},
+ "enabled": {
+ "title": "Enabled",
+ "type": "boolean"
+ },
+ "execution_mode": {
+ "enum": [
+ "autonomous",
+ "delegated",
+ "both"
+ ],
+ "title": "Execution Mode",
+ "type": "string"
+ },
"extra_headers": {
"anyOf": [
{
@@ -2392,6 +2405,16 @@
],
"title": "Extra Headers"
},
+ "identity": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/EntraIdentityConfig"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
"kill_switch": {
"anyOf": [
{
@@ -2470,8 +2493,7 @@
}
},
"required": [
- "agent_name",
- "agent_card_params"
+ "agent_name"
],
"title": "AgentConfig",
"type": "object"
@@ -2521,6 +2543,91 @@
"title": "AgentExtension",
"type": "object"
},
+ "AgentIdentityBinding": {
+ "properties": {
+ "active": {
+ "default": true,
+ "title": "Active",
+ "type": "boolean"
+ },
+ "agent_id": {
+ "title": "Agent Id",
+ "type": "string"
+ },
+ "client_id": {
+ "title": "Client Id",
+ "type": "string"
+ },
+ "issuer": {
+ "title": "Issuer",
+ "type": "string"
+ },
+ "last_authenticated_at": {
+ "anyOf": [
+ {
+ "format": "date-time",
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "Last Authenticated At"
+ },
+ "provider": {
+ "const": "microsoft_entra",
+ "title": "Provider",
+ "type": "string"
+ },
+ "required_roles": {
+ "default": [],
+ "items": {
+ "type": "string"
+ },
+ "title": "Required Roles",
+ "type": "array"
+ },
+ "required_scopes": {
+ "default": [
+ "user_impersonation"
+ ],
+ "items": {
+ "type": "string"
+ },
+ "title": "Required Scopes",
+ "type": "array"
+ },
+ "revision": {
+ "title": "Revision",
+ "type": "string"
+ },
+ "service_principal_id": {
+ "anyOf": [
+ {
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "Service Principal Id"
+ },
+ "tenant_id": {
+ "title": "Tenant Id",
+ "type": "string"
+ }
+ },
+ "required": [
+ "agent_id",
+ "provider",
+ "tenant_id",
+ "client_id",
+ "issuer",
+ "revision"
+ ],
+ "title": "AgentIdentityBinding",
+ "type": "object"
+ },
"AgentInterface": {
"description": "Declares a combination of a target URL and a transport protocol.",
"properties": {
@@ -2972,6 +3079,21 @@
],
"title": "Created By"
},
+ "enabled": {
+ "default": true,
+ "title": "Enabled",
+ "type": "boolean"
+ },
+ "execution_mode": {
+ "default": "autonomous",
+ "enum": [
+ "autonomous",
+ "delegated",
+ "both"
+ ],
+ "title": "Execution Mode",
+ "type": "string"
+ },
"extra_headers": {
"anyOf": [
{
@@ -2986,6 +3108,26 @@
],
"title": "Extra Headers"
},
+ "identity": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/AgentIdentityBinding"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
+ "identity_managed": {
+ "default": false,
+ "title": "Identity Managed",
+ "type": "boolean"
+ },
+ "jwt_auth_configured": {
+ "default": false,
+ "title": "Jwt Auth Configured",
+ "type": "boolean"
+ },
"keys": {
"anyOf": [
{
@@ -3441,6 +3583,61 @@
"title": "DailySpendMetadata",
"type": "object"
},
+ "EntraIdentityConfig": {
+ "additionalProperties": false,
+ "properties": {
+ "client_id": {
+ "title": "Client Id",
+ "type": "string"
+ },
+ "provider": {
+ "const": "microsoft_entra",
+ "title": "Provider",
+ "type": "string"
+ },
+ "required_roles": {
+ "default": [],
+ "items": {
+ "type": "string"
+ },
+ "title": "Required Roles",
+ "type": "array"
+ },
+ "required_scopes": {
+ "default": [
+ "user_impersonation"
+ ],
+ "description": "Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.",
+ "items": {
+ "type": "string"
+ },
+ "title": "Required Scopes",
+ "type": "array"
+ },
+ "service_principal_id": {
+ "anyOf": [
+ {
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "Service Principal Id"
+ },
+ "tenant_id": {
+ "title": "Tenant Id",
+ "type": "string"
+ }
+ },
+ "required": [
+ "provider",
+ "tenant_id",
+ "client_id"
+ ],
+ "title": "EntraIdentityConfig",
+ "type": "object"
+ },
"HTTPAuthSecurityScheme": {
"description": "Defines a security scheme using HTTP authentication.",
"properties": {
@@ -3590,6 +3787,54 @@
"title": "MakeAgentsPublicRequest",
"type": "object"
},
+ "ManagedAgentIdentityStatus": {
+ "properties": {
+ "enabled": {
+ "default": true,
+ "title": "Enabled",
+ "type": "boolean"
+ },
+ "execution_mode": {
+ "default": "autonomous",
+ "enum": [
+ "autonomous",
+ "delegated",
+ "both"
+ ],
+ "title": "Execution Mode",
+ "type": "string"
+ },
+ "identity": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/AgentIdentityBinding"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
+ "identity_managed": {
+ "default": false,
+ "title": "Identity Managed",
+ "type": "boolean"
+ },
+ "last_authenticated_at": {
+ "anyOf": [
+ {
+ "format": "date-time",
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "Last Authenticated At"
+ }
+ },
+ "title": "ManagedAgentIdentityStatus",
+ "type": "object"
+ },
"MetricWithMetadata": {
"properties": {
"api_key_breakdown": {
@@ -3790,6 +4035,19 @@
"title": "Agent Name",
"type": "string"
},
+ "enabled": {
+ "title": "Enabled",
+ "type": "boolean"
+ },
+ "execution_mode": {
+ "enum": [
+ "autonomous",
+ "delegated",
+ "both"
+ ],
+ "title": "Execution Mode",
+ "type": "string"
+ },
"extra_headers": {
"anyOf": [
{
@@ -3804,6 +4062,16 @@
],
"title": "Extra Headers"
},
+ "identity": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/EntraIdentityConfig"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
"kill_switch": {
"anyOf": [
{
@@ -4325,6 +4593,36 @@
]
}
},
+ "/v1/agents/identity/providers": {
+ "get": {
+ "operationId": "get_agent_identity_providers_v1_agents_identity_providers_get",
+ "responses": {
+ "200": {
+ "content": {
+ "application/json": {
+ "schema": {
+ "items": {
+ "type": "string"
+ },
+ "title": "Response Get Agent Identity Providers V1 Agents Identity Providers Get",
+ "type": "array"
+ }
+ }
+ },
+ "description": "Successful Response"
+ }
+ },
+ "security": [
+ {
+ "APIKeyHeader": []
+ }
+ ],
+ "summary": "Get Agent Identity Providers",
+ "tags": [
+ "agents"
+ ]
+ }
+ },
"/v1/agents/make_public": {
"post": {
"description": "Make multiple agents publicly discoverable\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/make_public\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"agent_ids\": [\"123e4567-e89b-12d3-a456-426614174000\", \"123e4567-e89b-12d3-a456-426614174001\"]\n }'\n```\n\nExample Response:\n```json\n{\n \"agent_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"agent_name\": \"my-custom-agent\",\n \"litellm_params\": {\n \"make_public\": true\n },\n \"agent_card_params\": {...},\n \"created_at\": \"2025-11-15T10:30:00Z\",\n \"updated_at\": \"2025-11-15T10:35:00Z\",\n \"created_by\": \"user123\",\n \"updated_by\": \"user123\"\n}\n```",
@@ -4576,6 +4874,53 @@
]
}
},
+ "/v1/agents/{agent_id}/identity": {
+ "get": {
+ "operationId": "get_agent_identity_status_v1_agents__agent_id__identity_get",
+ "parameters": [
+ {
+ "in": "path",
+ "name": "agent_id",
+ "required": true,
+ "schema": {
+ "title": "Agent Id",
+ "type": "string"
+ }
+ }
+ ],
+ "responses": {
+ "200": {
+ "content": {
+ "application/json": {
+ "schema": {
+ "$ref": "#/components/schemas/ManagedAgentIdentityStatus"
+ }
+ }
+ },
+ "description": "Successful Response"
+ },
+ "422": {
+ "content": {
+ "application/json": {
+ "schema": {
+ "$ref": "#/components/schemas/HTTPValidationError"
+ }
+ }
+ },
+ "description": "Validation Error"
+ }
+ },
+ "security": [
+ {
+ "APIKeyHeader": []
+ }
+ ],
+ "summary": "Get Agent Identity Status",
+ "tags": [
+ "agents"
+ ]
+ }
+ },
"/v1/agents/{agent_id}/kill_switch": {
"post": {
"description": "Fire the agent's configured kill switch webhook. Proxy admin only.\n\nLiteLLM only makes the configured HTTP call and reports what came back; it\ndoes not change the agent's state in LiteLLM. Returns 200 when the webhook\nanswered 2xx, 502 with the same result body otherwise. Every attempt is\nwritten to the audit log as a `kill_switch_fired` row against the agent.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch\" \\\n -H \"Authorization: Bearer \"\n```",
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 14aa42afefd..641be7078b7 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -27,7 +27,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
validate_langfuse_span_scope_value,
validate_no_callback_env_reference,
)
-from litellm.types.agents import AgentCaller
+from litellm.types.agents import AgentCaller, AgentResponse
from litellm.types.integrations.compression_interception import (
CompressionSavingsMetadata,
)
@@ -46,6 +46,7 @@ from litellm.types.mcp import (
MCPTransportType,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
+from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.types.proxy.carried_budget_state import (
OrgBudgetSnapshot,
TeamBudgetSnapshot,
@@ -569,6 +570,7 @@ class LiteLLMRoutes(enum.Enum):
"/agents",
"/a2a/{agent_id}",
"/a2a/{agent_id}/message/send",
+ "/v1/a2a/{agent_id}/message/send",
"/a2a/{agent_id}/message/stream",
"/a2a/{agent_id}/.well-known/agent-card.json",
)
@@ -3313,6 +3315,8 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union
# or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization.
mcp_admitted_user_subject: bool = Field(default=False, exclude=True)
+ requires_fresh_policy: bool = Field(default=False, exclude=True)
+ mcp_explicit_grants_only: bool = Field(default=False, exclude=True)
# team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP
# servers through several teams at once and therefore has no single team_id for the limiter to
# key off. Server-only and stripped from validated input for the same reason as the marker
@@ -3337,6 +3341,11 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
"user id."
),
)
+ invoked_agent_id: str | None = Field(default=None, exclude=True)
+ agent_invocation_cost: float | None = Field(default=None, exclude=True)
+ billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
+ managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
+ managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True)
agent_caller: AgentCaller | None = Field(
default=None,
exclude=True,
@@ -3374,11 +3383,18 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# path via post-construction assignment. Strip it from any validated input (constructor
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
+ values.pop("requires_fresh_policy", None)
+ values.pop("mcp_explicit_grants_only", None)
values.pop("mcp_source_team_rpm_limits", None)
values.pop("mcp_session_resource_server_id", None)
values.pop("mcp_toolset_id", None)
values.pop("via_virtual_key", None)
values.pop("agent_caller", None)
+ values.pop("managed_agent_context", None)
+ values.pop("managed_agent_policy", None)
+ values.pop("invoked_agent_id", None)
+ values.pop("agent_invocation_cost", None)
+ values.pop("billing_agent_policy", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):
@@ -4068,6 +4084,11 @@ class SpendLogsRouterMetadata(TypedDict):
class SpendLogsMetadata(TypedDict):
+ actor_agent_id: ReadOnly[NotRequired[str | None]]
+ target_agent_id: ReadOnly[NotRequired[str | None]]
+ billing_agent_id: ReadOnly[NotRequired[str | None]]
+ agent_execution_mode: ReadOnly[NotRequired[str | None]]
+ verified_human_user_id: ReadOnly[NotRequired[str | None]]
autorouter_baseline_observation: ReadOnly[str | None]
"""
Specific metadata k,v pairs logged to spendlogs for easier cost tracking
@@ -4131,6 +4152,7 @@ class SpendLogsPayload(TypedDict):
model_id: str | None
model_group: str | None
mcp_namespaced_tool_name: str | None
+ billing_agent_id: ReadOnly[NotRequired[str | None]]
agent_id: str | None
api_base: str
user: str
@@ -5053,6 +5075,7 @@ class JWTAuthBuilderResult(TypedDict):
org_id: str | None
team_membership: LiteLLM_TeamMembership | None
jwt_claims: dict # Decoded JWT token claims (avoids re-decoding)
+ managed_agent_context: ReadOnly[NotRequired[ManagedAgentContext | None]]
agent_id: ReadOnly[str | None]
diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py
index 2a189a76545..5589005cc0b 100644
--- a/litellm/proxy/agent_endpoints/a2a_endpoints.py
+++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py
@@ -723,6 +723,8 @@ async def invoke_agent_a2a(
detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.",
)
+ user_api_key_dict.invoked_agent_id = agent.agent_id
+
_enforce_inbound_trace_id(agent, request)
# Get backend URL and agent name
@@ -760,6 +762,8 @@ async def invoke_agent_a2a(
if "metadata" not in body:
body["metadata"] = {}
body["metadata"]["agent_id"] = agent.agent_id
+ body["metadata"]["model_group"] = f"a2a_agent/{agent_name}"
+ body["metadata"]["model_info"] = {"id": agent.agent_id}
body["agent_id"] = agent.agent_id
body.update(
@@ -863,6 +867,7 @@ async def invoke_agent_a2a(
# results written by the unified_guardrail hook are captured.
logging_obj._defer_async_logging = True
response = await asend_message(
+ model=f"a2a_agent/{agent_name}",
request=a2a_request,
api_base=agent_url,
litellm_params=litellm_params,
diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py
index 8a795214750..c57315ebc21 100644
--- a/litellm/proxy/agent_endpoints/a2a_routing.py
+++ b/litellm/proxy/agent_endpoints/a2a_routing.py
@@ -57,7 +57,7 @@ async def route_a2a_agent_request(
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
- if not is_admin:
+ if not is_admin or agent.identity_managed:
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
agent_id=agent.agent_id,
user_api_key_auth=user_api_key_dict,
diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py
index 3e775d7648e..ef7128a0e8e 100644
--- a/litellm/proxy/agent_endpoints/agent_registry.py
+++ b/litellm/proxy/agent_endpoints/agent_registry.py
@@ -6,6 +6,7 @@ from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, NamedTuple, Protocol, TypedDict
+from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly
@@ -14,13 +15,16 @@ from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy.agent_endpoints.kill_switch import restore_kill_switch
+from litellm.proxy.agent_endpoints.managed_identity import managed_write_fields, raise_identity_failure
from litellm.proxy.management_helpers.object_permission_utils import (
- handle_update_object_permission_common,
+ prepare_object_permission_upsert,
)
from litellm.proxy.utils import PrismaClient
+from litellm.repositories.base_repository import is_unique_violation
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository
from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest
+from litellm.types.proxy.agent_identity import AgentIdentityFailure
if TYPE_CHECKING:
from prisma import models as prisma_models
@@ -135,6 +139,39 @@ def object_permission_table(
return table
+class AgentPermissionWrite(TypedDict, total=False):
+ create: ReadOnly[Mapping[str, object]]
+ update: ReadOnly[Mapping[str, object]]
+
+
+async def _permission_write(
+ incoming: Mapping[str, object],
+ existing_id: str | None,
+ client: PrismaClient,
+) -> AgentPermissionWrite | None:
+ raw: Final = incoming.get("object_permission")
+ if raw is None:
+ return None
+ permission: Final = _AGENT_PARAMS_ADAPTER.validate_python(raw)
+ prepared: Final = await prepare_object_permission_upsert(permission, existing_id, client)
+ if existing_id is None:
+ created: Final[AgentPermissionWrite] = {"create": prepared.record}
+ return created
+ updated: Final[AgentPermissionWrite] = {"update": prepared.record}
+ return updated
+
+
+def _managed_fields(
+ incoming: Mapping[str, object],
+ existing: AgentResponse | None,
+ updated_by: str,
+) -> Mapping[str, object]:
+ result: Final = managed_write_fields(incoming, existing, updated_by)
+ if isinstance(result, AgentIdentityFailure):
+ raise_identity_failure(result, 400)
+ return result
+
+
def _dump_agent_params(raw: Mapping[str, object]) -> dict[str, object]:
model_dump: Final[Callable[[], dict[str, object]] | None] = getattr(raw, "model_dump", None)
if model_dump is not None:
@@ -552,11 +589,7 @@ class AgentRegistry:
agent_card_params_dict: Final[dict[str, object]] = _dump_agent_params(agent_card_params_obj)
agent_card_params: Final[str] = safe_dumps(agent_card_params_dict)
- # Handle object_permission (MCP tool access for agent)
- object_permission_id: str | None = None
- if agent.get("object_permission") is not None:
- agent_copy: Final = dict(agent)
- object_permission_id = await handle_update_object_permission_common(agent_copy, None, prisma_client)
+ permission_write: Final = await _permission_write(agent, None, prisma_client)
# Serialize static_headers
static_headers_obj: Final = agent.get("static_headers")
@@ -583,8 +616,8 @@ class AgentRegistry:
create_data["extra_headers"] = extra_headers_val
if access_group_ids_val is not None:
create_data["access_group_ids"] = tuple(dict.fromkeys(access_group_ids_val))
- if object_permission_id is not None:
- create_data["object_permission_id"] = object_permission_id
+ if permission_write is not None:
+ create_data["object_permission"] = permission_write
for rate_field in (
"tpm_limit",
@@ -598,31 +631,46 @@ class AgentRegistry:
# Create agent in DB
created_agent: Final = await agents_table(prisma_client).create(
- data=create_data,
- include={"object_permission": True},
+ data={**create_data, **_managed_fields(agent, None, created_by)},
+ include={"object_permission": True, "identity": True},
)
- created_agent_dict: Final = created_agent.model_dump()
- if created_agent.object_permission is not None:
- try:
- created_agent_dict["object_permission"] = created_agent.object_permission.model_dump()
- except Exception:
- created_agent_dict["object_permission"] = created_agent.object_permission.dict()
- return AgentResponse(**created_agent_dict)
+ return AgentResponse.model_validate(created_agent.model_dump())
+ except HTTPException:
+ raise
except Exception as e:
- raise Exception(f"Error adding agent to DB: {e}")
+ if is_unique_violation(e):
+ raise HTTPException(409, "Agent name or Entra application is already registered") from e
+ raise
async def delete_agent_from_db(self, agent_id: str, prisma_client: PrismaClient) -> Mapping[str, object]:
"""
Delete an agent from the database
"""
- try:
- deleted_agent: Final = await agents_table(prisma_client).delete(where={"agent_id": agent_id})
+ from prisma.types import (
+ LiteLLM_AgentsTableWhereUniqueInput,
+ LiteLLM_RetiredAgentCreateInput,
+ LiteLLM_RetiredAgentUpsertInput,
+ LiteLLM_RetiredAgentWhereUniqueInput,
+ LiteLLM_VerificationTokenWhereInput,
+ )
+
+ where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id}
+ async with prisma_client.tx() as tx:
+ existing: Final = await tx.litellm_agentstable.find_unique(where=where)
+ if existing is None:
+ raise ValueError(f"Agent not found, passed agent_id={agent_id}")
+ if existing.identity_managed:
+ history_where: Final[LiteLLM_RetiredAgentWhereUniqueInput] = {"original_agent_id": agent_id}
+ history_create: Final = LiteLLM_RetiredAgentCreateInput(original_agent_id=agent_id)
+ history_data: Final[LiteLLM_RetiredAgentUpsertInput] = {"create": history_create, "update": {}}
+ await tx.litellm_retiredagent.upsert(where=history_where, data=history_data)
+ keys_where: Final[LiteLLM_VerificationTokenWhereInput] = {"agent_id": agent_id}
+ await tx.litellm_verificationtoken.delete_many(where=keys_where)
+ deleted_agent: Final = await tx.litellm_agentstable.delete(where=where)
if deleted_agent is None:
raise ValueError(f"Agent not found, passed agent_id={agent_id}")
- return dict(deleted_agent)
- except Exception as e:
- raise Exception(f"Error deleting agent from DB: {e}")
+ return deleted_agent.model_dump()
async def patch_agent_in_db(
self,
@@ -646,7 +694,9 @@ class AgentRegistry:
The patched agent
"""
try:
- existing_record: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
+ existing_record: Final = await agents_table(prisma_client).find_unique(
+ where={"agent_id": agent_id}, include={"identity": True}
+ )
if existing_record is None:
raise Exception(f"Agent with ID {agent_id} not found")
existing_agent: Final[Mapping[str, object]] = dict(existing_record)
@@ -683,37 +733,31 @@ class AgentRegistry:
if "extra_headers" in agent:
extra_headers_value: Final = agent.get("extra_headers")
update_data["extra_headers"] = extra_headers_value if extra_headers_value is not None else []
- if agent.get("object_permission") is not None:
- agent_copy: Final = dict(augment_agent)
- existing_object_permission_id: Final = existing_record.object_permission_id
- object_permission_id: Final = await handle_update_object_permission_common(
- agent_copy,
- existing_object_permission_id,
- prisma_client,
- )
- if object_permission_id is not None:
- update_data["object_permission_id"] = object_permission_id
+ permission_write: Final = await _permission_write(
+ agent, existing_record.object_permission_id, prisma_client
+ )
+ if permission_write is not None:
+ update_data["object_permission"] = permission_write
# Patch agent in DB
patched_agent: Final = await agents_table(prisma_client).update(
where={"agent_id": agent_id},
data={
**update_data,
+ **_managed_fields(agent, AgentResponse.model_validate(existing_record.model_dump()), updated_by),
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
},
- include={"object_permission": True},
+ include={"object_permission": True, "identity": True},
)
if patched_agent is None:
raise ValueError(f"Agent not found, passed agent_id={agent_id}")
- patched_agent_dict: Final = patched_agent.model_dump()
- if patched_agent.object_permission is not None:
- try:
- patched_agent_dict["object_permission"] = patched_agent.object_permission.model_dump()
- except Exception:
- patched_agent_dict["object_permission"] = patched_agent.object_permission.dict()
- return AgentResponse(**patched_agent_dict)
+ return AgentResponse.model_validate(patched_agent.model_dump())
+ except HTTPException:
+ raise
except Exception as e:
- raise Exception(f"Error patching agent in DB: {e}")
+ if is_unique_violation(e):
+ raise HTTPException(409, "Agent name or Entra application is already registered") from e
+ raise
async def update_agent_in_db(
self,
@@ -733,7 +777,7 @@ class AgentRegistry:
# caller echoed back redacted (or omitted) rather than persisting
# the marker -- or nothing -- over the real stored credential.
existing_row: Final = await agents_table(prisma_client).find_unique(
- where={"agent_id": agent_id} # mutable-ok: prisma's query builder rejects a Mapping/MappingProxyType
+ where={"agent_id": agent_id}, include={"identity": True}
)
existing_litellm_params: Final = parse_agent_litellm_params(
existing_row.litellm_params if existing_row is not None else None
@@ -784,37 +828,35 @@ class AgentRegistry:
if _val is not None:
update_data[rate_field] = _val
- if agent.get("object_permission") is not None:
- existing_object_permission_id: Final = (
- existing_row.object_permission_id if existing_row is not None else None
- )
- agent_copy: Final = dict(agent)
- object_permission_id: Final = await handle_update_object_permission_common(
- agent_copy,
- existing_object_permission_id,
- prisma_client,
- )
- if object_permission_id is not None:
- update_data["object_permission_id"] = object_permission_id
+ permission_write: Final = await _permission_write(
+ agent, existing_row.object_permission_id if existing_row is not None else None, prisma_client
+ )
+ if permission_write is not None:
+ update_data["object_permission"] = permission_write
# Update agent in DB
updated_agent: Final = await agents_table(prisma_client).update(
where={"agent_id": agent_id},
- data=update_data,
- include={"object_permission": True},
+ data={
+ **update_data,
+ **_managed_fields(
+ agent,
+ AgentResponse.model_validate(existing_row.model_dump()) if existing_row else None,
+ updated_by,
+ ),
+ },
+ include={"object_permission": True, "identity": True},
)
if updated_agent is None:
raise ValueError(f"Agent not found, passed agent_id={agent_id}")
- updated_agent_dict: Final = updated_agent.model_dump()
- if updated_agent.object_permission is not None:
- try:
- updated_agent_dict["object_permission"] = updated_agent.object_permission.model_dump()
- except Exception:
- updated_agent_dict["object_permission"] = updated_agent.object_permission.dict()
- return AgentResponse(**updated_agent_dict)
+ return AgentResponse.model_validate(updated_agent.model_dump())
+ except HTTPException:
+ raise
except Exception as e:
- raise Exception(f"Error updating agent in DB: {e}")
+ if is_unique_violation(e):
+ raise HTTPException(409, "Agent name or Entra application is already registered") from e
+ raise
@staticmethod
async def get_all_agents_from_db(
@@ -826,12 +868,12 @@ class AgentRegistry:
try:
agents_from_db: Final = await agents_table(prisma_client).find_many(
order={"created_at": "desc"},
- include={"object_permission": True},
+ include={"object_permission": True, "identity": True},
)
agents: Final[list[dict[str, object]]] = []
for agent in agents_from_db:
- agent_dict = dict(agent)
+ agent_dict = agent.model_dump()
# object_permission is eagerly loaded via include above
if agent.object_permission is not None:
try:
diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py
index 49e5407ff88..6f0cc0393ad 100644
--- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py
+++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py
@@ -1,13 +1,16 @@
import asyncio
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
-from typing import Final, TypeAlias
+from typing import TYPE_CHECKING, Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_AccessGroupTable
+if TYPE_CHECKING:
+ from litellm.types.agents import AgentResponse
+
AccessGroupIds: TypeAlias = tuple[str, ...]
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
@@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ()) if agent is not None else ()
-async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
+async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
from litellm.proxy.auth.auth_checks import get_access_object
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
@@ -47,6 +50,7 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=check_db_only,
)
except HTTPException as e:
verbose_proxy_logger.warning(
@@ -73,3 +77,16 @@ async def resolve_agent_access_group_ceiling(
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
)
+
+
+async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
+ async def authoritative_group(group_id: str) -> LoadedAccessGroup:
+ return await _load_access_group(group_id, check_db_only=True)
+
+ async def manual_ids(_agent_id: str) -> AccessGroupIds:
+ return tuple(agent.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 ()
diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py
index 9fe74bfee3f..4b9238f1868 100644
--- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py
+++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py
@@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
+from types import MappingProxyType
from typing import Final, TypeAlias
+from fastapi import HTTPException
+
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
@@ -86,7 +89,9 @@ class AgentRequestHandler:
) -> AgentAccess:
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
- key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
+ if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
+ return await _managed_actor_agent_access(user_api_key_auth)
+ key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(user_api_key_auth)
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
@@ -104,13 +109,19 @@ class AgentRequestHandler:
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
@staticmethod
- async def _resolve_key_team_agent_access(
+ async def resolve_key_team_agent_access(
user_api_key_auth: UserAPIKeyAuth | None,
+ *,
+ strict: bool = False,
) -> AgentAccess:
try:
- key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
- team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
+ key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
+ team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
+ user_api_key_auth, strict=strict
+ )
except Exception as e:
+ if strict:
+ raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents: %s", e)
return UnrestrictedAgentAccess()
return _intersect_agent_access(key_access, team_access)
@@ -144,6 +155,33 @@ class AgentRequestHandler:
bool: True if agent is allowed, False otherwise
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
+ from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+ from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
+ from litellm.proxy.proxy_server import prisma_client
+ from litellm.types.proxy.agent_identity import AgentIdentityFailure
+
+ registered: Final = global_agent_registry.get_agent_by_id(agent_id)
+ if prisma_client is not None or (registered is not None and registered.identity_managed):
+ target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
+ if isinstance(target, AgentIdentityFailure):
+ raise_identity_failure(target)
+ if target is None and registered is not None and registered.identity_managed:
+ return False
+ if target is not None and target.identity_managed:
+ if (
+ not target.enabled
+ or target.identity is None
+ or not target.identity.active
+ or user_api_key_auth is None
+ ):
+ return False
+ fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True})
+ explicit: Final = await _granted_agent_ids(
+ fresh_auth,
+ _strict_agent_access,
+ build_effective_auth_contexts,
+ )
+ return target.agent_id in explicit
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
case UnrestrictedAgentAccess():
@@ -202,8 +240,10 @@ class AgentRequestHandler:
return team_obj.object_permission
@staticmethod
- async def _get_allowed_agents_for_key(
+ async def get_allowed_agents_for_key(
user_api_key_auth: UserAPIKeyAuth | None = None,
+ *,
+ strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a key.
@@ -237,24 +277,36 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
- tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
+ tuple(
+ await AgentRequestHandler._get_agents_from_access_groups(
+ list(declared_access_groups), check_db_only=strict
+ )
+ )
if declared_access_groups
else ()
)
unified_agents: Final = (
- tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
+ tuple(
+ await AgentRequestHandler._get_unified_access_group_agents(
+ list(key_access_group_ids), check_db_only=strict
+ )
+ )
if key_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
+ if strict:
+ raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
return UnrestrictedAgentAccess()
@staticmethod
async def _get_allowed_agents_for_team(
user_api_key_auth: UserAPIKeyAuth | None = None,
+ *,
+ strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a team.
@@ -263,7 +315,7 @@ class AgentRequestHandler:
2. Also includes agents from team's access_group_ids (unified access groups)
Fetches the team object once and reuses it for both permission sources.
- Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`.
+ Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`.
"""
if user_api_key_auth is None:
return UnrestrictedAgentAccess()
@@ -280,7 +332,7 @@ class AgentRequestHandler:
)
if not prisma_client:
- return UnrestrictedAgentAccess()
+ return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# Fetch the team object once for both permission sources
team_obj: Final = await get_team_object(
@@ -289,10 +341,11 @@ class AgentRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=strict,
)
if team_obj is None:
- return UnrestrictedAgentAccess()
+ return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# 1. Get agents from object_permission (native permissions)
object_permissions: Final = team_obj.object_permission
@@ -307,18 +360,28 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
- tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
+ tuple(
+ await AgentRequestHandler._get_agents_from_access_groups(
+ list(declared_access_groups), check_db_only=strict
+ )
+ )
if declared_access_groups
else ()
)
unified_agents: Final = (
- tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
+ tuple(
+ await AgentRequestHandler._get_unified_access_group_agents(
+ list(team_access_group_ids), check_db_only=strict
+ )
+ )
if team_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
+ if strict:
+ raise HTTPException(503, "Agent invocation policy is unavailable") from e
# litellm-dashboard is the default UI team and will never have agents;
# skip noisy warnings for it.
if user_api_key_auth.team_id != UI_TEAM_ID:
@@ -326,7 +389,9 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
@staticmethod
- def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
+ def _get_config_agent_ids_for_access_groups(
+ config_agents: Sequence[AgentResponse], access_groups: list[str]
+ ) -> set[str]:
"""
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
"""
@@ -339,7 +404,9 @@ class AgentRequestHandler:
return server_ids
@staticmethod
- async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
+ async def _get_db_agent_ids_for_access_groups(
+ prisma_client, access_groups: list[str], *, check_db_only: bool = False
+ ) -> set[str]:
"""
Helper to get agent_ids from DB agents that match any of the given access groups.
@@ -349,23 +416,27 @@ class AgentRequestHandler:
if not access_groups or prisma_client is None:
return set()
- agents: Final = await AgentsRepository(prisma_client).table.find_many(
+ agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
where={"agent_access_groups": {"hasSome": access_groups}}
)
return {agent.agent_id for agent in agents}
@staticmethod
- async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
+ async def _get_unified_access_group_agents(
+ access_group_ids: list[str], *, check_db_only: bool = False
+ ) -> list[str]:
"""
Resolve unified access group ids to agent IDs.
"""
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
- return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
+ return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
@staticmethod
async def _get_agents_from_access_groups(
access_groups: list[str],
+ *,
+ check_db_only: bool = False,
) -> list[str]:
"""
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
@@ -373,14 +444,13 @@ class AgentRequestHandler:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.proxy_server import prisma_client
- # Use the helper for config-loaded agents
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
global_agent_registry.agent_list, access_groups
)
# Use the helper for DB agents
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
- prisma_client, access_groups
+ prisma_client, access_groups, check_db_only=check_db_only
)
return list(config_agent_ids | db_agent_ids)
@@ -531,4 +601,58 @@ async def accessible_agents(
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
effective_contexts,
)
- return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
+ allowed: Final = await asyncio.gather(
+ *(
+ AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
+ for agent in agents
+ if agent.identity_managed
+ )
+ )
+ managed_ids: Final = frozenset(
+ agent.agent_id
+ for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
+ if permitted
+ )
+ return tuple(
+ agent
+ for agent in agents
+ if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
+ )
+
+
+async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
+ if auth.managed_agent_policy is not None:
+ return await _managed_actor_agent_access(auth)
+ return await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
+
+
+async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
+ agent: Final = auth.managed_agent_policy
+ if agent is None or not agent.object_permission:
+ return RestrictedAgentAccess(frozenset())
+ permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
+ own_auth: Final = UserAPIKeyAuth(object_permission=permission)
+ own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True))
+
+ from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
+
+ ceilings: Final = await resolve_managed_agent_ceilings(agent)
+ capped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
+ context: Final = auth.managed_agent_context
+ if context is None or context.mode == "autonomous":
+ return RestrictedAgentAccess(capped)
+ if context.user_id is None:
+ return RestrictedAgentAccess(frozenset())
+ human_ids: Final = await verified_human_agent_grants(context.user_id)
+ return RestrictedAgentAccess(capped.intersection(human_ids))
+
+
+async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]:
+ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
+
+ if user_id is None:
+ return frozenset()
+ human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
+ sources: Final = await MCPRequestHandler.admitted_subject_sources(human)
+ human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
+ return frozenset().union(*(_granted_ids(access) for access in human_access))
diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
new file mode 100644
index 00000000000..abcef353eef
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py
@@ -0,0 +1,232 @@
+from collections.abc import Mapping
+from itertools import product
+from types import MappingProxyType
+from typing import Annotated, Final
+
+from pydantic import Field, TypeAdapter, ValidationError
+
+from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
+from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext
+
+_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime"))
+_MANAGED_MODEL_ROUTES: Final = frozenset(
+ f"{prefix}/{operation}"
+ for prefix, operation in product(
+ ("", "/v1"),
+ (
+ "chat/completions",
+ "completions",
+ "embeddings",
+ "responses",
+ "messages",
+ "messages/count_tokens",
+ "images/generations",
+ "images/edits",
+ "audio/transcriptions",
+ "audio/speech",
+ "moderations",
+ "rerank",
+ "ocr",
+ ),
+ )
+) | frozenset(
+ (
+ "/openai/v1/responses",
+ "/v2/rerank",
+ "/claude_code_gateway/v1/messages",
+ "/claude_code_gateway/v1/messages/count_tokens",
+ "/cursor/chat/completions",
+ )
+)
+_MANAGED_MODEL_PATHS: Final = (
+ "/engines/{model:path}/chat/completions",
+ "/engines/{model:path}/completions",
+ "/engines/{model:path}/embeddings",
+ "/openai/deployments/{model:path}/chat/completions",
+ "/openai/deployments/{model:path}/completions",
+ "/openai/deployments/{model:path}/embeddings",
+ "/openai/deployments/{model:path}/images/generations",
+ "/openai/deployments/{model:path}/images/edits",
+ "/v1beta/models/{model_name:path}:countTokens",
+ "/v1beta/models/{model_name:path}:generateContent",
+ "/v1beta/models/{model_name:path}:streamGenerateContent",
+ "/models/{model_name:path}:countTokens",
+ "/models/{model_name:path}:generateContent",
+ "/models/{model_name:path}:streamGenerateContent",
+)
+_MANAGED_MCP_ROUTES: Final = tuple(
+ route for route in LiteLLMRoutes.mcp_inference_routes.value if route not in ("/token", "/introspect")
+)
+
+
+def managed_agent_route_allowed(route: str, method: str | None) -> bool:
+ from litellm.proxy.auth.route_checks import RouteChecks
+
+ if route in ("/agents", "/v1/agents"):
+ return method in (None, "GET", "HEAD")
+ if route in _MANAGED_REALTIME_ROUTES:
+ return method in (None, "GET")
+ if route in _MANAGED_MODEL_ROUTES or RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
+ return method in (None, "POST")
+ return RouteChecks.check_route_access(route, _MANAGED_MCP_ROUTES) or RouteChecks.check_route_access(
+ route, LiteLLMRoutes.agent_inference_routes.value
+ )
+
+
+def managed_inference_request(
+ route: str,
+ body: Mapping[str, object],
+ settings: Mapping[str, object],
+ cli_model: str | None,
+ path_model: object = None,
+ query_model: object = None,
+) -> dict[str, object]:
+ from litellm.proxy.auth.route_checks import RouteChecks
+
+ if route in _MANAGED_REALTIME_ROUTES:
+ model: Final = query_model or body.get("model")
+ if not isinstance(model, str) or not model:
+ raise_identity_failure(
+ AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
+ )
+ return {**body, "model": model}
+ if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
+ return dict(body)
+ from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
+
+ kind: Final = (
+ "image_generation"
+ if route.endswith("/images/generations")
+ else "image_edit"
+ if route.endswith("/images/edits")
+ else "moderation"
+ if route.endswith(("/moderations", "/audio/transcriptions"))
+ else "speech"
+ if route.endswith("/audio/speech")
+ else "body"
+ if route.endswith(("/rerank", "/messages/count_tokens"))
+ else "path"
+ if route.endswith(":countTokens")
+ else "completion"
+ )
+ endpoint_model: Final = path_model or (
+ query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
+ )
+ effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
+ if not isinstance(effective, str) or not effective:
+ raise_identity_failure(
+ AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
+ )
+ return {**body, "model": effective}
+
+
+async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None:
+ if auth.agent_id is None:
+ return
+ if store is None:
+ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
+
+ registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id)
+ if auth.managed_agent_context is not None or (
+ registered is not None and (registered.identity_managed or registered.identity is not None)
+ ):
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
+ )
+ return
+ agent: Final = await store.agent(auth.agent_id)
+ if isinstance(agent, AgentIdentityFailure):
+ raise_identity_failure(agent)
+ if agent is None:
+ retired: Final = await store.retired_agent(auth.agent_id)
+ if isinstance(retired, AgentIdentityFailure):
+ raise_identity_failure(retired)
+ if auth.managed_agent_context is not None or retired:
+ raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists"))
+ return
+ if not agent.identity_managed:
+ return
+ if auth.jwt_claims and auth.managed_agent_context is None:
+ raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity"))
+ failure: Final = actor_admission_failure(agent, auth.managed_agent_context)
+ if failure is not None:
+ raise_identity_failure(failure)
+ auth.managed_agent_policy = agent
+ auth.billing_agent_policy = agent
+ if auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated":
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
+
+ grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id)
+ if agent.agent_id not in grants:
+ raise_identity_failure(
+ AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent")
+ )
+
+
+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:
+ return AgentIdentityFailure(message="Agent execution is disabled")
+ if context is None:
+ return AgentIdentityFailure(message="This agent requires its bound identity provider token")
+ if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision:
+ return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
+ if agent.execution_mode not in (context.mode, "both"):
+ return AgentIdentityFailure(message="Agent is not enabled for this execution mode")
+ if context.mode == "delegated" and not context.user_id:
+ return AgentIdentityFailure(message="A verified human subject is required")
+ return None
+
+
+_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
+
+
+def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
+ model: Final = body.get("model")
+ if isinstance(model, str) and model.startswith("a2a/"):
+ return model.removeprefix("a2a/") or None
+ components: Final = tuple(route.strip("/").split("/"))
+ path: Final = components[1:] if components and components[0] == "v1" else components
+ return path[1] if len(path) >= 2 and path[0] == "a2a" else None
+
+
+async def prepare_agent_invocation(
+ auth: UserAPIKeyAuth, target_name: str, store: AgentIdentityStore | None, *, billable: bool = True
+) -> None:
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
+ from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
+
+ registered: Final = await get_agent_with_read_through(target_name)
+ if registered is None:
+ return
+ registered_managed: Final = registered.identity_managed or registered.identity is not None
+ if store is None and registered_managed:
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
+ )
+ target: Final = await store.agent(registered.agent_id) if store is not None else None
+ if isinstance(target, AgentIdentityFailure):
+ raise_identity_failure(target)
+ if target is None and registered_managed:
+ raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists"))
+ effective: Final = target if target is not None else registered
+ if not effective.identity_managed and auth.managed_agent_policy is None:
+ return
+ if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
+ raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
+ auth.invoked_agent_id = effective.agent_id
+ if auth.agent_id is None and effective.identity_managed:
+ auth.billing_agent_policy = effective
+ raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0
+ try:
+ fee: Final = _INVOCATION_COST.validate_python(raw_fee)
+ except ValidationError:
+ raise_identity_failure(
+ AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid")
+ )
+ auth.agent_invocation_cost = fee
diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py
index 28c82a715e0..e1b2ac63d51 100644
--- a/litellm/proxy/agent_endpoints/endpoints.py
+++ b/litellm/proxy/agent_endpoints/endpoints.py
@@ -16,6 +16,7 @@ from types import MappingProxyType
from typing import Annotated, Final, TypedDict
from fastapi import APIRouter, Depends, HTTPException, Query, Request
+from pydantic import ValidationError
from typing_extensions import ReadOnly, Required, assert_never
import litellm
@@ -47,6 +48,8 @@ from litellm.proxy.agent_endpoints.agent_search import (
search_agents,
)
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents
+from litellm.proxy.agent_endpoints.identity import reject_legacy_identity
+from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.kill_switch import (
KillSwitchAuditLogWriter,
KillSwitchHttpClient,
@@ -56,6 +59,7 @@ from litellm.proxy.agent_endpoints.kill_switch import (
fire_kill_switch,
redact_kill_switch,
)
+from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
@@ -72,6 +76,12 @@ from litellm.types.agents import (
PatchAgentRequest,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
+from litellm.types.proxy.agent_identity import (
+ AgentIdentityBinding,
+ AgentIdentityFailure,
+ EntraIdentityConfig,
+ ManagedAgentIdentityStatus,
+)
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
@@ -178,14 +188,21 @@ def _redact_sensitive_agent_fields(
virtual-key, header and kill-switch fields stripped entirely. The original
objects are not modified.
"""
+ from litellm.proxy.proxy_server import general_settings, jwt_handler
+
redacted: Final[list[AgentResponse]] = []
for agent in agents:
copy = agent.model_copy(deep=True)
+ copy.jwt_auth_configured = bool(
+ general_settings.get("enable_jwt_auth")
+ and (agent.identity is not None or jwt_handler.litellm_jwtauth.agent_id_jwt_field)
+ )
if not is_admin:
copy.static_headers = None
copy.extra_headers = None
copy.keys = None
copy.kill_switch = None
+ copy.identity = None
if copy.litellm_params:
copy.litellm_params = _redact_agent_litellm_params_dict(copy.litellm_params)
copy.kill_switch = redact_kill_switch(copy.kill_switch)
@@ -429,6 +446,71 @@ from litellm.proxy.agent_endpoints.agent_registry import (
)
+def _trusted_agent_issuers() -> tuple[str, ...]:
+ from litellm.proxy.proxy_server import general_settings, jwt_handler
+
+ if not general_settings.get("enable_jwt_auth"):
+ return ()
+ configured: Final = jwt_handler.litellm_jwtauth.issuers or ()
+ issuer: Final = os.getenv("JWT_ISSUER")
+ global_issuers: Final = (
+ (issuer,)
+ if issuer and os.getenv("JWT_AUDIENCE") and not any(item.issuer == issuer for item in configured)
+ else ()
+ )
+ return (
+ tuple(item.issuer for item in configured if item.audience and not item.disable_audience_validation)
+ + global_issuers
+ )
+
+
+def _validate_managed_identity_request(
+ request: AgentConfig | PatchAgentRequest, existing: AgentResponse | None = None
+) -> None:
+ raw: Final = request.get("identity") if "identity" in request else existing.identity if existing else None
+ if raw is None:
+ return
+ try:
+ identity: Final = raw if isinstance(raw, AgentIdentityBinding) else EntraIdentityConfig.model_validate(raw)
+ except ValidationError as exc:
+ raise HTTPException(400, "Invalid Entra identity configuration") from exc
+ if identity.issuer not in _trusted_agent_issuers():
+ raise HTTPException(400, "Configure trusted JWT issuer and audience validation for this Entra tenant first")
+ if request.get("execution_mode", existing.execution_mode if existing else "autonomous") != "autonomous":
+ if os.getenv("MICROSOFT_TENANT") != identity.tenant_id or not os.getenv("MICROSOFT_CLIENT_ID"):
+ raise HTTPException(400, "Delegated agents require Microsoft SSO for the same trusted tenant")
+
+
+@router.get("/v1/agents/identity/providers", response_model=tuple[str, ...], tags=("[beta] A2A Agents",))
+async def get_agent_identity_providers(
+ user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
+) -> tuple[str, ...]:
+ _check_agent_management_permission(user_api_key_dict)
+ return _trusted_agent_issuers()
+
+
+@router.get("/v1/agents/{agent_id}/identity", response_model=ManagedAgentIdentityStatus, tags=("[beta] A2A Agents",))
+async def get_agent_identity_status(
+ agent_id: str,
+ user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
+) -> ManagedAgentIdentityStatus:
+ from litellm.proxy.proxy_server import prisma_client
+
+ _check_agent_management_permission(user_api_key_dict)
+ agent: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
+ if isinstance(agent, AgentIdentityFailure):
+ raise_identity_failure(agent)
+ if agent is None:
+ raise HTTPException(404, "Agent not found")
+ return ManagedAgentIdentityStatus(
+ identity=agent.identity,
+ identity_managed=agent.identity_managed,
+ enabled=agent.enabled,
+ execution_mode=agent.execution_mode,
+ last_authenticated_at=agent.identity.last_authenticated_at if agent.identity else None,
+ )
+
+
@router.post(
"/v1/agents",
tags=["[beta] A2A Agents"],
@@ -490,6 +572,9 @@ async def create_agent(
# Get the user ID from the API key auth
created_by: Final = user_api_key_dict.user_id or "unknown"
+ _validate_managed_identity_request(request)
+ reject_legacy_identity(request.get("litellm_params"))
+
# check for naming conflicts
existing_agent: Final = AGENT_REGISTRY.get_agent_by_name(agent_name=request.get("agent_name"))
if existing_agent is not None:
@@ -591,7 +676,7 @@ async def get_agent_by_id(
if agent is None:
agent_row: Final = await agents_table(prisma_client).find_unique(
where={"agent_id": agent_id},
- include={"object_permission": True},
+ include={"object_permission": True, "identity": True},
)
if agent_row is not None:
agent_dict: Final = agent_row.model_dump()
@@ -680,13 +765,18 @@ async def update_agent(
try:
# Check if agent exists
- existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
+ existing_agent = await agents_table(prisma_client).find_unique(
+ where={"agent_id": agent_id}, include={"identity": True}
+ )
if existing_agent is not None:
- existing_agent = dict(existing_agent)
+ existing_agent = existing_agent.model_dump()
if existing_agent is None:
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
+ _validate_managed_identity_request(request, AgentResponse.model_validate(existing_agent))
+ reject_legacy_identity(request.get("litellm_params"))
+
# Get the user ID from the API key auth
updated_by: Final = user_api_key_dict.user_id or "unknown"
@@ -782,13 +872,18 @@ async def patch_agent(
try:
# Check if agent exists
- existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
+ existing_agent = await agents_table(prisma_client).find_unique(
+ where={"agent_id": agent_id}, include={"identity": True}
+ )
if existing_agent is not None:
- existing_agent = dict(existing_agent)
+ existing_agent = existing_agent.model_dump()
if existing_agent is None:
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
+ _validate_managed_identity_request(request, AgentResponse.model_validate(existing_agent))
+ reject_legacy_identity(request.get("litellm_params"))
+
# Get the user ID from the API key auth
updated_by: Final = user_api_key_dict.user_id or "unknown"
@@ -869,7 +964,9 @@ async def delete_agent(
try:
# Check if agent exists
- existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
+ existing_agent = await agents_table(prisma_client).find_unique(
+ where={"agent_id": agent_id}, include={"identity": True}
+ )
if existing_agent is not None:
existing_agent = dict[str, object](existing_agent)
diff --git a/litellm/proxy/agent_endpoints/identity.py b/litellm/proxy/agent_endpoints/identity.py
new file mode 100644
index 00000000000..c0e5a748144
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/identity.py
@@ -0,0 +1,17 @@
+from collections.abc import Mapping
+from typing import Final
+
+from fastapi import HTTPException
+
+LEGACY_IDENTITY_MESSAGE: Final = (
+ "litellm_params.identity is not supported: bind an Entra application through the top-level identity field"
+)
+
+
+def has_legacy_identity(params: Mapping[str, object] | None) -> bool:
+ return params is not None and "identity" in params
+
+
+def reject_legacy_identity(params: Mapping[str, object] | None) -> None:
+ if has_legacy_identity(params):
+ raise HTTPException(400, LEGACY_IDENTITY_MESSAGE)
diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py
new file mode 100644
index 00000000000..87512c779fe
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/identity_store.py
@@ -0,0 +1,224 @@
+from collections.abc import Mapping
+from datetime import datetime, timezone
+from typing import TYPE_CHECKING, Final
+
+from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject
+from litellm.repositories.table_repositories import (
+ AgentIdentityRepository,
+ AgentsRepository,
+ RetiredAgentIdentityRepository,
+ RetiredAgentRepository,
+ VerifiedSubjectRepository,
+)
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import (
+ AgentIdentityFailure,
+ ManagedAgentContext,
+ MicrosoftInteractiveSubject,
+ VerifiedHumanSubject,
+)
+
+if TYPE_CHECKING:
+ from prisma.models import LiteLLM_VerifiedSubject
+ from prisma.types import (
+ LiteLLM_AgentIdentityUpdateManyMutationInput,
+ LiteLLM_AgentIdentityWhereInput,
+ LiteLLM_AgentIdentityWhereUniqueInput,
+ LiteLLM_AgentsTableInclude,
+ LiteLLM_AgentsTableWhereUniqueInput,
+ LiteLLM_VerifiedSubjectCreateInput,
+ LiteLLM_VerifiedSubjectUpsertInput,
+ LiteLLM_VerifiedSubjectWhereUniqueInput,
+ )
+
+
+class AgentIdentityStore:
+ @classmethod
+ def from_client(cls, client: object) -> "AgentIdentityStore":
+ return cls(
+ AgentsRepository(client, use_writer=True),
+ AgentIdentityRepository(client, use_writer=True),
+ VerifiedSubjectRepository(client, use_writer=True),
+ RetiredAgentIdentityRepository(client, use_writer=True),
+ RetiredAgentRepository(client, use_writer=True),
+ )
+
+ def __init__(
+ self,
+ agents: AgentsRepository,
+ identities: AgentIdentityRepository,
+ humans: VerifiedSubjectRepository,
+ retired: RetiredAgentIdentityRepository | None = None,
+ retired_agents: RetiredAgentRepository | None = None,
+ ) -> None:
+ self.agents = agents
+ self.identities = identities
+ self.humans = humans
+ self.retired = retired
+ self.retired_agents = retired_agents
+
+ async def agent(self, agent_id: str) -> AgentResponse | AgentIdentityFailure | None:
+ try:
+ where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id}
+ include: Final[LiteLLM_AgentsTableInclude] = {
+ "identity": True,
+ "object_permission": True,
+ }
+ 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())
+ except Exception:
+ return AgentIdentityFailure(code="policy_unavailable", message="Agent policy could not be loaded")
+
+ async def unbound_client(self, where: "LiteLLM_AgentIdentityWhereUniqueInput") -> AgentIdentityFailure | None:
+ if self.retired is not None:
+ try:
+ retired: Final = await self.retired.table.find_unique(where=where)
+ except Exception:
+ return AgentIdentityFailure(
+ code="policy_unavailable", message="Retired agent identity could not be checked"
+ )
+ if retired is not None:
+ return AgentIdentityFailure(message="This agent identity binding has been retired")
+ return None
+
+ async def resolve_verified_claims(
+ self, claims: Mapping[str, object]
+ ) -> ManagedAgentContext | AgentIdentityFailure | None:
+ issuer: Final = claims.get("iss")
+ tenant: Final = claims.get("tid")
+ client: Final = claims.get("azp")
+ if not isinstance(issuer, str) or not isinstance(tenant, str) or not isinstance(client, str):
+ return None
+ where: Final[LiteLLM_AgentIdentityWhereUniqueInput] = {
+ "provider_tenant_id_client_id": {"provider": "microsoft_entra", "tenant_id": tenant, "client_id": client}
+ }
+ try:
+ row: Final = await self.identities.table.find_unique(where=where)
+ except Exception:
+ return AgentIdentityFailure(code="policy_unavailable", message="Agent identity could not be loaded")
+ if row is None:
+ return await self.unbound_client(where)
+ agent: Final = await self.agent(row.agent_id)
+ if isinstance(agent, AgentIdentityFailure):
+ return agent
+ if (
+ agent is None
+ or not agent.identity_managed
+ or not agent.enabled
+ or agent.identity is None
+ 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)
+ if isinstance(subject, AgentIdentityFailure):
+ return subject
+ if subject.kind == "application":
+ 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
+ and proven.kind == "human"
+ and proven.verified_via == "sso_interactive"
+ and proven.user_id is not None
+ else None
+ )
+ if human is None:
+ return AgentIdentityFailure(message="The delegated user must first sign in through trusted Microsoft SSO")
+ return ManagedAgentContext(
+ agent_id=agent.agent_id,
+ binding_revision=agent.identity.revision,
+ mode=subject.mode,
+ user_id=human.user_id,
+ subject_oid=subject.oid,
+ )
+
+ async def subject(
+ self, issuer: str, tenant_id: str, oid: object
+ ) -> "LiteLLM_VerifiedSubject | AgentIdentityFailure | None":
+ if not isinstance(oid, str):
+ return None
+ try:
+ where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = {
+ "issuer_tenant_id_oid": {"issuer": issuer, "tenant_id": tenant_id, "oid": oid}
+ }
+ return await self.humans.table.find_unique(where=where)
+ except Exception:
+ return AgentIdentityFailure(code="policy_unavailable", message="Subject classification is unavailable")
+
+ 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")
+ try:
+ return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None
+ except Exception:
+ return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")
+
+ async def record_authentication(self, context: ManagedAgentContext) -> AgentIdentityFailure | None:
+ try:
+ if context.binding_revision is None:
+ return AgentIdentityFailure(message="Agent authentication requires a binding revision")
+ where: Final[LiteLLM_AgentIdentityWhereInput] = {
+ "agent_id": context.agent_id,
+ "revision": context.binding_revision,
+ "active": True,
+ "agent": {"is": {"enabled": True, "identity_managed": True}},
+ }
+ data: Final[LiteLLM_AgentIdentityUpdateManyMutationInput] = {
+ "last_authenticated_at": datetime.now(timezone.utc)
+ }
+ count: Final = await self.identities.table.update_many(where=where, data=data)
+ if count != 1:
+ return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
+ return None
+ except Exception:
+ return AgentIdentityFailure(code="policy_unavailable", message="Agent authentication could not be recorded")
+
+ async def enroll_interactive_human(
+ self,
+ subject: MicrosoftInteractiveSubject,
+ user_id: str,
+ ) -> AgentIdentityFailure | None:
+ try:
+ where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = {
+ "issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": subject.tenant_id, "oid": subject.oid}
+ }
+ create_data: Final[LiteLLM_VerifiedSubjectCreateInput] = {
+ "issuer": subject.issuer,
+ "tenant_id": subject.tenant_id,
+ "oid": subject.oid,
+ "user_id": user_id,
+ "verified_via": "sso_interactive",
+ }
+ data: Final[LiteLLM_VerifiedSubjectUpsertInput] = {"create": create_data, "update": {}}
+ row: Final = await self.humans.table.upsert(where=where, data=data)
+ if row.kind != "human" or row.user_id != user_id or row.verified_via != "sso_interactive":
+ return AgentIdentityFailure(message="Microsoft subject is already bound to another local identity")
+ return None
+ except Exception:
+ return AgentIdentityFailure(
+ code="policy_unavailable", message="Microsoft subject enrollment is unavailable"
+ )
+
+
+async def resolve_managed_agent(
+ claims: Mapping[str, object],
+ client: object,
+) -> ManagedAgentContext | None:
+ from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
+
+ if client is None:
+ return None
+ result: Final = await AgentIdentityStore.from_client(client).resolve_verified_claims(claims)
+ if isinstance(result, AgentIdentityFailure):
+ raise_identity_failure(result)
+ return result
diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py
new file mode 100644
index 00000000000..260b74fcbd1
--- /dev/null
+++ b/litellm/proxy/agent_endpoints/managed_identity.py
@@ -0,0 +1,220 @@
+from collections.abc import Mapping
+from datetime import datetime
+from typing import Final, NoReturn, TypedDict
+from uuid import uuid4
+
+from fastapi import HTTPException
+from pydantic import TypeAdapter, ValidationError
+from typing_extensions import ReadOnly
+
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import (
+ AgentExecutionMode,
+ AgentIdentityBinding,
+ AgentIdentityFailure,
+ AgentSubject,
+ EntraIdentityConfig,
+)
+
+_MODE: Final = TypeAdapter(AgentExecutionMode)
+
+
+class IdentityFields(TypedDict, total=False):
+ provider: ReadOnly[str]
+ tenant_id: ReadOnly[str]
+ client_id: ReadOnly[str]
+ issuer: ReadOnly[str]
+ service_principal_id: ReadOnly[str | None]
+ required_roles: ReadOnly[tuple[str, ...]]
+ required_scopes: ReadOnly[tuple[str, ...]]
+ active: ReadOnly[bool]
+ revision: ReadOnly[str]
+ last_authenticated_at: ReadOnly[datetime | None]
+
+
+class IdentityUpsert(TypedDict):
+ create: ReadOnly[IdentityFields]
+ update: ReadOnly[IdentityFields]
+
+
+class IdentityRelationWrite(TypedDict, total=False):
+ create: ReadOnly[IdentityFields]
+ update: ReadOnly[IdentityFields]
+ upsert: ReadOnly[IdentityUpsert]
+
+
+class IdentityHistoryKey(TypedDict):
+ provider: ReadOnly[str]
+ tenant_id: ReadOnly[str]
+ client_id: ReadOnly[str]
+
+
+class IdentityHistoryWhere(TypedDict):
+ provider_tenant_id_client_id: ReadOnly[IdentityHistoryKey]
+
+
+class IdentityHistoryEntry(IdentityHistoryKey):
+ issuer: ReadOnly[str]
+
+
+class IdentityHistoryConnect(TypedDict):
+ where: ReadOnly[IdentityHistoryWhere]
+ create: ReadOnly[IdentityHistoryEntry]
+
+
+class IdentityHistoryWrite(TypedDict):
+ connectOrCreate: ReadOnly[IdentityHistoryConnect]
+
+
+class ManagedWriteFields(TypedDict, total=False):
+ enabled: ReadOnly[bool]
+ execution_mode: ReadOnly[AgentExecutionMode]
+ identity_managed: ReadOnly[bool]
+ identity: ReadOnly[IdentityRelationWrite]
+ retired_identities: ReadOnly[IdentityHistoryWrite]
+
+
+def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn:
+ raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message)
+
+
+def _configuration_failure(
+ identity: EntraIdentityConfig | AgentIdentityBinding | None,
+ mode: AgentExecutionMode,
+ enabling_without_binding: bool,
+) -> AgentIdentityFailure | None:
+ if identity is not None and mode != "delegated" and not identity.service_principal_id:
+ return AgentIdentityFailure(
+ message="Autonomous mode requires the Enterprise application service-principal object ID"
+ )
+ if enabling_without_binding and (
+ identity is None or isinstance(identity, AgentIdentityBinding) and not identity.active
+ ):
+ return AgentIdentityFailure(message="Bind an identity before enabling this managed agent")
+ return None
+
+
+def managed_write_fields(
+ incoming: Mapping[str, object],
+ existing: AgentResponse | None,
+ updated_by: str,
+) -> ManagedWriteFields | AgentIdentityFailure:
+ try:
+ identity: Final = (
+ EntraIdentityConfig.model_validate(incoming["identity"]) if incoming.get("identity") is not None else None
+ )
+ mode: Final = _MODE.validate_python(
+ incoming.get("execution_mode", existing.execution_mode if existing else "autonomous")
+ )
+ current_identity: Final = identity if "identity" in incoming else existing.identity if existing else None
+ failure: Final = _configuration_failure(
+ current_identity,
+ mode,
+ incoming.get("enabled") is True
+ and "identity" not in incoming
+ and bool(existing and existing.identity_managed),
+ )
+ if failure is not None:
+ return failure
+ empty: Final[ManagedWriteFields] = {}
+ identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty
+ result: Final[ManagedWriteFields] = {
+ **({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}),
+ **({"execution_mode": mode} if "execution_mode" in incoming else {}),
+ **identity_fields,
+ }
+ return result
+ except (ValidationError, ValueError) as exc:
+ return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}")
+
+
+def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields:
+ if identity is None:
+ unbind: Final[ManagedWriteFields] = {
+ **(
+ {"identity": {"update": {"active": False, "revision": str(uuid4()), "last_authenticated_at": None}}}
+ if existing and existing.identity
+ else {}
+ ),
+ **({"identity_managed": True, "enabled": False} if existing and existing.identity_managed else {}),
+ }
+ return unbind
+ if (
+ existing
+ and existing.identity
+ and existing.identity.active
+ and all(getattr(existing.identity, name) == value for name, value in identity.model_dump().items())
+ ):
+ unchanged: Final[ManagedWriteFields] = {}
+ return unchanged
+ binding: Final[IdentityFields] = {
+ "provider": identity.provider,
+ "tenant_id": identity.tenant_id,
+ "client_id": identity.client_id,
+ "service_principal_id": identity.service_principal_id,
+ "required_roles": identity.required_roles,
+ "required_scopes": identity.required_scopes,
+ "issuer": identity.issuer,
+ "active": True,
+ "revision": str(uuid4()),
+ "last_authenticated_at": None,
+ }
+ result: Final[ManagedWriteFields] = {
+ "retired_identities": {
+ "connectOrCreate": {
+ "where": {
+ "provider_tenant_id_client_id": {
+ "provider": identity.provider,
+ "tenant_id": identity.tenant_id,
+ "client_id": identity.client_id,
+ }
+ },
+ "create": {
+ "provider": identity.provider,
+ "issuer": identity.issuer,
+ "tenant_id": identity.tenant_id,
+ "client_id": identity.client_id,
+ },
+ }
+ },
+ "identity_managed": True,
+ "identity": {"upsert": {"create": binding, "update": binding}} if existing else {"create": binding},
+ }
+ return result
+
+
+def classify_agent_subject(
+ binding: AgentIdentityBinding,
+ claims: Mapping[str, object],
+ allowed_mode: AgentExecutionMode,
+) -> AgentSubject | AgentIdentityFailure:
+ if (claims.get("iss"), claims.get("tid"), claims.get("azp")) != (
+ binding.issuer,
+ binding.tenant_id,
+ binding.client_id,
+ ):
+ return AgentIdentityFailure(message="Token does not match the registered Entra application")
+ oid: Final = claims.get("oid")
+ if not isinstance(oid, str) or not oid:
+ return AgentIdentityFailure(message="Entra token must identify its object subject")
+ scope: Final = claims.get("scp")
+ 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")
+ if scope is not None and not isinstance(scope, str):
+ return AgentIdentityFailure(message="Invalid delegated scope claim")
+ if isinstance(scope, str) and scope:
+ if allowed_mode == "autonomous" or oid == binding.service_principal_id or claims.get("idtyp") == "app":
+ return AgentIdentityFailure(message="Delegated token contradicts the configured agent identity or mode")
+ granted_scopes: Final = frozenset(scope.split())
+ if not granted_scopes or not frozenset(binding.required_scopes).issubset(granted_scopes):
+ return AgentIdentityFailure(message="Token lacks the required delegated scopes")
+ return AgentSubject(kind="delegated_subject", oid=oid, mode="delegated")
+ if allowed_mode == "delegated" or oid != binding.service_principal_id or claims.get("idtyp") == "user":
+ return AgentIdentityFailure(message="Application token contradicts the configured agent identity or mode")
+ roles: Final = claims.get("roles", ())
+ if not isinstance(roles, (list, tuple)) or any(not isinstance(role, str) for role in roles):
+ return AgentIdentityFailure(message="Invalid application roles claim")
+ if not frozenset(binding.required_roles).issubset(roles):
+ return AgentIdentityFailure(message="Token lacks the required application roles")
+ return AgentSubject(kind="application", oid=oid, mode="autonomous")
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index 12d420141f1..11c7dbd64cf 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -1069,6 +1069,21 @@ async def common_checks(
code=status.HTTP_400_BAD_REQUEST,
)
+ if _model and valid_token is not None and valid_token.managed_agent_policy is not None:
+ managed_models: Final = (valid_token.managed_agent_policy.object_permission or MappingProxyType({})).get(
+ "models", ()
+ )
+ if not isinstance(managed_models, (list, tuple)) or not managed_models:
+ raise HTTPException(403, "This agent has no model grants")
+ _can_object_call_model(
+ model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router),
+ llm_router=llm_router,
+ models=list(managed_models),
+ team_id=valid_token.team_id,
+ object_type="agent",
+ key_model_aliases=key_model_aliases_for_auth_check(valid_token),
+ )
+
await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router)
await _check_agent_caller_model_access(
model=_model,
@@ -2653,7 +2668,7 @@ async def get_user_object(
)
if should_check_db:
- response = await _user_table(UserRepository(prisma_client)).find_unique(
+ response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
@@ -2691,7 +2706,7 @@ async def get_user_object(
budget_duration=new_user_params["budget_duration"]
)
- response = await _user_table(UserRepository(prisma_client)).create(
+ response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).create(
data=new_user_params,
include={"organization_memberships": True},
)
@@ -3116,9 +3131,9 @@ class TeamNotFoundError(HTTPException):
@log_db_metrics
async def _get_team_db_check(
- team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
+ team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None, *, use_writer: bool = False
) -> "_PrismaTeamRow | None":
- response = await _team_table(TeamRepository(prisma_client)).find_unique(
+ response = await _team_table(TeamRepository(prisma_client, use_writer=use_writer)).find_unique(
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
)
@@ -3152,6 +3167,7 @@ async def _get_team_object_from_user_api_key_cache(
proxy_logging_obj: ProxyLogging | None,
key: str,
team_id_upsert: bool | None = None,
+ use_writer: bool = False,
) -> LiteLLM_TeamTableCachedObj:
db_access_time_key: Final = key
should_check_db: Final = _should_check_db(
@@ -3160,7 +3176,9 @@ async def _get_team_object_from_user_api_key_cache(
db_cache_expiry=db_cache_expiry,
)
if should_check_db:
- response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert)
+ response = await _get_team_db_check(
+ team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert, use_writer=use_writer
+ )
# The database answered and the row is not there. Distinct from every
# other failure here, which leaves the team's grant unknown.
if response is None:
@@ -3182,8 +3200,11 @@ async def _get_team_object_from_user_api_key_cache(
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=use_writer,
)
except Exception as e:
+ if use_writer:
+ raise
verbose_proxy_logger.debug(
"Failed to load object_permission for team %s with object_permission_id=%s: %s",
team_id,
@@ -3273,6 +3294,7 @@ async def get_team_object(
db_cache_expiry=db_cache_expiry,
key=key,
team_id_upsert=team_id_upsert,
+ use_writer=bool(check_db_only),
)
except TeamNotFoundError:
raise
@@ -3318,16 +3340,15 @@ async def get_access_object(
prisma_client: DatabaseClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None = None,
+ *,
+ check_db_only: bool = False,
) -> LiteLLM_AccessGroupTable:
"""
- Check if access_group_id in proxy AccessGroupTable
- - Always checks cache first, then DB only when not found in cache
+ - Checks cache first unless authoritative writer admission is requested
- if valid, return LiteLLM_AccessGroupTable object
- if not, then raise an error
- Unlike get_team_object, this has no check_cache_only or check_db_only flags;
- it always follows cache-first-then-db semantics.
-
Raises:
- HTTPException: If access group doesn't exist in db or cache (status_code=404)
"""
@@ -3336,18 +3357,19 @@ async def get_access_object(
key: Final = f"access_group_id:{access_group_id}"
- cached_access_obj: Final = await user_api_key_cache.async_get_cache(
- key=key,
- model_type=LiteLLM_AccessGroupTable,
+ cached_access_obj: Final = (
+ None
+ if check_db_only
+ else await user_api_key_cache.async_get_cache(key=key, model_type=LiteLLM_AccessGroupTable)
)
if cached_access_obj is not None:
return cached_access_obj
# Not in cache - fetch from DB
try:
- response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique(
- where={"access_group_id": access_group_id}
- )
+ response: Final = await _dictable_table(
+ AccessGroupRepository(prisma_client, use_writer=check_db_only), "access_group"
+ ).find_unique(where={"access_group_id": access_group_id})
if response is None:
raise HTTPException(
@@ -3931,6 +3953,7 @@ async def get_object_permission(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
+ check_db_only: bool = False,
) -> LiteLLM_ObjectPermissionTable | None:
"""
- Check if object permission id in proxy ObjectPermissionTable
@@ -3942,9 +3965,13 @@ async def get_object_permission(
# check if in cache
key: Final = object_permission_cache_key(object_permission_id)
- deserialized_perm: Final = await user_api_key_cache.async_get_cache(
- key=key,
- model_type=LiteLLM_ObjectPermissionTable,
+ deserialized_perm: Final = (
+ None
+ if check_db_only
+ else await user_api_key_cache.async_get_cache(
+ key=key,
+ model_type=LiteLLM_ObjectPermissionTable,
+ )
)
if deserialized_perm is not None:
return deserialized_perm
@@ -3952,10 +3979,12 @@ async def get_object_permission(
# else, check db
try:
response: Final = await _dictable_table(
- ObjectPermissionRepository(prisma_client), "object_permission"
+ ObjectPermissionRepository(prisma_client, use_writer=check_db_only), "object_permission"
).find_unique(where={"object_permission_id": object_permission_id})
if response is None:
+ if check_db_only:
+ raise HTTPException(status_code=403, detail="Referenced object permission does not exist")
return None
_perm_obj: Final = LiteLLM_ObjectPermissionTable.model_validate(response.dict())
@@ -3968,6 +3997,8 @@ async def get_object_permission(
return _perm_obj
except Exception:
+ if check_db_only:
+ raise
return None
@@ -4177,6 +4208,7 @@ async def _get_resources_from_access_groups(
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
+ check_db_only: bool = False,
) -> list[str]:
"""
Fetch access groups by their IDs (from cache or DB) and collect
@@ -4219,6 +4251,7 @@ async def _get_resources_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=check_db_only,
)
resources.extend(getattr(ag, resource_field, []))
except Exception:
@@ -4254,6 +4287,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
+ check_db_only: bool = False,
) -> list[str]:
"""
Collect MCP server IDs from unified access groups.
@@ -4265,6 +4299,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=check_db_only,
)
@@ -4273,6 +4308,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
+ check_db_only: bool = False,
) -> list[str]:
"""
Collect agent IDs from unified access groups.
@@ -4284,6 +4320,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
+ check_db_only=check_db_only,
)
@@ -4483,26 +4520,36 @@ async def _check_agent_access_group_model_access(
"""Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows."""
if not model or valid_token is None or not valid_token.agent_id:
return True
- ceiling: Final = await resolve_ceiling(valid_token.agent_id)
- if ceiling is None:
- return True
- if not ceiling.models:
- raise ModelAccessDeniedProxyException(
- message=model_access_denied_client_message(model=model),
- internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models",
- type=ProxyErrorTypes.agent_model_access_denied,
- param="model",
- code=status.HTTP_403_FORBIDDEN,
- )
- dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
- return _can_object_call_model(
- model=dispatched,
- llm_router=llm_router,
- models=sorted(ceiling.models),
- team_id=valid_token.team_id,
- object_type="agent",
- key_model_aliases=key_model_aliases_for_auth_check(valid_token),
+
+ from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
+
+ unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if valid_token.managed_agent_policy is None else None
+ ceilings: Final = (
+ await resolve_managed_agent_ceilings(valid_token.managed_agent_policy)
+ if valid_token.managed_agent_policy is not None
+ else (unmanaged,)
+ if unmanaged is not None
+ else ()
)
+ dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
+ for ceiling in ceilings:
+ if not ceiling.models:
+ raise ModelAccessDeniedProxyException(
+ message=model_access_denied_client_message(model=model),
+ internal_message=f"agent {valid_token.agent_id} access groups grant no models",
+ type=ProxyErrorTypes.agent_model_access_denied,
+ param="model",
+ code=status.HTTP_403_FORBIDDEN,
+ )
+ _can_object_call_model(
+ model=dispatched,
+ llm_router=llm_router,
+ models=sorted(ceiling.models),
+ team_id=valid_token.team_id,
+ object_type="agent",
+ key_model_aliases=key_model_aliases_for_auth_check(valid_token),
+ )
+ return True
LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None
diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py
index 4f41b283a33..71b87b74915 100644
--- a/litellm/proxy/auth/handle_jwt.py
+++ b/litellm/proxy/auth/handle_jwt.py
@@ -52,6 +52,10 @@ from litellm.proxy._types import (
TeamMemberAddRequest,
UserAPIKeyAuth,
)
+from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
+from litellm.proxy.agent_endpoints.identity import has_legacy_identity
+from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
+from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.model_access_denied import (
ModelAccessDeniedHTTPException,
@@ -67,6 +71,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.user_repository import UserRepository
from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import AgentIdentityFailure
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from .auth_checks import (
@@ -157,6 +162,8 @@ class HeaderTeam:
class AgentLookup(Protocol):
"""The registered-agent lookups a JWT agent claim is matched against."""
+ def get_agent_list(self) -> Sequence[AgentResponse]: ...
+
def get_agent_by_id(self, agent_id: str) -> AgentResponse | None:
"""The agent registered under ``agent_id``, if any."""
@@ -167,6 +174,9 @@ class AgentLookup(Protocol):
class _NoRegisteredAgents:
"""The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches."""
+ def get_agent_list(self) -> tuple[AgentResponse, ...]:
+ return ()
+
def get_agent_by_id(self, agent_id: str) -> None:
return None
@@ -1096,6 +1106,15 @@ class JWTHandler:
"options": options or None,
}
+ def managed_issuer_is_trusted(self, issuer: object) -> bool:
+ if not isinstance(issuer, str):
+ return False
+ configured: Final = self.litellm_jwtauth.issuers or ()
+ for item in configured:
+ if item.issuer == issuer:
+ return bool(item.audience) and not item.disable_audience_validation
+ return issuer == os.getenv("JWT_ISSUER") and bool(os.getenv("JWT_AUDIENCE"))
+
def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
@@ -1488,7 +1507,12 @@ class JWTAuthManager:
agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
agent_name=agent_claim
)
- if agent is None:
+ if (
+ agent is None
+ or agent.identity_managed
+ or agent.identity is not None
+ or has_legacy_identity(agent.litellm_params)
+ ):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
@@ -2478,12 +2502,39 @@ class JWTAuthManager:
"""Resolve and authorize JWT context; only normal admission supplies provisioning."""
handler: Final = jwt_handler
jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
+ managed: Final = await resolve_managed_agent(jwt_valid_token, prisma_client)
+ if managed is not None:
+ if not handler.managed_issuer_is_trusted(jwt_valid_token.get("iss")):
+ raise HTTPException(403, "Managed agents require trusted JWT issuer and audience validation")
+ if not managed_agent_route_allowed(route, request_method):
+ raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
+ evidence: Final = await AgentIdentityStore.from_client(prisma_client).record_authentication(managed)
+ if isinstance(evidence, AgentIdentityFailure):
+ raise_identity_failure(evidence)
+ if managed.mode == "autonomous":
+ return JWTAuthBuilderResult(
+ is_proxy_admin=False,
+ team_id=None,
+ team_object=None,
+ user_id=None,
+ user_email=None,
+ user_object=None,
+ org_id=None,
+ org_object=None,
+ end_user_id=None,
+ end_user_object=None,
+ token=api_key,
+ team_membership=None,
+ jwt_claims=jwt_valid_token,
+ agent_id=managed.agent_id,
+ managed_agent_context=managed,
+ )
team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
model: Final = request_data.get("model")
requested_model: Final = model if isinstance(model, str) else None
# Check RBAC
- rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
+ rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if managed is None else None
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
# Check Scope Based Access
@@ -2499,7 +2550,11 @@ class JWTAuthManager:
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
# Get basic user info
- user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
+ user_id, user_email, valid_user_email = (
+ (managed.user_id, None, None)
+ if managed is not None
+ else await JWTAuthManager.get_user_info(handler, jwt_valid_token)
+ )
# Get IDs
org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
@@ -2514,23 +2569,31 @@ class JWTAuthManager:
elif rbac_role == LitellmUserRoles.INTERNAL_USER:
user_id = object_id
- agent_id: Final = JWTAuthManager.resolve_agent_id(
- jwt_handler=handler,
- jwt_valid_token=jwt_valid_token,
- agent_registry=handler.agent_lookup,
+ agent_id: Final = (
+ managed.agent_id
+ if managed is not None
+ else JWTAuthManager.resolve_agent_id(
+ jwt_handler=handler,
+ jwt_valid_token=jwt_valid_token,
+ agent_registry=handler.agent_lookup,
+ )
)
# Check admin access
- admin_result: Final = await JWTAuthManager.check_admin_access(
- handler,
- scopes,
- route,
- user_id,
- org_id,
- api_key,
- jwt_valid_token,
- user_email=user_email,
- agent_id=agent_id,
+ admin_result: Final = (
+ None
+ if managed is not None
+ else await JWTAuthManager.check_admin_access(
+ handler,
+ scopes,
+ route,
+ user_id,
+ org_id,
+ api_key,
+ jwt_valid_token,
+ user_email=user_email,
+ agent_id=agent_id,
+ )
)
if admin_result:
await JWTAuthManager._attach_team_from_header_for_admin(
@@ -2705,13 +2768,13 @@ class JWTAuthManager:
proxy_logging_obj=proxy_logging_obj,
route=route,
org_alias=org_alias,
- user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
+ user_id_upsert=provisioning.user_id_upsert if provisioning is not None and managed is None else False,
)
# Derive org_id from org_object if resolved by alias
resolved_org_id: Final = org_object.organization_id if org_object else org_id
- if provisioning is not None:
+ if provisioning is not None and managed is None:
await JWTAuthManager.sync_user_role_and_teams(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
@@ -2784,7 +2847,7 @@ class JWTAuthManager:
)
## MAP USER TO TEAMS
- if provisioning is not None:
+ if provisioning is not None and managed is None:
await JWTAuthManager.map_user_to_teams(
user_object=user_object,
team_object=team_object,
@@ -2799,7 +2862,9 @@ class JWTAuthManager:
)
# check if user is proxy admin
- is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN)
+ is_proxy_admin: Final = managed is None and bool(
+ user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN
+ )
return JWTAuthBuilderResult(
is_proxy_admin=is_proxy_admin,
@@ -2816,6 +2881,7 @@ class JWTAuthManager:
team_membership=team_membership_object,
jwt_claims=jwt_valid_token,
agent_id=agent_id,
+ managed_agent_context=managed,
)
@staticmethod
@@ -2826,11 +2892,13 @@ class JWTAuthManager:
"""Keep JWT identity and permission attribution identical across consumers."""
user: Final = result["user_object"]
admin: Final = result["is_proxy_admin"]
- return UserAPIKeyAuth(
+ auth: Final = UserAPIKeyAuth(
api_key=None,
user_role=(
LitellmUserRoles.PROXY_ADMIN
if admin
+ else LitellmUserRoles.INTERNAL_USER
+ if result.get("managed_agent_context") is not None
else LitellmUserRoles(user.user_role)
if user is not None and user.user_role is not None
else LitellmUserRoles.INTERNAL_USER
@@ -2852,3 +2920,5 @@ class JWTAuthManager:
user_id=result["user_id"],
),
)
+ auth.managed_agent_context = result.get("managed_agent_context")
+ return auth
diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py
index e3ce9bcd850..2829bc95b42 100644
--- a/litellm/proxy/auth/user_api_key_auth.py
+++ b/litellm/proxy/auth/user_api_key_auth.py
@@ -652,6 +652,8 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
# never reaches the fallback.
synthetic_scope: Final[dict[str, Any]] = {
"type": "http",
+ "method": "GET",
+ "query_string": ws_scope.get("query_string", b""),
"headers": scope_headers,
"path": ws_scope.get("path", ""),
"state": ws_scope.setdefault("state", {}), # mutable-ok: Starlette's socket state, shared with the request
@@ -1653,6 +1655,13 @@ async def _user_api_key_auth_builder(
else:
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
+ from litellm.proxy.agent_endpoints.identity_store import resolve_managed_agent
+
+ if jwt_claims and await resolve_managed_agent(jwt_claims, prisma_client) is not None:
+ raise HTTPException(
+ 403, "Managed agents require direct JWT authentication without virtual-key mapping"
+ )
+
resolve_result: Final = await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
@@ -3123,7 +3132,10 @@ async def _reserve_budget_after_common_checks(
end_user_id=end_user_id,
end_user_object=end_user_object,
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
- fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
+ fail_closed_budget_enforcement=(
+ general_settings.get("fail_closed_budget_enforcement") is True
+ or user_api_key_auth_obj.billing_agent_policy is not None
+ ),
raw_body=await read_raw_json_body(request=request),
)
if request is not None:
@@ -3197,10 +3209,48 @@ async def _authorize_authenticated_request(
# admin-only-route / model-access / budget checks) surface as
# ProxyException consistently with pre-refactor behavior.
try:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import (
+ admit_managed_actor,
+ invocation_target,
+ managed_agent_route_allowed,
+ managed_inference_request,
+ prepare_agent_invocation,
+ )
+ from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+ from litellm.proxy.proxy_server import general_settings, prisma_client, user_model
+
+ store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None
+ if user_api_key_auth_obj.agent_id is not None:
+ await admit_managed_actor(user_api_key_auth_obj, store)
+ if user_api_key_auth_obj.managed_agent_policy is not None and not managed_agent_route_allowed(
+ route, request.method
+ ):
+ raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
+ authorized_data: Final = (
+ managed_inference_request(
+ route,
+ request_data,
+ general_settings,
+ user_model,
+ request.path_params.get("model") or request.path_params.get("model_name"),
+ request.query_params.get("model"),
+ )
+ if user_api_key_auth_obj.managed_agent_policy is not None
+ else request_data
+ )
+ target_name: Final = invocation_target(route, authorized_data)
+ if target_name is not None:
+ await prepare_agent_invocation(
+ user_api_key_auth_obj,
+ target_name,
+ store,
+ billable=request_data.get("method")
+ in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
+ )
await _run_centralized_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request=request,
- request_data=request_data,
+ request_data=authorized_data,
route=route,
)
except Exception as e:
diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py
index 64b0c6c1967..b3be5fb78c9 100644
--- a/litellm/proxy/common_request_processing.py
+++ b/litellm/proxy/common_request_processing.py
@@ -93,6 +93,7 @@ from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_bod
from litellm.proxy.common_utils.http_parsing_utils import (
get_client_requested_model,
get_tags_from_request_body,
+ resolve_inference_model,
)
from litellm.proxy.common_utils.openai_error_payload import (
LITELLM_CALL_ID_HEADER,
@@ -2068,11 +2069,12 @@ class ProxyBaseLLMRequestProcessing:
if isinstance(model, str):
reject_url_valued_destination("model", model)
- self.data["model"] = (
- general_settings.get("completion_model", None) # server default
- or user_model # model name passed via cli args
- or model # for azure deployments
- or self.data.get("model", None) # default passed in http request
+ self.data["model"] = resolve_inference_model(
+ self.data.get("model"),
+ general_settings,
+ user_model,
+ model,
+ kind="image_edit" if route_type == "aimage_edit" else "completion",
)
# override with user settings, these are params passed via cli
diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py
index 1c2bd7ea217..be62448e3ed 100644
--- a/litellm/proxy/common_utils/http_parsing_utils.py
+++ b/litellm/proxy/common_utils/http_parsing_utils.py
@@ -2,7 +2,7 @@ import json
import re
from collections.abc import Collection, Mapping
from types import MappingProxyType, UnionType
-from typing import Annotated, Any, Final, Union, get_args, get_origin
+from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin
import orjson
from fastapi import Request, UploadFile, status
@@ -25,6 +25,39 @@ _FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form-
_ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required})
+def resolve_inference_model(
+ body_model: object,
+ settings: Mapping[str, object],
+ cli_model: str | None,
+ endpoint_model: object = None,
+ *,
+ kind: Literal[
+ "completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"
+ ] = "completion",
+) -> object:
+ match kind:
+ case "image_generation":
+ return cli_model or endpoint_model or settings.get("image_generation_model") or body_model
+ case "image_edit":
+ return (
+ settings.get("completion_model")
+ or cli_model
+ or endpoint_model
+ or settings.get("image_generation_model")
+ or body_model
+ )
+ case "moderation":
+ return cli_model or settings.get("moderation_model") or body_model
+ case "speech":
+ return cli_model or body_model
+ case "body":
+ return body_model
+ case "path":
+ return endpoint_model
+ case "completion":
+ return settings.get("completion_model") or cli_model or endpoint_model or body_model
+
+
def _normalize_media_type(content_type: str) -> str:
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
if not content_type:
diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py
index 63da5d15207..8a82e253c5c 100644
--- a/litellm/proxy/common_utils/registry_read_through.py
+++ b/litellm/proxy/common_utils/registry_read_through.py
@@ -174,7 +174,10 @@ async def _resync_agents(agent_id_or_name: str) -> bool:
table: Final = agents_table(prisma_client)
id_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id_or_name}
name_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_name": agent_id_or_name}
- include_permission: Final[LiteLLM_AgentsTableInclude] = {"object_permission": True}
+ include_permission: Final[LiteLLM_AgentsTableInclude] = {
+ "object_permission": True,
+ "identity": True,
+ }
async with AGENT_RECONCILE_LOCK:
if _agent_from_registry(agent_id_or_name) is not None:
return True
diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py
index 0178465739b..e735f4a51f0 100644
--- a/litellm/proxy/hooks/proxy_track_cost_callback.py
+++ b/litellm/proxy/hooks/proxy_track_cost_callback.py
@@ -358,6 +358,7 @@ class _ProxyDBLogger(CustomLogger):
team_id=team_id,
end_user_id=end_user_id,
call_type=call_type,
+ agent_id=metadata.get("billing_agent_id") or metadata.get("agent_id"),
):
## UPDATE DATABASE
charged: Final = await _update_database_and_spend_counters(
@@ -612,6 +613,7 @@ def _should_track_cost_callback(
team_id: str | None,
end_user_id: str | None,
call_type: str | None = None,
+ agent_id: str | None = None,
) -> bool:
"""
Determine if the cost callback should be tracked based on the kwargs
@@ -628,7 +630,13 @@ def _should_track_cost_callback(
if ProxyUpdateSpend.disable_spend_updates() is True:
return False
- if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None:
+ if (
+ agent_id is not None
+ or user_api_key is not None
+ or user_id is not None
+ or team_id is not None
+ or end_user_id is not None
+ ):
return True
return call_type in _UNATTRIBUTED_TRACKABLE_CALL_TYPES
diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py
index b9580ba3948..16dc38575da 100644
--- a/litellm/proxy/image_endpoints/endpoints.py
+++ b/litellm/proxy/image_endpoints/endpoints.py
@@ -21,6 +21,7 @@ from litellm.proxy.common_request_processing import (
from litellm.proxy.common_utils.http_parsing_utils import (
coerce_numeric_form_fields,
numeric_form_fields,
+ resolve_inference_model,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
@@ -118,14 +119,9 @@ async def image_generation(
if isinstance(model, str):
reject_url_valued_destination("model", model)
- data["model"] = (
- model
- or general_settings.get("image_generation_model", None) # server default
- or user_model # model name passed via cli args
- or data.get("model", None) # default passed in http request
+ data["model"] = resolve_inference_model(
+ data.get("model"), general_settings, user_model, model, kind="image_generation"
)
- if user_model:
- data["model"] = user_model
### MODEL ALIAS MAPPING ###
# check if model name in model alias map
@@ -324,12 +320,6 @@ async def image_edit_api(
if "prompt" not in data:
data["prompt"] = None
- data["model"] = (
- model
- or general_settings.get("image_generation_model", None) # server default
- or user_model # model name passed via cli args
- or data.get("model", None) # default passed in http request
- )
#########################################################
# Process request
#########################################################
@@ -346,7 +336,7 @@ async def image_edit_api(
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
- model=None,
+ model=model,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 56f647d5acc..a698881189a 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -1659,7 +1659,19 @@ class LiteLLMProxyRequestSetup:
_key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None)
_existing_agent_id: Final = data[_metadata_variable_name].get("agent_id")
_resolved_agent_id: Final = _key_agent_id or _existing_agent_id
- data[_metadata_variable_name]["agent_id"] = _resolved_agent_id
+ data[_metadata_variable_name]["agent_id"] = user_api_key_dict.invoked_agent_id or _resolved_agent_id
+ managed_context: Final = user_api_key_dict.managed_agent_context
+ data[_metadata_variable_name].update(
+ MappingProxyType(
+ {
+ "actor_agent_id": user_api_key_dict.agent_id,
+ "target_agent_id": user_api_key_dict.invoked_agent_id,
+ "billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id,
+ "agent_execution_mode": managed_context.mode if managed_context else None,
+ "verified_human_user_id": managed_context.user_id if managed_context else None,
+ }
+ )
+ )
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
user_api_key_dict, "end_user_max_budget", None
diff --git a/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py b/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py
new file mode 100644
index 00000000000..e548b9b7fa2
--- /dev/null
+++ b/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py
@@ -0,0 +1,49 @@
+from collections.abc import Mapping
+from types import MappingProxyType
+from typing import Final
+from uuid import UUID
+
+from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
+
+
+def microsoft_interactive_subject(
+ tenant: str | None,
+ response: Mapping[str, object],
+ endpoints: Mapping[str, str | None],
+) -> MicrosoftInteractiveSubject | None:
+ if tenant is None:
+ return None
+ try:
+ tenant_id: Final = str(UUID(tenant))
+ object_id: Final = response.get("id")
+ if not isinstance(object_id, str):
+ return None
+ oid: Final = str(UUID(object_id))
+ except ValueError:
+ return None
+ expected: Final = MappingProxyType(
+ {
+ "MICROSOFT_AUTHORIZATION_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/authorize",
+ "MICROSOFT_TOKEN_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token",
+ "MICROSOFT_USERINFO_ENDPOINT": "https://graph.microsoft.com/v1.0/me",
+ }
+ )
+ if any(value and value != expected.get(name) for name, value in endpoints.items()):
+ return None
+ return MicrosoftInteractiveSubject(
+ issuer=f"https://login.microsoftonline.com/{tenant_id}/v2.0",
+ tenant_id=tenant_id,
+ oid=oid,
+ )
+
+
+async def enroll_microsoft_subject(subject: object, user_id: object, client: object) -> None:
+ from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+ from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
+ from litellm.types.proxy.agent_identity import AgentIdentityFailure
+
+ if not isinstance(subject, MicrosoftInteractiveSubject) or not isinstance(user_id, str) or not user_id:
+ return
+ result: Final = await AgentIdentityStore.from_client(client).enroll_interactive_human(subject, user_id)
+ if isinstance(result, AgentIdentityFailure):
+ raise_identity_failure(result)
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index 618b200a14c..c995b5aedd4 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -3631,6 +3631,12 @@ class SSOAuthenticationHandler:
},
)
+ from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject
+
+ await enroll_microsoft_subject(
+ request.scope.get("litellm_microsoft_interactive_subject"), user_id, prisma_client
+ )
+
if isinstance(user_id, str) and user_id:
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
@@ -4300,6 +4306,22 @@ class MicrosoftSSOHandler:
original_msft_result["app_roles"] = app_roles
return original_msft_result or {}
+ from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import microsoft_interactive_subject
+
+ request.scope["litellm_microsoft_interactive_subject"] = microsoft_interactive_subject(
+ microsoft_tenant,
+ original_msft_result,
+ MappingProxyType(
+ {
+ name: os.getenv(name)
+ for name in (
+ "MICROSOFT_AUTHORIZATION_ENDPOINT",
+ "MICROSOFT_TOKEN_ENDPOINT",
+ "MICROSOFT_USERINFO_ENDPOINT",
+ )
+ }
+ ),
+ )
result: Final = MicrosoftSSOHandler.openid_from_response(
response=original_msft_result,
team_ids=user_team_ids,
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index f842f2e1e4a..b3b20a16234 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -421,6 +421,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
check_file_size_under_limit,
get_form_data,
+ resolve_inference_model,
)
from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
@@ -12186,13 +12187,7 @@ async def moderations(
proxy_config=proxy_config,
)
- data["model"] = (
- general_settings.get("moderation_model", None) # server default
- or user_model # model name passed via cli args
- or data.get("model") # default passed in http request
- )
- if user_model:
- data["model"] = user_model
+ data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
### CALL HOOKS ### - modify incoming data / reject request before calling the model
data = await proxy_logging_obj.pre_call_hook(
@@ -12446,13 +12441,7 @@ async def audio_transcriptions(
if data.get("user", None) is None and user_api_key_dict.user_id is not None:
data["user"] = user_api_key_dict.user_id
- data["model"] = (
- general_settings.get("moderation_model", None) # server default
- or user_model # model name passed via cli args
- or data.get("model", None) # default passed in http request
- )
- if user_model:
- data["model"] = user_model
+ data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
router_model_names: Final = llm_router.model_names if llm_router is not None else []
diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma
index 69c63d9ecd6..8535d1004f7 100644
--- a/litellm/proxy/schema.prisma
+++ b/litellm/proxy/schema.prisma
@@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
+ identity_managed Boolean @default(false)
+ enabled Boolean @default(true)
+ execution_mode String @default("autonomous")
+ identity LiteLLM_AgentIdentity?
+ retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
+model LiteLLM_AgentIdentity {
+ agent_id String @id
+ active Boolean @default(true)
+ agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ service_principal_id String?
+ required_roles String[] @default([])
+ required_scopes String[] @default(["user_impersonation"])
+ revision String @default(uuid())
+ last_authenticated_at DateTime?
+ @@unique([provider, tenant_id, client_id])
+ @@unique([issuer, service_principal_id])
+}
+
+model LiteLLM_RetiredAgentIdentity {
+ binding_id String @id @default(uuid())
+ agent_id String?
+ agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ @@unique([provider, tenant_id, client_id])
+}
+
+model LiteLLM_RetiredAgent {
+ original_agent_id String @id
+ retired_at DateTime @default(now())
+}
+
+model LiteLLM_VerifiedSubject {
+ subject_id String @id @default(uuid())
+ issuer String
+ tenant_id String
+ oid String
+ kind String @default("human")
+ user_id String?
+ user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
+ verified_via String @default("sso_interactive")
+ verified_at DateTime @default(now())
+ @@unique([issuer, tenant_id, oid])
+ @@index([user_id])
+}
+
+
+
+
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
+ verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@@ -674,6 +730,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
+ billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?
diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py
index c4ed8713f95..c3a1a7000c5 100644
--- a/litellm/proxy/spend_tracking/spend_management_endpoints.py
+++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py
@@ -77,6 +77,11 @@ _SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)"
_SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key"
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
+_SESSION_REPRESENTATIVE_ORDER_SQL: Final = (
+ f"(call_type = {_AGENT_CALL_TYPE_SQL}) DESC, "
+ f'CASE WHEN call_type = {_AGENT_CALL_TYPE_SQL} THEN "endTime" END DESC NULLS LAST, '
+ f'call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC, request_id'
+)
_BATCH_CALL_TYPES_SQL: Final = "('acreate_batch', 'create_batch', 'aretrieve_batch', 'retrieve_batch')"
_SPAN_TYPE_SQL_CONDITIONS: Final[Mapping[str, str]] = MappingProxyType(
{
@@ -2879,7 +2884,7 @@ async def ui_view_spend_logs(
p += 1
# Status filter
- if status_filter is not None:
+ if status_filter is not None and not (group_by_session is True and not is_search_lookup):
if status_filter == "success":
sql_conditions.append("(status = 'success' OR status IS NULL)")
else:
@@ -2925,6 +2930,23 @@ async def ui_view_spend_logs(
sql_params.append(f"%{error_message}%")
p += 1
+ if status_filter is not None and group_by_session is True and not is_search_lookup:
+ session_filter_conditions: Final = " AND ".join(sql_conditions) or "TRUE"
+ sql_conditions.append(
+ f"""({_SESSION_GROUP_KEY_SQL}) IN (
+ SELECT session_key, api_key FROM (
+ SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
+ {_SESSION_KEY_EXPR} AS session_key, api_key, status
+ FROM "LiteLLM_SpendLogs"
+ WHERE {session_filter_conditions}
+ ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
+ ) AS session_outcomes
+ WHERE COALESCE(status, 'success') = ${p}
+ )"""
+ )
+ sql_params.append(status_filter)
+ p += 1
+
if (
group_by_session is True
and not is_v2
@@ -2991,7 +3013,7 @@ async def ui_view_spend_logs(
{_SPEND_LOG_LIST_COLUMNS}
FROM "LiteLLM_SpendLogs"
WHERE {joined_conditions}
- ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
+ ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
) AS session_representatives
ORDER BY {exact_request_id_first}{_order_expr} {_sql_dir}{_nulls_clause}, request_id
LIMIT ${p} OFFSET ${p + 1}
@@ -3063,7 +3085,7 @@ async def _fetch_session_representatives(
next_param_index: int,
session_keys: Sequence[tuple[str, str]],
) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row
- """Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
+ """Fetch the final agent outcome, or newest non-MCP row, of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
rep_query: Final = f"""
SELECT * FROM (
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
@@ -3073,7 +3095,7 @@ async def _fetch_session_representatives(
AND ({_SESSION_GROUP_KEY_SQL}) IN (
SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[])
)
- ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
+ ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
) AS session_representatives
"""
rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place
@@ -3140,7 +3162,7 @@ async def _ui_session_grouped_spend_logs(
page_size``, trimmed to the end of the ``SPEND_LOGS_PAGINATION_COUNT_CAP``
window the capped ``total`` promises, so a page never runs past that total
and one starting at or past it returns no rows without a query. Each session is represented
- by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response``
+ by its final agent outcome (or newest non-MCP row), enriched by ``_build_ui_spend_logs_response``
exactly like the flat listing, and the response carries
``next_session_cursor`` / ``has_more`` while ``total`` counts sessions
(capped like the flat total). A page that runs out of sessions while still
diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py
index 1c51fb21d6e..e5001eeb32f 100644
--- a/litellm/proxy/spend_tracking/spend_tracking_utils.py
+++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py
@@ -795,6 +795,7 @@ def get_logging_payload(
model_id=_model_id,
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
agent_id=agent_id,
+ billing_agent_id=clean_metadata.get("billing_agent_id"),
requester_ip_address=clean_metadata.get("requester_ip_address", None),
custom_llm_provider=custom_llm_provider or "",
messages=_get_messages_for_spend_logs_payload(
diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py
index 7736939c696..b732d2ff94c 100644
--- a/litellm/repositories/object_permission_repository.py
+++ b/litellm/repositories/object_permission_repository.py
@@ -15,9 +15,14 @@ if TYPE_CHECKING:
class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
"""Repository for object permission database operations."""
+ def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
+ super().__init__(prisma_client)
+ self._use_writer = use_writer
+
@property
def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
- return self.prisma_client.db.litellm_objectpermissiontable
+ database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
+ return database.litellm_objectpermissiontable
@property
def model_class(self) -> type[LiteLLM_ObjectPermissionTable]:
diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py
index 1ad7a735d96..292694747f3 100644
--- a/litellm/repositories/table_repositories.py
+++ b/litellm/repositories/table_repositories.py
@@ -21,8 +21,9 @@ class PrismaTableRepository(Generic[RowT_co]):
table_name: str
- def __init__(self, prisma_client: object):
+ def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
self._prisma_client = prisma_client
+ self._use_writer = use_writer
@property
def prisma_client(self) -> Any:
@@ -32,7 +33,9 @@ class PrismaTableRepository(Generic[RowT_co]):
@property
def table(self) -> TableActions[RowT_co]:
- actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name)
+ actions: Final[TableActions[RowT_co]] = getattr(
+ self.prisma_client.writer_db if self._use_writer else self.prisma_client.db, self.table_name
+ )
return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name)
@@ -44,6 +47,18 @@ class AgentsRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentsTable"
table_name = "litellm_agentstable"
+class AgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentIdentity"]):
+ table_name = "litellm_agentidentity"
+
+
+class RetiredAgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgentIdentity"]):
+ table_name = "litellm_retiredagentidentity"
+
+
+class VerifiedSubjectRepository(PrismaTableRepository["prisma_models.LiteLLM_VerifiedSubject"]):
+ table_name = "litellm_verifiedsubject"
+
+
class ObjectPermissionRepository(PrismaTableRepository["prisma_models.LiteLLM_ObjectPermissionTable"]):
table_name = "litellm_objectpermissiontable"
@@ -246,3 +261,7 @@ class AuditLogRepository(PrismaTableRepository["prisma_models.LiteLLM_AuditLog"]
class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterSession"]):
table_name = "litellm_adaptiveroutersession"
+
+
+class RetiredAgentRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgent"]):
+ table_name = "litellm_retiredagent"
diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py
index cbe263699c9..57f6fd33c11 100644
--- a/litellm/repositories/team_repository.py
+++ b/litellm/repositories/team_repository.py
@@ -70,6 +70,9 @@ class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaTeamDb: ...
+ @property
+ def writer_db(self) -> _PrismaTeamDb: ...
+
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
_JSON_ENCODED_TEAM_FIELDS: Final = (
@@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = (
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
"""Repository for team database operations."""
+ def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
+ super().__init__(prisma_client)
+ self._use_writer = use_writer
+
@property
def _db(self) -> _PrismaTeamDb:
client: Final[_PrismaClientView] = self.prisma_client
- return client.db
+ return client.writer_db if self._use_writer else client.db
@property
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:
diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py
index 87eb45f262d..4a2aea46197 100644
--- a/litellm/repositories/user_repository.py
+++ b/litellm/repositories/user_repository.py
@@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
class UserRepository(BaseRepository[LiteLLM_UserTable]):
"""Repository for user database operations."""
+ def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
+ super().__init__(prisma_client)
+ self._use_writer = use_writer
+
@property
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
- return self.prisma_client.db.litellm_usertable
+ database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
+ return database.litellm_usertable
@property
def model_class(self) -> type[LiteLLM_UserTable]:
diff --git a/litellm/types/agents.py b/litellm/types/agents.py
index f7aef09fa29..94adb9f7c4a 100644
--- a/litellm/types/agents.py
+++ b/litellm/types/agents.py
@@ -7,6 +7,11 @@ from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field
from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
+from litellm.types.proxy.agent_identity import (
+ AgentExecutionMode,
+ AgentIdentityBinding,
+ EntraIdentityConfig,
+)
if TYPE_CHECKING:
from a2a.types import SendMessageResponse
@@ -248,8 +253,11 @@ class AgentKillSwitchResult(BaseModel):
class AgentConfig(TypedDict, total=False):
+ identity: ReadOnly[EntraIdentityConfig | None]
+ enabled: ReadOnly[bool]
+ execution_mode: ReadOnly[AgentExecutionMode]
agent_name: Required[str]
- agent_card_params: Required[AgentCard]
+ agent_card_params: ReadOnly[AgentCard]
litellm_params: dict[str, object] # allow for any future litellm params
object_permission: AgentObjectPermission
tpm_limit: int | None
@@ -263,6 +271,9 @@ class AgentConfig(TypedDict, total=False):
class PatchAgentRequest(TypedDict, total=False):
+ identity: ReadOnly[EntraIdentityConfig | None]
+ enabled: ReadOnly[bool]
+ execution_mode: ReadOnly[AgentExecutionMode]
agent_name: str
agent_card_params: AgentCard
litellm_params: dict[str, object]
@@ -301,6 +312,11 @@ class AgentKeySummary(BaseModel):
class AgentResponse(BaseModel):
+ identity: AgentIdentityBinding | None = None
+ identity_managed: bool = False
+ enabled: bool = True
+ execution_mode: AgentExecutionMode = "autonomous"
+ jwt_auth_configured: bool = False
agent_id: str
agent_name: str
litellm_params: dict[str, object] | None = None
diff --git a/litellm/types/proxy/agent_identity.py b/litellm/types/proxy/agent_identity.py
new file mode 100644
index 00000000000..a7fe0be37e1
--- /dev/null
+++ b/litellm/types/proxy/agent_identity.py
@@ -0,0 +1,96 @@
+from datetime import datetime
+from typing import Literal, TypeAlias
+from uuid import UUID
+
+from pydantic import BaseModel, ConfigDict, Field, field_validator
+
+AgentExecutionMode: TypeAlias = Literal["autonomous", "delegated", "both"]
+
+
+class EntraIdentityConfig(BaseModel):
+ model_config = ConfigDict(frozen=True, extra="forbid")
+
+ provider: Literal["microsoft_entra"]
+ tenant_id: str
+ client_id: str
+ service_principal_id: str | None = None
+ required_roles: tuple[str, ...] = ()
+ required_scopes: tuple[str, ...] = Field(
+ default=("user_impersonation",),
+ description="Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.",
+ )
+
+ @field_validator("tenant_id", "client_id", "service_principal_id")
+ @classmethod
+ def normalize_identifier(cls, value: str | None) -> str | None:
+ return str(UUID(value)) if value is not None else None
+
+ @property
+ def issuer(self) -> str:
+ return f"https://login.microsoftonline.com/{self.tenant_id}/v2.0"
+
+
+class AgentIdentityBinding(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ agent_id: str
+ active: bool = True
+ provider: Literal["microsoft_entra"]
+ tenant_id: str
+ client_id: str
+ service_principal_id: str | None = None
+ issuer: str
+ required_roles: tuple[str, ...] = ()
+ required_scopes: tuple[str, ...] = ("user_impersonation",)
+ revision: str
+ last_authenticated_at: datetime | None = None
+
+
+class AgentSubject(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ kind: Literal["application", "delegated_subject"]
+ oid: str
+ mode: Literal["autonomous", "delegated"]
+
+
+class AgentIdentityFailure(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ code: Literal["identity_denied", "policy_unavailable"] = "identity_denied"
+ message: str
+
+
+class ManagedAgentContext(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ agent_id: str
+ binding_revision: str | None = None
+ mode: Literal["autonomous", "delegated"]
+ user_id: str | None = None
+ subject_oid: str | None = None
+
+
+class VerifiedHumanSubject(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ issuer: str
+ tenant_id: str
+ oid: str
+ user_id: str
+
+
+class MicrosoftInteractiveSubject(BaseModel):
+ model_config = ConfigDict(frozen=True)
+
+ issuer: str
+ tenant_id: str
+ oid: str
+
+
+class ManagedAgentIdentityStatus(BaseModel):
+ identity: AgentIdentityBinding | None = None
+ identity_managed: bool = False
+ enabled: bool = True
+ execution_mode: AgentExecutionMode = "autonomous"
+ last_authenticated_at: datetime | None = None
diff --git a/schema.prisma b/schema.prisma
index 69c63d9ecd6..8535d1004f7 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
+ identity_managed Boolean @default(false)
+ enabled Boolean @default(true)
+ execution_mode String @default("autonomous")
+ identity LiteLLM_AgentIdentity?
+ retired_identities LiteLLM_RetiredAgentIdentity[]
tpm_limit Int?
rpm_limit Int?
session_tpm_limit Int?
@@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
updated_by String
}
+model LiteLLM_AgentIdentity {
+ agent_id String @id
+ active Boolean @default(true)
+ agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ service_principal_id String?
+ required_roles String[] @default([])
+ required_scopes String[] @default(["user_impersonation"])
+ revision String @default(uuid())
+ last_authenticated_at DateTime?
+ @@unique([provider, tenant_id, client_id])
+ @@unique([issuer, service_principal_id])
+}
+
+model LiteLLM_RetiredAgentIdentity {
+ binding_id String @id @default(uuid())
+ agent_id String?
+ agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
+ provider String
+ issuer String
+ tenant_id String
+ client_id String
+ @@unique([provider, tenant_id, client_id])
+}
+
+model LiteLLM_RetiredAgent {
+ original_agent_id String @id
+ retired_at DateTime @default(now())
+}
+
+model LiteLLM_VerifiedSubject {
+ subject_id String @id @default(uuid())
+ issuer String
+ tenant_id String
+ oid String
+ kind String @default("human")
+ user_id String?
+ user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
+ verified_via String @default("sso_interactive")
+ verified_at DateTime @default(now())
+ @@unique([issuer, tenant_id, oid])
+ @@index([user_id])
+}
+
+
+
+
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String
@@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
// Track spend, rate limit, budget Users
model LiteLLM_UserTable {
+ verified_subjects LiteLLM_VerifiedSubject[]
user_id String @id
user_alias String?
team_id String?
@@ -674,6 +730,7 @@ model LiteLLM_SpendLogs {
session_id String?
status String?
mcp_namespaced_tool_name String?
+ billing_agent_id String?
agent_id String?
proxy_server_request Json? @default("{}")
litellm_call_id String?
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
new file mode 100644
index 00000000000..b6ddd9cb7ac
--- /dev/null
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py
@@ -0,0 +1,301 @@
+from typing import Final
+from unittest.mock import AsyncMock, MagicMock
+
+import pytest
+from fastapi import HTTPException
+
+from litellm.proxy import proxy_server
+from litellm.proxy._experimental.mcp_server import mcp_server_manager
+from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
+from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth
+from litellm.proxy.auth import auth_checks
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import ManagedAgentContext
+
+
+def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth:
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="agent-permissions",
+ mcp_servers=["slack", "linear"],
+ mcp_tool_permissions={"slack": list(tools)} if tools is not None else None,
+ )
+ agent: Final = AgentResponse(
+ agent_id="publisher",
+ agent_name="Publisher",
+ agent_card_params={},
+ object_permission=permission.model_dump(),
+ identity_managed=True,
+ )
+ auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id)
+ auth.managed_agent_policy = agent
+ auth.managed_agent_context = ManagedAgentContext(
+ agent_id=agent.agent_id,
+ mode="delegated" if delegated else "autonomous",
+ user_id="human" if delegated else None,
+ )
+ return auth
+
+
+@pytest.fixture(autouse=True)
+def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager())
+ monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write")))
+async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None:
+ auth: Final = actor(tools)
+ assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
+ actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
+ assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None)
+ assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == []
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "agent_tools,user_tools,expected",
+ (
+ (None, ("read",), ("read",)),
+ (("read",), None, ("read",)),
+ (("read", "write"), ("read",), ("read",)),
+ (("read",), ("write",), ()),
+ ((), None, ()),
+ ),
+)
+async def test_delegated_server_and_tool_intersections(
+ monkeypatch: pytest.MonkeyPatch,
+ agent_tools: tuple[str, ...] | None,
+ user_tools: tuple[str, ...] | None,
+ expected: tuple[str, ...],
+) -> None:
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="user-permissions",
+ mcp_servers=["slack", "user-only"],
+ mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None,
+ )
+ user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
+ auth: Final = actor(agent_tools, delegated=True)
+ assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
+ assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected)
+ assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
+ assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == []
+
+
+@pytest.mark.asyncio
+async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable")))
+ with pytest.raises(HTTPException) as failure:
+ await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
+ assert failure.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ())))
+async def test_access_groups_cap_agent_servers_without_granting_new_ones(
+ monkeypatch: pytest.MonkeyPatch,
+ servers: tuple[str, ...],
+ expected: tuple[str, ...],
+) -> None:
+ from litellm.proxy._types import LiteLLM_AccessGroupTable
+
+ group: Final = LiteLLM_AccessGroupTable(
+ access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers)
+ )
+ monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group))
+ 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": ["group"]})
+ assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected
+ if "slack" not in expected:
+ assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
+async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
+ monkeypatch: pytest.MonkeyPatch, change: str
+) -> None:
+ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
+
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]}
+ )
+ user: Final = LiteLLM_UserTable(
+ user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant"
+ )
+ cache: Final = UserApiKeyCache()
+ cache.set_cache("human", user)
+ cache.set_cache(object_permission_cache_key("user-grant"), permission)
+ client: Final = MagicMock()
+ client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
+ client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
+ monkeypatch.setattr(proxy_server, "prisma_client", client)
+ monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
+ auth: Final = actor(("read", "write"), delegated=True)
+ assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"}
+ if change == "disabled":
+ client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy(
+ update={"metadata": {"scim_active": False}}
+ )
+ elif change == "outage":
+ client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable")
+ elif change == "servers":
+ client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
+ update={"mcp_servers": [], "mcp_tool_permissions": {}}
+ )
+ else:
+ client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
+ update={"mcp_tool_permissions": {"slack": ["read"]}}
+ )
+ if change in ("disabled", "outage"):
+ with pytest.raises(HTTPException):
+ await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
+ else:
+ assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (
+ ["read"] if change == "tools" else []
+ )
+ client.db.litellm_usertable.find_unique.assert_not_called()
+ client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
+@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
+@pytest.mark.parametrize("has_grant", [True, False])
+@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
+async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
+ monkeypatch: pytest.MonkeyPatch,
+ role: str,
+ open_channel: str,
+ has_grant: bool,
+ agent_tools: tuple[str, ...] | None,
+) -> None:
+ from litellm.proxy._types import LiteLLM_TeamTable
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ manager: Final = mcp_server_manager.global_mcp_server_manager
+ manager.registry = {
+ name: MCPServer(
+ server_id=name,
+ name=name,
+ transport="http",
+ url="https://example.com/mcp",
+ allow_all_keys=open_channel == "operator",
+ )
+ for name in ("slack", "linear")
+ }
+ from litellm.proxy._experimental.mcp_server import db
+
+ monkeypatch.setattr(
+ db,
+ "get_active_submitted_mcp_server_ids_for_user",
+ AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
+ )
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
+ )
+ user: Final = LiteLLM_UserTable(
+ user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
+ )
+ team: Final = LiteLLM_TeamTable(
+ team_id="team",
+ models=[],
+ members_with_roles=[{"user_id": "human", "role": "user"}],
+ object_permission_id="team-grant",
+ )
+ client: Final = MagicMock()
+ client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
+ client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
+ client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
+ monkeypatch.setattr(proxy_server, "prisma_client", client)
+ auth: Final = actor(agent_tools, delegated=True)
+ assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
+ assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
+ assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
+ admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
+ assert admitted.user_role == role
+
+
+@pytest.mark.asyncio
+async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ from litellm.proxy._experimental.mcp_server import db
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ manager: Final = mcp_server_manager.global_mcp_server_manager
+ manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
+ monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
+ auth: Final = UserAPIKeyAuth(user_id="human")
+ auth.mcp_explicit_grants_only = True
+ with pytest.MonkeyPatch.context() as patcher:
+ patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
+ assert await manager.get_allowed_mcp_servers(auth) == []
+ auth.mcp_explicit_grants_only = False
+ assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
+
+
+@pytest.mark.asyncio
+async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
+ from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
+
+ assert await managed_agent_servers(UserAPIKeyAuth()) == ()
+ auth: Final = actor(None, delegated=True)
+ assert auth.managed_agent_context is not None
+ auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
+ assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
+ assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
+
+
+@pytest.mark.asyncio
+async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
+ permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
+ user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")]))
+ with pytest.raises(HTTPException) as failure:
+ await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
+ assert failure.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
+@pytest.mark.parametrize("scoped", (False, True))
+async def test_manager_preserves_managed_server_grants_across_open_channels(
+ monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
+) -> None:
+ from litellm.proxy._experimental.mcp_server import db
+ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ manager: Final = mcp_server_manager.global_mcp_server_manager
+ manager.registry = {
+ "open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
+ "submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
+ "passthrough": MCPServer(
+ server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
+ ),
+ }
+ monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
+ auth: Final = actor(None)
+ auth.user_role = role
+ assert not auth.mcp_explicit_grants_only
+ access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
+ assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == ({"slack"} if scoped else {"slack", "linear"})
+
+
+@pytest.mark.asyncio
+async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ manager: Final = mcp_server_manager.global_mcp_server_manager
+ manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
+ with pytest.raises(HTTPException) as failure:
+ await manager.get_allowed_mcp_servers(actor(None, delegated=True))
+ assert failure.value.status_code == 503
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py
index fc4d7b45785..5aaa3d3c446 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py
@@ -4383,7 +4383,7 @@ class TestAgentMCPPermissions:
self._team_servers({"callers": ["server_2", "server_3"]}),
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
- MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
+ MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: neither the agent's owner nor the caller has a personal grant
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({})
@@ -4402,7 +4402,7 @@ class TestAgentMCPPermissions:
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
- MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
+ MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: same seam, keyed by which user is being asked about
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": ["server_1"]})
@@ -4421,7 +4421,7 @@ class TestAgentMCPPermissions:
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
- MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
+ MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: None is the resolver's own "entitlement unresolvable" signal
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": None})
@@ -4538,7 +4538,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
- with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
+ with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_1"]
@@ -4555,7 +4555,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
- with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
+ with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = [] # no agent-level restriction
@@ -4611,7 +4611,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
- with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
+ with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_2", "server_3"]
@@ -4637,7 +4637,7 @@ class TestAgentMCPPermissions:
):
with patch.object(
MCPRequestHandler,
- "_get_agent_tool_permissions_for_server",
+ "get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=["tool_a"],
) as mock_agent_tools:
@@ -4669,7 +4669,7 @@ class TestAgentMCPPermissions:
):
with patch.object(
MCPRequestHandler,
- "_get_agent_tool_permissions_for_server",
+ "get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=None,
):
@@ -4718,7 +4718,7 @@ class TestAgentMCPPermissions:
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
- result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
+ result = await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
assert sorted(result) == ["server-a", "server-direct"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
@@ -4760,7 +4760,7 @@ class TestAgentMCPPermissions:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
with pytest.raises(UnloadableEntitlementError):
- await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
+ await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler,
@@ -4789,13 +4789,13 @@ class TestAgentMCPPermissions:
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
- server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
+ server_a_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-a", user_api_key_auth
)
- server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
+ server_b_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-b", user_api_key_auth
)
- server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
+ server_c_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-c", user_api_key_auth
)
@@ -5833,7 +5833,7 @@ def test_expand_permission_list_does_not_honor_all_proxy_sentinel():
@pytest.mark.asyncio
-async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically():
+async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically(monkeypatch):
"""The TEAM resolver expands the all-proxy sentinel to every registered server and
picks up a server registered later, so a team scoped to all-proxy tracks the live
registry without any change to its stored permission. Reverting the team-side
@@ -5850,6 +5850,9 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
+ monkeypatch.setattr(global_mcp_server_manager, "registry", {})
+ monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {})
+
for sid in ("srv-x", "srv-y"):
global_mcp_server_manager.registry[sid] = MCPServer(
server_id=sid,
@@ -8305,7 +8308,7 @@ class TestUserSubjectTeamUnion:
) == ["t1"]
# An admitted subject never fans out HERE: it resolves one source per team first, and each of
# those pins a team_id, so this helper only ever answers the single-team question. The fan-out
- # itself is _admitted_subject_sources' job, asserted below.
+ # itself is admitted_subject_sources' job, asserted below.
with self._patch(teams_by_id={}, user_teams=["t2", "t3"]):
assert await MCPRequestHandler._team_ids_for_mcp_grant(_make_admitted_subject("u")) == []
# keyless, no user_id -> nothing
@@ -8868,7 +8871,7 @@ class TestUserSubjectTeamUnion:
teams["t-member"].organization_id = "org-a"
auth = _make_admitted_subject("sso-user")
with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]):
- sources = await MCPRequestHandler._admitted_subject_sources(auth)
+ sources = await MCPRequestHandler.admitted_subject_sources(auth)
assert [(s.team_id, s.org_id) for s in sources] == [(None, None), ("t-member", "org-a")]
# The user's own source carries their grants; a team source must NOT, or the team would be
@@ -9673,7 +9676,10 @@ class TestGetUserObjectPermission:
def _prisma_with_user(self, user_row):
prisma_client = MagicMock()
- prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
+ from litellm.proxy._types import LiteLLM_UserTable
+
+ row = LiteLLM_UserTable(user_id="human", object_permission_id=user_row.object_permission_id) if user_row is not None else None
+ prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row)
return prisma_client
async def test_resolves_through_the_shared_permission_cache(self):
@@ -9688,7 +9694,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
- patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch(
"litellm.proxy.auth.auth_checks.get_object_permission",
new_callable=AsyncMock,
@@ -9715,7 +9721,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
- patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch("litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock) as mock_get_perm,
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@@ -9734,7 +9740,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
- patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@@ -9748,7 +9754,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
- patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@@ -9765,7 +9771,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
- patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
+ patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch(
"litellm.proxy.auth.auth_checks.get_object_permission",
new_callable=AsyncMock,
@@ -10085,3 +10091,47 @@ class TestScopedSessionAdmission:
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None
+
+
+@pytest.mark.asyncio
+async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch):
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_UserTable
+
+ cached = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="revoked")
+ current = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="current")
+ cache = DualCache()
+ await cache.async_set_cache(key="fresh-human", value=cached)
+ database = MagicMock()
+ database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=current)
+ database.db.litellm_usertable.find_unique = AsyncMock(return_value=cached)
+ monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
+ assert await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) == "current"
+ database.db.litellm_usertable.find_unique.assert_not_awaited()
+ database.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("unavailable")
+ with pytest.raises(HTTPException) as denied:
+ await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True)
+ assert denied.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("operation", ["servers", "tools"])
+async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted_grant(monkeypatch, operation):
+ from litellm.proxy._experimental.mcp_server import mcp_server_manager
+ from litellm.types.agents import AgentResponse
+
+ auth = UserAPIKeyAuth(agent_id="managed")
+ auth.managed_agent_policy = AgentResponse(agent_id="managed", agent_name="Managed", agent_card_params={})
+ permission = LiteLLM_ObjectPermissionTable(object_permission_id="policy", mcp_toolsets=["unavailable"])
+ manager = MagicMock()
+ manager.expand_permission_list.return_value = []
+ manager.resolve_toolset_tool_permissions = AsyncMock(side_effect=RuntimeError("policy unavailable"))
+ monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
+ resolution = (
+ MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth, permission)
+ if operation == "servers"
+ else MCPRequestHandler.get_agent_tool_permissions_for_server("slack", auth, permission)
+ )
+ with pytest.raises(RuntimeError, match="policy unavailable"):
+ await resolution
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
index f9a0075e530..7791005a070 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
@@ -7847,7 +7847,7 @@ async def test_load_active_user_by_id_reads_the_row_from_the_database_not_the_ca
key="fresh-jwt-user", value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=[]), model_type=LiteLLM_UserTable
)
prisma = MagicMock()
- prisma.db.litellm_usertable.find_unique = AsyncMock(
+ prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=["team-a"])
)
proxy_globals.user_api_key_cache = cache
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py
index 03165bd0a4a..e94371a6056 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py
@@ -1,4 +1,5 @@
import logging
+from typing import Final
import pytest
from fastapi import HTTPException
@@ -235,3 +236,13 @@ async def test_a_database_fault_retrying_cannot_clear_is_not_reported_as_a_trans
assert refusal == SubjectTokenRefusal(error="temporarily_unavailable", description=SUBJECT_TOKEN_CHECK_FAULTED)
assert "retrying will not help" in refusal.description
assert "faulted: " in caplog.text and "query engine binary not found" in caplog.text
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("user_id", [None, "delegating-user"])
+async def test_agent_token_cannot_be_exchanged_for_a_user_identity(user_id: str | None) -> None:
+ authorizer: Final = _Authorizer({**_authorized(user_id=user_id), "agent_id": "managed-agent"})
+ result: Final = await _identity(authorizer)
+ assert isinstance(result, SubjectTokenRefusal)
+ assert result.error == "invalid_request"
+ assert "direct JWT authentication" in result.description
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
index 70ef4312f4c..e6b339f24be 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py
@@ -209,15 +209,9 @@ def _reload_mcp_manager_module():
manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"]
importlib.reload(utils_module)
reloaded = importlib.reload(manager_module)
- # After reload, server.py still holds a stale reference to the old
- # global_mcp_server_manager. Update it so tests that exercise server.py
- # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance.
- server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server")
- if server_module is not None and hasattr(server_module, "global_mcp_server_manager"):
- server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
- operations_module = sys.modules.get("litellm.proxy._experimental.mcp_server.operations")
- if operations_module is not None:
- operations_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
+ for name, module in tuple(sys.modules.items()):
+ if name.startswith("litellm.proxy._experimental.mcp_server.") and hasattr(module, "global_mcp_server_manager"):
+ module.global_mcp_server_manager = reloaded.global_mcp_server_manager
return reloaded
@@ -5583,9 +5577,7 @@ class TestMCPServerManager:
# Mock dependencies - set object_permission and object_permission_id to None
# so permission checks return None (no restrictions)
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
@@ -5652,9 +5644,7 @@ class TestMCPServerManager:
# Mock dependencies - set object_permission and object_permission_id to None
# so permission checks return None (no restrictions)
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
@@ -5721,9 +5711,7 @@ class TestMCPServerManager:
# Mock dependencies - set object_permission and object_permission_id to None
# so permission checks return None (no restrictions)
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
@@ -5758,9 +5746,7 @@ class TestMCPServerManager:
# Mock dependencies - set object_permission and object_permission_id to None
# so permission checks return None (no restrictions)
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
@@ -6836,9 +6822,7 @@ class TestMCPServerManager:
# Mock dependencies - set object_permission and object_permission_id to None
# so permission checks return None (no restrictions)
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
@@ -6922,9 +6906,7 @@ class TestMCPServerManager:
manager._create_mcp_client = AsyncMock(return_value=mock_client)
# Mock user auth with no restrictions
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
+ user_api_key_auth: Final = UserAPIKeyAuth()
# Mock proxy logging
proxy_logging_obj = MagicMock()
@@ -11491,12 +11473,8 @@ class TestDiscoveryFailureLogging:
assert "unresolved" in caplog.text
-def _unrestricted_auth() -> MagicMock:
- """A caller with no object_permission, so only server-level checks apply."""
- user_api_key_auth = MagicMock()
- user_api_key_auth.object_permission = None
- user_api_key_auth.object_permission_id = None
- return user_api_key_auth
+def _unrestricted_auth() -> UserAPIKeyAuth:
+ return UserAPIKeyAuth()
def _permissive_proxy_logging() -> MagicMock:
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py
index ed3e5f48516..8484dfdde72 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py
@@ -139,7 +139,7 @@ async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_r
key="stale-cache-user", value=_user(user_id="stale-cache-user", teams=[]), model_type=LiteLLM_UserTable
)
prisma = MagicMock()
- prisma.db.litellm_usertable.find_unique = AsyncMock(
+ prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=_user(user_id="stale-cache-user", teams=["team-a"])
)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
@@ -164,7 +164,7 @@ async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_the
key="deactivated-user", value=_user(user_id="deactivated-user", teams=["team-a"]), model_type=LiteLLM_UserTable
)
prisma = MagicMock()
- prisma.db.litellm_usertable.find_unique = AsyncMock(
+ prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False})
)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
index e82ab28bb4c..759014b54c5 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py
@@ -810,6 +810,12 @@ class TestTestConnection:
from litellm.proxy._types import LitellmUserRoles
from litellm.types.mcp_server.mcp_server_manager import MCPServer
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
+ from litellm.proxy.management_endpoints import mcp_management_endpoints
+
+ manager = MCPServerManager()
+ monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
+ monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
captured = self._capture_execute(monkeypatch)
saved = MCPServer(
server_id="saved-server-id",
@@ -1311,8 +1317,9 @@ class TestListToolsRestAPI:
session_auth = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="grant-user", user_role="internal_user")
admitted_auth = UserAPIKeyAuth(user_id="grant-user", org_id="admitted-org")
- async def fake_reload(user_id):
+ async def fake_reload(user_id, *, requires_fresh_policy=False):
assert user_id == "grant-user"
+ assert requires_fresh_policy is False
return admitted_auth
monkeypatch.setattr(
@@ -1480,9 +1487,12 @@ class TestListToolsRestAPI:
from mcp.types import Tool as MCPTool
import litellm.experimental_mcp_client.client as mcp_client_module
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._experimental.mcp_server.server import MCPServer
from litellm.types.mcp import MCPTransport
+ monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", MCPServerManager())
+
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
@@ -2414,6 +2424,7 @@ class TestCallToolRestAPI:
mock_server = MagicMock()
mock_server.server_id = "server-1"
+ mock_server.name = "Example server"
def fake_get_mcp_server_by_id(server_id):
return mock_server if server_id == "server-1" else None
@@ -2431,6 +2442,11 @@ class TestCallToolRestAPI:
raising=False,
)
+ failure_log = AsyncMock()
+ execute_tool = AsyncMock()
+ monkeypatch.setattr(rest_endpoints, "_safe_fire_mcp_tool_call_failure_logging", failure_log)
+ monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", execute_tool)
+
request_payload = {
"server_id": "server-1",
"name": "demo-tool",
@@ -2452,6 +2468,16 @@ class TestCallToolRestAPI:
assert exc_info.value.detail["error"] == "access_denied"
assert "server server-1" in exc_info.value.detail["message"]
+ execute_tool.assert_not_awaited()
+ failure_log.assert_awaited_once()
+ logged_data = failure_log.await_args.args[4]
+ assert logged_data["model"] == "MCP: demo-tool"
+ assert logged_data["metadata"]["model_group"] == "MCP: demo-tool"
+ logging_obj = failure_log.await_args.args[0]
+ assert logging_obj.model_call_details["mcp_tool_call_metadata"] == {
+ "name": "demo-tool", "mcp_server_name": "Example server",
+ }
+
async def test_executes_tool_when_allowed(self, monkeypatch):
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py
index a5f6994b1a7..816ccc5e7e6 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py
@@ -145,7 +145,7 @@ async def test_build_effective_auth_contexts_appends_admitted_user_context(monke
assert contexts[-1].user_id == "user-42" and contexts[-1].team_id is None
assert [ctx.team_id for ctx in contexts[:-1]] == ["team-one"]
- reload_mock.assert_awaited_once_with("user-42")
+ reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
@pytest.mark.asyncio
@@ -198,7 +198,7 @@ async def test_acting_user_auth_returns_admitted_subject_for_non_admin_sessions(
result = await acting_user_auth(user_auth)
assert result.user_id == "user-42" and result.team_id is None
- reload_mock.assert_awaited_once_with("user-42")
+ reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
@pytest.mark.asyncio
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 a87716375e8..e5dfeb369e2 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
@@ -67,7 +67,7 @@ class TestAgentRequestHandler:
# Case 1: Both key and team have agents - intersection
with patch.object(
- AgentRequestHandler, "_get_allowed_agents_for_key"
+ AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@@ -86,7 +86,7 @@ class TestAgentRequestHandler:
# Case 2: Team has agents, key has none - inherit from team
with patch.object(
- AgentRequestHandler, "_get_allowed_agents_for_key"
+ AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@@ -105,7 +105,7 @@ class TestAgentRequestHandler:
# Case 3: Key has agents, team has none - key restrictions stand
with patch.object(
- AgentRequestHandler, "_get_allowed_agents_for_key"
+ AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@@ -120,7 +120,7 @@ class TestAgentRequestHandler:
# Case 4: No grant anywhere - unrestricted (documented open-by-default)
with patch.object(
- AgentRequestHandler, "_get_allowed_agents_for_key"
+ AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@@ -141,7 +141,7 @@ class TestAgentRequestHandler:
api_key="test-key", user_id="test-user", team_id="test-team"
)
- with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
+ with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
mock_key.return_value = RestrictedAgentAccess(frozenset({"agent-alpha"}))
mock_team.return_value = RestrictedAgentAccess(frozenset({"agent-beta"}))
@@ -198,7 +198,7 @@ class TestAgentRequestHandler:
@staticmethod
def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock:
- async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess:
+ async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess:
assert user_api_key_auth is not None
return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess())
@@ -249,7 +249,6 @@ class TestAgentRequestHandler:
frozenset({"agent-alpha"})
)
-
async def test_agent_access_groups_intersect_with_key_grants(self):
agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent")
resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"}))
@@ -299,7 +298,7 @@ class TestAgentRequestHandler:
) as mock_groups:
mock_groups.return_value = []
- assert await AgentRequestHandler._get_allowed_agents_for_key(
+ assert await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
) == RestrictedAgentAccess(frozenset())
@@ -315,7 +314,7 @@ class TestAgentRequestHandler:
) as mock_groups:
mock_groups.side_effect = Exception("DB Error")
- assert await AgentRequestHandler._get_allowed_agents_for_key(
+ assert await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
) == UnrestrictedAgentAccess()
@@ -404,7 +403,7 @@ class TestAgentRequestHandler:
)
with patch.object(
- AgentRequestHandler, "_get_allowed_agents_for_key"
+ AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@@ -489,9 +488,9 @@ class TestAgentRequestHandler:
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
- async def test_get_allowed_agents_for_key_via_access_group_ids(self):
+ async def testget_allowed_agents_for_key_via_access_group_ids(self):
"""
- Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
+ Test that get_allowed_agents_for_key includes agents from key's access_group_ids
(unified access groups) when key has no native object_permission.
"""
mock_user_auth = UserAPIKeyAuth(
@@ -508,16 +507,16 @@ class TestAgentRequestHandler:
new_callable=AsyncMock,
return_value=["agent-from-ag-1", "agent-from-ag-2"],
):
- result = await AgentRequestHandler._get_allowed_agents_for_key(
+ result = await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert result == RestrictedAgentAccess(
frozenset({"agent-from-ag-1", "agent-from-ag-2"})
)
- async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self):
+ async def testget_allowed_agents_for_key_combines_native_and_access_groups(self):
"""
- Test that _get_allowed_agents_for_key combines agents from native object_permission
+ Test that get_allowed_agents_for_key combines agents from native object_permission
and key's access_group_ids (unified access groups).
"""
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
@@ -540,7 +539,7 @@ class TestAgentRequestHandler:
new_callable=AsyncMock,
return_value=["agent-from-ag"],
):
- result = await AgentRequestHandler._get_allowed_agents_for_key(
+ result = await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert result == RestrictedAgentAccess(
@@ -611,7 +610,7 @@ class TestAgentRequestHandler:
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
registry,
):
- with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
+ with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
for key_grant, team_grant in (
(
@@ -632,3 +631,257 @@ class TestAgentRequestHandler:
assert await AgentRequestHandler.resolve_agent_access(
user_api_key_auth=mock_user_auth
) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "state,allowed",
+ [
+ ({}, True),
+ ({"enabled": False}, False),
+ ],
+)
+async def test_managed_invocation_requires_local_and_directory_admission(
+ monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool
+) -> None:
+ from unittest.mock import MagicMock
+
+ from litellm.proxy import proxy_server
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ binding: Final = AgentIdentityBinding(
+ agent_id="target",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="client",
+ issuer="issuer",
+ revision="revision",
+ )
+ target: Final = AgentResponse(
+ agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True
+ ).model_copy(update=state)
+ client: Final = MagicMock()
+ client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
+ monkeypatch.setattr(proxy_server, "prisma_client", client)
+ permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"])
+ auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
+ assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("delegated", [True, False])
+async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
+ monkeypatch: pytest.MonkeyPatch, delegated: bool
+) -> None:
+ from unittest.mock import MagicMock
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_UserTable
+ from litellm.proxy.auth import auth_checks
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import ManagedAgentContext
+
+ database: Final = MagicMock()
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"])
+ human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"])
+ human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants)
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
+ auth: Final = UserAPIKeyAuth(agent_id="actor")
+ auth.managed_agent_policy = AgentResponse(
+ agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
+ )
+ auth.managed_agent_context = ManagedAgentContext(
+ agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None
+ )
+ access: Final = await AgentRequestHandler.resolve_agent_access(auth)
+ assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"}))
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"])
+async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches(
+ monkeypatch: pytest.MonkeyPatch, revoked: str
+) -> None:
+ from unittest.mock import MagicMock
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
+ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
+
+ direct: Final = revoked == "direct-grant"
+ grouped: Final = revoked == "access-group"
+ permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
+ human: Final = LiteLLM_UserTable(
+ user_id="human",
+ teams=[] if direct else ["team"],
+ organization_memberships=[],
+ object_permission_id="grant" if direct else None,
+ )
+ team: Final = LiteLLM_TeamTable(
+ team_id="team",
+ models=[],
+ members_with_roles=[{"user_id": "human", "role": "user"}],
+ object_permission_id=None if grouped else "grant",
+ access_group_ids=["group"] if grouped else [],
+ )
+ group: Final = LiteLLM_AccessGroupTable(
+ access_group_id="group", access_group_name="Group", access_agent_ids=["target"]
+ )
+ cache: Final = UserApiKeyCache()
+ cache.set_cache("human", human)
+ cache.set_cache("team_id:team", team)
+ cache.set_cache(object_permission_cache_key("grant"), permission)
+ cache.set_cache("access_group_id:group", group)
+ client: Final = MagicMock()
+ client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human)
+ client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
+ client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
+ client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
+ monkeypatch.setattr(proxy_server, "prisma_client", client)
+ monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
+ assert await verified_human_agent_grants("human") == frozenset({"target"})
+ client.writer_db.litellm_usertable.find_unique.return_value = (
+ human.model_copy(update={"teams": []}) if revoked == "user" else human
+ )
+ client.writer_db.litellm_teamtable.find_unique.return_value = (
+ team.model_copy(update={"members_with_roles": []})
+ if revoked == "team-member"
+ else team.model_copy(update={"object_permission_id": None})
+ if revoked == "team-grant"
+ else team
+ )
+ client.writer_db.litellm_objectpermissiontable.find_unique.return_value = (
+ permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission
+ )
+ client.writer_db.litellm_accessgrouptable.find_unique.return_value = (
+ group.model_copy(update={"access_agent_ids": []}) if grouped else group
+ )
+ assert await verified_human_agent_grants("human") == frozenset()
+ client.db.litellm_usertable.find_unique.assert_not_called()
+ client.db.litellm_teamtable.find_unique.assert_not_called()
+ client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
+ client.db.litellm_accessgrouptable.find_unique.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None:
+ from unittest.mock import MagicMock
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.types.agents import AgentResponse
+
+ stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={})
+ registry: Final = AgentRegistry()
+ registry.register_agent(stale)
+ database: Final = MagicMock()
+ database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
+ database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ auth: Final = UserAPIKeyAuth(
+ object_permission=LiteLLM_ObjectPermissionTable(
+ object_permission_id="permission", agent_access_groups=["group"]
+ )
+ )
+ assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
+ frozenset({"revoked"})
+ )
+ database.writer_db.litellm_agentstable.find_many.return_value = []
+ assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
+ frozenset()
+ )
+ database.db.litellm_agentstable.find_many.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("groups", [[], ["group"]])
+async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None:
+ assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("team", [False, True])
+async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all(
+ monkeypatch: pytest.MonkeyPatch, team: bool
+) -> None:
+ from fastapi import HTTPException
+ from unittest.mock import MagicMock
+ from litellm.proxy import proxy_server
+
+ database: Final = MagicMock()
+ database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable"))
+ database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable"))
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ auth: Final = UserAPIKeyAuth(
+ team_id="team" if team else None,
+ object_permission=None if team else LiteLLM_ObjectPermissionTable(
+ object_permission_id="grant", agent_access_groups=["group"]
+ ),
+ )
+ with pytest.raises(HTTPException, match="policy is unavailable") as denied:
+ await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
+ assert denied.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("available", [False, True])
+async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
+ from unittest.mock import MagicMock
+ from litellm.proxy import proxy_server
+ from litellm.proxy.auth import auth_checks
+
+ monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None)
+ monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None))
+ assert await AgentRequestHandler._get_allowed_agents_for_team(
+ UserAPIKeyAuth(team_id="missing"), strict=True
+ ) == RestrictedAgentAccess(frozenset())
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("outage", [False, True])
+async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy(
+ monkeypatch: pytest.MonkeyPatch, outage: bool
+) -> None:
+ from fastapi import HTTPException
+ from unittest.mock import MagicMock
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.types.agents import AgentResponse
+
+ registry: Final = AgentRegistry()
+ registry.register_agent(AgentResponse(
+ agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True
+ ))
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(
+ return_value=None, side_effect=ConnectionError("unavailable") if outage else None
+ )
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ if outage:
+ with pytest.raises(HTTPException, match="could not be loaded") as denied:
+ await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth())
+ assert denied.value.status_code == 503
+ else:
+ assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("grant", [False, True])
+async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None:
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import ManagedAgentContext
+ from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
+
+ auth: Final = UserAPIKeyAuth(agent_id="actor")
+ auth.managed_agent_policy = AgentResponse(
+ agent_id="actor", agent_name="Actor", agent_card_params={},
+ object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None,
+ )
+ auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated")
+ assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset())
+ assert await verified_human_agent_grants(None) == frozenset()
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
new file mode 100644
index 00000000000..2e37f0f4462
--- /dev/null
+++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py
@@ -0,0 +1,484 @@
+from collections.abc import Mapping
+from typing import Final
+from unittest.mock import AsyncMock, MagicMock
+
+import pytest
+from fastapi import HTTPException
+
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy.agent_endpoints.auth.managed_authorization import (
+ actor_admission_failure,
+ admit_managed_actor,
+ invocation_target,
+)
+from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext
+
+BINDING: Final = AgentIdentityBinding(
+ agent_id="agent",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="client",
+ service_principal_id="principal",
+ issuer="issuer",
+ revision="current",
+)
+
+
+def agent(**overrides: object) -> AgentResponse:
+ return AgentResponse.model_validate(
+ {
+ "agent_id": "agent",
+ "agent_name": "Agent",
+ "agent_card_params": {},
+ "identity": BINDING,
+ "identity_managed": True,
+ "execution_mode": "both",
+ **overrides,
+ }
+ )
+
+
+@pytest.mark.parametrize(
+ "state",
+ [
+ {"enabled": False},
+ {"identity": None},
+ {"identity": BINDING.model_copy(update={"active": False})},
+ {"execution_mode": "delegated"},
+ ],
+)
+def test_keys_cannot_bypass_lifecycle_or_delegated_only_mode(state: dict[str, object]) -> None:
+ assert isinstance(actor_admission_failure(agent(**state), None), AgentIdentityFailure)
+
+
+@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
+def test_keys_cannot_impersonate_an_entra_bound_agent(mode: str) -> None:
+ assert isinstance(actor_admission_failure(agent(execution_mode=mode), None), AgentIdentityFailure)
+
+
+@pytest.mark.parametrize(
+ "context",
+ [
+ ManagedAgentContext(agent_id="agent", binding_revision="previous", mode="autonomous"),
+ ManagedAgentContext(agent_id="another", binding_revision="current", mode="autonomous"),
+ ManagedAgentContext(agent_id="agent", binding_revision="current", mode="delegated"),
+ ],
+)
+def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context: ManagedAgentContext) -> None:
+ assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure)
+
+
+def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
+ context: Final = ManagedAgentContext(
+ agent_id="agent", binding_revision="current", mode="delegated", user_id="human"
+ )
+ auth: Final = UserAPIKeyAuth.model_validate(
+ {
+ "managed_agent_context": context,
+ "requires_fresh_policy": True,
+ "mcp_explicit_grants_only": True,
+ "managed_agent_policy": agent(),
+ "billing_agent_policy": agent(),
+ "invoked_agent_id": "forged-target",
+ "agent_invocation_cost": 0.0,
+ }
+ )
+ assert auth.requires_fresh_policy is False
+ assert auth.mcp_explicit_grants_only is False
+ assert "mcp_explicit_grants_only" not in auth.model_dump()
+ assert auth.managed_agent_context is None
+ assert auth.managed_agent_policy is None
+ assert auth.billing_agent_policy is None
+ assert auth.invoked_agent_id is None
+ assert auth.agent_invocation_cost is None
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("autonomous", (True, False))
+async def test_invocation_prepares_target_fee_for_the_correct_agent(
+ monkeypatch: pytest.MonkeyPatch,
+ autonomous: bool,
+) -> None:
+ from unittest.mock import AsyncMock, MagicMock
+
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+ from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
+
+ target: Final = agent(litellm_params={"cost_per_query": 0.25})
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(target)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"])
+ auth: Final = UserAPIKeyAuth(
+ agent_id="caller" if autonomous else None,
+ user_id=None if autonomous else "human",
+ object_permission=permission,
+ )
+ if autonomous:
+ caller: Final = agent(agent_id="caller", object_permission=permission.model_dump())
+ auth.managed_agent_policy = caller
+ auth.billing_agent_policy = caller
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
+ assert auth.agent_invocation_cost == pytest.approx(0.25)
+ assert auth.invoked_agent_id == "agent"
+ assert auth.billing_agent_policy is not None
+ assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent")
+
+
+@pytest.mark.asyncio
+async def test_deleted_agent_key_cannot_fall_back_to_unmanaged_authentication() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
+ database.writer_db.litellm_retiredagent.find_unique = AsyncMock(return_value={"original_agent_id": "deleted"})
+ with pytest.raises(HTTPException, match="Agent no longer exists"):
+ await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
+ database.writer_db.litellm_retiredagent.find_unique.return_value = None
+ auth: Final = UserAPIKeyAuth(agent_id="legacy-attribution-label")
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert auth.managed_agent_policy is None
+ database.db.litellm_agentstable.find_unique.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
+ database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("unavailable"))
+ with pytest.raises(HTTPException) as failure:
+ await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
+ assert failure.value.status_code == 503
+
+
+@pytest.mark.parametrize(
+ "route,body,expected",
+ [
+ ("/a2a/agent", {}, "agent"),
+ ("/v1/a2a/agent/", {}, "agent"),
+ ("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"),
+ ("/v1/chat/completions", {"model": "a2a/"}, None),
+ ("/v1/chat/completions", {"model": "ordinary-model"}, None),
+ ("/a2a", {}, None),
+ ],
+)
+def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, object], expected: str | None) -> None:
+ assert invocation_target(route, body) == expected
+
+
+@pytest.mark.asyncio
+async def test_agent_admission_database_outage_fails_closed() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(side_effect=RuntimeError("DB unavailable"))
+ with pytest.raises(HTTPException) as failure:
+ await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
+ assert failure.value.status_code == 503
+
+
+@pytest.mark.asyncio
+async def test_human_authentication_does_not_load_an_agent() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock()
+ await admit_managed_actor(UserAPIKeyAuth(user_id="human"), AgentIdentityStore.from_client(database))
+ database.writer_db.litellm_agentstable.find_unique.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_disabled_agent_key_is_rejected_at_admission() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(enabled=False))
+ with pytest.raises(HTTPException) as failure:
+ await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
+ assert failure.value.status_code == 403
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("permitted", [True, False])
+async def test_verified_human_still_needs_an_explicit_agent_invocation_grant(
+ monkeypatch: pytest.MonkeyPatch,
+ permitted: bool,
+) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable
+ from litellm.proxy.auth import auth_checks
+
+ policy: Final = agent()
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="human-grants",
+ agents=["agent"] if permitted else [],
+ )
+ human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
+ monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ auth.managed_agent_context = ManagedAgentContext(
+ agent_id="agent",
+ binding_revision="current",
+ mode="delegated",
+ user_id="human",
+ )
+ if permitted:
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert auth.managed_agent_policy == policy
+ assert auth.billing_agent_policy == policy
+ else:
+ with pytest.raises(HTTPException) as failure:
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert failure.value.status_code == 403
+
+
+def test_execution_mode_must_match_verified_token_mode() -> None:
+ context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
+ failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context)
+ assert isinstance(failure, AgentIdentityFailure)
+ assert "execution mode" in failure.message
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("state,status", [("missing", 403), ("outage", 503), ("denied", 403), ("invalid-fee", 503)])
+async def test_invocation_cannot_bypass_missing_policy_permission_or_invalid_price(
+ monkeypatch: pytest.MonkeyPatch, state: str, status: int
+) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ registered: Final = agent(litellm_params={"cost_per_query": -1 if state == "invalid-fee" else 0.25})
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(registered)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(
+ return_value=None if state == "missing" else registered,
+ side_effect=RuntimeError("unavailable") if state == "outage" else None,
+ )
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ permission: Final = LiteLLM_ObjectPermissionTable(
+ object_permission_id="grant", agents=[] if state == "denied" else ["agent"]
+ )
+ auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
+ with pytest.raises(HTTPException) as failure:
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
+ assert failure.value.status_code == status
+ assert auth.agent_invocation_cost is None
+
+
+@pytest.mark.asyncio
+async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(execution_mode="autonomous"))
+ auth: Final = UserAPIKeyAuth(agent_id="agent", jwt_claims={"agent": "agent", "sub": "unrelated-subject"})
+ with pytest.raises(HTTPException) as denied:
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert denied.value.status_code == 403
+ assert auth.managed_agent_policy is None
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("bound", [False, True])
+async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.MonkeyPatch, bound: bool) -> None:
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
+
+ registry: Final = AgentRegistry()
+ registry.register_agent(agent(identity_managed=bound, identity=BINDING if bound else None))
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ if not bound:
+ auth.managed_agent_context = ManagedAgentContext(
+ agent_id="agent", binding_revision="current", mode="autonomous"
+ )
+ with pytest.raises(HTTPException) as denied:
+ await admit_managed_actor(auth, None)
+ assert denied.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("managed_flag", [False, True])
+async def test_managed_invocation_requires_database(monkeypatch: pytest.MonkeyPatch, managed_flag: bool) -> None:
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ registry: Final = AgentRegistry()
+ registry.register_agent(agent(identity_managed=managed_flag))
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ with pytest.raises(HTTPException) as denied:
+ await prepare_agent_invocation(UserAPIKeyAuth(user_id="human"), "agent", None)
+ assert denied.value.status_code == 503
+
+
+@pytest.mark.asyncio
+async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> None:
+ policy: Final = agent(execution_mode="autonomous")
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ auth: Final = UserAPIKeyAuth(agent_id="agent", api_key="persisted-key")
+ with pytest.raises(HTTPException, match="bound identity provider token") as denied:
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert denied.value.status_code == 403
+ assert auth.managed_agent_policy is None
+ assert auth.billing_agent_policy is None
+
+
+@pytest.mark.asyncio
+async def test_unknown_invocation_target_leaves_billing_unset(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
+ monkeypatch.setattr(proxy_server, "prisma_client", None)
+ auth: Final = UserAPIKeyAuth(user_id="human")
+ await prepare_agent_invocation(auth, "missing", None)
+ assert auth.invoked_agent_id is None
+ assert auth.billing_agent_policy is None
+
+
+@pytest.mark.parametrize(
+ "route,method,allowed",
+ [
+ ("/v1/agents", "GET", True),
+ ("/v1/agents", "POST", False),
+ ("/v1/chat/completions", "POST", True),
+ ("/v1/chat/completions", "DELETE", False),
+ ("/openai/deployments/model/chat/completions", "POST", True),
+ ("/engines/openai/model/chat/completions", "POST", True),
+ ("/openai/deployments/openai/model/images/generations", "POST", True),
+ ("/openai/deployments/openai/model/images/edits", "POST", True),
+ ("/v1beta/models/gemini-model:generateContent", "POST", True),
+ ("/v1/realtime", "GET", True),
+ ("/v1/realtime", "POST", False),
+ ("/v1/realtime/client_secrets", "POST", False),
+ ("/mcp/tools/call", "POST", True),
+ ("/a2a/target/message/send", "POST", True),
+ ("/v1/a2a/target/message/send", "POST", True),
+ ("/v1/videos", "POST", False),
+ ("/v1/videos/other-video", "GET", False),
+ ("/v1/search", "POST", False),
+ ("/search", "POST", False),
+ ("/v1/agents/target", "PATCH", False),
+ ("/v1/responses/other-response", "GET", False),
+ ("/v1/files", "GET", False),
+ ("/v1/files", "POST", False),
+ ("/openai/v1/files", "GET", False),
+ ("/anthropic/v1/files", "GET", False),
+ ],
+)
+def test_managed_route_scope_excludes_provider_resources(route: str, method: str, allowed: bool) -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
+
+ assert managed_agent_route_allowed(route, method) is allowed
+
+
+@pytest.mark.parametrize(
+ "route,body,settings,cli_model,path_model,expected",
+ [
+ ("/v1/chat/completions", {"model": "body"}, {"completion_model": "default"}, "cli", "path", "default"),
+ ("/v1/moderations", {"model": "body"}, {"moderation_model": "default"}, "cli", None, "cli"),
+ ("/v1/audio/speech", {"model": "body"}, {"completion_model": "ignored"}, None, None, "body"),
+ ("/openai/deployments/path/embeddings", {"model": "body"}, {}, None, "path", "path"),
+ ("/v1/messages/count_tokens", {"model": "body"}, {"completion_model": "ignored"}, "cli", None, "body"),
+ ("/mcp/tools/call", {}, {"completion_model": "ignored"}, "cli", None, None),
+ ("/v1/images/generations", {"model": "image"}, {"completion_model": "text"}, None, None, "image"),
+ ("/v1/images/generations", {}, {"image_generation_model": "image"}, None, None, "image"),
+ ("/v1/images/edits", {}, {"image_generation_model": "image"}, None, None, "image"),
+ ("/v1/rerank", {"model": "reranker"}, {"completion_model": "text"}, "cli", None, "reranker"),
+ ("/v1beta/models/path:countTokens", {"model": "body"}, {"completion_model": "text"}, "cli", "path", "path"),
+ ],
+)
+def test_managed_inference_resolves_dispatch_precedence(
+ route: str,
+ body: Mapping[str, object],
+ settings: Mapping[str, object],
+ cli_model: str | None,
+ path_model: str | None,
+ expected: str | None,
+) -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ assert managed_inference_request(route, body, settings, cli_model, path_model).get("model") == expected
+
+
+def test_managed_inference_without_any_model_cannot_skip_model_grants():
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ with pytest.raises(HTTPException, match="explicit or configured model"):
+ managed_inference_request("/v1/moderations", {}, {}, None)
+
+
+@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/images/generations", "/v1/images/edits"])
+def test_managed_inference_query_model_takes_precedence_over_body(route: str):
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ assert managed_inference_request(route, {"model": "body"}, {}, None, query_model="query")["model"] == "query"
+
+
+def test_managed_inference_ignores_unsupported_query_model():
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ assert (
+ managed_inference_request("/v1/messages", {"model": "body"}, {}, None, query_model="query")["model"] == "body"
+ )
+
+
+@pytest.mark.parametrize("route", ["/realtime", "/v1/realtime", "/openai/v1/realtime"])
+def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route: str) -> None:
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
+
+ with pytest.raises(HTTPException, match="explicit or configured model"):
+ managed_inference_request(route, {}, {"completion_model": "allowed-default"}, "cli")
+ assert managed_inference_request(
+ route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli"
+ )["model"] == "requested"
+
+
+@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")])
+def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None:
+ context: Final = ManagedAgentContext.model_validate(
+ {"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user}
+ )
+ assert actor_admission_failure(agent(), context) is None
+
+
+@pytest.mark.asyncio
+async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
+
+ legacy: Final = agent(identity=None, identity_managed=False)
+ registry: Final = agent_registry.AgentRegistry()
+ registry.register_agent(legacy)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(proxy_server, "prisma_client", None)
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ await admit_managed_actor(auth, None)
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=legacy)
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
+ assert auth.managed_agent_policy is None
+ assert auth.billing_agent_policy is None
+ assert auth.invoked_agent_id is None
+
+
+@pytest.mark.asyncio
+async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent())
+ auth: Final = UserAPIKeyAuth(agent_id="agent")
+ auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
+ await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
+ assert auth.managed_agent_policy == agent()
+ assert auth.billing_agent_policy == agent()
+ assert auth.user_id is None
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
index 8a7ab0f0001..a5d0d0a3ecc 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py
@@ -59,6 +59,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
# Mock agent
mock_agent = MagicMock()
+ mock_agent.agent_id = "test-agent"
mock_agent.agent_card_params = {
"url": "http://backend-agent:10001",
"name": "Test Agent",
@@ -72,6 +73,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
"jsonrpc": "2.0",
"id": "test-id",
"method": "message/send",
+ "metadata": {"model_info": {"id": "caller-supplied-id"}},
"params": {
"message": {
"role": "user",
@@ -153,7 +155,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
"litellm.a2a_protocol.asend_message",
new_callable=AsyncMock,
return_value=mock_response,
- ),
+ ) as mock_send_message,
patch(
"litellm.proxy.proxy_server.general_settings",
{},
@@ -190,6 +192,9 @@ async def test_invoke_agent_a2a_adds_litellm_data():
mock_add_data.assert_called_once()
# Verify model and custom_llm_provider were set
+ assert mock_send_message.await_args.kwargs["model"] == "a2a_agent/Test Agent"
+ assert captured_data["metadata"]["model_group"] == "a2a_agent/Test Agent"
+ assert captured_data["metadata"]["model_info"] == {"id": mock_agent.agent_id}
assert captured_data.get("model") == "a2a_agent/Test Agent"
assert captured_data.get("custom_llm_provider") == "a2a_agent"
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
index ef20e88c368..ba160b8e1b4 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
@@ -2,11 +2,14 @@
import hashlib
import json
+from collections.abc import Mapping
+from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
+from prisma.models import LiteLLM_AgentsTable
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.proxy.agent_endpoints.agent_registry import (
@@ -451,11 +454,11 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update():
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=None)
+ return_value=_stored_agent_row(SimpleNamespace(litellm_params={}, object_permission_id=None))
)
mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None)
- with pytest.raises(Exception, match="Error updating agent in DB") as exc_info:
+ with pytest.raises(Exception, match="Agent not found") as exc_info:
await registry.update_agent_in_db(
agent_id="agent-123",
agent={
@@ -467,7 +470,7 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update():
updated_by="test-user",
)
- assert str(exc_info.value) == "Error updating agent in DB: Agent not found, passed agent_id=agent-123"
+ assert str(exc_info.value) == "Agent not found, passed agent_id=agent-123"
@pytest.mark.asyncio
@@ -476,11 +479,13 @@ async def test_patch_agent_in_db_raises_when_row_deleted_mid_update():
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None}
+ return_value=_stored_agent_row(
+ {"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None}
+ )
)
mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None)
- with pytest.raises(Exception, match="Error patching agent in DB") as exc_info:
+ with pytest.raises(Exception, match="Agent not found") as exc_info:
await registry.patch_agent_in_db(
agent_id="agent-123",
agent={"agent_name": "Patched Agent"},
@@ -488,20 +493,43 @@ async def test_patch_agent_in_db_raises_when_row_deleted_mid_update():
updated_by="test-user",
)
- assert str(exc_info.value) == "Error patching agent in DB: Agent not found, passed agent_id=agent-123"
+ assert str(exc_info.value) == "Agent not found, passed agent_id=agent-123"
@pytest.mark.asyncio
-async def test_delete_agent_from_db_raises_when_row_already_gone():
- """Prisma's delete returns None for a missing row, which dict() cannot consume."""
+async def test_delete_agent_from_db_raises_when_row_already_gone() -> None:
registry: Final = AgentRegistry()
- mock_prisma: Final = MagicMock()
- mock_prisma.db.litellm_agentstable.delete = AsyncMock(return_value=None)
+ database: Final = MagicMock()
+ tx: Final = database.tx.return_value.__aenter__.return_value
+ tx.litellm_agentstable.find_unique = AsyncMock(return_value=None)
+ with pytest.raises(ValueError, match="Agent not found, passed agent_id=agent-123"):
+ await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=database)
+ tx.litellm_verificationtoken.delete_many.assert_not_called()
- with pytest.raises(Exception, match="Error deleting agent from DB") as exc_info:
- await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=mock_prisma)
- assert str(exc_info.value) == "Error deleting agent from DB: Agent not found, passed agent_id=agent-123"
+@pytest.mark.asyncio
+@pytest.mark.parametrize("managed", [True, False])
+async def test_agent_deletion_revokes_managed_keys_and_keeps_identity_history(managed: bool) -> None:
+ registry: Final = AgentRegistry()
+ database: Final = MagicMock()
+ tx: Final = database.tx.return_value.__aenter__.return_value
+ row: Final = _stored_agent_row({"agent_id": "agent-123", "identity_managed": managed})
+ tx.litellm_agentstable.find_unique = AsyncMock(return_value=row)
+ tx.litellm_agentstable.delete = AsyncMock(return_value=row)
+ tx.litellm_verificationtoken.delete_many = AsyncMock(return_value=2)
+ tx.litellm_retiredagent.upsert = AsyncMock()
+ result: Final = await registry.delete_agent_from_db("agent-123", database)
+ assert result["agent_id"] == "agent-123"
+ tx.litellm_agentstable.delete.assert_awaited_once_with(where={"agent_id": "agent-123"})
+ if managed:
+ tx.litellm_retiredagent.upsert.assert_awaited_once_with(
+ where={"original_agent_id": "agent-123"},
+ data={"create": {"original_agent_id": "agent-123"}, "update": {}},
+ )
+ tx.litellm_verificationtoken.delete_many.assert_awaited_once_with(where={"agent_id": "agent-123"})
+ else:
+ tx.litellm_retiredagent.upsert.assert_not_awaited()
+ tx.litellm_verificationtoken.delete_many.assert_not_awaited()
# ---------- LIT-6736: agent litellm_params secret redaction ----------
@@ -729,14 +757,15 @@ async def test_update_agent_in_db_preserves_secret_when_echoed_back_redacted():
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(
- litellm_params={
- "aws_access_key_id": SENTINEL_AWS_ACCESS_KEY_ID,
- "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
- "model": "bedrock/agentcore/my-agent",
- },
- object_permission_id=None,
- kill_switch=None,
+ return_value=_stored_agent_row(
+ SimpleNamespace(
+ litellm_params={
+ "aws_access_key_id": SENTINEL_AWS_ACCESS_KEY_ID,
+ "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
+ "model": "bedrock/agentcore/my-agent",
+ },
+ object_permission_id=None,
+ )
)
)
updated_agent = MagicMock()
@@ -782,10 +811,11 @@ async def test_update_agent_in_db_preserves_secret_when_key_omitted_entirely():
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(
- litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
- object_permission_id=None,
- kill_switch=None,
+ return_value=_stored_agent_row(
+ SimpleNamespace(
+ litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
+ object_permission_id=None,
+ )
)
)
updated_agent = MagicMock()
@@ -824,15 +854,16 @@ async def test_update_agent_in_db_preserves_secret_nested_under_a_non_sensitive_
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(
- litellm_params={
- "provider_config": {
- "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
- "region": "us-east-1",
- }
- },
- object_permission_id=None,
- kill_switch=None,
+ return_value=_stored_agent_row(
+ SimpleNamespace(
+ litellm_params={
+ "provider_config": {
+ "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
+ "region": "us-east-1",
+ }
+ },
+ object_permission_id=None,
+ )
)
)
updated_agent = MagicMock()
@@ -878,10 +909,11 @@ async def test_update_agent_in_db_clears_secret_on_explicit_empty_value():
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(
- litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
- object_permission_id=None,
- kill_switch=None,
+ return_value=_stored_agent_row(
+ SimpleNamespace(
+ litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
+ object_permission_id=None,
+ )
)
)
updated_agent = MagicMock()
@@ -919,12 +951,14 @@ async def test_patch_agent_in_db_preserves_secret_when_litellm_params_omitted():
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Old Name",
- "litellm_params": {"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
- "object_permission_id": None,
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Old Name",
+ "litellm_params": {"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
+ "object_permission_id": None,
+ }
+ )
)
patched_agent = MagicMock()
patched_agent.model_dump.return_value = {
@@ -958,15 +992,17 @@ async def test_patch_agent_in_db_preserves_secret_when_echoed_back_redacted():
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Test Agent",
- "litellm_params": {
- "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
- "is_public": False,
- },
- "object_permission_id": None,
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Test Agent",
+ "litellm_params": {
+ "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY,
+ "is_public": False,
+ },
+ "object_permission_id": None,
+ }
+ )
)
patched_agent = MagicMock()
patched_agent.model_dump.return_value = {
@@ -997,6 +1033,48 @@ async def test_patch_agent_in_db_preserves_secret_when_echoed_back_redacted():
assert stored_params["is_public"] is True
+@pytest.mark.asyncio
+@pytest.mark.parametrize("operation", ["patch", "put"])
+async def test_runtime_update_drops_legacy_identity_and_keeps_agent_id(operation: str) -> None:
+ registry: Final = AgentRegistry()
+ prisma: Final = MagicMock()
+ identity: Final = {
+ "provider": "microsoft_entra",
+ "tenant_id": "11111111-1111-4111-8111-111111111111",
+ "client_id": "22222222-2222-4222-8222-222222222222",
+ }
+ existing_params: Final = {"identity": identity, "model": "old"}
+ existing: Final = (
+ SimpleNamespace(litellm_params=existing_params, object_permission_id=None)
+ if operation == "put"
+ else {"agent_name": "Readable agent", "litellm_params": existing_params}
+ )
+ prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row(existing))
+ saved: Final = MagicMock()
+ saved.object_permission = None
+ saved.model_dump.return_value = {
+ "agent_id": "unchanged-id",
+ "agent_name": "Renamed agent",
+ "agent_card_params": {},
+ "litellm_params": {"model": "new"},
+ }
+ prisma.db.litellm_agentstable.update = AsyncMock(return_value=saved)
+ update: Final = registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db
+ result: Final = await update(
+ agent_id="unchanged-id",
+ agent={"agent_name": "Renamed agent", "agent_card_params": {}, "litellm_params": {"model": "new"}},
+ prisma_client=prisma,
+ updated_by="admin",
+ )
+ stored: Final = prisma.db.litellm_agentstable.update.call_args.kwargs
+ assert stored["where"] == {"agent_id": "unchanged-id"}
+ assert json.loads(stored["data"]["litellm_params"]) == {"model": "new"}, (
+ "a stored litellm_params.identity must not be resurrected once the JWT path no longer honours it"
+ )
+ assert result.agent_id == "unchanged-id"
+ assert "object_permission_id" not in stored["data"]
+
+
def _agent_row_mock(access_group_ids: list[str]) -> MagicMock:
row: Final = MagicMock()
row.model_dump.return_value = {
@@ -1063,13 +1141,15 @@ async def test_patch_agent_in_db_replaces_access_group_ids_when_provided(
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Test Agent",
- "litellm_params": {},
- "object_permission_id": None,
- "access_group_ids": ["ag-1"],
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Test Agent",
+ "litellm_params": {},
+ "object_permission_id": None,
+ "access_group_ids": ["ag-1"],
+ }
+ )
)
mock_update = AsyncMock(return_value=_agent_row_mock(expected))
mock_prisma.db.litellm_agentstable.update = mock_update
@@ -1086,13 +1166,15 @@ async def test_patch_agent_in_db_keeps_access_group_ids_when_omitted():
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Old Name",
- "litellm_params": {},
- "object_permission_id": None,
- "access_group_ids": ["ag-1"],
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Old Name",
+ "litellm_params": {},
+ "object_permission_id": None,
+ "access_group_ids": ["ag-1"],
+ }
+ )
)
mock_update = AsyncMock(return_value=_agent_row_mock(["ag-1"]))
mock_prisma.db.litellm_agentstable.update = mock_update
@@ -1114,8 +1196,8 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(
- litellm_params={}, object_permission_id=None, kill_switch=None, access_group_ids=["ag-1"]
+ return_value=_stored_agent_row(
+ SimpleNamespace(litellm_params={}, object_permission_id=None, access_group_ids=["ag-1"])
)
)
mock_update = AsyncMock(return_value=_agent_row_mock(expected))
@@ -1134,6 +1216,34 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro
assert tuple(mock_update.call_args.kwargs["data"]["access_group_ids"]) == tuple(expected)
+def _stored_agent_row(values: Mapping[str, object] | SimpleNamespace) -> LiteLLM_AgentsTable:
+ fields: Final = vars(values) if isinstance(values, SimpleNamespace) else values
+ return LiteLLM_AgentsTable.model_validate(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Test Agent",
+ "agent_card_params": "{}",
+ "extra_headers": [],
+ "agent_access_groups": [],
+ "access_group_ids": [],
+ "created_at": datetime.now(timezone.utc),
+ "updated_at": datetime.now(timezone.utc),
+ "created_by": "admin",
+ "updated_by": "admin",
+ "spend": 0,
+ "identity_managed": False,
+ "enabled": True,
+ "execution_mode": "autonomous",
+ **{
+ key: json.dumps(value)
+ if key in ("litellm_params", "agent_card_params", "kill_switch") and not isinstance(value, str)
+ else value
+ for key, value in fields.items()
+ },
+ }
+ )
+
+
_KILL_SWITCH: Final = {
"url": "https://ops.example.com/kill",
"method": "POST",
@@ -1194,13 +1304,15 @@ async def test_patch_agent_in_db_keeps_kill_switch_when_omitted_and_clears_it_on
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Old",
- "litellm_params": {},
- "object_permission_id": None,
- "kill_switch": _KILL_SWITCH,
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "Old",
+ "litellm_params": {},
+ "object_permission_id": None,
+ "kill_switch": _KILL_SWITCH,
+ }
+ )
)
mock_update = AsyncMock(return_value=_agent_row_mock([]))
mock_prisma.db.litellm_agentstable.update = mock_update
@@ -1223,13 +1335,15 @@ async def test_patch_agent_in_db_restores_the_stored_kill_switch_secret_behind_t
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "A",
- "litellm_params": {},
- "object_permission_id": None,
- "kill_switch": _KILL_SWITCH,
- }
+ return_value=_stored_agent_row(
+ {
+ "agent_id": "agent-123",
+ "agent_name": "A",
+ "litellm_params": {},
+ "object_permission_id": None,
+ "kill_switch": _KILL_SWITCH,
+ }
+ )
)
mock_update = AsyncMock(return_value=_agent_row_mock([]))
mock_prisma.db.litellm_agentstable.update = mock_update
@@ -1258,7 +1372,9 @@ async def test_update_agent_in_db_clears_kill_switch_when_omitted_and_restores_s
registry: Final = AgentRegistry()
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH))
+ return_value=_stored_agent_row(
+ SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH))
+ )
)
mock_update = AsyncMock(return_value=_agent_row_mock([]))
mock_prisma.db.litellm_agentstable.update = mock_update
@@ -1284,3 +1400,145 @@ def test_load_agents_from_config_exposes_a_typed_kill_switch():
(agent,) = registry.get_agent_list()
assert agent.kill_switch is not None
assert agent.kill_switch.model_dump() == _KILL_SWITCH
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("bound", [False, True])
+async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) -> None:
+ from datetime import datetime, timezone
+
+ from prisma.models import LiteLLM_AgentIdentity, LiteLLM_AgentsTable
+
+ from litellm.types.agents import AgentResponse
+
+ binding: Final = LiteLLM_AgentIdentity(
+ agent_id="agent",
+ provider="microsoft_entra",
+ issuer="issuer",
+ tenant_id="tenant",
+ client_id="client",
+ active=True,
+ required_roles=[],
+ required_scopes=["user_impersonation"],
+ revision="revision",
+ )
+ row: Final = LiteLLM_AgentsTable(
+ agent_id="agent",
+ agent_name="Bound agent",
+ agent_card_params="{}",
+ identity_managed=bound,
+ identity=binding if bound else None,
+ enabled=True,
+ execution_mode="autonomous",
+ spend=0.0,
+ agent_access_groups=[],
+ access_group_ids=[],
+ extra_headers=[],
+ created_by="admin",
+ updated_by="admin",
+ created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
+ updated_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
+ )
+ client: Final = MagicMock()
+ client.db.litellm_agentstable.find_many = AsyncMock(return_value=[row])
+ listed: Final = await AgentRegistry.get_all_agents_from_db(client)
+ response: Final = AgentResponse.model_validate(listed[0])
+ if bound:
+ assert response.identity is not None
+ assert response.identity.client_id == binding.client_id
+ assert response.identity.revision == binding.revision
+ else:
+ assert response.identity is None
+ client.db.litellm_agentstable.find_many.assert_awaited_once_with(
+ order={"created_at": "desc"},
+ include={"object_permission": True, "identity": True},
+ )
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("operation", ["create", "patch", "put"])
+async def test_agent_permissions_are_written_atomically_with_the_registration(operation: str) -> None:
+ from litellm.proxy._types import LiteLLM_ObjectPermissionTable
+
+ registry: Final = AgentRegistry()
+ client: Final = MagicMock()
+ existing: Final = _stored_agent_row({"agent_id": "agent-123", "object_permission_id": "permissions"})
+ client.db.litellm_agentstable.find_unique = AsyncMock(return_value=existing)
+ client.db.litellm_agentstable.create = AsyncMock(return_value=existing)
+ client.db.litellm_agentstable.update = AsyncMock(return_value=existing)
+ client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
+ return_value=(
+ LiteLLM_ObjectPermissionTable(object_permission_id="permissions", models=["prior"], mcp_servers=["slack"])
+ if operation != "create"
+ else None
+ )
+ )
+ incoming: Final = {"agent_name": "Agent", "agent_card_params": {}, "object_permission": {"models": ["new"]}}
+ if operation == "create":
+ await registry.add_agent_to_db(incoming, client, created_by="admin")
+ else:
+ update: Final = registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db
+ await update("agent-123", incoming, client, updated_by="admin")
+ write: Final = (
+ client.db.litellm_agentstable.create if operation == "create" else client.db.litellm_agentstable.update
+ )
+ permission: Final = write.call_args.kwargs["data"]["object_permission"][
+ "create" if operation == "create" else "update"
+ ]
+ assert permission["models"] == ["new"]
+ if operation != "create":
+ assert permission["mcp_servers"] == ["slack"]
+ assert permission["object_permission_id"] == "permissions"
+ client.db.litellm_objectpermissiontable.update.assert_not_called()
+ client.db.litellm_objectpermissiontable.create.assert_not_called()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("operation", ["create", "patch", "put"])
+async def test_invalid_identity_fails_before_registration_is_written(operation: str) -> None:
+ from fastapi import HTTPException
+
+ registry: Final = AgentRegistry()
+ client: Final = MagicMock()
+ client.db.litellm_agentstable.create = AsyncMock()
+ client.db.litellm_agentstable.update = AsyncMock()
+ client.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row({"agent_id": "agent-123"}))
+ incoming: Final = {"agent_name": "Agent", "agent_card_params": {}, "identity": {"provider": "unknown"}}
+ write: Final = (
+ registry.add_agent_to_db(incoming, client, created_by="admin")
+ if operation == "create"
+ else (registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db)(
+ "agent-123", incoming, client, updated_by="admin"
+ )
+ )
+ with pytest.raises(HTTPException) as failure:
+ await write
+ assert failure.value.status_code == 400
+ client.db.litellm_agentstable.create.assert_not_awaited()
+ client.db.litellm_agentstable.update.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("operation", ["create", "patch", "put"])
+async def test_duplicate_agent_binding_returns_conflict_for_every_write(operation: str) -> None:
+ from fastapi import HTTPException
+ from prisma.errors import UniqueViolationError
+
+ registry: Final = AgentRegistry()
+ client: Final = MagicMock()
+ client.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row({"agent_id": "agent-123"}))
+ failure: Final = UniqueViolationError({"user_facing_error": {"message": "Unique constraint failed", "meta": {"target": ["client_id"]}, "error_code": "P2002"}})
+ client.db.litellm_agentstable.create = AsyncMock(side_effect=failure)
+ client.db.litellm_agentstable.update = AsyncMock(side_effect=failure)
+ incoming: Final = {"agent_name": "Agent", "agent_card_params": {}}
+ write: Final = (
+ registry.add_agent_to_db(incoming, client, created_by="admin")
+ if operation == "create"
+ else (registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db)(
+ "agent-123", incoming, client, updated_by="admin"
+ )
+ )
+ with pytest.raises(HTTPException) as denied:
+ await write
+ assert denied.value.status_code == 409
+ assert denied.value.detail == "Agent name or Entra application is already registered"
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
index 526f24c5221..4bdb066377a 100644
--- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
+++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py
@@ -1,11 +1,13 @@
import json
+from datetime import datetime, timezone
+
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
-from fastapi import FastAPI
+from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from litellm.constants import REDACTED_BY_LITELM_STRING
@@ -21,7 +23,8 @@ from litellm.proxy.agent_endpoints.endpoints import (
router,
user_api_key_auth,
)
-from litellm.types.agents import AgentResponse
+from litellm.types.agents import AgentResponse, PatchAgentRequest
+from litellm.types.proxy.agent_identity import AgentIdentityBinding
def _sample_agent_card_params() -> dict:
@@ -97,7 +100,7 @@ def test_update_agent_success(mock_prisma_client, mock_user_api_key_auth, monkey
"agent_card_params": _sample_agent_card_params(),
}
mock_prisma_client.db.litellm_agentstable.find_unique = AsyncMock(
- return_value=existing_agent
+ return_value=AgentResponse.model_validate(existing_agent)
)
mock_registry = MagicMock()
@@ -350,6 +353,7 @@ class TestAgentByIdKeyRedaction:
test_client = _make_app_with_role(role)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
+ mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value=None
)
@@ -412,6 +416,7 @@ class TestAgentRBACInternalUser:
return_value=_sample_agent_response()
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
+ mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
return_value=None
)
@@ -592,6 +597,24 @@ class TestAgentRBACProxyAdmin:
)
assert resp.status_code == 200
+ def test_create_agent_rejects_legacy_litellm_params_identity(self):
+ with patch("litellm.proxy.proxy_server.prisma_client"): # test-quality-ok: proxy_server module global is the endpoint's only injection point
+ self.mock_registry.get_agent_by_name = MagicMock(return_value=None)
+ self.mock_registry.add_agent_to_db = AsyncMock(return_value=_sample_agent_response())
+ config = _sample_agent_config()
+ config["litellm_params"] = {
+ **config["litellm_params"],
+ "identity": {
+ "provider": "microsoft_entra",
+ "tenant_id": "11111111-1111-4111-8111-111111111111",
+ "client_id": "22222222-2222-4222-8222-222222222222",
+ },
+ }
+ resp = self.admin_client.post("/v1/agents", json=config, headers={"Authorization": "Bearer k"})
+ assert resp.status_code == 400, resp.text
+ assert "top-level identity field" in resp.json()["detail"]
+ self.mock_registry.add_agent_to_db.assert_not_awaited()
+
def test_create_agent_applies_litellm_merge_to_stored_card(self):
"""The card stored in the DB must reflect the LiteLLM-fronting merge."""
with patch("litellm.proxy.proxy_server.prisma_client"):
@@ -663,11 +686,9 @@ class TestAgentRBACProxyAdmin:
"""LIT-6736: PUT /v1/agents/{id} must not echo the stored secret back."""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Existing Agent",
- "agent_card_params": _sample_agent_card_params(),
- }
+ return_value=AgentResponse(
+ agent_id="agent-123", agent_name="Existing Agent", agent_card_params=_sample_agent_card_params()
+ )
)
self.mock_registry.update_agent_in_db = AsyncMock(
return_value=AgentResponse(
@@ -698,11 +719,9 @@ class TestAgentRBACProxyAdmin:
"""LIT-6736: PATCH /v1/agents/{id} must not echo the stored secret back."""
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
- return_value={
- "agent_id": "agent-123",
- "agent_name": "Existing Agent",
- "agent_card_params": _sample_agent_card_params(),
- }
+ return_value=AgentResponse(
+ agent_id="agent-123", agent_name="Existing Agent", agent_card_params=_sample_agent_card_params()
+ )
)
self.mock_registry.patch_agent_in_db = AsyncMock(
return_value=AgentResponse(
@@ -1140,6 +1159,143 @@ def test_make_agent_public_rejects_an_agent_published_only_in_the_db(monkeypatch
assert "already in public agent groups" in duplicate.json()["detail"]
+@pytest.mark.parametrize("enabled, claim_field, expected", [(True, "azp", True), (False, "azp", False), (True, None, False)])
+def test_jwt_authentication_status_does_not_require_virtual_keys(
+ monkeypatch: pytest.MonkeyPatch, enabled: bool, claim_field: str | None, expected: bool
+) -> None:
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_JWTAuth
+ from litellm.proxy.auth.handle_jwt import JWTHandler
+
+ handler: Final = JWTHandler()
+ handler.update_environment(None, DualCache(), LiteLLM_JWTAuth(agent_id_jwt_field=claim_field))
+ monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": enabled})
+ monkeypatch.setattr(proxy_server, "jwt_handler", handler)
+ agent: Final = _sample_agent_response()
+ response: Final = agent_endpoints._redact_sensitive_agent_fields((agent,), is_admin=True)[0]
+ assert response.jwt_auth_configured is expected
+ assert agent.jwt_auth_configured is False
+
+
+def test_identity_providers_require_configured_issuer_and_audience(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import LiteLLM_JWTAuth
+ from litellm.proxy.auth.handle_jwt import JWTHandler
+
+ handler: Final = JWTHandler()
+ handler.update_environment(None, DualCache(), LiteLLM_JWTAuth())
+ monkeypatch.setattr(proxy_server, "jwt_handler", handler)
+ monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": True})
+ monkeypatch.setenv("JWT_ISSUER", "https://issuer.example")
+ monkeypatch.delenv("JWT_AUDIENCE", raising=False)
+ assert client.get("/v1/agents/identity/providers").json() == []
+ monkeypatch.setenv("JWT_AUDIENCE", "gateway")
+ response: Final = client.get("/v1/agents/identity/providers")
+ assert response.status_code == 200
+ assert response.json() == ["https://issuer.example"]
+ forbidden: Final = _make_app_with_role(LitellmUserRoles.INTERNAL_USER).get("/v1/agents/identity/providers")
+ assert forbidden.status_code == 403
+
+
+def test_identity_evidence_is_persisted_and_never_taken_from_runtime_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
+ from litellm.proxy import proxy_server
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ binding: Final = AgentIdentityBinding(
+ agent_id="bound",
+ provider="microsoft_entra",
+ tenant_id="11111111-1111-4111-8111-111111111111",
+ client_id="22222222-2222-4222-8222-222222222222",
+ issuer="https://issuer.example",
+ revision="revision-one",
+ )
+ bound: Final = AgentResponse(
+ agent_id="bound",
+ agent_name="Readable name",
+ agent_card_params={},
+ identity=binding,
+ identity_managed=True,
+ litellm_params={"last_authenticated_at": "forged-proof"},
+ )
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=bound)
+ monkeypatch.setattr(proxy_server, "prisma_client", database)
+ pending: Final = client.get("/v1/agents/bound/identity")
+ assert pending.status_code == 200
+ assert pending.json()["last_authenticated_at"] is None
+ verified_binding: Final = binding.model_copy(
+ update={"last_authenticated_at": datetime(2026, 1, 1, tzinfo=timezone.utc)}
+ )
+ database.writer_db.litellm_agentstable.find_unique.return_value = bound.model_copy(update={"identity": verified_binding})
+ verified: Final = client.get("/v1/agents/bound/identity")
+ assert verified.json()["last_authenticated_at"] == "2026-01-01T00:00:00Z"
+ assert verified.json()["identity"]["client_id"] == binding.client_id
+ database.writer_db.litellm_agentstable.find_unique.return_value = None
+ assert client.get("/v1/agents/missing/identity").status_code == 404
+ database.writer_db.litellm_agentstable.find_unique.side_effect = RuntimeError("unavailable")
+ assert client.get("/v1/agents/bound/identity").status_code == 503
+
+
+@pytest.mark.parametrize("enabled", [True, False])
+def test_identity_providers_honor_issuer_specific_audiences_and_global_fallback(
+ monkeypatch: pytest.MonkeyPatch, enabled: bool
+) -> None:
+ from litellm.caching.dual_cache import DualCache
+ from litellm.proxy import proxy_server
+ from litellm.proxy._types import JWTIssuerConfig, LiteLLM_JWTAuth
+ from litellm.proxy.auth.handle_jwt import JWTHandler
+
+ handler: Final = JWTHandler()
+ handler.update_environment(
+ None,
+ DualCache(),
+ LiteLLM_JWTAuth(
+ issuers=[
+ JWTIssuerConfig(issuer="https://scoped.example", audience="gateway"),
+ JWTIssuerConfig(issuer="https://unscoped.example", disable_audience_validation=True),
+ ]
+ ),
+ )
+ monkeypatch.setattr(proxy_server, "jwt_handler", handler)
+ monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": enabled})
+ monkeypatch.setenv("JWT_ISSUER", "https://global.example")
+ monkeypatch.setenv("JWT_AUDIENCE", "gateway")
+ assert client.get("/v1/agents/identity/providers").json() == (
+ ["https://scoped.example", "https://global.example"] if enabled else []
+ )
+ monkeypatch.setenv("JWT_ISSUER", "https://unscoped.example")
+ assert client.get("/v1/agents/identity/providers").json() == (["https://scoped.example"] if enabled else [])
+
+
+@pytest.mark.parametrize("change", ({"execution_mode": "delegated"}, {"execution_mode": "both"}))
+def test_mode_only_edit_requires_the_existing_identity_sso_tenant(
+ monkeypatch: pytest.MonkeyPatch, change: PatchAgentRequest
+) -> None:
+ from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING, TENANT, managed_agent
+
+ monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,))
+ monkeypatch.delenv("MICROSOFT_TENANT", raising=False)
+ monkeypatch.setenv("MICROSOFT_CLIENT_ID", "gateway-client")
+ with pytest.raises(HTTPException, match="Delegated agents require Microsoft SSO"):
+ agent_endpoints._validate_managed_identity_request(change, managed_agent())
+ monkeypatch.setenv("MICROSOFT_TENANT", TENANT)
+ agent_endpoints._validate_managed_identity_request(change, managed_agent())
+
+
+def test_identity_only_edit_preserves_delegated_mode_validation(monkeypatch: pytest.MonkeyPatch) -> None:
+ from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING, managed_agent
+
+ monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,))
+ monkeypatch.delenv("MICROSOFT_TENANT", raising=False)
+ configuration: Final = BINDING.model_dump(
+ exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
+ )
+ delegated: Final = managed_agent().model_copy(update={"execution_mode": "delegated"})
+ with pytest.raises(HTTPException, match="Delegated agents require Microsoft SSO"):
+ agent_endpoints._validate_managed_identity_request({"identity": configuration}, delegated)
+
_KILL_SWITCH: Final = {
"url": "https://ops.example.com/kill",
"method": "POST",
@@ -1342,6 +1498,7 @@ def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_other
def _get_as(role: LitellmUserRoles):
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
+ mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
return _make_app_with_role(role).get("/v1/agents/agent-123", headers={"Authorization": "Bearer k"})
@@ -1357,3 +1514,80 @@ def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_other
assert internal.status_code == 200, internal.text
assert internal.json()["kill_switch"] is None
assert "tok-real" not in internal.text
+
+
+@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
+@pytest.mark.parametrize("path", ["/v1/agents", "/v1/agents/agent-123"])
+def test_agent_identity_configuration_is_only_returned_to_admins(role, path, monkeypatch):
+ from litellm.proxy.agent_endpoints import agent_registry
+
+ binding = AgentIdentityBinding(
+ agent_id="agent-123", provider="microsoft_entra", tenant_id="tenant", client_id="client",
+ issuer="https://login.microsoftonline.com/tenant/v2.0", revision="revision",
+ )
+ agent = _sample_agent_response().model_copy(update={"identity": binding})
+ registry = MagicMock()
+ registry.get_agent_by_id.return_value = agent
+ registry.get_agent_list.return_value = [agent]
+ registry.ids_for_agent.return_value = frozenset({agent.agent_id})
+ monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ monkeypatch.setattr(
+ "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.resolve_agent_access",
+ AsyncMock(return_value=RestrictedAgentAccess(frozenset({agent.agent_id}))),
+ )
+ with patch("litellm.proxy.proxy_server.prisma_client") as prisma:
+ prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
+ prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[])
+ prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
+ prisma.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
+ response = _make_app_with_role(role).get(path, headers={"Authorization": "Bearer k"})
+ assert response.status_code == 200
+ payload = response.json()[0] if path == "/v1/agents" else response.json()
+ assert payload["identity"] == (binding.model_dump(mode="json") if role == LitellmUserRoles.PROXY_ADMIN else None)
+ assert agent.identity == binding
+
+
+@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.INTERNAL_USER])
+def test_agent_detail_cache_miss_preserves_admin_identity_visibility(role, monkeypatch):
+ binding = AgentIdentityBinding(
+ agent_id="agent-123", provider="microsoft_entra", tenant_id="tenant", client_id="client",
+ issuer="https://login.microsoftonline.com/tenant/v2.0", revision="revision",
+ )
+ agent = _sample_agent_response()
+ registry = MagicMock()
+ registry.get_agent_by_id.return_value = None
+ registry.ids_for_agent.return_value = frozenset({agent.agent_id})
+ monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
+ monkeypatch.setattr(
+ "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
+ AsyncMock(return_value=True),
+ )
+
+ async def load_row(*, where, include):
+ assert where == {"agent_id": agent.agent_id}
+ return agent.model_copy(update={"identity": binding if include.get("identity") else None})
+
+ with patch("litellm.proxy.proxy_server.prisma_client") as prisma:
+ prisma.db.litellm_agentstable.find_unique = AsyncMock(side_effect=load_row)
+ prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
+ response = _make_app_with_role(role).get("/v1/agents/agent-123")
+ assert response.status_code == 200
+ assert response.json()["identity"] == (binding.model_dump(mode="json") if role == LitellmUserRoles.PROXY_ADMIN else None)
+
+
+@pytest.mark.parametrize("trusted", [False, True])
+def test_invalid_identity_and_untrusted_tenant_cannot_be_registered(
+ monkeypatch: pytest.MonkeyPatch, trusted: bool
+) -> None:
+ from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING
+
+ configuration: Final = BINDING.model_dump(
+ exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
+ )
+ monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,) if trusted else ())
+ request: Final = {"identity": {**configuration, "client_id": "invalid"} if trusted else configuration}
+ message: Final = "Invalid Entra identity configuration" if trusted else "Configure trusted JWT issuer"
+ with pytest.raises(HTTPException, match=message) as failure:
+ agent_endpoints._validate_managed_identity_request(request)
+ assert failure.value.status_code == 400
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_identity.py
new file mode 100644
index 00000000000..c9d803fdae7
--- /dev/null
+++ b/tests/test_litellm/proxy/agent_endpoints/test_identity.py
@@ -0,0 +1,27 @@
+from collections.abc import Mapping
+
+import pytest
+from fastapi import HTTPException
+
+from litellm.proxy.agent_endpoints.identity import has_legacy_identity, reject_legacy_identity
+
+TENANT = "11111111-1111-4111-8111-111111111111"
+CLIENT = "22222222-2222-4222-8222-222222222222"
+
+
+@pytest.mark.parametrize("params", [None, {}, {"model": "gpt-4o", "api_key": "sk-test"}])
+def test_runtime_params_without_identity_are_accepted(params: Mapping[str, object] | None) -> None:
+ assert has_legacy_identity(params) is False
+ reject_legacy_identity(params)
+
+
+@pytest.mark.parametrize(
+ "identity", [None, {}, {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}]
+)
+def test_legacy_litellm_params_identity_is_rejected(identity: object) -> None:
+ params: Mapping[str, object] = {"model": "gpt-4o", "identity": identity}
+ assert has_legacy_identity(params) is True
+ with pytest.raises(HTTPException) as failure:
+ reject_legacy_identity(params)
+ assert failure.value.status_code == 400
+ assert "top-level identity field" in failure.value.detail
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py
new file mode 100644
index 00000000000..faace1ccf20
--- /dev/null
+++ b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py
@@ -0,0 +1,402 @@
+from datetime import datetime, timezone
+from types import SimpleNamespace
+from typing import Final
+from unittest.mock import AsyncMock, MagicMock
+
+import pytest
+from fastapi import HTTPException
+from prisma.models import LiteLLM_VerifiedSubject
+
+from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
+from litellm.repositories.table_repositories import (
+ AgentIdentityRepository,
+ AgentsRepository,
+ VerifiedSubjectRepository,
+)
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import (
+ AgentIdentityBinding,
+ AgentIdentityFailure,
+ ManagedAgentContext,
+ MicrosoftInteractiveSubject,
+)
+
+TENANT: Final = "11111111-1111-4111-8111-111111111111"
+CLIENT: Final = "22222222-2222-4222-8222-222222222222"
+PRINCIPAL: Final = "33333333-3333-4333-8333-333333333333"
+HUMAN: Final = "44444444-4444-4444-8444-444444444444"
+ISSUER: Final = f"https://login.microsoftonline.com/{TENANT}/v2.0"
+BINDING: Final = AgentIdentityBinding(
+ agent_id="agent-one",
+ provider="microsoft_entra",
+ tenant_id=TENANT,
+ client_id=CLIENT,
+ service_principal_id=PRINCIPAL,
+ issuer=ISSUER,
+ required_roles=("Agent.Invoke",),
+ revision="revision-one",
+)
+CLAIMS: Final = {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": PRINCIPAL, "roles": ["Agent.Invoke"]}
+
+
+def stored_agent(**overrides: object) -> AgentResponse:
+ return AgentResponse.model_validate(
+ {
+ "agent_id": "agent-one",
+ "agent_name": "Research",
+ "agent_card_params": {},
+ "identity": BINDING,
+ "identity_managed": True,
+ "execution_mode": "both",
+ **overrides,
+ }
+ )
+
+
+def setup_store(
+ agent: AgentResponse | None = stored_agent(),
+ human: LiteLLM_VerifiedSubject | None = None,
+) -> tuple[AgentIdentityStore, AsyncMock, AsyncMock, AsyncMock]:
+ agents: Final = AsyncMock()
+ identities: Final = AsyncMock()
+ humans: Final = AsyncMock()
+ agents.find_unique.return_value = agent
+ identities.find_unique.return_value = BINDING
+ identities.update_many.return_value = 1
+ humans.find_unique.return_value = human
+ db: Final = SimpleNamespace(
+ db=SimpleNamespace(
+ litellm_agentstable=agents,
+ litellm_agentidentity=identities,
+ litellm_verifiedsubject=humans,
+ )
+ )
+ return (
+ AgentIdentityStore(AgentsRepository(db), AgentIdentityRepository(db), VerifiedSubjectRepository(db)),
+ agents,
+ identities,
+ humans,
+ )
+
+
+@pytest.mark.asyncio
+async def test_application_authentication_has_no_fabricated_human() -> None:
+ store, _, _, humans = setup_store()
+ result: Final = await store.resolve_verified_claims(CLAIMS)
+ assert isinstance(result, ManagedAgentContext)
+ assert result.agent_id == "agent-one"
+ assert result.mode == "autonomous"
+ assert result.user_id is None
+ humans.upsert.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_lifecycle_is_read_on_every_request_without_cached_allow() -> None:
+ store, agents, _, _ = setup_store()
+ agents.find_unique.side_effect = [stored_agent(), stored_agent(enabled=False)]
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext)
+ denial: Final = await store.resolve_verified_claims(CLAIMS)
+ assert isinstance(denial, AgentIdentityFailure)
+ assert denial.code == "identity_denied"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("unavailable_table", ["agents", "identities", "humans"])
+async def test_identity_store_failure_never_becomes_a_legacy_allow(unavailable_table: str) -> None:
+ store, agents, identities, humans = setup_store()
+ table: Final = {"agents": agents, "identities": identities, "humans": humans}[unavailable_table]
+ table.find_unique.side_effect = RuntimeError("database unavailable")
+ result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "policy_unavailable"
+
+
+@pytest.mark.asyncio
+async def test_unclassified_delegated_subject_cannot_authenticate_as_a_user() -> None:
+ store, _, _, _ = setup_store()
+ result: Final = await store.resolve_verified_claims(
+ {**CLAIMS, "oid": HUMAN, "scp": "user_impersonation", "idtyp": "user"}
+ )
+ assert isinstance(result, AgentIdentityFailure)
+ assert "first sign in" in result.message
+
+
+@pytest.mark.asyncio
+async def test_delegated_subject_uses_canonical_sso_user_not_email_claim() -> None:
+ human: Final = LiteLLM_VerifiedSubject(
+ kind="human",
+ subject_id="subject-one",
+ issuer=ISSUER,
+ tenant_id=TENANT,
+ oid=HUMAN,
+ user_id="canonical-user",
+ verified_via="sso_interactive",
+ verified_at=datetime.now(timezone.utc),
+ )
+ store, _, _, humans = setup_store(human=human)
+ result: Final = await store.resolve_verified_claims(
+ {
+ **CLAIMS,
+ "oid": HUMAN,
+ "scp": "user_impersonation",
+ "email": "untrusted-alias@example.com",
+ }
+ )
+ assert isinstance(result, ManagedAgentContext)
+ assert result.mode == "delegated"
+ assert result.user_id == "canonical-user"
+ humans.find_unique.assert_awaited_once_with(
+ where={"issuer_tenant_id_oid": {"issuer": ISSUER, "tenant_id": TENANT, "oid": HUMAN}}
+ )
+
+
+@pytest.mark.asyncio
+async def test_rebinding_during_authentication_does_not_mark_new_identity_verified() -> None:
+ store, _, identities, _ = setup_store()
+ identities.update_many.return_value = 0
+ context: Final = ManagedAgentContext(agent_id="agent-one", binding_revision="old-revision", mode="autonomous")
+ result: Final = await store.record_authentication(context)
+ assert isinstance(result, AgentIdentityFailure)
+ assert "changed" in result.message
+ assert identities.update_many.call_args.kwargs["where"] == {
+ "agent_id": "agent-one",
+ "revision": "old-revision",
+ "active": True,
+ "agent": {"is": {"enabled": True, "identity_managed": True}},
+ }
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("agent", [None, stored_agent(identity=None), stored_agent(identity_managed=False)])
+async def test_stale_binding_cannot_bypass_lifecycle(agent: AgentResponse | None) -> None:
+ store, _, _, _ = setup_store(agent=agent)
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure)
+
+
+@pytest.mark.asyncio
+async def test_unrelated_non_entra_claims_do_not_query_identity_store() -> None:
+ store, agents, identities, _ = setup_store()
+ assert await store.resolve_verified_claims({"sub": "ordinary-user"}) is None
+ identities.find_unique.assert_not_awaited()
+ agents.find_unique.assert_not_awaited()
+
+
+HUMAN_CLAIMS: Final = {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": HUMAN, "scp": "user_impersonation"}
+
+
+@pytest.mark.asyncio
+async def test_bound_agents_and_policy_failures_are_never_served_from_the_miss_cache() -> None:
+ store, _, identities, _ = setup_store()
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext)
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext)
+ assert identities.find_unique.await_count == 2
+ identities.find_unique.side_effect = ConnectionError("database down")
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure)
+ assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure)
+ assert identities.find_unique.await_count == 4
+
+
+@pytest.mark.asyncio
+async def test_retired_client_cannot_fall_back_to_ordinary_user_authentication() -> None:
+ from prisma.models import LiteLLM_RetiredAgentIdentity
+
+ from litellm.repositories.table_repositories import RetiredAgentIdentityRepository
+
+ identities: Final = AsyncMock()
+ identities.find_unique.return_value = None
+ retired: Final = AsyncMock()
+ retired.find_unique.return_value = LiteLLM_RetiredAgentIdentity(
+ binding_id="retired",
+ agent_id="agent-one",
+ provider="microsoft_entra",
+ issuer=ISSUER,
+ tenant_id=TENANT,
+ client_id=CLIENT,
+ )
+ db: Final = SimpleNamespace(
+ db=SimpleNamespace(
+ litellm_agentidentity=identities,
+ litellm_retiredagentidentity=retired,
+ litellm_agentstable=AsyncMock(),
+ litellm_verifiedsubject=AsyncMock(),
+ )
+ )
+ store: Final = AgentIdentityStore(
+ AgentsRepository(db),
+ AgentIdentityRepository(db),
+ VerifiedSubjectRepository(db),
+ RetiredAgentIdentityRepository(db),
+ )
+ result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"})
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "identity_denied"
+ assert "retired" in result.message
+
+
+@pytest.mark.asyncio
+async def test_missing_revision_cannot_create_entra_authentication_evidence() -> None:
+ store, _, identities, _ = setup_store()
+ result: Final = await store.record_authentication(ManagedAgentContext(agent_id="agent-one", mode="autonomous"))
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "identity_denied"
+ identities.update_many.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_authentication_evidence_write_failure_is_not_success() -> None:
+ store, _, identities, _ = setup_store()
+ identities.update_many.side_effect = RuntimeError("writer unavailable")
+ result: Final = await store.record_authentication(
+ ManagedAgentContext(agent_id="agent-one", binding_revision="revision-one", mode="autonomous")
+ )
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "policy_unavailable"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("unavailable", [True, False])
+async def test_retired_binding_denies_and_history_outage_cannot_become_legacy_fallback(unavailable: bool) -> None:
+
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=None)
+ database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None)
+ database.writer_db.litellm_retiredagentidentity.find_unique = AsyncMock(
+ return_value={"client_id": CLIENT}, side_effect=RuntimeError("unavailable") if unavailable else None
+ )
+ result: Final = await AgentIdentityStore.from_client(database).resolve_verified_claims(CLAIMS)
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == ("policy_unavailable" if unavailable else "identity_denied")
+ assert result.message == (
+ "Retired agent identity could not be checked" if unavailable else "This agent identity binding has been retired"
+ )
+
+
+@pytest.mark.asyncio
+async def test_new_binding_is_enforced_after_another_worker_commits_it() -> None:
+ _, agents, identities, humans = setup_store()
+ identities.find_unique.return_value = None
+ retired: Final = AsyncMock()
+ retired.find_unique.return_value = None
+ db: Final = SimpleNamespace(
+ writer_db=SimpleNamespace(
+ litellm_agentstable=agents,
+ litellm_agentidentity=identities,
+ litellm_verifiedsubject=humans,
+ litellm_retiredagentidentity=retired,
+ litellm_retiredagent=retired,
+ )
+ )
+ worker: Final = AgentIdentityStore.from_client(db)
+ claims: Final = {**CLAIMS, "oid": "55555555-5555-4555-8555-555555555555"}
+ assert await worker.resolve_verified_claims(claims) is None
+ identities.find_unique.return_value = BINDING
+ denied: Final = await worker.resolve_verified_claims(claims)
+ assert isinstance(denied, AgentIdentityFailure)
+ assert denied.code == "identity_denied"
+ assert "Application token contradicts" in denied.message
+
+
+@pytest.mark.asyncio
+async def test_non_string_subject_does_not_query_directory_ownership() -> None:
+ store, _, _, humans = setup_store()
+ assert await store.subject(ISSUER, TENANT, None) is None
+ humans.find_unique.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("configured", [False, True])
+async def test_missing_or_unavailable_retirement_history_fails_closed(configured: bool) -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("history unavailable"))
+ store: Final = AgentIdentityStore.from_client(database) if configured else setup_store()[0]
+ result: Final = await store.retired_agent("deleted-agent")
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "policy_unavailable"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("owner", ["canonical-user", "another-user"])
+async def test_interactive_enrollment_preserves_existing_subject_ownership(owner: str) -> None:
+ store, _, _, humans = setup_store()
+ humans.upsert.return_value = LiteLLM_VerifiedSubject(
+ subject_id="subject-one",
+ issuer=ISSUER,
+ tenant_id=TENANT,
+ oid=HUMAN,
+ user_id=owner,
+ kind="human",
+ verified_via="sso_interactive",
+ verified_at=datetime.now(timezone.utc),
+ )
+ result: Final = await store.enroll_interactive_human(
+ MicrosoftInteractiveSubject(issuer=ISSUER, tenant_id=TENANT, oid=HUMAN), "canonical-user"
+ )
+ if owner == "canonical-user":
+ assert result is None
+ else:
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "identity_denied"
+ assert humans.upsert.call_args.kwargs["data"]["update"] == {}
+ assert humans.upsert.call_args.kwargs["data"]["create"]["user_id"] == "canonical-user"
+
+
+@pytest.mark.asyncio
+async def test_interactive_enrollment_outage_fails_closed() -> None:
+ store, _, _, humans = setup_store()
+ humans.upsert.side_effect = ConnectionError("writer unavailable")
+ result: Final = await store.enroll_interactive_human(
+ MicrosoftInteractiveSubject(issuer=ISSUER, tenant_id=TENANT, oid=HUMAN), "canonical-user"
+ )
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "policy_unavailable"
+
+
+@pytest.mark.asyncio
+async def test_matching_revision_records_successful_authentication() -> None:
+ store, _, identities, _ = setup_store()
+ assert (
+ await store.record_authentication(
+ ManagedAgentContext(agent_id="agent-one", binding_revision="revision-one", mode="autonomous")
+ )
+ is None
+ )
+ identities.update_many.assert_awaited_once()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("outage", [False, True])
+async def test_resolver_maps_denials_and_outages_to_public_errors(outage: bool) -> None:
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentidentity.find_unique = AsyncMock(
+ return_value=BINDING, side_effect=ConnectionError("unavailable") if outage else None
+ )
+ database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None)
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=stored_agent(enabled=False))
+ with pytest.raises(HTTPException) as exc:
+ await resolve_managed_agent(CLAIMS, database)
+ assert exc.value.status_code == (503 if outage else 403)
+
+
+@pytest.mark.asyncio
+async def test_resolver_preserves_unconfigured_and_unrelated_authentication() -> None:
+ assert await resolve_managed_agent(CLAIMS, None) is None
+ assert await resolve_managed_agent({"sub": "ordinary-user"}, MagicMock()) is None
+ store, _, identities, _ = setup_store()
+ identities.find_unique.return_value = None
+ assert await store.resolve_verified_claims(CLAIMS) is None
+
+
+@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:
+ store, _, identities, humans = setup_store()
+ identities.find_unique.return_value = BINDING if registered else None
+ humans.find_unique.side_effect = RuntimeError("subject database unavailable")
+ result: Final = await store.resolve_verified_claims(CLAIMS)
+ if registered:
+ assert isinstance(result, ManagedAgentContext)
+ assert result.mode == "autonomous"
+ assert result.user_id is None
+ else:
+ assert result is None
+ humans.find_unique.assert_not_awaited()
diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
new file mode 100644
index 00000000000..45fe4b0655f
--- /dev/null
+++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py
@@ -0,0 +1,257 @@
+from typing import Final
+
+import pytest
+
+from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject, managed_write_fields
+from litellm.types.agents import AgentResponse
+from litellm.types.proxy.agent_identity import (
+ AgentExecutionMode,
+ AgentIdentityBinding,
+ AgentIdentityFailure,
+ AgentSubject,
+)
+
+TENANT: Final = "11111111-1111-4111-8111-111111111111"
+CLIENT: Final = "22222222-2222-4222-8222-222222222222"
+PRINCIPAL: Final = "33333333-3333-4333-8333-333333333333"
+HUMAN: Final = "44444444-4444-4444-8444-444444444444"
+ISSUER: Final = f"https://login.microsoftonline.com/{TENANT}/v2.0"
+BINDING: Final = AgentIdentityBinding(
+ agent_id="agent-one",
+ provider="microsoft_entra",
+ tenant_id=TENANT,
+ client_id=CLIENT,
+ service_principal_id=PRINCIPAL,
+ issuer=ISSUER,
+ required_roles=("Agent.Invoke",),
+ required_scopes=("user_impersonation",),
+ revision="binding-one",
+)
+
+
+def claims(**overrides: object) -> dict[str, object]:
+ return {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": PRINCIPAL, "roles": ["Agent.Invoke"], **overrides}
+
+
+def test_autonomous_identity_needs_no_human_and_checks_the_pinned_principal() -> None:
+ result: Final = classify_agent_subject(BINDING, claims(), "autonomous")
+ assert result == AgentSubject(kind="application", oid=PRINCIPAL, mode="autonomous")
+ assert isinstance(classify_agent_subject(BINDING, claims(oid=HUMAN), "autonomous"), AgentIdentityFailure)
+
+
+@pytest.mark.parametrize(
+ "overrides",
+ [
+ {"iss": "https://untrusted.example"},
+ {"tid": CLIENT},
+ {"azp": TENANT},
+ {"roles": []},
+ {"idtyp": "user"},
+ {"scp": "user_impersonation"},
+ {"scp": 1},
+ {"oid": None},
+ ],
+)
+def test_application_rejects_mismatched_or_contradictory_verified_claims(overrides: dict[str, object]) -> None:
+ assert isinstance(classify_agent_subject(BINDING, claims(**overrides), "both"), AgentIdentityFailure)
+
+
+def test_delegated_profile_identifies_a_subject_without_asserting_that_it_is_human() -> None:
+ result: Final = classify_agent_subject(BINDING, claims(oid=HUMAN, scp="user_impersonation"), "delegated")
+ assert result == AgentSubject(kind="delegated_subject", oid=HUMAN, mode="delegated")
+
+
+@pytest.mark.parametrize(
+ "overrides",
+ [
+ {"scp": "unrelated"},
+ {"scp": ""},
+ {"idtyp": "app"},
+ {"xms_sub_fct": "2 13 15"},
+ {"xms_sub_fct": [13]},
+ ],
+)
+def test_delegated_profile_rejects_unknown_scope_and_known_nonhuman_subjects(overrides: dict[str, object]) -> None:
+ assert isinstance(
+ classify_agent_subject(BINDING, claims(**{"oid": HUMAN, "scp": "user_impersonation", **overrides}), "both"),
+ AgentIdentityFailure,
+ )
+
+
+def test_allowed_mode_cannot_be_selected_by_the_caller() -> None:
+ assert isinstance(classify_agent_subject(BINDING, claims(), "delegated"), AgentIdentityFailure)
+ assert isinstance(
+ classify_agent_subject(BINDING, claims(oid=HUMAN, scp="user_impersonation"), "autonomous"),
+ AgentIdentityFailure,
+ )
+
+
+def test_native_facet_absence_does_not_establish_human_identity() -> None:
+ result: Final = classify_agent_subject(
+ BINDING, claims(oid=HUMAN, scp="user_impersonation", xms_sub_fct="113"), "both"
+ )
+ assert isinstance(result, AgentSubject)
+ assert result.kind == "delegated_subject"
+
+
+def managed_agent() -> AgentResponse:
+ return AgentResponse(
+ agent_id="agent-one", agent_name="Research", agent_card_params={}, identity=BINDING, identity_managed=True
+ )
+
+
+def test_unbinding_keeps_managed_state_and_disables_agent() -> None:
+ result: Final = managed_write_fields({"identity": None, "enabled": True}, managed_agent(), "admin")
+ assert not isinstance(result, AgentIdentityFailure)
+ assert result["identity_managed"] is True
+ assert result["enabled"] is False
+ assert result["identity"]["update"]["active"] is False
+ assert result["identity"]["update"]["last_authenticated_at"] is None
+ assert result["identity"]["update"]["revision"] != BINDING.revision
+
+
+def test_rename_does_not_rewrite_binding_or_evidence() -> None:
+ assert managed_write_fields({"agent_name": "Renamed"}, managed_agent(), "admin") == {}
+
+
+def test_autonomous_binding_requires_enterprise_application_object_id() -> None:
+ result: Final = managed_write_fields(
+ {"identity": {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}}, None, "admin"
+ )
+ assert isinstance(result, AgentIdentityFailure)
+ assert "service-principal" in result.message
+
+
+def test_rebinding_clears_evidence_and_uses_atomic_nested_write() -> None:
+ result: Final = managed_write_fields(
+ {
+ "identity": {
+ "provider": "microsoft_entra",
+ "tenant_id": TENANT,
+ "client_id": CLIENT,
+ "service_principal_id": PRINCIPAL,
+ }
+ },
+ managed_agent(),
+ "admin",
+ )
+ assert not isinstance(result, AgentIdentityFailure)
+ assert result["identity_managed"] is True
+ assert "upsert" in result["identity"]
+ assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision
+ assert result["identity"]["upsert"]["update"]["last_authenticated_at"] is None
+
+
+def test_unbound_identity_can_be_reactivated_with_the_same_application() -> None:
+ disabled: Final = managed_agent().model_copy(
+ update={"identity": BINDING.model_copy(update={"active": False}), "enabled": False}
+ )
+ configuration: Final = BINDING.model_dump(
+ exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
+ )
+ result: Final = managed_write_fields({"identity": configuration, "enabled": True}, disabled, "admin")
+ assert not isinstance(result, AgentIdentityFailure)
+ assert result["enabled"] is True
+ assert result["identity"]["upsert"]["update"]["active"] is True
+ assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision
+
+
+def test_each_application_binding_records_its_history_atomically() -> None:
+ configuration: Final = BINDING.model_dump(
+ exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
+ )
+ created: Final = managed_write_fields({"identity": configuration}, None, "admin")
+ assert not isinstance(created, AgentIdentityFailure)
+ assert created["retired_identities"]["connectOrCreate"]["create"]["client_id"] == CLIENT
+ replacement: Final = managed_write_fields(
+ {"identity": {**configuration, "client_id": HUMAN}}, managed_agent(), "admin"
+ )
+ assert not isinstance(replacement, AgentIdentityFailure)
+ assert replacement["retired_identities"]["connectOrCreate"]["create"]["client_id"] == HUMAN
+
+
+def test_unchanged_binding_preserves_revision_and_authentication_evidence() -> None:
+ configuration: Final = BINDING.model_dump(
+ exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"}
+ )
+ assert managed_write_fields({"identity": configuration}, managed_agent(), "admin") == {}
+
+
+@pytest.mark.parametrize("identity", [None, BINDING.model_copy(update={"active": False})])
+def test_enabling_unbound_or_inactive_identity_requires_rebinding(identity: AgentIdentityBinding | None) -> None:
+ agent: Final = managed_agent().model_copy(update={"identity": identity, "enabled": False})
+ result: Final = managed_write_fields({"enabled": True}, agent, "admin")
+ assert isinstance(result, AgentIdentityFailure)
+ assert "Bind an identity" in result.message
+
+
+@pytest.mark.parametrize("mode", ["delegated", "both"])
+def test_explicit_empty_scope_requirements_can_be_registered_and_preserved(mode: str) -> None:
+ from litellm.types.proxy.agent_identity import EntraIdentityConfig
+
+ configuration: Final = EntraIdentityConfig(
+ provider="microsoft_entra",
+ tenant_id=TENANT,
+ client_id=CLIENT,
+ service_principal_id=PRINCIPAL,
+ required_scopes=(),
+ )
+ created: Final = managed_write_fields(
+ {"identity": configuration.model_dump(), "execution_mode": mode}, None, "admin"
+ )
+ assert not isinstance(created, AgentIdentityFailure)
+ assert created["identity"]["create"]["required_scopes"] == ()
+ agent: Final = managed_agent().model_copy(update={"identity": BINDING.model_copy(update={"required_scopes": ()})})
+ updated: Final = managed_write_fields({"execution_mode": mode}, agent, "admin")
+ assert not isinstance(updated, AgentIdentityFailure)
+ assert updated["execution_mode"] == mode
+
+
+@pytest.mark.parametrize(
+ "incoming",
+ [
+ {"identity": {"provider": "microsoft_entra", "tenant_id": "invalid", "client_id": CLIENT}},
+ {"execution_mode": "unknown"},
+ ],
+)
+def test_invalid_identity_configuration_returns_a_public_validation_failure(incoming: dict[str, object]) -> None:
+ result: Final = managed_write_fields(incoming, None, "admin")
+ assert isinstance(result, AgentIdentityFailure)
+ assert result.code == "identity_denied"
+ assert result.message.startswith("Invalid agent identity configuration:")
+
+
+@pytest.mark.parametrize("roles", ["Agent.Invoke", [42], None])
+def test_malformed_application_roles_are_rejected(roles: object) -> None:
+ result: Final = classify_agent_subject(BINDING, claims(roles=roles), "autonomous")
+ assert isinstance(result, AgentIdentityFailure)
+ assert "Invalid application roles" in result.message
+
+
+def test_entra_binding_normalizes_identifiers_and_rejects_invalid_configuration() -> None:
+ from pydantic import ValidationError
+
+ from litellm.types.proxy.agent_identity import EntraIdentityConfig
+
+ identifier = "ABCDEF00-1234-4234-9234-123456789ABC"
+ config = EntraIdentityConfig(provider="microsoft_entra", tenant_id=identifier, client_id=identifier)
+ assert config.tenant_id == identifier.lower()
+ assert config.client_id == identifier.lower()
+ assert config.service_principal_id is None
+ assert config.issuer == f"https://login.microsoftonline.com/{config.tenant_id}/v2.0"
+ with pytest.raises(ValidationError):
+ EntraIdentityConfig(provider="microsoft_entra", tenant_id="invalid", client_id=identifier)
+
+
+@pytest.mark.parametrize("mode", ["delegated", "both"])
+def test_empty_required_scopes_allow_valid_delegated_scope(mode: AgentExecutionMode) -> None:
+ binding: Final = BINDING.model_copy(update={"required_scopes": ()})
+ result: Final = classify_agent_subject(binding, claims(oid=HUMAN, scp="custom_scope"), mode)
+ assert result == AgentSubject(kind="delegated_subject", oid=HUMAN, mode="delegated")
+
+
+@pytest.mark.parametrize("scope", [None, "", " \t ", 42])
+def test_empty_requirements_do_not_make_a_scope_less_human_token_valid(scope: object) -> None:
+ binding: Final = BINDING.model_copy(update={"required_scopes": ()})
+ result: Final = classify_agent_subject(binding, claims(oid=HUMAN, scp=scope), "both")
+ assert isinstance(result, AgentIdentityFailure)
diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py
index f014e9c26d1..88833f38270 100644
--- a/tests/test_litellm/proxy/auth/test_auth_checks.py
+++ b/tests/test_litellm/proxy/auth/test_auth_checks.py
@@ -1091,7 +1091,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
monkeypatch.setitem(auth_checks.last_db_access_time, f"user_id:{user_id}", (None, time.time()))
db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user")
mock_prisma_client = MagicMock()
- mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
+ mock_prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
result = await get_user_object(
user_id=user_id,
@@ -1103,7 +1103,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
assert result is not None
assert result.user_id == user_id
- mock_prisma_client.db.litellm_usertable.find_unique.assert_awaited_once()
+ mock_prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_once()
@pytest.mark.asyncio
@@ -3058,7 +3058,7 @@ async def test_get_team_object_raises_404_when_not_found():
mock_prisma_client = MagicMock()
mock_db = AsyncMock()
mock_prisma_client.db = mock_db
- mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
+ mock_prisma_client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
@@ -3076,11 +3076,40 @@ async def test_get_team_object_raises_404_when_not_found():
assert "Team doesn't exist in db" in str(exc_info.value.detail)
+@pytest.mark.asyncio
+async def test_get_team_object_check_db_only_reads_writer_through_the_shared_loader():
+ """Management endpoints mock ``_get_team_object_from_user_api_key_cache`` and expect
+ ``check_db_only`` to still flow through it; only the table it reads moves to the writer."""
+ from unittest.mock import AsyncMock, MagicMock
+
+ from litellm.proxy.auth import auth_checks
+ from litellm.proxy.auth.auth_checks import get_team_object
+
+ row = {"team_id": "team-writer", "models": ["gpt-4o"], "object_permission_id": None}
+ prisma = MagicMock()
+ prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
+ prisma.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
+ cache = MagicMock()
+ cache.async_get_cache = AsyncMock(return_value=None)
+ cache.async_set_cache = AsyncMock()
+ shared_loader = AsyncMock(wraps=auth_checks._get_team_object_from_user_api_key_cache)
+
+ with patch.object(auth_checks, "_get_team_object_from_user_api_key_cache", shared_loader):
+ team = await get_team_object("team-writer", prisma, cache, check_db_only=True)
+
+ assert team.team_id == "team-writer"
+ assert shared_loader.await_args.kwargs["use_writer"] is True
+ prisma.writer_db.litellm_teamtable.find_unique.assert_awaited_once()
+ prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
+ cache.async_set_cache.assert_awaited_once()
+
+
def _mock_prisma_for_team_lookup(find_unique):
from unittest.mock import MagicMock
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teamtable.find_unique = find_unique
+ mock_prisma_client.writer_db.litellm_teamtable.find_unique = find_unique
return mock_prisma_client
@@ -9979,3 +10008,130 @@ def test_can_object_call_model_allows_listed_model_for_key():
)
assert result is True
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("allowed", [True, False])
+async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache(allowed: bool) -> None:
+ from litellm.proxy._types import LiteLLM_AccessGroupTable
+ from litellm.proxy.auth.auth_checks import get_access_object
+
+ stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"])
+ current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []})
+ client: Final = MagicMock()
+ client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current)
+ client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=stale)
+ cache: Final = MagicMock()
+ cache.async_get_cache = AsyncMock(return_value=stale)
+ cache.async_set_cache = AsyncMock()
+ result: Final = await get_access_object("group", client, cache, check_db_only=True)
+ assert result.access_model_names == (["new"] if allowed else [])
+ cache.async_get_cache.assert_not_awaited()
+ client.db.litellm_accessgrouptable.find_unique.assert_not_awaited()
+ client.writer_db.litellm_accessgrouptable.find_unique.assert_awaited_once_with(where={"access_group_id": "group"})
+
+
+@pytest.mark.asyncio
+async def test_authoritative_access_group_outage_does_not_use_cached_grants() -> None:
+ from fastapi import HTTPException
+
+ from litellm.proxy.auth.auth_checks import get_access_object
+
+ client: Final = MagicMock()
+ client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable"))
+ cache: Final = MagicMock()
+ cache.async_get_cache = AsyncMock()
+ with pytest.raises(HTTPException) as failure:
+ await get_access_object("group", client, cache, check_db_only=True)
+ assert failure.value.status_code == 404
+ cache.async_get_cache.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_authoritative_team_permission_outage_cannot_drop_the_teams_restrictions() -> None:
+ from fastapi import HTTPException
+
+ from litellm.proxy.auth.auth_checks import get_team_object
+
+ row: Final = LiteLLM_TeamTable(team_id="team-policy-outage", object_permission_id="team-permission")
+ client: Final = MagicMock()
+ client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=row)
+ client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(side_effect=RuntimeError("unavailable"))
+ cache: Final = MagicMock()
+ cache.async_get_cache = AsyncMock()
+ cache.async_set_cache = AsyncMock()
+ with pytest.raises(HTTPException) as failure:
+ await get_team_object(row.team_id, client, cache, check_db_only=True)
+ assert failure.value.status_code == 404
+ client.writer_db.litellm_objectpermissiontable.find_unique.assert_awaited_once()
+ cache.async_set_cache.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("strict", [True, False])
+@pytest.mark.parametrize("missing", [True, False])
+async def test_referenced_permission_failures_preserve_legacy_behavior_and_deny_strict_reads(strict, missing):
+ from fastapi import HTTPException
+
+ from litellm.proxy.auth.auth_checks import get_object_permission
+
+ client = MagicMock()
+ lookup = AsyncMock(return_value=None, side_effect=None if missing else RuntimeError("unavailable"))
+ client.writer_db.litellm_objectpermissiontable.find_unique = lookup
+ client.db.litellm_objectpermissiontable.find_unique = lookup
+ cache = MagicMock()
+ cache.async_get_cache = AsyncMock(return_value=None)
+ if strict:
+ with pytest.raises(HTTPException if missing else RuntimeError):
+ await get_object_permission("referenced", client, cache, check_db_only=True)
+ cache.async_get_cache.assert_not_awaited()
+ else:
+ assert await get_object_permission("referenced", client, cache) is None
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "models,key_aliases,team_aliases,allowed",
+ [
+ (["fast"], {}, {}, True),
+ ([], {}, {}, False),
+ (["other"], {}, {}, False),
+ (["target"], {"fast": "target"}, {}, True),
+ (["target"], {}, {"fast": "target"}, True),
+ (["fast"], {}, {"fast": "forbidden"}, False),
+ ],
+)
+async def test_managed_agent_model_policy_checks_dispatched_model(
+ models: list[str], key_aliases: dict[str, str], team_aliases: dict[str, str], allowed: bool
+) -> None:
+ from fastapi import HTTPException
+
+ from litellm.proxy.auth.auth_checks import common_checks
+ from litellm.types.agents import AgentResponse
+
+ agent: Final = AgentResponse(
+ agent_id="managed", agent_name="Managed", agent_card_params={}, object_permission={"models": models}
+ )
+ auth: Final = UserAPIKeyAuth(
+ token="test-token", team_id="team", aliases=key_aliases, team_model_aliases=team_aliases
+ )
+ auth.managed_agent_policy = agent
+ checks: Final = common_checks(
+ request_body={"model": "fast", "messages": [{"role": "user", "content": "hi"}]},
+ team_object=None,
+ user_object=None,
+ end_user_object=None,
+ global_proxy_spend=None,
+ general_settings={},
+ route="/chat/completions",
+ llm_router=None,
+ proxy_logging_obj=MagicMock(),
+ valid_token=auth,
+ request=MagicMock(spec=Request),
+ )
+ if allowed:
+ assert await checks is True
+ else:
+ with pytest.raises((HTTPException, ModelAccessDeniedProxyException)) as failure:
+ await checks
+ assert str(getattr(failure.value, "status_code", getattr(failure.value, "code", None))) == "403"
diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py
index b1622e0dff0..0280a570292 100644
--- a/tests/test_litellm/proxy/auth/test_handle_jwt.py
+++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py
@@ -2,15 +2,14 @@ import asyncio
import re
import time
from collections.abc import Mapping, Sequence
-from typing import Final, Optional
+from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
-from fastapi import HTTPException
import httpx
import pytest
+from fastapi import HTTPException
-import litellm
-
+from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import (
DEFAULT_JWKS_STALE_TTL,
JWTLiteLLMRoleMap,
@@ -26,7 +25,6 @@ from litellm.proxy._types import (
RoleBasedPermissions,
ScopeMapping,
)
-from litellm.caching.dual_cache import DualCache
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.auth.auth_checks import TeamNotFoundError
from litellm.proxy.auth.handle_jwt import (
@@ -1637,7 +1635,6 @@ async def test_auth_builder_returns_team_membership_object():
@pytest.mark.asyncio
async def test_auth_builder_with_oidc_userinfo_enabled():
"""Test that auth_builder uses OIDC UserInfo endpoint when enabled"""
- from unittest.mock import MagicMock
from litellm.caching import DualCache
from litellm.proxy.utils import ProxyLogging
@@ -1648,9 +1645,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
general_settings = {"enforce_rbac": False}
route = "/chat/completions"
- user_object = LiteLLM_UserTable(
- user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
- )
+ user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
# Create JWT handler with OIDC UserInfo enabled
jwt_handler = JWTHandler()
@@ -1677,18 +1672,12 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
# Mock all the dependencies
with (
- patch.object(
- jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
- ) as mock_get_userinfo,
+ patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
- patch.object(
- JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
- ) as mock_check_rbac,
+ patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
- patch.object(
- jwt_handler, "get_object_id", return_value=None
- ) as mock_get_object_id,
+ patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
patch.object(
JWTAuthManager,
"get_user_info",
@@ -1696,9 +1685,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
return_value=("test_user_1", "test@example.com", True),
) as mock_get_user_info,
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
- patch.object(
- jwt_handler, "get_end_user_id", return_value=None
- ) as mock_get_end_user_id,
+ patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
patch.object(
JWTAuthManager,
"check_admin_access",
@@ -1711,9 +1698,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
new_callable=AsyncMock,
return_value=(None, None),
) as mock_find_team,
- patch.object(
- JWTAuthManager, "get_all_team_ids", return_value=set()
- ) as mock_get_all_team_ids,
+ patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
patch.object(
JWTAuthManager,
"find_team_with_model_access",
@@ -1726,15 +1711,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
new_callable=AsyncMock,
return_value=(user_object, None, None, None, user_object.user_id),
) as mock_get_objects,
- patch.object(
- JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
- ) as mock_map_user,
- patch.object(
- JWTAuthManager, "validate_object_id", return_value=True
- ) as mock_validate_object,
- patch.object(
- JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
- ) as mock_sync_user,
+ patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
+ patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
+ patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
):
# Set up mock return values
mock_get_userinfo.return_value = userinfo_response
@@ -1764,7 +1743,6 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
@pytest.mark.asyncio
async def test_auth_builder_with_oidc_userinfo_disabled():
"""Test that auth_builder uses JWT validation when OIDC UserInfo is disabled"""
- from unittest.mock import MagicMock
from litellm.caching import DualCache
from litellm.proxy.utils import ProxyLogging
@@ -1775,9 +1753,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
general_settings = {"enforce_rbac": False}
route = "/chat/completions"
- user_object = LiteLLM_UserTable(
- user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
- )
+ user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
# Create JWT handler with OIDC UserInfo disabled
jwt_handler = JWTHandler()
@@ -1801,18 +1777,12 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
# Mock all the dependencies
with (
- patch.object(
- jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
- ) as mock_get_userinfo,
+ patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
- patch.object(
- JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
- ) as mock_check_rbac,
+ patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
- patch.object(
- jwt_handler, "get_object_id", return_value=None
- ) as mock_get_object_id,
+ patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
patch.object(
JWTAuthManager,
"get_user_info",
@@ -1820,9 +1790,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
return_value=("test_user_1", None, None),
) as mock_get_user_info,
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
- patch.object(
- jwt_handler, "get_end_user_id", return_value=None
- ) as mock_get_end_user_id,
+ patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
patch.object(
JWTAuthManager,
"check_admin_access",
@@ -1835,9 +1803,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
new_callable=AsyncMock,
return_value=(None, None),
) as mock_find_team,
- patch.object(
- JWTAuthManager, "get_all_team_ids", return_value=set()
- ) as mock_get_all_team_ids,
+ patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
patch.object(
JWTAuthManager,
"find_team_with_model_access",
@@ -1850,15 +1816,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
new_callable=AsyncMock,
return_value=(user_object, None, None, None, user_object.user_id),
) as mock_get_objects,
- patch.object(
- JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
- ) as mock_map_user,
- patch.object(
- JWTAuthManager, "validate_object_id", return_value=True
- ) as mock_validate_object,
- patch.object(
- JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
- ) as mock_sync_user,
+ patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
+ patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
+ patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
):
# Set up mock return values
mock_auth_jwt.return_value = jwt_response
@@ -2631,7 +2591,6 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
"""
Test that find_and_validate_specific_team_id resolves team by name when team_id is not found
"""
- from unittest.mock import MagicMock
from litellm.caching import DualCache
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
@@ -2654,9 +2613,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
# Mock team object returned by get_team_object_by_alias
team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team")
- with patch(
- "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock
- ) as mock_get_by_alias:
+ with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias:
mock_get_by_alias.return_value = team_object
team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id(
@@ -2685,7 +2642,6 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
"""
Test that team_id_jwt_field takes precedence over team_alias_jwt_field
"""
- from unittest.mock import MagicMock
from litellm.caching import DualCache
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
@@ -2699,9 +2655,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=user_api_key_cache,
- litellm_jwtauth=LiteLLM_JWTAuth(
- team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"
- ),
+ litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"),
)
# Token with both team_id and team name
@@ -2711,9 +2665,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
team_object = LiteLLM_TeamTable(team_id="direct-team-id")
with (
- patch(
- "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock
- ) as mock_get_by_id,
+ patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id,
patch(
"litellm.proxy.auth.handle_jwt.get_team_object_by_alias",
new_callable=AsyncMock,
@@ -2890,7 +2842,6 @@ async def test_get_objects_resolves_org_by_name():
@pytest.mark.asyncio
async def test_resolve_jwks_url_passthrough_for_direct_jwks_url():
"""Non-discovery URLs are returned unchanged."""
- from unittest.mock import AsyncMock, MagicMock
from litellm.caching.dual_cache import DualCache
@@ -3143,7 +3094,7 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
When team_id_jwt_field is a normal field name (no dot-notation) the
error message should not contain a spurious bracket-notation hint.
"""
- from unittest.mock import AsyncMock, MagicMock
+ from unittest.mock import MagicMock
from litellm.caching.dual_cache import DualCache
@@ -3230,8 +3181,8 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
user_id: str,
user_teams: list,
- get_team_object_return: Optional[str],
- expected_team_id: Optional[str],
+ get_team_object_return: str | None,
+ expected_team_id: str | None,
expect_get_team_called: bool,
expect_get_membership_called: bool,
) -> None:
@@ -3244,9 +3195,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
if len(user_teams) == 1 and get_team_object_return == "resolved_row":
only = user_teams[0]
team_table = LiteLLM_TeamTable(team_id=only)
- membership = LiteLLM_TeamMembership(
- user_id=user_id, team_id=only, litellm_budget_table=None
- )
+ membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None)
get_team_return_value = team_table
membership_return_value = membership
else:
@@ -3305,9 +3254,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
),
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
- patch.object(
- JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
- ),
+ patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
patch(
"litellm.proxy.auth.handle_jwt.get_team_object",
new_callable=AsyncMock,
@@ -3324,9 +3271,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
code = 404 if get_team_object_return == "http_404" else 500
mock_get_team.side_effect = HTTPException(
status_code=code,
- detail={
- "error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."
- },
+ detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."},
)
else:
mock_get_team.return_value = get_team_return_value
@@ -4047,7 +3992,7 @@ def _encode_rsa_jwt(
issuer: str,
audience: str,
kid: str,
- extra_claims: Optional[dict] = None,
+ extra_claims: dict | None = None,
) -> str:
import time
@@ -4743,12 +4688,9 @@ async def test_get_objects_team_membership_uses_rebound_user_id():
async def fake_get_team_membership(user_id, team_id, *args, **kwargs):
captured["user_id"] = user_id
captured["team_id"] = team_id
- return None
jwt_handler = JWTHandler()
- jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
- user_id_jwt_field="email", user_id_upsert=True
- )
+ jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True)
with (
patch(
@@ -5389,7 +5331,7 @@ async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback(
assert team_object is None
-def _db_fallback_handler(litellm_jwtauth: Optional[LiteLLM_JWTAuth] = None) -> JWTHandler:
+def _db_fallback_handler(litellm_jwtauth: LiteLLM_JWTAuth | None = None) -> JWTHandler:
handler = JWTHandler()
handler.litellm_jwtauth = litellm_jwtauth or LiteLLM_JWTAuth()
return handler
@@ -5447,9 +5389,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
"expect_403",
),
[
- pytest.param(
- True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"
- ),
+ pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"),
pytest.param(
True,
["team_a", "team_b"],
@@ -5497,8 +5437,8 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
fallback_to_db_teams: bool,
user_teams: list,
- header_team_id: Optional[str],
- expected_team_id: Optional[str],
+ header_team_id: str | None,
+ expected_team_id: str | None,
expect_403: bool,
) -> None:
"""End-to-end auth_builder behavior with no JWT team claims.
@@ -5527,9 +5467,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
async def call_auth_builder():
with (
- patch.object(
- jwt_handler, "auth_jwt", new_callable=AsyncMock
- ) as mock_auth_jwt,
+ patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock),
patch.object(jwt_handler, "get_rbac_role", return_value=None),
patch.object(jwt_handler, "get_scopes", return_value=[]),
@@ -5569,9 +5507,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
),
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
- patch.object(
- JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
- ),
+ patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
patch(
"litellm.proxy.auth.handle_jwt.get_team_object",
new_callable=AsyncMock,
@@ -6765,7 +6701,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
team_id_upsert=True,
)
- upsert_by_team: dict[str, Optional[bool]] = {}
+ upsert_by_team: dict[str, bool | None] = {}
async def spy_get_team(team_id, **kwargs):
upsert_by_team[team_id] = kwargs.get("team_id_upsert")
@@ -6800,9 +6736,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
),
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
- patch.object(
- JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
- ),
+ patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
patch(
"litellm.proxy.auth.handle_jwt.get_team_object",
new_callable=AsyncMock,
@@ -7806,6 +7740,58 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc
assert result["team_id"] is None
+def _explicit_identity_registry() -> AgentRegistry:
+ registry: Final = AgentRegistry()
+ registry.register_agent(AgentResponse(
+ agent_id="explicit-agent-id",
+ agent_name="Readable agent name",
+ agent_card_params={},
+ litellm_params={"identity": {
+ "provider": "microsoft_entra",
+ "tenant_id": "11111111-1111-4111-8111-111111111111",
+ "client_id": "22222222-2222-4222-8222-222222222222",
+ }},
+ ))
+ return registry
+
+
+@pytest.mark.parametrize("claim_field", ["azp", None])
+def test_runtime_json_cannot_establish_a_managed_identity(claim_field: str | None) -> None:
+ registry: Final = _explicit_identity_registry()
+ handler: Final = _entra_agent_jwt_handler(claim_field)
+ claims: Final = {
+ "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
+ "tid": "11111111-1111-4111-8111-111111111111",
+ "azp": "22222222-2222-4222-8222-222222222222",
+ }
+ if claim_field is None:
+ assert JWTAuthManager.resolve_agent_id(handler, claims, registry) is None
+ else:
+ with pytest.raises(HTTPException) as failure:
+ JWTAuthManager.resolve_agent_id(handler, claims, registry)
+ assert failure.value.status_code == 403
+
+
+@pytest.mark.parametrize("override", [
+ {"iss": "https://attacker.example"},
+ {"tid": "33333333-3333-4333-8333-333333333333"},
+ {"azp": "33333333-3333-4333-8333-333333333333"},
+ {"azp": "explicit-agent-id"},
+ {"azp": "Readable agent name"},
+])
+def test_explicit_entra_identity_cannot_be_claimed_via_legacy_lookup(override: Mapping[str, object]) -> None:
+ registry: Final = _explicit_identity_registry()
+ handler: Final = _entra_agent_jwt_handler("azp")
+ with pytest.raises(HTTPException) as failure:
+ JWTAuthManager.resolve_agent_id(handler, {
+ "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
+ "tid": "11111111-1111-4111-8111-111111111111",
+ "azp": "22222222-2222-4222-8222-222222222222",
+ **override,
+ }, registry)
+ assert failure.value.status_code == 403
+
+
@pytest.mark.asyncio
@pytest.mark.parametrize("existing_user", [False, True])
@pytest.mark.parametrize("warm_cache", [False, True])
@@ -7853,3 +7839,216 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning
users.create.assert_not_awaited()
if existing_user:
assert users.find_unique.await_count == (0 if warm_cache else 1)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
+@pytest.mark.parametrize("audience_validation", (True, False))
+@pytest.mark.parametrize(
+ "route,allowed",
+ [
+ ("/chat/completions", True), ("/v1/messages", True), ("/v1/responses", True),
+ ("/mcp-rest/tools/call", True), ("/a2a/target", True),
+ ("/v1/files", False), ("/v1/batches", False), ("/v1/vector_stores", False),
+ ("/v1/containers", False), ("/openai/v1/files", False),
+ ("/v1/responses/other-response", False), ("/v1/realtime/client_secrets", False),
+ ],
+)
+async def test_managed_application_uses_persisted_identity_without_provisioning_human(
+ monkeypatch: pytest.MonkeyPatch, mode: str, audience_validation: bool, route: str, allowed: bool
+) -> None:
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ tenant: Final = "11111111-1111-4111-8111-111111111111"
+ client_id: Final = "22222222-2222-4222-8222-222222222222"
+ principal: Final = "33333333-3333-4333-8333-333333333333"
+ issuer: Final = f"https://login.microsoftonline.com/{tenant}/v2.0"
+ jwks_url: Final = "https://login.microsoftonline.test/managed-keys"
+ monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
+ monkeypatch.setenv("JWT_ISSUER", issuer)
+ monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
+ private_key, jwk = _get_rsa_key_and_jwk(kid="managed-key")
+ cache: Final = DualCache()
+ cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
+ handler: Final = JWTHandler()
+ handler.update_environment(None, cache, LiteLLM_JWTAuth(user_id_upsert=True))
+ binding: Final = AgentIdentityBinding(
+ agent_id="stable-id",
+ provider="microsoft_entra",
+ issuer=issuer,
+ tenant_id=tenant,
+ client_id=client_id,
+ service_principal_id=principal,
+ revision="revision-one",
+ required_roles=("Agent.Invoke",),
+ )
+ agent: Final = AgentResponse.model_validate(
+ {
+ "agent_id": "stable-id",
+ "agent_name": "A readable name",
+ "agent_card_params": {},
+ "identity": binding,
+ "identity_managed": True,
+ "execution_mode": mode,
+ }
+ )
+ database: Final = MagicMock()
+ database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
+ database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1)
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
+ database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None)
+ database.db.litellm_usertable.upsert = AsyncMock()
+ token: Final = _encode_rsa_jwt(
+ private_key,
+ issuer=issuer,
+ audience="api://gateway",
+ kid="managed-key",
+ extra_claims={
+ "tid": tenant,
+ "azp": client_id,
+ "oid": principal,
+ "roles": ["Agent.Invoke"],
+ "idtyp": "app",
+ },
+ )
+ arguments: Final = dict(
+ api_key=token,
+ jwt_handler=handler,
+ request_data={},
+ general_settings={},
+ route=route,
+ prisma_client=database,
+ user_api_key_cache=cache,
+ parent_otel_span=None,
+ proxy_logging_obj=MagicMock(),
+ )
+ if not audience_validation:
+ monkeypatch.delenv("JWT_AUDIENCE")
+ if mode == "delegated" or not audience_validation or not allowed:
+ with pytest.raises(HTTPException) as failure:
+ await JWTAuthManager.auth_builder(**arguments)
+ assert failure.value.status_code == 403
+ else:
+ result: Final = await JWTAuthManager.auth_builder(**arguments)
+ auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
+ assert auth.agent_id == "stable-id"
+ assert auth.user_id is None
+ assert auth.team_id is None
+ assert auth.managed_agent_context is not None
+ assert auth.managed_agent_context.mode == "autonomous"
+ assert result["is_proxy_admin"] is False
+ database.db.litellm_usertable.upsert.assert_not_awaited()
+
+
+@pytest.mark.parametrize("claim_value", ["managed", "Readable managed agent"])
+def test_legacy_claim_cannot_select_a_top_level_entra_binding(claim_value: str) -> None:
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ registry: Final = AgentRegistry()
+ registry.register_agent(
+ AgentResponse(
+ agent_id="managed",
+ agent_name="Readable managed agent",
+ agent_card_params={},
+ identity_managed=True,
+ identity=AgentIdentityBinding(
+ agent_id="managed",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="client",
+ service_principal_id="principal",
+ issuer="issuer",
+ revision="revision",
+ ),
+ )
+ )
+ with pytest.raises(HTTPException) as denied:
+ JWTAuthManager.resolve_agent_id(_entra_agent_jwt_handler("agent"), {"agent": claim_value}, registry)
+ assert denied.value.status_code == 403
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("kind", ["human", "config-agent", "managed-agent"])
+async def test_database_free_jwt_admission_with_entra_shaped_claims(monkeypatch: pytest.MonkeyPatch, kind: str) -> None:
+ issuer: Final = "https://login.microsoftonline.com/test-tenant/v2.0"
+ jwks_url: Final = "https://login.microsoftonline.test/config-only-keys"
+ monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
+ monkeypatch.setenv("JWT_ISSUER", issuer)
+ monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
+ private_key, jwk = _get_rsa_key_and_jwk(kid="config-key")
+ cache: Final = DualCache()
+ cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
+ registry: Final = AgentRegistry()
+ registry.register_agent(
+ AgentResponse(
+ agent_id="configured",
+ agent_name="Configured",
+ agent_card_params={},
+ identity_managed=kind == "managed-agent",
+ )
+ )
+ handler: Final = JWTHandler()
+ handler.update_environment(None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"]))
+ handler.bind_agent_lookup(registry)
+ token: Final = _encode_rsa_jwt(
+ private_key,
+ issuer=issuer,
+ audience="api://gateway",
+ kid="config-key",
+ extra_claims={
+ "tid": "test-tenant",
+ "azp": "application",
+ "scope": "litellm_proxy_admin",
+ **({"agent": "configured"} if kind != "human" else {}),
+ },
+ )
+ arguments: Final = dict(
+ api_key=token,
+ jwt_handler=handler,
+ request_data={},
+ general_settings={},
+ route="/chat/completions",
+ prisma_client=None,
+ user_api_key_cache=cache,
+ parent_otel_span=None,
+ proxy_logging_obj=MagicMock(),
+ )
+ if kind == "managed-agent":
+ with pytest.raises(HTTPException) as denied:
+ await JWTAuthManager.auth_builder(**arguments)
+ assert denied.value.status_code == 403
+ else:
+ result: Final = await JWTAuthManager.auth_builder(**arguments)
+ auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
+ assert auth.agent_id == ("configured" if kind == "config-agent" else None)
+ assert auth.managed_agent_context is None
+ assert result["is_proxy_admin"] is True
+
+
+@pytest.mark.parametrize(
+ "issuer,audience,disabled,expected",
+ [
+ (None, "gateway", False, False),
+ ("trusted", "gateway", False, True),
+ ("trusted", None, True, False),
+ ("other", "gateway", False, False),
+ ],
+)
+def test_managed_issuer_requires_configured_audience_validation(
+ monkeypatch: pytest.MonkeyPatch, issuer: str | None, audience: str | None, disabled: bool, expected: bool
+) -> None:
+ from litellm.proxy._types import JWTIssuerConfig
+
+ monkeypatch.delenv("JWT_ISSUER", raising=False)
+ monkeypatch.delenv("JWT_AUDIENCE", raising=False)
+ handler: Final = JWTHandler()
+ handler.update_environment(
+ None,
+ DualCache(),
+ LiteLLM_JWTAuth(
+ issuers=[
+ JWTIssuerConfig(issuer="trusted", audience=audience, disable_audience_validation=disabled),
+ ]
+ ),
+ )
+ assert handler.managed_issuer_is_trusted(issuer) is expected
diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
index 470db99108a..268b70c9596 100644
--- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
+++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py
@@ -9278,6 +9278,8 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
)
async def auth_that_reserves(request, api_key):
+ assert request.method == "GET"
+ assert request.query_params.get("model") == "gpt-realtime"
request.state.budget_reservation = reservation
return UserAPIKeyAuth(token="hashed", budget_reservation=reservation)
@@ -9290,3 +9292,201 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
assert result.budget_reservation == reservation
assert websocket.state.budget_reservation is reservation
assert websocket.scope["state"]["budget_reservation"] is reservation
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("invoke", [False, True])
+async def test_centralized_authorization_preserves_database_free_config_agents(monkeypatch, invoke: bool):
+ from litellm.proxy import proxy_server
+ from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
+ from litellm.proxy.agent_endpoints import agent_registry
+ from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
+
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "prisma_client": None,
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ registry = AgentRegistry()
+ registry.load_agents_from_config(
+ [{"agent_name": "config-agent", "agent_card_params": {"name": "Config", "url": "http://localhost:9999"}}]
+ )
+ monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
+ registered = registry.get_agent_by_name("config-agent")
+ model = "a2a/config-agent" if invoke else "test-model"
+ auth = UserAPIKeyAuth(agent_id=registered.agent_id, jwt_claims={"agent": "config-agent"}, models=[model])
+ data = {"model": model, "messages": [{"role": "user", "content": "hi"}]}
+ assert (
+ await _authorize_authenticated_request(
+ auth, _alias_request("/v1/chat/completions", data), data, "/v1/chat/completions", "jwt-token"
+ )
+ is None
+ )
+ assert auth.managed_agent_policy is None
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("verified_identity", [False, True])
+async def test_managed_actor_cannot_access_provider_resource_routes(monkeypatch, verified_identity: bool):
+ from litellm.proxy import proxy_server
+ from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ policy = AgentResponse(
+ agent_id="managed",
+ agent_name="Managed",
+ agent_card_params={},
+ identity_managed=True,
+ identity=AgentIdentityBinding(
+ agent_id="managed",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="application",
+ service_principal_id="principal",
+ issuer="issuer",
+ revision="revision",
+ ),
+ object_permission={"models": ["test-model"]},
+ )
+ database = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "prisma_client": database,
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ request = _alias_request("/v1/files", {})
+ request.scope["method"] = "GET"
+ from litellm.types.proxy.agent_identity import ManagedAgentContext
+
+ auth = UserAPIKeyAuth(
+ agent_id="managed", api_key="persisted-key", models=["test-model"],
+ managed_agent_context=(
+ ManagedAgentContext(agent_id="managed", binding_revision="revision", mode="autonomous")
+ if verified_identity else None
+ ),
+ )
+ with pytest.raises(ProxyException) as denied:
+ await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key")
+ assert denied.value.code == "403"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("requested", [None, "test-model"])
+@pytest.mark.parametrize("grant_default", [False, True])
+@pytest.mark.parametrize(
+ "route,settings,cli_model",
+ [
+ ("/v1/chat/completions", {"completion_model": "forbidden-model"}, None),
+ ("/v1/responses", {"completion_model": "forbidden-model"}, None),
+ ("/v1/messages", {"completion_model": "forbidden-model"}, None),
+ ("/v1/moderations", {"moderation_model": "forbidden-model"}, None),
+ ("/v1/audio/transcriptions", {"moderation_model": "forbidden-model"}, None),
+ ("/v1/audio/speech", {}, "forbidden-model"),
+ ("/v1/chat/completions", {}, "forbidden-model"),
+ ("/v1/images/generations", {"image_generation_model": "forbidden-model"}, None),
+ ("/v1/images/edits", {"image_generation_model": "forbidden-model"}, None),
+ ],
+)
+async def test_managed_agent_cannot_bypass_grants_with_server_default(
+ monkeypatch, requested, route, settings, cli_model, grant_default
+):
+ from litellm.proxy import proxy_server
+ from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext
+
+ policy = AgentResponse(
+ agent_id="managed",
+ agent_name="Managed",
+ agent_card_params={},
+ identity_managed=True,
+ identity=AgentIdentityBinding(
+ agent_id="managed",
+ provider="microsoft_entra",
+ tenant_id="tenant",
+ client_id="application",
+ service_principal_id="principal",
+ issuer="issuer",
+ revision="revision",
+ ),
+ object_permission={"models": ["test-model", "forbidden-model"] if grant_default else ["test-model"]},
+ )
+ database = MagicMock()
+ database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "prisma_client": database,
+ "general_settings": settings,
+ "user_model": cli_model,
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ data = {"messages": [{"role": "user", "content": "hi"}], **({"model": requested} if requested else {})}
+ auth = UserAPIKeyAuth(agent_id="managed")
+ auth.managed_agent_context = ManagedAgentContext(
+ agent_id="managed", binding_revision="revision", mode="autonomous"
+ )
+ if not grant_default:
+ with pytest.raises(ProxyException) as denied:
+ await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
+ assert denied.value.code == "403"
+ assert "forbidden-model" in denied.value.message
+ return
+ with patch(
+ "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
+ new_callable=AsyncMock,
+ ) as reserve:
+ reserve.return_value = None
+ assert (
+ await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
+ is None
+ )
+ reserve.assert_awaited_once()
+ assert reserve.call_args.kwargs["request_body"]["model"] == "forbidden-model"
+
+
+@pytest.mark.asyncio
+async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeypatch: pytest.MonkeyPatch) -> None:
+ from typing import Final
+
+ from litellm.proxy import proxy_server
+ from litellm.types.agents import AgentResponse
+ from litellm.types.proxy.agent_identity import AgentIdentityBinding
+
+ binding: Final = AgentIdentityBinding(
+ agent_id="managed", provider="microsoft_entra", issuer="issuer", tenant_id="tenant",
+ client_id="client", service_principal_id="principal", revision="current",
+ )
+ agent: Final = AgentResponse(
+ agent_id="managed", agent_name="Managed", agent_card_params={},
+ identity_managed=True, identity=binding, execution_mode="autonomous",
+ )
+ client: Final = MagicMock()
+ client.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
+ client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
+ handler: Final = MagicMock()
+ handler.is_jwt.return_value = True
+ handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub")
+ handler.auth_jwt = AsyncMock(return_value={
+ "iss": "issuer", "tid": "tenant", "azp": "client", "oid": "principal", "sub": "mapped-key",
+ })
+ for name, value in {
+ **_proxy_attrs_for_centralized_checks(),
+ "general_settings": {"enable_jwt_auth": True}, "premium_user": True,
+ "prisma_client": client, "jwt_handler": handler, "user_api_key_cache": DualCache(),
+ "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
+ }.items():
+ monkeypatch.setattr(proxy_server, name, value)
+ with pytest.raises(ProxyException) as failure:
+ await _user_api_key_auth_builder(
+ request=_alias_request("/v1/chat/completions", {}), api_key="Bearer verified.jwt.token",
+ azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
+ azure_apim_header=None, request_data={},
+ )
+ assert failure.value.code == "403"
+ assert "without virtual-key mapping" in failure.value.message
+ client.writer_db.litellm_agentstable.find_unique.assert_awaited_once()
diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py
index 7929a0b21af..6087aa1f570 100644
--- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py
+++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py
@@ -1,6 +1,7 @@
import io
import json
-from typing import get_type_hints
+from collections.abc import Mapping
+from typing import Literal, get_type_hints
from unittest.mock import AsyncMock, MagicMock, patch
import orjson
@@ -1210,3 +1211,42 @@ class TestCoerceNumericFormFields:
numeric_fields=self.numeric_fields,
)
assert result == {"n": 3, "temperature": None, "image": buffer}
+
+
+@pytest.mark.parametrize(
+ "kind,settings,cli,path,body,expected",
+ [
+ ("completion", {"completion_model": "default"}, "cli", "path", "body", "default"),
+ ("completion", {}, "cli", "path", "body", "cli"),
+ ("completion", {}, None, "path", "body", "path"),
+ ("completion", {}, None, None, "body", "body"),
+ (
+ "image_generation",
+ {"completion_model": "text", "image_generation_model": "image"},
+ None,
+ None,
+ "body",
+ "image",
+ ),
+ ("image_generation", {"image_generation_model": "image"}, "cli", "path", "body", "cli"),
+ ("image_generation", {"image_generation_model": "image"}, None, "path", "body", "path"),
+ ("image_edit", {"completion_model": "text", "image_generation_model": "image"}, None, None, "body", "text"),
+ ("image_edit", {"image_generation_model": "image"}, None, "path", "body", "path"),
+ ("image_edit", {"image_generation_model": "image"}, None, None, "body", "image"),
+ ("moderation", {"moderation_model": "mod"}, "cli", None, "body", "cli"),
+ ("speech", {"completion_model": "text"}, None, None, "body", "body"),
+ ("body", {"completion_model": "text"}, "cli", None, "body", "body"),
+ ("path", {"completion_model": "text"}, "cli", "path", "body", "path"),
+ ],
+)
+def test_shared_inference_model_selection_preserves_handler_precedence(
+ kind: Literal["completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"],
+ settings: Mapping[str, object],
+ cli: str | None,
+ path: str | None,
+ body: str,
+ expected: str,
+) -> None:
+ from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
+
+ assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected
diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
index ca2ff8bcce1..9e20386bf3d 100644
--- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
+++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py
@@ -177,7 +177,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep
assert agent.agent_id == agent_id
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once_with(
where={"agent_id": agent_id},
- include={"object_permission": True},
+ include={"object_permission": True, "identity": True},
)
@@ -202,7 +202,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re
assert agent.agent_name == agent_name
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_with(
where={"agent_name": agent_name},
- include={"object_permission": True},
+ include={"object_permission": True, "identity": True},
)
@@ -521,3 +521,33 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra
assert await resync_task is True
assert len(clean_agent_registry.agent_list) == 1
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("lookup", ["agent-id", "Agent name"])
+async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch):
+ from types import SimpleNamespace
+ from unittest.mock import AsyncMock, MagicMock
+
+ from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
+
+ binding = {
+ "agent_id": "agent-id", "provider": "microsoft_entra", "tenant_id": "tenant", "client_id": "client",
+ "issuer": "https://login.microsoftonline.com/tenant/v2.0", "revision": "revision",
+ }
+
+ async def load_row(*, where, include):
+ if where == {"agent_id": "Agent name"}:
+ return None
+ row = FakeAgentRow("agent-id", "Agent name").model_dump()
+ return SimpleNamespace(model_dump=lambda: {**row, "identity": binding if include.get("identity") else None})
+
+ prisma = MagicMock()
+ prisma.db.litellm_agentstable.find_unique = AsyncMock(side_effect=load_row)
+ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
+ monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
+ agent = await get_agent_with_read_through(lookup)
+ assert agent is not None
+ assert agent.identity is not None
+ assert agent.identity.model_dump(include=set(binding)) == binding
+ assert clean_agent_registry.get_agent_by_id(agent_id="agent-id").identity == agent.identity
diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py
index b5e594db701..d7401554d58 100644
--- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py
+++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py
@@ -2712,3 +2712,34 @@ async def test_track_cost_callback_failure_alert_never_carries_request_metadata_
assert "headers" in failure_debug_lines[0]
else:
assert failure_debug_lines == []
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("identity_field", ["agent_id", "billing_agent_id"])
+async def test_autonomous_llm_callback_persists_without_human_or_key(identity_field: str) -> None: # test-quality-ok: verifies anonymous-agent charges reach the persistence boundary; no injection seam
+ kwargs: Final = {
+ "call_type": "acompletion",
+ "model": "test-model",
+ "response_cost": 0.01,
+ "litellm_params": {"metadata": {identity_field: "autonomous-agent"}},
+ }
+ with patch(
+ "litellm.proxy.hooks.proxy_track_cost_callback._update_database_and_spend_counters",
+ new_callable=AsyncMock,
+ return_value=False,
+ ) as persist:
+ await _ProxyDBLogger()._PROXY_track_cost_callback(
+ kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now()
+ )
+ persist.assert_awaited_once()
+ assert persist.call_args.kwargs["response_cost"] == 0.01
+ assert persist.call_args.kwargs["user_id"] is None
+ assert persist.call_args.kwargs["user_api_key"] is None
+ assert persist.call_args.kwargs["kwargs"]["litellm_params"]["metadata"][identity_field] == "autonomous-agent"
+
+
+@pytest.mark.parametrize("agent_id,expected", [(None, False), ("autonomous-agent", True)])
+def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id: str | None, expected: bool) -> None:
+ assert _should_track_cost_callback(
+ user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id
+ ) is expected
diff --git a/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py b/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py
new file mode 100644
index 00000000000..68fe77cc76a
--- /dev/null
+++ b/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py
@@ -0,0 +1,114 @@
+from types import SimpleNamespace
+from typing import Final
+from unittest.mock import AsyncMock
+
+import pytest
+from fastapi import HTTPException
+
+from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import (
+ enroll_microsoft_subject,
+ microsoft_interactive_subject,
+)
+
+TENANT: Final = "11111111-1111-4111-8111-111111111111"
+OID: Final = "22222222-2222-4222-8222-222222222222"
+
+
+def test_enrollment_uses_provider_object_id_and_configured_tenant() -> None:
+ subject: Final = microsoft_interactive_subject(
+ TENANT, {"id": OID, "mail": "alias@example.com", "tid": "untrusted"}, {}
+ )
+ assert subject is not None
+ assert subject.oid == OID
+ assert subject.tenant_id == TENANT
+ assert subject.issuer == f"https://login.microsoftonline.com/{TENANT}/v2.0"
+
+
+@pytest.mark.parametrize("tenant", [None, "common", "organizations", "invalid"])
+def test_multitenant_sso_does_not_guess_the_subject_tenant(tenant: str | None) -> None:
+ assert microsoft_interactive_subject(tenant, {"id": OID, "tid": TENANT}, {}) is None
+
+
+@pytest.mark.parametrize("response", [{"mail": "user@example.com"}, {"id": "user@example.com"}, {"id": 42}])
+def test_email_and_configurable_aliases_are_not_human_subject_proof(response: dict[str, object]) -> None:
+ assert microsoft_interactive_subject(TENANT, response, {}) is None
+
+
+@pytest.mark.parametrize(
+ "endpoint", ["MICROSOFT_USERINFO_ENDPOINT", "MICROSOFT_TOKEN_ENDPOINT", "MICROSOFT_AUTHORIZATION_ENDPOINT"]
+)
+def test_custom_provider_endpoints_do_not_enroll_trusted_microsoft_subjects(endpoint: str) -> None:
+ assert microsoft_interactive_subject(TENANT, {"id": OID}, {endpoint: "https://custom.example"}) is None
+
+
+@pytest.mark.asyncio
+async def test_interactive_enrollment_preserves_the_canonical_local_user() -> None:
+ table: Final = AsyncMock()
+ table.upsert.return_value = SimpleNamespace(kind="human", user_id="canonical", verified_via="sso_interactive")
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ subject: Final = microsoft_interactive_subject(TENANT, {"id": OID}, {})
+ assert subject is not None
+ await enroll_microsoft_subject(subject, "canonical", client)
+ table.upsert.assert_awaited_once_with(
+ where={"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": TENANT, "oid": OID}},
+ data={
+ "create": {
+ "issuer": subject.issuer,
+ "tenant_id": TENANT,
+ "oid": OID,
+ "user_id": "canonical",
+ "verified_via": "sso_interactive",
+ },
+ "update": {},
+ },
+ )
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("user_id,verified_via", [("another-user", "sso_interactive"), ("canonical", "untrusted")])
+async def test_interactive_enrollment_does_not_reassign_an_existing_subject(user_id: str, verified_via: str) -> None:
+ table: Final = AsyncMock()
+ table.upsert.return_value = SimpleNamespace(kind="human", user_id=user_id, verified_via=verified_via)
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ with pytest.raises(HTTPException) as failure:
+ await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
+ assert failure.value.status_code == 403
+ assert table.upsert.call_args.kwargs["data"]["update"] == {}
+
+
+@pytest.mark.asyncio
+async def test_enrollment_storage_failure_is_not_a_successful_login() -> None:
+ table: Final = AsyncMock()
+ table.upsert.side_effect = RuntimeError("database unavailable")
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ with pytest.raises(HTTPException) as failure:
+ await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
+ assert failure.value.status_code == 503
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("user_id", [None, "", 42])
+async def test_enrollment_requires_a_canonical_local_user(user_id: object) -> None:
+ table: Final = AsyncMock()
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), user_id, client)
+ table.upsert.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_untrusted_metadata_cannot_enroll_a_human() -> None:
+ table: Final = AsyncMock()
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ await enroll_microsoft_subject({"issuer": "forged", "tenant_id": TENANT, "oid": OID}, "canonical", client)
+ table.upsert.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_scim_agent_subject_cannot_be_enrolled_as_a_human() -> None:
+ table: Final = AsyncMock()
+ table.upsert.return_value = SimpleNamespace(kind="agent_user", user_id=None, verified_via="scim")
+ client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
+ with pytest.raises(HTTPException) as failure:
+ await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
+ assert failure.value.status_code == 403
+ assert table.upsert.call_args.kwargs["data"]["update"] == {}
diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
index 11b3dcf54bc..ce11a604192 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
@@ -7659,7 +7659,7 @@ class TestConnectedAppViewAnnotation:
flags = {server.server_id: server.connected_app_reachable for server in result}
assert flags == {"server-1": True, "server-2": False}
- reload_mock.assert_awaited_once_with("test_user_id")
+ reload_mock.assert_awaited_once_with("test_user_id", requires_fresh_policy=False)
mock_manager.get_allowed_mcp_servers.assert_awaited_once_with(admitted_auth)
@pytest.mark.asyncio
@@ -9411,7 +9411,7 @@ class TestMCPServerResolutionCharacterization:
server_id: str,
) -> tuple[MagicMock, MCPServerManager, UserAPIKeyAuth]:
team_id: Final = UI_SESSION_TOKEN_TEAM_ID if grant_route == "direct user object_permission" else "lit3974_team"
- user_id: Final = "lit3974_direct_user"
+ user_id: Final = f"{server_id}:{grant_route}:user"
key_permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id=f"lit3974_{grant_route}_key_permission",
mcp_servers=None,
diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
index 6d902cb7fec..54827f094f0 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
@@ -4355,7 +4355,7 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs(
prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
prisma_client.db.litellm_teamtable.count = AsyncMock(side_effect=count)
prisma_client.db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
- prisma_client.db.litellm_usertable.find_unique = AsyncMock(
+ prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=LiteLLM_UserTable(
user_id="org_admin_user",
teams=["team_in_org_A", "team_in_org_B"],
@@ -4394,11 +4394,11 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs(
assert await list_teams(None) == own_view
assert await list_teams("org_admin_user", search="team_in_org_B") == ["team_in_org_B"]
assert await list_teams("other_user") == ["other_team_in_org_A"]
- prisma_client.db.litellm_usertable.find_unique.assert_awaited_with(
+ prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_with(
where={"user_id": "org_admin_user"}, include={"organization_memberships": True}
)
- prisma_client.db.litellm_usertable.find_unique.side_effect = RuntimeError("db down")
+ prisma_client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("db down")
with pytest.raises(ValueError, match="db down"):
await list_teams("org_admin_user")
@@ -16025,7 +16025,7 @@ async def test_get_team_spend_by_user_team_admin_sees_every_member(mock_db_clien
alpha = _team_spend_by_user_team("team-alpha", "Team Alpha", Member(user_id="alice", role="admin"), [])
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha])
mock_db_client.db.query_raw = AsyncMock(return_value=[])
- mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
+ mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=_team_spend_by_user_caller("alice", ["team-alpha"])
)
@@ -16047,7 +16047,7 @@ async def test_get_team_spend_by_user_plain_member_only_sees_own_row(mock_db_cli
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha])
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_db_client.db.query_raw = AsyncMock(return_value=[_team_spend_by_user_db_row("team-alpha", "bob", 0.25, 2)])
- mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
+ mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=_team_spend_by_user_caller("bob", ["team-alpha"])
)
@@ -16068,7 +16068,7 @@ async def test_get_team_spend_by_user_member_of_other_team_gets_404(mock_db_clie
caller = UserAPIKeyAuth(user_id="bob", user_role=LitellmUserRoles.INTERNAL_USER)
mock_db_client.db.query_raw = AsyncMock(return_value=[])
- mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
+ mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=_team_spend_by_user_caller("bob", ["team-alpha"])
)
diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
index 21c0f565486..971019160f9 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
@@ -206,6 +206,7 @@ def test_microsoft_sso_handler_openid_from_response_with_custom_attributes():
def test_get_microsoft_callback_response():
# Arrange
mock_request = MagicMock(spec=Request)
+ mock_request.scope = {}
mock_response = {
"mail": "microsoft_user@example.com",
"displayName": "Microsoft User",
@@ -8751,6 +8752,7 @@ async def test_redirect_from_openid_persists_assertion_under_canonical_user_id()
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
assert assertion is not None
mock_request = MagicMock(spec=Request)
+ mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
@@ -8989,6 +8991,7 @@ async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplo
"""Wiring: the browser login path must reach the diagnostic, not just define it."""
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
mock_request = MagicMock(spec=Request)
+ mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
mock_request.cookies = {}
diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
index 7378564f7a8..0157200ed5c 100644
--- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
+++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py
@@ -4699,24 +4699,28 @@ def _config_agent(agent_name: str) -> Dict[str, Any]:
}
-class _FakeAgentRow:
- """Stand-in for a prisma agent record: supports dict() and .object_permission."""
+def _agent_db_row(agent_id: str, agent_name: str):
+ import json
+ from datetime import datetime, timezone
- def __init__(self, agent_id: str, agent_name: str) -> None:
- self.agent_id = agent_id
- self.agent_name = agent_name
- self.object_permission = None
- self.spend = 0.0
+ from prisma.models import LiteLLM_AgentsTable
- def __iter__(self):
- return iter(
- {
- "agent_id": self.agent_id,
- "agent_name": self.agent_name,
- "agent_card_params": {"name": self.agent_name, "url": "http://db-agent"},
- "litellm_params": {},
- }.items()
- )
+ return LiteLLM_AgentsTable(
+ agent_id=agent_id,
+ agent_name=agent_name,
+ agent_card_params=json.dumps({"name": agent_name, "url": "http://db-agent"}),
+ extra_headers=[],
+ agent_access_groups=[],
+ access_group_ids=[],
+ spend=0.0,
+ identity_managed=False,
+ enabled=True,
+ execution_mode="autonomous",
+ created_at=datetime.now(timezone.utc),
+ updated_at=datetime.now(timezone.utc),
+ created_by="admin",
+ updated_by="admin",
+ )
@pytest.mark.asyncio
@@ -4740,7 +4744,7 @@ async def test_ProxyConfig__init_agents_in_db_keeps_config_defined_agents(clean_
)
prisma_client = MagicMock()
- prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_FakeAgentRow("db-id", "db-agent")])
+ prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_agent_db_row("db-id", "db-agent")])
await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client)
@@ -4777,7 +4781,7 @@ async def test_ProxyStartupEvent_jwt_auth_resolves_agent_claims_against_live_reg
elif agents_source == "db":
prisma_client = MagicMock()
prisma_client.db.litellm_agentstable.find_many = AsyncMock(
- return_value=[_FakeAgentRow("db-id", "loaded-agent")]
+ return_value=[_agent_db_row("db-id", "loaded-agent")]
)
await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client)
else:
diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
index 347adc421a2..8cf34ce51d5 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
@@ -4,6 +4,7 @@ import datetime
import hashlib
import json
import re
+import sqlite3
from datetime import timezone
from unittest.mock import AsyncMock, MagicMock, patch
@@ -3762,7 +3763,7 @@ class TestSpendLogsPayload:
"model": "gpt-4o",
"user": "",
"team_id": "",
- "metadata": '{"applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
+ "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
"cache_key": "Cache OFF",
"spend": 0.00022500000000000002,
"total_tokens": 30,
@@ -3781,6 +3782,7 @@ class TestSpendLogsPayload:
"status": "success",
"mcp_namespaced_tool_name": None,
"agent_id": None,
+ "billing_agent_id": None,
}
)
@@ -6586,9 +6588,7 @@ def test_key_spend_report_scopes_to_caller_key(client, monkeypatch):
def test_key_spend_report_scopes_a_cli_session_to_the_per_user_alias_not_the_login_token(client, monkeypatch):
- mock_prisma = _spend_report_mock_prisma(
- query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}]
- )
+ mock_prisma = _spend_report_mock_prisma(query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
@@ -7138,9 +7138,8 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc
rep_call = emitted[2]
assert f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in rep_call[0]
assert (
- f"ORDER BY {SESSION_GROUP_KEY_SQL}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
- in rep_call[0]
- ), "the session representative must prefer the newest non-MCP call"
+ f"ORDER BY {SESSION_GROUP_KEY_SQL}, " + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
+ ) in rep_call[0]
assert rep_call[-2] == ["sess-1", "req-solo"]
assert rep_call[-1] == ["hashed-key", "hashed-key"]
finally:
@@ -7630,6 +7629,39 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo
app.dependency_overrides.pop(ps.user_api_key_auth, None)
+@pytest.mark.parametrize(
+ ("parent_status", "child_status", "expected"),
+ [("failure", "success", "failure"), ("success", "failure", "success")],
+)
+def test_session_representative_uses_completed_agent_outcome(parent_status, child_status, expected):
+ with sqlite3.connect(":memory:") as connection:
+ connection.execute(
+ 'CREATE TABLE logs (request_id TEXT, call_type TEXT, status TEXT, "startTime" TEXT, "endTime" TEXT)'
+ )
+ connection.executemany(
+ "INSERT INTO logs VALUES (?, ?, ?, ?, ?)",
+ (
+ ("parent", "asend_message", parent_status, "10:00:00", "10:00:05"),
+ ("nested-agent", "asend_message", child_status, "10:00:01", "10:00:03"),
+ ("llm", "acompletion", "success", "10:00:02", "10:00:04"),
+ ("tool", "call_mcp_tool", child_status, "10:00:04", "10:00:04"),
+ ),
+ )
+ result = connection.execute(
+ "SELECT request_id, status FROM logs ORDER BY "
+ + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
+ + " LIMIT 1"
+ ).fetchone()
+ assert result == ("parent", expected)
+ connection.execute("DELETE FROM logs WHERE call_type = 'asend_message'")
+ fallback = connection.execute(
+ "SELECT request_id, status FROM logs ORDER BY "
+ + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
+ + " LIMIT 1"
+ ).fetchone()
+ assert fallback == ("llm", "success")
+
+
@pytest.mark.asyncio
async def test_calculate_spend_unpriced_model_returns_400():
model = "openrouter/unit-test-unpriced-model"
diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py
index 54e5a6d5385..6752c91e9f2 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py
@@ -539,7 +539,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
async def mock_query_raw(sql_query, *params):
if "COUNT(*) AS total_count" in sql_query:
return [{"total_count": 60}]
- if "DISTINCT ON" in sql_query:
+ if "AS session_representatives" in sql_query:
return representative_rows
return session_rows
@@ -584,9 +584,11 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
rep_sql = emitted[2][0]
assert f"DISTINCT ON ({group_key})" in rep_sql, f"page must return one row per session. SQL was:\n{rep_sql}"
- assert f"ORDER BY {group_key}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in rep_sql, (
- "the session representative must prefer the newest non-MCP call"
- )
+ assert (
+ f"ORDER BY {group_key}, (call_type = 'asend_message') DESC, "
+ "CASE WHEN call_type = 'asend_message' THEN \"endTime\" END DESC NULLS LAST, "
+ "call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
+ ) in rep_sql, "the session representative must prefer the final agent outcome, then the newest non-MCP call"
assert "COUNT(*) OVER ()" not in rep_sql
assert [row["request_id"] for row in response["data"]] == ["req-1", "req-2"]
diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
index 00223f192ec..cef60e5c8da 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
@@ -5156,6 +5156,24 @@ def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_recei
)
+def test_failed_agent_request_keeps_registered_display_name():
+ agent_model: Final = "a2a_agent/Research Agent"
+ payload: Final = get_logging_payload(
+ kwargs={
+ "model": agent_model,
+ "call_type": "asend_message",
+ "litellm_params": {
+ "metadata": {"model_group": agent_model, "model_info": {"id": "registered-agent"}, "status": "failure"}
+ },
+ },
+ response_obj=ValueError("Agent action denied"),
+ start_time=datetime.datetime.now(timezone.utc),
+ end_time=datetime.datetime.now(timezone.utc),
+ )
+ assert payload["model"] == agent_model
+ assert payload["status"] == "failure"
+ assert payload["model_id"] == "registered-agent"
+
_CLI_SESSION_ALIAS: Final = "cli-session-alice"
_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
@@ -5281,3 +5299,21 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None:
assert result["autorouter_savings_estimate"] == recorded
absent: Final = _get_spend_logs_metadata({"autorouter_savings_estimate": supplied}) # mutable-ok: legacy metadata helper accepts dicts
assert absent["autorouter_savings_estimate"] is None
+
+
+@pytest.mark.parametrize("billing_agent", [None, "authenticated-agent"])
+def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None:
+ kwargs = {
+ "model": "gpt-4",
+ "litellm_params": {"metadata": {
+ "user_api_key": "test-key",
+ "agent_id": "header-selected-agent",
+ "billing_agent_id": billing_agent,
+ }},
+ }
+ payload = get_logging_payload(
+ kwargs=kwargs, response_obj={"id": "request"},
+ start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc),
+ )
+ assert payload["agent_id"] == "header-selected-agent"
+ assert payload["billing_agent_id"] == billing_agent
diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py
index 35308474949..0429dd97a1c 100644
--- a/tests/test_litellm/proxy/test_route_a2a_models.py
+++ b/tests/test_litellm/proxy/test_route_a2a_models.py
@@ -149,6 +149,9 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re
prisma_client.db.litellm_agentstable.find_unique = AsyncMock(
side_effect=[None, _DbAgentRow("a2a-sibling-replica-agent-id", agent_name)]
)
+ prisma_client.writer_db.litellm_agentstable.find_unique = AsyncMock(
+ return_value=_DbAgentRow("a2a-sibling-replica-agent-id", agent_name)
+ )
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx
new file mode 100644
index 00000000000..ce54ab78d9c
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx
@@ -0,0 +1,43 @@
+import { screen } from "@testing-library/react";
+import { beforeEach, describe, expect, it, vi } from "vitest";
+import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils";
+import { apiClient } from "@/components/networking";
+import { AgentIdentityDetails } from "./AgentIdentityDetails";
+
+vi.mock("@/components/networking", () => ({ apiClient: { get: vi.fn() } }));
+
+const identity = {
+ provider: "microsoft_entra",
+ tenant_id: "11111111-1111-4111-8111-111111111111",
+ client_id: "22222222-2222-4222-8222-222222222222",
+};
+
+const status = {
+ enabled: true,
+ execution_mode: "autonomous",
+ last_authenticated_at: "2026-09-24T12:00:00Z",
+};
+
+describe("agent identity evidence", () => {
+ beforeEach(() => {
+ vi.clearAllMocks();
+ testQueryClient.clear();
+ });
+
+ it("shows persisted application identity evidence and links to the current logs route", async () => {
+ vi.mocked(apiClient.get).mockResolvedValue(status);
+ renderWithProviders( );
+ expect(await screen.findByText(/Last authenticated identity match:/)).toBeInTheDocument();
+ expect(screen.getByText(/Application \(Client\) ID:/)).toBeInTheDocument();
+ expect(screen.getByRole("link", { name: "View request logs" })).toHaveAttribute("href", "/ui/logs/");
+ expect(apiClient.get).toHaveBeenCalledWith("/v1/agents/native/identity", { accessToken: "admin" });
+ });
+
+ it("does not request or show administrator identity evidence to ordinary users", () => {
+ renderWithProviders(
+ ,
+ );
+ expect(screen.queryByRole("region", { name: "Agent Identity" })).not.toBeInTheDocument();
+ expect(apiClient.get).not.toHaveBeenCalled();
+ });
+});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx
new file mode 100644
index 00000000000..12465c6e861
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx
@@ -0,0 +1,81 @@
+import React from "react";
+import type { components } from "@/lib/http/schema";
+import { useQuery } from "@tanstack/react-query";
+import { apiClient } from "@/components/networking";
+import { Button } from "@/components/ui/button";
+import { readAgentIdentity } from "./agent_identity";
+
+const authenticationMessage = (error: boolean, lastAuthenticated?: string | null): string => {
+ if (error) return "Could not load authentication evidence";
+ if (lastAuthenticated) return `Last authenticated identity match: ${new Date(lastAuthenticated).toLocaleString()}`;
+ return "Configured, awaiting an authenticated request";
+};
+
+export const AgentIdentityDetails = ({
+ agentId,
+ identity: value,
+ accessToken,
+ isAdmin,
+}: {
+ agentId: string;
+ identity: unknown;
+ accessToken: string | null;
+ isAdmin: boolean;
+}) => {
+ const identity = readAgentIdentity(value);
+ const { data, isError, isFetching, refetch } = useQuery({
+ queryKey: ["agent-identity", agentId, identity],
+ queryFn: () =>
+ apiClient.get(
+ `/v1/agents/${encodeURIComponent(agentId)}/identity`,
+ {
+ accessToken: accessToken ?? "",
+ },
+ ),
+ enabled: Boolean(isAdmin && accessToken && identity),
+ });
+
+ if (!identity || !isAdmin) return null;
+ const executionLabel = data?.enabled ? "Enabled" : "Disabled";
+ return (
+
+ Agent Identity: Microsoft Entra ID
+
+ Tenant: {identity.tenant_id}
+
+ <>
+
+ Application (Client) ID: {identity.client_id}
+
+ Enterprise application Object ID: {identity.service_principal_id || "Not configured"}
+ >
+
+ Execution: {data ? executionLabel : "Loading"} · Mode: {data?.execution_mode ?? "Loading"}
+
+
+ {data?.identity?.active === false
+ ? "Identity unbound; execution is disabled"
+ : authenticationMessage(isError, data?.last_authenticated_at)}
+
+
+ Recent evidence comes from a validated Entra token matching this binding. It is persisted across restarts and
+ cleared when the binding changes. Tool and model permissions are checked separately.
+
+
+
+ );
+};
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
new file mode 100644
index 00000000000..43538784145
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx
@@ -0,0 +1,259 @@
+import React, { useEffect, useState } from "react";
+import { useWatch } from "react-hook-form";
+import { apiClient } from "@/components/networking";
+import { Input } from "@/components/ui/input";
+import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
+import { AgentFormField, type AgentFormValues } from "./AgentFormKit";
+import { entraTenantFromIssuer, IDENTITY_UUID_PATTERN } from "./agent_identity";
+
+const PROVIDER_OPTIONS = [
+ { value: "none", label: "No explicit identity binding" },
+ { value: "microsoft_entra", label: "Microsoft Entra ID" },
+];
+const EXECUTION_MODE_OPTIONS = [
+ { value: "autonomous", label: "Autonomous" },
+ { value: "delegated", label: "On behalf of a user" },
+ { value: "both", label: "Both" },
+];
+const EXECUTION_OPTIONS = [
+ { value: "enabled", label: "Enabled" },
+ { value: "disabled", label: "Disabled" },
+];
+
+export const AgentIdentityFields = ({ accessToken }: { accessToken: string | null }) => {
+ const provider = useWatch({ name: "identity_provider" });
+ const mode = useWatch({ name: "execution_mode" });
+ const showScopes = mode !== "autonomous" && mode !== undefined;
+ const [tenants, setTenants] = useState([]);
+ const [error, setError] = useState(null);
+
+ useEffect(() => {
+ if (!accessToken || provider !== "microsoft_entra") return;
+ let active = true;
+ apiClient
+ .get("/v1/agents/identity/providers", { accessToken })
+ .then((issuers) => {
+ if (active)
+ setTenants(
+ issuers.flatMap((issuer) => {
+ const tenant = entraTenantFromIssuer(issuer);
+ return tenant ? [tenant] : [];
+ }),
+ );
+ })
+ .catch(() => {
+ if (active) setError("Could not load the gateway's trusted identity providers");
+ });
+ return () => {
+ active = false;
+ };
+ }, [accessToken, provider]);
+
+ return (
+ <>
+
+
+
Agent Identity
+
+ Connect an existing identity provider application to this agent. Its name and runtime address can change
+ independently.
+
+
+
+ {({ value, onChange, id }) => (
+
+
+
+
+
+ {PROVIDER_OPTIONS.map((option) => (
+
+ {option.label}
+
+ ))}
+
+
+ )}
+
+ {provider === "microsoft_entra" && (
+ <>
+
+ {({ value, onChange, id }) => (
+
+
+
+
+
+ {tenants.map((tenant) => (
+
+ {tenant}
+
+ ))}
+
+
+ )}
+
+ {error && (
+
+ {error}
+
+ )}
+ {!error && tenants.length === 0 && (
+
+ No trusted Entra tenant is available. Configure JWT issuer and audience validation on the gateway first.
+ Dashboard Microsoft SSO is configured separately.
+
+ )}
+
+ Find this under{" "}
+
+ Entra App registrations
+
+ , select your agent application, then Overview. No client secret is required here.
+ >
+ }
+ >
+ {({ value, onChange, ref, ...control }) => (
+
+ )}
+
+
+ {({ value, onChange, id }) => (
+
+
+
+
+
+ {EXECUTION_MODE_OPTIONS.map((option) => (
+
+ {option.label}
+
+ ))}
+
+
+ )}
+
+
+ Open{" "}
+
+ Entra Enterprise applications
+
+ , select this application, and copy its Object ID. The App registrations Object ID is a different
+ value.
+ >
+ }
+ >
+ {({ value, onChange, ref, ...control }) => (
+
+ )}
+
+
+
+ {({ value, onChange, ref, ...control }) => (
+
+ )}
+
+
+ {showScopes && (
+ <>
+
+ {({ value, onChange, ref, ...control }) => (
+
+ )}
+
+
+ Users must first sign in through this gateway's Microsoft SSO. Subsequent delegated calls must
+ satisfy both user and agent permissions.
+
+ >
+ )}
+
+ {({ value, onChange, id }) => (
+ onChange(next === "enabled")}
+ >
+
+
+
+
+ {EXECUTION_OPTIONS.map((option) => (
+
+ {option.label}
+
+ ))}
+
+
+ )}
+
+
+ LiteLLM verifies the agent's Entra token before matching this identity. Saving these fields
+ configures the binding; an authenticated request provides verification. Runtime authentication headers are
+ configured separately.
+
+ >
+ )}
+
+ >
+ );
+};
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx
index f53b03b6a08..b392e270d33 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx
@@ -145,10 +145,10 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams
- Why do agents need keys?
+ How do agents authenticate?
- Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from
- the Virtual Keys page.
+ Agents can authenticate with a virtual key or a trusted identity provider using JWT. Configure an identity
+ binding when adding or editing an agent. JWT authentication does not require a virtual key.
{isAdmin && (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx
index bef938cd31c..95fcb516c55 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx
@@ -62,6 +62,12 @@ describe("AgentsTable", () => {
expect(within(keylessRow).getByText("Needs Setup")).toBeInTheDocument();
});
+ it("shows JWT configured for agents without a virtual key", () => {
+ render( );
+ expect(screen.getByText("JWT configured")).toBeInTheDocument();
+ expect(screen.queryByText("Needs Setup")).not.toBeInTheDocument();
+ });
+
it("deletes an agent through the ⋯ actions menu", async () => {
const user = userEvent.setup();
const onDeleteClick = vi.fn();
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx
index a8fe3973a42..932bd5c513f 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx
@@ -136,6 +136,7 @@ export const getAgentsTableColumns = ({
enableSorting: false,
cell: ({ row }) => {
const hasKeys = (row.original.keys?.length ?? 0) > 0;
+ if (row.original.jwt_auth_configured) return ;
return hasKeys ? (
) : (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx
index 457ee656415..b7585e8f8b2 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx
@@ -1,5 +1,5 @@
import React from "react";
-import { screen, waitFor, within } from "@testing-library/react";
+import { fireEvent, screen, waitFor, within } from "@testing-library/react";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import AddAgentForm from "./add_agent_form";
@@ -8,6 +8,7 @@ import type { AgentCreateInfo } from "@/components/networking";
import { chooseSelectOption, renderWithProviders as render } from "../../../../../tests/test-utils";
vi.mock("@/components/networking", () => ({
+ apiClient: { get: vi.fn() },
createAgentCall: vi.fn(),
getAgentCreateMetadata: vi.fn(),
getAgentsList: vi.fn(),
@@ -95,6 +96,51 @@ describe("AddAgentForm submit payload", () => {
.mockResolvedValue({} as never);
});
+ it("registers a readable agent with an explicit Entra identity and no virtual key", async () => {
+ const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
+ const tenant = "11111111-1111-4111-8111-111111111111";
+ const clientId = "22222222-2222-4222-8222-222222222222";
+ vi.mocked(networking.apiClient.get).mockResolvedValue([`https://login.microsoftonline.com/${tenant}/v2.0`]);
+ renderForm();
+ fireEvent.change(await screen.findByLabelText("Agent Name"), { target: { value: "Readable agent" } });
+ fireEvent.change(screen.getByLabelText("URL"), { target: { value: "https://runtime.example/a2a" } });
+ fireEvent.change(screen.getByLabelText("Display Name"), { target: { value: "Readable agent" } });
+ fireEvent.change(screen.getByPlaceholderText("Describe what this agent does..."), {
+ target: { value: "Test agent" },
+ });
+ await user.click(screen.getByLabelText("Identity Provider"));
+ await user.click(await screen.findByRole("option", { name: "Microsoft Entra ID" }));
+ await user.click(screen.getByLabelText("Trusted Entra Tenant"));
+ await user.click(await screen.findByRole("option", { name: tenant }));
+ fireEvent.change(screen.getByLabelText("Application (Client) ID"), { target: { value: clientId } });
+ fireEvent.change(screen.getByLabelText("Enterprise Application Object ID"), {
+ target: { value: "33333333-3333-4333-8333-333333333333" },
+ });
+ await user.click(screen.getByRole("button", { name: /^Next/ }));
+ await user.click(screen.getByRole("button", { name: /^Next/ }));
+ await user.click(screen.getByRole("button", { name: /^Next/ }));
+ await user.click(screen.getByRole("button", { name: "Use Entra JWT authentication" }));
+ await user.click(screen.getByRole("button", { name: /Create Agent/ }));
+ await waitFor(() => expect(networking.createAgentCall).toHaveBeenCalledTimes(1));
+ expect(createdPayload().agent_name).toBe("Readable agent");
+ const expectedIdentity = {
+ provider: "microsoft_entra",
+ tenant_id: tenant,
+ client_id: clientId,
+ service_principal_id: "33333333-3333-4333-8333-333333333333",
+ required_roles: [],
+ required_scopes: ["user_impersonation"],
+ };
+ expect(createdPayload().identity).toEqual(expectedIdentity);
+ expect(createdPayload()).not.toHaveProperty("litellm_params.identity");
+ expect(networking.keyCreateForAgentCall).not.toHaveBeenCalled();
+ expect(
+ screen.getByText(
+ "Microsoft Entra ID is configured. Send an authenticated agent request to verify the connection.",
+ ),
+ ).toBeInTheDocument();
+ });
+
it("sends every a2a field the user filled across all collapsible panels", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderForm();
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx
index fbf5cf8c1fb..5e8ba145396 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx
@@ -83,7 +83,9 @@ describe("AddAgentForm logos", () => {
expect(titleLogo).toBeInstanceOf(HTMLImageElement);
expect(titleLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
- const selectionLogo = within(await screen.findByRole("combobox")).getByAltText("A2A Agent logo");
+ const selectionLogo = within(await screen.findByRole("combobox", { name: "Agent Type" })).getByAltText(
+ "A2A Agent logo",
+ );
expect(selectionLogo).toBeInstanceOf(HTMLImageElement);
expect(selectionLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png"));
});
@@ -93,14 +95,14 @@ describe("AddAgentForm logos", () => {
await screen.findByAltText("A2A Agent logo");
- expect(screen.getByLabelText("Agent Type")).toBe(screen.getByRole("combobox"));
+ expect(screen.getByLabelText("Agent Type")).toBe(screen.getByRole("combobox", { name: "Agent Type" }));
});
it("renders the option logo when the agent type dropdown is opened", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderForm();
- const trigger = await screen.findByRole("combobox");
+ const trigger = await screen.findByRole("combobox", { name: "Agent Type" });
await within(trigger).findByAltText("A2A Agent logo");
await user.click(trigger);
@@ -123,7 +125,7 @@ describe("AddAgentForm logos", () => {
expect(screen.queryByAltText("Agent logo")).not.toBeInTheDocument();
expect(within(header).getByText("A")).toBeInTheDocument();
- const trigger = screen.getByRole("combobox");
+ const trigger = screen.getByRole("combobox", { name: "Agent Type" });
fireEvent.error(within(trigger).getByAltText("A2A Agent logo"));
expect(within(trigger).queryByAltText("A2A Agent logo")).not.toBeInTheDocument();
expect(warnSpy).toHaveBeenCalledTimes(2);
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx
index 5bd6ea9b83a..b243d9d1601 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx
@@ -1,3 +1,5 @@
+import { AgentIdentityFields } from "./AgentIdentityFields";
+import { withAgentIdentity } from "./agent_identity";
import React, { useState, useEffect } from "react";
import { FormProvider, useForm, useWatch } from "react-hook-form";
import { toast } from "@/lib/toast";
@@ -287,6 +289,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
const buildAgentData = (values: AgentFormValues): AgentRequestPayload | null => {
if (agentType === CUSTOM_AGENT_TYPE) {
+ if (values.identity_provider === "microsoft_entra") return { agent_name: values.agent_name };
return {
agent_name: values.agent_name,
agent_card_params: {
@@ -353,12 +356,13 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
return;
}
const values = form.getValues();
- const agentData = buildAgentData(values);
- if (!agentData) {
+ const built = buildAgentData(values);
+ if (!built) {
toast.error("Failed to build agent data");
setIsSubmitting(false);
return;
}
+ const agentData = withAgentIdentity(built, values);
// Build object_permission from MCP Tools step (allowed_mcp_servers_and_groups, mcp_tool_permissions)
const mcpServersAndGroups = values.allowed_mcp_servers_and_groups ?? {};
@@ -792,7 +796,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
- For agents that don't follow a standard protocol, just needs a virtual key
+ For outbound agents using an identity provider or virtual key
@@ -801,6 +805,8 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
+
+
{agentType === CUSTOM_AGENT_TYPE ? (
@@ -910,7 +916,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
name="team_id"
label={labelWithHint(
"Assign to Team",
- "Optionally assign this agent to a team. The agent and its key will belong to the selected team.",
+ "Optionally select a team for the virtual key. The agent identity and its permissions are managed separately.",
)}
>
{({ value, onChange }) => (
@@ -920,6 +926,11 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
+ {form.getValues("identity_provider") === "microsoft_entra" && (
+
+ This agent will authenticate with Microsoft Entra ID. You can skip virtual key creation.
+
+ )}
setKeyAssignOption(value as "create_new" | "existing_key" | "skip")}
@@ -1004,7 +1015,9 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
className="text-sm text-muted-foreground underline hover:text-foreground"
onClick={() => setKeyAssignOption("skip")}
>
- Skip for now — I'll assign a key later
+ {form.getValues("identity_provider") === "microsoft_entra"
+ ? "Use Entra JWT authentication"
+ : "Skip for now, I’ll assign a key later"}
@@ -1033,7 +1046,9 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
)}
{!createdKeyValue && !assignedKeyAlias && keyAssignOption === "skip" && (
- No key assigned. You can create one from the Virtual Keys page.
+ {form.getValues("identity_provider") === "microsoft_entra"
+ ? "Microsoft Entra ID is configured. Send an authenticated agent request to verify the connection."
+ : "No key assigned. You can create one from the Virtual Keys page."}
)}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts
index 16ce6848402..6ec3c3181f7 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts
@@ -1,3 +1,4 @@
+import { parseIdentityForForm } from "./agent_identity";
/**
* Shared configuration for agent form fields
* Used across create, view, and update operations
@@ -57,7 +58,7 @@ export const AGENT_FORM_CONFIG: {
name: "description",
label: "Description",
type: "textarea",
- required: true,
+ required: false,
placeholder: "Describe what this agent does...",
rows: 3,
},
@@ -340,6 +341,7 @@ export const parseAccessGroupIdsForForm = (agent: { access_group_ids?: string[]
});
export const parseMcpPermissionsForForm = (agent: any) => ({
+ ...parseIdentityForForm(agent),
allowed_mcp_servers_and_groups: {
servers: agent.object_permission?.mcp_servers ?? [],
accessGroups: agent.object_permission?.mcp_access_groups ?? [],
@@ -363,8 +365,9 @@ export const buildMcpObjectPermission = (values: any) => ({
* Parse agent data for form fields
*/
export const parseAgentForForm = (agent: any) => {
+ const card = agent.agent_card_params ?? {};
const skills =
- agent.agent_card_params?.skills?.map((skill: any) => ({
+ card.skills?.map((skill: any) => ({
...skill,
tags: skill.tags,
examples: skill.examples || [],
@@ -372,18 +375,18 @@ export const parseAgentForForm = (agent: any) => {
return {
agent_name: agent.agent_name,
- name: agent.agent_card_params?.name,
- description: agent.agent_card_params?.description,
- url: agent.agent_card_params?.url,
- version: agent.agent_card_params?.version,
- protocolVersion: agent.agent_card_params?.protocolVersion,
- streaming: agent.agent_card_params?.capabilities?.streaming,
- pushNotifications: agent.agent_card_params?.capabilities?.pushNotifications,
- stateTransitionHistory: agent.agent_card_params?.capabilities?.stateTransitionHistory,
+ name: card.name || agent.agent_name,
+ description: card.description,
+ url: card.url,
+ version: card.version,
+ protocolVersion: card.protocolVersion,
+ streaming: card.capabilities?.streaming,
+ pushNotifications: card.capabilities?.pushNotifications,
+ stateTransitionHistory: card.capabilities?.stateTransitionHistory,
skills: skills,
- iconUrl: agent.agent_card_params?.iconUrl,
- documentationUrl: agent.agent_card_params?.documentationUrl,
- supportsAuthenticatedExtendedCard: agent.agent_card_params?.supportsAuthenticatedExtendedCard,
+ iconUrl: card.iconUrl,
+ documentationUrl: card.documentationUrl,
+ supportsAuthenticatedExtendedCard: card.supportsAuthenticatedExtendedCard,
model: agent.litellm_params?.model,
make_public: agent.litellm_params?.make_public,
cost_per_query: agent.litellm_params?.cost_per_query,
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
new file mode 100644
index 00000000000..0639e7a6dd4
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts
@@ -0,0 +1,83 @@
+import { describe, expect, it } from "vitest";
+import {
+ buildIdentityParams,
+ entraTenantFromIssuer,
+ parseIdentityForForm,
+ readAgentIdentity,
+ withAgentIdentity,
+} from "./agent_identity";
+
+const identity = {
+ provider: "microsoft_entra",
+ tenant_id: "11111111-1111-4111-8111-111111111111",
+ client_id: "22222222-2222-4222-8222-222222222222",
+ service_principal_id: "33333333-3333-4333-8333-333333333333",
+ required_roles: ["Agent.Invoke"],
+ required_scopes: ["user_impersonation"],
+} satisfies import("./agent_identity").EntraAgentIdentity;
+
+describe("agent identity configuration", () => {
+ it("round trips an existing binding independently of the agent name and runtime", () => {
+ const values = {
+ ...parseIdentityForForm({
+ identity: { ...identity, agent_id: "stable", active: true, revision: "rev", issuer: "https://issuer.example" },
+ }),
+ agent_name: "Renamed",
+ url: "https://new-runtime.example",
+ };
+ expect(buildIdentityParams(values)).toEqual({ identity });
+ });
+ it("preserves untouched bindings and explicitly clears a removed binding", () => {
+ expect(buildIdentityParams({ agent_name: "legacy" })).toEqual({});
+ expect(buildIdentityParams({ identity_provider: "none" }, identity)).toEqual({ identity: null });
+ expect(parseIdentityForForm({}).identity_provider).toBe("none");
+ });
+ it.each([
+ null,
+ {},
+ "invalid",
+ { ...identity, client_id: "bad" },
+ { ...identity, tenant_id: 3 },
+ { ...identity, provider: "other" },
+ ])("rejects malformed bindings: %j", (value) => {
+ expect(readAgentIdentity(value)).toBeNull();
+ });
+ it("rejects incomplete submissions", () => {
+ expect(() => buildIdentityParams({ identity_provider: "microsoft_entra" })).toThrow("Enter valid Entra");
+ });
+ it("submits identity as top-level settings without changing runtime parameters", () => {
+ const formValues = {
+ identity_provider: "microsoft_entra",
+ identity_tenant_id: identity.tenant_id,
+ identity_client_id: identity.client_id,
+ identity_service_principal_id: identity.service_principal_id,
+ execution_mode: "both",
+ enabled: false,
+ };
+ const payload = withAgentIdentity({ litellm_params: { model: "runtime" } }, formValues);
+ expect(payload.litellm_params).toEqual({ model: "runtime" });
+ expect(payload.identity).toMatchObject({
+ client_id: identity.client_id,
+ service_principal_id: identity.service_principal_id,
+ });
+ expect(payload.execution_mode).toBe("both");
+ expect(payload.enabled).toBe(false);
+ });
+ it("requires a service principal for autonomous execution", () => {
+ const values = {
+ identity_provider: "microsoft_entra",
+ identity_tenant_id: identity.tenant_id,
+ identity_client_id: identity.client_id,
+ execution_mode: "autonomous",
+ };
+ expect(() => buildIdentityParams(values)).toThrow("Enterprise application Object ID");
+ });
+
+ it("only offers tenant-specific Microsoft issuers", () => {
+ expect(entraTenantFromIssuer(`https://login.microsoftonline.com/${identity.tenant_id}/v2.0`)).toBe(
+ identity.tenant_id,
+ );
+ expect(entraTenantFromIssuer("https://attacker.example/tenant/v2.0")).toBeNull();
+ expect(entraTenantFromIssuer("https://login.microsoftonline.com/common/v2.0")).toBeNull();
+ });
+});
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
new file mode 100644
index 00000000000..34986f75679
--- /dev/null
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts
@@ -0,0 +1,101 @@
+import { z } from "zod";
+import type { components } from "@/lib/http/schema";
+import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit";
+
+export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"];
+type AgentIdentityState = Pick;
+
+export const IDENTITY_UUID_PATTERN = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i;
+
+const stringGrants = (fallback: string[]) =>
+ z
+ .unknown()
+ .transform((value) =>
+ Array.isArray(value) ? value.filter((entry): entry is string => typeof entry === "string") : fallback,
+ );
+
+const identityShape = {
+ provider: z.literal("microsoft_entra"),
+ tenant_id: z.string().regex(IDENTITY_UUID_PATTERN),
+ client_id: z.string().regex(IDENTITY_UUID_PATTERN),
+ service_principal_id: z.string().regex(IDENTITY_UUID_PATTERN).nullable().default(null),
+ required_roles: stringGrants([]),
+ required_scopes: stringGrants(["user_impersonation"]),
+};
+const identitySchema = z.object(identityShape);
+
+export const readAgentIdentity = (value: unknown): EntraAgentIdentity | null => {
+ const parsed = identitySchema.safeParse(value);
+ return parsed.success ? parsed.data : null;
+};
+
+const identityFormFields = (identity: EntraAgentIdentity | null): AgentFormValues => ({
+ identity_provider: identity?.provider ?? "none",
+ identity_tenant_id: identity?.tenant_id ?? "",
+ identity_client_id: identity?.client_id ?? "",
+ identity_service_principal_id: identity?.service_principal_id ?? "",
+ identity_required_roles: identity?.required_roles?.join(", ") ?? "",
+ identity_required_scopes: identity?.required_scopes?.join(", ") ?? "user_impersonation",
+});
+
+export const parseIdentityForForm = (agent?: Partial | null): AgentFormValues => {
+ const identity = agent?.identity?.active === false ? null : readAgentIdentity(agent?.identity);
+ return {
+ ...identityFormFields(identity),
+ execution_mode: agent?.execution_mode ?? "autonomous",
+ enabled: agent?.enabled ?? true,
+ };
+};
+
+const splitGrants = (value: unknown, fallback: string[]): string[] =>
+ typeof value === "string"
+ ? value
+ .split(",")
+ .map((item) => item.trim())
+ .filter(Boolean)
+ : fallback;
+
+export const buildIdentityParams = (
+ values: AgentFormValues,
+ existingIdentity?: unknown,
+): { identity?: EntraAgentIdentity | null } => {
+ if (values.identity_provider === undefined) return {};
+ if (values.identity_provider !== "microsoft_entra")
+ return readAgentIdentity(existingIdentity) ? { identity: null } : {};
+ const candidate: EntraAgentIdentity = {
+ provider: "microsoft_entra",
+ tenant_id: typeof values.identity_tenant_id === "string" ? values.identity_tenant_id.trim().toLowerCase() : "",
+ client_id: typeof values.identity_client_id === "string" ? values.identity_client_id.trim().toLowerCase() : "",
+ service_principal_id:
+ typeof values.identity_service_principal_id === "string" && values.identity_service_principal_id.trim()
+ ? values.identity_service_principal_id.trim().toLowerCase()
+ : null,
+ required_roles: splitGrants(values.identity_required_roles, []),
+ required_scopes: splitGrants(values.identity_required_scopes, ["user_impersonation"]),
+ };
+ const identity = readAgentIdentity(candidate);
+ if (!identity) throw new Error("Enter valid Entra tenant, application client and service principal IDs");
+ if (values.execution_mode !== "delegated" && !identity.service_principal_id)
+ throw new Error("Autonomous agents require the Enterprise application Object ID");
+ return { identity };
+};
+
+export const entraTenantFromIssuer = (issuer: string): string | null => {
+ const match = /^https:\/\/login\.microsoftonline\.com\/([^/]+)\/v2\.0$/.exec(issuer);
+ return match && IDENTITY_UUID_PATTERN.test(match[1]) ? match[1] : null;
+};
+
+export const withAgentIdentity = (
+ payload: AgentRequestPayload,
+ values: AgentFormValues,
+ existing?: Partial,
+): AgentRequestPayload => {
+ const identityFields = buildIdentityParams(values, existing?.identity);
+ const managed = values.identity_provider === "microsoft_entra" || Boolean(readAgentIdentity(existing?.identity));
+ return {
+ ...payload,
+ ...identityFields,
+ ...(managed && values.execution_mode !== undefined ? { execution_mode: values.execution_mode } : {}),
+ ...(managed && values.enabled !== undefined ? { enabled: values.enabled } : {}),
+ };
+};
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
index 37e00766a75..5f5eb4e020c 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx
@@ -8,6 +8,7 @@ import * as networking from "@/components/networking";
import type { AgentCreateInfo } from "@/components/networking";
vi.mock("@/components/networking", () => ({
+ apiClient: { get: vi.fn() },
getAgentInfo: vi.fn(),
patchAgentCall: vi.fn(),
getAgentCreateMetadata: vi.fn(),
@@ -155,6 +156,50 @@ describe("AgentInfoView update payload", () => {
.mockResolvedValue({} as never);
});
+ it.each(["complete", "empty"])("preserves the Entra binding while renaming an agent with a %s card", async (card) => {
+ const user = setup();
+ const identity = {
+ provider: "microsoft_entra",
+ tenant_id: "11111111-1111-4111-8111-111111111111",
+ client_id: "22222222-2222-4222-8222-222222222222",
+ service_principal_id: "33333333-3333-4333-8333-333333333333",
+ };
+ const params = { ...A2A_AGENT.litellm_params, require_trace_id_on_calls_by_agent: true };
+ vi.mocked(networking.getAgentInfo).mockResolvedValue({
+ ...A2A_AGENT,
+ agent_card_params: card === "empty" ? {} : A2A_AGENT.agent_card_params,
+ litellm_params: params,
+ identity: { ...identity, agent_id: "agent-1", issuer: "https://issuer.example", revision: "rev", active: true },
+ identity_managed: true,
+ execution_mode: "autonomous",
+ enabled: true,
+ access_group_ids: ["ag-entra"],
+ } as never);
+ vi.mocked(networking.apiClient.get).mockImplementation(async (path) =>
+ path.endsWith("/providers")
+ ? [`https://login.microsoftonline.com/${identity.tenant_id}/v2.0`]
+ : { last_authenticated_at: null },
+ );
+ renderView();
+ expect(await screen.findByText("Configured, awaiting an authenticated request")).toBeInTheDocument();
+ await openEditor(user);
+ expect(screen.getByLabelText("Application (Client) ID")).toHaveValue(identity.client_id);
+ expect(screen.getByRole("combobox", { name: "Identity Provider" })).toHaveTextContent("Microsoft Entra ID");
+ expect(screen.getByRole("combobox", { name: "Execution Mode" })).toHaveTextContent("Autonomous");
+ expect(screen.getByRole("combobox", { name: /^Execution$/ })).toHaveTextContent("Enabled");
+ fireEvent.change(screen.getByLabelText("Agent Name"), { target: { value: "Renamed agent" } });
+ await save(user);
+ expect(patchedPayload().agent_name).toBe("Renamed agent");
+ expect(patchedPayload()).not.toHaveProperty("litellm_params");
+ expect(patchedPayload().identity).toMatchObject(identity);
+ expect(patchedPayload().access_group_ids).toEqual(["ag-entra"]);
+ expect(networking.patchAgentCall).toHaveBeenCalledWith(
+ "tok",
+ "agent-1",
+ expect.objectContaining({ agent_name: "Renamed agent" }),
+ );
+ });
+
it("sends only the fields whose panel has been opened, dropping the rest", async () => {
const user = setup();
renderView();
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx
index 19b1ee8ca48..4eb3534ebe3 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx
@@ -2,6 +2,7 @@ import React from "react";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import AgentInfoView from "./agent_info";
+import AgentFormFields from "./agent_form_fields";
import * as networking from "@/components/networking";
import type { Agent } from "@/components/agents/types";
@@ -16,12 +17,16 @@ vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({
useKeys: () => ({ data: { keys: [] }, isLoading: false, refetch: vi.fn() }),
}));
+vi.mock("./AgentIdentityDetails", () => ({
+ AgentIdentityDetails: () => null,
+}));
+
vi.mock("./agent_card_discovery", () => ({
default: () =>
,
}));
vi.mock("./agent_form_fields", () => ({
- default: () =>
,
+ default: vi.fn(() =>
),
unmountedA2AFieldNames: () => [],
}));
@@ -77,6 +82,9 @@ const agent = {
describe("AgentInfoView settings", () => {
beforeEach(() => {
vi.restoreAllMocks();
+ vi.mocked(AgentFormFields)
+ .mockReset()
+ .mockImplementation(() =>
);
vi.mocked(networking.getAgentInfo).mockReset().mockResolvedValue(agent);
vi.mocked(networking.getAgentCreateMetadata).mockReset().mockResolvedValue([]);
vi.mocked(networking.patchAgentCall).mockReset().mockResolvedValue({});
@@ -104,6 +112,23 @@ describe("AgentInfoView settings", () => {
expect(payload.access_group_ids).toEqual([]);
});
+ it("saves unrelated settings when the existing card has no description", async () => {
+ const actual = await vi.importActual("./agent_form_fields");
+ vi.mocked(AgentFormFields).mockImplementation(actual.default);
+ const { description: _description, ...card } = agent.agent_card_params ?? {};
+ vi.mocked(networking.getAgentInfo).mockResolvedValue({ ...agent, agent_card_params: card });
+ render( );
+ fireEvent.click(await screen.findByRole("tab", { name: "Settings" }));
+ fireEvent.click(screen.getByRole("button", { name: "Edit Settings" }));
+ expect(await screen.findByLabelText("Description")).toHaveValue("");
+ fireEvent.change(screen.getByLabelText("TPM Limit"), { target: { value: "42" } });
+ fireEvent.click(screen.getByRole("button", { name: /Save Changes/ }));
+ await waitFor(() => expect(networking.patchAgentCall).toHaveBeenCalledOnce());
+ const [, , payload] = vi.mocked(networking.patchAgentCall).mock.calls[0];
+ expect(payload.tpm_limit).toBe(42);
+ expect(payload.agent_card_params?.description).toBe("");
+ });
+
it("sends the newly attached access group in the update payload", async () => {
render( );
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
index 6e7389a3fa1..d4c05d0ebac 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx
@@ -1,3 +1,6 @@
+import { AgentIdentityFields } from "./AgentIdentityFields";
+import { AgentIdentityDetails } from "./AgentIdentityDetails";
+import { withAgentIdentity } from "./agent_identity";
import React, { useState, useEffect, useMemo } from "react";
import { cx } from "@/lib/cva.config";
import { FormProvider, useForm, useWatch } from "react-hook-form";
@@ -237,7 +240,7 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT
: built;
await patchAgentCall(accessToken, agentId, {
- ...updateData,
+ ...withAgentIdentity(updateData, values, agent),
object_permission: buildMcpObjectPermission(values),
access_group_ids: values.access_group_ids ?? [],
});
@@ -337,6 +340,12 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT
{/* Overview Panel */}
+
{agent.agent_id}
{agent.agent_name}
@@ -505,6 +514,8 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT
)}
+
+
{discoveryRequest && (
{
expect(screen.getByRole("option", { name: /group-b/ })).toBeInTheDocument();
});
+ it("offers only individual agents when legacy groups are disabled", async () => {
+ const user = userEvent.setup();
+ const onChange = vi.fn();
+ render( );
+ await user.click(screen.getByRole("combobox"));
+ await user.click(await screen.findByRole("option", { name: /Agent One/ }));
+ expect(screen.queryByRole("option", { name: /group-a/ })).not.toBeInTheDocument();
+ expect(onChange).toHaveBeenCalledWith({ agents: ["agent-1"], accessGroups: [] });
+ });
+
+ it("preserves a saved legacy group while adding an agent with legacy groups disabled", async () => {
+ const user = userEvent.setup();
+ const onChange = vi.fn();
+ render(
+ ,
+ );
+ await user.click(screen.getByRole("combobox"));
+ expect(await screen.findByRole("option", { name: /retired-group/ })).toHaveTextContent(
+ "Existing legacy agent group",
+ );
+ await user.click(await screen.findByRole("option", { name: /Agent One/ }));
+ expect(onChange).toHaveBeenCalledWith({ agents: ["agent-1"], accessGroups: ["retired-group"] });
+ });
+
it("respects disabled prop", () => {
render( );
expect(screen.getByRole("combobox")).toBeDisabled();
diff --git a/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx b/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx
index d26772d5c9a..7215bcf25e8 100644
--- a/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx
+++ b/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx
@@ -19,6 +19,7 @@ interface AgentSelectorProps {
accessToken: string;
placeholder?: string;
disabled?: boolean;
+ allowAccessGroups?: boolean;
}
const AgentSelector: React.FC = ({
@@ -28,6 +29,7 @@ const AgentSelector: React.FC = ({
accessToken,
placeholder = "Select agents",
disabled = false,
+ allowAccessGroups = true,
}) => {
const [agents, setAgents] = useState([]);
const [accessGroups, setAccessGroups] = useState([]);
@@ -60,12 +62,15 @@ const AgentSelector: React.FC = ({
fetchData();
}, [accessToken]);
- // Combine options, access groups first
+ const selectableGroups = allowAccessGroups
+ ? Array.from(new Set([...accessGroups, ...(value?.accessGroups ?? [])]))
+ : value?.accessGroups ?? [];
+
const options: MultiSelectOption[] = [
- ...accessGroups.map((group) => ({
+ ...selectableGroups.map((group) => ({
label: group,
value: `group:${group}`,
- description: "Access Group",
+ description: allowAccessGroups ? "Access Group" : "Existing legacy agent group",
})),
...agents.map((agent) => ({
label: `${agent.agent_name || agent.agent_id}`,
diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts
index 469ecf36f06..92d946c19e1 100644
--- a/ui/litellm-dashboard/src/components/agents/types.ts
+++ b/ui/litellm-dashboard/src/components/agents/types.ts
@@ -11,6 +11,11 @@ export type AgentKillSwitchConfig = components["schemas"]["AgentKillSwitchConfig
export type AgentKillSwitchResult = components["schemas"]["AgentKillSwitchResult"];
export interface Agent {
+ identity?: components["schemas"]["AgentIdentityBinding"] | null;
+ identity_managed?: boolean;
+ enabled?: boolean;
+ execution_mode?: components["schemas"]["AgentResponse"]["execution_mode"];
+ jwt_auth_configured?: boolean;
agent_id: string;
agent_name: string;
litellm_params: {
diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx
index 47fb5643b22..66748d3cae7 100644
--- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx
+++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx
@@ -170,3 +170,41 @@ describe("MCPServerSelector all-proxy-mcpservers option", () => {
expect(optionByLabel("Server One")).toHaveAttribute("aria-disabled", "true");
});
});
+
+describe("MCPServerSelector unified group flow", () => {
+ beforeEach(() => {
+ vi.clearAllMocks();
+ setupMcpMocks();
+ mockUseMCPAccessGroups.mockReturnValue({ data: ["legacy-group"], isLoading: false } as ReturnType<
+ typeof useMCPAccessGroups
+ >);
+ });
+
+ it("offers servers without legacy groups when disabled", async () => {
+ const user = userEvent.setup();
+ const onChange = vi.fn();
+ renderWithProviders( );
+ await openSelector(user);
+ expect(optionByLabel("legacy-group")).toBeUndefined();
+ await user.click(optionByLabel("Server One")!);
+ expect(onChange).toHaveBeenCalledWith({ servers: ["srv-1"], accessGroups: [], toolsets: [] });
+ });
+
+ it("preserves a saved legacy group even if discovery no longer returns it", async () => {
+ const user = userEvent.setup();
+ const onChange = vi.fn();
+ renderWithProviders(
+ ,
+ );
+ await openSelector(user);
+ expect(optionByLabel("retired-group")).toHaveTextContent("Existing legacy MCP group");
+ expect(optionByLabel("legacy-group")).toBeUndefined();
+ await user.click(optionByLabel("Server One")!);
+ expect(onChange).toHaveBeenCalledWith({ servers: ["srv-1"], accessGroups: ["retired-group"], toolsets: [] });
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx
index 969db57462f..fe23fd87b76 100644
--- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx
+++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx
@@ -18,11 +18,15 @@ interface MCPServerSelectorProps {
disabled?: boolean;
teamId?: string | null;
allowNoMcpServers?: boolean;
+ allowAccessGroups?: boolean;
allowAllProxyMcpServers?: boolean;
}
const TOOLSET_PREFIX = "toolset:";
+const selectableLegacyGroups = (available: string[], selected: string[] = [], allowNew: boolean): string[] =>
+ allowNew ? Array.from(new Set([...available, ...selected])) : selected;
+
const MCPServerSelector: React.FC = ({
onChange,
value,
@@ -32,22 +36,24 @@ const MCPServerSelector: React.FC = ({
disabled = false,
teamId,
allowNoMcpServers = false,
+ allowAccessGroups = true,
allowAllProxyMcpServers = false,
}) => {
const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(teamId);
const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups();
const { data: toolsets = [], isLoading: toolsetsLoading } = useMCPToolsets();
- const loading = serversLoading || groupsLoading || toolsetsLoading;
+ const loading = [serversLoading, groupsLoading, toolsetsLoading].some(Boolean);
- const accessGroupSet = new Set(accessGroups);
+ const selectableGroups = selectableLegacyGroups(accessGroups, value?.accessGroups, allowAccessGroups);
+ const accessGroupSet = new Set(selectableGroups);
// Combine options: access groups + servers + toolsets
const options = [
- ...accessGroups.map((group) => ({
+ ...selectableGroups.map((group) => ({
label: group,
value: group,
- description: "Access Group",
+ description: allowAccessGroups ? "Access Group" : "Existing legacy MCP group",
})),
...mcpServers.map((server) => ({
label: `${server.server_name || server.server_id} (${server.server_id})`,
diff --git a/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx b/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx
index d1ca25c975b..0d5525d7d78 100644
--- a/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx
+++ b/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx
@@ -67,7 +67,7 @@ export function AgentPermissions({
-
Agents
+
Allowed agents to call
{totalCount}
diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx
index a693ee971d4..bbabcf2458e 100644
--- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx
+++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx
@@ -2249,9 +2249,9 @@ describe("TeamInfoView - the exact bytes the update call sends", () => {
await openEditorWithAgents(user);
await user.click(within(screen.getByLabelText("agent-1")).getByRole("button"));
- await user.click(within(screen.getByLabelText("group:group-a")).getByRole("button"));
+ await user.click(within(screen.getByLabelText("group-a")).getByRole("button"));
expect(screen.queryByLabelText("agent-1")).not.toBeInTheDocument();
- expect(screen.queryByLabelText("group:group-a")).not.toBeInTheDocument();
+ expect(screen.queryByLabelText("group-a")).not.toBeInTheDocument();
const payload = await save(user);
diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx
index 3845f94593d..e03a08dfabd 100644
--- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx
+++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx
@@ -2039,13 +2039,14 @@ const TeamInfoView: React.FC
= ({
)}
-
+
{({ value, onChange }) => (
)}
@@ -2062,13 +2063,14 @@ const TeamInfoView: React.FC = ({
/>
-
+
{({ value, onChange }) => (
)}
diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx
index 4d451a9c7b9..4d980d1b018 100644
--- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx
@@ -375,3 +375,17 @@ describe("TTFT column", () => {
expect(screen.getByText("1.00")).toBeInTheDocument();
});
});
+
+describe("Request outcome", () => {
+ it("shows a failed agent outcome even when metadata has no status", () => {
+ renderRows([logEntry({ call_type: "asend_message", status: "failure", session_total_count: 4 })]);
+ expect(screen.getByText("Failure")).toBeInTheDocument();
+ expect(screen.queryByText("Success")).not.toBeInTheDocument();
+ });
+
+ it("prefers the recorded outcome over stale metadata", () => {
+ renderRows([logEntry({ status: "success", metadata: { status: "failure" } })]);
+ expect(screen.getByText("Success")).toBeInTheDocument();
+ expect(screen.queryByText("Failure")).not.toBeInTheDocument();
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx
index 5705f41f3de..6b35ceb81a2 100644
--- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx
@@ -120,7 +120,7 @@ export const getRequestLogsTableColumns = ({
enableSorting: false,
meta: { skeleton: "badge" },
cell: ({ row }) => {
- const status = readMetaString(row.original.metadata, "status") ?? "Success";
+ const status = row.original.status || readMetaString(row.original.metadata, "status") || "Success";
const isSuccess = status.toLowerCase() !== "failure";
const batchCounts = isSuccess ? getBatchRequestCounts(row.original.metadata) : undefined;
if (batchCounts && batchCounts.failed > 0) {
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index b5d515bd1fb..51d0165330b 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -18071,6 +18071,23 @@ export interface paths {
patch?: never;
trace?: never;
};
+ "/v1/agents/identity/providers": {
+ parameters: {
+ query?: never;
+ header?: never;
+ path?: never;
+ cookie?: never;
+ };
+ /** Get Agent Identity Providers */
+ get: operations["get_agent_identity_providers_v1_agents_identity_providers_get"];
+ put?: never;
+ post?: never;
+ delete?: never;
+ options?: never;
+ head?: never;
+ patch?: never;
+ trace?: never;
+ };
"/v1/agents/make_public": {
parameters: {
query?: never;
@@ -18220,6 +18237,23 @@ export interface paths {
patch: operations["patch_agent_v1_agents__agent_id__patch"];
trace?: never;
};
+ "/v1/agents/{agent_id}/identity": {
+ parameters: {
+ query?: never;
+ header?: never;
+ path?: never;
+ cookie?: never;
+ };
+ /** Get Agent Identity Status */
+ get: operations["get_agent_identity_status_v1_agents__agent_id__identity_get"];
+ put?: never;
+ post?: never;
+ delete?: never;
+ options?: never;
+ head?: never;
+ patch?: never;
+ trace?: never;
+ };
"/v1/agents/{agent_id}/kill_switch": {
parameters: {
query?: never;
@@ -24206,11 +24240,19 @@ export interface components {
AgentConfig: {
/** Access Group Ids */
access_group_ids?: string[] | null;
- agent_card_params: components["schemas"]["AgentCard"];
+ agent_card_params?: components["schemas"]["AgentCard"];
/** Agent Name */
agent_name: string;
+ /** Enabled */
+ enabled?: boolean;
+ /**
+ * Execution Mode
+ * @enum {string}
+ */
+ execution_mode?: "autonomous" | "delegated" | "both";
/** Extra Headers */
extra_headers?: string[] | null;
+ identity?: components["schemas"]["EntraIdentityConfig"] | null;
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
/** Litellm Params */
litellm_params?: {
@@ -24297,6 +24339,45 @@ export interface components {
/** Uri */
uri?: string;
};
+ /** AgentIdentityBinding */
+ AgentIdentityBinding: {
+ /**
+ * Active
+ * @default true
+ */
+ active: boolean;
+ /** Agent Id */
+ agent_id: string;
+ /** Client Id */
+ client_id: string;
+ /** Issuer */
+ issuer: string;
+ /** Last Authenticated At */
+ last_authenticated_at?: string | null;
+ /**
+ * Provider
+ * @constant
+ */
+ provider: "microsoft_entra";
+ /**
+ * Required Roles
+ * @default []
+ */
+ required_roles: string[];
+ /**
+ * Required Scopes
+ * @default [
+ * "user_impersonation"
+ * ]
+ */
+ required_scopes: string[];
+ /** Revision */
+ revision: string;
+ /** Service Principal Id */
+ service_principal_id?: string | null;
+ /** Tenant Id */
+ tenant_id: string;
+ };
/**
* AgentInterface
* @description Declares a combination of a target URL and a transport protocol.
@@ -24452,8 +24533,30 @@ export interface components {
created_at?: string | null;
/** Created By */
created_by?: string | null;
+ /**
+ * Enabled
+ * @default true
+ */
+ enabled: boolean;
+ /**
+ * Execution Mode
+ * @default autonomous
+ * @enum {string}
+ */
+ execution_mode: "autonomous" | "delegated" | "both";
/** Extra Headers */
extra_headers?: string[] | null;
+ identity?: components["schemas"]["AgentIdentityBinding"] | null;
+ /**
+ * Identity Managed
+ * @default false
+ */
+ identity_managed: boolean;
+ /**
+ * Jwt Auth Configured
+ * @default false
+ */
+ jwt_auth_configured: boolean;
/** Keys */
keys?: components["schemas"]["AgentKeySummary"][] | null;
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
@@ -30137,6 +30240,33 @@ export interface components {
/** Template Id */
template_id: string;
};
+ /** EntraIdentityConfig */
+ EntraIdentityConfig: {
+ /** Client Id */
+ client_id: string;
+ /**
+ * Provider
+ * @constant
+ */
+ provider: "microsoft_entra";
+ /**
+ * Required Roles
+ * @default []
+ */
+ required_roles: string[];
+ /**
+ * Required Scopes
+ * @description Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.
+ * @default [
+ * "user_impersonation"
+ * ]
+ */
+ required_scopes: string[];
+ /** Service Principal Id */
+ service_principal_id?: string | null;
+ /** Tenant Id */
+ tenant_id: string;
+ };
/** EnvironmentReport */
EnvironmentReport: {
/** Config Lines */
@@ -35544,6 +35674,28 @@ export interface components {
/** Mcp Server Ids */
mcp_server_ids: string[];
};
+ /** ManagedAgentIdentityStatus */
+ ManagedAgentIdentityStatus: {
+ /**
+ * Enabled
+ * @default true
+ */
+ enabled: boolean;
+ /**
+ * Execution Mode
+ * @default autonomous
+ * @enum {string}
+ */
+ execution_mode: "autonomous" | "delegated" | "both";
+ identity?: components["schemas"]["AgentIdentityBinding"] | null;
+ /**
+ * Identity Managed
+ * @default false
+ */
+ identity_managed: boolean;
+ /** Last Authenticated At */
+ last_authenticated_at?: string | null;
+ };
/**
* Mcp
* @description Give the model access to additional tools via remote Model Context Protocol
@@ -37618,8 +37770,16 @@ export interface components {
agent_card_params?: components["schemas"]["AgentCard"];
/** Agent Name */
agent_name?: string;
+ /** Enabled */
+ enabled?: boolean;
+ /**
+ * Execution Mode
+ * @enum {string}
+ */
+ execution_mode?: "autonomous" | "delegated" | "both";
/** Extra Headers */
extra_headers?: string[] | null;
+ identity?: components["schemas"]["EntraIdentityConfig"] | null;
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
/** Litellm Params */
litellm_params?: {
@@ -70077,6 +70237,26 @@ export interface operations {
};
};
};
+ get_agent_identity_providers_v1_agents_identity_providers_get: {
+ parameters: {
+ query?: never;
+ header?: never;
+ path?: never;
+ cookie?: never;
+ };
+ requestBody?: never;
+ responses: {
+ /** @description Successful Response */
+ 200: {
+ headers: {
+ [name: string]: unknown;
+ };
+ content: {
+ "application/json": string[];
+ };
+ };
+ };
+ };
make_agents_public_v1_agents_make_public_post: {
parameters: {
query?: never;
@@ -70242,6 +70422,37 @@ export interface operations {
};
};
};
+ get_agent_identity_status_v1_agents__agent_id__identity_get: {
+ parameters: {
+ query?: never;
+ header?: never;
+ path: {
+ agent_id: string;
+ };
+ cookie?: never;
+ };
+ requestBody?: never;
+ responses: {
+ /** @description Successful Response */
+ 200: {
+ headers: {
+ [name: string]: unknown;
+ };
+ content: {
+ "application/json": components["schemas"]["ManagedAgentIdentityStatus"];
+ };
+ };
+ /** @description Validation Error */
+ 422: {
+ headers: {
+ [name: string]: unknown;
+ };
+ content: {
+ "application/json": components["schemas"]["HTTPValidationError"];
+ };
+ };
+ };
+ };
trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post: {
parameters: {
query?: never;