feat(agents): authoritative permissions

This commit is contained in:
Joshua Valluru 2026-09-28 16:23:54 -07:00
parent 6684256136
commit 18576ee4f2
17 changed files with 1498 additions and 205 deletions

View file

@ -0,0 +1,67 @@
from types import MappingProxyType
from typing import Final
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.types.proxy.agent_identity import AgentIdentityFailure
async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
agent: Final = auth.managed_agent_policy
if agent is None:
return ()
try:
base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
ceilings: Final = await resolve_managed_agent_ceilings(agent)
expanded: Final = tuple(
frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
for ceiling in ceilings
)
own: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return tuple(sorted(own))
if context.user_id is None:
return ()
human: Final = await _delegated_resource_subject(context.user_id)
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human)
return tuple(sorted(own.intersection(allowed)))
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
)
async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
if server_id not in await managed_agent_servers(auth):
return []
try:
own: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return own
if context.user_id is None:
return []
human: Final = await _delegated_resource_subject(context.user_id)
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human)
if own is None:
return human_tools
return own if human_tools is None else sorted(frozenset(own).intersection(human_tools))
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
)

View file

@ -67,7 +67,6 @@ from litellm.repositories.table_repositories import (
AgentsRepository,
MCPServerRepository,
)
from litellm.repositories.user_repository import UserRepository
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
@ -1086,7 +1085,7 @@ class MCPRequestHandler:
assert_never(identity.subject_type)
@staticmethod
async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
@ -1111,6 +1110,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=requires_fresh_policy,
)
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
@ -1119,6 +1119,7 @@ class MCPRequestHandler:
if user_object is not None and object_permission is None and user_object.object_permission_id:
object_permission = await get_object_permission(
object_permission_id=user_object.object_permission_id,
check_db_only=requires_fresh_policy,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
@ -1147,6 +1148,7 @@ class MCPRequestHandler:
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
admitted.mcp_admitted_user_subject = True
admitted.requires_fresh_policy = requires_fresh_policy
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
# several teams under its own identity, so without this a cross-team user outruns every team's
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
@ -1597,6 +1599,11 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
try:
@ -1606,7 +1613,7 @@ class MCPRequestHandler:
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return MCPServerAccess(
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
)
# Get allowed servers from key and team
@ -1703,7 +1710,7 @@ class MCPRequestHandler:
if user_api_key_auth and user_api_key_auth.agent_id:
agent_capped: Final = _agent_capped_servers(
allowed_mcp_servers,
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
)
if agent_capped is not None:
@ -1829,10 +1836,12 @@ class MCPRequestHandler:
scoped.object_permission = auth.object_permission
scoped.object_permission_id = auth.object_permission_id
scoped.access_group_ids = auth.access_group_ids
scoped.requires_fresh_policy = auth.requires_fresh_policy
scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
return scoped
@staticmethod
async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
direct grants, plus every team they are a live roster member of.
@ -1886,6 +1895,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(auth and auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
@ -1941,11 +1951,11 @@ class MCPRequestHandler:
roster instead of by grant charged unrelated teams' buckets)."""
return [
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
for source in await MCPRequestHandler._admitted_subject_sources(auth)
for source in await MCPRequestHandler.admitted_subject_sources(auth)
]
@staticmethod
async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
"""Union of what each of the admitted subject's sources reaches, each answered by the
canonical resolver so no rule is reimplemented for this caller shape."""
reachable: Final[set[str]] = set()
@ -2007,7 +2017,7 @@ class MCPRequestHandler:
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
@staticmethod
async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
sources that actually grant that server.
@ -2088,6 +2098,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not team_obj:
@ -2098,6 +2109,8 @@ class MCPRequestHandler:
@staticmethod
async def _toolset_tool_permissions(
object_permission: LiteLLM_ObjectPermissionTable | None,
*,
requires_fresh_policy: bool = False,
) -> Mapping[str, Sequence[str]]:
"""The ``server_id -> tool names`` grants of this permission row's toolsets, empty when it
declares none. The shared resolver for the team, org, and internal-user levels, so a toolset
@ -2114,7 +2127,8 @@ class MCPRequestHandler:
if object_permission is None or not object_permission.mcp_toolsets:
return _EMPTY_TOOLSET_GRANTS
resolved: Final = await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=object_permission.mcp_toolsets
toolset_ids=object_permission.mcp_toolsets,
requires_fresh_policy=requires_fresh_policy,
)
if not resolved:
raise UnloadableEntitlementError(
@ -2126,10 +2140,15 @@ class MCPRequestHandler:
async def _toolset_tools_for_server(
object_permission: LiteLLM_ObjectPermissionTable | None,
server_id: str,
*,
requires_fresh_policy: bool = False,
) -> Sequence[str] | None:
"""Tool names this row's toolsets grant on ``server_id``, ``None`` when its toolsets place
no restriction on that server (it declares no toolsets, or none of them name it)."""
return (await MCPRequestHandler._toolset_tool_permissions(object_permission)).get(server_id)
grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permission, requires_fresh_policy=requires_fresh_policy
)
return grants.get(server_id)
@staticmethod
def _union_tool_grants(
@ -2171,6 +2190,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 +2239,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)
@ -2249,9 +2274,12 @@ class MCPRequestHandler:
# tool-level check sees the key's full effective tool scope
key_toolset_ids: Final = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
key_toolset_tools: Final = (
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
server_id
)
(
await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=key_toolset_ids,
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
).get(server_id)
if key_toolset_ids
else None
)
@ -2265,7 +2293,9 @@ class MCPRequestHandler:
# Tools granted through the team's toolsets restrict this server exactly
# as the team's direct tool permissions do, mirroring the key path above
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
team_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
# Apply same inheritance logic as get_allowed_mcp_servers
@ -2334,7 +2364,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,
@ -2365,7 +2395,9 @@ class MCPRequestHandler:
if org_obj_perm and org_obj_perm.mcp_tool_permissions
else None
)
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(org_obj_perm, server_id)
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
org_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
org_tools: Final = MCPRequestHandler._union_tool_grants(org_direct_tools, org_toolset_tools)
if org_tools is not None:
allowed_tools = (
@ -2456,6 +2488,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 +2535,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 []
@ -2518,7 +2552,8 @@ class MCPRequestHandler:
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
key_object_permission.mcp_access_groups or []
key_object_permission.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
# servers referenced in tool permissions should also be accessible
@ -2531,7 +2566,14 @@ class MCPRequestHandler:
# ceilings as any other key-level grant
toolset_ids: Final = key_object_permission.mcp_toolsets or []
toolset_servers: Final = (
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
list(
(
await global_mcp_server_manager.resolve_toolset_tool_permissions(
toolset_ids=toolset_ids,
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
).keys()
)
if toolset_ids
else []
)
@ -2550,7 +2592,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 +2610,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 +2638,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)
@ -2605,7 +2648,12 @@ class MCPRequestHandler:
return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID))
@staticmethod
async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]:
async def _team_granted_servers(
team_obj: LiteLLM_TeamTable,
team_access_group_servers: list[str],
*,
requires_fresh_policy: bool = False,
) -> set[str]:
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
@ -2620,13 +2668,17 @@ class MCPRequestHandler:
if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []):
return set(global_mcp_server_manager.get_registry().keys())
legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=requires_fresh_policy,
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=requires_fresh_policy
)
return (
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
| set(legacy_access_group_servers)
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
| toolset_grants.keys()
| set(team_access_group_servers)
)
@ -2667,6 +2719,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,9 +2733,14 @@ 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)
servers: Final = await MCPRequestHandler._team_granted_servers(
team_obj,
team_access_group_servers,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
return list(servers)
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
@ -2716,6 +2774,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
@ -2811,7 +2870,8 @@ class MCPRequestHandler:
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
tool_perm_servers: Final = list(
@ -2820,7 +2880,10 @@ class MCPRequestHandler:
# servers referenced by the org's toolset grants are part of the org ceiling,
# exactly as servers referenced by its inline tool permissions are
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
all_servers: Final = tuple(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
@ -2912,7 +2975,8 @@ class MCPRequestHandler:
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permission.mcp_access_groups or []
object_permission.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
# servers referenced in tool permissions should also be accessible
@ -2961,7 +3025,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 +3037,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 +3046,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 +3057,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 +3082,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
@ -3031,13 +3109,17 @@ class MCPRequestHandler:
return []
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
object_permissions.mcp_access_groups or [],
requires_fresh_policy=fresh,
)
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=fresh
)
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
@ -3119,9 +3201,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 = (
@ -3167,7 +3253,11 @@ class MCPRequestHandler:
user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
object_permissions,
server_id,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
user_tools: Final = MCPRequestHandler._union_tool_grants(user_direct_tools, user_toolset_tools)
if user_tools is None:
return allowed_tools
@ -3196,7 +3286,9 @@ class MCPRequestHandler:
return allowed_tools
try:
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
team_obj_perm, server_id, requires_fresh_policy=caller_auth.requires_fresh_policy
)
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
verbose_logger.warning(
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
@ -3241,7 +3333,11 @@ class MCPRequestHandler:
end_user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
object_permissions,
server_id,
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
end_user_tools: Final = MCPRequestHandler._union_tool_grants(end_user_direct_tools, end_user_toolset_tools)
if end_user_tools is None:
return allowed_tools
@ -3302,6 +3398,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 +3419,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]:
@ -3358,12 +3458,15 @@ class MCPRequestHandler:
obj_perm.mcp_servers or []
)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
obj_perm.mcp_access_groups or []
obj_perm.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
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 +3493,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,
@ -3430,11 +3533,13 @@ class MCPRequestHandler:
if obj_perm.mcp_tool_permissions
else None
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
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
@ -3452,28 +3557,38 @@ class MCPRequestHandler:
return server_ids
@staticmethod
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
async def _get_db_server_ids_for_access_groups(
prisma_client,
access_groups: list[str],
*,
use_writer: bool = False,
) -> set[str]:
"""
Helper to get server_ids from DB servers that match any of the given access groups.
"""
server_ids: Final[set[str]] = set()
if access_groups and prisma_client is not None:
try:
mcp_servers: Final = await MCPServerRepository(prisma_client).table.find_many(
mcp_servers: Final = await MCPServerRepository(prisma_client, use_writer=use_writer).table.find_many(
where={"mcp_access_groups": {"hasSome": access_groups}}
)
for server in mcp_servers:
server_ids.add(server.server_id)
except Exception as e:
if use_writer:
raise
verbose_logger.debug("Error getting MCP servers from access groups: %s", e)
return server_ids
@staticmethod
async def _get_mcp_servers_from_access_groups(
access_groups: list[str],
*,
requires_fresh_policy: bool = False,
) -> list[str]:
"""
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
"""
from litellm.proxy.proxy_server import prisma_client
@ -3489,11 +3604,15 @@ class MCPRequestHandler:
)
# Use the new helper for DB servers
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(prisma_client, access_groups)
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
prisma_client, access_groups, use_writer=requires_fresh_policy
)
server_ids.update(db_server_ids)
return list(server_ids)
except Exception as e:
if requires_fresh_policy:
raise
verbose_logger.warning("Failed to get MCP servers from access groups: %s", e)
return []
@ -3548,6 +3667,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 +3711,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
verbose_logger.debug("team_obj is None")

View file

@ -3428,7 +3428,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()
@ -3477,9 +3479,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
@ -3511,7 +3518,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)
)
@ -3580,12 +3587,14 @@ 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(
self,
toolset_ids: list[str],
*,
requires_fresh_policy: bool = False,
) -> dict[str, list[str]]:
"""
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
@ -3595,6 +3604,10 @@ class MCPServerManager:
Redis-backed ``DualCache`` in production) so that cache entries are
shared across workers and cold-cache DB hits are minimised.
``requires_fresh_policy`` bypasses the cache and reads the writer so a
revocation is honoured on the very next request; a read fault then
propagates instead of resolving to no grants.
A row names a tool on the server identified by ``server_id``, so the
stored name is the tool's own name and is used as written. It is never
reduced by the server's wire prefix: that prefix is added on the way out
@ -3609,12 +3622,16 @@ class MCPServerManager:
return {}
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final[dict[str, list[str]] | None] = (
None if requires_fresh_policy else await user_api_key_cache.async_get_cache(key=cache_key)
)
if cached is not None:
return cached
try:
toolsets: Final = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
toolsets: Final = await list_mcp_toolsets(
prisma_client, toolset_ids=toolset_ids, use_writer=requires_fresh_policy
)
tool_permissions: Final[dict[str, list[str]]] = {}
for toolset in toolsets:
for tool in toolset.tools:
@ -3628,6 +3645,8 @@ class MCPServerManager:
)
return tool_permissions
except Exception as e:
if requires_fresh_policy:
raise
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
return {}

