From 320a688bae823b4c334256d4eb37d619a484092f Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:12:00 -0700 Subject: [PATCH] feat(agents): integrate Entra identity registration and authorization --- .../migration.sql | 97 ++++ .../litellm_proxy_extras/schema.prisma | 57 +++ .../mcp_server/auth/managed_agent_access.py | 67 +++ .../mcp_server/auth/user_api_key_auth_mcp.py | 95 +++- .../mcp_server/idp_token_exchange.py | 5 + .../mcp_server/mcp_server_manager.py | 13 +- .../mcp_server/rest_endpoints.py | 33 +- .../mcp_server/ui_session_utils.py | 5 +- litellm/proxy/_lazy_openapi_snapshot.json | 349 ++++++++++++- litellm/proxy/_types.py | 25 +- .../proxy/agent_endpoints/a2a_endpoints.py | 5 + litellm/proxy/agent_endpoints/a2a_routing.py | 2 +- .../proxy/agent_endpoints/agent_registry.py | 178 ++++--- .../auth/agent_access_groups.py | 21 +- .../auth/agent_permission_handler.py | 164 +++++- .../auth/managed_authorization.py | 232 +++++++++ litellm/proxy/agent_endpoints/endpoints.py | 109 +++- litellm/proxy/agent_endpoints/identity.py | 17 + .../proxy/agent_endpoints/identity_store.py | 224 ++++++++ .../proxy/agent_endpoints/managed_identity.py | 220 ++++++++ litellm/proxy/auth/auth_checks.py | 123 +++-- litellm/proxy/auth/handle_jwt.py | 114 ++++- litellm/proxy/auth/user_api_key_auth.py | 54 +- litellm/proxy/common_request_processing.py | 12 +- .../proxy/common_utils/http_parsing_utils.py | 35 +- .../common_utils/registry_read_through.py | 5 +- .../proxy/hooks/proxy_track_cost_callback.py | 10 +- litellm/proxy/image_endpoints/endpoints.py | 18 +- litellm/proxy/litellm_pre_call_utils.py | 14 +- .../sso/agent_subject_enrollment.py | 49 ++ litellm/proxy/management_endpoints/ui_sso.py | 22 + litellm/proxy/proxy_server.py | 17 +- litellm/proxy/schema.prisma | 57 +++ .../spend_management_endpoints.py | 32 +- .../spend_tracking/spend_tracking_utils.py | 1 + .../object_permission_repository.py | 7 +- litellm/repositories/table_repositories.py | 23 +- litellm/repositories/team_repository.py | 9 +- litellm/repositories/user_repository.py | 7 +- litellm/types/agents.py | 18 +- litellm/types/proxy/agent_identity.py | 96 ++++ schema.prisma | 57 +++ .../auth/test_managed_agent_access.py | 301 +++++++++++ .../auth/test_user_api_key_auth_mcp.py | 94 +++- .../mcp_server/test_discoverable_endpoints.py | 2 +- .../mcp_server/test_idp_token_exchange.py | 11 + .../mcp_server/test_mcp_server_manager.py | 44 +- .../mcp_server/test_proxy_api_credentials.py | 4 +- .../mcp_server/test_rest_endpoints.py | 28 +- .../mcp_server/test_ui_session_utils.py | 4 +- .../auth/test_agent_permission_handler.py | 287 ++++++++++- .../auth/test_managed_authorization.py | 484 ++++++++++++++++++ .../agent_endpoints/test_a2a_endpoints.py | 7 +- .../agent_endpoints/test_agent_registry.py | 426 ++++++++++++--- .../proxy/agent_endpoints/test_endpoints.py | 260 +++++++++- .../proxy/agent_endpoints/test_identity.py | 27 + .../agent_endpoints/test_identity_store.py | 402 +++++++++++++++ .../agent_endpoints/test_managed_identity.py | 257 ++++++++++ .../proxy/auth/test_auth_checks.py | 162 +++++- .../proxy/auth/test_handle_jwt.py | 411 +++++++++++---- .../proxy/auth/test_user_api_key_auth.py | 200 ++++++++ .../common_utils/test_http_parsing_utils.py | 42 +- .../test_registry_read_through.py | 34 +- .../hooks/test_proxy_track_cost_callback.py | 31 ++ .../sso/test_agent_subject_enrollment.py | 114 +++++ .../test_mcp_management_endpoints.py | 4 +- .../test_team_endpoints.py | 12 +- .../proxy/management_endpoints/test_ui_sso.py | 3 + .../proxy/proxy_server/test_proxy_config.py | 40 +- .../test_spend_management_endpoints.py | 46 +- .../test_spend_query_optimization.py | 10 +- .../test_spend_tracking_utils.py | 36 ++ .../proxy/test_route_a2a_models.py | 3 + .../_components/AgentIdentityDetails.test.tsx | 43 ++ .../_components/AgentIdentityDetails.tsx | 81 +++ .../_components/AgentIdentityFields.tsx | 259 ++++++++++ .../agents/_components/AgentsPanel.tsx | 6 +- .../agents/_components/AgentsTable.test.tsx | 6 + .../agents/_components/AgentsTableColumns.tsx | 1 + .../add_agent_form.integration.test.tsx | 48 +- .../_components/add_agent_form.test.tsx | 10 +- .../agents/_components/add_agent_form.tsx | 27 +- .../agents/_components/agent_config.ts | 29 +- .../agents/_components/agent_identity.test.ts | 83 +++ .../agents/_components/agent_identity.ts | 101 ++++ .../agent_info.integration.test.tsx | 45 ++ .../agents/_components/agent_info.test.tsx | 27 +- .../agents/_components/agent_info.tsx | 13 +- .../agent_management/AgentSelector.test.tsx | 29 ++ .../agent_management/AgentSelector.tsx | 11 +- .../src/components/agents/types.ts | 5 + .../MCPServerSelector.test.tsx | 38 ++ .../MCPServerSelector.tsx | 14 +- .../permissions/AgentPermissions.tsx | 2 +- .../src/components/team/TeamInfo.test.tsx | 4 +- .../src/components/team/TeamInfo.tsx | 10 +- .../RequestLogsTableColumns.test.tsx | 14 + .../view_logs/RequestLogsTableColumns.tsx | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 213 +++++++- 99 files changed, 6995 insertions(+), 610 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql create mode 100644 litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py create mode 100644 litellm/proxy/agent_endpoints/auth/managed_authorization.py create mode 100644 litellm/proxy/agent_endpoints/identity.py create mode 100644 litellm/proxy/agent_endpoints/identity_store.py create mode 100644 litellm/proxy/agent_endpoints/managed_identity.py create mode 100644 litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py create mode 100644 litellm/types/proxy/agent_identity.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_identity.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_identity_store.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py create mode 100644 tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql new file mode 100644 index 00000000000..06cf03b26b5 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260921190000_agent_identity/migration.sql @@ -0,0 +1,97 @@ +-- AlterTable +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "enabled" BOOLEAN NOT NULL DEFAULT true, +ADD COLUMN IF NOT EXISTS "execution_mode" TEXT NOT NULL DEFAULT 'autonomous', +ADD COLUMN IF NOT EXISTS "identity_managed" BOOLEAN NOT NULL DEFAULT false; + +-- AlterTable +ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "billing_agent_id" TEXT; + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_AgentIdentity" ( + "agent_id" TEXT NOT NULL, + "active" BOOLEAN NOT NULL DEFAULT true, + "provider" TEXT NOT NULL, + "issuer" TEXT NOT NULL, + "tenant_id" TEXT NOT NULL, + "client_id" TEXT NOT NULL, + "service_principal_id" TEXT, + "required_roles" TEXT[] DEFAULT ARRAY[]::TEXT[], + "required_scopes" TEXT[] DEFAULT ARRAY['user_impersonation']::TEXT[], + "revision" TEXT NOT NULL, + "last_authenticated_at" TIMESTAMP(3), + + CONSTRAINT "LiteLLM_AgentIdentity_pkey" PRIMARY KEY ("agent_id") +); + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgentIdentity" ( + "binding_id" TEXT NOT NULL, + "agent_id" TEXT, + "provider" TEXT NOT NULL, + "issuer" TEXT NOT NULL, + "tenant_id" TEXT NOT NULL, + "client_id" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_RetiredAgentIdentity_pkey" PRIMARY KEY ("binding_id") +); + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgent" ( + "original_agent_id" TEXT NOT NULL, + "retired_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_RetiredAgent_pkey" PRIMARY KEY ("original_agent_id") +); + +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_VerifiedSubject" ( + "subject_id" TEXT NOT NULL, + "issuer" TEXT NOT NULL, + "tenant_id" TEXT NOT NULL, + "oid" TEXT NOT NULL, + "kind" TEXT NOT NULL DEFAULT 'human', + "user_id" TEXT, + "verified_via" TEXT NOT NULL DEFAULT 'sso_interactive', + "verified_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_VerifiedSubject_pkey" PRIMARY KEY ("subject_id") +); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_AgentIdentity"("provider", "tenant_id", "client_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_issuer_service_principal_id_key" ON "LiteLLM_AgentIdentity"("issuer", "service_principal_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_RetiredAgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_RetiredAgentIdentity"("provider", "tenant_id", "client_id"); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_user_id_idx" ON "LiteLLM_VerifiedSubject"("user_id"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_issuer_tenant_id_oid_key" ON "LiteLLM_VerifiedSubject"("issuer", "tenant_id", "oid"); + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentIdentity_agent_id_fkey') THEN + ALTER TABLE "LiteLLM_AgentIdentity" ADD CONSTRAINT "LiteLLM_AgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE; + END IF; +END $$; + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_RetiredAgentIdentity_agent_id_fkey') THEN + ALTER TABLE "LiteLLM_RetiredAgentIdentity" ADD CONSTRAINT "LiteLLM_RetiredAgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE; + END IF; +END $$; + +-- AddForeignKey +DO $$ +BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_user_id_fkey') THEN + ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE CASCADE ON UPDATE CASCADE; + END IF; +END $$; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 69c63d9ecd6..8535d1004f7 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -78,6 +78,11 @@ model LiteLLM_AgentsTable { object_permission_id String? object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) spend Float @default(0.0) + identity_managed Boolean @default(false) + enabled Boolean @default(true) + execution_mode String @default("autonomous") + identity LiteLLM_AgentIdentity? + retired_identities LiteLLM_RetiredAgentIdentity[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -88,6 +93,56 @@ model LiteLLM_AgentsTable { updated_by String } +model LiteLLM_AgentIdentity { + agent_id String @id + active Boolean @default(true) + agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade) + provider String + issuer String + tenant_id String + client_id String + service_principal_id String? + required_roles String[] @default([]) + required_scopes String[] @default(["user_impersonation"]) + revision String @default(uuid()) + last_authenticated_at DateTime? + @@unique([provider, tenant_id, client_id]) + @@unique([issuer, service_principal_id]) +} + +model LiteLLM_RetiredAgentIdentity { + binding_id String @id @default(uuid()) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + provider String + issuer String + tenant_id String + client_id String + @@unique([provider, tenant_id, client_id]) +} + +model LiteLLM_RetiredAgent { + original_agent_id String @id + retired_at DateTime @default(now()) +} + +model LiteLLM_VerifiedSubject { + subject_id String @id @default(uuid()) + issuer String + tenant_id String + oid String + kind String @default("human") + user_id String? + user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + verified_via String @default("sso_interactive") + verified_at DateTime @default(now()) + @@unique([issuer, tenant_id, oid]) + @@index([user_id]) +} + + + + model LiteLLM_OrganizationTable { organization_id String @id @default(uuid()) organization_alias String @@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? team_id String? @@ -674,6 +730,7 @@ model LiteLLM_SpendLogs { session_id String? status String? mcp_namespaced_tool_name String? + billing_agent_id String? agent_id String? proxy_server_request Json? @default("{}") litellm_call_id String? diff --git a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py new file mode 100644 index 00000000000..dd7390a784d --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py @@ -0,0 +1,67 @@ +from types import MappingProxyType +from typing import Final + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings +from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure +from litellm.types.proxy.agent_identity import AgentIdentityFailure + + +async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True) + return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True})) + + +async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + agent: Final = auth.managed_agent_policy + if agent is None: + return () + + try: + base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth)) + ceilings: Final = await resolve_managed_agent_ceilings(agent) + expanded: Final = tuple( + frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids))) + for ceiling in ceilings + ) + own: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded)) + context: Final = auth.managed_agent_context + if context is None or context.mode == "autonomous": + return tuple(sorted(own)) + if context.user_id is None: + return () + human: Final = await _delegated_resource_subject(context.user_id) + allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human) + return tuple(sorted(own.intersection(allowed))) + except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable") + ) + + +async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + if server_id not in await managed_agent_servers(auth): + return [] + try: + own: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth) + context: Final = auth.managed_agent_context + if context is None or context.mode == "autonomous": + return own + if context.user_id is None: + return [] + human: Final = await _delegated_resource_subject(context.user_id) + human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human) + if own is None: + return human_tools + return own if human_tools is None else sorted(frozenset(own).intersection(human_tools)) + except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable") + ) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a93ffaeac9f..ec92bbba677 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -67,7 +67,6 @@ from litellm.repositories.table_repositories import ( AgentsRepository, MCPServerRepository, ) -from litellm.repositories.user_repository import UserRepository from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: @@ -1086,7 +1085,7 @@ class MCPRequestHandler: assert_never(identity.subject_type) @staticmethod - async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth: + async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth: """Reload the live user an interactively-minted envelope references and admit them as themselves. The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the @@ -1111,6 +1110,7 @@ class MCPRequestHandler: prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, user_id_upsert=False, + check_db_only=requires_fresh_policy, ) # Resolve the user's own MCP object permission (get_user_object does not load it) so the shared # get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same @@ -1119,6 +1119,7 @@ class MCPRequestHandler: if user_object is not None and object_permission is None and user_object.object_permission_id: object_permission = await get_object_permission( object_permission_id=user_object.object_permission_id, + check_db_only=requires_fresh_policy, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) @@ -1147,6 +1148,7 @@ class MCPRequestHandler: # Server-only marker, set AFTER construction: the before-validator strips it from any validated # input, so caller-supplied data (key metadata, JWT claims) can never forge it. admitted.mcp_admitted_user_subject = True + admitted.requires_fresh_policy = requires_fresh_policy # Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through # several teams under its own identity, so without this a cross-team user outruns every team's # limit. Resolved from the same roster-checked sources as the grant union, so a team throttles @@ -1597,6 +1599,11 @@ class MCPRequestHandler: """ from litellm.proxy.proxy_server import general_settings + if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None: + from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers + + return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped") + key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth) try: @@ -1606,7 +1613,7 @@ class MCPRequestHandler: # independent; an opt-out silences only its own source, inside the recursive call). if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: return MCPServerAccess( - server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)), + server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)), ) # Get allowed servers from key and team @@ -1703,7 +1710,7 @@ class MCPRequestHandler: if user_api_key_auth and user_api_key_auth.agent_id: agent_capped: Final = _agent_capped_servers( allowed_mcp_servers, - await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth), + await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth), await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth), ) if agent_capped is not None: @@ -1829,10 +1836,12 @@ class MCPRequestHandler: scoped.object_permission = auth.object_permission scoped.object_permission_id = auth.object_permission_id scoped.access_group_ids = auth.access_group_ids + scoped.requires_fresh_policy = auth.requires_fresh_policy + scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only return scoped @staticmethod - async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: """The independent sources a keyless admitted subject reaches MCP servers through: their own direct grants, plus every team they are a live roster member of. @@ -1886,6 +1895,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(auth and auth.requires_fresh_policy), ) except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others # Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for @@ -1941,11 +1951,11 @@ class MCPRequestHandler: roster instead of by grant charged unrelated teams' buckets).""" return [ (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) - for source in await MCPRequestHandler._admitted_subject_sources(auth) + for source in await MCPRequestHandler.admitted_subject_sources(auth) ] @staticmethod - async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]: + async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]: """Union of what each of the admitted subject's sources reaches, each answered by the canonical resolver so no rule is reimplemented for this caller shape.""" reachable: Final[set[str]] = set() @@ -2007,7 +2017,7 @@ class MCPRequestHandler: return min((source for source, _ in granting), key=lambda s: s.team_id or "") @staticmethod - async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: + async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: """Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the sources that actually grant that server. @@ -2088,6 +2098,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if not team_obj: @@ -2171,6 +2182,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) @staticmethod @@ -2219,12 +2231,17 @@ class MCPRequestHandler: if not user_api_key_auth: return None + if user_api_key_auth.managed_agent_policy is not None: + from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools + + return await managed_agent_tools(server_id, user_api_key_auth) + try: # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per # source and shares nothing with the single-credential prelude below. Ordering is the invariant: # sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant. if _is_mcp_admitted_user_subject(user_api_key_auth): - return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth) + return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth) # Get key and team object permissions (already loaded in main auth flow) key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth) @@ -2334,7 +2351,7 @@ class MCPRequestHandler: if user_api_key_auth.agent_id: # Pre-fetch agent object_permission once to avoid a duplicate DB query. agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) - agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server( + agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server( server_id=server_id, user_api_key_auth=user_api_key_auth, agent_object_permission=agent_obj_perm, @@ -2456,6 +2473,7 @@ class MCPRequestHandler: prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if not raw_server_ids: return [] @@ -2502,6 +2520,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if key_object_permission is None: return [] @@ -2550,7 +2569,7 @@ class MCPRequestHandler: """Get allowed MCP servers a caller inherits from the team it is pinned to. Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not - fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``, + fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``, and each of those sources pins a single ``team_id`` before reaching this point. Keeping the fan-out here as well would be a second multi-team path to drift from that one. """ @@ -2568,7 +2587,7 @@ class MCPRequestHandler: which must NOT silently gain the union across every team the user belongs to), and it covers each single-source auth an admitted subject fans out into — those pin a team_id, so they land on the first branch. The admitted subject itself never reaches here: it resolves per source - in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel + in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel resolves to no teams exactly as before.""" if user_api_key_auth is None or not user_api_key_auth.team_id: return [] @@ -2596,6 +2615,7 @@ class MCPRequestHandler: user_id_upsert=False, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e) @@ -2667,6 +2687,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if team_obj is None: return [] @@ -2680,6 +2701,7 @@ class MCPRequestHandler: prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers) @@ -2716,6 +2738,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with raise unloadable from e @@ -2961,7 +2984,9 @@ class MCPRequestHandler: return None user_id: Final = user_api_key_auth.user_id - object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client) + object_permission_id: Final = await MCPRequestHandler._user_object_permission_id( + user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy + ) if object_permission_id is None: return None @@ -2971,6 +2996,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if object_permission is None: raise ValueError( @@ -2979,7 +3005,9 @@ class MCPRequestHandler: return object_permission @staticmethod - async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None: + async def _user_object_permission_id( + user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False + ) -> str | None: """The permission row this human's user row links to, or None when they link none. Caches the link (with a sentinel for "links none") so a human without an entitlement costs no @@ -2988,16 +3016,23 @@ class MCPRequestHandler: whether someone is entitled is the state that existed before this level, so it places no ceiling. Only a link we DID resolve can make the caller deny. """ + from litellm.proxy.auth.auth_checks import get_user_object from litellm.proxy.proxy_server import user_api_key_cache cache_key: Final = user_object_permission_id_cache_key(user_id) try: - cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) + cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key) if cached == USER_NO_MCP_PERMISSION_SENTINEL: return None if isinstance(cached, str) and cached: return cached - user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) + user_row: Final = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + check_db_only=check_db_only, + ) linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None object_permission_id: Final = linked if isinstance(linked, str) and linked else None await user_api_key_cache.async_set_cache( @@ -3006,7 +3041,9 @@ class MCPRequestHandler: ttl=get_management_object_ttl(user_api_key_cache), ) return object_permission_id - except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before + except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior + if check_db_only: + raise HTTPException(503, "User policy is unavailable") from e verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e) return None @@ -3119,9 +3156,13 @@ class MCPRequestHandler: (any non-empty entitlement, or an unresolved one, disqualifies), exactly as ``operator_open_server_ids`` reads the same row. The one owner of this predicate: the server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open - channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot + channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot disagree.""" - if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth): + if ( + user_api_key_auth is None + or user_api_key_auth.mcp_explicit_grants_only + or not user_api_key_has_admin_view(user_api_key_auth) + ): return False object_permission: Final = user_api_key_auth.object_permission credential_scoped: Final = ( @@ -3302,6 +3343,10 @@ class MCPRequestHandler: if not user_api_key_auth or not user_api_key_auth.agent_id: return None + if user_api_key_auth.managed_agent_policy is not None: + permission: Final = user_api_key_auth.managed_agent_policy.object_permission + return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None + if prisma_client is None: verbose_logger.debug("prisma_client is None") return None @@ -3319,7 +3364,7 @@ class MCPRequestHandler: ) @staticmethod - async def _get_allowed_mcp_servers_for_agent( + async def get_allowed_mcp_servers_for_agent( user_api_key_auth: UserAPIKeyAuth | None = None, agent_object_permission: LiteLLM_ObjectPermissionTable | None = None, ) -> list[str]: @@ -3363,7 +3408,7 @@ class MCPRequestHandler: toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm) return list({*expanded_direct_servers, *access_group_servers, *toolset_grants}) except Exception as e: - if isinstance(e, UnloadableEntitlementError): + if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError): raise verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e) return [] @@ -3390,7 +3435,7 @@ class MCPRequestHandler: return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids))) @staticmethod - async def _get_agent_tool_permissions_for_server( + async def get_agent_tool_permissions_for_server( server_id: str, user_api_key_auth: UserAPIKeyAuth | None = None, agent_object_permission: LiteLLM_ObjectPermissionTable | None = None, @@ -3432,9 +3477,9 @@ class MCPRequestHandler: ) toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id) agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools) - return list(agent_tools) if agent_tools else None + return list(agent_tools) if agent_tools is not None else None except Exception as e: - if isinstance(e, UnloadableEntitlementError): + if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError): raise verbose_logger.warning("Failed to get agent tool permissions for server: %s", e) return None @@ -3548,6 +3593,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if key_object_permission is None: return [] @@ -3591,6 +3637,7 @@ class MCPRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) if team_obj is None: verbose_logger.debug("team_obj is None") diff --git a/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py b/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py index a437df17e6a..15632eb4783 100644 --- a/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py +++ b/litellm/proxy/_experimental/mcp_server/idp_token_exchange.py @@ -170,6 +170,11 @@ async def identity_from_subject_token( return _refusal_for(denied, denied.message) except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures return _refusal_for(denied, denied) + if result.get("agent_id") is not None: + return SubjectTokenRefusal( + error="invalid_request", + description="Agent tokens require direct JWT authentication; this exchange supports users only", + ) user_id: Final = result["user_id"] if user_id is None: return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows") diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 31896d9ddc5..cdd1dcb64c5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3418,7 +3418,9 @@ class MCPServerManager: ``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union, which precomputes both for its fallback path, does not compute them twice.""" - if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None: + if user_api_key_auth is not None and ( + user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only + ): return set() if allow_all_server_ids is None: allow_all_server_ids = self.get_allow_all_keys_server_ids() @@ -3467,9 +3469,14 @@ class MCPServerManager: 2. If admin and no object_permission, return all servers 3. Otherwise, use standard permission checks """ + if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None: + managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth) + return managed if access is None else [server for server in managed if server in access.server_ids] + from litellm.proxy.proxy_server import general_settings as proxy_general_settings resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings + explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only) allow_all_server_ids: Final = self.get_allow_all_keys_server_ids() # A keyless admitted subject is resolved per grant source, and channel decisions that are @@ -3501,7 +3508,7 @@ class MCPServerManager: # only keys without their own mcp_servers list get submitted servers unioned in. submitted_server_ids: Final = ( [] - if has_explicit_object_permission + if has_explicit_object_permission or explicit_grants_only else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth) ) @@ -3570,7 +3577,7 @@ class MCPServerManager: return [ server_id for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids) - if scope is None or server_id == scope + if not explicit_grants_only and (scope is None or server_id == scope) ] async def resolve_toolset_tool_permissions( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 02694f110b1..bae1043691b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -5,7 +5,7 @@ from dataclasses import dataclass from datetime import datetime from traceback import walk_tb from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict from uuid import uuid4 import anyio @@ -14,6 +14,7 @@ import httpx2 from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from pydantic import ValidationError from starlette.datastructures import Headers +from typing_extensions import ReadOnly from litellm._logging import verbose_logger from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT @@ -62,7 +63,27 @@ if TYPE_CHECKING: from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm.types.mcp import MCPAuth -from litellm.types.utils import CallTypes +from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall + + +class _MCPModelMetadata(TypedDict): + model_group: ReadOnly[str] + + +def _stamp_mcp_tool_metadata(logging_obj: "LiteLLMLoggingObj | None", server_id: str, tool_name: str) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + if logging_obj is None: + return + server: Final = global_mcp_server_manager.get_mcp_server_by_id( + server_id + ) or global_mcp_server_manager.get_mcp_server_by_name(server_id) + metadata: Final[StandardLoggingMCPToolCall] = { + "name": tool_name, + "mcp_server_name": server.name if server is not None else server_id, + } + logging_obj.model_call_details["mcp_tool_call_metadata"] = metadata + MCP_AVAILABLE: bool = True try: @@ -1143,6 +1164,12 @@ if MCP_AVAILABLE: }, ) + data["model"] = f"MCP: {tool_name}" + model_metadata: Final[_MCPModelMetadata] = { + **(data.get("metadata") or MappingProxyType({})), + "model_group": f"MCP: {tool_name}", + } + data["metadata"] = model_metadata proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) _request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below try: @@ -1176,6 +1203,8 @@ if MCP_AVAILABLE: if "metadata" in data and "user_api_key_auth" in data["metadata"]: data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"] + _stamp_mcp_tool_metadata(logging_obj, server_id, tool_name) + # Resolve allowed MCP servers with IP filtering ( allowed_mcp_servers, diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index 107a4818de1..901259c18ad 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, user_id_upsert=False, + check_db_only=user_api_key_auth.requires_fresh_policy, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) @@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey ) try: - admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id) + admitted: Final = await MCPRequestHandler.reload_admitted_user( + user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy + ) except HTTPException as e: verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail) return None diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index ded4db6d2aa..f374ba09d07 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -2378,6 +2378,19 @@ "title": "Agent Name", "type": "string" }, + "enabled": { + "title": "Enabled", + "type": "boolean" + }, + "execution_mode": { + "enum": [ + "autonomous", + "delegated", + "both" + ], + "title": "Execution Mode", + "type": "string" + }, "extra_headers": { "anyOf": [ { @@ -2392,6 +2405,16 @@ ], "title": "Extra Headers" }, + "identity": { + "anyOf": [ + { + "$ref": "#/components/schemas/EntraIdentityConfig" + }, + { + "type": "null" + } + ] + }, "kill_switch": { "anyOf": [ { @@ -2470,8 +2493,7 @@ } }, "required": [ - "agent_name", - "agent_card_params" + "agent_name" ], "title": "AgentConfig", "type": "object" @@ -2521,6 +2543,91 @@ "title": "AgentExtension", "type": "object" }, + "AgentIdentityBinding": { + "properties": { + "active": { + "default": true, + "title": "Active", + "type": "boolean" + }, + "agent_id": { + "title": "Agent Id", + "type": "string" + }, + "client_id": { + "title": "Client Id", + "type": "string" + }, + "issuer": { + "title": "Issuer", + "type": "string" + }, + "last_authenticated_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Last Authenticated At" + }, + "provider": { + "const": "microsoft_entra", + "title": "Provider", + "type": "string" + }, + "required_roles": { + "default": [], + "items": { + "type": "string" + }, + "title": "Required Roles", + "type": "array" + }, + "required_scopes": { + "default": [ + "user_impersonation" + ], + "items": { + "type": "string" + }, + "title": "Required Scopes", + "type": "array" + }, + "revision": { + "title": "Revision", + "type": "string" + }, + "service_principal_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Service Principal Id" + }, + "tenant_id": { + "title": "Tenant Id", + "type": "string" + } + }, + "required": [ + "agent_id", + "provider", + "tenant_id", + "client_id", + "issuer", + "revision" + ], + "title": "AgentIdentityBinding", + "type": "object" + }, "AgentInterface": { "description": "Declares a combination of a target URL and a transport protocol.", "properties": { @@ -2972,6 +3079,21 @@ ], "title": "Created By" }, + "enabled": { + "default": true, + "title": "Enabled", + "type": "boolean" + }, + "execution_mode": { + "default": "autonomous", + "enum": [ + "autonomous", + "delegated", + "both" + ], + "title": "Execution Mode", + "type": "string" + }, "extra_headers": { "anyOf": [ { @@ -2986,6 +3108,26 @@ ], "title": "Extra Headers" }, + "identity": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentIdentityBinding" + }, + { + "type": "null" + } + ] + }, + "identity_managed": { + "default": false, + "title": "Identity Managed", + "type": "boolean" + }, + "jwt_auth_configured": { + "default": false, + "title": "Jwt Auth Configured", + "type": "boolean" + }, "keys": { "anyOf": [ { @@ -3441,6 +3583,61 @@ "title": "DailySpendMetadata", "type": "object" }, + "EntraIdentityConfig": { + "additionalProperties": false, + "properties": { + "client_id": { + "title": "Client Id", + "type": "string" + }, + "provider": { + "const": "microsoft_entra", + "title": "Provider", + "type": "string" + }, + "required_roles": { + "default": [], + "items": { + "type": "string" + }, + "title": "Required Roles", + "type": "array" + }, + "required_scopes": { + "default": [ + "user_impersonation" + ], + "description": "Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.", + "items": { + "type": "string" + }, + "title": "Required Scopes", + "type": "array" + }, + "service_principal_id": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Service Principal Id" + }, + "tenant_id": { + "title": "Tenant Id", + "type": "string" + } + }, + "required": [ + "provider", + "tenant_id", + "client_id" + ], + "title": "EntraIdentityConfig", + "type": "object" + }, "HTTPAuthSecurityScheme": { "description": "Defines a security scheme using HTTP authentication.", "properties": { @@ -3590,6 +3787,54 @@ "title": "MakeAgentsPublicRequest", "type": "object" }, + "ManagedAgentIdentityStatus": { + "properties": { + "enabled": { + "default": true, + "title": "Enabled", + "type": "boolean" + }, + "execution_mode": { + "default": "autonomous", + "enum": [ + "autonomous", + "delegated", + "both" + ], + "title": "Execution Mode", + "type": "string" + }, + "identity": { + "anyOf": [ + { + "$ref": "#/components/schemas/AgentIdentityBinding" + }, + { + "type": "null" + } + ] + }, + "identity_managed": { + "default": false, + "title": "Identity Managed", + "type": "boolean" + }, + "last_authenticated_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Last Authenticated At" + } + }, + "title": "ManagedAgentIdentityStatus", + "type": "object" + }, "MetricWithMetadata": { "properties": { "api_key_breakdown": { @@ -3790,6 +4035,19 @@ "title": "Agent Name", "type": "string" }, + "enabled": { + "title": "Enabled", + "type": "boolean" + }, + "execution_mode": { + "enum": [ + "autonomous", + "delegated", + "both" + ], + "title": "Execution Mode", + "type": "string" + }, "extra_headers": { "anyOf": [ { @@ -3804,6 +4062,16 @@ ], "title": "Extra Headers" }, + "identity": { + "anyOf": [ + { + "$ref": "#/components/schemas/EntraIdentityConfig" + }, + { + "type": "null" + } + ] + }, "kill_switch": { "anyOf": [ { @@ -4325,6 +4593,36 @@ ] } }, + "/v1/agents/identity/providers": { + "get": { + "operationId": "get_agent_identity_providers_v1_agents_identity_providers_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "items": { + "type": "string" + }, + "title": "Response Get Agent Identity Providers V1 Agents Identity Providers Get", + "type": "array" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Identity Providers", + "tags": [ + "agents" + ] + } + }, "/v1/agents/make_public": { "post": { "description": "Make multiple agents publicly discoverable\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/make_public\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"agent_ids\": [\"123e4567-e89b-12d3-a456-426614174000\", \"123e4567-e89b-12d3-a456-426614174001\"]\n }'\n```\n\nExample Response:\n```json\n{\n \"agent_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"agent_name\": \"my-custom-agent\",\n \"litellm_params\": {\n \"make_public\": true\n },\n \"agent_card_params\": {...},\n \"created_at\": \"2025-11-15T10:30:00Z\",\n \"updated_at\": \"2025-11-15T10:35:00Z\",\n \"created_by\": \"user123\",\n \"updated_by\": \"user123\"\n}\n```", @@ -4576,6 +4874,53 @@ ] } }, + "/v1/agents/{agent_id}/identity": { + "get": { + "operationId": "get_agent_identity_status_v1_agents__agent_id__identity_get", + "parameters": [ + { + "in": "path", + "name": "agent_id", + "required": true, + "schema": { + "title": "Agent Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ManagedAgentIdentityStatus" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Identity Status", + "tags": [ + "agents" + ] + } + }, "/v1/agents/{agent_id}/kill_switch": { "post": { "description": "Fire the agent's configured kill switch webhook. Proxy admin only.\n\nLiteLLM only makes the configured HTTP call and reports what came back; it\ndoes not change the agent's state in LiteLLM. Returns 200 when the webhook\nanswered 2xx, 502 with the same result body otherwise. Every attempt is\nwritten to the audit log as a `kill_switch_fired` row against the agent.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch\" \\\n -H \"Authorization: Bearer \"\n```", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 14aa42afefd..641be7078b7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -27,7 +27,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( validate_langfuse_span_scope_value, validate_no_callback_env_reference, ) -from litellm.types.agents import AgentCaller +from litellm.types.agents import AgentCaller, AgentResponse from litellm.types.integrations.compression_interception import ( CompressionSavingsMetadata, ) @@ -46,6 +46,7 @@ from litellm.types.mcp import ( MCPTransportType, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo +from litellm.types.proxy.agent_identity import ManagedAgentContext from litellm.types.proxy.carried_budget_state import ( OrgBudgetSnapshot, TeamBudgetSnapshot, @@ -569,6 +570,7 @@ class LiteLLMRoutes(enum.Enum): "/agents", "/a2a/{agent_id}", "/a2a/{agent_id}/message/send", + "/v1/a2a/{agent_id}/message/send", "/a2a/{agent_id}/message/stream", "/a2a/{agent_id}/.well-known/agent-card.json", ) @@ -3313,6 +3315,8 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union # or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization. mcp_admitted_user_subject: bool = Field(default=False, exclude=True) + requires_fresh_policy: bool = Field(default=False, exclude=True) + mcp_explicit_grants_only: bool = Field(default=False, exclude=True) # team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP # servers through several teams at once and therefore has no single team_id for the limiter to # key off. Server-only and stripped from validated input for the same reason as the marker @@ -3337,6 +3341,11 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob "user id." ), ) + invoked_agent_id: str | None = Field(default=None, exclude=True) + agent_invocation_cost: float | None = Field(default=None, exclude=True) + billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True) + managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True) + managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True) agent_caller: AgentCaller | None = Field( default=None, exclude=True, @@ -3374,11 +3383,18 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # path via post-construction assignment. Strip it from any validated input (constructor # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. values.pop("mcp_admitted_user_subject", None) + values.pop("requires_fresh_policy", None) + values.pop("mcp_explicit_grants_only", None) values.pop("mcp_source_team_rpm_limits", None) values.pop("mcp_session_resource_server_id", None) values.pop("mcp_toolset_id", None) values.pop("via_virtual_key", None) values.pop("agent_caller", None) + values.pop("managed_agent_context", None) + values.pop("managed_agent_policy", None) + values.pop("invoked_agent_id", None) + values.pop("agent_invocation_cost", None) + values.pop("billing_agent_policy", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) if isinstance(values.get("api_key"), str): @@ -4068,6 +4084,11 @@ class SpendLogsRouterMetadata(TypedDict): class SpendLogsMetadata(TypedDict): + actor_agent_id: ReadOnly[NotRequired[str | None]] + target_agent_id: ReadOnly[NotRequired[str | None]] + billing_agent_id: ReadOnly[NotRequired[str | None]] + agent_execution_mode: ReadOnly[NotRequired[str | None]] + verified_human_user_id: ReadOnly[NotRequired[str | None]] autorouter_baseline_observation: ReadOnly[str | None] """ Specific metadata k,v pairs logged to spendlogs for easier cost tracking @@ -4131,6 +4152,7 @@ class SpendLogsPayload(TypedDict): model_id: str | None model_group: str | None mcp_namespaced_tool_name: str | None + billing_agent_id: ReadOnly[NotRequired[str | None]] agent_id: str | None api_base: str user: str @@ -5053,6 +5075,7 @@ class JWTAuthBuilderResult(TypedDict): org_id: str | None team_membership: LiteLLM_TeamMembership | None jwt_claims: dict # Decoded JWT token claims (avoids re-decoding) + managed_agent_context: ReadOnly[NotRequired[ManagedAgentContext | None]] agent_id: ReadOnly[str | None] diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 2a189a76545..5589005cc0b 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -723,6 +723,8 @@ async def invoke_agent_a2a( detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", ) + user_api_key_dict.invoked_agent_id = agent.agent_id + _enforce_inbound_trace_id(agent, request) # Get backend URL and agent name @@ -760,6 +762,8 @@ async def invoke_agent_a2a( if "metadata" not in body: body["metadata"] = {} body["metadata"]["agent_id"] = agent.agent_id + body["metadata"]["model_group"] = f"a2a_agent/{agent_name}" + body["metadata"]["model_info"] = {"id": agent.agent_id} body["agent_id"] = agent.agent_id body.update( @@ -863,6 +867,7 @@ async def invoke_agent_a2a( # results written by the unified_guardrail hook are captured. logging_obj._defer_async_logging = True response = await asend_message( + model=f"a2a_agent/{agent_name}", request=a2a_request, api_base=agent_url, litellm_params=litellm_params, diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 8a795214750..c57315ebc21 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -57,7 +57,7 @@ async def route_a2a_agent_request( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value ) - if not is_admin: + if not is_admin or agent.identity_managed: is_allowed: Final = await AgentRequestHandler.is_agent_allowed( agent_id=agent.agent_id, user_api_key_auth=user_api_key_dict, diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 3e775d7648e..ef7128a0e8e 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -6,6 +6,7 @@ from datetime import datetime, timezone from types import MappingProxyType from typing import TYPE_CHECKING, Final, NamedTuple, Protocol, TypedDict +from fastapi import HTTPException from pydantic import TypeAdapter, ValidationError from typing_extensions import ReadOnly @@ -14,13 +15,16 @@ from litellm.constants import REDACTED_BY_LITELM_STRING from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy.agent_endpoints.kill_switch import restore_kill_switch +from litellm.proxy.agent_endpoints.managed_identity import managed_write_fields, raise_identity_failure from litellm.proxy.management_helpers.object_permission_utils import ( - handle_update_object_permission_common, + prepare_object_permission_upsert, ) from litellm.proxy.utils import PrismaClient +from litellm.repositories.base_repository import is_unique_violation from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest +from litellm.types.proxy.agent_identity import AgentIdentityFailure if TYPE_CHECKING: from prisma import models as prisma_models @@ -135,6 +139,39 @@ def object_permission_table( return table +class AgentPermissionWrite(TypedDict, total=False): + create: ReadOnly[Mapping[str, object]] + update: ReadOnly[Mapping[str, object]] + + +async def _permission_write( + incoming: Mapping[str, object], + existing_id: str | None, + client: PrismaClient, +) -> AgentPermissionWrite | None: + raw: Final = incoming.get("object_permission") + if raw is None: + return None + permission: Final = _AGENT_PARAMS_ADAPTER.validate_python(raw) + prepared: Final = await prepare_object_permission_upsert(permission, existing_id, client) + if existing_id is None: + created: Final[AgentPermissionWrite] = {"create": prepared.record} + return created + updated: Final[AgentPermissionWrite] = {"update": prepared.record} + return updated + + +def _managed_fields( + incoming: Mapping[str, object], + existing: AgentResponse | None, + updated_by: str, +) -> Mapping[str, object]: + result: Final = managed_write_fields(incoming, existing, updated_by) + if isinstance(result, AgentIdentityFailure): + raise_identity_failure(result, 400) + return result + + def _dump_agent_params(raw: Mapping[str, object]) -> dict[str, object]: model_dump: Final[Callable[[], dict[str, object]] | None] = getattr(raw, "model_dump", None) if model_dump is not None: @@ -552,11 +589,7 @@ class AgentRegistry: agent_card_params_dict: Final[dict[str, object]] = _dump_agent_params(agent_card_params_obj) agent_card_params: Final[str] = safe_dumps(agent_card_params_dict) - # Handle object_permission (MCP tool access for agent) - object_permission_id: str | None = None - if agent.get("object_permission") is not None: - agent_copy: Final = dict(agent) - object_permission_id = await handle_update_object_permission_common(agent_copy, None, prisma_client) + permission_write: Final = await _permission_write(agent, None, prisma_client) # Serialize static_headers static_headers_obj: Final = agent.get("static_headers") @@ -583,8 +616,8 @@ class AgentRegistry: create_data["extra_headers"] = extra_headers_val if access_group_ids_val is not None: create_data["access_group_ids"] = tuple(dict.fromkeys(access_group_ids_val)) - if object_permission_id is not None: - create_data["object_permission_id"] = object_permission_id + if permission_write is not None: + create_data["object_permission"] = permission_write for rate_field in ( "tpm_limit", @@ -598,31 +631,46 @@ class AgentRegistry: # Create agent in DB created_agent: Final = await agents_table(prisma_client).create( - data=create_data, - include={"object_permission": True}, + data={**create_data, **_managed_fields(agent, None, created_by)}, + include={"object_permission": True, "identity": True}, ) - created_agent_dict: Final = created_agent.model_dump() - if created_agent.object_permission is not None: - try: - created_agent_dict["object_permission"] = created_agent.object_permission.model_dump() - except Exception: - created_agent_dict["object_permission"] = created_agent.object_permission.dict() - return AgentResponse(**created_agent_dict) + return AgentResponse.model_validate(created_agent.model_dump()) + except HTTPException: + raise except Exception as e: - raise Exception(f"Error adding agent to DB: {e}") + if is_unique_violation(e): + raise HTTPException(409, "Agent name or Entra application is already registered") from e + raise async def delete_agent_from_db(self, agent_id: str, prisma_client: PrismaClient) -> Mapping[str, object]: """ Delete an agent from the database """ - try: - deleted_agent: Final = await agents_table(prisma_client).delete(where={"agent_id": agent_id}) + from prisma.types import ( + LiteLLM_AgentsTableWhereUniqueInput, + LiteLLM_RetiredAgentCreateInput, + LiteLLM_RetiredAgentUpsertInput, + LiteLLM_RetiredAgentWhereUniqueInput, + LiteLLM_VerificationTokenWhereInput, + ) + + where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id} + async with prisma_client.tx() as tx: + existing: Final = await tx.litellm_agentstable.find_unique(where=where) + if existing is None: + raise ValueError(f"Agent not found, passed agent_id={agent_id}") + if existing.identity_managed: + history_where: Final[LiteLLM_RetiredAgentWhereUniqueInput] = {"original_agent_id": agent_id} + history_create: Final = LiteLLM_RetiredAgentCreateInput(original_agent_id=agent_id) + history_data: Final[LiteLLM_RetiredAgentUpsertInput] = {"create": history_create, "update": {}} + await tx.litellm_retiredagent.upsert(where=history_where, data=history_data) + keys_where: Final[LiteLLM_VerificationTokenWhereInput] = {"agent_id": agent_id} + await tx.litellm_verificationtoken.delete_many(where=keys_where) + deleted_agent: Final = await tx.litellm_agentstable.delete(where=where) if deleted_agent is None: raise ValueError(f"Agent not found, passed agent_id={agent_id}") - return dict(deleted_agent) - except Exception as e: - raise Exception(f"Error deleting agent from DB: {e}") + return deleted_agent.model_dump() async def patch_agent_in_db( self, @@ -646,7 +694,9 @@ class AgentRegistry: The patched agent """ try: - existing_record: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id}) + existing_record: Final = await agents_table(prisma_client).find_unique( + where={"agent_id": agent_id}, include={"identity": True} + ) if existing_record is None: raise Exception(f"Agent with ID {agent_id} not found") existing_agent: Final[Mapping[str, object]] = dict(existing_record) @@ -683,37 +733,31 @@ class AgentRegistry: if "extra_headers" in agent: extra_headers_value: Final = agent.get("extra_headers") update_data["extra_headers"] = extra_headers_value if extra_headers_value is not None else [] - if agent.get("object_permission") is not None: - agent_copy: Final = dict(augment_agent) - existing_object_permission_id: Final = existing_record.object_permission_id - object_permission_id: Final = await handle_update_object_permission_common( - agent_copy, - existing_object_permission_id, - prisma_client, - ) - if object_permission_id is not None: - update_data["object_permission_id"] = object_permission_id + permission_write: Final = await _permission_write( + agent, existing_record.object_permission_id, prisma_client + ) + if permission_write is not None: + update_data["object_permission"] = permission_write # Patch agent in DB patched_agent: Final = await agents_table(prisma_client).update( where={"agent_id": agent_id}, data={ **update_data, + **_managed_fields(agent, AgentResponse.model_validate(existing_record.model_dump()), updated_by), "updated_by": updated_by, "updated_at": datetime.now(timezone.utc), }, - include={"object_permission": True}, + include={"object_permission": True, "identity": True}, ) if patched_agent is None: raise ValueError(f"Agent not found, passed agent_id={agent_id}") - patched_agent_dict: Final = patched_agent.model_dump() - if patched_agent.object_permission is not None: - try: - patched_agent_dict["object_permission"] = patched_agent.object_permission.model_dump() - except Exception: - patched_agent_dict["object_permission"] = patched_agent.object_permission.dict() - return AgentResponse(**patched_agent_dict) + return AgentResponse.model_validate(patched_agent.model_dump()) + except HTTPException: + raise except Exception as e: - raise Exception(f"Error patching agent in DB: {e}") + if is_unique_violation(e): + raise HTTPException(409, "Agent name or Entra application is already registered") from e + raise async def update_agent_in_db( self, @@ -733,7 +777,7 @@ class AgentRegistry: # caller echoed back redacted (or omitted) rather than persisting # the marker -- or nothing -- over the real stored credential. existing_row: Final = await agents_table(prisma_client).find_unique( - where={"agent_id": agent_id} # mutable-ok: prisma's query builder rejects a Mapping/MappingProxyType + where={"agent_id": agent_id}, include={"identity": True} ) existing_litellm_params: Final = parse_agent_litellm_params( existing_row.litellm_params if existing_row is not None else None @@ -784,37 +828,35 @@ class AgentRegistry: if _val is not None: update_data[rate_field] = _val - if agent.get("object_permission") is not None: - existing_object_permission_id: Final = ( - existing_row.object_permission_id if existing_row is not None else None - ) - agent_copy: Final = dict(agent) - object_permission_id: Final = await handle_update_object_permission_common( - agent_copy, - existing_object_permission_id, - prisma_client, - ) - if object_permission_id is not None: - update_data["object_permission_id"] = object_permission_id + permission_write: Final = await _permission_write( + agent, existing_row.object_permission_id if existing_row is not None else None, prisma_client + ) + if permission_write is not None: + update_data["object_permission"] = permission_write # Update agent in DB updated_agent: Final = await agents_table(prisma_client).update( where={"agent_id": agent_id}, - data=update_data, - include={"object_permission": True}, + data={ + **update_data, + **_managed_fields( + agent, + AgentResponse.model_validate(existing_row.model_dump()) if existing_row else None, + updated_by, + ), + }, + include={"object_permission": True, "identity": True}, ) if updated_agent is None: raise ValueError(f"Agent not found, passed agent_id={agent_id}") - updated_agent_dict: Final = updated_agent.model_dump() - if updated_agent.object_permission is not None: - try: - updated_agent_dict["object_permission"] = updated_agent.object_permission.model_dump() - except Exception: - updated_agent_dict["object_permission"] = updated_agent.object_permission.dict() - return AgentResponse(**updated_agent_dict) + return AgentResponse.model_validate(updated_agent.model_dump()) + except HTTPException: + raise except Exception as e: - raise Exception(f"Error updating agent in DB: {e}") + if is_unique_violation(e): + raise HTTPException(409, "Agent name or Entra application is already registered") from e + raise @staticmethod async def get_all_agents_from_db( @@ -826,12 +868,12 @@ class AgentRegistry: try: agents_from_db: Final = await agents_table(prisma_client).find_many( order={"created_at": "desc"}, - include={"object_permission": True}, + include={"object_permission": True, "identity": True}, ) agents: Final[list[dict[str, object]]] = [] for agent in agents_from_db: - agent_dict = dict(agent) + agent_dict = agent.model_dump() # object_permission is eagerly loaded via include above if agent.object_permission is not None: try: diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py index 49e5407ff88..6f0cc0393ad 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py +++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py @@ -1,13 +1,16 @@ import asyncio from collections.abc import Awaitable, Callable from dataclasses import dataclass -from typing import Final, TypeAlias +from typing import TYPE_CHECKING, Final, TypeAlias from fastapi import HTTPException from litellm._logging import verbose_proxy_logger from litellm.proxy._types import LiteLLM_AccessGroupTable +if TYPE_CHECKING: + from litellm.types.agents import AgentResponse + AccessGroupIds: TypeAlias = tuple[str, ...] AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None @@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds: return tuple(agent.access_group_ids or ()) if agent is not None else () -async def _load_access_group(access_group_id: str) -> LoadedAccessGroup: +async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup: from litellm.proxy.auth.auth_checks import get_access_object from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache @@ -47,6 +50,7 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup: prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) except HTTPException as e: verbose_proxy_logger.warning( @@ -73,3 +77,16 @@ async def resolve_agent_access_group_ceiling( mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids), agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids), ) + + +async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]: + async def authoritative_group(group_id: str) -> LoadedAccessGroup: + return await _load_access_group(group_id, check_db_only=True) + + async def manual_ids(_agent_id: str) -> AccessGroupIds: + return tuple(agent.access_group_ids or ()) + + manual: Final = await resolve_agent_access_group_ceiling( + agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group + ) + return (manual,) if manual is not None else () diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 9fe74bfee3f..4b9238f1868 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling. import asyncio from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass +from types import MappingProxyType from typing import Final, TypeAlias +from fastapi import HTTPException + from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts from litellm.proxy._types import ( @@ -86,7 +89,9 @@ class AgentRequestHandler: ) -> AgentAccess: """Agents the key may reach: key and team grants, intersected with the agent's access group ceiling and, for an agent key acting on behalf of an invoking user, with that user's team grants.""" - key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth) + if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None: + return await _managed_actor_agent_access(user_api_key_auth) + key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(user_api_key_auth) caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth) own_access: Final = _intersect_agent_access(key_team_access, caller_access) agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling) @@ -104,13 +109,19 @@ class AgentRequestHandler: return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth) @staticmethod - async def _resolve_key_team_agent_access( + async def resolve_key_team_agent_access( user_api_key_auth: UserAPIKeyAuth | None, + *, + strict: bool = False, ) -> AgentAccess: try: - key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth) - team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth) + key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict) + team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team( + user_api_key_auth, strict=strict + ) except Exception as e: + if strict: + raise HTTPException(503, "Agent invocation policy is unavailable") from e verbose_logger.warning("Failed to get allowed agents: %s", e) return UnrestrictedAgentAccess() return _intersect_agent_access(key_access, team_access) @@ -144,6 +155,33 @@ class AgentRequestHandler: bool: True if agent is allowed, False otherwise """ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure + from litellm.proxy.proxy_server import prisma_client + from litellm.types.proxy.agent_identity import AgentIdentityFailure + + registered: Final = global_agent_registry.get_agent_by_id(agent_id) + if prisma_client is not None or (registered is not None and registered.identity_managed): + target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id) + if isinstance(target, AgentIdentityFailure): + raise_identity_failure(target) + if target is None and registered is not None and registered.identity_managed: + return False + if target is not None and target.identity_managed: + if ( + not target.enabled + or target.identity is None + or not target.identity.active + or user_api_key_auth is None + ): + return False + fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True}) + explicit: Final = await _granted_agent_ids( + fresh_auth, + _strict_agent_access, + build_effective_auth_contexts, + ) + return target.agent_id in explicit match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling): case UnrestrictedAgentAccess(): @@ -202,8 +240,10 @@ class AgentRequestHandler: return team_obj.object_permission @staticmethod - async def _get_allowed_agents_for_key( + async def get_allowed_agents_for_key( user_api_key_auth: UserAPIKeyAuth | None = None, + *, + strict: bool = False, ) -> AgentAccess: """ Get allowed agents for a key. @@ -237,24 +277,36 @@ class AgentRequestHandler: return UnrestrictedAgentAccess() access_group_agents: Final = ( - tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups))) + tuple( + await AgentRequestHandler._get_agents_from_access_groups( + list(declared_access_groups), check_db_only=strict + ) + ) if declared_access_groups else () ) unified_agents: Final = ( - tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids))) + tuple( + await AgentRequestHandler._get_unified_access_group_agents( + list(key_access_group_ids), check_db_only=strict + ) + ) if key_access_group_ids else () ) return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents)) except Exception as e: + if strict: + raise HTTPException(503, "Agent invocation policy is unavailable") from e verbose_logger.warning("Failed to get allowed agents for key: %s", e) return UnrestrictedAgentAccess() @staticmethod async def _get_allowed_agents_for_team( user_api_key_auth: UserAPIKeyAuth | None = None, + *, + strict: bool = False, ) -> AgentAccess: """ Get allowed agents for a team. @@ -263,7 +315,7 @@ class AgentRequestHandler: 2. Also includes agents from team's access_group_ids (unified access groups) Fetches the team object once and reuses it for both permission sources. - Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`. + Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`. """ if user_api_key_auth is None: return UnrestrictedAgentAccess() @@ -280,7 +332,7 @@ class AgentRequestHandler: ) if not prisma_client: - return UnrestrictedAgentAccess() + return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess() # Fetch the team object once for both permission sources team_obj: Final = await get_team_object( @@ -289,10 +341,11 @@ class AgentRequestHandler: user_api_key_cache=user_api_key_cache, parent_otel_span=user_api_key_auth.parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=strict, ) if team_obj is None: - return UnrestrictedAgentAccess() + return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess() # 1. Get agents from object_permission (native permissions) object_permissions: Final = team_obj.object_permission @@ -307,18 +360,28 @@ class AgentRequestHandler: return UnrestrictedAgentAccess() access_group_agents: Final = ( - tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups))) + tuple( + await AgentRequestHandler._get_agents_from_access_groups( + list(declared_access_groups), check_db_only=strict + ) + ) if declared_access_groups else () ) unified_agents: Final = ( - tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids))) + tuple( + await AgentRequestHandler._get_unified_access_group_agents( + list(team_access_group_ids), check_db_only=strict + ) + ) if team_access_group_ids else () ) return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents)) except Exception as e: + if strict: + raise HTTPException(503, "Agent invocation policy is unavailable") from e # litellm-dashboard is the default UI team and will never have agents; # skip noisy warnings for it. if user_api_key_auth.team_id != UI_TEAM_ID: @@ -326,7 +389,9 @@ class AgentRequestHandler: return UnrestrictedAgentAccess() @staticmethod - def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]: + def _get_config_agent_ids_for_access_groups( + config_agents: Sequence[AgentResponse], access_groups: list[str] + ) -> set[str]: """ Helper to get agent_ids from config-loaded agents that match any of the given access groups. """ @@ -339,7 +404,9 @@ class AgentRequestHandler: return server_ids @staticmethod - async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]: + async def _get_db_agent_ids_for_access_groups( + prisma_client, access_groups: list[str], *, check_db_only: bool = False + ) -> set[str]: """ Helper to get agent_ids from DB agents that match any of the given access groups. @@ -349,23 +416,27 @@ class AgentRequestHandler: if not access_groups or prisma_client is None: return set() - agents: Final = await AgentsRepository(prisma_client).table.find_many( + agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many( where={"agent_access_groups": {"hasSome": access_groups}} ) return {agent.agent_id for agent in agents} @staticmethod - async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]: + async def _get_unified_access_group_agents( + access_group_ids: list[str], *, check_db_only: bool = False + ) -> list[str]: """ Resolve unified access group ids to agent IDs. """ from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups - return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids) + return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only) @staticmethod async def _get_agents_from_access_groups( access_groups: list[str], + *, + check_db_only: bool = False, ) -> list[str]: """ Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents. @@ -373,14 +444,13 @@ class AgentRequestHandler: from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry from litellm.proxy.proxy_server import prisma_client - # Use the helper for config-loaded agents config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups( global_agent_registry.agent_list, access_groups ) # Use the helper for DB agents db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups( - prisma_client, access_groups + prisma_client, access_groups, check_db_only=check_db_only ) return list(config_agent_ids | db_agent_ids) @@ -531,4 +601,58 @@ async def accessible_agents( AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access, effective_contexts, ) - return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids) + allowed: Final = await asyncio.gather( + *( + AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth) + for agent in agents + if agent.identity_managed + ) + ) + managed_ids: Final = frozenset( + agent.agent_id + for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed) + if permitted + ) + return tuple( + agent + for agent in agents + if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids) + ) + + +async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess: + if auth.managed_agent_policy is not None: + return await _managed_actor_agent_access(auth) + return await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True) + + +async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess: + agent: Final = auth.managed_agent_policy + if agent is None or not agent.object_permission: + return RestrictedAgentAccess(frozenset()) + permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({})) + own_auth: Final = UserAPIKeyAuth(object_permission=permission) + own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True)) + + from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings + + ceilings: Final = await resolve_managed_agent_ceilings(agent) + capped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings)) + context: Final = auth.managed_agent_context + if context is None or context.mode == "autonomous": + return RestrictedAgentAccess(capped) + if context.user_id is None: + return RestrictedAgentAccess(frozenset()) + human_ids: Final = await verified_human_agent_grants(context.user_id) + return RestrictedAgentAccess(capped.intersection(human_ids)) + + +async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + if user_id is None: + return frozenset() + human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True) + sources: Final = await MCPRequestHandler.admitted_subject_sources(human) + human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources)) + return frozenset().union(*(_granted_ids(access) for access in human_access)) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py new file mode 100644 index 00000000000..abcef353eef --- /dev/null +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -0,0 +1,232 @@ +from collections.abc import Mapping +from itertools import product +from types import MappingProxyType +from typing import Annotated, Final + +from pydantic import Field, TypeAdapter, ValidationError + +from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore +from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext + +_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime")) +_MANAGED_MODEL_ROUTES: Final = frozenset( + f"{prefix}/{operation}" + for prefix, operation in product( + ("", "/v1"), + ( + "chat/completions", + "completions", + "embeddings", + "responses", + "messages", + "messages/count_tokens", + "images/generations", + "images/edits", + "audio/transcriptions", + "audio/speech", + "moderations", + "rerank", + "ocr", + ), + ) +) | frozenset( + ( + "/openai/v1/responses", + "/v2/rerank", + "/claude_code_gateway/v1/messages", + "/claude_code_gateway/v1/messages/count_tokens", + "/cursor/chat/completions", + ) +) +_MANAGED_MODEL_PATHS: Final = ( + "/engines/{model:path}/chat/completions", + "/engines/{model:path}/completions", + "/engines/{model:path}/embeddings", + "/openai/deployments/{model:path}/chat/completions", + "/openai/deployments/{model:path}/completions", + "/openai/deployments/{model:path}/embeddings", + "/openai/deployments/{model:path}/images/generations", + "/openai/deployments/{model:path}/images/edits", + "/v1beta/models/{model_name:path}:countTokens", + "/v1beta/models/{model_name:path}:generateContent", + "/v1beta/models/{model_name:path}:streamGenerateContent", + "/models/{model_name:path}:countTokens", + "/models/{model_name:path}:generateContent", + "/models/{model_name:path}:streamGenerateContent", +) +_MANAGED_MCP_ROUTES: Final = tuple( + route for route in LiteLLMRoutes.mcp_inference_routes.value if route not in ("/token", "/introspect") +) + + +def managed_agent_route_allowed(route: str, method: str | None) -> bool: + from litellm.proxy.auth.route_checks import RouteChecks + + if route in ("/agents", "/v1/agents"): + return method in (None, "GET", "HEAD") + if route in _MANAGED_REALTIME_ROUTES: + return method in (None, "GET") + if route in _MANAGED_MODEL_ROUTES or RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS): + return method in (None, "POST") + return RouteChecks.check_route_access(route, _MANAGED_MCP_ROUTES) or RouteChecks.check_route_access( + route, LiteLLMRoutes.agent_inference_routes.value + ) + + +def managed_inference_request( + route: str, + body: Mapping[str, object], + settings: Mapping[str, object], + cli_model: str | None, + path_model: object = None, + query_model: object = None, +) -> dict[str, object]: + from litellm.proxy.auth.route_checks import RouteChecks + + if route in _MANAGED_REALTIME_ROUTES: + model: Final = query_model or body.get("model") + if not isinstance(model, str) or not model: + raise_identity_failure( + AgentIdentityFailure(message="Managed inference requires an explicit or configured model") + ) + return {**body, "model": model} + if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS): + return dict(body) + from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model + + kind: Final = ( + "image_generation" + if route.endswith("/images/generations") + else "image_edit" + if route.endswith("/images/edits") + else "moderation" + if route.endswith(("/moderations", "/audio/transcriptions")) + else "speech" + if route.endswith("/audio/speech") + else "body" + if route.endswith(("/rerank", "/messages/count_tokens")) + else "path" + if route.endswith(":countTokens") + else "completion" + ) + endpoint_model: Final = path_model or ( + query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None + ) + effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind) + if not isinstance(effective, str) or not effective: + raise_identity_failure( + AgentIdentityFailure(message="Managed inference requires an explicit or configured model") + ) + return {**body, "model": effective} + + +async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None: + if auth.agent_id is None: + return + if store is None: + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id) + if auth.managed_agent_context is not None or ( + registered is not None and (registered.identity_managed or registered.identity is not None) + ): + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database") + ) + return + agent: Final = await store.agent(auth.agent_id) + if isinstance(agent, AgentIdentityFailure): + raise_identity_failure(agent) + if agent is None: + retired: Final = await store.retired_agent(auth.agent_id) + if isinstance(retired, AgentIdentityFailure): + raise_identity_failure(retired) + if auth.managed_agent_context is not None or retired: + raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists")) + return + if not agent.identity_managed: + return + if auth.jwt_claims and auth.managed_agent_context is None: + raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity")) + failure: Final = actor_admission_failure(agent, auth.managed_agent_context) + if failure is not None: + raise_identity_failure(failure) + auth.managed_agent_policy = agent + auth.billing_agent_policy = agent + if auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated": + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants + + grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id) + if agent.agent_id not in grants: + raise_identity_failure( + AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent") + ) + + +def actor_admission_failure( + agent: AgentResponse, + context: ManagedAgentContext | None, +) -> AgentIdentityFailure | None: + if not agent.enabled or agent.identity is None or not agent.identity.active: + return AgentIdentityFailure(message="Agent execution is disabled") + if context is None: + return AgentIdentityFailure(message="This agent requires its bound identity provider token") + if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision: + return AgentIdentityFailure(message="Agent identity changed during authentication; retry") + if agent.execution_mode not in (context.mode, "both"): + return AgentIdentityFailure(message="Agent is not enabled for this execution mode") + if context.mode == "delegated" and not context.user_id: + return AgentIdentityFailure(message="A verified human subject is required") + return None + + +_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)]) + + +def invocation_target(route: str, body: Mapping[str, object]) -> str | None: + model: Final = body.get("model") + if isinstance(model, str) and model.startswith("a2a/"): + return model.removeprefix("a2a/") or None + components: Final = tuple(route.strip("/").split("/")) + path: Final = components[1:] if components and components[0] == "v1" else components + return path[1] if len(path) >= 2 and path[0] == "a2a" else None + + +async def prepare_agent_invocation( + auth: UserAPIKeyAuth, target_name: str, store: AgentIdentityStore | None, *, billable: bool = True +) -> None: + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler + from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through + + registered: Final = await get_agent_with_read_through(target_name) + if registered is None: + return + registered_managed: Final = registered.identity_managed or registered.identity is not None + if store is None and registered_managed: + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database") + ) + target: Final = await store.agent(registered.agent_id) if store is not None else None + if isinstance(target, AgentIdentityFailure): + raise_identity_failure(target) + if target is None and registered_managed: + raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists")) + effective: Final = target if target is not None else registered + if not effective.identity_managed and auth.managed_agent_policy is None: + return + if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth): + raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent")) + auth.invoked_agent_id = effective.agent_id + if auth.agent_id is None and effective.identity_managed: + auth.billing_agent_policy = effective + raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0 + try: + fee: Final = _INVOCATION_COST.validate_python(raw_fee) + except ValidationError: + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid") + ) + auth.agent_invocation_cost = fee diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 28c82a715e0..e1b2ac63d51 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -16,6 +16,7 @@ from types import MappingProxyType from typing import Annotated, Final, TypedDict from fastapi import APIRouter, Depends, HTTPException, Query, Request +from pydantic import ValidationError from typing_extensions import ReadOnly, Required, assert_never import litellm @@ -47,6 +48,8 @@ from litellm.proxy.agent_endpoints.agent_search import ( search_agents, ) from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents +from litellm.proxy.agent_endpoints.identity import reject_legacy_identity +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore from litellm.proxy.agent_endpoints.kill_switch import ( KillSwitchAuditLogWriter, KillSwitchHttpClient, @@ -56,6 +59,7 @@ from litellm.proxy.agent_endpoints.kill_switch import ( fire_kill_switch, redact_kill_switch, ) +from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity @@ -72,6 +76,12 @@ from litellm.types.agents import ( PatchAgentRequest, ) from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.types.proxy.agent_identity import ( + AgentIdentityBinding, + AgentIdentityFailure, + EntraIdentityConfig, + ManagedAgentIdentityStatus, +) from litellm.types.proxy.management_endpoints.common_daily_activity import ( DailySpendMetadata, SpendAnalyticsPaginatedResponse, @@ -178,14 +188,21 @@ def _redact_sensitive_agent_fields( virtual-key, header and kill-switch fields stripped entirely. The original objects are not modified. """ + from litellm.proxy.proxy_server import general_settings, jwt_handler + redacted: Final[list[AgentResponse]] = [] for agent in agents: copy = agent.model_copy(deep=True) + copy.jwt_auth_configured = bool( + general_settings.get("enable_jwt_auth") + and (agent.identity is not None or jwt_handler.litellm_jwtauth.agent_id_jwt_field) + ) if not is_admin: copy.static_headers = None copy.extra_headers = None copy.keys = None copy.kill_switch = None + copy.identity = None if copy.litellm_params: copy.litellm_params = _redact_agent_litellm_params_dict(copy.litellm_params) copy.kill_switch = redact_kill_switch(copy.kill_switch) @@ -429,6 +446,71 @@ from litellm.proxy.agent_endpoints.agent_registry import ( ) +def _trusted_agent_issuers() -> tuple[str, ...]: + from litellm.proxy.proxy_server import general_settings, jwt_handler + + if not general_settings.get("enable_jwt_auth"): + return () + configured: Final = jwt_handler.litellm_jwtauth.issuers or () + issuer: Final = os.getenv("JWT_ISSUER") + global_issuers: Final = ( + (issuer,) + if issuer and os.getenv("JWT_AUDIENCE") and not any(item.issuer == issuer for item in configured) + else () + ) + return ( + tuple(item.issuer for item in configured if item.audience and not item.disable_audience_validation) + + global_issuers + ) + + +def _validate_managed_identity_request( + request: AgentConfig | PatchAgentRequest, existing: AgentResponse | None = None +) -> None: + raw: Final = request.get("identity") if "identity" in request else existing.identity if existing else None + if raw is None: + return + try: + identity: Final = raw if isinstance(raw, AgentIdentityBinding) else EntraIdentityConfig.model_validate(raw) + except ValidationError as exc: + raise HTTPException(400, "Invalid Entra identity configuration") from exc + if identity.issuer not in _trusted_agent_issuers(): + raise HTTPException(400, "Configure trusted JWT issuer and audience validation for this Entra tenant first") + if request.get("execution_mode", existing.execution_mode if existing else "autonomous") != "autonomous": + if os.getenv("MICROSOFT_TENANT") != identity.tenant_id or not os.getenv("MICROSOFT_CLIENT_ID"): + raise HTTPException(400, "Delegated agents require Microsoft SSO for the same trusted tenant") + + +@router.get("/v1/agents/identity/providers", response_model=tuple[str, ...], tags=("[beta] A2A Agents",)) +async def get_agent_identity_providers( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> tuple[str, ...]: + _check_agent_management_permission(user_api_key_dict) + return _trusted_agent_issuers() + + +@router.get("/v1/agents/{agent_id}/identity", response_model=ManagedAgentIdentityStatus, tags=("[beta] A2A Agents",)) +async def get_agent_identity_status( + agent_id: str, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> ManagedAgentIdentityStatus: + from litellm.proxy.proxy_server import prisma_client + + _check_agent_management_permission(user_api_key_dict) + agent: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id) + if isinstance(agent, AgentIdentityFailure): + raise_identity_failure(agent) + if agent is None: + raise HTTPException(404, "Agent not found") + return ManagedAgentIdentityStatus( + identity=agent.identity, + identity_managed=agent.identity_managed, + enabled=agent.enabled, + execution_mode=agent.execution_mode, + last_authenticated_at=agent.identity.last_authenticated_at if agent.identity else None, + ) + + @router.post( "/v1/agents", tags=["[beta] A2A Agents"], @@ -490,6 +572,9 @@ async def create_agent( # Get the user ID from the API key auth created_by: Final = user_api_key_dict.user_id or "unknown" + _validate_managed_identity_request(request) + reject_legacy_identity(request.get("litellm_params")) + # check for naming conflicts existing_agent: Final = AGENT_REGISTRY.get_agent_by_name(agent_name=request.get("agent_name")) if existing_agent is not None: @@ -591,7 +676,7 @@ async def get_agent_by_id( if agent is None: agent_row: Final = await agents_table(prisma_client).find_unique( where={"agent_id": agent_id}, - include={"object_permission": True}, + include={"object_permission": True, "identity": True}, ) if agent_row is not None: agent_dict: Final = agent_row.model_dump() @@ -680,13 +765,18 @@ async def update_agent( try: # Check if agent exists - existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id}) + existing_agent = await agents_table(prisma_client).find_unique( + where={"agent_id": agent_id}, include={"identity": True} + ) if existing_agent is not None: - existing_agent = dict(existing_agent) + existing_agent = existing_agent.model_dump() if existing_agent is None: raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found") + _validate_managed_identity_request(request, AgentResponse.model_validate(existing_agent)) + reject_legacy_identity(request.get("litellm_params")) + # Get the user ID from the API key auth updated_by: Final = user_api_key_dict.user_id or "unknown" @@ -782,13 +872,18 @@ async def patch_agent( try: # Check if agent exists - existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id}) + existing_agent = await agents_table(prisma_client).find_unique( + where={"agent_id": agent_id}, include={"identity": True} + ) if existing_agent is not None: - existing_agent = dict(existing_agent) + existing_agent = existing_agent.model_dump() if existing_agent is None: raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found") + _validate_managed_identity_request(request, AgentResponse.model_validate(existing_agent)) + reject_legacy_identity(request.get("litellm_params")) + # Get the user ID from the API key auth updated_by: Final = user_api_key_dict.user_id or "unknown" @@ -869,7 +964,9 @@ async def delete_agent( try: # Check if agent exists - existing_agent = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id}) + existing_agent = await agents_table(prisma_client).find_unique( + where={"agent_id": agent_id}, include={"identity": True} + ) if existing_agent is not None: existing_agent = dict[str, object](existing_agent) diff --git a/litellm/proxy/agent_endpoints/identity.py b/litellm/proxy/agent_endpoints/identity.py new file mode 100644 index 00000000000..c0e5a748144 --- /dev/null +++ b/litellm/proxy/agent_endpoints/identity.py @@ -0,0 +1,17 @@ +from collections.abc import Mapping +from typing import Final + +from fastapi import HTTPException + +LEGACY_IDENTITY_MESSAGE: Final = ( + "litellm_params.identity is not supported: bind an Entra application through the top-level identity field" +) + + +def has_legacy_identity(params: Mapping[str, object] | None) -> bool: + return params is not None and "identity" in params + + +def reject_legacy_identity(params: Mapping[str, object] | None) -> None: + if has_legacy_identity(params): + raise HTTPException(400, LEGACY_IDENTITY_MESSAGE) diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py new file mode 100644 index 00000000000..87512c779fe --- /dev/null +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -0,0 +1,224 @@ +from collections.abc import Mapping +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Final + +from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject +from litellm.repositories.table_repositories import ( + AgentIdentityRepository, + AgentsRepository, + RetiredAgentIdentityRepository, + RetiredAgentRepository, + VerifiedSubjectRepository, +) +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import ( + AgentIdentityFailure, + ManagedAgentContext, + MicrosoftInteractiveSubject, + VerifiedHumanSubject, +) + +if TYPE_CHECKING: + from prisma.models import LiteLLM_VerifiedSubject + from prisma.types import ( + LiteLLM_AgentIdentityUpdateManyMutationInput, + LiteLLM_AgentIdentityWhereInput, + LiteLLM_AgentIdentityWhereUniqueInput, + LiteLLM_AgentsTableInclude, + LiteLLM_AgentsTableWhereUniqueInput, + LiteLLM_VerifiedSubjectCreateInput, + LiteLLM_VerifiedSubjectUpsertInput, + LiteLLM_VerifiedSubjectWhereUniqueInput, + ) + + +class AgentIdentityStore: + @classmethod + def from_client(cls, client: object) -> "AgentIdentityStore": + return cls( + AgentsRepository(client, use_writer=True), + AgentIdentityRepository(client, use_writer=True), + VerifiedSubjectRepository(client, use_writer=True), + RetiredAgentIdentityRepository(client, use_writer=True), + RetiredAgentRepository(client, use_writer=True), + ) + + def __init__( + self, + agents: AgentsRepository, + identities: AgentIdentityRepository, + humans: VerifiedSubjectRepository, + retired: RetiredAgentIdentityRepository | None = None, + retired_agents: RetiredAgentRepository | None = None, + ) -> None: + self.agents = agents + self.identities = identities + self.humans = humans + self.retired = retired + self.retired_agents = retired_agents + + async def agent(self, agent_id: str) -> AgentResponse | AgentIdentityFailure | None: + try: + where: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id} + include: Final[LiteLLM_AgentsTableInclude] = { + "identity": True, + "object_permission": True, + } + row: Final = await self.agents.table.find_unique(where=where, include=include) + if row is None: + return None + return AgentResponse.model_validate(row.model_dump()) + except Exception: + return AgentIdentityFailure(code="policy_unavailable", message="Agent policy could not be loaded") + + async def unbound_client(self, where: "LiteLLM_AgentIdentityWhereUniqueInput") -> AgentIdentityFailure | None: + if self.retired is not None: + try: + retired: Final = await self.retired.table.find_unique(where=where) + except Exception: + return AgentIdentityFailure( + code="policy_unavailable", message="Retired agent identity could not be checked" + ) + if retired is not None: + return AgentIdentityFailure(message="This agent identity binding has been retired") + return None + + async def resolve_verified_claims( + self, claims: Mapping[str, object] + ) -> ManagedAgentContext | AgentIdentityFailure | None: + issuer: Final = claims.get("iss") + tenant: Final = claims.get("tid") + client: Final = claims.get("azp") + if not isinstance(issuer, str) or not isinstance(tenant, str) or not isinstance(client, str): + return None + where: Final[LiteLLM_AgentIdentityWhereUniqueInput] = { + "provider_tenant_id_client_id": {"provider": "microsoft_entra", "tenant_id": tenant, "client_id": client} + } + try: + row: Final = await self.identities.table.find_unique(where=where) + except Exception: + return AgentIdentityFailure(code="policy_unavailable", message="Agent identity could not be loaded") + if row is None: + return await self.unbound_client(where) + agent: Final = await self.agent(row.agent_id) + if isinstance(agent, AgentIdentityFailure): + return agent + if ( + agent is None + or not agent.identity_managed + or not agent.enabled + or agent.identity is None + or not agent.identity.active + ): + return AgentIdentityFailure(message="Agent is disabled or no longer bound to an identity") + subject: Final = classify_agent_subject(agent.identity, claims, agent.execution_mode) + if isinstance(subject, AgentIdentityFailure): + return subject + if subject.kind == "application": + return ManagedAgentContext( + agent_id=agent.agent_id, + binding_revision=agent.identity.revision, + mode=subject.mode, + subject_oid=subject.oid, + ) + proven: Final = await self.subject(issuer, tenant, claims.get("oid")) + if isinstance(proven, AgentIdentityFailure): + return proven + human: Final = ( + VerifiedHumanSubject.model_validate(proven.model_dump()) + if proven is not None + and proven.kind == "human" + and proven.verified_via == "sso_interactive" + and proven.user_id is not None + else None + ) + if human is None: + return AgentIdentityFailure(message="The delegated user must first sign in through trusted Microsoft SSO") + return ManagedAgentContext( + agent_id=agent.agent_id, + binding_revision=agent.identity.revision, + mode=subject.mode, + user_id=human.user_id, + subject_oid=subject.oid, + ) + + async def subject( + self, issuer: str, tenant_id: str, oid: object + ) -> "LiteLLM_VerifiedSubject | AgentIdentityFailure | None": + if not isinstance(oid, str): + return None + try: + where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = { + "issuer_tenant_id_oid": {"issuer": issuer, "tenant_id": tenant_id, "oid": oid} + } + return await self.humans.table.find_unique(where=where) + except Exception: + return AgentIdentityFailure(code="policy_unavailable", message="Subject classification is unavailable") + + async def retired_agent(self, agent_id: str) -> bool | AgentIdentityFailure: + if self.retired_agents is None: + return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable") + try: + return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None + except Exception: + return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable") + + async def record_authentication(self, context: ManagedAgentContext) -> AgentIdentityFailure | None: + try: + if context.binding_revision is None: + return AgentIdentityFailure(message="Agent authentication requires a binding revision") + where: Final[LiteLLM_AgentIdentityWhereInput] = { + "agent_id": context.agent_id, + "revision": context.binding_revision, + "active": True, + "agent": {"is": {"enabled": True, "identity_managed": True}}, + } + data: Final[LiteLLM_AgentIdentityUpdateManyMutationInput] = { + "last_authenticated_at": datetime.now(timezone.utc) + } + count: Final = await self.identities.table.update_many(where=where, data=data) + if count != 1: + return AgentIdentityFailure(message="Agent identity changed during authentication; retry") + return None + except Exception: + return AgentIdentityFailure(code="policy_unavailable", message="Agent authentication could not be recorded") + + async def enroll_interactive_human( + self, + subject: MicrosoftInteractiveSubject, + user_id: str, + ) -> AgentIdentityFailure | None: + try: + where: Final[LiteLLM_VerifiedSubjectWhereUniqueInput] = { + "issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": subject.tenant_id, "oid": subject.oid} + } + create_data: Final[LiteLLM_VerifiedSubjectCreateInput] = { + "issuer": subject.issuer, + "tenant_id": subject.tenant_id, + "oid": subject.oid, + "user_id": user_id, + "verified_via": "sso_interactive", + } + data: Final[LiteLLM_VerifiedSubjectUpsertInput] = {"create": create_data, "update": {}} + row: Final = await self.humans.table.upsert(where=where, data=data) + if row.kind != "human" or row.user_id != user_id or row.verified_via != "sso_interactive": + return AgentIdentityFailure(message="Microsoft subject is already bound to another local identity") + return None + except Exception: + return AgentIdentityFailure( + code="policy_unavailable", message="Microsoft subject enrollment is unavailable" + ) + + +async def resolve_managed_agent( + claims: Mapping[str, object], + client: object, +) -> ManagedAgentContext | None: + from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure + + if client is None: + return None + result: Final = await AgentIdentityStore.from_client(client).resolve_verified_claims(claims) + if isinstance(result, AgentIdentityFailure): + raise_identity_failure(result) + return result diff --git a/litellm/proxy/agent_endpoints/managed_identity.py b/litellm/proxy/agent_endpoints/managed_identity.py new file mode 100644 index 00000000000..260b74fcbd1 --- /dev/null +++ b/litellm/proxy/agent_endpoints/managed_identity.py @@ -0,0 +1,220 @@ +from collections.abc import Mapping +from datetime import datetime +from typing import Final, NoReturn, TypedDict +from uuid import uuid4 + +from fastapi import HTTPException +from pydantic import TypeAdapter, ValidationError +from typing_extensions import ReadOnly + +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import ( + AgentExecutionMode, + AgentIdentityBinding, + AgentIdentityFailure, + AgentSubject, + EntraIdentityConfig, +) + +_MODE: Final = TypeAdapter(AgentExecutionMode) + + +class IdentityFields(TypedDict, total=False): + provider: ReadOnly[str] + tenant_id: ReadOnly[str] + client_id: ReadOnly[str] + issuer: ReadOnly[str] + service_principal_id: ReadOnly[str | None] + required_roles: ReadOnly[tuple[str, ...]] + required_scopes: ReadOnly[tuple[str, ...]] + active: ReadOnly[bool] + revision: ReadOnly[str] + last_authenticated_at: ReadOnly[datetime | None] + + +class IdentityUpsert(TypedDict): + create: ReadOnly[IdentityFields] + update: ReadOnly[IdentityFields] + + +class IdentityRelationWrite(TypedDict, total=False): + create: ReadOnly[IdentityFields] + update: ReadOnly[IdentityFields] + upsert: ReadOnly[IdentityUpsert] + + +class IdentityHistoryKey(TypedDict): + provider: ReadOnly[str] + tenant_id: ReadOnly[str] + client_id: ReadOnly[str] + + +class IdentityHistoryWhere(TypedDict): + provider_tenant_id_client_id: ReadOnly[IdentityHistoryKey] + + +class IdentityHistoryEntry(IdentityHistoryKey): + issuer: ReadOnly[str] + + +class IdentityHistoryConnect(TypedDict): + where: ReadOnly[IdentityHistoryWhere] + create: ReadOnly[IdentityHistoryEntry] + + +class IdentityHistoryWrite(TypedDict): + connectOrCreate: ReadOnly[IdentityHistoryConnect] + + +class ManagedWriteFields(TypedDict, total=False): + enabled: ReadOnly[bool] + execution_mode: ReadOnly[AgentExecutionMode] + identity_managed: ReadOnly[bool] + identity: ReadOnly[IdentityRelationWrite] + retired_identities: ReadOnly[IdentityHistoryWrite] + + +def raise_identity_failure(failure: AgentIdentityFailure, status_code: int = 403) -> NoReturn: + raise HTTPException(503 if failure.code == "policy_unavailable" else status_code, failure.message) + + +def _configuration_failure( + identity: EntraIdentityConfig | AgentIdentityBinding | None, + mode: AgentExecutionMode, + enabling_without_binding: bool, +) -> AgentIdentityFailure | None: + if identity is not None and mode != "delegated" and not identity.service_principal_id: + return AgentIdentityFailure( + message="Autonomous mode requires the Enterprise application service-principal object ID" + ) + if enabling_without_binding and ( + identity is None or isinstance(identity, AgentIdentityBinding) and not identity.active + ): + return AgentIdentityFailure(message="Bind an identity before enabling this managed agent") + return None + + +def managed_write_fields( + incoming: Mapping[str, object], + existing: AgentResponse | None, + updated_by: str, +) -> ManagedWriteFields | AgentIdentityFailure: + try: + identity: Final = ( + EntraIdentityConfig.model_validate(incoming["identity"]) if incoming.get("identity") is not None else None + ) + mode: Final = _MODE.validate_python( + incoming.get("execution_mode", existing.execution_mode if existing else "autonomous") + ) + current_identity: Final = identity if "identity" in incoming else existing.identity if existing else None + failure: Final = _configuration_failure( + current_identity, + mode, + incoming.get("enabled") is True + and "identity" not in incoming + and bool(existing and existing.identity_managed), + ) + if failure is not None: + return failure + empty: Final[ManagedWriteFields] = {} + identity_fields: Final = _identity_write(identity, existing) if "identity" in incoming else empty + result: Final[ManagedWriteFields] = { + **({"enabled": incoming["enabled"] is True} if "enabled" in incoming else {}), + **({"execution_mode": mode} if "execution_mode" in incoming else {}), + **identity_fields, + } + return result + except (ValidationError, ValueError) as exc: + return AgentIdentityFailure(message=f"Invalid agent identity configuration: {exc}") + + +def _identity_write(identity: EntraIdentityConfig | None, existing: AgentResponse | None) -> ManagedWriteFields: + if identity is None: + unbind: Final[ManagedWriteFields] = { + **( + {"identity": {"update": {"active": False, "revision": str(uuid4()), "last_authenticated_at": None}}} + if existing and existing.identity + else {} + ), + **({"identity_managed": True, "enabled": False} if existing and existing.identity_managed else {}), + } + return unbind + if ( + existing + and existing.identity + and existing.identity.active + and all(getattr(existing.identity, name) == value for name, value in identity.model_dump().items()) + ): + unchanged: Final[ManagedWriteFields] = {} + return unchanged + binding: Final[IdentityFields] = { + "provider": identity.provider, + "tenant_id": identity.tenant_id, + "client_id": identity.client_id, + "service_principal_id": identity.service_principal_id, + "required_roles": identity.required_roles, + "required_scopes": identity.required_scopes, + "issuer": identity.issuer, + "active": True, + "revision": str(uuid4()), + "last_authenticated_at": None, + } + result: Final[ManagedWriteFields] = { + "retired_identities": { + "connectOrCreate": { + "where": { + "provider_tenant_id_client_id": { + "provider": identity.provider, + "tenant_id": identity.tenant_id, + "client_id": identity.client_id, + } + }, + "create": { + "provider": identity.provider, + "issuer": identity.issuer, + "tenant_id": identity.tenant_id, + "client_id": identity.client_id, + }, + } + }, + "identity_managed": True, + "identity": {"upsert": {"create": binding, "update": binding}} if existing else {"create": binding}, + } + return result + + +def classify_agent_subject( + binding: AgentIdentityBinding, + claims: Mapping[str, object], + allowed_mode: AgentExecutionMode, +) -> AgentSubject | AgentIdentityFailure: + if (claims.get("iss"), claims.get("tid"), claims.get("azp")) != ( + binding.issuer, + binding.tenant_id, + binding.client_id, + ): + return AgentIdentityFailure(message="Token does not match the registered Entra application") + oid: Final = claims.get("oid") + if not isinstance(oid, str) or not oid: + return AgentIdentityFailure(message="Entra token must identify its object subject") + scope: Final = claims.get("scp") + facets: Final = claims.get("xms_sub_fct") + if facets is not None and (not isinstance(facets, str) or "13" in facets.split()): + return AgentIdentityFailure(message="Native agent-user authentication is not supported by this binding") + if scope is not None and not isinstance(scope, str): + return AgentIdentityFailure(message="Invalid delegated scope claim") + if isinstance(scope, str) and scope: + if allowed_mode == "autonomous" or oid == binding.service_principal_id or claims.get("idtyp") == "app": + return AgentIdentityFailure(message="Delegated token contradicts the configured agent identity or mode") + granted_scopes: Final = frozenset(scope.split()) + if not granted_scopes or not frozenset(binding.required_scopes).issubset(granted_scopes): + return AgentIdentityFailure(message="Token lacks the required delegated scopes") + return AgentSubject(kind="delegated_subject", oid=oid, mode="delegated") + if allowed_mode == "delegated" or oid != binding.service_principal_id or claims.get("idtyp") == "user": + return AgentIdentityFailure(message="Application token contradicts the configured agent identity or mode") + roles: Final = claims.get("roles", ()) + if not isinstance(roles, (list, tuple)) or any(not isinstance(role, str) for role in roles): + return AgentIdentityFailure(message="Invalid application roles claim") + if not frozenset(binding.required_roles).issubset(roles): + return AgentIdentityFailure(message="Token lacks the required application roles") + return AgentSubject(kind="application", oid=oid, mode="autonomous") diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 12d420141f1..11c7dbd64cf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1069,6 +1069,21 @@ async def common_checks( code=status.HTTP_400_BAD_REQUEST, ) + if _model and valid_token is not None and valid_token.managed_agent_policy is not None: + managed_models: Final = (valid_token.managed_agent_policy.object_permission or MappingProxyType({})).get( + "models", () + ) + if not isinstance(managed_models, (list, tuple)) or not managed_models: + raise HTTPException(403, "This agent has no model grants") + _can_object_call_model( + model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router), + llm_router=llm_router, + models=list(managed_models), + team_id=valid_token.team_id, + object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), + ) + await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router) await _check_agent_caller_model_access( model=_model, @@ -2653,7 +2668,7 @@ async def get_user_object( ) if should_check_db: - response = await _user_table(UserRepository(prisma_client)).find_unique( + response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).find_unique( where={"user_id": user_id}, include={"organization_memberships": True} ) @@ -2691,7 +2706,7 @@ async def get_user_object( budget_duration=new_user_params["budget_duration"] ) - response = await _user_table(UserRepository(prisma_client)).create( + response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).create( data=new_user_params, include={"organization_memberships": True}, ) @@ -3116,9 +3131,9 @@ class TeamNotFoundError(HTTPException): @log_db_metrics async def _get_team_db_check( - team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None + team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None, *, use_writer: bool = False ) -> "_PrismaTeamRow | None": - response = await _team_table(TeamRepository(prisma_client)).find_unique( + response = await _team_table(TeamRepository(prisma_client, use_writer=use_writer)).find_unique( where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS ) @@ -3152,6 +3167,7 @@ async def _get_team_object_from_user_api_key_cache( proxy_logging_obj: ProxyLogging | None, key: str, team_id_upsert: bool | None = None, + use_writer: bool = False, ) -> LiteLLM_TeamTableCachedObj: db_access_time_key: Final = key should_check_db: Final = _should_check_db( @@ -3160,7 +3176,9 @@ async def _get_team_object_from_user_api_key_cache( db_cache_expiry=db_cache_expiry, ) if should_check_db: - response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert) + response = await _get_team_db_check( + team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert, use_writer=use_writer + ) # The database answered and the row is not there. Distinct from every # other failure here, which leaves the team's grant unknown. if response is None: @@ -3182,8 +3200,11 @@ async def _get_team_object_from_user_api_key_cache( user_api_key_cache=user_api_key_cache, parent_otel_span=None, proxy_logging_obj=proxy_logging_obj, + check_db_only=use_writer, ) except Exception as e: + if use_writer: + raise verbose_proxy_logger.debug( "Failed to load object_permission for team %s with object_permission_id=%s: %s", team_id, @@ -3273,6 +3294,7 @@ async def get_team_object( db_cache_expiry=db_cache_expiry, key=key, team_id_upsert=team_id_upsert, + use_writer=bool(check_db_only), ) except TeamNotFoundError: raise @@ -3318,16 +3340,15 @@ async def get_access_object( prisma_client: DatabaseClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None = None, + *, + check_db_only: bool = False, ) -> LiteLLM_AccessGroupTable: """ - Check if access_group_id in proxy AccessGroupTable - - Always checks cache first, then DB only when not found in cache + - Checks cache first unless authoritative writer admission is requested - if valid, return LiteLLM_AccessGroupTable object - if not, then raise an error - Unlike get_team_object, this has no check_cache_only or check_db_only flags; - it always follows cache-first-then-db semantics. - Raises: - HTTPException: If access group doesn't exist in db or cache (status_code=404) """ @@ -3336,18 +3357,19 @@ async def get_access_object( key: Final = f"access_group_id:{access_group_id}" - cached_access_obj: Final = await user_api_key_cache.async_get_cache( - key=key, - model_type=LiteLLM_AccessGroupTable, + cached_access_obj: Final = ( + None + if check_db_only + else await user_api_key_cache.async_get_cache(key=key, model_type=LiteLLM_AccessGroupTable) ) if cached_access_obj is not None: return cached_access_obj # Not in cache - fetch from DB try: - response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique( - where={"access_group_id": access_group_id} - ) + response: Final = await _dictable_table( + AccessGroupRepository(prisma_client, use_writer=check_db_only), "access_group" + ).find_unique(where={"access_group_id": access_group_id}) if response is None: raise HTTPException( @@ -3931,6 +3953,7 @@ async def get_object_permission( user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None = None, proxy_logging_obj: ProxyLogging | None = None, + check_db_only: bool = False, ) -> LiteLLM_ObjectPermissionTable | None: """ - Check if object permission id in proxy ObjectPermissionTable @@ -3942,9 +3965,13 @@ async def get_object_permission( # check if in cache key: Final = object_permission_cache_key(object_permission_id) - deserialized_perm: Final = await user_api_key_cache.async_get_cache( - key=key, - model_type=LiteLLM_ObjectPermissionTable, + deserialized_perm: Final = ( + None + if check_db_only + else await user_api_key_cache.async_get_cache( + key=key, + model_type=LiteLLM_ObjectPermissionTable, + ) ) if deserialized_perm is not None: return deserialized_perm @@ -3952,10 +3979,12 @@ async def get_object_permission( # else, check db try: response: Final = await _dictable_table( - ObjectPermissionRepository(prisma_client), "object_permission" + ObjectPermissionRepository(prisma_client, use_writer=check_db_only), "object_permission" ).find_unique(where={"object_permission_id": object_permission_id}) if response is None: + if check_db_only: + raise HTTPException(status_code=403, detail="Referenced object permission does not exist") return None _perm_obj: Final = LiteLLM_ObjectPermissionTable.model_validate(response.dict()) @@ -3968,6 +3997,8 @@ async def get_object_permission( return _perm_obj except Exception: + if check_db_only: + raise return None @@ -4177,6 +4208,7 @@ async def _get_resources_from_access_groups( prisma_client: DatabaseClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, proxy_logging_obj: ProxyLogging | None = None, + check_db_only: bool = False, ) -> list[str]: """ Fetch access groups by their IDs (from cache or DB) and collect @@ -4219,6 +4251,7 @@ async def _get_resources_from_access_groups( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) resources.extend(getattr(ag, resource_field, [])) except Exception: @@ -4254,6 +4287,7 @@ async def _get_mcp_server_ids_from_access_groups( prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, proxy_logging_obj: ProxyLogging | None = None, + check_db_only: bool = False, ) -> list[str]: """ Collect MCP server IDs from unified access groups. @@ -4265,6 +4299,7 @@ async def _get_mcp_server_ids_from_access_groups( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) @@ -4273,6 +4308,7 @@ async def _get_agent_ids_from_access_groups( prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, proxy_logging_obj: ProxyLogging | None = None, + check_db_only: bool = False, ) -> list[str]: """ Collect agent IDs from unified access groups. @@ -4284,6 +4320,7 @@ async def _get_agent_ids_from_access_groups( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) @@ -4483,26 +4520,36 @@ async def _check_agent_access_group_model_access( """Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows.""" if not model or valid_token is None or not valid_token.agent_id: return True - ceiling: Final = await resolve_ceiling(valid_token.agent_id) - if ceiling is None: - return True - if not ceiling.models: - raise ModelAccessDeniedProxyException( - message=model_access_denied_client_message(model=model), - internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models", - type=ProxyErrorTypes.agent_model_access_denied, - param="model", - code=status.HTTP_403_FORBIDDEN, - ) - dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) - return _can_object_call_model( - model=dispatched, - llm_router=llm_router, - models=sorted(ceiling.models), - team_id=valid_token.team_id, - object_type="agent", - key_model_aliases=key_model_aliases_for_auth_check(valid_token), + + from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings + + unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if valid_token.managed_agent_policy is None else None + ceilings: Final = ( + await resolve_managed_agent_ceilings(valid_token.managed_agent_policy) + if valid_token.managed_agent_policy is not None + else (unmanaged,) + if unmanaged is not None + else () ) + dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) + for ceiling in ceilings: + if not ceiling.models: + raise ModelAccessDeniedProxyException( + message=model_access_denied_client_message(model=model), + internal_message=f"agent {valid_token.agent_id} access groups grant no models", + type=ProxyErrorTypes.agent_model_access_denied, + param="model", + code=status.HTTP_403_FORBIDDEN, + ) + _can_object_call_model( + model=dispatched, + llm_router=llm_router, + models=sorted(ceiling.models), + team_id=valid_token.team_id, + object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), + ) + return True LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 4f41b283a33..71b87b74915 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -52,6 +52,10 @@ from litellm.proxy._types import ( TeamMemberAddRequest, UserAPIKeyAuth, ) +from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed +from litellm.proxy.agent_endpoints.identity import has_legacy_identity +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent +from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure from litellm.proxy.auth.auth_checks import can_team_access_model from litellm.proxy.auth.model_access_denied import ( ModelAccessDeniedHTTPException, @@ -67,6 +71,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.user_repository import UserRepository from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import AgentIdentityFailure from litellm.types.proxy.auth.auth_checks import UserNotFoundError from .auth_checks import ( @@ -157,6 +162,8 @@ class HeaderTeam: class AgentLookup(Protocol): """The registered-agent lookups a JWT agent claim is matched against.""" + def get_agent_list(self) -> Sequence[AgentResponse]: ... + def get_agent_by_id(self, agent_id: str) -> AgentResponse | None: """The agent registered under ``agent_id``, if any.""" @@ -167,6 +174,9 @@ class AgentLookup(Protocol): class _NoRegisteredAgents: """The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches.""" + def get_agent_list(self) -> tuple[AgentResponse, ...]: + return () + def get_agent_by_id(self, agent_id: str) -> None: return None @@ -1096,6 +1106,15 @@ class JWTHandler: "options": options or None, } + def managed_issuer_is_trusted(self, issuer: object) -> bool: + if not isinstance(issuer, str): + return False + configured: Final = self.litellm_jwtauth.issuers or () + for item in configured: + if item.issuer == issuer: + return bool(item.audience) and not item.disable_audience_validation + return issuer == os.getenv("JWT_ISSUER") and bool(os.getenv("JWT_AUDIENCE")) + def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None: litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None) if litellm_jwtauth is None: @@ -1488,7 +1507,12 @@ class JWTAuthManager: agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name( agent_name=agent_claim ) - if agent is None: + if ( + agent is None + or agent.identity_managed + or agent.identity is not None + or has_legacy_identity(agent.litellm_params) + ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}", @@ -2478,12 +2502,39 @@ class JWTAuthManager: """Resolve and authorize JWT context; only normal admission supplies provisioning.""" handler: Final = jwt_handler jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler) + managed: Final = await resolve_managed_agent(jwt_valid_token, prisma_client) + if managed is not None: + if not handler.managed_issuer_is_trusted(jwt_valid_token.get("iss")): + raise HTTPException(403, "Managed agents require trusted JWT issuer and audience validation") + if not managed_agent_route_allowed(route, request_method): + raise HTTPException(403, "Agent identities can only access inference and agent discovery routes") + evidence: Final = await AgentIdentityStore.from_client(prisma_client).record_authentication(managed) + if isinstance(evidence, AgentIdentityFailure): + raise_identity_failure(evidence) + if managed.mode == "autonomous": + return JWTAuthBuilderResult( + is_proxy_admin=False, + team_id=None, + team_object=None, + user_id=None, + user_email=None, + user_object=None, + org_id=None, + org_object=None, + end_user_id=None, + end_user_object=None, + token=api_key, + team_membership=None, + jwt_claims=jwt_valid_token, + agent_id=managed.agent_id, + managed_agent_context=managed, + ) team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False model: Final = request_data.get("model") requested_model: Final = model if isinstance(model, str) else None # Check RBAC - rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) + rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if managed is None else None await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role) # Check Scope Based Access @@ -2499,7 +2550,11 @@ class JWTAuthManager: object_id = handler.get_object_id(token=jwt_valid_token, default_value=None) # Get basic user info - user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token) + user_id, user_email, valid_user_email = ( + (managed.user_id, None, None) + if managed is not None + else await JWTAuthManager.get_user_info(handler, jwt_valid_token) + ) # Get IDs org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None) @@ -2514,23 +2569,31 @@ class JWTAuthManager: elif rbac_role == LitellmUserRoles.INTERNAL_USER: user_id = object_id - agent_id: Final = JWTAuthManager.resolve_agent_id( - jwt_handler=handler, - jwt_valid_token=jwt_valid_token, - agent_registry=handler.agent_lookup, + agent_id: Final = ( + managed.agent_id + if managed is not None + else JWTAuthManager.resolve_agent_id( + jwt_handler=handler, + jwt_valid_token=jwt_valid_token, + agent_registry=handler.agent_lookup, + ) ) # Check admin access - admin_result: Final = await JWTAuthManager.check_admin_access( - handler, - scopes, - route, - user_id, - org_id, - api_key, - jwt_valid_token, - user_email=user_email, - agent_id=agent_id, + admin_result: Final = ( + None + if managed is not None + else await JWTAuthManager.check_admin_access( + handler, + scopes, + route, + user_id, + org_id, + api_key, + jwt_valid_token, + user_email=user_email, + agent_id=agent_id, + ) ) if admin_result: await JWTAuthManager._attach_team_from_header_for_admin( @@ -2705,13 +2768,13 @@ class JWTAuthManager: proxy_logging_obj=proxy_logging_obj, route=route, org_alias=org_alias, - user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False, + user_id_upsert=provisioning.user_id_upsert if provisioning is not None and managed is None else False, ) # Derive org_id from org_object if resolved by alias resolved_org_id: Final = org_object.organization_id if org_object else org_id - if provisioning is not None: + if provisioning is not None and managed is None: await JWTAuthManager.sync_user_role_and_teams( jwt_handler=handler, jwt_valid_token=jwt_valid_token, @@ -2784,7 +2847,7 @@ class JWTAuthManager: ) ## MAP USER TO TEAMS - if provisioning is not None: + if provisioning is not None and managed is None: await JWTAuthManager.map_user_to_teams( user_object=user_object, team_object=team_object, @@ -2799,7 +2862,9 @@ class JWTAuthManager: ) # check if user is proxy admin - is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN) + is_proxy_admin: Final = managed is None and bool( + user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN + ) return JWTAuthBuilderResult( is_proxy_admin=is_proxy_admin, @@ -2816,6 +2881,7 @@ class JWTAuthManager: team_membership=team_membership_object, jwt_claims=jwt_valid_token, agent_id=agent_id, + managed_agent_context=managed, ) @staticmethod @@ -2826,11 +2892,13 @@ class JWTAuthManager: """Keep JWT identity and permission attribution identical across consumers.""" user: Final = result["user_object"] admin: Final = result["is_proxy_admin"] - return UserAPIKeyAuth( + auth: Final = UserAPIKeyAuth( api_key=None, user_role=( LitellmUserRoles.PROXY_ADMIN if admin + else LitellmUserRoles.INTERNAL_USER + if result.get("managed_agent_context") is not None else LitellmUserRoles(user.user_role) if user is not None and user.user_role is not None else LitellmUserRoles.INTERNAL_USER @@ -2852,3 +2920,5 @@ class JWTAuthManager: user_id=result["user_id"], ), ) + auth.managed_agent_context = result.get("managed_agent_context") + return auth diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..2829bc95b42 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -652,6 +652,8 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str # never reaches the fallback. synthetic_scope: Final[dict[str, Any]] = { "type": "http", + "method": "GET", + "query_string": ws_scope.get("query_string", b""), "headers": scope_headers, "path": ws_scope.get("path", ""), "state": ws_scope.setdefault("state", {}), # mutable-ok: Starlette's socket state, shared with the request @@ -1653,6 +1655,13 @@ async def _user_api_key_auth_builder( else: jwt_claims = await jwt_handler.auth_jwt(token=api_key) + from litellm.proxy.agent_endpoints.identity_store import resolve_managed_agent + + if jwt_claims and await resolve_managed_agent(jwt_claims, prisma_client) is not None: + raise HTTPException( + 403, "Managed agents require direct JWT authentication without virtual-key mapping" + ) + resolve_result: Final = await _resolve_jwt_to_virtual_key( jwt_claims=jwt_claims, jwt_handler=jwt_handler, @@ -3123,7 +3132,10 @@ async def _reserve_budget_after_common_checks( end_user_id=end_user_id, end_user_object=end_user_object, apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True, - fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True, + fail_closed_budget_enforcement=( + general_settings.get("fail_closed_budget_enforcement") is True + or user_api_key_auth_obj.billing_agent_policy is not None + ), raw_body=await read_raw_json_body(request=request), ) if request is not None: @@ -3197,10 +3209,48 @@ async def _authorize_authenticated_request( # admin-only-route / model-access / budget checks) surface as # ProxyException consistently with pre-refactor behavior. try: + from litellm.proxy.agent_endpoints.auth.managed_authorization import ( + admit_managed_actor, + invocation_target, + managed_agent_route_allowed, + managed_inference_request, + prepare_agent_invocation, + ) + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + from litellm.proxy.proxy_server import general_settings, prisma_client, user_model + + store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None + if user_api_key_auth_obj.agent_id is not None: + await admit_managed_actor(user_api_key_auth_obj, store) + if user_api_key_auth_obj.managed_agent_policy is not None and not managed_agent_route_allowed( + route, request.method + ): + raise HTTPException(403, "Agent identities can only access inference and agent discovery routes") + authorized_data: Final = ( + managed_inference_request( + route, + request_data, + general_settings, + user_model, + request.path_params.get("model") or request.path_params.get("model_name"), + request.query_params.get("model"), + ) + if user_api_key_auth_obj.managed_agent_policy is not None + else request_data + ) + target_name: Final = invocation_target(route, authorized_data) + if target_name is not None: + await prepare_agent_invocation( + user_api_key_auth_obj, + target_name, + store, + billable=request_data.get("method") + in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), + ) await _run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, request=request, - request_data=request_data, + request_data=authorized_data, route=route, ) except Exception as e: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 64b0c6c1967..b3be5fb78c9 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -93,6 +93,7 @@ from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_bod from litellm.proxy.common_utils.http_parsing_utils import ( get_client_requested_model, get_tags_from_request_body, + resolve_inference_model, ) from litellm.proxy.common_utils.openai_error_payload import ( LITELLM_CALL_ID_HEADER, @@ -2068,11 +2069,12 @@ class ProxyBaseLLMRequestProcessing: if isinstance(model, str): reject_url_valued_destination("model", model) - self.data["model"] = ( - general_settings.get("completion_model", None) # server default - or user_model # model name passed via cli args - or model # for azure deployments - or self.data.get("model", None) # default passed in http request + self.data["model"] = resolve_inference_model( + self.data.get("model"), + general_settings, + user_model, + model, + kind="image_edit" if route_type == "aimage_edit" else "completion", ) # override with user settings, these are params passed via cli diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 1c2bd7ea217..be62448e3ed 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -2,7 +2,7 @@ import json import re from collections.abc import Collection, Mapping from types import MappingProxyType, UnionType -from typing import Annotated, Any, Final, Union, get_args, get_origin +from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin import orjson from fastapi import Request, UploadFile, status @@ -25,6 +25,39 @@ _FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form- _ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required}) +def resolve_inference_model( + body_model: object, + settings: Mapping[str, object], + cli_model: str | None, + endpoint_model: object = None, + *, + kind: Literal[ + "completion", "image_generation", "image_edit", "moderation", "speech", "body", "path" + ] = "completion", +) -> object: + match kind: + case "image_generation": + return cli_model or endpoint_model or settings.get("image_generation_model") or body_model + case "image_edit": + return ( + settings.get("completion_model") + or cli_model + or endpoint_model + or settings.get("image_generation_model") + or body_model + ) + case "moderation": + return cli_model or settings.get("moderation_model") or body_model + case "speech": + return cli_model or body_model + case "body": + return body_model + case "path": + return endpoint_model + case "completion": + return settings.get("completion_model") or cli_model or endpoint_model or body_model + + def _normalize_media_type(content_type: str) -> str: """Return the bare media type per RFC 7231: strip params, trim, lowercase.""" if not content_type: diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index 63da5d15207..8a82e253c5c 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -174,7 +174,10 @@ async def _resync_agents(agent_id_or_name: str) -> bool: table: Final = agents_table(prisma_client) id_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id_or_name} name_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_name": agent_id_or_name} - include_permission: Final[LiteLLM_AgentsTableInclude] = {"object_permission": True} + include_permission: Final[LiteLLM_AgentsTableInclude] = { + "object_permission": True, + "identity": True, + } async with AGENT_RECONCILE_LOCK: if _agent_from_registry(agent_id_or_name) is not None: return True diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 0178465739b..e735f4a51f0 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -358,6 +358,7 @@ class _ProxyDBLogger(CustomLogger): team_id=team_id, end_user_id=end_user_id, call_type=call_type, + agent_id=metadata.get("billing_agent_id") or metadata.get("agent_id"), ): ## UPDATE DATABASE charged: Final = await _update_database_and_spend_counters( @@ -612,6 +613,7 @@ def _should_track_cost_callback( team_id: str | None, end_user_id: str | None, call_type: str | None = None, + agent_id: str | None = None, ) -> bool: """ Determine if the cost callback should be tracked based on the kwargs @@ -628,7 +630,13 @@ def _should_track_cost_callback( if ProxyUpdateSpend.disable_spend_updates() is True: return False - if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None: + if ( + agent_id is not None + or user_api_key is not None + or user_id is not None + or team_id is not None + or end_user_id is not None + ): return True return call_type in _UNATTRIBUTED_TRACKABLE_CALL_TYPES diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index b9580ba3948..16dc38575da 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -21,6 +21,7 @@ from litellm.proxy.common_request_processing import ( from litellm.proxy.common_utils.http_parsing_utils import ( coerce_numeric_form_fields, numeric_form_fields, + resolve_inference_model, ) from litellm.proxy.common_utils.openai_error_payload import ( error_status_code, @@ -118,14 +119,9 @@ async def image_generation( if isinstance(model, str): reject_url_valued_destination("model", model) - data["model"] = ( - model - or general_settings.get("image_generation_model", None) # server default - or user_model # model name passed via cli args - or data.get("model", None) # default passed in http request + data["model"] = resolve_inference_model( + data.get("model"), general_settings, user_model, model, kind="image_generation" ) - if user_model: - data["model"] = user_model ### MODEL ALIAS MAPPING ### # check if model name in model alias map @@ -324,12 +320,6 @@ async def image_edit_api( if "prompt" not in data: data["prompt"] = None - data["model"] = ( - model - or general_settings.get("image_generation_model", None) # server default - or user_model # model name passed via cli args - or data.get("model", None) # default passed in http request - ) ######################################################### # Process request ######################################################### @@ -346,7 +336,7 @@ async def image_edit_api( general_settings=general_settings, proxy_config=proxy_config, select_data_generator=select_data_generator, - model=None, + model=model, user_model=user_model, user_temperature=user_temperature, user_request_timeout=user_request_timeout, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 56f647d5acc..a698881189a 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1659,7 +1659,19 @@ class LiteLLMProxyRequestSetup: _key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None) _existing_agent_id: Final = data[_metadata_variable_name].get("agent_id") _resolved_agent_id: Final = _key_agent_id or _existing_agent_id - data[_metadata_variable_name]["agent_id"] = _resolved_agent_id + data[_metadata_variable_name]["agent_id"] = user_api_key_dict.invoked_agent_id or _resolved_agent_id + managed_context: Final = user_api_key_dict.managed_agent_context + data[_metadata_variable_name].update( + MappingProxyType( + { + "actor_agent_id": user_api_key_dict.agent_id, + "target_agent_id": user_api_key_dict.invoked_agent_id, + "billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id, + "agent_execution_mode": managed_context.mode if managed_context else None, + "verified_human_user_id": managed_context.user_id if managed_context else None, + } + ) + ) data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr( user_api_key_dict, "end_user_max_budget", None diff --git a/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py b/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py new file mode 100644 index 00000000000..e548b9b7fa2 --- /dev/null +++ b/litellm/proxy/management_endpoints/sso/agent_subject_enrollment.py @@ -0,0 +1,49 @@ +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final +from uuid import UUID + +from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject + + +def microsoft_interactive_subject( + tenant: str | None, + response: Mapping[str, object], + endpoints: Mapping[str, str | None], +) -> MicrosoftInteractiveSubject | None: + if tenant is None: + return None + try: + tenant_id: Final = str(UUID(tenant)) + object_id: Final = response.get("id") + if not isinstance(object_id, str): + return None + oid: Final = str(UUID(object_id)) + except ValueError: + return None + expected: Final = MappingProxyType( + { + "MICROSOFT_AUTHORIZATION_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/authorize", + "MICROSOFT_TOKEN_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token", + "MICROSOFT_USERINFO_ENDPOINT": "https://graph.microsoft.com/v1.0/me", + } + ) + if any(value and value != expected.get(name) for name, value in endpoints.items()): + return None + return MicrosoftInteractiveSubject( + issuer=f"https://login.microsoftonline.com/{tenant_id}/v2.0", + tenant_id=tenant_id, + oid=oid, + ) + + +async def enroll_microsoft_subject(subject: object, user_id: object, client: object) -> None: + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure + from litellm.types.proxy.agent_identity import AgentIdentityFailure + + if not isinstance(subject, MicrosoftInteractiveSubject) or not isinstance(user_id, str) or not user_id: + return + result: Final = await AgentIdentityStore.from_client(client).enroll_interactive_human(subject, user_id) + if isinstance(result, AgentIdentityFailure): + raise_identity_failure(result) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 618b200a14c..c995b5aedd4 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -3631,6 +3631,12 @@ class SSOAuthenticationHandler: }, ) + from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject + + await enroll_microsoft_subject( + request.scope.get("litellm_microsoft_interactive_subject"), user_id, prisma_client + ) + if isinstance(user_id, str) and user_id: await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion) await warn_if_id_jag_assertion_uncaptured(sso_assertion) @@ -4300,6 +4306,22 @@ class MicrosoftSSOHandler: original_msft_result["app_roles"] = app_roles return original_msft_result or {} + from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import microsoft_interactive_subject + + request.scope["litellm_microsoft_interactive_subject"] = microsoft_interactive_subject( + microsoft_tenant, + original_msft_result, + MappingProxyType( + { + name: os.getenv(name) + for name in ( + "MICROSOFT_AUTHORIZATION_ENDPOINT", + "MICROSOFT_TOKEN_ENDPOINT", + "MICROSOFT_USERINFO_ENDPOINT", + ) + } + ), + ) result: Final = MicrosoftSSOHandler.openid_from_response( response=original_msft_result, team_ids=user_team_ids, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f842f2e1e4a..b3b20a16234 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -421,6 +421,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, check_file_size_under_limit, get_form_data, + resolve_inference_model, ) from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations @@ -12186,13 +12187,7 @@ async def moderations( proxy_config=proxy_config, ) - data["model"] = ( - general_settings.get("moderation_model", None) # server default - or user_model # model name passed via cli args - or data.get("model") # default passed in http request - ) - if user_model: - data["model"] = user_model + data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation") ### CALL HOOKS ### - modify incoming data / reject request before calling the model data = await proxy_logging_obj.pre_call_hook( @@ -12446,13 +12441,7 @@ async def audio_transcriptions( if data.get("user", None) is None and user_api_key_dict.user_id is not None: data["user"] = user_api_key_dict.user_id - data["model"] = ( - general_settings.get("moderation_model", None) # server default - or user_model # model name passed via cli args - or data.get("model", None) # default passed in http request - ) - if user_model: - data["model"] = user_model + data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation") router_model_names: Final = llm_router.model_names if llm_router is not None else [] diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 69c63d9ecd6..8535d1004f7 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -78,6 +78,11 @@ model LiteLLM_AgentsTable { object_permission_id String? object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) spend Float @default(0.0) + identity_managed Boolean @default(false) + enabled Boolean @default(true) + execution_mode String @default("autonomous") + identity LiteLLM_AgentIdentity? + retired_identities LiteLLM_RetiredAgentIdentity[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -88,6 +93,56 @@ model LiteLLM_AgentsTable { updated_by String } +model LiteLLM_AgentIdentity { + agent_id String @id + active Boolean @default(true) + agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade) + provider String + issuer String + tenant_id String + client_id String + service_principal_id String? + required_roles String[] @default([]) + required_scopes String[] @default(["user_impersonation"]) + revision String @default(uuid()) + last_authenticated_at DateTime? + @@unique([provider, tenant_id, client_id]) + @@unique([issuer, service_principal_id]) +} + +model LiteLLM_RetiredAgentIdentity { + binding_id String @id @default(uuid()) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + provider String + issuer String + tenant_id String + client_id String + @@unique([provider, tenant_id, client_id]) +} + +model LiteLLM_RetiredAgent { + original_agent_id String @id + retired_at DateTime @default(now()) +} + +model LiteLLM_VerifiedSubject { + subject_id String @id @default(uuid()) + issuer String + tenant_id String + oid String + kind String @default("human") + user_id String? + user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + verified_via String @default("sso_interactive") + verified_at DateTime @default(now()) + @@unique([issuer, tenant_id, oid]) + @@index([user_id]) +} + + + + model LiteLLM_OrganizationTable { organization_id String @id @default(uuid()) organization_alias String @@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? team_id String? @@ -674,6 +730,7 @@ model LiteLLM_SpendLogs { session_id String? status String? mcp_namespaced_tool_name String? + billing_agent_id String? agent_id String? proxy_server_request Json? @default("{}") litellm_call_id String? diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index c4ed8713f95..c3a1a7000c5 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -77,6 +77,11 @@ _SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)" _SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key" _MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')" _AGENT_CALL_TYPE_SQL: Final = "'asend_message'" +_SESSION_REPRESENTATIVE_ORDER_SQL: Final = ( + f"(call_type = {_AGENT_CALL_TYPE_SQL}) DESC, " + f'CASE WHEN call_type = {_AGENT_CALL_TYPE_SQL} THEN "endTime" END DESC NULLS LAST, ' + f'call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC, request_id' +) _BATCH_CALL_TYPES_SQL: Final = "('acreate_batch', 'create_batch', 'aretrieve_batch', 'retrieve_batch')" _SPAN_TYPE_SQL_CONDITIONS: Final[Mapping[str, str]] = MappingProxyType( { @@ -2879,7 +2884,7 @@ async def ui_view_spend_logs( p += 1 # Status filter - if status_filter is not None: + if status_filter is not None and not (group_by_session is True and not is_search_lookup): if status_filter == "success": sql_conditions.append("(status = 'success' OR status IS NULL)") else: @@ -2925,6 +2930,23 @@ async def ui_view_spend_logs( sql_params.append(f"%{error_message}%") p += 1 + if status_filter is not None and group_by_session is True and not is_search_lookup: + session_filter_conditions: Final = " AND ".join(sql_conditions) or "TRUE" + sql_conditions.append( + f"""({_SESSION_GROUP_KEY_SQL}) IN ( + SELECT session_key, api_key FROM ( + SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL}) + {_SESSION_KEY_EXPR} AS session_key, api_key, status + FROM "LiteLLM_SpendLogs" + WHERE {session_filter_conditions} + ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL} + ) AS session_outcomes + WHERE COALESCE(status, 'success') = ${p} + )""" + ) + sql_params.append(status_filter) + p += 1 + if ( group_by_session is True and not is_v2 @@ -2991,7 +3013,7 @@ async def ui_view_spend_logs( {_SPEND_LOG_LIST_COLUMNS} FROM "LiteLLM_SpendLogs" WHERE {joined_conditions} - ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC + ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL} ) AS session_representatives ORDER BY {exact_request_id_first}{_order_expr} {_sql_dir}{_nulls_clause}, request_id LIMIT ${p} OFFSET ${p + 1} @@ -3063,7 +3085,7 @@ async def _fetch_session_representatives( next_param_index: int, session_keys: Sequence[tuple[str, str]], ) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row - """Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order.""" + """Fetch the final agent outcome, or newest non-MCP row, of each ``(session_key, api_key)`` session, in ``session_keys`` order.""" rep_query: Final = f""" SELECT * FROM ( SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL}) @@ -3073,7 +3095,7 @@ async def _fetch_session_representatives( AND ({_SESSION_GROUP_KEY_SQL}) IN ( SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[]) ) - ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC + ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL} ) AS session_representatives """ rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place @@ -3140,7 +3162,7 @@ async def _ui_session_grouped_spend_logs( page_size``, trimmed to the end of the ``SPEND_LOGS_PAGINATION_COUNT_CAP`` window the capped ``total`` promises, so a page never runs past that total and one starting at or past it returns no rows without a query. Each session is represented - by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response`` + by its final agent outcome (or newest non-MCP row), enriched by ``_build_ui_spend_logs_response`` exactly like the flat listing, and the response carries ``next_session_cursor`` / ``has_more`` while ``total`` counts sessions (capped like the flat total). A page that runs out of sessions while still diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 1c51fb21d6e..e5001eeb32f 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -795,6 +795,7 @@ def get_logging_payload( model_id=_model_id, mcp_namespaced_tool_name=mcp_namespaced_tool_name, agent_id=agent_id, + billing_agent_id=clean_metadata.get("billing_agent_id"), requester_ip_address=clean_metadata.get("requester_ip_address", None), custom_llm_provider=custom_llm_provider or "", messages=_get_messages_for_spend_logs_payload( diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 7736939c696..b732d2ff94c 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -15,9 +15,14 @@ if TYPE_CHECKING: class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): """Repository for object permission database operations.""" + def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None: + super().__init__(prisma_client) + self._use_writer = use_writer + @property def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: - return self.prisma_client.db.litellm_objectpermissiontable + database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + return database.litellm_objectpermissiontable @property def model_class(self) -> type[LiteLLM_ObjectPermissionTable]: diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 1ad7a735d96..292694747f3 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -21,8 +21,9 @@ class PrismaTableRepository(Generic[RowT_co]): table_name: str - def __init__(self, prisma_client: object): + def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None: self._prisma_client = prisma_client + self._use_writer = use_writer @property def prisma_client(self) -> Any: @@ -32,7 +33,9 @@ class PrismaTableRepository(Generic[RowT_co]): @property def table(self) -> TableActions[RowT_co]: - actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name) + actions: Final[TableActions[RowT_co]] = getattr( + self.prisma_client.writer_db if self._use_writer else self.prisma_client.db, self.table_name + ) return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name) @@ -44,6 +47,18 @@ class AgentsRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentsTable" table_name = "litellm_agentstable" +class AgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentIdentity"]): + table_name = "litellm_agentidentity" + + +class RetiredAgentIdentityRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgentIdentity"]): + table_name = "litellm_retiredagentidentity" + + +class VerifiedSubjectRepository(PrismaTableRepository["prisma_models.LiteLLM_VerifiedSubject"]): + table_name = "litellm_verifiedsubject" + + class ObjectPermissionRepository(PrismaTableRepository["prisma_models.LiteLLM_ObjectPermissionTable"]): table_name = "litellm_objectpermissiontable" @@ -246,3 +261,7 @@ class AuditLogRepository(PrismaTableRepository["prisma_models.LiteLLM_AuditLog"] class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterSession"]): table_name = "litellm_adaptiveroutersession" + + +class RetiredAgentRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgent"]): + table_name = "litellm_retiredagent" diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index cbe263699c9..57f6fd33c11 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -70,6 +70,9 @@ class _PrismaClientView(Protocol): @property def db(self) -> _PrismaTeamDb: ... + @property + def writer_db(self) -> _PrismaTeamDb: ... + _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( @@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = ( class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" + def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None: + super().__init__(prisma_client) + self._use_writer = use_writer + @property def _db(self) -> _PrismaTeamDb: client: Final[_PrismaClientView] = self.prisma_client - return client.db + return client.writer_db if self._use_writer else client.db @property def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 87eb45f262d..4a2aea46197 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...]) class UserRepository(BaseRepository[LiteLLM_UserTable]): """Repository for user database operations.""" + def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None: + super().__init__(prisma_client) + self._use_writer = use_writer + @property def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: - return self.prisma_client.db.litellm_usertable + database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + return database.litellm_usertable @property def model_class(self) -> type[LiteLLM_UserTable]: diff --git a/litellm/types/agents.py b/litellm/types/agents.py index f7aef09fa29..94adb9f7c4a 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -7,6 +7,11 @@ from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field from typing_extensions import ReadOnly, Required, TypedDict from litellm.types.llms.base import LiteLLMPydanticObjectBase +from litellm.types.proxy.agent_identity import ( + AgentExecutionMode, + AgentIdentityBinding, + EntraIdentityConfig, +) if TYPE_CHECKING: from a2a.types import SendMessageResponse @@ -248,8 +253,11 @@ class AgentKillSwitchResult(BaseModel): class AgentConfig(TypedDict, total=False): + identity: ReadOnly[EntraIdentityConfig | None] + enabled: ReadOnly[bool] + execution_mode: ReadOnly[AgentExecutionMode] agent_name: Required[str] - agent_card_params: Required[AgentCard] + agent_card_params: ReadOnly[AgentCard] litellm_params: dict[str, object] # allow for any future litellm params object_permission: AgentObjectPermission tpm_limit: int | None @@ -263,6 +271,9 @@ class AgentConfig(TypedDict, total=False): class PatchAgentRequest(TypedDict, total=False): + identity: ReadOnly[EntraIdentityConfig | None] + enabled: ReadOnly[bool] + execution_mode: ReadOnly[AgentExecutionMode] agent_name: str agent_card_params: AgentCard litellm_params: dict[str, object] @@ -301,6 +312,11 @@ class AgentKeySummary(BaseModel): class AgentResponse(BaseModel): + identity: AgentIdentityBinding | None = None + identity_managed: bool = False + enabled: bool = True + execution_mode: AgentExecutionMode = "autonomous" + jwt_auth_configured: bool = False agent_id: str agent_name: str litellm_params: dict[str, object] | None = None diff --git a/litellm/types/proxy/agent_identity.py b/litellm/types/proxy/agent_identity.py new file mode 100644 index 00000000000..a7fe0be37e1 --- /dev/null +++ b/litellm/types/proxy/agent_identity.py @@ -0,0 +1,96 @@ +from datetime import datetime +from typing import Literal, TypeAlias +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, field_validator + +AgentExecutionMode: TypeAlias = Literal["autonomous", "delegated", "both"] + + +class EntraIdentityConfig(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + provider: Literal["microsoft_entra"] + tenant_id: str + client_id: str + service_principal_id: str | None = None + required_roles: tuple[str, ...] = () + required_scopes: tuple[str, ...] = Field( + default=("user_impersonation",), + description="Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway.", + ) + + @field_validator("tenant_id", "client_id", "service_principal_id") + @classmethod + def normalize_identifier(cls, value: str | None) -> str | None: + return str(UUID(value)) if value is not None else None + + @property + def issuer(self) -> str: + return f"https://login.microsoftonline.com/{self.tenant_id}/v2.0" + + +class AgentIdentityBinding(BaseModel): + model_config = ConfigDict(frozen=True) + + agent_id: str + active: bool = True + provider: Literal["microsoft_entra"] + tenant_id: str + client_id: str + service_principal_id: str | None = None + issuer: str + required_roles: tuple[str, ...] = () + required_scopes: tuple[str, ...] = ("user_impersonation",) + revision: str + last_authenticated_at: datetime | None = None + + +class AgentSubject(BaseModel): + model_config = ConfigDict(frozen=True) + + kind: Literal["application", "delegated_subject"] + oid: str + mode: Literal["autonomous", "delegated"] + + +class AgentIdentityFailure(BaseModel): + model_config = ConfigDict(frozen=True) + + code: Literal["identity_denied", "policy_unavailable"] = "identity_denied" + message: str + + +class ManagedAgentContext(BaseModel): + model_config = ConfigDict(frozen=True) + + agent_id: str + binding_revision: str | None = None + mode: Literal["autonomous", "delegated"] + user_id: str | None = None + subject_oid: str | None = None + + +class VerifiedHumanSubject(BaseModel): + model_config = ConfigDict(frozen=True) + + issuer: str + tenant_id: str + oid: str + user_id: str + + +class MicrosoftInteractiveSubject(BaseModel): + model_config = ConfigDict(frozen=True) + + issuer: str + tenant_id: str + oid: str + + +class ManagedAgentIdentityStatus(BaseModel): + identity: AgentIdentityBinding | None = None + identity_managed: bool = False + enabled: bool = True + execution_mode: AgentExecutionMode = "autonomous" + last_authenticated_at: datetime | None = None diff --git a/schema.prisma b/schema.prisma index 69c63d9ecd6..8535d1004f7 100644 --- a/schema.prisma +++ b/schema.prisma @@ -78,6 +78,11 @@ model LiteLLM_AgentsTable { object_permission_id String? object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) spend Float @default(0.0) + identity_managed Boolean @default(false) + enabled Boolean @default(true) + execution_mode String @default("autonomous") + identity LiteLLM_AgentIdentity? + retired_identities LiteLLM_RetiredAgentIdentity[] tpm_limit Int? rpm_limit Int? session_tpm_limit Int? @@ -88,6 +93,56 @@ model LiteLLM_AgentsTable { updated_by String } +model LiteLLM_AgentIdentity { + agent_id String @id + active Boolean @default(true) + agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade) + provider String + issuer String + tenant_id String + client_id String + service_principal_id String? + required_roles String[] @default([]) + required_scopes String[] @default(["user_impersonation"]) + revision String @default(uuid()) + last_authenticated_at DateTime? + @@unique([provider, tenant_id, client_id]) + @@unique([issuer, service_principal_id]) +} + +model LiteLLM_RetiredAgentIdentity { + binding_id String @id @default(uuid()) + agent_id String? + agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull) + provider String + issuer String + tenant_id String + client_id String + @@unique([provider, tenant_id, client_id]) +} + +model LiteLLM_RetiredAgent { + original_agent_id String @id + retired_at DateTime @default(now()) +} + +model LiteLLM_VerifiedSubject { + subject_id String @id @default(uuid()) + issuer String + tenant_id String + oid String + kind String @default("human") + user_id String? + user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade) + verified_via String @default("sso_interactive") + verified_at DateTime @default(now()) + @@unique([issuer, tenant_id, oid]) + @@index([user_id]) +} + + + + model LiteLLM_OrganizationTable { organization_id String @id @default(uuid()) organization_alias String @@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable { // Track spend, rate limit, budget Users model LiteLLM_UserTable { + verified_subjects LiteLLM_VerifiedSubject[] user_id String @id user_alias String? team_id String? @@ -674,6 +730,7 @@ model LiteLLM_SpendLogs { session_id String? status String? mcp_namespaced_tool_name String? + billing_agent_id String? agent_id String? proxy_server_request Json? @default("{}") litellm_call_id String? diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py new file mode 100644 index 00000000000..b6ddd9cb7ac --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py @@ -0,0 +1,301 @@ +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy import proxy_server +from litellm.proxy._experimental.mcp_server import mcp_server_manager +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth +from litellm.proxy.auth import auth_checks +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import ManagedAgentContext + + +def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth: + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="agent-permissions", + mcp_servers=["slack", "linear"], + mcp_tool_permissions={"slack": list(tools)} if tools is not None else None, + ) + agent: Final = AgentResponse( + agent_id="publisher", + agent_name="Publisher", + agent_card_params={}, + object_permission=permission.model_dump(), + identity_managed=True, + ) + auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id) + auth.managed_agent_policy = agent + auth.managed_agent_context = ManagedAgentContext( + agent_id=agent.agent_id, + mode="delegated" if delegated else "autonomous", + user_id="human" if delegated else None, + ) + return auth + + +@pytest.fixture(autouse=True) +def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager()) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write"))) +async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None: + auth: Final = actor(tools) + assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"} + actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) + assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None) + assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "agent_tools,user_tools,expected", + ( + (None, ("read",), ("read",)), + (("read",), None, ("read",)), + (("read", "write"), ("read",), ("read",)), + (("read",), ("write",), ()), + ((), None, ()), + ), +) +async def test_delegated_server_and_tool_intersections( + monkeypatch: pytest.MonkeyPatch, + agent_tools: tuple[str, ...] | None, + user_tools: tuple[str, ...] | None, + expected: tuple[str, ...], +) -> None: + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="user-permissions", + mcp_servers=["slack", "user-only"], + mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None, + ) + user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission) + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user)) + auth: Final = actor(agent_tools, delegated=True) + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"] + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected) + assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == [] + assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == [] + + +@pytest.mark.asyncio +async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable"))) + with pytest.raises(HTTPException) as failure: + await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True)) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ()))) +async def test_access_groups_cap_agent_servers_without_granting_new_ones( + monkeypatch: pytest.MonkeyPatch, + servers: tuple[str, ...], + expected: tuple[str, ...], +) -> None: + from litellm.proxy._types import LiteLLM_AccessGroupTable + + group: Final = LiteLLM_AccessGroupTable( + access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers) + ) + monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group)) + auth: Final = actor(None) + assert auth.managed_agent_policy is not None + auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"access_group_ids": ["group"]}) + assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected + if "slack" not in expected: + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"]) +async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]} + ) + user: Final = LiteLLM_UserTable( + user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant" + ) + cache: Final = UserApiKeyCache() + cache.set_cache("human", user) + cache.set_cache(object_permission_cache_key("user-grant"), permission) + client: Final = MagicMock() + client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user) + client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) + monkeypatch.setattr(proxy_server, "prisma_client", client) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + auth: Final = actor(("read", "write"), delegated=True) + assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"} + if change == "disabled": + client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy( + update={"metadata": {"scim_active": False}} + ) + elif change == "outage": + client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable") + elif change == "servers": + client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy( + update={"mcp_servers": [], "mcp_tool_permissions": {}} + ) + else: + client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy( + update={"mcp_tool_permissions": {"slack": ["read"]}} + ) + if change in ("disabled", "outage"): + with pytest.raises(HTTPException): + await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) + else: + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ( + ["read"] if change == "tools" else [] + ) + client.db.litellm_usertable.find_unique.assert_not_called() + client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"]) +@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"]) +@pytest.mark.parametrize("has_grant", [True, False]) +@pytest.mark.parametrize("agent_tools", [("read", "write"), None]) +async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins( + monkeypatch: pytest.MonkeyPatch, + role: str, + open_channel: str, + has_grant: bool, + agent_tools: tuple[str, ...] | None, +) -> None: + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager: Final = mcp_server_manager.global_mcp_server_manager + manager.registry = { + name: MCPServer( + server_id=name, + name=name, + transport="http", + url="https://example.com/mcp", + allow_all_keys=open_channel == "operator", + ) + for name in ("slack", "linear") + } + from litellm.proxy._experimental.mcp_server import db + + monkeypatch.setattr( + db, + "get_active_submitted_mcp_server_ids_for_user", + AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []), + ) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]} + ) + user: Final = LiteLLM_UserTable( + user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[] + ) + team: Final = LiteLLM_TeamTable( + team_id="team", + models=[], + members_with_roles=[{"user_id": "human", "role": "user"}], + object_permission_id="team-grant", + ) + client: Final = MagicMock() + client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user) + client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team) + client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) + monkeypatch.setattr(proxy_server, "prisma_client", client) + auth: Final = actor(agent_tools, delegated=True) + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else []) + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else []) + assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == [] + admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True) + assert admitted.user_role == role + + +@pytest.mark.asyncio +async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy._experimental.mcp_server import db + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager: Final = mcp_server_manager.global_mcp_server_manager + manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)} + monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"])) + auth: Final = UserAPIKeyAuth(user_id="human") + auth.mcp_explicit_grants_only = True + with pytest.MonkeyPatch.context() as patcher: + patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable"))) + assert await manager.get_allowed_mcp_servers(auth) == [] + auth.mcp_explicit_grants_only = False + assert await manager.get_allowed_mcp_servers(auth) == ["slack"] + + +@pytest.mark.asyncio +async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None: + from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers + + assert await managed_agent_servers(UserAPIKeyAuth()) == () + auth: Final = actor(None, delegated=True) + assert auth.managed_agent_context is not None + auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None}) + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == [] + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == [] + + +@pytest.mark.asyncio +async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None: + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"]) + user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission) + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")])) + with pytest.raises(HTTPException) as failure: + await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True)) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user")) +@pytest.mark.parametrize("scoped", (False, True)) +async def test_manager_preserves_managed_server_grants_across_open_channels( + monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool +) -> None: + from litellm.proxy._experimental.mcp_server import db + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager: Final = mcp_server_manager.global_mcp_server_manager + manager.registry = { + "open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True), + "submitted": MCPServer(server_id="submitted", name="submitted", transport="http"), + "passthrough": MCPServer( + server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough" + ), + } + monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"])) + auth: Final = actor(None) + auth.user_role = role + assert not auth.mcp_explicit_grants_only + access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None + assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == ({"slack"} if scoped else {"slack", "linear"}) + + +@pytest.mark.asyncio +async def test_manager_does_not_replace_managed_policy_failure_with_open_servers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager: Final = mcp_server_manager.global_mcp_server_manager + manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)} + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable"))) + with pytest.raises(HTTPException) as failure: + await manager.get_allowed_mcp_servers(actor(None, delegated=True)) + assert failure.value.status_code == 503 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index fc4d7b45785..5aaa3d3c446 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -4383,7 +4383,7 @@ class TestAgentMCPPermissions: self._team_servers({"callers": ["server_2", "server_3"]}), ), patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) + MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: neither the agent's owner nor the caller has a personal grant MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({}) @@ -4402,7 +4402,7 @@ class TestAgentMCPPermissions: MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({}) ), patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) + MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: same seam, keyed by which user is being asked about MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": ["server_1"]}) @@ -4421,7 +4421,7 @@ class TestAgentMCPPermissions: MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({}) ), patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here - MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) + MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: None is the resolver's own "entitlement unresolvable" signal MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": None}) @@ -4538,7 +4538,7 @@ class TestAgentMCPPermissions: ) with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: - with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = ["server_1"] @@ -4555,7 +4555,7 @@ class TestAgentMCPPermissions: ) with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: - with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = [] # no agent-level restriction @@ -4611,7 +4611,7 @@ class TestAgentMCPPermissions: ) with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key: with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team: - with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent: + with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent: mock_key.return_value = ["server_1", "server_2"] mock_team.return_value = [] mock_agent.return_value = ["server_2", "server_3"] @@ -4637,7 +4637,7 @@ class TestAgentMCPPermissions: ): with patch.object( MCPRequestHandler, - "_get_agent_tool_permissions_for_server", + "get_agent_tool_permissions_for_server", new_callable=AsyncMock, return_value=["tool_a"], ) as mock_agent_tools: @@ -4669,7 +4669,7 @@ class TestAgentMCPPermissions: ): with patch.object( MCPRequestHandler, - "_get_agent_tool_permissions_for_server", + "get_agent_tool_permissions_for_server", new_callable=AsyncMock, return_value=None, ): @@ -4718,7 +4718,7 @@ class TestAgentMCPPermissions: with contextlib.ExitStack() as stack: for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): stack.enter_context(patcher) - result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth) + result = await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth) assert sorted(result) == ["server-a", "server-direct"] mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"]) @@ -4760,7 +4760,7 @@ class TestAgentMCPPermissions: for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): stack.enter_context(patcher) with pytest.raises(UnloadableEntitlementError): - await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth) + await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth) stack.enter_context( patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here MCPRequestHandler, @@ -4789,13 +4789,13 @@ class TestAgentMCPPermissions: with contextlib.ExitStack() as stack: for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): stack.enter_context(patcher) - server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_a_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server( "server-a", user_api_key_auth ) - server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_b_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server( "server-b", user_api_key_auth ) - server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_c_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server( "server-c", user_api_key_auth ) @@ -5833,7 +5833,7 @@ def test_expand_permission_list_does_not_honor_all_proxy_sentinel(): @pytest.mark.asyncio -async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically(): +async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically(monkeypatch): """The TEAM resolver expands the all-proxy sentinel to every registered server and picks up a server registered later, so a team scoped to all-proxy tracks the live registry without any change to its stored permission. Reverting the team-side @@ -5850,6 +5850,9 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam from litellm.types.mcp import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer + monkeypatch.setattr(global_mcp_server_manager, "registry", {}) + monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {}) + for sid in ("srv-x", "srv-y"): global_mcp_server_manager.registry[sid] = MCPServer( server_id=sid, @@ -8305,7 +8308,7 @@ class TestUserSubjectTeamUnion: ) == ["t1"] # An admitted subject never fans out HERE: it resolves one source per team first, and each of # those pins a team_id, so this helper only ever answers the single-team question. The fan-out - # itself is _admitted_subject_sources' job, asserted below. + # itself is admitted_subject_sources' job, asserted below. with self._patch(teams_by_id={}, user_teams=["t2", "t3"]): assert await MCPRequestHandler._team_ids_for_mcp_grant(_make_admitted_subject("u")) == [] # keyless, no user_id -> nothing @@ -8868,7 +8871,7 @@ class TestUserSubjectTeamUnion: teams["t-member"].organization_id = "org-a" auth = _make_admitted_subject("sso-user") with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]): - sources = await MCPRequestHandler._admitted_subject_sources(auth) + sources = await MCPRequestHandler.admitted_subject_sources(auth) assert [(s.team_id, s.org_id) for s in sources] == [(None, None), ("t-member", "org-a")] # The user's own source carries their grants; a team source must NOT, or the team would be @@ -9673,7 +9676,10 @@ class TestGetUserObjectPermission: def _prisma_with_user(self, user_row): prisma_client = MagicMock() - prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + from litellm.proxy._types import LiteLLM_UserTable + + row = LiteLLM_UserTable(user_id="human", object_permission_id=user_row.object_permission_id) if user_row is not None else None + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row) return prisma_client async def test_resolves_through_the_shared_permission_cache(self): @@ -9688,7 +9694,7 @@ class TestGetUserObjectPermission: with ( patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))), patch( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, @@ -9715,7 +9721,7 @@ class TestGetUserObjectPermission: with ( patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))), patch("litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock) as mock_get_perm, ): assert await MCPRequestHandler._get_user_object_permission(auth) is None @@ -9734,7 +9740,7 @@ class TestGetUserObjectPermission: with ( patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))), ): assert await MCPRequestHandler._get_user_object_permission(auth) is None @@ -9748,7 +9754,7 @@ class TestGetUserObjectPermission: with ( patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))), ): assert await MCPRequestHandler._get_user_object_permission(auth) is None @@ -9765,7 +9771,7 @@ class TestGetUserObjectPermission: with ( patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))), patch( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, @@ -10085,3 +10091,47 @@ class TestScopedSessionAdmission: def test_scope_field_cannot_be_forged_through_construction(self): forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server") assert forged.mcp_session_resource_server_id is None + + +@pytest.mark.asyncio +async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch): + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_UserTable + + cached = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="revoked") + current = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="current") + cache = DualCache() + await cache.async_set_cache(key="fresh-human", value=cached) + database = MagicMock() + database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=current) + database.db.litellm_usertable.find_unique = AsyncMock(return_value=cached) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + assert await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) == "current" + database.db.litellm_usertable.find_unique.assert_not_awaited() + database.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("unavailable") + with pytest.raises(HTTPException) as denied: + await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) + assert denied.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["servers", "tools"]) +async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted_grant(monkeypatch, operation): + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.types.agents import AgentResponse + + auth = UserAPIKeyAuth(agent_id="managed") + auth.managed_agent_policy = AgentResponse(agent_id="managed", agent_name="Managed", agent_card_params={}) + permission = LiteLLM_ObjectPermissionTable(object_permission_id="policy", mcp_toolsets=["unavailable"]) + manager = MagicMock() + manager.expand_permission_list.return_value = [] + manager.resolve_toolset_tool_permissions = AsyncMock(side_effect=RuntimeError("policy unavailable")) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + resolution = ( + MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth, permission) + if operation == "servers" + else MCPRequestHandler.get_agent_tool_permissions_for_server("slack", auth, permission) + ) + with pytest.raises(RuntimeError, match="policy unavailable"): + await resolution diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..7791005a070 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7847,7 +7847,7 @@ async def test_load_active_user_by_id_reads_the_row_from_the_database_not_the_ca key="fresh-jwt-user", value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=[]), model_type=LiteLLM_UserTable ) prisma = MagicMock() - prisma.db.litellm_usertable.find_unique = AsyncMock( + prisma.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=["team-a"]) ) proxy_globals.user_api_key_cache = cache diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py index 03165bd0a4a..e94371a6056 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_idp_token_exchange.py @@ -1,4 +1,5 @@ import logging +from typing import Final import pytest from fastapi import HTTPException @@ -235,3 +236,13 @@ async def test_a_database_fault_retrying_cannot_clear_is_not_reported_as_a_trans assert refusal == SubjectTokenRefusal(error="temporarily_unavailable", description=SUBJECT_TOKEN_CHECK_FAULTED) assert "retrying will not help" in refusal.description assert "faulted: " in caplog.text and "query engine binary not found" in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("user_id", [None, "delegating-user"]) +async def test_agent_token_cannot_be_exchanged_for_a_user_identity(user_id: str | None) -> None: + authorizer: Final = _Authorizer({**_authorized(user_id=user_id), "agent_id": "managed-agent"}) + result: Final = await _identity(authorizer) + assert isinstance(result, SubjectTokenRefusal) + assert result.error == "invalid_request" + assert "direct JWT authentication" in result.description diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 70ef4312f4c..e6b339f24be 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -209,15 +209,9 @@ def _reload_mcp_manager_module(): manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"] importlib.reload(utils_module) reloaded = importlib.reload(manager_module) - # After reload, server.py still holds a stale reference to the old - # global_mcp_server_manager. Update it so tests that exercise server.py - # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance. - server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server") - if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): - server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager - operations_module = sys.modules.get("litellm.proxy._experimental.mcp_server.operations") - if operations_module is not None: - operations_module.global_mcp_server_manager = reloaded.global_mcp_server_manager + for name, module in tuple(sys.modules.items()): + if name.startswith("litellm.proxy._experimental.mcp_server.") and hasattr(module, "global_mcp_server_manager"): + module.global_mcp_server_manager = reloaded.global_mcp_server_manager return reloaded @@ -5583,9 +5577,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls @@ -5652,9 +5644,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls @@ -5721,9 +5711,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls @@ -5758,9 +5746,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls @@ -6836,9 +6822,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() proxy_logging_obj = MagicMock() # Mock the async methods that pre_call_tool_check calls @@ -6922,9 +6906,7 @@ class TestMCPServerManager: manager._create_mcp_client = AsyncMock(return_value=mock_client) # Mock user auth with no restrictions - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None + user_api_key_auth: Final = UserAPIKeyAuth() # Mock proxy logging proxy_logging_obj = MagicMock() @@ -11491,12 +11473,8 @@ class TestDiscoveryFailureLogging: assert "unresolved" in caplog.text -def _unrestricted_auth() -> MagicMock: - """A caller with no object_permission, so only server-level checks apply.""" - user_api_key_auth = MagicMock() - user_api_key_auth.object_permission = None - user_api_key_auth.object_permission_id = None - return user_api_key_auth +def _unrestricted_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth() def _permissive_proxy_logging() -> MagicMock: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py index ed3e5f48516..8484dfdde72 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py @@ -139,7 +139,7 @@ async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_r key="stale-cache-user", value=_user(user_id="stale-cache-user", teams=[]), model_type=LiteLLM_UserTable ) prisma = MagicMock() - prisma.db.litellm_usertable.find_unique = AsyncMock( + prisma.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=_user(user_id="stale-cache-user", teams=["team-a"]) ) monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) @@ -164,7 +164,7 @@ async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_the key="deactivated-user", value=_user(user_id="deactivated-user", teams=["team-a"]), model_type=LiteLLM_UserTable ) prisma = MagicMock() - prisma.db.litellm_usertable.find_unique = AsyncMock( + prisma.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False}) ) monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index e82ab28bb4c..759014b54c5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -810,6 +810,12 @@ class TestTestConnection: from litellm.proxy._types import LitellmUserRoles from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy.management_endpoints import mcp_management_endpoints + + manager = MCPServerManager() + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager) captured = self._capture_execute(monkeypatch) saved = MCPServer( server_id="saved-server-id", @@ -1311,8 +1317,9 @@ class TestListToolsRestAPI: session_auth = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="grant-user", user_role="internal_user") admitted_auth = UserAPIKeyAuth(user_id="grant-user", org_id="admitted-org") - async def fake_reload(user_id): + async def fake_reload(user_id, *, requires_fresh_policy=False): assert user_id == "grant-user" + assert requires_fresh_policy is False return admitted_auth monkeypatch.setattr( @@ -1480,9 +1487,12 @@ class TestListToolsRestAPI: from mcp.types import Tool as MCPTool import litellm.experimental_mcp_client.client as mcp_client_module + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.proxy._experimental.mcp_server.server import MCPServer from litellm.types.mcp import MCPTransport + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", MCPServerManager()) + async def fake_contexts(user_api_key_auth): return [user_api_key_auth] @@ -2414,6 +2424,7 @@ class TestCallToolRestAPI: mock_server = MagicMock() mock_server.server_id = "server-1" + mock_server.name = "Example server" def fake_get_mcp_server_by_id(server_id): return mock_server if server_id == "server-1" else None @@ -2431,6 +2442,11 @@ class TestCallToolRestAPI: raising=False, ) + failure_log = AsyncMock() + execute_tool = AsyncMock() + monkeypatch.setattr(rest_endpoints, "_safe_fire_mcp_tool_call_failure_logging", failure_log) + monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", execute_tool) + request_payload = { "server_id": "server-1", "name": "demo-tool", @@ -2452,6 +2468,16 @@ class TestCallToolRestAPI: assert exc_info.value.detail["error"] == "access_denied" assert "server server-1" in exc_info.value.detail["message"] + execute_tool.assert_not_awaited() + failure_log.assert_awaited_once() + logged_data = failure_log.await_args.args[4] + assert logged_data["model"] == "MCP: demo-tool" + assert logged_data["metadata"]["model_group"] == "MCP: demo-tool" + logging_obj = failure_log.await_args.args[0] + assert logging_obj.model_call_details["mcp_tool_call_metadata"] == { + "name": "demo-tool", "mcp_server_name": "Example server", + } + async def test_executes_tool_when_allowed(self, monkeypatch): async def fake_contexts(user_api_key_auth): return [user_api_key_auth] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py index a5f6994b1a7..816ccc5e7e6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py @@ -145,7 +145,7 @@ async def test_build_effective_auth_contexts_appends_admitted_user_context(monke assert contexts[-1].user_id == "user-42" and contexts[-1].team_id is None assert [ctx.team_id for ctx in contexts[:-1]] == ["team-one"] - reload_mock.assert_awaited_once_with("user-42") + reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False) @pytest.mark.asyncio @@ -198,7 +198,7 @@ async def test_acting_user_auth_returns_admitted_subject_for_non_admin_sessions( result = await acting_user_auth(user_auth) assert result.user_id == "user-42" and result.team_id is None - reload_mock.assert_awaited_once_with("user-42") + reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False) @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a87716375e8..e5dfeb369e2 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -67,7 +67,7 @@ class TestAgentRequestHandler: # Case 1: Both key and team have agents - intersection with patch.object( - AgentRequestHandler, "_get_allowed_agents_for_key" + AgentRequestHandler, "get_allowed_agents_for_key" ) as mock_key: with patch.object( AgentRequestHandler, "_get_allowed_agents_for_team" @@ -86,7 +86,7 @@ class TestAgentRequestHandler: # Case 2: Team has agents, key has none - inherit from team with patch.object( - AgentRequestHandler, "_get_allowed_agents_for_key" + AgentRequestHandler, "get_allowed_agents_for_key" ) as mock_key: with patch.object( AgentRequestHandler, "_get_allowed_agents_for_team" @@ -105,7 +105,7 @@ class TestAgentRequestHandler: # Case 3: Key has agents, team has none - key restrictions stand with patch.object( - AgentRequestHandler, "_get_allowed_agents_for_key" + AgentRequestHandler, "get_allowed_agents_for_key" ) as mock_key: with patch.object( AgentRequestHandler, "_get_allowed_agents_for_team" @@ -120,7 +120,7 @@ class TestAgentRequestHandler: # Case 4: No grant anywhere - unrestricted (documented open-by-default) with patch.object( - AgentRequestHandler, "_get_allowed_agents_for_key" + AgentRequestHandler, "get_allowed_agents_for_key" ) as mock_key: with patch.object( AgentRequestHandler, "_get_allowed_agents_for_team" @@ -141,7 +141,7 @@ class TestAgentRequestHandler: api_key="test-key", user_id="test-user", team_id="test-team" ) - with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key: + with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key: with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team: mock_key.return_value = RestrictedAgentAccess(frozenset({"agent-alpha"})) mock_team.return_value = RestrictedAgentAccess(frozenset({"agent-beta"})) @@ -198,7 +198,7 @@ class TestAgentRequestHandler: @staticmethod def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock: - async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess: + async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess: assert user_api_key_auth is not None return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess()) @@ -249,7 +249,6 @@ class TestAgentRequestHandler: frozenset({"agent-alpha"}) ) - async def test_agent_access_groups_intersect_with_key_grants(self): agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent") resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"})) @@ -299,7 +298,7 @@ class TestAgentRequestHandler: ) as mock_groups: mock_groups.return_value = [] - assert await AgentRequestHandler._get_allowed_agents_for_key( + assert await AgentRequestHandler.get_allowed_agents_for_key( user_api_key_auth=mock_user_auth ) == RestrictedAgentAccess(frozenset()) @@ -315,7 +314,7 @@ class TestAgentRequestHandler: ) as mock_groups: mock_groups.side_effect = Exception("DB Error") - assert await AgentRequestHandler._get_allowed_agents_for_key( + assert await AgentRequestHandler.get_allowed_agents_for_key( user_api_key_auth=mock_user_auth ) == UnrestrictedAgentAccess() @@ -404,7 +403,7 @@ class TestAgentRequestHandler: ) with patch.object( - AgentRequestHandler, "_get_allowed_agents_for_key" + AgentRequestHandler, "get_allowed_agents_for_key" ) as mock_key: with patch.object( AgentRequestHandler, "_get_allowed_agents_for_team" @@ -489,9 +488,9 @@ class TestAgentRequestHandler: listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts) assert {agent.agent_name for agent in listed} == {"alpha", "beta"} - async def test_get_allowed_agents_for_key_via_access_group_ids(self): + async def testget_allowed_agents_for_key_via_access_group_ids(self): """ - Test that _get_allowed_agents_for_key includes agents from key's access_group_ids + Test that get_allowed_agents_for_key includes agents from key's access_group_ids (unified access groups) when key has no native object_permission. """ mock_user_auth = UserAPIKeyAuth( @@ -508,16 +507,16 @@ class TestAgentRequestHandler: new_callable=AsyncMock, return_value=["agent-from-ag-1", "agent-from-ag-2"], ): - result = await AgentRequestHandler._get_allowed_agents_for_key( + result = await AgentRequestHandler.get_allowed_agents_for_key( user_api_key_auth=mock_user_auth ) assert result == RestrictedAgentAccess( frozenset({"agent-from-ag-1", "agent-from-ag-2"}) ) - async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self): + async def testget_allowed_agents_for_key_combines_native_and_access_groups(self): """ - Test that _get_allowed_agents_for_key combines agents from native object_permission + Test that get_allowed_agents_for_key combines agents from native object_permission and key's access_group_ids (unified access groups). """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable @@ -540,7 +539,7 @@ class TestAgentRequestHandler: new_callable=AsyncMock, return_value=["agent-from-ag"], ): - result = await AgentRequestHandler._get_allowed_agents_for_key( + result = await AgentRequestHandler.get_allowed_agents_for_key( user_api_key_auth=mock_user_auth ) assert result == RestrictedAgentAccess( @@ -611,7 +610,7 @@ class TestAgentRequestHandler: "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry, ): - with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key: + with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key: with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team: for key_grant, team_grant in ( ( @@ -632,3 +631,257 @@ class TestAgentRequestHandler: assert await AgentRequestHandler.resolve_agent_access( user_api_key_auth=mock_user_auth ) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "state,allowed", + [ + ({}, True), + ({"enabled": False}, False), + ], +) +async def test_managed_invocation_requires_local_and_directory_admission( + monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool +) -> None: + from unittest.mock import MagicMock + + from litellm.proxy import proxy_server + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + binding: Final = AgentIdentityBinding( + agent_id="target", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + issuer="issuer", + revision="revision", + ) + target: Final = AgentResponse( + agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True + ).model_copy(update=state) + client: Final = MagicMock() + client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + monkeypatch.setattr(proxy_server, "prisma_client", client) + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"]) + auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission) + assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed + + +@pytest.mark.asyncio +@pytest.mark.parametrize("delegated", [True, False]) +async def test_managed_agent_invocation_grants_intersect_verified_user_grants( + monkeypatch: pytest.MonkeyPatch, delegated: bool +) -> None: + from unittest.mock import MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.auth import auth_checks + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import ManagedAgentContext + + database: Final = MagicMock() + monkeypatch.setattr(proxy_server, "prisma_client", database) + own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"]) + human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"]) + human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants) + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human)) + auth: Final = UserAPIKeyAuth(agent_id="actor") + auth.managed_agent_policy = AgentResponse( + agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump() + ) + auth.managed_agent_context = ManagedAgentContext( + agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None + ) + access: Final = await AgentRequestHandler.resolve_agent_access(auth) + assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"})) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"]) +async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches( + monkeypatch: pytest.MonkeyPatch, revoked: str +) -> None: + from unittest.mock import MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + + direct: Final = revoked == "direct-grant" + grouped: Final = revoked == "access-group" + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"]) + human: Final = LiteLLM_UserTable( + user_id="human", + teams=[] if direct else ["team"], + organization_memberships=[], + object_permission_id="grant" if direct else None, + ) + team: Final = LiteLLM_TeamTable( + team_id="team", + models=[], + members_with_roles=[{"user_id": "human", "role": "user"}], + object_permission_id=None if grouped else "grant", + access_group_ids=["group"] if grouped else [], + ) + group: Final = LiteLLM_AccessGroupTable( + access_group_id="group", access_group_name="Group", access_agent_ids=["target"] + ) + cache: Final = UserApiKeyCache() + cache.set_cache("human", human) + cache.set_cache("team_id:team", team) + cache.set_cache(object_permission_cache_key("grant"), permission) + cache.set_cache("access_group_id:group", group) + client: Final = MagicMock() + client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human) + client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team) + client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) + client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group) + monkeypatch.setattr(proxy_server, "prisma_client", client) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + assert await verified_human_agent_grants("human") == frozenset({"target"}) + client.writer_db.litellm_usertable.find_unique.return_value = ( + human.model_copy(update={"teams": []}) if revoked == "user" else human + ) + client.writer_db.litellm_teamtable.find_unique.return_value = ( + team.model_copy(update={"members_with_roles": []}) + if revoked == "team-member" + else team.model_copy(update={"object_permission_id": None}) + if revoked == "team-grant" + else team + ) + client.writer_db.litellm_objectpermissiontable.find_unique.return_value = ( + permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission + ) + client.writer_db.litellm_accessgrouptable.find_unique.return_value = ( + group.model_copy(update={"access_agent_ids": []}) if grouped else group + ) + assert await verified_human_agent_grants("human") == frozenset() + client.db.litellm_usertable.find_unique.assert_not_called() + client.db.litellm_teamtable.find_unique.assert_not_called() + client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + client.db.litellm_accessgrouptable.find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.types.agents import AgentResponse + + stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={}) + registry: Final = AgentRegistry() + registry.register_agent(stale) + database: Final = MagicMock() + database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale]) + database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale]) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(proxy_server, "prisma_client", database) + auth: Final = UserAPIKeyAuth( + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="permission", agent_access_groups=["group"] + ) + ) + assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess( + frozenset({"revoked"}) + ) + database.writer_db.litellm_agentstable.find_many.return_value = [] + assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess( + frozenset() + ) + database.db.litellm_agentstable.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("groups", [[], ["group"]]) +async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None: + assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("team", [False, True]) +async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all( + monkeypatch: pytest.MonkeyPatch, team: bool +) -> None: + from fastapi import HTTPException + from unittest.mock import MagicMock + from litellm.proxy import proxy_server + + database: Final = MagicMock() + database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable")) + database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable")) + monkeypatch.setattr(proxy_server, "prisma_client", database) + auth: Final = UserAPIKeyAuth( + team_id="team" if team else None, + object_permission=None if team else LiteLLM_ObjectPermissionTable( + object_permission_id="grant", agent_access_groups=["group"] + ), + ) + with pytest.raises(HTTPException, match="policy is unavailable") as denied: + await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True) + assert denied.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("available", [False, True]) +async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None: + from unittest.mock import MagicMock + from litellm.proxy import proxy_server + from litellm.proxy.auth import auth_checks + + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None) + monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None)) + assert await AgentRequestHandler._get_allowed_agents_for_team( + UserAPIKeyAuth(team_id="missing"), strict=True + ) == RestrictedAgentAccess(frozenset()) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outage", [False, True]) +async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy( + monkeypatch: pytest.MonkeyPatch, outage: bool +) -> None: + from fastapi import HTTPException + from unittest.mock import MagicMock + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.types.agents import AgentResponse + + registry: Final = AgentRegistry() + registry.register_agent(AgentResponse( + agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True + )) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock( + return_value=None, side_effect=ConnectionError("unavailable") if outage else None + ) + monkeypatch.setattr(proxy_server, "prisma_client", database) + if outage: + with pytest.raises(HTTPException, match="could not be loaded") as denied: + await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) + assert denied.value.status_code == 503 + else: + assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("grant", [False, True]) +async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None: + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import ManagedAgentContext + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants + + auth: Final = UserAPIKeyAuth(agent_id="actor") + auth.managed_agent_policy = AgentResponse( + agent_id="actor", agent_name="Actor", agent_card_params={}, + object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None, + ) + auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated") + assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset()) + assert await verified_human_agent_grants(None) == frozenset() diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py new file mode 100644 index 00000000000..2e37f0f4462 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -0,0 +1,484 @@ +from collections.abc import Mapping +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.auth.managed_authorization import ( + actor_admission_failure, + admit_managed_actor, + invocation_target, +) +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext + +BINDING: Final = AgentIdentityBinding( + agent_id="agent", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + issuer="issuer", + revision="current", +) + + +def agent(**overrides: object) -> AgentResponse: + return AgentResponse.model_validate( + { + "agent_id": "agent", + "agent_name": "Agent", + "agent_card_params": {}, + "identity": BINDING, + "identity_managed": True, + "execution_mode": "both", + **overrides, + } + ) + + +@pytest.mark.parametrize( + "state", + [ + {"enabled": False}, + {"identity": None}, + {"identity": BINDING.model_copy(update={"active": False})}, + {"execution_mode": "delegated"}, + ], +) +def test_keys_cannot_bypass_lifecycle_or_delegated_only_mode(state: dict[str, object]) -> None: + assert isinstance(actor_admission_failure(agent(**state), None), AgentIdentityFailure) + + +@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"]) +def test_keys_cannot_impersonate_an_entra_bound_agent(mode: str) -> None: + assert isinstance(actor_admission_failure(agent(execution_mode=mode), None), AgentIdentityFailure) + + +@pytest.mark.parametrize( + "context", + [ + ManagedAgentContext(agent_id="agent", binding_revision="previous", mode="autonomous"), + ManagedAgentContext(agent_id="another", binding_revision="current", mode="autonomous"), + ManagedAgentContext(agent_id="agent", binding_revision="current", mode="delegated"), + ], +) +def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context: ManagedAgentContext) -> None: + assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure) + + +def test_caller_cannot_construct_trusted_subject_or_policy() -> None: + context: Final = ManagedAgentContext( + agent_id="agent", binding_revision="current", mode="delegated", user_id="human" + ) + auth: Final = UserAPIKeyAuth.model_validate( + { + "managed_agent_context": context, + "requires_fresh_policy": True, + "mcp_explicit_grants_only": True, + "managed_agent_policy": agent(), + "billing_agent_policy": agent(), + "invoked_agent_id": "forged-target", + "agent_invocation_cost": 0.0, + } + ) + assert auth.requires_fresh_policy is False + assert auth.mcp_explicit_grants_only is False + assert "mcp_explicit_grants_only" not in auth.model_dump() + assert auth.managed_agent_context is None + assert auth.managed_agent_policy is None + assert auth.billing_agent_policy is None + assert auth.invoked_agent_id is None + assert auth.agent_invocation_cost is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("autonomous", (True, False)) +async def test_invocation_prepares_target_fee_for_the_correct_agent( + monkeypatch: pytest.MonkeyPatch, + autonomous: bool, +) -> None: + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + + target: Final = agent(litellm_params={"cost_per_query": 0.25}) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(target) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"]) + auth: Final = UserAPIKeyAuth( + agent_id="caller" if autonomous else None, + user_id=None if autonomous else "human", + object_permission=permission, + ) + if autonomous: + caller: Final = agent(agent_id="caller", object_permission=permission.model_dump()) + auth.managed_agent_policy = caller + auth.billing_agent_policy = caller + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) + assert auth.agent_invocation_cost == pytest.approx(0.25) + assert auth.invoked_agent_id == "agent" + assert auth.billing_agent_policy is not None + assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent") + + +@pytest.mark.asyncio +async def test_deleted_agent_key_cannot_fall_back_to_unmanaged_authentication() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_retiredagent.find_unique = AsyncMock(return_value={"original_agent_id": "deleted"}) + with pytest.raises(HTTPException, match="Agent no longer exists"): + await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database)) + database.writer_db.litellm_retiredagent.find_unique.return_value = None + auth: Final = UserAPIKeyAuth(agent_id="legacy-attribution-label") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy is None + database.db.litellm_agentstable.find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("unavailable")) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 503 + + +@pytest.mark.parametrize( + "route,body,expected", + [ + ("/a2a/agent", {}, "agent"), + ("/v1/a2a/agent/", {}, "agent"), + ("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"), + ("/v1/chat/completions", {"model": "a2a/"}, None), + ("/v1/chat/completions", {"model": "ordinary-model"}, None), + ("/a2a", {}, None), + ], +) +def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, object], expected: str | None) -> None: + assert invocation_target(route, body) == expected + + +@pytest.mark.asyncio +async def test_agent_admission_database_outage_fails_closed() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(side_effect=RuntimeError("DB unavailable")) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_human_authentication_does_not_load_an_agent() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock() + await admit_managed_actor(UserAPIKeyAuth(user_id="human"), AgentIdentityStore.from_client(database)) + database.writer_db.litellm_agentstable.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_disabled_agent_key_is_rejected_at_admission() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(enabled=False)) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("permitted", [True, False]) +async def test_verified_human_still_needs_an_explicit_agent_invocation_grant( + monkeypatch: pytest.MonkeyPatch, + permitted: bool, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable + from litellm.proxy.auth import auth_checks + + policy: Final = agent() + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="human-grants", + agents=["agent"] if permitted else [], + ) + human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission) + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human)) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.managed_agent_context = ManagedAgentContext( + agent_id="agent", + binding_revision="current", + mode="delegated", + user_id="human", + ) + if permitted: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy == policy + assert auth.billing_agent_policy == policy + else: + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 403 + + +def test_execution_mode_must_match_verified_token_mode() -> None: + context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context) + assert isinstance(failure, AgentIdentityFailure) + assert "execution mode" in failure.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state,status", [("missing", 403), ("outage", 503), ("denied", 403), ("invalid-fee", 503)]) +async def test_invocation_cannot_bypass_missing_policy_permission_or_invalid_price( + monkeypatch: pytest.MonkeyPatch, state: str, status: int +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + registered: Final = agent(litellm_params={"cost_per_query": -1 if state == "invalid-fee" else 0.25}) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(registered) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock( + return_value=None if state == "missing" else registered, + side_effect=RuntimeError("unavailable") if state == "outage" else None, + ) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="grant", agents=[] if state == "denied" else ["agent"] + ) + auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission) + with pytest.raises(HTTPException) as failure: + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) + assert failure.value.status_code == status + assert auth.agent_invocation_cost is None + + +@pytest.mark.asyncio +async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(execution_mode="autonomous")) + auth: Final = UserAPIKeyAuth(agent_id="agent", jwt_claims={"agent": "agent", "sub": "unrelated-subject"}) + with pytest.raises(HTTPException) as denied: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert denied.value.status_code == 403 + assert auth.managed_agent_policy is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bound", [False, True]) +async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.MonkeyPatch, bound: bool) -> None: + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + + registry: Final = AgentRegistry() + registry.register_agent(agent(identity_managed=bound, identity=BINDING if bound else None)) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(agent_id="agent") + if not bound: + auth.managed_agent_context = ManagedAgentContext( + agent_id="agent", binding_revision="current", mode="autonomous" + ) + with pytest.raises(HTTPException) as denied: + await admit_managed_actor(auth, None) + assert denied.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("managed_flag", [False, True]) +async def test_managed_invocation_requires_database(monkeypatch: pytest.MonkeyPatch, managed_flag: bool) -> None: + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + registry: Final = AgentRegistry() + registry.register_agent(agent(identity_managed=managed_flag)) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + with pytest.raises(HTTPException) as denied: + await prepare_agent_invocation(UserAPIKeyAuth(user_id="human"), "agent", None) + assert denied.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> None: + policy: Final = agent(execution_mode="autonomous") + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + auth: Final = UserAPIKeyAuth(agent_id="agent", api_key="persisted-key") + with pytest.raises(HTTPException, match="bound identity provider token") as denied: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert denied.value.status_code == 403 + assert auth.managed_agent_policy is None + assert auth.billing_agent_policy is None + + +@pytest.mark.asyncio +async def test_unknown_invocation_target_leaves_billing_unset(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry()) + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(user_id="human") + await prepare_agent_invocation(auth, "missing", None) + assert auth.invoked_agent_id is None + assert auth.billing_agent_policy is None + + +@pytest.mark.parametrize( + "route,method,allowed", + [ + ("/v1/agents", "GET", True), + ("/v1/agents", "POST", False), + ("/v1/chat/completions", "POST", True), + ("/v1/chat/completions", "DELETE", False), + ("/openai/deployments/model/chat/completions", "POST", True), + ("/engines/openai/model/chat/completions", "POST", True), + ("/openai/deployments/openai/model/images/generations", "POST", True), + ("/openai/deployments/openai/model/images/edits", "POST", True), + ("/v1beta/models/gemini-model:generateContent", "POST", True), + ("/v1/realtime", "GET", True), + ("/v1/realtime", "POST", False), + ("/v1/realtime/client_secrets", "POST", False), + ("/mcp/tools/call", "POST", True), + ("/a2a/target/message/send", "POST", True), + ("/v1/a2a/target/message/send", "POST", True), + ("/v1/videos", "POST", False), + ("/v1/videos/other-video", "GET", False), + ("/v1/search", "POST", False), + ("/search", "POST", False), + ("/v1/agents/target", "PATCH", False), + ("/v1/responses/other-response", "GET", False), + ("/v1/files", "GET", False), + ("/v1/files", "POST", False), + ("/openai/v1/files", "GET", False), + ("/anthropic/v1/files", "GET", False), + ], +) +def test_managed_route_scope_excludes_provider_resources(route: str, method: str, allowed: bool) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed + + assert managed_agent_route_allowed(route, method) is allowed + + +@pytest.mark.parametrize( + "route,body,settings,cli_model,path_model,expected", + [ + ("/v1/chat/completions", {"model": "body"}, {"completion_model": "default"}, "cli", "path", "default"), + ("/v1/moderations", {"model": "body"}, {"moderation_model": "default"}, "cli", None, "cli"), + ("/v1/audio/speech", {"model": "body"}, {"completion_model": "ignored"}, None, None, "body"), + ("/openai/deployments/path/embeddings", {"model": "body"}, {}, None, "path", "path"), + ("/v1/messages/count_tokens", {"model": "body"}, {"completion_model": "ignored"}, "cli", None, "body"), + ("/mcp/tools/call", {}, {"completion_model": "ignored"}, "cli", None, None), + ("/v1/images/generations", {"model": "image"}, {"completion_model": "text"}, None, None, "image"), + ("/v1/images/generations", {}, {"image_generation_model": "image"}, None, None, "image"), + ("/v1/images/edits", {}, {"image_generation_model": "image"}, None, None, "image"), + ("/v1/rerank", {"model": "reranker"}, {"completion_model": "text"}, "cli", None, "reranker"), + ("/v1beta/models/path:countTokens", {"model": "body"}, {"completion_model": "text"}, "cli", "path", "path"), + ], +) +def test_managed_inference_resolves_dispatch_precedence( + route: str, + body: Mapping[str, object], + settings: Mapping[str, object], + cli_model: str | None, + path_model: str | None, + expected: str | None, +) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + assert managed_inference_request(route, body, settings, cli_model, path_model).get("model") == expected + + +def test_managed_inference_without_any_model_cannot_skip_model_grants(): + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + with pytest.raises(HTTPException, match="explicit or configured model"): + managed_inference_request("/v1/moderations", {}, {}, None) + + +@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/images/generations", "/v1/images/edits"]) +def test_managed_inference_query_model_takes_precedence_over_body(route: str): + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + assert managed_inference_request(route, {"model": "body"}, {}, None, query_model="query")["model"] == "query" + + +def test_managed_inference_ignores_unsupported_query_model(): + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + assert ( + managed_inference_request("/v1/messages", {"model": "body"}, {}, None, query_model="query")["model"] == "body" + ) + + +@pytest.mark.parametrize("route", ["/realtime", "/v1/realtime", "/openai/v1/realtime"]) +def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route: str) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request + + with pytest.raises(HTTPException, match="explicit or configured model"): + managed_inference_request(route, {}, {"completion_model": "allowed-default"}, "cli") + assert managed_inference_request( + route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli" + )["model"] == "requested" + + +@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")]) +def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None: + context: Final = ManagedAgentContext.model_validate( + {"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user} + ) + assert actor_admission_failure(agent(), context) is None + + +@pytest.mark.asyncio +async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation + + legacy: Final = agent(identity=None, identity_managed=False) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(legacy) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(agent_id="agent") + await admit_managed_actor(auth, None) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=legacy) + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy is None + assert auth.billing_agent_policy is None + assert auth.invoked_agent_id is None + + +@pytest.mark.asyncio +async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent()) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy == agent() + assert auth.billing_agent_policy == agent() + assert auth.user_id is None diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 8a7ab0f0001..a5d0d0a3ecc 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -59,6 +59,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): # Mock agent mock_agent = MagicMock() + mock_agent.agent_id = "test-agent" mock_agent.agent_card_params = { "url": "http://backend-agent:10001", "name": "Test Agent", @@ -72,6 +73,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "jsonrpc": "2.0", "id": "test-id", "method": "message/send", + "metadata": {"model_info": {"id": "caller-supplied-id"}}, "params": { "message": { "role": "user", @@ -153,7 +155,7 @@ async def test_invoke_agent_a2a_adds_litellm_data(): "litellm.a2a_protocol.asend_message", new_callable=AsyncMock, return_value=mock_response, - ), + ) as mock_send_message, patch( "litellm.proxy.proxy_server.general_settings", {}, @@ -190,6 +192,9 @@ async def test_invoke_agent_a2a_adds_litellm_data(): mock_add_data.assert_called_once() # Verify model and custom_llm_provider were set + assert mock_send_message.await_args.kwargs["model"] == "a2a_agent/Test Agent" + assert captured_data["metadata"]["model_group"] == "a2a_agent/Test Agent" + assert captured_data["metadata"]["model_info"] == {"id": mock_agent.agent_id} assert captured_data.get("model") == "a2a_agent/Test Agent" assert captured_data.get("custom_llm_provider") == "a2a_agent" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index ef20e88c368..ba160b8e1b4 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -2,11 +2,14 @@ import hashlib import json +from collections.abc import Mapping +from datetime import datetime, timezone from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +from prisma.models import LiteLLM_AgentsTable from litellm.constants import REDACTED_BY_LITELM_STRING from litellm.proxy.agent_endpoints.agent_registry import ( @@ -451,11 +454,11 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update(): registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=None) + return_value=_stored_agent_row(SimpleNamespace(litellm_params={}, object_permission_id=None)) ) mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) - with pytest.raises(Exception, match="Error updating agent in DB") as exc_info: + with pytest.raises(Exception, match="Agent not found") as exc_info: await registry.update_agent_in_db( agent_id="agent-123", agent={ @@ -467,7 +470,7 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update(): updated_by="test-user", ) - assert str(exc_info.value) == "Error updating agent in DB: Agent not found, passed agent_id=agent-123" + assert str(exc_info.value) == "Agent not found, passed agent_id=agent-123" @pytest.mark.asyncio @@ -476,11 +479,13 @@ async def test_patch_agent_in_db_raises_when_row_deleted_mid_update(): registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None} + return_value=_stored_agent_row( + {"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None} + ) ) mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) - with pytest.raises(Exception, match="Error patching agent in DB") as exc_info: + with pytest.raises(Exception, match="Agent not found") as exc_info: await registry.patch_agent_in_db( agent_id="agent-123", agent={"agent_name": "Patched Agent"}, @@ -488,20 +493,43 @@ async def test_patch_agent_in_db_raises_when_row_deleted_mid_update(): updated_by="test-user", ) - assert str(exc_info.value) == "Error patching agent in DB: Agent not found, passed agent_id=agent-123" + assert str(exc_info.value) == "Agent not found, passed agent_id=agent-123" @pytest.mark.asyncio -async def test_delete_agent_from_db_raises_when_row_already_gone(): - """Prisma's delete returns None for a missing row, which dict() cannot consume.""" +async def test_delete_agent_from_db_raises_when_row_already_gone() -> None: registry: Final = AgentRegistry() - mock_prisma: Final = MagicMock() - mock_prisma.db.litellm_agentstable.delete = AsyncMock(return_value=None) + database: Final = MagicMock() + tx: Final = database.tx.return_value.__aenter__.return_value + tx.litellm_agentstable.find_unique = AsyncMock(return_value=None) + with pytest.raises(ValueError, match="Agent not found, passed agent_id=agent-123"): + await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=database) + tx.litellm_verificationtoken.delete_many.assert_not_called() - with pytest.raises(Exception, match="Error deleting agent from DB") as exc_info: - await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=mock_prisma) - assert str(exc_info.value) == "Error deleting agent from DB: Agent not found, passed agent_id=agent-123" +@pytest.mark.asyncio +@pytest.mark.parametrize("managed", [True, False]) +async def test_agent_deletion_revokes_managed_keys_and_keeps_identity_history(managed: bool) -> None: + registry: Final = AgentRegistry() + database: Final = MagicMock() + tx: Final = database.tx.return_value.__aenter__.return_value + row: Final = _stored_agent_row({"agent_id": "agent-123", "identity_managed": managed}) + tx.litellm_agentstable.find_unique = AsyncMock(return_value=row) + tx.litellm_agentstable.delete = AsyncMock(return_value=row) + tx.litellm_verificationtoken.delete_many = AsyncMock(return_value=2) + tx.litellm_retiredagent.upsert = AsyncMock() + result: Final = await registry.delete_agent_from_db("agent-123", database) + assert result["agent_id"] == "agent-123" + tx.litellm_agentstable.delete.assert_awaited_once_with(where={"agent_id": "agent-123"}) + if managed: + tx.litellm_retiredagent.upsert.assert_awaited_once_with( + where={"original_agent_id": "agent-123"}, + data={"create": {"original_agent_id": "agent-123"}, "update": {}}, + ) + tx.litellm_verificationtoken.delete_many.assert_awaited_once_with(where={"agent_id": "agent-123"}) + else: + tx.litellm_retiredagent.upsert.assert_not_awaited() + tx.litellm_verificationtoken.delete_many.assert_not_awaited() # ---------- LIT-6736: agent litellm_params secret redaction ---------- @@ -729,14 +757,15 @@ async def test_update_agent_in_db_preserves_secret_when_echoed_back_redacted(): mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace( - litellm_params={ - "aws_access_key_id": SENTINEL_AWS_ACCESS_KEY_ID, - "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, - "model": "bedrock/agentcore/my-agent", - }, - object_permission_id=None, - kill_switch=None, + return_value=_stored_agent_row( + SimpleNamespace( + litellm_params={ + "aws_access_key_id": SENTINEL_AWS_ACCESS_KEY_ID, + "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, + "model": "bedrock/agentcore/my-agent", + }, + object_permission_id=None, + ) ) ) updated_agent = MagicMock() @@ -782,10 +811,11 @@ async def test_update_agent_in_db_preserves_secret_when_key_omitted_entirely(): mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace( - litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, - object_permission_id=None, - kill_switch=None, + return_value=_stored_agent_row( + SimpleNamespace( + litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, + object_permission_id=None, + ) ) ) updated_agent = MagicMock() @@ -824,15 +854,16 @@ async def test_update_agent_in_db_preserves_secret_nested_under_a_non_sensitive_ mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace( - litellm_params={ - "provider_config": { - "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, - "region": "us-east-1", - } - }, - object_permission_id=None, - kill_switch=None, + return_value=_stored_agent_row( + SimpleNamespace( + litellm_params={ + "provider_config": { + "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, + "region": "us-east-1", + } + }, + object_permission_id=None, + ) ) ) updated_agent = MagicMock() @@ -878,10 +909,11 @@ async def test_update_agent_in_db_clears_secret_on_explicit_empty_value(): mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace( - litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, - object_permission_id=None, - kill_switch=None, + return_value=_stored_agent_row( + SimpleNamespace( + litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, + object_permission_id=None, + ) ) ) updated_agent = MagicMock() @@ -919,12 +951,14 @@ async def test_patch_agent_in_db_preserves_secret_when_litellm_params_omitted(): mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Old Name", - "litellm_params": {"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, - "object_permission_id": None, - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "Old Name", + "litellm_params": {"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY}, + "object_permission_id": None, + } + ) ) patched_agent = MagicMock() patched_agent.model_dump.return_value = { @@ -958,15 +992,17 @@ async def test_patch_agent_in_db_preserves_secret_when_echoed_back_redacted(): mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Test Agent", - "litellm_params": { - "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, - "is_public": False, - }, - "object_permission_id": None, - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "Test Agent", + "litellm_params": { + "aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY, + "is_public": False, + }, + "object_permission_id": None, + } + ) ) patched_agent = MagicMock() patched_agent.model_dump.return_value = { @@ -997,6 +1033,48 @@ async def test_patch_agent_in_db_preserves_secret_when_echoed_back_redacted(): assert stored_params["is_public"] is True +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["patch", "put"]) +async def test_runtime_update_drops_legacy_identity_and_keeps_agent_id(operation: str) -> None: + registry: Final = AgentRegistry() + prisma: Final = MagicMock() + identity: Final = { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + } + existing_params: Final = {"identity": identity, "model": "old"} + existing: Final = ( + SimpleNamespace(litellm_params=existing_params, object_permission_id=None) + if operation == "put" + else {"agent_name": "Readable agent", "litellm_params": existing_params} + ) + prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row(existing)) + saved: Final = MagicMock() + saved.object_permission = None + saved.model_dump.return_value = { + "agent_id": "unchanged-id", + "agent_name": "Renamed agent", + "agent_card_params": {}, + "litellm_params": {"model": "new"}, + } + prisma.db.litellm_agentstable.update = AsyncMock(return_value=saved) + update: Final = registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db + result: Final = await update( + agent_id="unchanged-id", + agent={"agent_name": "Renamed agent", "agent_card_params": {}, "litellm_params": {"model": "new"}}, + prisma_client=prisma, + updated_by="admin", + ) + stored: Final = prisma.db.litellm_agentstable.update.call_args.kwargs + assert stored["where"] == {"agent_id": "unchanged-id"} + assert json.loads(stored["data"]["litellm_params"]) == {"model": "new"}, ( + "a stored litellm_params.identity must not be resurrected once the JWT path no longer honours it" + ) + assert result.agent_id == "unchanged-id" + assert "object_permission_id" not in stored["data"] + + def _agent_row_mock(access_group_ids: list[str]) -> MagicMock: row: Final = MagicMock() row.model_dump.return_value = { @@ -1063,13 +1141,15 @@ async def test_patch_agent_in_db_replaces_access_group_ids_when_provided( registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Test Agent", - "litellm_params": {}, - "object_permission_id": None, - "access_group_ids": ["ag-1"], - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "Test Agent", + "litellm_params": {}, + "object_permission_id": None, + "access_group_ids": ["ag-1"], + } + ) ) mock_update = AsyncMock(return_value=_agent_row_mock(expected)) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1086,13 +1166,15 @@ async def test_patch_agent_in_db_keeps_access_group_ids_when_omitted(): registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Old Name", - "litellm_params": {}, - "object_permission_id": None, - "access_group_ids": ["ag-1"], - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "Old Name", + "litellm_params": {}, + "object_permission_id": None, + "access_group_ids": ["ag-1"], + } + ) ) mock_update = AsyncMock(return_value=_agent_row_mock(["ag-1"])) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1114,8 +1196,8 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace( - litellm_params={}, object_permission_id=None, kill_switch=None, access_group_ids=["ag-1"] + return_value=_stored_agent_row( + SimpleNamespace(litellm_params={}, object_permission_id=None, access_group_ids=["ag-1"]) ) ) mock_update = AsyncMock(return_value=_agent_row_mock(expected)) @@ -1134,6 +1216,34 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro assert tuple(mock_update.call_args.kwargs["data"]["access_group_ids"]) == tuple(expected) +def _stored_agent_row(values: Mapping[str, object] | SimpleNamespace) -> LiteLLM_AgentsTable: + fields: Final = vars(values) if isinstance(values, SimpleNamespace) else values + return LiteLLM_AgentsTable.model_validate( + { + "agent_id": "agent-123", + "agent_name": "Test Agent", + "agent_card_params": "{}", + "extra_headers": [], + "agent_access_groups": [], + "access_group_ids": [], + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": "admin", + "updated_by": "admin", + "spend": 0, + "identity_managed": False, + "enabled": True, + "execution_mode": "autonomous", + **{ + key: json.dumps(value) + if key in ("litellm_params", "agent_card_params", "kill_switch") and not isinstance(value, str) + else value + for key, value in fields.items() + }, + } + ) + + _KILL_SWITCH: Final = { "url": "https://ops.example.com/kill", "method": "POST", @@ -1194,13 +1304,15 @@ async def test_patch_agent_in_db_keeps_kill_switch_when_omitted_and_clears_it_on registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Old", - "litellm_params": {}, - "object_permission_id": None, - "kill_switch": _KILL_SWITCH, - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "Old", + "litellm_params": {}, + "object_permission_id": None, + "kill_switch": _KILL_SWITCH, + } + ) ) mock_update = AsyncMock(return_value=_agent_row_mock([])) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1223,13 +1335,15 @@ async def test_patch_agent_in_db_restores_the_stored_kill_switch_secret_behind_t registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "A", - "litellm_params": {}, - "object_permission_id": None, - "kill_switch": _KILL_SWITCH, - } + return_value=_stored_agent_row( + { + "agent_id": "agent-123", + "agent_name": "A", + "litellm_params": {}, + "object_permission_id": None, + "kill_switch": _KILL_SWITCH, + } + ) ) mock_update = AsyncMock(return_value=_agent_row_mock([])) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1258,7 +1372,9 @@ async def test_update_agent_in_db_clears_kill_switch_when_omitted_and_restores_s registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH)) + return_value=_stored_agent_row( + SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH)) + ) ) mock_update = AsyncMock(return_value=_agent_row_mock([])) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1284,3 +1400,145 @@ def test_load_agents_from_config_exposes_a_typed_kill_switch(): (agent,) = registry.get_agent_list() assert agent.kill_switch is not None assert agent.kill_switch.model_dump() == _KILL_SWITCH + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bound", [False, True]) +async def test_agent_listing_preserves_stored_identity_bindings(bound: bool) -> None: + from datetime import datetime, timezone + + from prisma.models import LiteLLM_AgentIdentity, LiteLLM_AgentsTable + + from litellm.types.agents import AgentResponse + + binding: Final = LiteLLM_AgentIdentity( + agent_id="agent", + provider="microsoft_entra", + issuer="issuer", + tenant_id="tenant", + client_id="client", + active=True, + required_roles=[], + required_scopes=["user_impersonation"], + revision="revision", + ) + row: Final = LiteLLM_AgentsTable( + agent_id="agent", + agent_name="Bound agent", + agent_card_params="{}", + identity_managed=bound, + identity=binding if bound else None, + enabled=True, + execution_mode="autonomous", + spend=0.0, + agent_access_groups=[], + access_group_ids=[], + extra_headers=[], + created_by="admin", + updated_by="admin", + created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + updated_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + ) + client: Final = MagicMock() + client.db.litellm_agentstable.find_many = AsyncMock(return_value=[row]) + listed: Final = await AgentRegistry.get_all_agents_from_db(client) + response: Final = AgentResponse.model_validate(listed[0]) + if bound: + assert response.identity is not None + assert response.identity.client_id == binding.client_id + assert response.identity.revision == binding.revision + else: + assert response.identity is None + client.db.litellm_agentstable.find_many.assert_awaited_once_with( + order={"created_at": "desc"}, + include={"object_permission": True, "identity": True}, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "patch", "put"]) +async def test_agent_permissions_are_written_atomically_with_the_registration(operation: str) -> None: + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + registry: Final = AgentRegistry() + client: Final = MagicMock() + existing: Final = _stored_agent_row({"agent_id": "agent-123", "object_permission_id": "permissions"}) + client.db.litellm_agentstable.find_unique = AsyncMock(return_value=existing) + client.db.litellm_agentstable.create = AsyncMock(return_value=existing) + client.db.litellm_agentstable.update = AsyncMock(return_value=existing) + client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=( + LiteLLM_ObjectPermissionTable(object_permission_id="permissions", models=["prior"], mcp_servers=["slack"]) + if operation != "create" + else None + ) + ) + incoming: Final = {"agent_name": "Agent", "agent_card_params": {}, "object_permission": {"models": ["new"]}} + if operation == "create": + await registry.add_agent_to_db(incoming, client, created_by="admin") + else: + update: Final = registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db + await update("agent-123", incoming, client, updated_by="admin") + write: Final = ( + client.db.litellm_agentstable.create if operation == "create" else client.db.litellm_agentstable.update + ) + permission: Final = write.call_args.kwargs["data"]["object_permission"][ + "create" if operation == "create" else "update" + ] + assert permission["models"] == ["new"] + if operation != "create": + assert permission["mcp_servers"] == ["slack"] + assert permission["object_permission_id"] == "permissions" + client.db.litellm_objectpermissiontable.update.assert_not_called() + client.db.litellm_objectpermissiontable.create.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "patch", "put"]) +async def test_invalid_identity_fails_before_registration_is_written(operation: str) -> None: + from fastapi import HTTPException + + registry: Final = AgentRegistry() + client: Final = MagicMock() + client.db.litellm_agentstable.create = AsyncMock() + client.db.litellm_agentstable.update = AsyncMock() + client.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row({"agent_id": "agent-123"})) + incoming: Final = {"agent_name": "Agent", "agent_card_params": {}, "identity": {"provider": "unknown"}} + write: Final = ( + registry.add_agent_to_db(incoming, client, created_by="admin") + if operation == "create" + else (registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db)( + "agent-123", incoming, client, updated_by="admin" + ) + ) + with pytest.raises(HTTPException) as failure: + await write + assert failure.value.status_code == 400 + client.db.litellm_agentstable.create.assert_not_awaited() + client.db.litellm_agentstable.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "patch", "put"]) +async def test_duplicate_agent_binding_returns_conflict_for_every_write(operation: str) -> None: + from fastapi import HTTPException + from prisma.errors import UniqueViolationError + + registry: Final = AgentRegistry() + client: Final = MagicMock() + client.db.litellm_agentstable.find_unique = AsyncMock(return_value=_stored_agent_row({"agent_id": "agent-123"})) + failure: Final = UniqueViolationError({"user_facing_error": {"message": "Unique constraint failed", "meta": {"target": ["client_id"]}, "error_code": "P2002"}}) + client.db.litellm_agentstable.create = AsyncMock(side_effect=failure) + client.db.litellm_agentstable.update = AsyncMock(side_effect=failure) + incoming: Final = {"agent_name": "Agent", "agent_card_params": {}} + write: Final = ( + registry.add_agent_to_db(incoming, client, created_by="admin") + if operation == "create" + else (registry.patch_agent_in_db if operation == "patch" else registry.update_agent_in_db)( + "agent-123", incoming, client, updated_by="admin" + ) + ) + with pytest.raises(HTTPException) as denied: + await write + assert denied.value.status_code == 409 + assert denied.value.detail == "Agent name or Entra application is already registered" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 526f24c5221..4bdb066377a 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -1,11 +1,13 @@ import json +from datetime import datetime, timezone + from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from fastapi import FastAPI +from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient from litellm.constants import REDACTED_BY_LITELM_STRING @@ -21,7 +23,8 @@ from litellm.proxy.agent_endpoints.endpoints import ( router, user_api_key_auth, ) -from litellm.types.agents import AgentResponse +from litellm.types.agents import AgentResponse, PatchAgentRequest +from litellm.types.proxy.agent_identity import AgentIdentityBinding def _sample_agent_card_params() -> dict: @@ -97,7 +100,7 @@ def test_update_agent_success(mock_prisma_client, mock_user_api_key_auth, monkey "agent_card_params": _sample_agent_card_params(), } mock_prisma_client.db.litellm_agentstable.find_unique = AsyncMock( - return_value=existing_agent + return_value=AgentResponse.model_validate(existing_agent) ) mock_registry = MagicMock() @@ -350,6 +353,7 @@ class TestAgentByIdKeyRedaction: test_client = _make_app_with_role(role) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.writer_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=None ) @@ -412,6 +416,7 @@ class TestAgentRBACInternalUser: return_value=_sample_agent_response() ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.writer_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=None ) @@ -592,6 +597,24 @@ class TestAgentRBACProxyAdmin: ) assert resp.status_code == 200 + def test_create_agent_rejects_legacy_litellm_params_identity(self): + with patch("litellm.proxy.proxy_server.prisma_client"): # test-quality-ok: proxy_server module global is the endpoint's only injection point + self.mock_registry.get_agent_by_name = MagicMock(return_value=None) + self.mock_registry.add_agent_to_db = AsyncMock(return_value=_sample_agent_response()) + config = _sample_agent_config() + config["litellm_params"] = { + **config["litellm_params"], + "identity": { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + }, + } + resp = self.admin_client.post("/v1/agents", json=config, headers={"Authorization": "Bearer k"}) + assert resp.status_code == 400, resp.text + assert "top-level identity field" in resp.json()["detail"] + self.mock_registry.add_agent_to_db.assert_not_awaited() + def test_create_agent_applies_litellm_merge_to_stored_card(self): """The card stored in the DB must reflect the LiteLLM-fronting merge.""" with patch("litellm.proxy.proxy_server.prisma_client"): @@ -663,11 +686,9 @@ class TestAgentRBACProxyAdmin: """LIT-6736: PUT /v1/agents/{id} must not echo the stored secret back.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Existing Agent", - "agent_card_params": _sample_agent_card_params(), - } + return_value=AgentResponse( + agent_id="agent-123", agent_name="Existing Agent", agent_card_params=_sample_agent_card_params() + ) ) self.mock_registry.update_agent_in_db = AsyncMock( return_value=AgentResponse( @@ -698,11 +719,9 @@ class TestAgentRBACProxyAdmin: """LIT-6736: PATCH /v1/agents/{id} must not echo the stored secret back.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( - return_value={ - "agent_id": "agent-123", - "agent_name": "Existing Agent", - "agent_card_params": _sample_agent_card_params(), - } + return_value=AgentResponse( + agent_id="agent-123", agent_name="Existing Agent", agent_card_params=_sample_agent_card_params() + ) ) self.mock_registry.patch_agent_in_db = AsyncMock( return_value=AgentResponse( @@ -1140,6 +1159,143 @@ def test_make_agent_public_rejects_an_agent_published_only_in_the_db(monkeypatch assert "already in public agent groups" in duplicate.json()["detail"] +@pytest.mark.parametrize("enabled, claim_field, expected", [(True, "azp", True), (False, "azp", False), (True, None, False)]) +def test_jwt_authentication_status_does_not_require_virtual_keys( + monkeypatch: pytest.MonkeyPatch, enabled: bool, claim_field: str | None, expected: bool +) -> None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.handle_jwt import JWTHandler + + handler: Final = JWTHandler() + handler.update_environment(None, DualCache(), LiteLLM_JWTAuth(agent_id_jwt_field=claim_field)) + monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": enabled}) + monkeypatch.setattr(proxy_server, "jwt_handler", handler) + agent: Final = _sample_agent_response() + response: Final = agent_endpoints._redact_sensitive_agent_fields((agent,), is_admin=True)[0] + assert response.jwt_auth_configured is expected + assert agent.jwt_auth_configured is False + + +def test_identity_providers_require_configured_issuer_and_audience(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.handle_jwt import JWTHandler + + handler: Final = JWTHandler() + handler.update_environment(None, DualCache(), LiteLLM_JWTAuth()) + monkeypatch.setattr(proxy_server, "jwt_handler", handler) + monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": True}) + monkeypatch.setenv("JWT_ISSUER", "https://issuer.example") + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + assert client.get("/v1/agents/identity/providers").json() == [] + monkeypatch.setenv("JWT_AUDIENCE", "gateway") + response: Final = client.get("/v1/agents/identity/providers") + assert response.status_code == 200 + assert response.json() == ["https://issuer.example"] + forbidden: Final = _make_app_with_role(LitellmUserRoles.INTERNAL_USER).get("/v1/agents/identity/providers") + assert forbidden.status_code == 403 + + +def test_identity_evidence_is_persisted_and_never_taken_from_runtime_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + binding: Final = AgentIdentityBinding( + agent_id="bound", + provider="microsoft_entra", + tenant_id="11111111-1111-4111-8111-111111111111", + client_id="22222222-2222-4222-8222-222222222222", + issuer="https://issuer.example", + revision="revision-one", + ) + bound: Final = AgentResponse( + agent_id="bound", + agent_name="Readable name", + agent_card_params={}, + identity=binding, + identity_managed=True, + litellm_params={"last_authenticated_at": "forged-proof"}, + ) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=bound) + monkeypatch.setattr(proxy_server, "prisma_client", database) + pending: Final = client.get("/v1/agents/bound/identity") + assert pending.status_code == 200 + assert pending.json()["last_authenticated_at"] is None + verified_binding: Final = binding.model_copy( + update={"last_authenticated_at": datetime(2026, 1, 1, tzinfo=timezone.utc)} + ) + database.writer_db.litellm_agentstable.find_unique.return_value = bound.model_copy(update={"identity": verified_binding}) + verified: Final = client.get("/v1/agents/bound/identity") + assert verified.json()["last_authenticated_at"] == "2026-01-01T00:00:00Z" + assert verified.json()["identity"]["client_id"] == binding.client_id + database.writer_db.litellm_agentstable.find_unique.return_value = None + assert client.get("/v1/agents/missing/identity").status_code == 404 + database.writer_db.litellm_agentstable.find_unique.side_effect = RuntimeError("unavailable") + assert client.get("/v1/agents/bound/identity").status_code == 503 + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_identity_providers_honor_issuer_specific_audiences_and_global_fallback( + monkeypatch: pytest.MonkeyPatch, enabled: bool +) -> None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy import proxy_server + from litellm.proxy._types import JWTIssuerConfig, LiteLLM_JWTAuth + from litellm.proxy.auth.handle_jwt import JWTHandler + + handler: Final = JWTHandler() + handler.update_environment( + None, + DualCache(), + LiteLLM_JWTAuth( + issuers=[ + JWTIssuerConfig(issuer="https://scoped.example", audience="gateway"), + JWTIssuerConfig(issuer="https://unscoped.example", disable_audience_validation=True), + ] + ), + ) + monkeypatch.setattr(proxy_server, "jwt_handler", handler) + monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": enabled}) + monkeypatch.setenv("JWT_ISSUER", "https://global.example") + monkeypatch.setenv("JWT_AUDIENCE", "gateway") + assert client.get("/v1/agents/identity/providers").json() == ( + ["https://scoped.example", "https://global.example"] if enabled else [] + ) + monkeypatch.setenv("JWT_ISSUER", "https://unscoped.example") + assert client.get("/v1/agents/identity/providers").json() == (["https://scoped.example"] if enabled else []) + + +@pytest.mark.parametrize("change", ({"execution_mode": "delegated"}, {"execution_mode": "both"})) +def test_mode_only_edit_requires_the_existing_identity_sso_tenant( + monkeypatch: pytest.MonkeyPatch, change: PatchAgentRequest +) -> None: + from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING, TENANT, managed_agent + + monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,)) + monkeypatch.delenv("MICROSOFT_TENANT", raising=False) + monkeypatch.setenv("MICROSOFT_CLIENT_ID", "gateway-client") + with pytest.raises(HTTPException, match="Delegated agents require Microsoft SSO"): + agent_endpoints._validate_managed_identity_request(change, managed_agent()) + monkeypatch.setenv("MICROSOFT_TENANT", TENANT) + agent_endpoints._validate_managed_identity_request(change, managed_agent()) + + +def test_identity_only_edit_preserves_delegated_mode_validation(monkeypatch: pytest.MonkeyPatch) -> None: + from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING, managed_agent + + monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,)) + monkeypatch.delenv("MICROSOFT_TENANT", raising=False) + configuration: Final = BINDING.model_dump( + exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"} + ) + delegated: Final = managed_agent().model_copy(update={"execution_mode": "delegated"}) + with pytest.raises(HTTPException, match="Delegated agents require Microsoft SSO"): + agent_endpoints._validate_managed_identity_request({"identity": configuration}, delegated) + _KILL_SWITCH: Final = { "url": "https://ops.example.com/kill", "method": "POST", @@ -1342,6 +1498,7 @@ def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_other def _get_as(role: LitellmUserRoles): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.writer_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) return _make_app_with_role(role).get("/v1/agents/agent-123", headers={"Authorization": "Bearer k"}) @@ -1357,3 +1514,80 @@ def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_other assert internal.status_code == 200, internal.text assert internal.json()["kill_switch"] is None assert "tok-real" not in internal.text + + +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +@pytest.mark.parametrize("path", ["/v1/agents", "/v1/agents/agent-123"]) +def test_agent_identity_configuration_is_only_returned_to_admins(role, path, monkeypatch): + from litellm.proxy.agent_endpoints import agent_registry + + binding = AgentIdentityBinding( + agent_id="agent-123", provider="microsoft_entra", tenant_id="tenant", client_id="client", + issuer="https://login.microsoftonline.com/tenant/v2.0", revision="revision", + ) + agent = _sample_agent_response().model_copy(update={"identity": binding}) + registry = MagicMock() + registry.get_agent_by_id.return_value = agent + registry.get_agent_list.return_value = [agent] + registry.ids_for_agent.return_value = frozenset({agent.agent_id}) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.resolve_agent_access", + AsyncMock(return_value=RestrictedAgentAccess(frozenset({agent.agent_id}))), + ) + with patch("litellm.proxy.proxy_server.prisma_client") as prisma: + prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + prisma.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + response = _make_app_with_role(role).get(path, headers={"Authorization": "Bearer k"}) + assert response.status_code == 200 + payload = response.json()[0] if path == "/v1/agents" else response.json() + assert payload["identity"] == (binding.model_dump(mode="json") if role == LitellmUserRoles.PROXY_ADMIN else None) + assert agent.identity == binding + + +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.INTERNAL_USER]) +def test_agent_detail_cache_miss_preserves_admin_identity_visibility(role, monkeypatch): + binding = AgentIdentityBinding( + agent_id="agent-123", provider="microsoft_entra", tenant_id="tenant", client_id="client", + issuer="https://login.microsoftonline.com/tenant/v2.0", revision="revision", + ) + agent = _sample_agent_response() + registry = MagicMock() + registry.get_agent_by_id.return_value = None + registry.ids_for_agent.return_value = frozenset({agent.agent_id}) + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry) + monkeypatch.setattr( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + AsyncMock(return_value=True), + ) + + async def load_row(*, where, include): + assert where == {"agent_id": agent.agent_id} + return agent.model_copy(update={"identity": binding if include.get("identity") else None}) + + with patch("litellm.proxy.proxy_server.prisma_client") as prisma: + prisma.db.litellm_agentstable.find_unique = AsyncMock(side_effect=load_row) + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + response = _make_app_with_role(role).get("/v1/agents/agent-123") + assert response.status_code == 200 + assert response.json()["identity"] == (binding.model_dump(mode="json") if role == LitellmUserRoles.PROXY_ADMIN else None) + + +@pytest.mark.parametrize("trusted", [False, True]) +def test_invalid_identity_and_untrusted_tenant_cannot_be_registered( + monkeypatch: pytest.MonkeyPatch, trusted: bool +) -> None: + from tests.test_litellm.proxy.agent_endpoints.test_managed_identity import BINDING + + configuration: Final = BINDING.model_dump( + exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"} + ) + monkeypatch.setattr(agent_endpoints, "_trusted_agent_issuers", lambda: (BINDING.issuer,) if trusted else ()) + request: Final = {"identity": {**configuration, "client_id": "invalid"} if trusted else configuration} + message: Final = "Invalid Entra identity configuration" if trusted else "Configure trusted JWT issuer" + with pytest.raises(HTTPException, match=message) as failure: + agent_endpoints._validate_managed_identity_request(request) + assert failure.value.status_code == 400 diff --git a/tests/test_litellm/proxy/agent_endpoints/test_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_identity.py new file mode 100644 index 00000000000..c9d803fdae7 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_identity.py @@ -0,0 +1,27 @@ +from collections.abc import Mapping + +import pytest +from fastapi import HTTPException + +from litellm.proxy.agent_endpoints.identity import has_legacy_identity, reject_legacy_identity + +TENANT = "11111111-1111-4111-8111-111111111111" +CLIENT = "22222222-2222-4222-8222-222222222222" + + +@pytest.mark.parametrize("params", [None, {}, {"model": "gpt-4o", "api_key": "sk-test"}]) +def test_runtime_params_without_identity_are_accepted(params: Mapping[str, object] | None) -> None: + assert has_legacy_identity(params) is False + reject_legacy_identity(params) + + +@pytest.mark.parametrize( + "identity", [None, {}, {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}] +) +def test_legacy_litellm_params_identity_is_rejected(identity: object) -> None: + params: Mapping[str, object] = {"model": "gpt-4o", "identity": identity} + assert has_legacy_identity(params) is True + with pytest.raises(HTTPException) as failure: + reject_legacy_identity(params) + assert failure.value.status_code == 400 + assert "top-level identity field" in failure.value.detail diff --git a/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py new file mode 100644 index 00000000000..faace1ccf20 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_identity_store.py @@ -0,0 +1,402 @@ +from datetime import datetime, timezone +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException +from prisma.models import LiteLLM_VerifiedSubject + +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent +from litellm.repositories.table_repositories import ( + AgentIdentityRepository, + AgentsRepository, + VerifiedSubjectRepository, +) +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import ( + AgentIdentityBinding, + AgentIdentityFailure, + ManagedAgentContext, + MicrosoftInteractiveSubject, +) + +TENANT: Final = "11111111-1111-4111-8111-111111111111" +CLIENT: Final = "22222222-2222-4222-8222-222222222222" +PRINCIPAL: Final = "33333333-3333-4333-8333-333333333333" +HUMAN: Final = "44444444-4444-4444-8444-444444444444" +ISSUER: Final = f"https://login.microsoftonline.com/{TENANT}/v2.0" +BINDING: Final = AgentIdentityBinding( + agent_id="agent-one", + provider="microsoft_entra", + tenant_id=TENANT, + client_id=CLIENT, + service_principal_id=PRINCIPAL, + issuer=ISSUER, + required_roles=("Agent.Invoke",), + revision="revision-one", +) +CLAIMS: Final = {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": PRINCIPAL, "roles": ["Agent.Invoke"]} + + +def stored_agent(**overrides: object) -> AgentResponse: + return AgentResponse.model_validate( + { + "agent_id": "agent-one", + "agent_name": "Research", + "agent_card_params": {}, + "identity": BINDING, + "identity_managed": True, + "execution_mode": "both", + **overrides, + } + ) + + +def setup_store( + agent: AgentResponse | None = stored_agent(), + human: LiteLLM_VerifiedSubject | None = None, +) -> tuple[AgentIdentityStore, AsyncMock, AsyncMock, AsyncMock]: + agents: Final = AsyncMock() + identities: Final = AsyncMock() + humans: Final = AsyncMock() + agents.find_unique.return_value = agent + identities.find_unique.return_value = BINDING + identities.update_many.return_value = 1 + humans.find_unique.return_value = human + db: Final = SimpleNamespace( + db=SimpleNamespace( + litellm_agentstable=agents, + litellm_agentidentity=identities, + litellm_verifiedsubject=humans, + ) + ) + return ( + AgentIdentityStore(AgentsRepository(db), AgentIdentityRepository(db), VerifiedSubjectRepository(db)), + agents, + identities, + humans, + ) + + +@pytest.mark.asyncio +async def test_application_authentication_has_no_fabricated_human() -> None: + store, _, _, humans = setup_store() + result: Final = await store.resolve_verified_claims(CLAIMS) + assert isinstance(result, ManagedAgentContext) + assert result.agent_id == "agent-one" + assert result.mode == "autonomous" + assert result.user_id is None + humans.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_lifecycle_is_read_on_every_request_without_cached_allow() -> None: + store, agents, _, _ = setup_store() + agents.find_unique.side_effect = [stored_agent(), stored_agent(enabled=False)] + assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext) + denial: Final = await store.resolve_verified_claims(CLAIMS) + assert isinstance(denial, AgentIdentityFailure) + assert denial.code == "identity_denied" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("unavailable_table", ["agents", "identities", "humans"]) +async def test_identity_store_failure_never_becomes_a_legacy_allow(unavailable_table: str) -> None: + store, agents, identities, humans = setup_store() + table: Final = {"agents": agents, "identities": identities, "humans": humans}[unavailable_table] + table.find_unique.side_effect = RuntimeError("database unavailable") + result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"}) + assert isinstance(result, AgentIdentityFailure) + assert result.code == "policy_unavailable" + + +@pytest.mark.asyncio +async def test_unclassified_delegated_subject_cannot_authenticate_as_a_user() -> None: + store, _, _, _ = setup_store() + result: Final = await store.resolve_verified_claims( + {**CLAIMS, "oid": HUMAN, "scp": "user_impersonation", "idtyp": "user"} + ) + assert isinstance(result, AgentIdentityFailure) + assert "first sign in" in result.message + + +@pytest.mark.asyncio +async def test_delegated_subject_uses_canonical_sso_user_not_email_claim() -> None: + human: Final = LiteLLM_VerifiedSubject( + kind="human", + subject_id="subject-one", + issuer=ISSUER, + tenant_id=TENANT, + oid=HUMAN, + user_id="canonical-user", + verified_via="sso_interactive", + verified_at=datetime.now(timezone.utc), + ) + store, _, _, humans = setup_store(human=human) + result: Final = await store.resolve_verified_claims( + { + **CLAIMS, + "oid": HUMAN, + "scp": "user_impersonation", + "email": "untrusted-alias@example.com", + } + ) + assert isinstance(result, ManagedAgentContext) + assert result.mode == "delegated" + assert result.user_id == "canonical-user" + humans.find_unique.assert_awaited_once_with( + where={"issuer_tenant_id_oid": {"issuer": ISSUER, "tenant_id": TENANT, "oid": HUMAN}} + ) + + +@pytest.mark.asyncio +async def test_rebinding_during_authentication_does_not_mark_new_identity_verified() -> None: + store, _, identities, _ = setup_store() + identities.update_many.return_value = 0 + context: Final = ManagedAgentContext(agent_id="agent-one", binding_revision="old-revision", mode="autonomous") + result: Final = await store.record_authentication(context) + assert isinstance(result, AgentIdentityFailure) + assert "changed" in result.message + assert identities.update_many.call_args.kwargs["where"] == { + "agent_id": "agent-one", + "revision": "old-revision", + "active": True, + "agent": {"is": {"enabled": True, "identity_managed": True}}, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("agent", [None, stored_agent(identity=None), stored_agent(identity_managed=False)]) +async def test_stale_binding_cannot_bypass_lifecycle(agent: AgentResponse | None) -> None: + store, _, _, _ = setup_store(agent=agent) + assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure) + + +@pytest.mark.asyncio +async def test_unrelated_non_entra_claims_do_not_query_identity_store() -> None: + store, agents, identities, _ = setup_store() + assert await store.resolve_verified_claims({"sub": "ordinary-user"}) is None + identities.find_unique.assert_not_awaited() + agents.find_unique.assert_not_awaited() + + +HUMAN_CLAIMS: Final = {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": HUMAN, "scp": "user_impersonation"} + + +@pytest.mark.asyncio +async def test_bound_agents_and_policy_failures_are_never_served_from_the_miss_cache() -> None: + store, _, identities, _ = setup_store() + assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext) + assert isinstance(await store.resolve_verified_claims(CLAIMS), ManagedAgentContext) + assert identities.find_unique.await_count == 2 + identities.find_unique.side_effect = ConnectionError("database down") + assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure) + assert isinstance(await store.resolve_verified_claims(CLAIMS), AgentIdentityFailure) + assert identities.find_unique.await_count == 4 + + +@pytest.mark.asyncio +async def test_retired_client_cannot_fall_back_to_ordinary_user_authentication() -> None: + from prisma.models import LiteLLM_RetiredAgentIdentity + + from litellm.repositories.table_repositories import RetiredAgentIdentityRepository + + identities: Final = AsyncMock() + identities.find_unique.return_value = None + retired: Final = AsyncMock() + retired.find_unique.return_value = LiteLLM_RetiredAgentIdentity( + binding_id="retired", + agent_id="agent-one", + provider="microsoft_entra", + issuer=ISSUER, + tenant_id=TENANT, + client_id=CLIENT, + ) + db: Final = SimpleNamespace( + db=SimpleNamespace( + litellm_agentidentity=identities, + litellm_retiredagentidentity=retired, + litellm_agentstable=AsyncMock(), + litellm_verifiedsubject=AsyncMock(), + ) + ) + store: Final = AgentIdentityStore( + AgentsRepository(db), + AgentIdentityRepository(db), + VerifiedSubjectRepository(db), + RetiredAgentIdentityRepository(db), + ) + result: Final = await store.resolve_verified_claims({**CLAIMS, "oid": HUMAN, "scp": "user_impersonation"}) + assert isinstance(result, AgentIdentityFailure) + assert result.code == "identity_denied" + assert "retired" in result.message + + +@pytest.mark.asyncio +async def test_missing_revision_cannot_create_entra_authentication_evidence() -> None: + store, _, identities, _ = setup_store() + result: Final = await store.record_authentication(ManagedAgentContext(agent_id="agent-one", mode="autonomous")) + assert isinstance(result, AgentIdentityFailure) + assert result.code == "identity_denied" + identities.update_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_authentication_evidence_write_failure_is_not_success() -> None: + store, _, identities, _ = setup_store() + identities.update_many.side_effect = RuntimeError("writer unavailable") + result: Final = await store.record_authentication( + ManagedAgentContext(agent_id="agent-one", binding_revision="revision-one", mode="autonomous") + ) + assert isinstance(result, AgentIdentityFailure) + assert result.code == "policy_unavailable" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("unavailable", [True, False]) +async def test_retired_binding_denies_and_history_outage_cannot_become_legacy_fallback(unavailable: bool) -> None: + + database: Final = MagicMock() + database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_retiredagentidentity.find_unique = AsyncMock( + return_value={"client_id": CLIENT}, side_effect=RuntimeError("unavailable") if unavailable else None + ) + result: Final = await AgentIdentityStore.from_client(database).resolve_verified_claims(CLAIMS) + assert isinstance(result, AgentIdentityFailure) + assert result.code == ("policy_unavailable" if unavailable else "identity_denied") + assert result.message == ( + "Retired agent identity could not be checked" if unavailable else "This agent identity binding has been retired" + ) + + +@pytest.mark.asyncio +async def test_new_binding_is_enforced_after_another_worker_commits_it() -> None: + _, agents, identities, humans = setup_store() + identities.find_unique.return_value = None + retired: Final = AsyncMock() + retired.find_unique.return_value = None + db: Final = SimpleNamespace( + writer_db=SimpleNamespace( + litellm_agentstable=agents, + litellm_agentidentity=identities, + litellm_verifiedsubject=humans, + litellm_retiredagentidentity=retired, + litellm_retiredagent=retired, + ) + ) + worker: Final = AgentIdentityStore.from_client(db) + claims: Final = {**CLAIMS, "oid": "55555555-5555-4555-8555-555555555555"} + assert await worker.resolve_verified_claims(claims) is None + identities.find_unique.return_value = BINDING + denied: Final = await worker.resolve_verified_claims(claims) + assert isinstance(denied, AgentIdentityFailure) + assert denied.code == "identity_denied" + assert "Application token contradicts" in denied.message + + +@pytest.mark.asyncio +async def test_non_string_subject_does_not_query_directory_ownership() -> None: + store, _, _, humans = setup_store() + assert await store.subject(ISSUER, TENANT, None) is None + humans.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("configured", [False, True]) +async def test_missing_or_unavailable_retirement_history_fails_closed(configured: bool) -> None: + database: Final = MagicMock() + database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("history unavailable")) + store: Final = AgentIdentityStore.from_client(database) if configured else setup_store()[0] + result: Final = await store.retired_agent("deleted-agent") + assert isinstance(result, AgentIdentityFailure) + assert result.code == "policy_unavailable" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("owner", ["canonical-user", "another-user"]) +async def test_interactive_enrollment_preserves_existing_subject_ownership(owner: str) -> None: + store, _, _, humans = setup_store() + humans.upsert.return_value = LiteLLM_VerifiedSubject( + subject_id="subject-one", + issuer=ISSUER, + tenant_id=TENANT, + oid=HUMAN, + user_id=owner, + kind="human", + verified_via="sso_interactive", + verified_at=datetime.now(timezone.utc), + ) + result: Final = await store.enroll_interactive_human( + MicrosoftInteractiveSubject(issuer=ISSUER, tenant_id=TENANT, oid=HUMAN), "canonical-user" + ) + if owner == "canonical-user": + assert result is None + else: + assert isinstance(result, AgentIdentityFailure) + assert result.code == "identity_denied" + assert humans.upsert.call_args.kwargs["data"]["update"] == {} + assert humans.upsert.call_args.kwargs["data"]["create"]["user_id"] == "canonical-user" + + +@pytest.mark.asyncio +async def test_interactive_enrollment_outage_fails_closed() -> None: + store, _, _, humans = setup_store() + humans.upsert.side_effect = ConnectionError("writer unavailable") + result: Final = await store.enroll_interactive_human( + MicrosoftInteractiveSubject(issuer=ISSUER, tenant_id=TENANT, oid=HUMAN), "canonical-user" + ) + assert isinstance(result, AgentIdentityFailure) + assert result.code == "policy_unavailable" + + +@pytest.mark.asyncio +async def test_matching_revision_records_successful_authentication() -> None: + store, _, identities, _ = setup_store() + assert ( + await store.record_authentication( + ManagedAgentContext(agent_id="agent-one", binding_revision="revision-one", mode="autonomous") + ) + is None + ) + identities.update_many.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outage", [False, True]) +async def test_resolver_maps_denials_and_outages_to_public_errors(outage: bool) -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentidentity.find_unique = AsyncMock( + return_value=BINDING, side_effect=ConnectionError("unavailable") if outage else None + ) + database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=stored_agent(enabled=False)) + with pytest.raises(HTTPException) as exc: + await resolve_managed_agent(CLAIMS, database) + assert exc.value.status_code == (503 if outage else 403) + + +@pytest.mark.asyncio +async def test_resolver_preserves_unconfigured_and_unrelated_authentication() -> None: + assert await resolve_managed_agent(CLAIMS, None) is None + assert await resolve_managed_agent({"sub": "ordinary-user"}, MagicMock()) is None + store, _, identities, _ = setup_store() + identities.find_unique.return_value = None + assert await store.resolve_verified_claims(CLAIMS) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("registered", [True, False]) +async def test_application_and_unregistered_clients_do_not_depend_on_human_subject_storage(registered: bool) -> None: + store, _, identities, humans = setup_store() + identities.find_unique.return_value = BINDING if registered else None + humans.find_unique.side_effect = RuntimeError("subject database unavailable") + result: Final = await store.resolve_verified_claims(CLAIMS) + if registered: + assert isinstance(result, ManagedAgentContext) + assert result.mode == "autonomous" + assert result.user_id is None + else: + assert result is None + humans.find_unique.assert_not_awaited() diff --git a/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py new file mode 100644 index 00000000000..45fe4b0655f --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_managed_identity.py @@ -0,0 +1,257 @@ +from typing import Final + +import pytest + +from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject, managed_write_fields +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import ( + AgentExecutionMode, + AgentIdentityBinding, + AgentIdentityFailure, + AgentSubject, +) + +TENANT: Final = "11111111-1111-4111-8111-111111111111" +CLIENT: Final = "22222222-2222-4222-8222-222222222222" +PRINCIPAL: Final = "33333333-3333-4333-8333-333333333333" +HUMAN: Final = "44444444-4444-4444-8444-444444444444" +ISSUER: Final = f"https://login.microsoftonline.com/{TENANT}/v2.0" +BINDING: Final = AgentIdentityBinding( + agent_id="agent-one", + provider="microsoft_entra", + tenant_id=TENANT, + client_id=CLIENT, + service_principal_id=PRINCIPAL, + issuer=ISSUER, + required_roles=("Agent.Invoke",), + required_scopes=("user_impersonation",), + revision="binding-one", +) + + +def claims(**overrides: object) -> dict[str, object]: + return {"iss": ISSUER, "tid": TENANT, "azp": CLIENT, "oid": PRINCIPAL, "roles": ["Agent.Invoke"], **overrides} + + +def test_autonomous_identity_needs_no_human_and_checks_the_pinned_principal() -> None: + result: Final = classify_agent_subject(BINDING, claims(), "autonomous") + assert result == AgentSubject(kind="application", oid=PRINCIPAL, mode="autonomous") + assert isinstance(classify_agent_subject(BINDING, claims(oid=HUMAN), "autonomous"), AgentIdentityFailure) + + +@pytest.mark.parametrize( + "overrides", + [ + {"iss": "https://untrusted.example"}, + {"tid": CLIENT}, + {"azp": TENANT}, + {"roles": []}, + {"idtyp": "user"}, + {"scp": "user_impersonation"}, + {"scp": 1}, + {"oid": None}, + ], +) +def test_application_rejects_mismatched_or_contradictory_verified_claims(overrides: dict[str, object]) -> None: + assert isinstance(classify_agent_subject(BINDING, claims(**overrides), "both"), AgentIdentityFailure) + + +def test_delegated_profile_identifies_a_subject_without_asserting_that_it_is_human() -> None: + result: Final = classify_agent_subject(BINDING, claims(oid=HUMAN, scp="user_impersonation"), "delegated") + assert result == AgentSubject(kind="delegated_subject", oid=HUMAN, mode="delegated") + + +@pytest.mark.parametrize( + "overrides", + [ + {"scp": "unrelated"}, + {"scp": ""}, + {"idtyp": "app"}, + {"xms_sub_fct": "2 13 15"}, + {"xms_sub_fct": [13]}, + ], +) +def test_delegated_profile_rejects_unknown_scope_and_known_nonhuman_subjects(overrides: dict[str, object]) -> None: + assert isinstance( + classify_agent_subject(BINDING, claims(**{"oid": HUMAN, "scp": "user_impersonation", **overrides}), "both"), + AgentIdentityFailure, + ) + + +def test_allowed_mode_cannot_be_selected_by_the_caller() -> None: + assert isinstance(classify_agent_subject(BINDING, claims(), "delegated"), AgentIdentityFailure) + assert isinstance( + classify_agent_subject(BINDING, claims(oid=HUMAN, scp="user_impersonation"), "autonomous"), + AgentIdentityFailure, + ) + + +def test_native_facet_absence_does_not_establish_human_identity() -> None: + result: Final = classify_agent_subject( + BINDING, claims(oid=HUMAN, scp="user_impersonation", xms_sub_fct="113"), "both" + ) + assert isinstance(result, AgentSubject) + assert result.kind == "delegated_subject" + + +def managed_agent() -> AgentResponse: + return AgentResponse( + agent_id="agent-one", agent_name="Research", agent_card_params={}, identity=BINDING, identity_managed=True + ) + + +def test_unbinding_keeps_managed_state_and_disables_agent() -> None: + result: Final = managed_write_fields({"identity": None, "enabled": True}, managed_agent(), "admin") + assert not isinstance(result, AgentIdentityFailure) + assert result["identity_managed"] is True + assert result["enabled"] is False + assert result["identity"]["update"]["active"] is False + assert result["identity"]["update"]["last_authenticated_at"] is None + assert result["identity"]["update"]["revision"] != BINDING.revision + + +def test_rename_does_not_rewrite_binding_or_evidence() -> None: + assert managed_write_fields({"agent_name": "Renamed"}, managed_agent(), "admin") == {} + + +def test_autonomous_binding_requires_enterprise_application_object_id() -> None: + result: Final = managed_write_fields( + {"identity": {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}}, None, "admin" + ) + assert isinstance(result, AgentIdentityFailure) + assert "service-principal" in result.message + + +def test_rebinding_clears_evidence_and_uses_atomic_nested_write() -> None: + result: Final = managed_write_fields( + { + "identity": { + "provider": "microsoft_entra", + "tenant_id": TENANT, + "client_id": CLIENT, + "service_principal_id": PRINCIPAL, + } + }, + managed_agent(), + "admin", + ) + assert not isinstance(result, AgentIdentityFailure) + assert result["identity_managed"] is True + assert "upsert" in result["identity"] + assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision + assert result["identity"]["upsert"]["update"]["last_authenticated_at"] is None + + +def test_unbound_identity_can_be_reactivated_with_the_same_application() -> None: + disabled: Final = managed_agent().model_copy( + update={"identity": BINDING.model_copy(update={"active": False}), "enabled": False} + ) + configuration: Final = BINDING.model_dump( + exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"} + ) + result: Final = managed_write_fields({"identity": configuration, "enabled": True}, disabled, "admin") + assert not isinstance(result, AgentIdentityFailure) + assert result["enabled"] is True + assert result["identity"]["upsert"]["update"]["active"] is True + assert result["identity"]["upsert"]["update"]["revision"] != BINDING.revision + + +def test_each_application_binding_records_its_history_atomically() -> None: + configuration: Final = BINDING.model_dump( + exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"} + ) + created: Final = managed_write_fields({"identity": configuration}, None, "admin") + assert not isinstance(created, AgentIdentityFailure) + assert created["retired_identities"]["connectOrCreate"]["create"]["client_id"] == CLIENT + replacement: Final = managed_write_fields( + {"identity": {**configuration, "client_id": HUMAN}}, managed_agent(), "admin" + ) + assert not isinstance(replacement, AgentIdentityFailure) + assert replacement["retired_identities"]["connectOrCreate"]["create"]["client_id"] == HUMAN + + +def test_unchanged_binding_preserves_revision_and_authentication_evidence() -> None: + configuration: Final = BINDING.model_dump( + exclude={"agent_id", "issuer", "revision", "last_authenticated_at", "active"} + ) + assert managed_write_fields({"identity": configuration}, managed_agent(), "admin") == {} + + +@pytest.mark.parametrize("identity", [None, BINDING.model_copy(update={"active": False})]) +def test_enabling_unbound_or_inactive_identity_requires_rebinding(identity: AgentIdentityBinding | None) -> None: + agent: Final = managed_agent().model_copy(update={"identity": identity, "enabled": False}) + result: Final = managed_write_fields({"enabled": True}, agent, "admin") + assert isinstance(result, AgentIdentityFailure) + assert "Bind an identity" in result.message + + +@pytest.mark.parametrize("mode", ["delegated", "both"]) +def test_explicit_empty_scope_requirements_can_be_registered_and_preserved(mode: str) -> None: + from litellm.types.proxy.agent_identity import EntraIdentityConfig + + configuration: Final = EntraIdentityConfig( + provider="microsoft_entra", + tenant_id=TENANT, + client_id=CLIENT, + service_principal_id=PRINCIPAL, + required_scopes=(), + ) + created: Final = managed_write_fields( + {"identity": configuration.model_dump(), "execution_mode": mode}, None, "admin" + ) + assert not isinstance(created, AgentIdentityFailure) + assert created["identity"]["create"]["required_scopes"] == () + agent: Final = managed_agent().model_copy(update={"identity": BINDING.model_copy(update={"required_scopes": ()})}) + updated: Final = managed_write_fields({"execution_mode": mode}, agent, "admin") + assert not isinstance(updated, AgentIdentityFailure) + assert updated["execution_mode"] == mode + + +@pytest.mark.parametrize( + "incoming", + [ + {"identity": {"provider": "microsoft_entra", "tenant_id": "invalid", "client_id": CLIENT}}, + {"execution_mode": "unknown"}, + ], +) +def test_invalid_identity_configuration_returns_a_public_validation_failure(incoming: dict[str, object]) -> None: + result: Final = managed_write_fields(incoming, None, "admin") + assert isinstance(result, AgentIdentityFailure) + assert result.code == "identity_denied" + assert result.message.startswith("Invalid agent identity configuration:") + + +@pytest.mark.parametrize("roles", ["Agent.Invoke", [42], None]) +def test_malformed_application_roles_are_rejected(roles: object) -> None: + result: Final = classify_agent_subject(BINDING, claims(roles=roles), "autonomous") + assert isinstance(result, AgentIdentityFailure) + assert "Invalid application roles" in result.message + + +def test_entra_binding_normalizes_identifiers_and_rejects_invalid_configuration() -> None: + from pydantic import ValidationError + + from litellm.types.proxy.agent_identity import EntraIdentityConfig + + identifier = "ABCDEF00-1234-4234-9234-123456789ABC" + config = EntraIdentityConfig(provider="microsoft_entra", tenant_id=identifier, client_id=identifier) + assert config.tenant_id == identifier.lower() + assert config.client_id == identifier.lower() + assert config.service_principal_id is None + assert config.issuer == f"https://login.microsoftonline.com/{config.tenant_id}/v2.0" + with pytest.raises(ValidationError): + EntraIdentityConfig(provider="microsoft_entra", tenant_id="invalid", client_id=identifier) + + +@pytest.mark.parametrize("mode", ["delegated", "both"]) +def test_empty_required_scopes_allow_valid_delegated_scope(mode: AgentExecutionMode) -> None: + binding: Final = BINDING.model_copy(update={"required_scopes": ()}) + result: Final = classify_agent_subject(binding, claims(oid=HUMAN, scp="custom_scope"), mode) + assert result == AgentSubject(kind="delegated_subject", oid=HUMAN, mode="delegated") + + +@pytest.mark.parametrize("scope", [None, "", " \t ", 42]) +def test_empty_requirements_do_not_make_a_scope_less_human_token_valid(scope: object) -> None: + binding: Final = BINDING.model_copy(update={"required_scopes": ()}) + result: Final = classify_agent_subject(binding, claims(oid=HUMAN, scp=scope), "both") + assert isinstance(result, AgentIdentityFailure) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index f014e9c26d1..88833f38270 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1091,7 +1091,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch): monkeypatch.setitem(auth_checks.last_db_access_time, f"user_id:{user_id}", (None, time.time())) db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user") mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row) + mock_prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=db_row) result = await get_user_object( user_id=user_id, @@ -1103,7 +1103,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch): assert result is not None assert result.user_id == user_id - mock_prisma_client.db.litellm_usertable.find_unique.assert_awaited_once() + mock_prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_once() @pytest.mark.asyncio @@ -3058,7 +3058,7 @@ async def test_get_team_object_raises_404_when_not_found(): mock_prisma_client = MagicMock() mock_db = AsyncMock() mock_prisma_client.db = mock_db - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=None) mock_cache = MagicMock() mock_cache.async_get_cache = AsyncMock(return_value=None) @@ -3076,11 +3076,40 @@ async def test_get_team_object_raises_404_when_not_found(): assert "Team doesn't exist in db" in str(exc_info.value.detail) +@pytest.mark.asyncio +async def test_get_team_object_check_db_only_reads_writer_through_the_shared_loader(): + """Management endpoints mock ``_get_team_object_from_user_api_key_cache`` and expect + ``check_db_only`` to still flow through it; only the table it reads moves to the writer.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.auth import auth_checks + from litellm.proxy.auth.auth_checks import get_team_object + + row = {"team_id": "team-writer", "models": ["gpt-4o"], "object_permission_id": None} + prisma = MagicMock() + prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row)) + prisma.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row)) + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=None) + cache.async_set_cache = AsyncMock() + shared_loader = AsyncMock(wraps=auth_checks._get_team_object_from_user_api_key_cache) + + with patch.object(auth_checks, "_get_team_object_from_user_api_key_cache", shared_loader): + team = await get_team_object("team-writer", prisma, cache, check_db_only=True) + + assert team.team_id == "team-writer" + assert shared_loader.await_args.kwargs["use_writer"] is True + prisma.writer_db.litellm_teamtable.find_unique.assert_awaited_once() + prisma.db.litellm_teamtable.find_unique.assert_not_awaited() + cache.async_set_cache.assert_awaited_once() + + def _mock_prisma_for_team_lookup(find_unique): from unittest.mock import MagicMock mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_teamtable.find_unique = find_unique + mock_prisma_client.writer_db.litellm_teamtable.find_unique = find_unique return mock_prisma_client @@ -9979,3 +10008,130 @@ def test_can_object_call_model_allows_listed_model_for_key(): ) assert result is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("allowed", [True, False]) +async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache(allowed: bool) -> None: + from litellm.proxy._types import LiteLLM_AccessGroupTable + from litellm.proxy.auth.auth_checks import get_access_object + + stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"]) + current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []}) + client: Final = MagicMock() + client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current) + client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=stale) + cache: Final = MagicMock() + cache.async_get_cache = AsyncMock(return_value=stale) + cache.async_set_cache = AsyncMock() + result: Final = await get_access_object("group", client, cache, check_db_only=True) + assert result.access_model_names == (["new"] if allowed else []) + cache.async_get_cache.assert_not_awaited() + client.db.litellm_accessgrouptable.find_unique.assert_not_awaited() + client.writer_db.litellm_accessgrouptable.find_unique.assert_awaited_once_with(where={"access_group_id": "group"}) + + +@pytest.mark.asyncio +async def test_authoritative_access_group_outage_does_not_use_cached_grants() -> None: + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import get_access_object + + client: Final = MagicMock() + client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable")) + cache: Final = MagicMock() + cache.async_get_cache = AsyncMock() + with pytest.raises(HTTPException) as failure: + await get_access_object("group", client, cache, check_db_only=True) + assert failure.value.status_code == 404 + cache.async_get_cache.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_authoritative_team_permission_outage_cannot_drop_the_teams_restrictions() -> None: + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import get_team_object + + row: Final = LiteLLM_TeamTable(team_id="team-policy-outage", object_permission_id="team-permission") + client: Final = MagicMock() + client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=row) + client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(side_effect=RuntimeError("unavailable")) + cache: Final = MagicMock() + cache.async_get_cache = AsyncMock() + cache.async_set_cache = AsyncMock() + with pytest.raises(HTTPException) as failure: + await get_team_object(row.team_id, client, cache, check_db_only=True) + assert failure.value.status_code == 404 + client.writer_db.litellm_objectpermissiontable.find_unique.assert_awaited_once() + cache.async_set_cache.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("strict", [True, False]) +@pytest.mark.parametrize("missing", [True, False]) +async def test_referenced_permission_failures_preserve_legacy_behavior_and_deny_strict_reads(strict, missing): + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import get_object_permission + + client = MagicMock() + lookup = AsyncMock(return_value=None, side_effect=None if missing else RuntimeError("unavailable")) + client.writer_db.litellm_objectpermissiontable.find_unique = lookup + client.db.litellm_objectpermissiontable.find_unique = lookup + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=None) + if strict: + with pytest.raises(HTTPException if missing else RuntimeError): + await get_object_permission("referenced", client, cache, check_db_only=True) + cache.async_get_cache.assert_not_awaited() + else: + assert await get_object_permission("referenced", client, cache) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "models,key_aliases,team_aliases,allowed", + [ + (["fast"], {}, {}, True), + ([], {}, {}, False), + (["other"], {}, {}, False), + (["target"], {"fast": "target"}, {}, True), + (["target"], {}, {"fast": "target"}, True), + (["fast"], {}, {"fast": "forbidden"}, False), + ], +) +async def test_managed_agent_model_policy_checks_dispatched_model( + models: list[str], key_aliases: dict[str, str], team_aliases: dict[str, str], allowed: bool +) -> None: + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import common_checks + from litellm.types.agents import AgentResponse + + agent: Final = AgentResponse( + agent_id="managed", agent_name="Managed", agent_card_params={}, object_permission={"models": models} + ) + auth: Final = UserAPIKeyAuth( + token="test-token", team_id="team", aliases=key_aliases, team_model_aliases=team_aliases + ) + auth.managed_agent_policy = agent + checks: Final = common_checks( + request_body={"model": "fast", "messages": [{"role": "user", "content": "hi"}]}, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/chat/completions", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=auth, + request=MagicMock(spec=Request), + ) + if allowed: + assert await checks is True + else: + with pytest.raises((HTTPException, ModelAccessDeniedProxyException)) as failure: + await checks + assert str(getattr(failure.value, "status_code", getattr(failure.value, "code", None))) == "403" diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index b1622e0dff0..0280a570292 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -2,15 +2,14 @@ import asyncio import re import time from collections.abc import Mapping, Sequence -from typing import Final, Optional +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch -from fastapi import HTTPException import httpx import pytest +from fastapi import HTTPException -import litellm - +from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( DEFAULT_JWKS_STALE_TTL, JWTLiteLLMRoleMap, @@ -26,7 +25,6 @@ from litellm.proxy._types import ( RoleBasedPermissions, ScopeMapping, ) -from litellm.caching.dual_cache import DualCache from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry from litellm.proxy.auth.auth_checks import TeamNotFoundError from litellm.proxy.auth.handle_jwt import ( @@ -1637,7 +1635,6 @@ async def test_auth_builder_returns_team_membership_object(): @pytest.mark.asyncio async def test_auth_builder_with_oidc_userinfo_enabled(): """Test that auth_builder uses OIDC UserInfo endpoint when enabled""" - from unittest.mock import MagicMock from litellm.caching import DualCache from litellm.proxy.utils import ProxyLogging @@ -1648,9 +1645,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo enabled jwt_handler = JWTHandler() @@ -1677,18 +1672,12 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1696,9 +1685,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1711,9 +1698,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1726,15 +1711,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up mock return values mock_get_userinfo.return_value = userinfo_response @@ -1764,7 +1743,6 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): @pytest.mark.asyncio async def test_auth_builder_with_oidc_userinfo_disabled(): """Test that auth_builder uses JWT validation when OIDC UserInfo is disabled""" - from unittest.mock import MagicMock from litellm.caching import DualCache from litellm.proxy.utils import ProxyLogging @@ -1775,9 +1753,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo disabled jwt_handler = JWTHandler() @@ -1801,18 +1777,12 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1820,9 +1790,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): return_value=("test_user_1", None, None), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1835,9 +1803,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1850,15 +1816,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up mock return values mock_auth_jwt.return_value = jwt_response @@ -2631,7 +2591,6 @@ async def test_find_and_validate_specific_team_id_with_team_alias(): """ Test that find_and_validate_specific_team_id resolves team by name when team_id is not found """ - from unittest.mock import MagicMock from litellm.caching import DualCache from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable @@ -2654,9 +2613,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias(): # Mock team object returned by get_team_object_by_alias team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team") - with patch( - "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = team_object team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id( @@ -2685,7 +2642,6 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): """ Test that team_id_jwt_field takes precedence over team_alias_jwt_field """ - from unittest.mock import MagicMock from litellm.caching import DualCache from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable @@ -2699,9 +2655,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): jwt_handler.update_environment( prisma_client=None, user_api_key_cache=user_api_key_cache, - litellm_jwtauth=LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_alias" - ), + litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"), ) # Token with both team_id and team name @@ -2711,9 +2665,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): team_object = LiteLLM_TeamTable(team_id="direct-team-id") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -2890,7 +2842,6 @@ async def test_get_objects_resolves_org_by_name(): @pytest.mark.asyncio async def test_resolve_jwks_url_passthrough_for_direct_jwks_url(): """Non-discovery URLs are returned unchanged.""" - from unittest.mock import AsyncMock, MagicMock from litellm.caching.dual_cache import DualCache @@ -3143,7 +3094,7 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field(): When team_id_jwt_field is a normal field name (no dot-notation) the error message should not contain a spurious bracket-notation hint. """ - from unittest.mock import AsyncMock, MagicMock + from unittest.mock import MagicMock from litellm.caching.dual_cache import DualCache @@ -3230,8 +3181,8 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field(): async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( user_id: str, user_teams: list, - get_team_object_return: Optional[str], - expected_team_id: Optional[str], + get_team_object_return: str | None, + expected_team_id: str | None, expect_get_team_called: bool, expect_get_membership_called: bool, ) -> None: @@ -3244,9 +3195,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( if len(user_teams) == 1 and get_team_object_return == "resolved_row": only = user_teams[0] team_table = LiteLLM_TeamTable(team_id=only) - membership = LiteLLM_TeamMembership( - user_id=user_id, team_id=only, litellm_budget_table=None - ) + membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None) get_team_return_value = team_table membership_return_value = membership else: @@ -3305,9 +3254,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3324,9 +3271,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( code = 404 if get_team_object_return == "http_404" else 500 mock_get_team.side_effect = HTTPException( status_code=code, - detail={ - "error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call." - }, + detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."}, ) else: mock_get_team.return_value = get_team_return_value @@ -4047,7 +3992,7 @@ def _encode_rsa_jwt( issuer: str, audience: str, kid: str, - extra_claims: Optional[dict] = None, + extra_claims: dict | None = None, ) -> str: import time @@ -4743,12 +4688,9 @@ async def test_get_objects_team_membership_uses_rebound_user_id(): async def fake_get_team_membership(user_id, team_id, *args, **kwargs): captured["user_id"] = user_id captured["team_id"] = team_id - return None jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_id_jwt_field="email", user_id_upsert=True - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True) with ( patch( @@ -5389,7 +5331,7 @@ async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback( assert team_object is None -def _db_fallback_handler(litellm_jwtauth: Optional[LiteLLM_JWTAuth] = None) -> JWTHandler: +def _db_fallback_handler(litellm_jwtauth: LiteLLM_JWTAuth | None = None) -> JWTHandler: handler = JWTHandler() handler.litellm_jwtauth = litellm_jwtauth or LiteLLM_JWTAuth() return handler @@ -5447,9 +5389,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership(): "expect_403", ), [ - pytest.param( - True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team" - ), + pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"), pytest.param( True, ["team_a", "team_b"], @@ -5497,8 +5437,8 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership(): async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( fallback_to_db_teams: bool, user_teams: list, - header_team_id: Optional[str], - expected_team_id: Optional[str], + header_team_id: str | None, + expected_team_id: str | None, expect_403: bool, ) -> None: """End-to-end auth_builder behavior with no JWT team claims. @@ -5527,9 +5467,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( async def call_auth_builder(): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5569,9 +5507,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6765,7 +6701,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted(): team_id_upsert=True, ) - upsert_by_team: dict[str, Optional[bool]] = {} + upsert_by_team: dict[str, bool | None] = {} async def spy_get_team(team_id, **kwargs): upsert_by_team[team_id] = kwargs.get("team_id_upsert") @@ -6800,9 +6736,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -7806,6 +7740,58 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc assert result["team_id"] is None +def _explicit_identity_registry() -> AgentRegistry: + registry: Final = AgentRegistry() + registry.register_agent(AgentResponse( + agent_id="explicit-agent-id", + agent_name="Readable agent name", + agent_card_params={}, + litellm_params={"identity": { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + }}, + )) + return registry + + +@pytest.mark.parametrize("claim_field", ["azp", None]) +def test_runtime_json_cannot_establish_a_managed_identity(claim_field: str | None) -> None: + registry: Final = _explicit_identity_registry() + handler: Final = _entra_agent_jwt_handler(claim_field) + claims: Final = { + "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0", + "tid": "11111111-1111-4111-8111-111111111111", + "azp": "22222222-2222-4222-8222-222222222222", + } + if claim_field is None: + assert JWTAuthManager.resolve_agent_id(handler, claims, registry) is None + else: + with pytest.raises(HTTPException) as failure: + JWTAuthManager.resolve_agent_id(handler, claims, registry) + assert failure.value.status_code == 403 + + +@pytest.mark.parametrize("override", [ + {"iss": "https://attacker.example"}, + {"tid": "33333333-3333-4333-8333-333333333333"}, + {"azp": "33333333-3333-4333-8333-333333333333"}, + {"azp": "explicit-agent-id"}, + {"azp": "Readable agent name"}, +]) +def test_explicit_entra_identity_cannot_be_claimed_via_legacy_lookup(override: Mapping[str, object]) -> None: + registry: Final = _explicit_identity_registry() + handler: Final = _entra_agent_jwt_handler("azp") + with pytest.raises(HTTPException) as failure: + JWTAuthManager.resolve_agent_id(handler, { + "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0", + "tid": "11111111-1111-4111-8111-111111111111", + "azp": "22222222-2222-4222-8222-222222222222", + **override, + }, registry) + assert failure.value.status_code == 403 + + @pytest.mark.asyncio @pytest.mark.parametrize("existing_user", [False, True]) @pytest.mark.parametrize("warm_cache", [False, True]) @@ -7853,3 +7839,216 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning users.create.assert_not_awaited() if existing_user: assert users.find_unique.await_count == (0 if warm_cache else 1) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"]) +@pytest.mark.parametrize("audience_validation", (True, False)) +@pytest.mark.parametrize( + "route,allowed", + [ + ("/chat/completions", True), ("/v1/messages", True), ("/v1/responses", True), + ("/mcp-rest/tools/call", True), ("/a2a/target", True), + ("/v1/files", False), ("/v1/batches", False), ("/v1/vector_stores", False), + ("/v1/containers", False), ("/openai/v1/files", False), + ("/v1/responses/other-response", False), ("/v1/realtime/client_secrets", False), + ], +) +async def test_managed_application_uses_persisted_identity_without_provisioning_human( + monkeypatch: pytest.MonkeyPatch, mode: str, audience_validation: bool, route: str, allowed: bool +) -> None: + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + tenant: Final = "11111111-1111-4111-8111-111111111111" + client_id: Final = "22222222-2222-4222-8222-222222222222" + principal: Final = "33333333-3333-4333-8333-333333333333" + issuer: Final = f"https://login.microsoftonline.com/{tenant}/v2.0" + jwks_url: Final = "https://login.microsoftonline.test/managed-keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + monkeypatch.setenv("JWT_ISSUER", issuer) + monkeypatch.setenv("JWT_AUDIENCE", "api://gateway") + private_key, jwk = _get_rsa_key_and_jwk(kid="managed-key") + cache: Final = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + handler: Final = JWTHandler() + handler.update_environment(None, cache, LiteLLM_JWTAuth(user_id_upsert=True)) + binding: Final = AgentIdentityBinding( + agent_id="stable-id", + provider="microsoft_entra", + issuer=issuer, + tenant_id=tenant, + client_id=client_id, + service_principal_id=principal, + revision="revision-one", + required_roles=("Agent.Invoke",), + ) + agent: Final = AgentResponse.model_validate( + { + "agent_id": "stable-id", + "agent_name": "A readable name", + "agent_card_params": {}, + "identity": binding, + "identity_managed": True, + "execution_mode": mode, + } + ) + database: Final = MagicMock() + database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding) + database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1) + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent) + database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None) + database.db.litellm_usertable.upsert = AsyncMock() + token: Final = _encode_rsa_jwt( + private_key, + issuer=issuer, + audience="api://gateway", + kid="managed-key", + extra_claims={ + "tid": tenant, + "azp": client_id, + "oid": principal, + "roles": ["Agent.Invoke"], + "idtyp": "app", + }, + ) + arguments: Final = dict( + api_key=token, + jwt_handler=handler, + request_data={}, + general_settings={}, + route=route, + prisma_client=database, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + if not audience_validation: + monkeypatch.delenv("JWT_AUDIENCE") + if mode == "delegated" or not audience_validation or not allowed: + with pytest.raises(HTTPException) as failure: + await JWTAuthManager.auth_builder(**arguments) + assert failure.value.status_code == 403 + else: + result: Final = await JWTAuthManager.auth_builder(**arguments) + auth: Final = JWTAuthManager.user_api_key_auth_from_result(result) + assert auth.agent_id == "stable-id" + assert auth.user_id is None + assert auth.team_id is None + assert auth.managed_agent_context is not None + assert auth.managed_agent_context.mode == "autonomous" + assert result["is_proxy_admin"] is False + database.db.litellm_usertable.upsert.assert_not_awaited() + + +@pytest.mark.parametrize("claim_value", ["managed", "Readable managed agent"]) +def test_legacy_claim_cannot_select_a_top_level_entra_binding(claim_value: str) -> None: + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + registry: Final = AgentRegistry() + registry.register_agent( + AgentResponse( + agent_id="managed", + agent_name="Readable managed agent", + agent_card_params={}, + identity_managed=True, + identity=AgentIdentityBinding( + agent_id="managed", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + issuer="issuer", + revision="revision", + ), + ) + ) + with pytest.raises(HTTPException) as denied: + JWTAuthManager.resolve_agent_id(_entra_agent_jwt_handler("agent"), {"agent": claim_value}, registry) + assert denied.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["human", "config-agent", "managed-agent"]) +async def test_database_free_jwt_admission_with_entra_shaped_claims(monkeypatch: pytest.MonkeyPatch, kind: str) -> None: + issuer: Final = "https://login.microsoftonline.com/test-tenant/v2.0" + jwks_url: Final = "https://login.microsoftonline.test/config-only-keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + monkeypatch.setenv("JWT_ISSUER", issuer) + monkeypatch.setenv("JWT_AUDIENCE", "api://gateway") + private_key, jwk = _get_rsa_key_and_jwk(kid="config-key") + cache: Final = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + registry: Final = AgentRegistry() + registry.register_agent( + AgentResponse( + agent_id="configured", + agent_name="Configured", + agent_card_params={}, + identity_managed=kind == "managed-agent", + ) + ) + handler: Final = JWTHandler() + handler.update_environment(None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"])) + handler.bind_agent_lookup(registry) + token: Final = _encode_rsa_jwt( + private_key, + issuer=issuer, + audience="api://gateway", + kid="config-key", + extra_claims={ + "tid": "test-tenant", + "azp": "application", + "scope": "litellm_proxy_admin", + **({"agent": "configured"} if kind != "human" else {}), + }, + ) + arguments: Final = dict( + api_key=token, + jwt_handler=handler, + request_data={}, + general_settings={}, + route="/chat/completions", + prisma_client=None, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + if kind == "managed-agent": + with pytest.raises(HTTPException) as denied: + await JWTAuthManager.auth_builder(**arguments) + assert denied.value.status_code == 403 + else: + result: Final = await JWTAuthManager.auth_builder(**arguments) + auth: Final = JWTAuthManager.user_api_key_auth_from_result(result) + assert auth.agent_id == ("configured" if kind == "config-agent" else None) + assert auth.managed_agent_context is None + assert result["is_proxy_admin"] is True + + +@pytest.mark.parametrize( + "issuer,audience,disabled,expected", + [ + (None, "gateway", False, False), + ("trusted", "gateway", False, True), + ("trusted", None, True, False), + ("other", "gateway", False, False), + ], +) +def test_managed_issuer_requires_configured_audience_validation( + monkeypatch: pytest.MonkeyPatch, issuer: str | None, audience: str | None, disabled: bool, expected: bool +) -> None: + from litellm.proxy._types import JWTIssuerConfig + + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + handler: Final = JWTHandler() + handler.update_environment( + None, + DualCache(), + LiteLLM_JWTAuth( + issuers=[ + JWTIssuerConfig(issuer="trusted", audience=audience, disable_audience_validation=disabled), + ] + ), + ) + assert handler.managed_issuer_is_trusted(issuer) is expected diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 470db99108a..268b70c9596 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -9278,6 +9278,8 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state(): ) async def auth_that_reserves(request, api_key): + assert request.method == "GET" + assert request.query_params.get("model") == "gpt-realtime" request.state.budget_reservation = reservation return UserAPIKeyAuth(token="hashed", budget_reservation=reservation) @@ -9290,3 +9292,201 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state(): assert result.budget_reservation == reservation assert websocket.state.budget_reservation is reservation assert websocket.scope["state"]["budget_reservation"] is reservation + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invoke", [False, True]) +async def test_centralized_authorization_preserves_database_free_config_agents(monkeypatch, invoke: bool): + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": None, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + registry = AgentRegistry() + registry.load_agents_from_config( + [{"agent_name": "config-agent", "agent_card_params": {"name": "Config", "url": "http://localhost:9999"}}] + ) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + registered = registry.get_agent_by_name("config-agent") + model = "a2a/config-agent" if invoke else "test-model" + auth = UserAPIKeyAuth(agent_id=registered.agent_id, jwt_claims={"agent": "config-agent"}, models=[model]) + data = {"model": model, "messages": [{"role": "user", "content": "hi"}]} + assert ( + await _authorize_authenticated_request( + auth, _alias_request("/v1/chat/completions", data), data, "/v1/chat/completions", "jwt-token" + ) + is None + ) + assert auth.managed_agent_policy is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("verified_identity", [False, True]) +async def test_managed_actor_cannot_access_provider_resource_routes(monkeypatch, verified_identity: bool): + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + policy = AgentResponse( + agent_id="managed", + agent_name="Managed", + agent_card_params={}, + identity_managed=True, + identity=AgentIdentityBinding( + agent_id="managed", + provider="microsoft_entra", + tenant_id="tenant", + client_id="application", + service_principal_id="principal", + issuer="issuer", + revision="revision", + ), + object_permission={"models": ["test-model"]}, + ) + database = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + request = _alias_request("/v1/files", {}) + request.scope["method"] = "GET" + from litellm.types.proxy.agent_identity import ManagedAgentContext + + auth = UserAPIKeyAuth( + agent_id="managed", api_key="persisted-key", models=["test-model"], + managed_agent_context=( + ManagedAgentContext(agent_id="managed", binding_revision="revision", mode="autonomous") + if verified_identity else None + ), + ) + with pytest.raises(ProxyException) as denied: + await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key") + assert denied.value.code == "403" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("requested", [None, "test-model"]) +@pytest.mark.parametrize("grant_default", [False, True]) +@pytest.mark.parametrize( + "route,settings,cli_model", + [ + ("/v1/chat/completions", {"completion_model": "forbidden-model"}, None), + ("/v1/responses", {"completion_model": "forbidden-model"}, None), + ("/v1/messages", {"completion_model": "forbidden-model"}, None), + ("/v1/moderations", {"moderation_model": "forbidden-model"}, None), + ("/v1/audio/transcriptions", {"moderation_model": "forbidden-model"}, None), + ("/v1/audio/speech", {}, "forbidden-model"), + ("/v1/chat/completions", {}, "forbidden-model"), + ("/v1/images/generations", {"image_generation_model": "forbidden-model"}, None), + ("/v1/images/edits", {"image_generation_model": "forbidden-model"}, None), + ], +) +async def test_managed_agent_cannot_bypass_grants_with_server_default( + monkeypatch, requested, route, settings, cli_model, grant_default +): + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext + + policy = AgentResponse( + agent_id="managed", + agent_name="Managed", + agent_card_params={}, + identity_managed=True, + identity=AgentIdentityBinding( + agent_id="managed", + provider="microsoft_entra", + tenant_id="tenant", + client_id="application", + service_principal_id="principal", + issuer="issuer", + revision="revision", + ), + object_permission={"models": ["test-model", "forbidden-model"] if grant_default else ["test-model"]}, + ) + database = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "prisma_client": database, + "general_settings": settings, + "user_model": cli_model, + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + data = {"messages": [{"role": "user", "content": "hi"}], **({"model": requested} if requested else {})} + auth = UserAPIKeyAuth(agent_id="managed") + auth.managed_agent_context = ManagedAgentContext( + agent_id="managed", binding_revision="revision", mode="autonomous" + ) + if not grant_default: + with pytest.raises(ProxyException) as denied: + await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key") + assert denied.value.code == "403" + assert "forbidden-model" in denied.value.message + return + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new_callable=AsyncMock, + ) as reserve: + reserve.return_value = None + assert ( + await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key") + is None + ) + reserve.assert_awaited_once() + assert reserve.call_args.kwargs["request_body"]["model"] == "forbidden-model" + + +@pytest.mark.asyncio +async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeypatch: pytest.MonkeyPatch) -> None: + from typing import Final + + from litellm.proxy import proxy_server + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + binding: Final = AgentIdentityBinding( + agent_id="managed", provider="microsoft_entra", issuer="issuer", tenant_id="tenant", + client_id="client", service_principal_id="principal", revision="current", + ) + agent: Final = AgentResponse( + agent_id="managed", agent_name="Managed", agent_card_params={}, + identity_managed=True, identity=binding, execution_mode="autonomous", + ) + client: Final = MagicMock() + client.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding) + client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent) + handler: Final = MagicMock() + handler.is_jwt.return_value = True + handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub") + handler.auth_jwt = AsyncMock(return_value={ + "iss": "issuer", "tid": "tenant", "azp": "client", "oid": "principal", "sub": "mapped-key", + }) + for name, value in { + **_proxy_attrs_for_centralized_checks(), + "general_settings": {"enable_jwt_auth": True}, "premium_user": True, + "prisma_client": client, "jwt_handler": handler, "user_api_key_cache": DualCache(), + "proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)), + }.items(): + monkeypatch.setattr(proxy_server, name, value) + with pytest.raises(ProxyException) as failure: + await _user_api_key_auth_builder( + request=_alias_request("/v1/chat/completions", {}), api_key="Bearer verified.jwt.token", + azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, + azure_apim_header=None, request_data={}, + ) + assert failure.value.code == "403" + assert "without virtual-key mapping" in failure.value.message + client.writer_db.litellm_agentstable.find_unique.assert_awaited_once() diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index 7929a0b21af..6087aa1f570 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -1,6 +1,7 @@ import io import json -from typing import get_type_hints +from collections.abc import Mapping +from typing import Literal, get_type_hints from unittest.mock import AsyncMock, MagicMock, patch import orjson @@ -1210,3 +1211,42 @@ class TestCoerceNumericFormFields: numeric_fields=self.numeric_fields, ) assert result == {"n": 3, "temperature": None, "image": buffer} + + +@pytest.mark.parametrize( + "kind,settings,cli,path,body,expected", + [ + ("completion", {"completion_model": "default"}, "cli", "path", "body", "default"), + ("completion", {}, "cli", "path", "body", "cli"), + ("completion", {}, None, "path", "body", "path"), + ("completion", {}, None, None, "body", "body"), + ( + "image_generation", + {"completion_model": "text", "image_generation_model": "image"}, + None, + None, + "body", + "image", + ), + ("image_generation", {"image_generation_model": "image"}, "cli", "path", "body", "cli"), + ("image_generation", {"image_generation_model": "image"}, None, "path", "body", "path"), + ("image_edit", {"completion_model": "text", "image_generation_model": "image"}, None, None, "body", "text"), + ("image_edit", {"image_generation_model": "image"}, None, "path", "body", "path"), + ("image_edit", {"image_generation_model": "image"}, None, None, "body", "image"), + ("moderation", {"moderation_model": "mod"}, "cli", None, "body", "cli"), + ("speech", {"completion_model": "text"}, None, None, "body", "body"), + ("body", {"completion_model": "text"}, "cli", None, "body", "body"), + ("path", {"completion_model": "text"}, "cli", "path", "body", "path"), + ], +) +def test_shared_inference_model_selection_preserves_handler_precedence( + kind: Literal["completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"], + settings: Mapping[str, object], + cli: str | None, + path: str | None, + body: str, + expected: str, +) -> None: + from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model + + assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py index ca2ff8bcce1..9e20386bf3d 100644 --- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -177,7 +177,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep assert agent.agent_id == agent_id prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once_with( where={"agent_id": agent_id}, - include={"object_permission": True}, + include={"object_permission": True, "identity": True}, ) @@ -202,7 +202,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re assert agent.agent_name == agent_name prisma_client.db.litellm_agentstable.find_unique.assert_awaited_with( where={"agent_name": agent_name}, - include={"object_permission": True}, + include={"object_permission": True, "identity": True}, ) @@ -521,3 +521,33 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra assert await resync_task is True assert len(clean_agent_registry.agent_list) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lookup", ["agent-id", "Agent name"]) +async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch): + from types import SimpleNamespace + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through + + binding = { + "agent_id": "agent-id", "provider": "microsoft_entra", "tenant_id": "tenant", "client_id": "client", + "issuer": "https://login.microsoftonline.com/tenant/v2.0", "revision": "revision", + } + + async def load_row(*, where, include): + if where == {"agent_id": "Agent name"}: + return None + row = FakeAgentRow("agent-id", "Agent name").model_dump() + return SimpleNamespace(model_dump=lambda: {**row, "identity": binding if include.get("identity") else None}) + + prisma = MagicMock() + prisma.db.litellm_agentstable.find_unique = AsyncMock(side_effect=load_row) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + agent = await get_agent_with_read_through(lookup) + assert agent is not None + assert agent.identity is not None + assert agent.identity.model_dump(include=set(binding)) == binding + assert clean_agent_registry.get_agent_by_id(agent_id="agent-id").identity == agent.identity diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index b5e594db701..d7401554d58 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -2712,3 +2712,34 @@ async def test_track_cost_callback_failure_alert_never_carries_request_metadata_ assert "headers" in failure_debug_lines[0] else: assert failure_debug_lines == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("identity_field", ["agent_id", "billing_agent_id"]) +async def test_autonomous_llm_callback_persists_without_human_or_key(identity_field: str) -> None: # test-quality-ok: verifies anonymous-agent charges reach the persistence boundary; no injection seam + kwargs: Final = { + "call_type": "acompletion", + "model": "test-model", + "response_cost": 0.01, + "litellm_params": {"metadata": {identity_field: "autonomous-agent"}}, + } + with patch( + "litellm.proxy.hooks.proxy_track_cost_callback._update_database_and_spend_counters", + new_callable=AsyncMock, + return_value=False, + ) as persist: + await _ProxyDBLogger()._PROXY_track_cost_callback( + kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now() + ) + persist.assert_awaited_once() + assert persist.call_args.kwargs["response_cost"] == 0.01 + assert persist.call_args.kwargs["user_id"] is None + assert persist.call_args.kwargs["user_api_key"] is None + assert persist.call_args.kwargs["kwargs"]["litellm_params"]["metadata"][identity_field] == "autonomous-agent" + + +@pytest.mark.parametrize("agent_id,expected", [(None, False), ("autonomous-agent", True)]) +def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id: str | None, expected: bool) -> None: + assert _should_track_cost_callback( + user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id + ) is expected diff --git a/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py b/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py new file mode 100644 index 00000000000..68fe77cc76a --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/sso/test_agent_subject_enrollment.py @@ -0,0 +1,114 @@ +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import ( + enroll_microsoft_subject, + microsoft_interactive_subject, +) + +TENANT: Final = "11111111-1111-4111-8111-111111111111" +OID: Final = "22222222-2222-4222-8222-222222222222" + + +def test_enrollment_uses_provider_object_id_and_configured_tenant() -> None: + subject: Final = microsoft_interactive_subject( + TENANT, {"id": OID, "mail": "alias@example.com", "tid": "untrusted"}, {} + ) + assert subject is not None + assert subject.oid == OID + assert subject.tenant_id == TENANT + assert subject.issuer == f"https://login.microsoftonline.com/{TENANT}/v2.0" + + +@pytest.mark.parametrize("tenant", [None, "common", "organizations", "invalid"]) +def test_multitenant_sso_does_not_guess_the_subject_tenant(tenant: str | None) -> None: + assert microsoft_interactive_subject(tenant, {"id": OID, "tid": TENANT}, {}) is None + + +@pytest.mark.parametrize("response", [{"mail": "user@example.com"}, {"id": "user@example.com"}, {"id": 42}]) +def test_email_and_configurable_aliases_are_not_human_subject_proof(response: dict[str, object]) -> None: + assert microsoft_interactive_subject(TENANT, response, {}) is None + + +@pytest.mark.parametrize( + "endpoint", ["MICROSOFT_USERINFO_ENDPOINT", "MICROSOFT_TOKEN_ENDPOINT", "MICROSOFT_AUTHORIZATION_ENDPOINT"] +) +def test_custom_provider_endpoints_do_not_enroll_trusted_microsoft_subjects(endpoint: str) -> None: + assert microsoft_interactive_subject(TENANT, {"id": OID}, {endpoint: "https://custom.example"}) is None + + +@pytest.mark.asyncio +async def test_interactive_enrollment_preserves_the_canonical_local_user() -> None: + table: Final = AsyncMock() + table.upsert.return_value = SimpleNamespace(kind="human", user_id="canonical", verified_via="sso_interactive") + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + subject: Final = microsoft_interactive_subject(TENANT, {"id": OID}, {}) + assert subject is not None + await enroll_microsoft_subject(subject, "canonical", client) + table.upsert.assert_awaited_once_with( + where={"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": TENANT, "oid": OID}}, + data={ + "create": { + "issuer": subject.issuer, + "tenant_id": TENANT, + "oid": OID, + "user_id": "canonical", + "verified_via": "sso_interactive", + }, + "update": {}, + }, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("user_id,verified_via", [("another-user", "sso_interactive"), ("canonical", "untrusted")]) +async def test_interactive_enrollment_does_not_reassign_an_existing_subject(user_id: str, verified_via: str) -> None: + table: Final = AsyncMock() + table.upsert.return_value = SimpleNamespace(kind="human", user_id=user_id, verified_via=verified_via) + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + with pytest.raises(HTTPException) as failure: + await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client) + assert failure.value.status_code == 403 + assert table.upsert.call_args.kwargs["data"]["update"] == {} + + +@pytest.mark.asyncio +async def test_enrollment_storage_failure_is_not_a_successful_login() -> None: + table: Final = AsyncMock() + table.upsert.side_effect = RuntimeError("database unavailable") + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + with pytest.raises(HTTPException) as failure: + await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("user_id", [None, "", 42]) +async def test_enrollment_requires_a_canonical_local_user(user_id: object) -> None: + table: Final = AsyncMock() + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), user_id, client) + table.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_untrusted_metadata_cannot_enroll_a_human() -> None: + table: Final = AsyncMock() + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + await enroll_microsoft_subject({"issuer": "forged", "tenant_id": TENANT, "oid": OID}, "canonical", client) + table.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_scim_agent_subject_cannot_be_enrolled_as_a_human() -> None: + table: Final = AsyncMock() + table.upsert.return_value = SimpleNamespace(kind="agent_user", user_id=None, verified_via="scim") + client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table)) + with pytest.raises(HTTPException) as failure: + await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client) + assert failure.value.status_code == 403 + assert table.upsert.call_args.kwargs["data"]["update"] == {} diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 11b3dcf54bc..ce11a604192 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -7659,7 +7659,7 @@ class TestConnectedAppViewAnnotation: flags = {server.server_id: server.connected_app_reachable for server in result} assert flags == {"server-1": True, "server-2": False} - reload_mock.assert_awaited_once_with("test_user_id") + reload_mock.assert_awaited_once_with("test_user_id", requires_fresh_policy=False) mock_manager.get_allowed_mcp_servers.assert_awaited_once_with(admitted_auth) @pytest.mark.asyncio @@ -9411,7 +9411,7 @@ class TestMCPServerResolutionCharacterization: server_id: str, ) -> tuple[MagicMock, MCPServerManager, UserAPIKeyAuth]: team_id: Final = UI_SESSION_TOKEN_TEAM_ID if grant_route == "direct user object_permission" else "lit3974_team" - user_id: Final = "lit3974_direct_user" + user_id: Final = f"{server_id}:{grant_route}:user" key_permission: Final = LiteLLM_ObjectPermissionTable( object_permission_id=f"lit3974_{grant_route}_key_permission", mcp_servers=None, diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 6d902cb7fec..54827f094f0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -4355,7 +4355,7 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs( prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many) prisma_client.db.litellm_teamtable.count = AsyncMock(side_effect=count) prisma_client.db.litellm_verificationtoken.group_by = AsyncMock(return_value=[]) - prisma_client.db.litellm_usertable.find_unique = AsyncMock( + prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable( user_id="org_admin_user", teams=["team_in_org_A", "team_in_org_B"], @@ -4394,11 +4394,11 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs( assert await list_teams(None) == own_view assert await list_teams("org_admin_user", search="team_in_org_B") == ["team_in_org_B"] assert await list_teams("other_user") == ["other_team_in_org_A"] - prisma_client.db.litellm_usertable.find_unique.assert_awaited_with( + prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_with( where={"user_id": "org_admin_user"}, include={"organization_memberships": True} ) - prisma_client.db.litellm_usertable.find_unique.side_effect = RuntimeError("db down") + prisma_client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("db down") with pytest.raises(ValueError, match="db down"): await list_teams("org_admin_user") @@ -16025,7 +16025,7 @@ async def test_get_team_spend_by_user_team_admin_sees_every_member(mock_db_clien alpha = _team_spend_by_user_team("team-alpha", "Team Alpha", Member(user_id="alice", role="admin"), []) mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha]) mock_db_client.db.query_raw = AsyncMock(return_value=[]) - mock_db_client.db.litellm_usertable.find_unique = AsyncMock( + mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=_team_spend_by_user_caller("alice", ["team-alpha"]) ) @@ -16047,7 +16047,7 @@ async def test_get_team_spend_by_user_plain_member_only_sees_own_row(mock_db_cli mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha]) mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_db_client.db.query_raw = AsyncMock(return_value=[_team_spend_by_user_db_row("team-alpha", "bob", 0.25, 2)]) - mock_db_client.db.litellm_usertable.find_unique = AsyncMock( + mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=_team_spend_by_user_caller("bob", ["team-alpha"]) ) @@ -16068,7 +16068,7 @@ async def test_get_team_spend_by_user_member_of_other_team_gets_404(mock_db_clie caller = UserAPIKeyAuth(user_id="bob", user_role=LitellmUserRoles.INTERNAL_USER) mock_db_client.db.query_raw = AsyncMock(return_value=[]) - mock_db_client.db.litellm_usertable.find_unique = AsyncMock( + mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock( return_value=_team_spend_by_user_caller("bob", ["team-alpha"]) ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 21c0f565486..971019160f9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -206,6 +206,7 @@ def test_microsoft_sso_handler_openid_from_response_with_custom_attributes(): def test_get_microsoft_callback_response(): # Arrange mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_response = { "mail": "microsoft_user@example.com", "displayName": "Microsoft User", @@ -8751,6 +8752,7 @@ async def test_redirect_from_openid_persists_assertion_under_canonical_user_id() assertion = assertion_from_sso_login(_ema_id_token(), "rt_1") assert assertion is not None mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://localhost:4000/" mock_request.cookies = {} @@ -8989,6 +8991,7 @@ async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplo """Wiring: the browser login path must reach the diagnostic, not just define it.""" monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid") mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://localhost:4000/" mock_request.cookies = {} diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 7378564f7a8..0157200ed5c 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -4699,24 +4699,28 @@ def _config_agent(agent_name: str) -> Dict[str, Any]: } -class _FakeAgentRow: - """Stand-in for a prisma agent record: supports dict() and .object_permission.""" +def _agent_db_row(agent_id: str, agent_name: str): + import json + from datetime import datetime, timezone - def __init__(self, agent_id: str, agent_name: str) -> None: - self.agent_id = agent_id - self.agent_name = agent_name - self.object_permission = None - self.spend = 0.0 + from prisma.models import LiteLLM_AgentsTable - def __iter__(self): - return iter( - { - "agent_id": self.agent_id, - "agent_name": self.agent_name, - "agent_card_params": {"name": self.agent_name, "url": "http://db-agent"}, - "litellm_params": {}, - }.items() - ) + return LiteLLM_AgentsTable( + agent_id=agent_id, + agent_name=agent_name, + agent_card_params=json.dumps({"name": agent_name, "url": "http://db-agent"}), + extra_headers=[], + agent_access_groups=[], + access_group_ids=[], + spend=0.0, + identity_managed=False, + enabled=True, + execution_mode="autonomous", + created_at=datetime.now(timezone.utc), + updated_at=datetime.now(timezone.utc), + created_by="admin", + updated_by="admin", + ) @pytest.mark.asyncio @@ -4740,7 +4744,7 @@ async def test_ProxyConfig__init_agents_in_db_keeps_config_defined_agents(clean_ ) prisma_client = MagicMock() - prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_FakeAgentRow("db-id", "db-agent")]) + prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_agent_db_row("db-id", "db-agent")]) await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client) @@ -4777,7 +4781,7 @@ async def test_ProxyStartupEvent_jwt_auth_resolves_agent_claims_against_live_reg elif agents_source == "db": prisma_client = MagicMock() prisma_client.db.litellm_agentstable.find_many = AsyncMock( - return_value=[_FakeAgentRow("db-id", "loaded-agent")] + return_value=[_agent_db_row("db-id", "loaded-agent")] ) await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client) else: diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 347adc421a2..8cf34ce51d5 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -4,6 +4,7 @@ import datetime import hashlib import json import re +import sqlite3 from datetime import timezone from unittest.mock import AsyncMock, MagicMock, patch @@ -3762,7 +3763,7 @@ class TestSpendLogsPayload: "model": "gpt-4o", "user": "", "team_id": "", - "metadata": '{"applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', + "metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}', "cache_key": "Cache OFF", "spend": 0.00022500000000000002, "total_tokens": 30, @@ -3781,6 +3782,7 @@ class TestSpendLogsPayload: "status": "success", "mcp_namespaced_tool_name": None, "agent_id": None, + "billing_agent_id": None, } ) @@ -6586,9 +6588,7 @@ def test_key_spend_report_scopes_to_caller_key(client, monkeypatch): def test_key_spend_report_scopes_a_cli_session_to_the_per_user_alias_not_the_login_token(client, monkeypatch): - mock_prisma = _spend_report_mock_prisma( - query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}] - ) + mock_prisma = _spend_report_mock_prisma(query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -7138,9 +7138,8 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc rep_call = emitted[2] assert f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in rep_call[0] assert ( - f"ORDER BY {SESSION_GROUP_KEY_SQL}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" - in rep_call[0] - ), "the session representative must prefer the newest non-MCP call" + f"ORDER BY {SESSION_GROUP_KEY_SQL}, " + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL + ) in rep_call[0] assert rep_call[-2] == ["sess-1", "req-solo"] assert rep_call[-1] == ["hashed-key", "hashed-key"] finally: @@ -7630,6 +7629,39 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo app.dependency_overrides.pop(ps.user_api_key_auth, None) +@pytest.mark.parametrize( + ("parent_status", "child_status", "expected"), + [("failure", "success", "failure"), ("success", "failure", "success")], +) +def test_session_representative_uses_completed_agent_outcome(parent_status, child_status, expected): + with sqlite3.connect(":memory:") as connection: + connection.execute( + 'CREATE TABLE logs (request_id TEXT, call_type TEXT, status TEXT, "startTime" TEXT, "endTime" TEXT)' + ) + connection.executemany( + "INSERT INTO logs VALUES (?, ?, ?, ?, ?)", + ( + ("parent", "asend_message", parent_status, "10:00:00", "10:00:05"), + ("nested-agent", "asend_message", child_status, "10:00:01", "10:00:03"), + ("llm", "acompletion", "success", "10:00:02", "10:00:04"), + ("tool", "call_mcp_tool", child_status, "10:00:04", "10:00:04"), + ), + ) + result = connection.execute( + "SELECT request_id, status FROM logs ORDER BY " + + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL + + " LIMIT 1" + ).fetchone() + assert result == ("parent", expected) + connection.execute("DELETE FROM logs WHERE call_type = 'asend_message'") + fallback = connection.execute( + "SELECT request_id, status FROM logs ORDER BY " + + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL + + " LIMIT 1" + ).fetchone() + assert fallback == ("llm", "success") + + @pytest.mark.asyncio async def test_calculate_spend_unpriced_model_returns_400(): model = "openrouter/unit-test-unpriced-model" diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py index 54e5a6d5385..6752c91e9f2 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py @@ -539,7 +539,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): async def mock_query_raw(sql_query, *params): if "COUNT(*) AS total_count" in sql_query: return [{"total_count": 60}] - if "DISTINCT ON" in sql_query: + if "AS session_representatives" in sql_query: return representative_rows return session_rows @@ -584,9 +584,11 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): rep_sql = emitted[2][0] assert f"DISTINCT ON ({group_key})" in rep_sql, f"page must return one row per session. SQL was:\n{rep_sql}" - assert f"ORDER BY {group_key}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in rep_sql, ( - "the session representative must prefer the newest non-MCP call" - ) + assert ( + f"ORDER BY {group_key}, (call_type = 'asend_message') DESC, " + "CASE WHEN call_type = 'asend_message' THEN \"endTime\" END DESC NULLS LAST, " + "call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" + ) in rep_sql, "the session representative must prefer the final agent outcome, then the newest non-MCP call" assert "COUNT(*) OVER ()" not in rep_sql assert [row["request_id"] for row in response["data"]] == ["req-1", "req-2"] diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 00223f192ec..cef60e5c8da 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -5156,6 +5156,24 @@ def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_recei ) +def test_failed_agent_request_keeps_registered_display_name(): + agent_model: Final = "a2a_agent/Research Agent" + payload: Final = get_logging_payload( + kwargs={ + "model": agent_model, + "call_type": "asend_message", + "litellm_params": { + "metadata": {"model_group": agent_model, "model_info": {"id": "registered-agent"}, "status": "failure"} + }, + }, + response_obj=ValueError("Agent action denied"), + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model"] == agent_model + assert payload["status"] == "failure" + assert payload["model_id"] == "registered-agent" + _CLI_SESSION_ALIAS: Final = "cli-session-alice" _CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" @@ -5281,3 +5299,21 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None: assert result["autorouter_savings_estimate"] == recorded absent: Final = _get_spend_logs_metadata({"autorouter_savings_estimate": supplied}) # mutable-ok: legacy metadata helper accepts dicts assert absent["autorouter_savings_estimate"] is None + + +@pytest.mark.parametrize("billing_agent", [None, "authenticated-agent"]) +def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None: + kwargs = { + "model": "gpt-4", + "litellm_params": {"metadata": { + "user_api_key": "test-key", + "agent_id": "header-selected-agent", + "billing_agent_id": billing_agent, + }}, + } + payload = get_logging_payload( + kwargs=kwargs, response_obj={"id": "request"}, + start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["agent_id"] == "header-selected-agent" + assert payload["billing_agent_id"] == billing_agent diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 35308474949..0429dd97a1c 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -149,6 +149,9 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re prisma_client.db.litellm_agentstable.find_unique = AsyncMock( side_effect=[None, _DbAgentRow("a2a-sibling-replica-agent-id", agent_name)] ) + prisma_client.writer_db.litellm_agentstable.find_unique = AsyncMock( + return_value=_DbAgentRow("a2a-sibling-replica-agent-id", agent_name) + ) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx new file mode 100644 index 00000000000..ce54ab78d9c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.test.tsx @@ -0,0 +1,43 @@ +import { screen } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { renderWithProviders, testQueryClient } from "../../../../../tests/test-utils"; +import { apiClient } from "@/components/networking"; +import { AgentIdentityDetails } from "./AgentIdentityDetails"; + +vi.mock("@/components/networking", () => ({ apiClient: { get: vi.fn() } })); + +const identity = { + provider: "microsoft_entra", + tenant_id: "11111111-1111-4111-8111-111111111111", + client_id: "22222222-2222-4222-8222-222222222222", +}; + +const status = { + enabled: true, + execution_mode: "autonomous", + last_authenticated_at: "2026-09-24T12:00:00Z", +}; + +describe("agent identity evidence", () => { + beforeEach(() => { + vi.clearAllMocks(); + testQueryClient.clear(); + }); + + it("shows persisted application identity evidence and links to the current logs route", async () => { + vi.mocked(apiClient.get).mockResolvedValue(status); + renderWithProviders(); + expect(await screen.findByText(/Last authenticated identity match:/)).toBeInTheDocument(); + expect(screen.getByText(/Application \(Client\) ID:/)).toBeInTheDocument(); + expect(screen.getByRole("link", { name: "View request logs" })).toHaveAttribute("href", "/ui/logs/"); + expect(apiClient.get).toHaveBeenCalledWith("/v1/agents/native/identity", { accessToken: "admin" }); + }); + + it("does not request or show administrator identity evidence to ordinary users", () => { + renderWithProviders( + , + ); + expect(screen.queryByRole("region", { name: "Agent Identity" })).not.toBeInTheDocument(); + expect(apiClient.get).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx new file mode 100644 index 00000000000..12465c6e861 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityDetails.tsx @@ -0,0 +1,81 @@ +import React from "react"; +import type { components } from "@/lib/http/schema"; +import { useQuery } from "@tanstack/react-query"; +import { apiClient } from "@/components/networking"; +import { Button } from "@/components/ui/button"; +import { readAgentIdentity } from "./agent_identity"; + +const authenticationMessage = (error: boolean, lastAuthenticated?: string | null): string => { + if (error) return "Could not load authentication evidence"; + if (lastAuthenticated) return `Last authenticated identity match: ${new Date(lastAuthenticated).toLocaleString()}`; + return "Configured, awaiting an authenticated request"; +}; + +export const AgentIdentityDetails = ({ + agentId, + identity: value, + accessToken, + isAdmin, +}: { + agentId: string; + identity: unknown; + accessToken: string | null; + isAdmin: boolean; +}) => { + const identity = readAgentIdentity(value); + const { data, isError, isFetching, refetch } = useQuery({ + queryKey: ["agent-identity", agentId, identity], + queryFn: () => + apiClient.get( + `/v1/agents/${encodeURIComponent(agentId)}/identity`, + { + accessToken: accessToken ?? "", + }, + ), + enabled: Boolean(isAdmin && accessToken && identity), + }); + + if (!identity || !isAdmin) return null; + const executionLabel = data?.enabled ? "Enabled" : "Disabled"; + return ( +
+

Agent Identity: Microsoft Entra ID

+

+ Tenant: {identity.tenant_id} +

+ <> +

+ Application (Client) ID: {identity.client_id} +

+

Enterprise application Object ID: {identity.service_principal_id || "Not configured"}

+ +

+ Execution: {data ? executionLabel : "Loading"} · Mode: {data?.execution_mode ?? "Loading"} +

+

+ {data?.identity?.active === false + ? "Identity unbound; execution is disabled" + : authenticationMessage(isError, data?.last_authenticated_at)} +

+

+ Recent evidence comes from a validated Entra token matching this binding. It is persisted across restarts and + cleared when the binding changes. Tool and model permissions are checked separately. +

+
+ + + View request logs + +
+
+ ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx new file mode 100644 index 00000000000..43538784145 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentIdentityFields.tsx @@ -0,0 +1,259 @@ +import React, { useEffect, useState } from "react"; +import { useWatch } from "react-hook-form"; +import { apiClient } from "@/components/networking"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { AgentFormField, type AgentFormValues } from "./AgentFormKit"; +import { entraTenantFromIssuer, IDENTITY_UUID_PATTERN } from "./agent_identity"; + +const PROVIDER_OPTIONS = [ + { value: "none", label: "No explicit identity binding" }, + { value: "microsoft_entra", label: "Microsoft Entra ID" }, +]; +const EXECUTION_MODE_OPTIONS = [ + { value: "autonomous", label: "Autonomous" }, + { value: "delegated", label: "On behalf of a user" }, + { value: "both", label: "Both" }, +]; +const EXECUTION_OPTIONS = [ + { value: "enabled", label: "Enabled" }, + { value: "disabled", label: "Disabled" }, +]; + +export const AgentIdentityFields = ({ accessToken }: { accessToken: string | null }) => { + const provider = useWatch({ name: "identity_provider" }); + const mode = useWatch({ name: "execution_mode" }); + const showScopes = mode !== "autonomous" && mode !== undefined; + const [tenants, setTenants] = useState([]); + const [error, setError] = useState(null); + + useEffect(() => { + if (!accessToken || provider !== "microsoft_entra") return; + let active = true; + apiClient + .get("/v1/agents/identity/providers", { accessToken }) + .then((issuers) => { + if (active) + setTenants( + issuers.flatMap((issuer) => { + const tenant = entraTenantFromIssuer(issuer); + return tenant ? [tenant] : []; + }), + ); + }) + .catch(() => { + if (active) setError("Could not load the gateway's trusted identity providers"); + }); + return () => { + active = false; + }; + }, [accessToken, provider]); + + return ( + <> +
+
+

Agent Identity

+

+ Connect an existing identity provider application to this agent. Its name and runtime address can change + independently. +

+
+ + {({ value, onChange, id }) => ( + + )} + + {provider === "microsoft_entra" && ( + <> + + {({ value, onChange, id }) => ( + + )} + + {error && ( +

+ {error} +

+ )} + {!error && tenants.length === 0 && ( +

+ No trusted Entra tenant is available. Configure JWT issuer and audience validation on the gateway first. + Dashboard Microsoft SSO is configured separately. +

+ )} + + Find this under{" "} + + Entra App registrations + + , select your agent application, then Overview. No client secret is required here. + + } + > + {({ value, onChange, ref, ...control }) => ( + + )} + + + {({ value, onChange, id }) => ( + + )} + + + Open{" "} + + Entra Enterprise applications + + , select this application, and copy its Object ID. The App registrations Object ID is a different + value. + + } + > + {({ value, onChange, ref, ...control }) => ( + + )} + + + + {({ value, onChange, ref, ...control }) => ( + + )} + + + {showScopes && ( + <> + + {({ value, onChange, ref, ...control }) => ( + + )} + +

+ Users must first sign in through this gateway's Microsoft SSO. Subsequent delegated calls must + satisfy both user and agent permissions. +

+ + )} + + {({ value, onChange, id }) => ( + + )} + +

+ LiteLLM verifies the agent's Entra token before matching this identity. Saving these fields + configures the binding; an authenticated request provides verification. Runtime authentication headers are + configured separately. +

+ + )} +
+ + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx index f53b03b6a08..b392e270d33 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx @@ -145,10 +145,10 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams

- Why do agents need keys? + How do agents authenticate? - Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from - the Virtual Keys page. + Agents can authenticate with a virtual key or a trusted identity provider using JWT. Configure an identity + binding when adding or editing an agent. JWT authentication does not require a virtual key. {isAdmin && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx index bef938cd31c..95fcb516c55 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx @@ -62,6 +62,12 @@ describe("AgentsTable", () => { expect(within(keylessRow).getByText("Needs Setup")).toBeInTheDocument(); }); + it("shows JWT configured for agents without a virtual key", () => { + render(); + expect(screen.getByText("JWT configured")).toBeInTheDocument(); + expect(screen.queryByText("Needs Setup")).not.toBeInTheDocument(); + }); + it("deletes an agent through the ⋯ actions menu", async () => { const user = userEvent.setup(); const onDeleteClick = vi.fn(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx index a8fe3973a42..932bd5c513f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx @@ -136,6 +136,7 @@ export const getAgentsTableColumns = ({ enableSorting: false, cell: ({ row }) => { const hasKeys = (row.original.keys?.length ?? 0) > 0; + if (row.original.jwt_auth_configured) return ; return hasKeys ? ( ) : ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx index 457ee656415..b7585e8f8b2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.integration.test.tsx @@ -1,5 +1,5 @@ import React from "react"; -import { screen, waitFor, within } from "@testing-library/react"; +import { fireEvent, screen, waitFor, within } from "@testing-library/react"; import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; import { describe, it, expect, vi, beforeEach } from "vitest"; import AddAgentForm from "./add_agent_form"; @@ -8,6 +8,7 @@ import type { AgentCreateInfo } from "@/components/networking"; import { chooseSelectOption, renderWithProviders as render } from "../../../../../tests/test-utils"; vi.mock("@/components/networking", () => ({ + apiClient: { get: vi.fn() }, createAgentCall: vi.fn(), getAgentCreateMetadata: vi.fn(), getAgentsList: vi.fn(), @@ -95,6 +96,51 @@ describe("AddAgentForm submit payload", () => { .mockResolvedValue({} as never); }); + it("registers a readable agent with an explicit Entra identity and no virtual key", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const tenant = "11111111-1111-4111-8111-111111111111"; + const clientId = "22222222-2222-4222-8222-222222222222"; + vi.mocked(networking.apiClient.get).mockResolvedValue([`https://login.microsoftonline.com/${tenant}/v2.0`]); + renderForm(); + fireEvent.change(await screen.findByLabelText("Agent Name"), { target: { value: "Readable agent" } }); + fireEvent.change(screen.getByLabelText("URL"), { target: { value: "https://runtime.example/a2a" } }); + fireEvent.change(screen.getByLabelText("Display Name"), { target: { value: "Readable agent" } }); + fireEvent.change(screen.getByPlaceholderText("Describe what this agent does..."), { + target: { value: "Test agent" }, + }); + await user.click(screen.getByLabelText("Identity Provider")); + await user.click(await screen.findByRole("option", { name: "Microsoft Entra ID" })); + await user.click(screen.getByLabelText("Trusted Entra Tenant")); + await user.click(await screen.findByRole("option", { name: tenant })); + fireEvent.change(screen.getByLabelText("Application (Client) ID"), { target: { value: clientId } }); + fireEvent.change(screen.getByLabelText("Enterprise Application Object ID"), { + target: { value: "33333333-3333-4333-8333-333333333333" }, + }); + await user.click(screen.getByRole("button", { name: /^Next/ })); + await user.click(screen.getByRole("button", { name: /^Next/ })); + await user.click(screen.getByRole("button", { name: /^Next/ })); + await user.click(screen.getByRole("button", { name: "Use Entra JWT authentication" })); + await user.click(screen.getByRole("button", { name: /Create Agent/ })); + await waitFor(() => expect(networking.createAgentCall).toHaveBeenCalledTimes(1)); + expect(createdPayload().agent_name).toBe("Readable agent"); + const expectedIdentity = { + provider: "microsoft_entra", + tenant_id: tenant, + client_id: clientId, + service_principal_id: "33333333-3333-4333-8333-333333333333", + required_roles: [], + required_scopes: ["user_impersonation"], + }; + expect(createdPayload().identity).toEqual(expectedIdentity); + expect(createdPayload()).not.toHaveProperty("litellm_params.identity"); + expect(networking.keyCreateForAgentCall).not.toHaveBeenCalled(); + expect( + screen.getByText( + "Microsoft Entra ID is configured. Send an authenticated agent request to verify the connection.", + ), + ).toBeInTheDocument(); + }); + it("sends every a2a field the user filled across all collapsible panels", async () => { const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx index fbf5cf8c1fb..5e8ba145396 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx @@ -83,7 +83,9 @@ describe("AddAgentForm logos", () => { expect(titleLogo).toBeInstanceOf(HTMLImageElement); expect(titleLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png")); - const selectionLogo = within(await screen.findByRole("combobox")).getByAltText("A2A Agent logo"); + const selectionLogo = within(await screen.findByRole("combobox", { name: "Agent Type" })).getByAltText( + "A2A Agent logo", + ); expect(selectionLogo).toBeInstanceOf(HTMLImageElement); expect(selectionLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png")); }); @@ -93,14 +95,14 @@ describe("AddAgentForm logos", () => { await screen.findByAltText("A2A Agent logo"); - expect(screen.getByLabelText("Agent Type")).toBe(screen.getByRole("combobox")); + expect(screen.getByLabelText("Agent Type")).toBe(screen.getByRole("combobox", { name: "Agent Type" })); }); it("renders the option logo when the agent type dropdown is opened", async () => { const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); renderForm(); - const trigger = await screen.findByRole("combobox"); + const trigger = await screen.findByRole("combobox", { name: "Agent Type" }); await within(trigger).findByAltText("A2A Agent logo"); await user.click(trigger); @@ -123,7 +125,7 @@ describe("AddAgentForm logos", () => { expect(screen.queryByAltText("Agent logo")).not.toBeInTheDocument(); expect(within(header).getByText("A")).toBeInTheDocument(); - const trigger = screen.getByRole("combobox"); + const trigger = screen.getByRole("combobox", { name: "Agent Type" }); fireEvent.error(within(trigger).getByAltText("A2A Agent logo")); expect(within(trigger).queryByAltText("A2A Agent logo")).not.toBeInTheDocument(); expect(warnSpy).toHaveBeenCalledTimes(2); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx index 5bd6ea9b83a..b243d9d1601 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx @@ -1,3 +1,5 @@ +import { AgentIdentityFields } from "./AgentIdentityFields"; +import { withAgentIdentity } from "./agent_identity"; import React, { useState, useEffect } from "react"; import { FormProvider, useForm, useWatch } from "react-hook-form"; import { toast } from "@/lib/toast"; @@ -287,6 +289,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok const buildAgentData = (values: AgentFormValues): AgentRequestPayload | null => { if (agentType === CUSTOM_AGENT_TYPE) { + if (values.identity_provider === "microsoft_entra") return { agent_name: values.agent_name }; return { agent_name: values.agent_name, agent_card_params: { @@ -353,12 +356,13 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok return; } const values = form.getValues(); - const agentData = buildAgentData(values); - if (!agentData) { + const built = buildAgentData(values); + if (!built) { toast.error("Failed to build agent data"); setIsSubmitting(false); return; } + const agentData = withAgentIdentity(built, values); // Build object_permission from MCP Tools step (allowed_mcp_servers_and_groups, mcp_tool_permissions) const mcpServersAndGroups = values.allowed_mcp_servers_and_groups ?? {}; @@ -792,7 +796,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok - For agents that don't follow a standard protocol, just needs a virtual key + For outbound agents using an identity provider or virtual key @@ -801,6 +805,8 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok + +
{agentType === CUSTOM_AGENT_TYPE ? ( @@ -910,7 +916,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok name="team_id" label={labelWithHint( "Assign to Team", - "Optionally assign this agent to a team. The agent and its key will belong to the selected team.", + "Optionally select a team for the virtual key. The agent identity and its permissions are managed separately.", )} > {({ value, onChange }) => ( @@ -920,6 +926,11 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok + {form.getValues("identity_provider") === "microsoft_entra" && ( +

+ This agent will authenticate with Microsoft Entra ID. You can skip virtual key creation. +

+ )} setKeyAssignOption(value as "create_new" | "existing_key" | "skip")} @@ -1004,7 +1015,9 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok className="text-sm text-muted-foreground underline hover:text-foreground" onClick={() => setKeyAssignOption("skip")} > - Skip for now — I'll assign a key later + {form.getValues("identity_provider") === "microsoft_entra" + ? "Use Entra JWT authentication" + : "Skip for now, I’ll assign a key later"}
@@ -1033,7 +1046,9 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok )} {!createdKeyValue && !assignedKeyAlias && keyAssignOption === "skip" && (

- No key assigned. You can create one from the Virtual Keys page. + {form.getValues("identity_provider") === "microsoft_entra" + ? "Microsoft Entra ID is configured. Send an authenticated agent request to verify the connection." + : "No key assigned. You can create one from the Virtual Keys page."}

)} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts index 16ce6848402..6ec3c3181f7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_config.ts @@ -1,3 +1,4 @@ +import { parseIdentityForForm } from "./agent_identity"; /** * Shared configuration for agent form fields * Used across create, view, and update operations @@ -57,7 +58,7 @@ export const AGENT_FORM_CONFIG: { name: "description", label: "Description", type: "textarea", - required: true, + required: false, placeholder: "Describe what this agent does...", rows: 3, }, @@ -340,6 +341,7 @@ export const parseAccessGroupIdsForForm = (agent: { access_group_ids?: string[] }); export const parseMcpPermissionsForForm = (agent: any) => ({ + ...parseIdentityForForm(agent), allowed_mcp_servers_and_groups: { servers: agent.object_permission?.mcp_servers ?? [], accessGroups: agent.object_permission?.mcp_access_groups ?? [], @@ -363,8 +365,9 @@ export const buildMcpObjectPermission = (values: any) => ({ * Parse agent data for form fields */ export const parseAgentForForm = (agent: any) => { + const card = agent.agent_card_params ?? {}; const skills = - agent.agent_card_params?.skills?.map((skill: any) => ({ + card.skills?.map((skill: any) => ({ ...skill, tags: skill.tags, examples: skill.examples || [], @@ -372,18 +375,18 @@ export const parseAgentForForm = (agent: any) => { return { agent_name: agent.agent_name, - name: agent.agent_card_params?.name, - description: agent.agent_card_params?.description, - url: agent.agent_card_params?.url, - version: agent.agent_card_params?.version, - protocolVersion: agent.agent_card_params?.protocolVersion, - streaming: agent.agent_card_params?.capabilities?.streaming, - pushNotifications: agent.agent_card_params?.capabilities?.pushNotifications, - stateTransitionHistory: agent.agent_card_params?.capabilities?.stateTransitionHistory, + name: card.name || agent.agent_name, + description: card.description, + url: card.url, + version: card.version, + protocolVersion: card.protocolVersion, + streaming: card.capabilities?.streaming, + pushNotifications: card.capabilities?.pushNotifications, + stateTransitionHistory: card.capabilities?.stateTransitionHistory, skills: skills, - iconUrl: agent.agent_card_params?.iconUrl, - documentationUrl: agent.agent_card_params?.documentationUrl, - supportsAuthenticatedExtendedCard: agent.agent_card_params?.supportsAuthenticatedExtendedCard, + iconUrl: card.iconUrl, + documentationUrl: card.documentationUrl, + supportsAuthenticatedExtendedCard: card.supportsAuthenticatedExtendedCard, model: agent.litellm_params?.model, make_public: agent.litellm_params?.make_public, cost_per_query: agent.litellm_params?.cost_per_query, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts new file mode 100644 index 00000000000..0639e7a6dd4 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.test.ts @@ -0,0 +1,83 @@ +import { describe, expect, it } from "vitest"; +import { + buildIdentityParams, + entraTenantFromIssuer, + parseIdentityForForm, + readAgentIdentity, + withAgentIdentity, +} from "./agent_identity"; + +const identity = { + provider: "microsoft_entra", + tenant_id: "11111111-1111-4111-8111-111111111111", + client_id: "22222222-2222-4222-8222-222222222222", + service_principal_id: "33333333-3333-4333-8333-333333333333", + required_roles: ["Agent.Invoke"], + required_scopes: ["user_impersonation"], +} satisfies import("./agent_identity").EntraAgentIdentity; + +describe("agent identity configuration", () => { + it("round trips an existing binding independently of the agent name and runtime", () => { + const values = { + ...parseIdentityForForm({ + identity: { ...identity, agent_id: "stable", active: true, revision: "rev", issuer: "https://issuer.example" }, + }), + agent_name: "Renamed", + url: "https://new-runtime.example", + }; + expect(buildIdentityParams(values)).toEqual({ identity }); + }); + it("preserves untouched bindings and explicitly clears a removed binding", () => { + expect(buildIdentityParams({ agent_name: "legacy" })).toEqual({}); + expect(buildIdentityParams({ identity_provider: "none" }, identity)).toEqual({ identity: null }); + expect(parseIdentityForForm({}).identity_provider).toBe("none"); + }); + it.each([ + null, + {}, + "invalid", + { ...identity, client_id: "bad" }, + { ...identity, tenant_id: 3 }, + { ...identity, provider: "other" }, + ])("rejects malformed bindings: %j", (value) => { + expect(readAgentIdentity(value)).toBeNull(); + }); + it("rejects incomplete submissions", () => { + expect(() => buildIdentityParams({ identity_provider: "microsoft_entra" })).toThrow("Enter valid Entra"); + }); + it("submits identity as top-level settings without changing runtime parameters", () => { + const formValues = { + identity_provider: "microsoft_entra", + identity_tenant_id: identity.tenant_id, + identity_client_id: identity.client_id, + identity_service_principal_id: identity.service_principal_id, + execution_mode: "both", + enabled: false, + }; + const payload = withAgentIdentity({ litellm_params: { model: "runtime" } }, formValues); + expect(payload.litellm_params).toEqual({ model: "runtime" }); + expect(payload.identity).toMatchObject({ + client_id: identity.client_id, + service_principal_id: identity.service_principal_id, + }); + expect(payload.execution_mode).toBe("both"); + expect(payload.enabled).toBe(false); + }); + it("requires a service principal for autonomous execution", () => { + const values = { + identity_provider: "microsoft_entra", + identity_tenant_id: identity.tenant_id, + identity_client_id: identity.client_id, + execution_mode: "autonomous", + }; + expect(() => buildIdentityParams(values)).toThrow("Enterprise application Object ID"); + }); + + it("only offers tenant-specific Microsoft issuers", () => { + expect(entraTenantFromIssuer(`https://login.microsoftonline.com/${identity.tenant_id}/v2.0`)).toBe( + identity.tenant_id, + ); + expect(entraTenantFromIssuer("https://attacker.example/tenant/v2.0")).toBeNull(); + expect(entraTenantFromIssuer("https://login.microsoftonline.com/common/v2.0")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts new file mode 100644 index 00000000000..34986f75679 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_identity.ts @@ -0,0 +1,101 @@ +import { z } from "zod"; +import type { components } from "@/lib/http/schema"; +import type { AgentFormValues, AgentRequestPayload } from "./AgentFormKit"; + +export type EntraAgentIdentity = components["schemas"]["EntraIdentityConfig"]; +type AgentIdentityState = Pick; + +export const IDENTITY_UUID_PATTERN = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i; + +const stringGrants = (fallback: string[]) => + z + .unknown() + .transform((value) => + Array.isArray(value) ? value.filter((entry): entry is string => typeof entry === "string") : fallback, + ); + +const identityShape = { + provider: z.literal("microsoft_entra"), + tenant_id: z.string().regex(IDENTITY_UUID_PATTERN), + client_id: z.string().regex(IDENTITY_UUID_PATTERN), + service_principal_id: z.string().regex(IDENTITY_UUID_PATTERN).nullable().default(null), + required_roles: stringGrants([]), + required_scopes: stringGrants(["user_impersonation"]), +}; +const identitySchema = z.object(identityShape); + +export const readAgentIdentity = (value: unknown): EntraAgentIdentity | null => { + const parsed = identitySchema.safeParse(value); + return parsed.success ? parsed.data : null; +}; + +const identityFormFields = (identity: EntraAgentIdentity | null): AgentFormValues => ({ + identity_provider: identity?.provider ?? "none", + identity_tenant_id: identity?.tenant_id ?? "", + identity_client_id: identity?.client_id ?? "", + identity_service_principal_id: identity?.service_principal_id ?? "", + identity_required_roles: identity?.required_roles?.join(", ") ?? "", + identity_required_scopes: identity?.required_scopes?.join(", ") ?? "user_impersonation", +}); + +export const parseIdentityForForm = (agent?: Partial | null): AgentFormValues => { + const identity = agent?.identity?.active === false ? null : readAgentIdentity(agent?.identity); + return { + ...identityFormFields(identity), + execution_mode: agent?.execution_mode ?? "autonomous", + enabled: agent?.enabled ?? true, + }; +}; + +const splitGrants = (value: unknown, fallback: string[]): string[] => + typeof value === "string" + ? value + .split(",") + .map((item) => item.trim()) + .filter(Boolean) + : fallback; + +export const buildIdentityParams = ( + values: AgentFormValues, + existingIdentity?: unknown, +): { identity?: EntraAgentIdentity | null } => { + if (values.identity_provider === undefined) return {}; + if (values.identity_provider !== "microsoft_entra") + return readAgentIdentity(existingIdentity) ? { identity: null } : {}; + const candidate: EntraAgentIdentity = { + provider: "microsoft_entra", + tenant_id: typeof values.identity_tenant_id === "string" ? values.identity_tenant_id.trim().toLowerCase() : "", + client_id: typeof values.identity_client_id === "string" ? values.identity_client_id.trim().toLowerCase() : "", + service_principal_id: + typeof values.identity_service_principal_id === "string" && values.identity_service_principal_id.trim() + ? values.identity_service_principal_id.trim().toLowerCase() + : null, + required_roles: splitGrants(values.identity_required_roles, []), + required_scopes: splitGrants(values.identity_required_scopes, ["user_impersonation"]), + }; + const identity = readAgentIdentity(candidate); + if (!identity) throw new Error("Enter valid Entra tenant, application client and service principal IDs"); + if (values.execution_mode !== "delegated" && !identity.service_principal_id) + throw new Error("Autonomous agents require the Enterprise application Object ID"); + return { identity }; +}; + +export const entraTenantFromIssuer = (issuer: string): string | null => { + const match = /^https:\/\/login\.microsoftonline\.com\/([^/]+)\/v2\.0$/.exec(issuer); + return match && IDENTITY_UUID_PATTERN.test(match[1]) ? match[1] : null; +}; + +export const withAgentIdentity = ( + payload: AgentRequestPayload, + values: AgentFormValues, + existing?: Partial, +): AgentRequestPayload => { + const identityFields = buildIdentityParams(values, existing?.identity); + const managed = values.identity_provider === "microsoft_entra" || Boolean(readAgentIdentity(existing?.identity)); + return { + ...payload, + ...identityFields, + ...(managed && values.execution_mode !== undefined ? { execution_mode: values.execution_mode } : {}), + ...(managed && values.enabled !== undefined ? { enabled: values.enabled } : {}), + }; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx index 37e00766a75..5f5eb4e020c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.integration.test.tsx @@ -8,6 +8,7 @@ import * as networking from "@/components/networking"; import type { AgentCreateInfo } from "@/components/networking"; vi.mock("@/components/networking", () => ({ + apiClient: { get: vi.fn() }, getAgentInfo: vi.fn(), patchAgentCall: vi.fn(), getAgentCreateMetadata: vi.fn(), @@ -155,6 +156,50 @@ describe("AgentInfoView update payload", () => { .mockResolvedValue({} as never); }); + it.each(["complete", "empty"])("preserves the Entra binding while renaming an agent with a %s card", async (card) => { + const user = setup(); + const identity = { + provider: "microsoft_entra", + tenant_id: "11111111-1111-4111-8111-111111111111", + client_id: "22222222-2222-4222-8222-222222222222", + service_principal_id: "33333333-3333-4333-8333-333333333333", + }; + const params = { ...A2A_AGENT.litellm_params, require_trace_id_on_calls_by_agent: true }; + vi.mocked(networking.getAgentInfo).mockResolvedValue({ + ...A2A_AGENT, + agent_card_params: card === "empty" ? {} : A2A_AGENT.agent_card_params, + litellm_params: params, + identity: { ...identity, agent_id: "agent-1", issuer: "https://issuer.example", revision: "rev", active: true }, + identity_managed: true, + execution_mode: "autonomous", + enabled: true, + access_group_ids: ["ag-entra"], + } as never); + vi.mocked(networking.apiClient.get).mockImplementation(async (path) => + path.endsWith("/providers") + ? [`https://login.microsoftonline.com/${identity.tenant_id}/v2.0`] + : { last_authenticated_at: null }, + ); + renderView(); + expect(await screen.findByText("Configured, awaiting an authenticated request")).toBeInTheDocument(); + await openEditor(user); + expect(screen.getByLabelText("Application (Client) ID")).toHaveValue(identity.client_id); + expect(screen.getByRole("combobox", { name: "Identity Provider" })).toHaveTextContent("Microsoft Entra ID"); + expect(screen.getByRole("combobox", { name: "Execution Mode" })).toHaveTextContent("Autonomous"); + expect(screen.getByRole("combobox", { name: /^Execution$/ })).toHaveTextContent("Enabled"); + fireEvent.change(screen.getByLabelText("Agent Name"), { target: { value: "Renamed agent" } }); + await save(user); + expect(patchedPayload().agent_name).toBe("Renamed agent"); + expect(patchedPayload()).not.toHaveProperty("litellm_params"); + expect(patchedPayload().identity).toMatchObject(identity); + expect(patchedPayload().access_group_ids).toEqual(["ag-entra"]); + expect(networking.patchAgentCall).toHaveBeenCalledWith( + "tok", + "agent-1", + expect.objectContaining({ agent_name: "Renamed agent" }), + ); + }); + it("sends only the fields whose panel has been opened, dropping the rest", async () => { const user = setup(); renderView(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx index 19b1ee8ca48..4eb3534ebe3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.test.tsx @@ -2,6 +2,7 @@ import React from "react"; import { fireEvent, render, screen, waitFor } from "@testing-library/react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import AgentInfoView from "./agent_info"; +import AgentFormFields from "./agent_form_fields"; import * as networking from "@/components/networking"; import type { Agent } from "@/components/agents/types"; @@ -16,12 +17,16 @@ vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({ useKeys: () => ({ data: { keys: [] }, isLoading: false, refetch: vi.fn() }), })); +vi.mock("./AgentIdentityDetails", () => ({ + AgentIdentityDetails: () => null, +})); + vi.mock("./agent_card_discovery", () => ({ default: () =>
, })); vi.mock("./agent_form_fields", () => ({ - default: () =>
, + default: vi.fn(() =>
), unmountedA2AFieldNames: () => [], })); @@ -77,6 +82,9 @@ const agent = { describe("AgentInfoView settings", () => { beforeEach(() => { vi.restoreAllMocks(); + vi.mocked(AgentFormFields) + .mockReset() + .mockImplementation(() =>
); vi.mocked(networking.getAgentInfo).mockReset().mockResolvedValue(agent); vi.mocked(networking.getAgentCreateMetadata).mockReset().mockResolvedValue([]); vi.mocked(networking.patchAgentCall).mockReset().mockResolvedValue({}); @@ -104,6 +112,23 @@ describe("AgentInfoView settings", () => { expect(payload.access_group_ids).toEqual([]); }); + it("saves unrelated settings when the existing card has no description", async () => { + const actual = await vi.importActual("./agent_form_fields"); + vi.mocked(AgentFormFields).mockImplementation(actual.default); + const { description: _description, ...card } = agent.agent_card_params ?? {}; + vi.mocked(networking.getAgentInfo).mockResolvedValue({ ...agent, agent_card_params: card }); + render(); + fireEvent.click(await screen.findByRole("tab", { name: "Settings" })); + fireEvent.click(screen.getByRole("button", { name: "Edit Settings" })); + expect(await screen.findByLabelText("Description")).toHaveValue(""); + fireEvent.change(screen.getByLabelText("TPM Limit"), { target: { value: "42" } }); + fireEvent.click(screen.getByRole("button", { name: /Save Changes/ })); + await waitFor(() => expect(networking.patchAgentCall).toHaveBeenCalledOnce()); + const [, , payload] = vi.mocked(networking.patchAgentCall).mock.calls[0]; + expect(payload.tpm_limit).toBe(42); + expect(payload.agent_card_params?.description).toBe(""); + }); + it("sends the newly attached access group in the update payload", async () => { render(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx index 6e7389a3fa1..d4c05d0ebac 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/agent_info.tsx @@ -1,3 +1,6 @@ +import { AgentIdentityFields } from "./AgentIdentityFields"; +import { AgentIdentityDetails } from "./AgentIdentityDetails"; +import { withAgentIdentity } from "./agent_identity"; import React, { useState, useEffect, useMemo } from "react"; import { cx } from "@/lib/cva.config"; import { FormProvider, useForm, useWatch } from "react-hook-form"; @@ -237,7 +240,7 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT : built; await patchAgentCall(accessToken, agentId, { - ...updateData, + ...withAgentIdentity(updateData, values, agent), object_permission: buildMcpObjectPermission(values), access_group_ids: values.access_group_ids ?? [], }); @@ -337,6 +340,12 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT
{/* Overview Panel */} + {agent.agent_id} {agent.agent_name} @@ -505,6 +514,8 @@ const AgentInfoView: React.FC = ({ agentId, onClose, accessT )} + + {discoveryRequest && (
{ expect(screen.getByRole("option", { name: /group-b/ })).toBeInTheDocument(); }); + it("offers only individual agents when legacy groups are disabled", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + render(); + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByRole("option", { name: /Agent One/ })); + expect(screen.queryByRole("option", { name: /group-a/ })).not.toBeInTheDocument(); + expect(onChange).toHaveBeenCalledWith({ agents: ["agent-1"], accessGroups: [] }); + }); + + it("preserves a saved legacy group while adding an agent with legacy groups disabled", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + render( + , + ); + await user.click(screen.getByRole("combobox")); + expect(await screen.findByRole("option", { name: /retired-group/ })).toHaveTextContent( + "Existing legacy agent group", + ); + await user.click(await screen.findByRole("option", { name: /Agent One/ })); + expect(onChange).toHaveBeenCalledWith({ agents: ["agent-1"], accessGroups: ["retired-group"] }); + }); + it("respects disabled prop", () => { render(); expect(screen.getByRole("combobox")).toBeDisabled(); diff --git a/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx b/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx index d26772d5c9a..7215bcf25e8 100644 --- a/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx +++ b/ui/litellm-dashboard/src/components/agent_management/AgentSelector.tsx @@ -19,6 +19,7 @@ interface AgentSelectorProps { accessToken: string; placeholder?: string; disabled?: boolean; + allowAccessGroups?: boolean; } const AgentSelector: React.FC = ({ @@ -28,6 +29,7 @@ const AgentSelector: React.FC = ({ accessToken, placeholder = "Select agents", disabled = false, + allowAccessGroups = true, }) => { const [agents, setAgents] = useState([]); const [accessGroups, setAccessGroups] = useState([]); @@ -60,12 +62,15 @@ const AgentSelector: React.FC = ({ fetchData(); }, [accessToken]); - // Combine options, access groups first + const selectableGroups = allowAccessGroups + ? Array.from(new Set([...accessGroups, ...(value?.accessGroups ?? [])])) + : value?.accessGroups ?? []; + const options: MultiSelectOption[] = [ - ...accessGroups.map((group) => ({ + ...selectableGroups.map((group) => ({ label: group, value: `group:${group}`, - description: "Access Group", + description: allowAccessGroups ? "Access Group" : "Existing legacy agent group", })), ...agents.map((agent) => ({ label: `${agent.agent_name || agent.agent_id}`, diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts index 469ecf36f06..92d946c19e1 100644 --- a/ui/litellm-dashboard/src/components/agents/types.ts +++ b/ui/litellm-dashboard/src/components/agents/types.ts @@ -11,6 +11,11 @@ export type AgentKillSwitchConfig = components["schemas"]["AgentKillSwitchConfig export type AgentKillSwitchResult = components["schemas"]["AgentKillSwitchResult"]; export interface Agent { + identity?: components["schemas"]["AgentIdentityBinding"] | null; + identity_managed?: boolean; + enabled?: boolean; + execution_mode?: components["schemas"]["AgentResponse"]["execution_mode"]; + jwt_auth_configured?: boolean; agent_id: string; agent_name: string; litellm_params: { diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx index 47fb5643b22..66748d3cae7 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.test.tsx @@ -170,3 +170,41 @@ describe("MCPServerSelector all-proxy-mcpservers option", () => { expect(optionByLabel("Server One")).toHaveAttribute("aria-disabled", "true"); }); }); + +describe("MCPServerSelector unified group flow", () => { + beforeEach(() => { + vi.clearAllMocks(); + setupMcpMocks(); + mockUseMCPAccessGroups.mockReturnValue({ data: ["legacy-group"], isLoading: false } as ReturnType< + typeof useMCPAccessGroups + >); + }); + + it("offers servers without legacy groups when disabled", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + renderWithProviders(); + await openSelector(user); + expect(optionByLabel("legacy-group")).toBeUndefined(); + await user.click(optionByLabel("Server One")!); + expect(onChange).toHaveBeenCalledWith({ servers: ["srv-1"], accessGroups: [], toolsets: [] }); + }); + + it("preserves a saved legacy group even if discovery no longer returns it", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + renderWithProviders( + , + ); + await openSelector(user); + expect(optionByLabel("retired-group")).toHaveTextContent("Existing legacy MCP group"); + expect(optionByLabel("legacy-group")).toBeUndefined(); + await user.click(optionByLabel("Server One")!); + expect(onChange).toHaveBeenCalledWith({ servers: ["srv-1"], accessGroups: ["retired-group"], toolsets: [] }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx index 969db57462f..fe23fd87b76 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx @@ -18,11 +18,15 @@ interface MCPServerSelectorProps { disabled?: boolean; teamId?: string | null; allowNoMcpServers?: boolean; + allowAccessGroups?: boolean; allowAllProxyMcpServers?: boolean; } const TOOLSET_PREFIX = "toolset:"; +const selectableLegacyGroups = (available: string[], selected: string[] = [], allowNew: boolean): string[] => + allowNew ? Array.from(new Set([...available, ...selected])) : selected; + const MCPServerSelector: React.FC = ({ onChange, value, @@ -32,22 +36,24 @@ const MCPServerSelector: React.FC = ({ disabled = false, teamId, allowNoMcpServers = false, + allowAccessGroups = true, allowAllProxyMcpServers = false, }) => { const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(teamId); const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups(); const { data: toolsets = [], isLoading: toolsetsLoading } = useMCPToolsets(); - const loading = serversLoading || groupsLoading || toolsetsLoading; + const loading = [serversLoading, groupsLoading, toolsetsLoading].some(Boolean); - const accessGroupSet = new Set(accessGroups); + const selectableGroups = selectableLegacyGroups(accessGroups, value?.accessGroups, allowAccessGroups); + const accessGroupSet = new Set(selectableGroups); // Combine options: access groups + servers + toolsets const options = [ - ...accessGroups.map((group) => ({ + ...selectableGroups.map((group) => ({ label: group, value: group, - description: "Access Group", + description: allowAccessGroups ? "Access Group" : "Existing legacy MCP group", })), ...mcpServers.map((server) => ({ label: `${server.server_name || server.server_id} (${server.server_id})`, diff --git a/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx b/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx index d1ca25c975b..0d5525d7d78 100644 --- a/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx +++ b/ui/litellm-dashboard/src/components/permissions/AgentPermissions.tsx @@ -67,7 +67,7 @@ export function AgentPermissions({
-

Agents

+

Allowed agents to call

{totalCount}
diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index a693ee971d4..bbabcf2458e 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -2249,9 +2249,9 @@ describe("TeamInfoView - the exact bytes the update call sends", () => { await openEditorWithAgents(user); await user.click(within(screen.getByLabelText("agent-1")).getByRole("button")); - await user.click(within(screen.getByLabelText("group:group-a")).getByRole("button")); + await user.click(within(screen.getByLabelText("group-a")).getByRole("button")); expect(screen.queryByLabelText("agent-1")).not.toBeInTheDocument(); - expect(screen.queryByLabelText("group:group-a")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("group-a")).not.toBeInTheDocument(); const payload = await save(user); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 3845f94593d..e03a08dfabd 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -2039,13 +2039,14 @@ const TeamInfoView: React.FC = ({ )} - + {({ value, onChange }) => ( )} @@ -2062,13 +2063,14 @@ const TeamInfoView: React.FC = ({ />
- + {({ value, onChange }) => ( )} diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx index 4d451a9c7b9..4d980d1b018 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx @@ -375,3 +375,17 @@ describe("TTFT column", () => { expect(screen.getByText("1.00")).toBeInTheDocument(); }); }); + +describe("Request outcome", () => { + it("shows a failed agent outcome even when metadata has no status", () => { + renderRows([logEntry({ call_type: "asend_message", status: "failure", session_total_count: 4 })]); + expect(screen.getByText("Failure")).toBeInTheDocument(); + expect(screen.queryByText("Success")).not.toBeInTheDocument(); + }); + + it("prefers the recorded outcome over stale metadata", () => { + renderRows([logEntry({ status: "success", metadata: { status: "failure" } })]); + expect(screen.getByText("Success")).toBeInTheDocument(); + expect(screen.queryByText("Failure")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx index 5705f41f3de..6b35ceb81a2 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx @@ -120,7 +120,7 @@ export const getRequestLogsTableColumns = ({ enableSorting: false, meta: { skeleton: "badge" }, cell: ({ row }) => { - const status = readMetaString(row.original.metadata, "status") ?? "Success"; + const status = row.original.status || readMetaString(row.original.metadata, "status") || "Success"; const isSuccess = status.toLowerCase() !== "failure"; const batchCounts = isSuccess ? getBatchRequestCounts(row.original.metadata) : undefined; if (batchCounts && batchCounts.failed > 0) { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b5d515bd1fb..51d0165330b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -18071,6 +18071,23 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/agents/identity/providers": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Agent Identity Providers */ + get: operations["get_agent_identity_providers_v1_agents_identity_providers_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/agents/make_public": { parameters: { query?: never; @@ -18220,6 +18237,23 @@ export interface paths { patch: operations["patch_agent_v1_agents__agent_id__patch"]; trace?: never; }; + "/v1/agents/{agent_id}/identity": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Agent Identity Status */ + get: operations["get_agent_identity_status_v1_agents__agent_id__identity_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/agents/{agent_id}/kill_switch": { parameters: { query?: never; @@ -24206,11 +24240,19 @@ export interface components { AgentConfig: { /** Access Group Ids */ access_group_ids?: string[] | null; - agent_card_params: components["schemas"]["AgentCard"]; + agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name: string; + /** Enabled */ + enabled?: boolean; + /** + * Execution Mode + * @enum {string} + */ + execution_mode?: "autonomous" | "delegated" | "both"; /** Extra Headers */ extra_headers?: string[] | null; + identity?: components["schemas"]["EntraIdentityConfig"] | null; kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null; /** Litellm Params */ litellm_params?: { @@ -24297,6 +24339,45 @@ export interface components { /** Uri */ uri?: string; }; + /** AgentIdentityBinding */ + AgentIdentityBinding: { + /** + * Active + * @default true + */ + active: boolean; + /** Agent Id */ + agent_id: string; + /** Client Id */ + client_id: string; + /** Issuer */ + issuer: string; + /** Last Authenticated At */ + last_authenticated_at?: string | null; + /** + * Provider + * @constant + */ + provider: "microsoft_entra"; + /** + * Required Roles + * @default [] + */ + required_roles: string[]; + /** + * Required Scopes + * @default [ + * "user_impersonation" + * ] + */ + required_scopes: string[]; + /** Revision */ + revision: string; + /** Service Principal Id */ + service_principal_id?: string | null; + /** Tenant Id */ + tenant_id: string; + }; /** * AgentInterface * @description Declares a combination of a target URL and a transport protocol. @@ -24452,8 +24533,30 @@ export interface components { created_at?: string | null; /** Created By */ created_by?: string | null; + /** + * Enabled + * @default true + */ + enabled: boolean; + /** + * Execution Mode + * @default autonomous + * @enum {string} + */ + execution_mode: "autonomous" | "delegated" | "both"; /** Extra Headers */ extra_headers?: string[] | null; + identity?: components["schemas"]["AgentIdentityBinding"] | null; + /** + * Identity Managed + * @default false + */ + identity_managed: boolean; + /** + * Jwt Auth Configured + * @default false + */ + jwt_auth_configured: boolean; /** Keys */ keys?: components["schemas"]["AgentKeySummary"][] | null; kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null; @@ -30137,6 +30240,33 @@ export interface components { /** Template Id */ template_id: string; }; + /** EntraIdentityConfig */ + EntraIdentityConfig: { + /** Client Id */ + client_id: string; + /** + * Provider + * @constant + */ + provider: "microsoft_entra"; + /** + * Required Roles + * @default [] + */ + required_roles: string[]; + /** + * Required Scopes + * @description Required delegated scopes. An empty list accepts any nonempty scope granted for this gateway. + * @default [ + * "user_impersonation" + * ] + */ + required_scopes: string[]; + /** Service Principal Id */ + service_principal_id?: string | null; + /** Tenant Id */ + tenant_id: string; + }; /** EnvironmentReport */ EnvironmentReport: { /** Config Lines */ @@ -35544,6 +35674,28 @@ export interface components { /** Mcp Server Ids */ mcp_server_ids: string[]; }; + /** ManagedAgentIdentityStatus */ + ManagedAgentIdentityStatus: { + /** + * Enabled + * @default true + */ + enabled: boolean; + /** + * Execution Mode + * @default autonomous + * @enum {string} + */ + execution_mode: "autonomous" | "delegated" | "both"; + identity?: components["schemas"]["AgentIdentityBinding"] | null; + /** + * Identity Managed + * @default false + */ + identity_managed: boolean; + /** Last Authenticated At */ + last_authenticated_at?: string | null; + }; /** * Mcp * @description Give the model access to additional tools via remote Model Context Protocol @@ -37618,8 +37770,16 @@ export interface components { agent_card_params?: components["schemas"]["AgentCard"]; /** Agent Name */ agent_name?: string; + /** Enabled */ + enabled?: boolean; + /** + * Execution Mode + * @enum {string} + */ + execution_mode?: "autonomous" | "delegated" | "both"; /** Extra Headers */ extra_headers?: string[] | null; + identity?: components["schemas"]["EntraIdentityConfig"] | null; kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null; /** Litellm Params */ litellm_params?: { @@ -70077,6 +70237,26 @@ export interface operations { }; }; }; + get_agent_identity_providers_v1_agents_identity_providers_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": string[]; + }; + }; + }; + }; make_agents_public_v1_agents_make_public_post: { parameters: { query?: never; @@ -70242,6 +70422,37 @@ export interface operations { }; }; }; + get_agent_identity_status_v1_agents__agent_id__identity_get: { + parameters: { + query?: never; + header?: never; + path: { + agent_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ManagedAgentIdentityStatus"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post: { parameters: { query?: never;