mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
feat(agents): integrate Entra identity registration and authorization
This commit is contained in:
parent
5d777c16d9
commit
320a688bae
99 changed files with 6995 additions and 610 deletions
|
|
@ -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 $$;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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```",
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 ()
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
232
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
232
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
17
litellm/proxy/agent_endpoints/identity.py
Normal file
17
litellm/proxy/agent_endpoints/identity.py
Normal 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)
|
||||
224
litellm/proxy/agent_endpoints/identity_store.py
Normal file
224
litellm/proxy/agent_endpoints/identity_store.py
Normal 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
|
||||
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal file
220
litellm/proxy/agent_endpoints/managed_identity.py
Normal 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")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
96
litellm/types/proxy/agent_identity.py
Normal file
96
litellm/types/proxy/agent_identity.py
Normal 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
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
27
tests/test_litellm/proxy/agent_endpoints/test_identity.py
Normal file
27
tests/test_litellm/proxy/agent_endpoints/test_identity.py
Normal 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
|
||||
402
tests/test_litellm/proxy/agent_endpoints/test_identity_store.py
Normal file
402
tests/test_litellm/proxy/agent_endpoints/test_identity_store.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"] == {}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
@ -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'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'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>
|
||||
</>
|
||||
);
|
||||
};
|
||||
|
|
@ -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 && (
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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" />
|
||||
) : (
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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'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'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>
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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 } : {}),
|
||||
};
|
||||
};
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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} />);
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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}`,
|
||||
|
|
|
|||
|
|
@ -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: {
|
||||
|
|
|
|||
|
|
@ -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: [] });
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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})`,
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
213
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
213
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue