feat(agents): integrate Entra identity registration and authorization
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

This commit is contained in:
Joshua Valluru 2026-09-26 17:12:00 -07:00
parent 5d777c16d9
commit 320a688bae
99 changed files with 6995 additions and 610 deletions

View file

@ -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 $$;

View file

@ -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?

View file

@ -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")
)

View file

@ -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")

View file

@ -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")

View file

@ -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(

View file

@ -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,

View file

@ -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

View file

@ -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 <your_api_key>\" \\\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 <your_api_key>\"\n```",

View file

@ -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]

View file

@ -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,

View file

@ -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,

View file

@ -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:

View file

@ -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 ()

View file

@ -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))

View file

@ -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

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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,

View file

@ -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 []

View file

@ -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?

View file

@ -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

View file

@ -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(

View file

@ -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]:

View file

@ -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"

View file

@ -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"]:

View file

@ -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]:

View file

@ -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

View file

@ -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

View file

@ -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?

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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)

View file

@ -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]

View file

@ -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

View file

@ -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()

View file

@ -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

View file

@ -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"

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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)

View file

@ -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"

View file

@ -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

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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"] == {}

View file

@ -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,

View file

@ -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"])
)

View file

@ -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 = {}

View file

@ -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:

View file

@ -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"

View file

@ -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"]

View file

@ -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

View file

@ -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)

View file

@ -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(<AgentIdentityDetails agentId="native" identity={identity} accessToken="admin" isAdmin />);
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(
<AgentIdentityDetails agentId="native" identity={identity} accessToken="user" isAdmin={false} />,
);
expect(screen.queryByRole("region", { name: "Agent Identity" })).not.toBeInTheDocument();
expect(apiClient.get).not.toHaveBeenCalled();
});
});

View file

@ -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<components["schemas"]["ManagedAgentIdentityStatus"]>(
`/v1/agents/${encodeURIComponent(agentId)}/identity`,
{
accessToken: accessToken ?? "",
},
),
enabled: Boolean(isAdmin && accessToken && identity),
});
if (!identity || !isAdmin) return null;
const executionLabel = data?.enabled ? "Enabled" : "Disabled";
return (
<section aria-label="Agent Identity" className="mb-6 space-y-2 rounded-lg border border-border p-4">
<h3 className="font-medium">Agent Identity: Microsoft Entra ID</h3>
<p className="text-sm">
Tenant: <span className="font-mono">{identity.tenant_id}</span>
</p>
<>
<p className="text-sm">
Application (Client) ID: <span className="font-mono">{identity.client_id}</span>
</p>
<p className="text-sm">Enterprise application Object ID: {identity.service_principal_id || "Not configured"}</p>
</>
<p className="text-sm">
Execution: {data ? executionLabel : "Loading"} · Mode: {data?.execution_mode ?? "Loading"}
</p>
<p className="text-sm">
{data?.identity?.active === false
? "Identity unbound; execution is disabled"
: authenticationMessage(isError, data?.last_authenticated_at)}
</p>
<p className="text-xs text-muted-foreground">
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.
</p>
<div className="flex items-center gap-4">
<Button
variant="outline"
size="sm"
disabled={isFetching}
onClick={() => {
void refetch();
}}
>
Refresh authentication evidence
</Button>
<a className="text-sm underline" href="/ui/logs/">
View request logs
</a>
</div>
</section>
);
};

View file