View file

@ -65,9 +65,9 @@ class MCPToolsetTable(Protocol):
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
def _toolset_table(prisma_client: PrismaClient, *, use_writer: bool = False) -> MCPToolsetTable:
"""The toolset table actions of the prisma client."""
return MCPToolsetRepository(prisma_client).table
return MCPToolsetRepository(prisma_client, use_writer=use_writer).table
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
@ -107,12 +107,16 @@ async def get_mcp_toolset(
async def list_mcp_toolsets(
prisma_client: PrismaClient,
toolset_ids: Sequence[str] | None = None,
*,
use_writer: bool = False,
) -> Sequence[MCPToolset]:
try:
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
rows: Final = await _toolset_table(prisma_client, use_writer=use_writer).find_many(where=where)
return [_toolset_from_row(r) for r in rows]
except Exception as e:
if use_writer:
raise
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
return []

View file

@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=user_api_key_auth.requires_fresh_policy,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey
)
try:
admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id)
admitted: Final = await MCPRequestHandler.reload_admitted_user(
user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
except HTTPException as e:
verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail)
return None

View file

@ -1,13 +1,16 @@
import asyncio
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Final, TypeAlias
from typing import TYPE_CHECKING, Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_AccessGroupTable
if TYPE_CHECKING:
from litellm.types.agents import AgentResponse
AccessGroupIds: TypeAlias = tuple[str, ...]
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ()) if agent is not None else ()
async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
from litellm.proxy.auth.auth_checks import get_access_object
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
@ -47,6 +50,7 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
except HTTPException as e:
verbose_proxy_logger.warning(
@ -73,3 +77,16 @@ async def resolve_agent_access_group_ceiling(
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
)
async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
async def authoritative_group(group_id: str) -> LoadedAccessGroup:
return await _load_access_group(group_id, check_db_only=True)
async def manual_ids(_agent_id: str) -> AccessGroupIds:
return tuple(agent.access_group_ids or ())
manual: Final = await resolve_agent_access_group_ceiling(
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
)
return (manual,) if manual is not None else ()

View file

@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
@ -86,7 +89,9 @@ class AgentRequestHandler:
) -> AgentAccess:
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
return await _managed_actor_agent_access(user_api_key_auth)
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(user_api_key_auth)
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
@ -104,13 +109,19 @@ class AgentRequestHandler:
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
@staticmethod
async def _resolve_key_team_agent_access(
async def resolve_key_team_agent_access(
user_api_key_auth: UserAPIKeyAuth | None,
*,
strict: bool = False,
) -> AgentAccess:
try:
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
user_api_key_auth, strict=strict
)
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents: %s", e)
return UnrestrictedAgentAccess()
return _intersect_agent_access(key_access, team_access)
@ -144,6 +155,33 @@ class AgentRequestHandler:
bool: True if agent is allowed, False otherwise
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.proxy_server import prisma_client
from litellm.types.proxy.agent_identity import AgentIdentityFailure
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
if prisma_client is not None or (registered is not None and registered.identity_managed):
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
if isinstance(target, AgentIdentityFailure):
raise_identity_failure(target)
if target is None and registered is not None and registered.identity_managed:
return False
if target is not None and target.identity_managed:
if (
not target.enabled
or target.identity is None
or not target.identity.active
or user_api_key_auth is None
):
return False
fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True})
explicit: Final = await _granted_agent_ids(
fresh_auth,
_strict_agent_access,
build_effective_auth_contexts,
)
return target.agent_id in explicit
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
case UnrestrictedAgentAccess():
@ -202,8 +240,10 @@ class AgentRequestHandler:
return team_obj.object_permission
@staticmethod
async def _get_allowed_agents_for_key(
async def get_allowed_agents_for_key(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a key.
@ -237,24 +277,36 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(key_access_group_ids), check_db_only=strict
)
)
if key_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
return UnrestrictedAgentAccess()
@staticmethod
async def _get_allowed_agents_for_team(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a team.
@ -263,7 +315,7 @@ class AgentRequestHandler:
2. Also includes agents from team's access_group_ids (unified access groups)
Fetches the team object once and reuses it for both permission sources.
Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`.
Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`.
"""
if user_api_key_auth is None:
return UnrestrictedAgentAccess()
@ -280,7 +332,7 @@ class AgentRequestHandler:
)
if not prisma_client:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# Fetch the team object once for both permission sources
team_obj: Final = await get_team_object(
@ -289,10 +341,11 @@ class AgentRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=strict,
)
if team_obj is None:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# 1. Get agents from object_permission (native permissions)
object_permissions: Final = team_obj.object_permission
@ -307,18 +360,28 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(team_access_group_ids), check_db_only=strict
)
)
if team_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
# litellm-dashboard is the default UI team and will never have agents;
# skip noisy warnings for it.
if user_api_key_auth.team_id != UI_TEAM_ID:
@ -326,7 +389,9 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
@staticmethod
def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
def _get_config_agent_ids_for_access_groups(
config_agents: Sequence[AgentResponse], access_groups: list[str]
) -> set[str]:
"""
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
"""
@ -339,7 +404,9 @@ class AgentRequestHandler:
return server_ids
@staticmethod
async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
async def _get_db_agent_ids_for_access_groups(
prisma_client, access_groups: list[str], *, check_db_only: bool = False
) -> set[str]:
"""
Helper to get agent_ids from DB agents that match any of the given access groups.
@ -349,23 +416,27 @@ class AgentRequestHandler:
if not access_groups or prisma_client is None:
return set()
agents: Final = await AgentsRepository(prisma_client).table.find_many(
agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
where={"agent_access_groups": {"hasSome": access_groups}}
)
return {agent.agent_id for agent in agents}
@staticmethod
async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
async def _get_unified_access_group_agents(
access_group_ids: list[str], *, check_db_only: bool = False
) -> list[str]:
"""
Resolve unified access group ids to agent IDs.
"""
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
@staticmethod
async def _get_agents_from_access_groups(
access_groups: list[str],
*,
check_db_only: bool = False,
) -> list[str]:
"""
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
@ -373,14 +444,13 @@ class AgentRequestHandler:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.proxy_server import prisma_client
# Use the helper for config-loaded agents
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
global_agent_registry.agent_list, access_groups
)
# Use the helper for DB agents
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
prisma_client, access_groups
prisma_client, access_groups, check_db_only=check_db_only
)
return list(config_agent_ids | db_agent_ids)
@ -531,4 +601,58 @@ async def accessible_agents(
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
effective_contexts,
)
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
allowed: Final = await asyncio.gather(
*(
AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
for agent in agents
if agent.identity_managed
)
)
managed_ids: Final = frozenset(
agent.agent_id
for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
if permitted
)
return tuple(
agent
for agent in agents
if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
)
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
if auth.managed_agent_policy is not None:
return await _managed_actor_agent_access(auth)
return await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
agent: Final = auth.managed_agent_policy
if agent is None or not agent.object_permission:
return RestrictedAgentAccess(frozenset())
permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
own_auth: Final = UserAPIKeyAuth(object_permission=permission)
own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True))
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
ceilings: Final = await resolve_managed_agent_ceilings(agent)
capped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return RestrictedAgentAccess(capped)
if context.user_id is None:
return RestrictedAgentAccess(frozenset())
human_ids: Final = await verified_human_agent_grants(context.user_id)
return RestrictedAgentAccess(capped.intersection(human_ids))
async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
if user_id is None:
return frozenset()
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
sources: Final = await MCPRequestHandler.admitted_subject_sources(human)
human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
return frozenset().union(*(_granted_ids(access) for access in human_access))

View file

@ -1057,6 +1057,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,
@ -2642,7 +2657,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}
)
@ -2680,7 +2695,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},
)
@ -3126,9 +3141,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
)
@ -3162,6 +3177,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(
@ -3170,7 +3186,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:
@ -3192,8 +3210,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,
@ -3283,6 +3304,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
@ -3328,16 +3350,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)
"""
@ -3346,18 +3367,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(
@ -3941,6 +3963,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
@ -3952,9 +3975,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
@ -3962,10 +3989,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())
@ -3978,6 +4007,8 @@ async def get_object_permission(
return _perm_obj
except Exception:
if check_db_only:
raise
return None
@ -4187,6 +4218,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
@ -4229,6 +4261,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:
@ -4264,6 +4297,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.
@ -4275,6 +4309,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,
)
@ -4283,6 +4318,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.
@ -4294,6 +4330,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,
)
@ -4493,26 +4530,36 @@ async def _check_agent_access_group_model_access(
"""Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows."""
if not model or valid_token is None or not valid_token.agent_id:
return True
ceiling: Final = await resolve_ceiling(valid_token.agent_id)
if ceiling is None:
return True
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
return _can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if valid_token.managed_agent_policy is None else None
ceilings: Final = (
await resolve_managed_agent_ceilings(valid_token.managed_agent_policy)
if valid_token.managed_agent_policy is not None
else (unmanaged,)
if unmanaged is not None
else ()
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
for ceiling in ceilings:
if not ceiling.models:
raise ModelAccessDeniedProxyException(
message=model_access_denied_client_message(model=model),
internal_message=f"agent {valid_token.agent_id} access groups grant no models",
type=ProxyErrorTypes.agent_model_access_denied,
param="model",
code=status.HTTP_403_FORBIDDEN,
)
_can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
return True
LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None

View file

@ -15,9 +15,14 @@ if TYPE_CHECKING:
class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
"""Repository for object permission database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
return self.prisma_client.db.litellm_objectpermissiontable
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
return database.litellm_objectpermissiontable
@property
def model_class(self) -> type[LiteLLM_ObjectPermissionTable]:

View file

@ -70,6 +70,9 @@ class _PrismaClientView(Protocol):
@property
def db(self) -> _PrismaTeamDb: ...
@property
def writer_db(self) -> _PrismaTeamDb: ...
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
_JSON_ENCODED_TEAM_FIELDS: Final = (
@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = (
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
"""Repository for team database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def _db(self) -> _PrismaTeamDb:
client: Final[_PrismaClientView] = self.prisma_client
return client.db
return client.writer_db if self._use_writer else client.db
@property
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:

View file

@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
class UserRepository(BaseRepository[LiteLLM_UserTable]):
"""Repository for user database operations."""
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
super().__init__(prisma_client)
self._use_writer = use_writer
@property
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
return self.prisma_client.db.litellm_usertable
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
return database.litellm_usertable
@property
def model_class(self) -> type[LiteLLM_UserTable]:

View file

@ -0,0 +1,367 @@
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()
def _server_row(server_id: str, access_groups: tuple[str, ...]) -> MagicMock:
row: Final = MagicMock()
row.server_id = server_id
row.mcp_access_groups = list(access_groups)
return row
def _toolset_row(server_id: str, tool_name: str) -> MagicMock:
row: Final = MagicMock()
row.tools = [{"server_id": server_id, "tool_name": tool_name}]
return row
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["tool", "server", "outage"])
async def test_autonomous_agent_toolset_and_access_group_revocations_bind_on_the_next_request(
monkeypatch: pytest.MonkeyPatch, change: str
) -> None:
"""The agent's entitlements are read through the shared toolset and access-group resolvers. Once the
writer revokes a tool or drops the server from the group, the next managed request must be denied
even though the legacy cache still holds the warm grant and the replica still shows the old rows"""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server import toolset_db
warm_toolset: Final = _toolset_row("slack", "read")
list_toolsets: Final = AsyncMock(return_value=[warm_toolset])
monkeypatch.setattr(toolset_db, "list_mcp_toolsets", list_toolsets)
client: Final = MagicMock()
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="agent-permissions", mcp_toolsets=["ts"], mcp_access_groups=["grp"]
)
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
update={"object_permission": permission.model_dump()}
)
auth.requires_fresh_policy = True
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
if change == "tool":
list_toolsets.return_value = [_toolset_row("slack", "other")]
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["other"]
elif change == "server":
client.writer_db.litellm_mcpservertable.find_many.return_value = []
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
else:
list_toolsets.side_effect = RuntimeError("writer unavailable")
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
assert failure.value.status_code == 503
for call in list_toolsets.await_args_list:
assert call.kwargs["use_writer"] is True, "managed agent toolsets must be read from the writer"
client.db.litellm_mcpservertable.find_many.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
@pytest.mark.parametrize("has_grant", [True, False])
@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
monkeypatch: pytest.MonkeyPatch,
role: str,
open_channel: str,
has_grant: bool,
agent_tools: tuple[str, ...] | None,
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {
name: MCPServer(
server_id=name,
name=name,
transport="http",
url="https://example.com/mcp",
allow_all_keys=open_channel == "operator",
)
for name in ("slack", "linear")
}
from litellm.proxy._experimental.mcp_server import db
monkeypatch.setattr(
db,
"get_active_submitted_mcp_server_ids_for_user",
AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
)
user: Final = LiteLLM_UserTable(
user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
)
team: Final = LiteLLM_TeamTable(
team_id="team",
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission_id="team-grant",
)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
monkeypatch.setattr(proxy_server, "prisma_client", client)
auth: Final = actor(agent_tools, delegated=True)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
assert admitted.user_role == role
@pytest.mark.asyncio
async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy._experimental.mcp_server import db
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
auth: Final = UserAPIKeyAuth(user_id="human")
auth.mcp_explicit_grants_only = True
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
assert await manager.get_allowed_mcp_servers(auth) == []
auth.mcp_explicit_grants_only = False
assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
@pytest.mark.asyncio
async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
assert await managed_agent_servers(UserAPIKeyAuth()) == ()
auth: Final = actor(None, delegated=True)
assert auth.managed_agent_context is not None
auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(
auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")])
)
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
assert failure.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
@pytest.mark.parametrize("scoped", (False, True))
async def test_manager_preserves_managed_server_grants_across_open_channels(
monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
) -> None:
from litellm.proxy._experimental.mcp_server import db
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {
"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
"submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
"passthrough": MCPServer(
server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
),
}
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
auth: Final = actor(None)
auth.user_role = role
assert not auth.mcp_explicit_grants_only
access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == (
{"slack"} if scoped else {"slack", "linear"}
)
@pytest.mark.asyncio
async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
with pytest.raises(HTTPException) as failure:
await manager.get_allowed_mcp_servers(actor(None, delegated=True))
assert failure.value.status_code == 503

