feat(agents): enforce current model and MCP permission grants

This commit is contained in:
Joshua Valluru 2026-09-26 11:46:15 -07:00
parent 1dd343ae18
commit d97e5a90e4
16 changed files with 750 additions and 130 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,6 +1836,8 @@ 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
@ -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
@ -1945,7 +1955,7 @@ class MCPRequestHandler:
]
@staticmethod
async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
"""Union of what each of the admitted subject's sources reaches, each answered by the
canonical resolver so no rule is reimplemented for this caller shape."""
reachable: Final[set[str]] = set()
@ -2007,7 +2017,7 @@ class MCPRequestHandler:
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
@staticmethod
async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
sources that actually grant that server.
@ -2088,6 +2098,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not team_obj:
@ -2171,6 +2182,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
@staticmethod
@ -2219,12 +2231,17 @@ class MCPRequestHandler:
if not user_api_key_auth:
return None
if user_api_key_auth.managed_agent_policy is not None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools
return await managed_agent_tools(server_id, user_api_key_auth)
try:
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
if _is_mcp_admitted_user_subject(user_api_key_auth):
return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
# Get key and team object permissions (already loaded in main auth flow)
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
@ -2334,7 +2351,7 @@ class MCPRequestHandler:
if user_api_key_auth.agent_id:
# Pre-fetch agent object_permission once to avoid a duplicate DB query.
agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server(
agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(
server_id=server_id,
user_api_key_auth=user_api_key_auth,
agent_object_permission=agent_obj_perm,
@ -2456,6 +2473,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if not raw_server_ids:
return []
@ -2502,6 +2520,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@ -2550,7 +2569,7 @@ class MCPRequestHandler:
"""Get allowed MCP servers a caller inherits from the team it is pinned to.
Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not
fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``,
fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``,
and each of those sources pins a single ``team_id`` before reaching this point. Keeping the
fan-out here as well would be a second multi-team path to drift from that one.
"""
@ -2568,7 +2587,7 @@ class MCPRequestHandler:
which must NOT silently gain the union across every team the user belongs to), and it covers
each single-source auth an admitted subject fans out into — those pin a team_id, so they land
on the first branch. The admitted subject itself never reaches here: it resolves per source
in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
resolves to no teams exactly as before."""
if user_api_key_auth is None or not user_api_key_auth.team_id:
return []
@ -2596,6 +2615,7 @@ class MCPRequestHandler:
user_id_upsert=False,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises
verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e)
@ -2667,6 +2687,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
return []
@ -2680,6 +2701,7 @@ class MCPRequestHandler:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers)
@ -2716,6 +2738,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
raise unloadable from e
@ -2961,7 +2984,9 @@ class MCPRequestHandler:
return None
user_id: Final = user_api_key_auth.user_id
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(
user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy
)
if object_permission_id is None:
return None
@ -2971,6 +2996,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if object_permission is None:
raise ValueError(
@ -2979,7 +3005,9 @@ class MCPRequestHandler:
return object_permission
@staticmethod
async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
async def _user_object_permission_id(
user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
) -> str | None:
"""The permission row this human's user row links to, or None when they link none.
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
@ -2988,16 +3016,23 @@ class MCPRequestHandler:
whether someone is entitled is the state that existed before this level, so it places no
ceiling. Only a link we DID resolve can make the caller deny.
"""
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import user_api_key_cache
cache_key: Final = user_object_permission_id_cache_key(user_id)
try:
cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key)
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
return None
if isinstance(cached, str) and cached:
return cached
user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
user_row: Final = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=check_db_only,
)
linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None
object_permission_id: Final = linked if isinstance(linked, str) and linked else None
await user_api_key_cache.async_set_cache(
@ -3006,7 +3041,9 @@ class MCPRequestHandler:
ttl=get_management_object_ttl(user_api_key_cache),
)
return object_permission_id
except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
if check_db_only:
raise HTTPException(503, "User policy is unavailable") from e
verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e)
return None
@ -3119,9 +3156,13 @@ class MCPRequestHandler:
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot
disagree."""
if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
if (
user_api_key_auth is None
or user_api_key_auth.mcp_explicit_grants_only
or not user_api_key_has_admin_view(user_api_key_auth)
):
return False
object_permission: Final = user_api_key_auth.object_permission
credential_scoped: Final = (
@ -3302,6 +3343,10 @@ class MCPRequestHandler:
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
if user_api_key_auth.managed_agent_policy is not None:
permission: Final = user_api_key_auth.managed_agent_policy.object_permission
return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
@ -3319,7 +3364,7 @@ class MCPRequestHandler:
)
@staticmethod
async def _get_allowed_mcp_servers_for_agent(
async def get_allowed_mcp_servers_for_agent(
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> list[str]:
@ -3363,7 +3408,7 @@ class MCPRequestHandler:
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
return []
@ -3390,7 +3435,7 @@ class MCPRequestHandler:
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
@staticmethod
async def _get_agent_tool_permissions_for_server(
async def get_agent_tool_permissions_for_server(
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
@ -3432,9 +3477,9 @@ class MCPRequestHandler:
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
return list(agent_tools) if agent_tools else None
return list(agent_tools) if agent_tools is not None else None
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
return None
@ -3548,6 +3593,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if key_object_permission is None:
return []
@ -3591,6 +3637,7 @@ class MCPRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
if team_obj is None:
verbose_logger.debug("team_obj is None")

View file

@ -3383,7 +3383,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()
@ -3435,6 +3437,7 @@ class MCPServerManager:
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
@ -3466,7 +3469,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)
)
@ -3535,7 +3538,7 @@ class MCPServerManager:
return [
server_id
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
if scope is None or server_id == scope
if not explicit_grants_only and (scope is None or server_id == scope)
]
async def resolve_toolset_tool_permissions(

View file

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

@ -27,7 +27,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
validate_langfuse_span_scope_value,
validate_no_callback_env_reference,
)
from litellm.types.agents import AgentCaller
from litellm.types.agents import AgentCaller, AgentResponse
from litellm.types.integrations.compression_interception import (
CompressionSavingsMetadata,
)
@ -46,6 +46,7 @@ from litellm.types.mcp import (
MCPTransportType,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.types.proxy.carried_budget_state import (
OrgBudgetSnapshot,
TeamBudgetSnapshot,
@ -3305,6 +3306,8 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union
# or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization.
mcp_admitted_user_subject: bool = Field(default=False, exclude=True)
requires_fresh_policy: bool = Field(default=False, exclude=True)
mcp_explicit_grants_only: bool = Field(default=False, exclude=True)
# team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP
# servers through several teams at once and therefore has no single team_id for the limiter to
# key off. Server-only and stripped from validated input for the same reason as the marker
@ -3329,6 +3332,11 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
"user id."
),
)
invoked_agent_id: str | None = Field(default=None, exclude=True)
agent_invocation_cost: float | None = Field(default=None, exclude=True)
billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True)
agent_caller: AgentCaller | None = Field(
default=None,
exclude=True,
@ -3366,11 +3374,18 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# path via post-construction assignment. Strip it from any validated input (constructor
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
values.pop("requires_fresh_policy", None)
values.pop("mcp_explicit_grants_only", None)
values.pop("mcp_source_team_rpm_limits", None)
values.pop("mcp_session_resource_server_id", None)
values.pop("mcp_toolset_id", None)
values.pop("via_virtual_key", None)
values.pop("agent_caller", None)
values.pop("managed_agent_context", None)
values.pop("managed_agent_policy", None)
values.pop("invoked_agent_id", None)
values.pop("agent_invocation_cost", None)
values.pop("billing_agent_policy", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):

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

@ -1069,6 +1069,21 @@ async def common_checks(
code=status.HTTP_400_BAD_REQUEST,
)
if _model and valid_token is not None and valid_token.managed_agent_policy is not None:
managed_models: Final = (valid_token.managed_agent_policy.object_permission or MappingProxyType({})).get(
"models", ()
)
if not isinstance(managed_models, (list, tuple)) or not managed_models:
raise HTTPException(403, "This agent has no model grants")
_can_object_call_model(
model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router),
llm_router=llm_router,
models=list(managed_models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router)
await _check_agent_caller_model_access(
model=_model,
@ -2653,7 +2668,7 @@ async def get_user_object(
)
if should_check_db:
response = await _user_table(UserRepository(prisma_client)).find_unique(
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
@ -2691,7 +2706,7 @@ async def get_user_object(
budget_duration=new_user_params["budget_duration"]
)
response = await _user_table(UserRepository(prisma_client)).create(
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).create(
data=new_user_params,
include={"organization_memberships": True},
)
@ -3116,9 +3131,9 @@ class TeamNotFoundError(HTTPException):
@log_db_metrics
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None, *, use_writer: bool = False
) -> "_PrismaTeamRow | None":
response = await _team_table(TeamRepository(prisma_client)).find_unique(
response = await _team_table(TeamRepository(prisma_client, use_writer=use_writer)).find_unique(
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
)
@ -3152,6 +3167,7 @@ async def _get_team_object_from_user_api_key_cache(
proxy_logging_obj: ProxyLogging | None,
key: str,
team_id_upsert: bool | None = None,
use_writer: bool = False,
) -> LiteLLM_TeamTableCachedObj:
db_access_time_key: Final = key
should_check_db: Final = _should_check_db(
@ -3160,7 +3176,9 @@ async def _get_team_object_from_user_api_key_cache(
db_cache_expiry=db_cache_expiry,
)
if should_check_db:
response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert)
response = await _get_team_db_check(
team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert, use_writer=use_writer
)
# The database answered and the row is not there. Distinct from every
# other failure here, which leaves the team's grant unknown.
if response is None:
@ -3182,8 +3200,11 @@ async def _get_team_object_from_user_api_key_cache(
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
check_db_only=use_writer,
)
except Exception as e:
if use_writer:
raise
verbose_proxy_logger.debug(
"Failed to load object_permission for team %s with object_permission_id=%s: %s",
team_id,
@ -3273,6 +3294,7 @@ async def get_team_object(
db_cache_expiry=db_cache_expiry,
key=key,
team_id_upsert=team_id_upsert,
use_writer=bool(check_db_only),
)
except TeamNotFoundError:
raise
@ -3318,16 +3340,15 @@ async def get_access_object(
prisma_client: DatabaseClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None = None,
*,
check_db_only: bool = False,
) -> LiteLLM_AccessGroupTable:
"""
- Check if access_group_id in proxy AccessGroupTable
- Always checks cache first, then DB only when not found in cache
- Checks cache first unless authoritative writer admission is requested
- if valid, return LiteLLM_AccessGroupTable object
- if not, then raise an error
Unlike get_team_object, this has no check_cache_only or check_db_only flags;
it always follows cache-first-then-db semantics.
Raises:
- HTTPException: If access group doesn't exist in db or cache (status_code=404)
"""
@ -3336,18 +3357,19 @@ async def get_access_object(
key: Final = f"access_group_id:{access_group_id}"
cached_access_obj: Final = await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_AccessGroupTable,
cached_access_obj: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(key=key, model_type=LiteLLM_AccessGroupTable)
)
if cached_access_obj is not None:
return cached_access_obj
# Not in cache - fetch from DB
try:
response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique(
where={"access_group_id": access_group_id}
)
response: Final = await _dictable_table(
AccessGroupRepository(prisma_client, use_writer=check_db_only), "access_group"
).find_unique(where={"access_group_id": access_group_id})
if response is None:
raise HTTPException(
@ -3931,6 +3953,7 @@ async def get_object_permission(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> LiteLLM_ObjectPermissionTable | None:
"""
- Check if object permission id in proxy ObjectPermissionTable
@ -3942,9 +3965,13 @@ async def get_object_permission(
# check if in cache
key: Final = object_permission_cache_key(object_permission_id)
deserialized_perm: Final = await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_ObjectPermissionTable,
deserialized_perm: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_ObjectPermissionTable,
)
)
if deserialized_perm is not None:
return deserialized_perm
@ -3952,7 +3979,7 @@ 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:
@ -4177,6 +4204,7 @@ async def _get_resources_from_access_groups(
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Fetch access groups by their IDs (from cache or DB) and collect
@ -4219,6 +4247,7 @@ async def _get_resources_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
resources.extend(getattr(ag, resource_field, []))
except Exception:
@ -4254,6 +4283,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect MCP server IDs from unified access groups.
@ -4265,6 +4295,7 @@ async def _get_mcp_server_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -4273,6 +4304,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect agent IDs from unified access groups.
@ -4284,6 +4316,7 @@ async def _get_agent_ids_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -4483,26 +4516,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,261 @@
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth
from litellm.proxy.auth import auth_checks
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth:
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="agent-permissions",
mcp_servers=["slack", "linear"],
mcp_tool_permissions={"slack": list(tools)} if tools is not None else None,
)
agent: Final = AgentResponse(
agent_id="publisher",
agent_name="Publisher",
agent_card_params={},
object_permission=permission.model_dump(),
identity_managed=True,
)
auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id)
auth.managed_agent_policy = agent
auth.managed_agent_context = ManagedAgentContext(
agent_id=agent.agent_id,
mode="delegated" if delegated else "autonomous",
user_id="human" if delegated else None,
)
return auth
@pytest.fixture(autouse=True)
def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager())
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
@pytest.mark.asyncio
@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write")))
async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None:
auth: Final = actor(tools)
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None)
assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize(
"agent_tools,user_tools,expected",
(
(None, ("read",), ("read",)),
(("read",), None, ("read",)),
(("read", "write"), ("read",), ("read",)),
(("read",), ("write",), ()),
((), None, ()),
),
)
async def test_delegated_server_and_tool_intersections(
monkeypatch: pytest.MonkeyPatch,
agent_tools: tuple[str, ...] | None,
user_tools: tuple[str, ...] | None,
expected: tuple[str, ...],
) -> None:
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="user-permissions",
mcp_servers=["slack", "user-only"],
mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None,
)
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
auth: Final = actor(agent_tools, delegated=True)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected)
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == []
@pytest.mark.asyncio
async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable")))
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
assert failure.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ())))
async def test_access_groups_cap_agent_servers_without_granting_new_ones(
monkeypatch: pytest.MonkeyPatch,
servers: tuple[str, ...],
expected: tuple[str, ...],
) -> None:
from litellm.proxy._types import LiteLLM_AccessGroupTable
group: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers)
)
monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group))
auth: Final = actor(None)
assert auth.managed_agent_policy is not None
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"access_group_ids": ["group"]})
assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected
if "slack" not in expected:
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
monkeypatch: pytest.MonkeyPatch, change: str
) -> None:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]}
)
user: Final = LiteLLM_UserTable(
user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant"
)
cache: Final = UserApiKeyCache()
cache.set_cache("human", user)
cache.set_cache(object_permission_cache_key("user-grant"), permission)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
auth: Final = actor(("read", "write"), delegated=True)
assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"}
if change == "disabled":
client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy(
update={"metadata": {"scim_active": False}}
)
elif change == "outage":
client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable")
elif change == "servers":
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
update={"mcp_servers": [], "mcp_tool_permissions": {}}
)
else:
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
update={"mcp_tool_permissions": {"slack": ["read"]}}
)
if change in ("disabled", "outage"):
with pytest.raises(HTTPException):
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
else:
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (
["read"] if change == "tools" else []
)
client.db.litellm_usertable.find_unique.assert_not_called()
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
@pytest.mark.parametrize("has_grant", [True, False])
@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
monkeypatch: pytest.MonkeyPatch,
role: str,
open_channel: str,
has_grant: bool,
agent_tools: tuple[str, ...] | None,
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {
name: MCPServer(
server_id=name,
name=name,
transport="http",
url="https://example.com/mcp",
allow_all_keys=open_channel == "operator",
)
for name in ("slack", "linear")
}
from litellm.proxy._experimental.mcp_server import db
monkeypatch.setattr(
db,
"get_active_submitted_mcp_server_ids_for_user",
AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
)
user: Final = LiteLLM_UserTable(
user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
)
team: Final = LiteLLM_TeamTable(
team_id="team",
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission_id="team-grant",
)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
monkeypatch.setattr(proxy_server, "prisma_client", client)
auth: Final = actor(agent_tools, delegated=True)
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
assert admitted.user_role == role
@pytest.mark.asyncio
async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy._experimental.mcp_server import db
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager: Final = mcp_server_manager.global_mcp_server_manager
manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
auth: Final = UserAPIKeyAuth(user_id="human")
auth.mcp_explicit_grants_only = True
with pytest.MonkeyPatch.context() as patcher:
patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
assert await manager.get_allowed_mcp_servers(auth) == []
auth.mcp_explicit_grants_only = False
assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
@pytest.mark.asyncio
async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
assert await managed_agent_servers(UserAPIKeyAuth()) == ()
auth: Final = actor(None, delegated=True)
assert auth.managed_agent_context is not None
auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
@pytest.mark.asyncio
async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")]))
with pytest.raises(HTTPException) as failure:
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
assert failure.value.status_code == 503

View file

@ -4383,7 +4383,7 @@ class TestAgentMCPPermissions:
self._team_servers({"callers": ["server_2", "server_3"]}),
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: neither the agent's owner nor the caller has a personal grant
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({})
@ -4402,7 +4402,7 @@ class TestAgentMCPPermissions:
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: same seam, keyed by which user is being asked about
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": ["server_1"]})
@ -4421,7 +4421,7 @@ class TestAgentMCPPermissions:
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
),
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: None is the resolver's own "entitlement unresolvable" signal
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": None})
@ -4538,7 +4538,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_1"]
@ -4555,7 +4555,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = [] # no agent-level restriction
@ -4611,7 +4611,7 @@ class TestAgentMCPPermissions:
)
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_2", "server_3"]
@ -4637,7 +4637,7 @@ class TestAgentMCPPermissions:
):
with patch.object(
MCPRequestHandler,
"_get_agent_tool_permissions_for_server",
"get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=["tool_a"],
) as mock_agent_tools:
@ -4669,7 +4669,7 @@ class TestAgentMCPPermissions:
):
with patch.object(
MCPRequestHandler,
"_get_agent_tool_permissions_for_server",
"get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=None,
):
@ -4708,7 +4708,7 @@ class TestAgentMCPPermissions:
),
)
async def test_get_allowed_mcp_servers_for_agent_includes_toolset_servers(self):
async def testget_allowed_mcp_servers_for_agent_includes_toolset_servers(self):
"""An agent granted only mcp_toolsets reaches the toolset's servers, exactly as a
key, team, or org granted only toolsets does"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
@ -4718,7 +4718,7 @@ class TestAgentMCPPermissions:
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
result = await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
assert sorted(result) == ["server-a", "server-direct"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
@ -4760,7 +4760,7 @@ class TestAgentMCPPermissions:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
with pytest.raises(UnloadableEntitlementError):
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler,
@ -4777,7 +4777,7 @@ class TestAgentMCPPermissions:
assert result == []
async def test_get_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self):
async def testget_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self):
"""The agent's tool ceiling on a server is its direct tool grants plus the tools its
toolsets grant there, and None only when neither names the server"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
@ -4789,13 +4789,13 @@ class TestAgentMCPPermissions:
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
server_a_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-a", user_api_key_auth
)
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
server_b_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-b", user_api_key_auth
)
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
server_c_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
"server-c", user_api_key_auth
)
@ -5833,7 +5833,7 @@ def test_expand_permission_list_does_not_honor_all_proxy_sentinel():
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically():
async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically(monkeypatch):
"""The TEAM resolver expands the all-proxy sentinel to every registered server and
picks up a server registered later, so a team scoped to all-proxy tracks the live
registry without any change to its stored permission. Reverting the team-side
@ -5850,6 +5850,9 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setattr(global_mcp_server_manager, "registry", {})
monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {})
for sid in ("srv-x", "srv-y"):
global_mcp_server_manager.registry[sid] = MCPServer(
server_id=sid,
@ -9673,7 +9676,10 @@ class TestGetUserObjectPermission:
def _prisma_with_user(self, user_row):
prisma_client = MagicMock()
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
from litellm.proxy._types import LiteLLM_UserTable
row = LiteLLM_UserTable(user_id="human", object_permission_id=user_row.object_permission_id) if user_row is not None else None
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row)
return prisma_client
async def test_resolves_through_the_shared_permission_cache(self):
@ -9688,7 +9694,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch(
"litellm.proxy.auth.auth_checks.get_object_permission",
new_callable=AsyncMock,
@ -9715,7 +9721,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch("litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock) as mock_get_perm,
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@ -9734,7 +9740,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@ -9748,7 +9754,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
):
assert await MCPRequestHandler._get_user_object_permission(auth) is None
@ -9765,7 +9771,7 @@ class TestGetUserObjectPermission:
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
patch(
"litellm.proxy.auth.auth_checks.get_object_permission",
new_callable=AsyncMock,
@ -10085,3 +10091,25 @@ 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

View file

@ -208,15 +208,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
@ -5405,9 +5399,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
@ -5474,9 +5466,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
@ -5543,9 +5533,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
@ -5580,9 +5568,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
@ -6658,9 +6644,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
@ -6744,9 +6728,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()
@ -11256,12 +11238,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

@ -3,6 +3,8 @@ from typing import Final
import pytest
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
actor_admission_failure,
invocation_target,
@ -182,3 +184,28 @@ def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route
]
== "requested"
)
def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
context: Final = ManagedAgentContext(
agent_id="agent", binding_revision="current", mode="delegated", user_id="human"
)
auth: Final = UserAPIKeyAuth.model_validate(
{
"managed_agent_context": context,
"requires_fresh_policy": True,
"mcp_explicit_grants_only": True,
"managed_agent_policy": agent(),
"billing_agent_policy": agent(),
"invoked_agent_id": "forged-target",
"agent_invocation_cost": 0.0,
}
)
assert auth.requires_fresh_policy is False
assert auth.mcp_explicit_grants_only is False
assert "mcp_explicit_grants_only" not in auth.model_dump()
assert auth.managed_agent_context is None
assert auth.managed_agent_policy is None
assert auth.billing_agent_policy is None
assert auth.invoked_agent_id is None
assert auth.agent_invocation_cost is None

View file

@ -1091,7 +1091,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
monkeypatch.setitem(auth_checks.last_db_access_time, f"user_id:{user_id}", (None, time.time()))
db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user")
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
mock_prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
result = await get_user_object(
user_id=user_id,
@ -1103,7 +1103,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
assert result is not None
assert result.user_id == user_id
mock_prisma_client.db.litellm_usertable.find_unique.assert_awaited_once()
mock_prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_once()
@pytest.mark.asyncio
@ -3058,7 +3058,7 @@ async def test_get_team_object_raises_404_when_not_found():
mock_prisma_client = MagicMock()
mock_db = AsyncMock()
mock_prisma_client.db = mock_db
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
@ -3076,11 +3076,40 @@ async def test_get_team_object_raises_404_when_not_found():
assert "Team doesn't exist in db" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_get_team_object_check_db_only_reads_writer_through_the_shared_loader():
"""Management endpoints mock ``_get_team_object_from_user_api_key_cache`` and expect
``check_db_only`` to still flow through it; only the table it reads moves to the writer."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.auth import auth_checks
from litellm.proxy.auth.auth_checks import get_team_object
row = {"team_id": "team-writer", "models": ["gpt-4o"], "object_permission_id": None}
prisma = MagicMock()
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
prisma.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
cache = MagicMock()
cache.async_get_cache = AsyncMock(return_value=None)
cache.async_set_cache = AsyncMock()
shared_loader = AsyncMock(wraps=auth_checks._get_team_object_from_user_api_key_cache)
with patch.object(auth_checks, "_get_team_object_from_user_api_key_cache", shared_loader):
team = await get_team_object("team-writer", prisma, cache, check_db_only=True)
assert team.team_id == "team-writer"
assert shared_loader.await_args.kwargs["use_writer"] is True
prisma.writer_db.litellm_teamtable.find_unique.assert_awaited_once()
prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
cache.async_set_cache.assert_awaited_once()
def _mock_prisma_for_team_lookup(find_unique):
from unittest.mock import MagicMock
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_teamtable.find_unique = find_unique
mock_prisma_client.writer_db.litellm_teamtable.find_unique = find_unique
return mock_prisma_client
@ -9979,3 +10008,88 @@ 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
@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"