@ -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<AgentFormValues>({ name: "identity_provider" });
const mode = useWatch<AgentFormValues>({ name: "execution_mode" });
const showScopes = mode !== "autonomous" && mode !== undefined;
const [tenants, setTenants] = useState<string[]>([]);
const [error, setError] = useState<string | null>(null);
useEffect(() => {
if (!accessToken || provider !== "microsoft_entra") return;
let active = true;
apiClient
.get<string[]>("/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 (
<>
<section aria-label="Agent Identity" className="my-6 space-y-4 rounded-lg border border-border p-4">
<div>
<h3 className="font-medium">Agent Identity</h3>
<p className="mt-1 text-sm text-muted-foreground">
Connect an existing identity provider application to this agent. Its name and runtime address can change
independently.
</p>
</div>
<AgentFormField name="identity_provider" label="Identity Provider" defaultValue="none">
{({ value, onChange, id }) => (
<Select
items={PROVIDER_OPTIONS}
value={typeof value === "string" ? value : "none"}
onValueChange={onChange}
>
<SelectTrigger id={id}>
<SelectValue />
</SelectTrigger>
<SelectContent>
{PROVIDER_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</AgentFormField>
{provider === "microsoft_entra" && (
<>
<AgentFormField
name="identity_tenant_id"
label="Trusted Entra Tenant"
rules={{ required: "Select a trusted tenant" }}
>
{({ value, onChange, id }) => (
<Select value={typeof value === "string" ? value : ""} onValueChange={onChange}>
<SelectTrigger id={id}>
<SelectValue placeholder="Select the gateway's trusted tenant" />
</SelectTrigger>
<SelectContent>
{tenants.map((tenant) => (
<SelectItem key={tenant} value={tenant}>
{tenant}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</AgentFormField>
{error && (
<p role="alert" className="text-sm text-destructive">
{error}
</p>
)}
{!error && tenants.length === 0 && (
<p className="text-sm text-muted-foreground">
No trusted Entra tenant is available. Configure JWT issuer and audience validation on the gateway first.
Dashboard Microsoft SSO is configured separately.
</p>
)}
<AgentFormField
name="identity_client_id"
label="Application (Client) ID"
rules={{
required: "Enter the Entra application client ID",
pattern: { value: IDENTITY_UUID_PATTERN, message: "Enter a valid application client UUID" },
}}
description={
<>
Find this under{" "}
<a className="underline" href="https://entra.microsoft.com/" target="_blank" rel="noreferrer">
Entra App registrations
</a>
, select your agent application, then Overview. No client secret is required here.
</>
}
>
{({ value, onChange, ref, ...control }) => (
<Input
{...control}
ref={ref}
value={typeof value === "string" ? value : ""}
onChange={onChange}
placeholder="xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx"
/>
)}
</AgentFormField>
<AgentFormField name="execution_mode" label="Execution Mode" defaultValue="autonomous">
{({ value, onChange, id }) => (
<Select
items={EXECUTION_MODE_OPTIONS}
value={typeof value === "string" ? value : "autonomous"}
onValueChange={onChange}
>
<SelectTrigger id={id}>
<SelectValue />
</SelectTrigger>
<SelectContent>
{EXECUTION_MODE_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</AgentFormField>
<AgentFormField
name="identity_service_principal_id"
label="Enterprise Application Object ID"
rules={{
required: mode !== "delegated" ? "Enter the service principal Object ID" : false,
pattern: { value: IDENTITY_UUID_PATTERN, message: "Enter a valid service principal UUID" },
}}
description={
<>
Open{" "}
<a
className="underline"
href="https://entra.microsoft.com/#view/Microsoft_AAD_IAM/StartboardApplicationsMenuBlade/~/AppAppsPreview"
target="_blank"
rel="noreferrer"
>
Entra Enterprise applications
</a>
, select this application, and copy its Object ID. The App registrations Object ID is a different
value.
</>
}
>
{({ value, onChange, ref, ...control }) => (
<Input {...control} ref={ref} value={typeof value === "string" ? value : ""} onChange={onChange} />
)}
</AgentFormField>
<AgentFormField
name="identity_required_roles"
label="Required Application Roles"
description="Comma-separated role values required on autonomous application tokens"
>
{({ value, onChange, ref, ...control }) => (
<Input
{...control}
ref={ref}
value={typeof value === "string" ? value : ""}
onChange={onChange}
placeholder="Agent.Invoke"
/>
)}
</AgentFormField>
{showScopes && (
<>
<AgentFormField
name="identity_required_scopes"
label="Required Delegated Scopes"
defaultValue="user_impersonation"
rules={{ required: "Enter a delegated scope" }}
>
{({ value, onChange, ref, ...control }) => (
<Input
{...control}
ref={ref}
value={typeof value === "string" ? value : "user_impersonation"}
onChange={onChange}
/>
)}
</AgentFormField>
<p className="text-sm text-muted-foreground">
Users must first sign in through this gateway&apos;s Microsoft SSO. Subsequent delegated calls must
satisfy both user and agent permissions.
</p>
</>
)}
<AgentFormField name="enabled" label="Execution" defaultValue={true}>
{({ value, onChange, id }) => (
<Select
items={EXECUTION_OPTIONS}
value={value === false ? "disabled" : "enabled"}
onValueChange={(next) => onChange(next === "enabled")}
>
<SelectTrigger id={id}>
<SelectValue />
</SelectTrigger>
<SelectContent>
{EXECUTION_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</AgentFormField>
<p className="text-sm text-muted-foreground">
LiteLLM verifies the agent&apos;s Entra token before matching this identity. Saving these fields
configures the binding; an authenticated request provides verification. Runtime authentication headers are
configured separately.
</p>
</>
)}
</section>
</>
);
};

View file

@ -145,10 +145,10 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole, teams
</p>
<Alert className="mb-3">
<Info />
<AlertTitle>Why do agents need keys?</AlertTitle>
<AlertTitle>How do agents authenticate?</AlertTitle>
<AlertDescription>
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.
</AlertDescription>
</Alert>
{isAdmin && (

View file

@ -62,6 +62,12 @@ describe("AgentsTable", () => {
expect(within(keylessRow).getByText("Needs Setup")).toBeInTheDocument();
});
it("shows JWT configured for agents without a virtual key", () => {
render(<AgentsTable agents={[makeAgent({ keys: [], jwt_auth_configured: true })]} {...baseProps} />);
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();

View file

@ -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 <StatusBadge tone="success" label="JWT configured" />;
return hasKeys ? (
<StatusBadge tone="success" label="Active" />
) : (

View file

@ -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();

View file

@ -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);

View file

@ -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<AddAgentFormProps> = ({ 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<AddAgentFormProps> = ({ 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<AddAgentFormProps> = ({ visible, onClose, accessTok
<StatusBadge tone="warning" label="GENERIC" className="h-4 px-1 text-[10px]" />
</span>
<span className="block text-xs whitespace-normal text-warning">
For agents that don&apos;t follow a standard protocol, just needs a virtual key
For outbound agents using an identity provider or virtual key
</span>
</span>
</span>
@ -801,6 +805,8 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
</Select>
</Field>
<AgentIdentityFields accessToken={accessToken} />
<div className="mt-4">
{agentType === CUSTOM_AGENT_TYPE ? (
<FieldGroup>
@ -910,7 +916,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ 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<AddAgentFormProps> = ({ visible, onClose, accessTok
<Separator className="my-4" />
{form.getValues("identity_provider") === "microsoft_entra" && (
<p className="mb-4 text-sm text-muted-foreground">
This agent will authenticate with Microsoft Entra ID. You can skip virtual key creation.
</p>
)}
<RadioGroup
value={keyAssignOption}
onValueChange={(value) => setKeyAssignOption(value as "create_new" | "existing_key" | "skip")}
@ -1004,7 +1015,9 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
className="text-sm text-muted-foreground underline hover:text-foreground"
onClick={() => setKeyAssignOption("skip")}
>
Skip for now — I&apos;ll assign a key later
{form.getValues("identity_provider") === "microsoft_entra"
? "Use Entra JWT authentication"
: "Skip for now, I’ll assign a key later"}
</button>
</div>
</div>
@ -1033,7 +1046,9 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
)}
{!createdKeyValue && !assignedKeyAlias && keyAssignOption === "skip" && (
<p className="mt-2 text-sm text-muted-foreground">
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."}
</p>
)}
</div>

View file

@ -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,

View file

@ -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();
});
});

View file

@ -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<components["schemas"]["AgentResponse"], "identity" | "enabled" | "execution_mode">;
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<AgentIdentityState> | 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<AgentIdentityState>,
): 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 } : {}),
};
};

View file

@ -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();

View file

@ -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: () => <div data-testid="agent-card-discovery" />,
}));
vi.mock("./agent_form_fields", () => ({
default: () => <div data-testid="agent-form-fields" />,
default: vi.fn(() => <div data-testid="agent-form-fields" />),
unmountedA2AFieldNames: () => [],
}));
@ -77,6 +82,9 @@ const agent = {
describe("AgentInfoView settings", () => {
beforeEach(() => {
vi.restoreAllMocks();
vi.mocked(AgentFormFields)
.mockReset()
.mockImplementation(() => <div data-testid="agent-form-fields" />);
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<typeof import("./agent_form_fields")>("./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(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);
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(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);

View file

@ -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<AgentInfoViewProps> = ({ 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<AgentInfoViewProps> = ({ agentId, onClose, accessT
<div>
{/* Overview Panel */}
<TabsContent value="overview" keepMounted>
<AgentIdentityDetails
agentId={agentId}
identity={agent.identity}
accessToken={accessToken}
isAdmin={isAdmin}
/>
<DetailList>
<DetailItem label="Agent ID">{agent.agent_id}</DetailItem>
<DetailItem label="Agent Name">{agent.agent_name}</DetailItem>
@ -505,6 +514,8 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
<AgentFormFields showAgentName={true} panels={panels} />
)}
<AgentIdentityFields accessToken={accessToken} />
{discoveryRequest && (
<div className="mt-4">
<AgentCardDiscovery

View file

@ -82,6 +82,35 @@ describe("AgentSelector", () => {
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(<AgentSelector {...defaultProps} onChange={onChange} allowAccessGroups={false} />);
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(
<AgentSelector
{...defaultProps}
onChange={onChange}
allowAccessGroups={false}
value={{ agents: [], accessGroups: ["retired-group"] }}
/>,
);
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(<AgentSelector {...defaultProps} disabled />);
expect(screen.getByRole("combobox")).toBeDisabled();

View file

@ -19,6 +19,7 @@ interface AgentSelectorProps {
accessToken: string;
placeholder?: string;
disabled?: boolean;
allowAccessGroups?: boolean;
}
const AgentSelector: React.FC<AgentSelectorProps> = ({
@ -28,6 +29,7 @@ const AgentSelector: React.FC<AgentSelectorProps> = ({
accessToken,
placeholder = "Select agents",
disabled = false,
allowAccessGroups = true,
}) => {
const [agents, setAgents] = useState<Agent[]>([]);
const [accessGroups, setAccessGroups] = useState<string[]>([]);
@ -60,12 +62,15 @@ const AgentSelector: React.FC<AgentSelectorProps> = ({
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}`,

View file

@ -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: {

View file

@ -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(<MCPServerSelector accessToken="tok" onChange={onChange} allowAccessGroups={false} />);
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(
<MCPServerSelector
accessToken="tok"
onChange={onChange}
allowAccessGroups={false}
value={{ servers: [], accessGroups: ["retired-group"] }}
/>,
);
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: [] });
});
});

View file

@ -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<MCPServerSelectorProps> = ({
onChange,
value,
@ -32,22 +36,24 @@ const MCPServerSelector: React.FC<MCPServerSelectorProps> = ({
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})`,

View file

@ -67,7 +67,7 @@ export function AgentPermissions({
<div className="space-y-3">
<div className="flex items-center gap-2">
<UserGroupIcon className="h-4 w-4 text-purple-600" />
<p className="text-sm font-semibold text-foreground">Agents</p>
<p className="text-sm font-semibold text-foreground">Allowed agents to call</p>
<Badge variant="secondary">{totalCount}</Badge>
</div>

View file

@ -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);

View file

@ -2039,13 +2039,14 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
)}
</FormField>
<FormField control={form.control} name="mcp_servers_and_groups" label="MCP Servers / Access Groups">
<FormField control={form.control} name="mcp_servers_and_groups" label="MCP Servers">
{({ value, onChange }) => (
<MCPServerSelector
allowAccessGroups={false}
onChange={onChange}
value={value}
accessToken={accessToken || ""}
placeholder="Select MCP servers or access groups (optional)"
placeholder="Select MCP servers or toolsets (optional)"
allowAllProxyMcpServers={is_proxy_admin}
/>
)}
@ -2062,13 +2063,14 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
/>
</div>
<FormField control={form.control} name="agents_and_groups" label="Agents / Access Groups">
<FormField control={form.control} name="agents_and_groups" label="Agents">
{({ value, onChange }) => (
<AgentSelector
allowAccessGroups={false}
onChange={onChange}
value={value}
accessToken={accessToken || ""}
placeholder="Select agents or access groups (optional)"
placeholder="Select agents (optional)"
/>
)}
</FormField>

View file

@ -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();
});
});

View file

@ -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) {

View file

@ -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;