View file

@ -369,7 +369,9 @@ class TestMCPRequestHandler:
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
assert result == ["server-a"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(
toolset_ids=["toolset-1"], requires_fresh_policy=False
)
async def test_get_allowed_mcp_servers_for_key_skips_toolset_resolution_when_none_granted(self):
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
@ -4147,7 +4149,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
"group-server2",
}
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
mock_get_access_group_servers.assert_called_once_with(["dev-group"], requires_fresh_policy=False)
finally:
for sid in ("direct-server1", "direct-server2"):
global_mcp_server_manager.registry.pop(sid, None)
@ -4316,7 +4318,7 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
assert set(result) == {"direct-server", "group-server"}
mock_get_perm.assert_not_called()
mock_access_groups.assert_called_once_with(["grp-alpha"])
mock_access_groups.assert_called_once_with(["grp-alpha"], requires_fresh_policy=False)
finally:
global_mcp_server_manager.registry.pop("direct-server", None)
@ -4383,7 +4385,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 +4404,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 +4423,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 +4540,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 +4557,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 +4613,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 +4639,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 +4671,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,10 +4720,12 @@ 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"])
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(
toolset_ids=["toolset-1"], requires_fresh_policy=False
)
async def test_get_allowed_mcp_servers_toolset_only_agent_caps_key_servers(self):
"""Regression: an agent whose only grant is a toolset used to resolve to [] and place
@ -4760,7 +4764,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 +4793,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 +5837,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 +5854,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 +8312,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 +8875,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 +9680,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 +9698,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 +9725,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 +9744,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 +9758,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 +9775,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 +10095,47 @@ class TestScopedSessionAdmission:
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None
@pytest.mark.asyncio
async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch):
from litellm.caching.dual_cache import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_UserTable
cached = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="revoked")
current = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="current")
cache = DualCache()
await cache.async_set_cache(key="fresh-human", value=cached)
database = MagicMock()
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=current)
database.db.litellm_usertable.find_unique = AsyncMock(return_value=cached)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
assert await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) == "current"
database.db.litellm_usertable.find_unique.assert_not_awaited()
database.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("unavailable")
with pytest.raises(HTTPException) as denied:
await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True)
assert denied.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ["servers", "tools"])
async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted_grant(monkeypatch, operation):
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.types.agents import AgentResponse
auth = UserAPIKeyAuth(agent_id="managed")
auth.managed_agent_policy = AgentResponse(agent_id="managed", agent_name="Managed", agent_card_params={})
permission = LiteLLM_ObjectPermissionTable(object_permission_id="policy", mcp_toolsets=["unavailable"])
manager = MagicMock()
manager.expand_permission_list.return_value = []
manager.resolve_toolset_tool_permissions = AsyncMock(side_effect=RuntimeError("policy unavailable"))
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
resolution = (
MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth, permission)
if operation == "servers"
else MCPRequestHandler.get_agent_tool_permissions_for_server("slack", auth, permission)
)
with pytest.raises(RuntimeError, match="policy unavailable"):
await resolution

View file

@ -211,15 +211,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
@ -5585,9 +5579,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
@ -5654,9 +5646,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
@ -5723,9 +5713,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
@ -5760,9 +5748,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
@ -6838,9 +6824,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
@ -6924,9 +6908,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()
@ -11175,6 +11157,72 @@ async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks():
list_toolsets_mock.assert_awaited_once()
@pytest.mark.asyncio
async def test_resolve_toolset_tool_permissions_fresh_policy_sees_writer_revocation_past_warm_cache():
"""A managed agent's tool grant revoked in the writer DB must be gone on the very next fresh
request even though the legacy cache still holds the old grant, and the fresh read must go to
the writer, not the replica"""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
manager = MCPServerManager()
granted = MagicMock()
granted.tools = [{"server_id": "server-a", "tool_name": "echo"}]
revoked = MagicMock()
revoked.tools = [{"server_id": "server-a", "tool_name": "other"}]
list_toolsets_mock = AsyncMock(side_effect=[[granted], [revoked]])
with (
patch(
"litellm.proxy._experimental.mcp_server.toolset_db.list_mcp_toolsets",
list_toolsets_mock,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
):
warm = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
legacy_after_revoke = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
fresh_after_revoke = await manager.resolve_toolset_tool_permissions(
toolset_ids=["ts-1"], requires_fresh_policy=True
)
assert warm == {"server-a": ["echo"]}
assert legacy_after_revoke == warm, "legacy callers keep the cached grant by design"
assert fresh_after_revoke == {"server-a": ["other"]}
assert list_toolsets_mock.await_count == 2
assert list_toolsets_mock.await_args_list[0].kwargs["use_writer"] is False
assert list_toolsets_mock.await_args_list[1].kwargs["use_writer"] is True
@pytest.mark.asyncio
async def test_resolve_toolset_tool_permissions_fresh_policy_propagates_db_fault_instead_of_no_grants():
"""A fresh read that fails must raise so the managed-agent boundary fails closed; the legacy
path keeps its swallow-to-empty behaviour"""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
manager = MCPServerManager()
list_toolsets_mock = AsyncMock(side_effect=RuntimeError("relation does not exist"))
with (
patch(
"litellm.proxy._experimental.mcp_server.toolset_db.list_mcp_toolsets",
list_toolsets_mock,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
):
legacy = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
with pytest.raises(RuntimeError, match="relation does not exist"):
await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"], requires_fresh_policy=True)
assert legacy == {}
class TestMaterializeAuthHeaders:
"""_materialize_auth_headers drives one step of a resolved httpx.Auth's own flow to turn it
into a header dict for the OpenAPI egress arm, which sends plain headers and cannot carry an
@ -11496,12 +11544,8 @@ class TestDiscoveryFailureLogging:
assert "unresolved" in caplog.text
def _unrestricted_auth() -> MagicMock:
"""A caller with no object_permission, so only server-level checks apply."""
user_api_key_auth = MagicMock()
user_api_key_auth.object_permission = None
user_api_key_auth.object_permission_id = None
return user_api_key_auth
def _unrestricted_auth() -> UserAPIKeyAuth:
return UserAPIKeyAuth()
def _permissive_proxy_logging() -> MagicMock:

View file

@ -145,7 +145,7 @@ async def test_build_effective_auth_contexts_appends_admitted_user_context(monke
assert contexts[-1].user_id == "user-42" and contexts[-1].team_id is None
assert [ctx.team_id for ctx in contexts[:-1]] == ["team-one"]
reload_mock.assert_awaited_once_with("user-42")
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
@pytest.mark.asyncio
@ -198,7 +198,7 @@ async def test_acting_user_auth_returns_admitted_subject_for_non_admin_sessions(
result = await acting_user_auth(user_auth)
assert result.user_id == "user-42" and result.team_id is None
reload_mock.assert_awaited_once_with("user-42")
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
@pytest.mark.asyncio

View file

@ -67,7 +67,7 @@ class TestAgentRequestHandler:
# Case 1: Both key and team have agents - intersection
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_key"
AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@ -86,7 +86,7 @@ class TestAgentRequestHandler:
# Case 2: Team has agents, key has none - inherit from team
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_key"
AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@ -105,7 +105,7 @@ class TestAgentRequestHandler:
# Case 3: Key has agents, team has none - key restrictions stand
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_key"
AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@ -120,7 +120,7 @@ class TestAgentRequestHandler:
# Case 4: No grant anywhere - unrestricted (documented open-by-default)
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_key"
AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@ -141,7 +141,7 @@ class TestAgentRequestHandler:
api_key="test-key", user_id="test-user", team_id="test-team"
)
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
mock_key.return_value = RestrictedAgentAccess(frozenset({"agent-alpha"}))
mock_team.return_value = RestrictedAgentAccess(frozenset({"agent-beta"}))
@ -198,7 +198,7 @@ class TestAgentRequestHandler:
@staticmethod
def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock:
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess:
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess:
assert user_api_key_auth is not None
return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess())
@ -249,7 +249,6 @@ class TestAgentRequestHandler:
frozenset({"agent-alpha"})
)
async def test_agent_access_groups_intersect_with_key_grants(self):
agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent")
resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"}))
@ -299,7 +298,7 @@ class TestAgentRequestHandler:
) as mock_groups:
mock_groups.return_value = []
assert await AgentRequestHandler._get_allowed_agents_for_key(
assert await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
) == RestrictedAgentAccess(frozenset())
@ -315,7 +314,7 @@ class TestAgentRequestHandler:
) as mock_groups:
mock_groups.side_effect = Exception("DB Error")
assert await AgentRequestHandler._get_allowed_agents_for_key(
assert await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
) == UnrestrictedAgentAccess()
@ -404,7 +403,7 @@ class TestAgentRequestHandler:
)
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_key"
AgentRequestHandler, "get_allowed_agents_for_key"
) as mock_key:
with patch.object(
AgentRequestHandler, "_get_allowed_agents_for_team"
@ -489,9 +488,9 @@ class TestAgentRequestHandler:
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
async def testget_allowed_agents_for_key_via_access_group_ids(self):
"""
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
Test that get_allowed_agents_for_key includes agents from key's access_group_ids
(unified access groups) when key has no native object_permission.
"""
mock_user_auth = UserAPIKeyAuth(
@ -508,16 +507,16 @@ class TestAgentRequestHandler:
new_callable=AsyncMock,
return_value=["agent-from-ag-1", "agent-from-ag-2"],
):
result = await AgentRequestHandler._get_allowed_agents_for_key(
result = await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert result == RestrictedAgentAccess(
frozenset({"agent-from-ag-1", "agent-from-ag-2"})
)
async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self):
async def testget_allowed_agents_for_key_combines_native_and_access_groups(self):
"""
Test that _get_allowed_agents_for_key combines agents from native object_permission
Test that get_allowed_agents_for_key combines agents from native object_permission
and key's access_group_ids (unified access groups).
"""
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
@ -540,7 +539,7 @@ class TestAgentRequestHandler:
new_callable=AsyncMock,
return_value=["agent-from-ag"],
):
result = await AgentRequestHandler._get_allowed_agents_for_key(
result = await AgentRequestHandler.get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert result == RestrictedAgentAccess(
@ -611,7 +610,7 @@ class TestAgentRequestHandler:
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
registry,
):
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
for key_grant, team_grant in (
(
@ -632,3 +631,257 @@ class TestAgentRequestHandler:
assert await AgentRequestHandler.resolve_agent_access(
user_api_key_auth=mock_user_auth
) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"state,allowed",
[
({}, True),
({"enabled": False}, False),
],
)
async def test_managed_invocation_requires_local_and_directory_admission(
monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityBinding
binding: Final = AgentIdentityBinding(
agent_id="target",
provider="microsoft_entra",
tenant_id="tenant",
client_id="client",
issuer="issuer",
revision="revision",
)
target: Final = AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True
).model_copy(update=state)
client: Final = MagicMock()
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
monkeypatch.setattr(proxy_server, "prisma_client", client)
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"])
auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
@pytest.mark.asyncio
@pytest.mark.parametrize("delegated", [True, False])
async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
monkeypatch: pytest.MonkeyPatch, delegated: bool
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth import auth_checks
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
database: Final = MagicMock()
monkeypatch.setattr(proxy_server, "prisma_client", database)
own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"])
human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"])
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
auth: Final = UserAPIKeyAuth(agent_id="actor")
auth.managed_agent_policy = AgentResponse(
agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
)
auth.managed_agent_context = ManagedAgentContext(
agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None
)
access: Final = await AgentRequestHandler.resolve_agent_access(auth)
assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"}))
@pytest.mark.asyncio
@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"])
async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches(
monkeypatch: pytest.MonkeyPatch, revoked: str
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
direct: Final = revoked == "direct-grant"
grouped: Final = revoked == "access-group"
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
human: Final = LiteLLM_UserTable(
user_id="human",
teams=[] if direct else ["team"],
organization_memberships=[],
object_permission_id="grant" if direct else None,
)
team: Final = LiteLLM_TeamTable(
team_id="team",
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission_id=None if grouped else "grant",
access_group_ids=["group"] if grouped else [],
)
group: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="Group", access_agent_ids=["target"]
)
cache: Final = UserApiKeyCache()
cache.set_cache("human", human)
cache.set_cache("team_id:team", team)
cache.set_cache(object_permission_cache_key("grant"), permission)
cache.set_cache("access_group_id:group", group)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human)
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
assert await verified_human_agent_grants("human") == frozenset({"target"})
client.writer_db.litellm_usertable.find_unique.return_value = (
human.model_copy(update={"teams": []}) if revoked == "user" else human
)
client.writer_db.litellm_teamtable.find_unique.return_value = (
team.model_copy(update={"members_with_roles": []})
if revoked == "team-member"
else team.model_copy(update={"object_permission_id": None})
if revoked == "team-grant"
else team
)
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = (
permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission
)
client.writer_db.litellm_accessgrouptable.find_unique.return_value = (
group.model_copy(update={"access_agent_ids": []}) if grouped else group
)
assert await verified_human_agent_grants("human") == frozenset()
client.db.litellm_usertable.find_unique.assert_not_called()
client.db.litellm_teamtable.find_unique.assert_not_called()
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
client.db.litellm_accessgrouptable.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.types.agents import AgentResponse
stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={})
registry: Final = AgentRegistry()
registry.register_agent(stale)
database: Final = MagicMock()
database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
monkeypatch.setattr(proxy_server, "prisma_client", database)
auth: Final = UserAPIKeyAuth(
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="permission", agent_access_groups=["group"]
)
)
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
frozenset({"revoked"})
)
database.writer_db.litellm_agentstable.find_many.return_value = []
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
frozenset()
)
database.db.litellm_agentstable.find_many.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", [[], ["group"]])
async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None:
assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set()
@pytest.mark.asyncio
@pytest.mark.parametrize("team", [False, True])
async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all(
monkeypatch: pytest.MonkeyPatch, team: bool
) -> None:
from fastapi import HTTPException
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
database: Final = MagicMock()
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable"))
database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", database)
auth: Final = UserAPIKeyAuth(
team_id="team" if team else None,
object_permission=None if team else LiteLLM_ObjectPermissionTable(
object_permission_id="grant", agent_access_groups=["group"]
),
)
with pytest.raises(HTTPException, match="policy is unavailable") as denied:
await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
assert denied.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("available", [False, True])
async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.auth import auth_checks
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None)
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None))
assert await AgentRequestHandler._get_allowed_agents_for_team(
UserAPIKeyAuth(team_id="missing"), strict=True
) == RestrictedAgentAccess(frozenset())
@pytest.mark.asyncio
@pytest.mark.parametrize("outage", [False, True])
async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy(
monkeypatch: pytest.MonkeyPatch, outage: bool
) -> None:
from fastapi import HTTPException
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.types.agents import AgentResponse
registry: Final = AgentRegistry()
registry.register_agent(AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True
))
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
return_value=None, side_effect=ConnectionError("unavailable") if outage else None
)
monkeypatch.setattr(proxy_server, "prisma_client", database)
if outage:
with pytest.raises(HTTPException, match="could not be loaded") as denied:
await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth())
assert denied.value.status_code == 503
else:
assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("grant", [False, True])
async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None:
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
auth: Final = UserAPIKeyAuth(agent_id="actor")
auth.managed_agent_policy = AgentResponse(
agent_id="actor", agent_name="Actor", agent_card_params={},
object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None,
)
auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated")
assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset())
assert await verified_human_agent_grants(None) == frozenset()

View file

@ -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
@ -10004,3 +10033,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"