mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(agents): authoritative permissions
This commit is contained in:
parent
6684256136
commit
18576ee4f2
17 changed files with 1498 additions and 205 deletions
|
|
@ -0,0 +1,67 @@
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
|
||||
async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
|
||||
|
||||
|
||||
async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
agent: Final = auth.managed_agent_policy
|
||||
if agent is None:
|
||||
return ()
|
||||
|
||||
try:
|
||||
base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
|
||||
ceilings: Final = await resolve_managed_agent_ceilings(agent)
|
||||
expanded: Final = tuple(
|
||||
frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
for ceiling in ceilings
|
||||
)
|
||||
own: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return tuple(sorted(own))
|
||||
if context.user_id is None:
|
||||
return ()
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human)
|
||||
return tuple(sorted(own.intersection(allowed)))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
|
||||
)
|
||||
|
||||
|
||||
async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if server_id not in await managed_agent_servers(auth):
|
||||
return []
|
||||
try:
|
||||
own: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return own
|
||||
if context.user_id is None:
|
||||
return []
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human)
|
||||
if own is None:
|
||||
return human_tools
|
||||
return own if human_tools is None else sorted(frozenset(own).intersection(human_tools))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
|
||||
)
|
||||
|
|
@ -67,7 +67,6 @@ from litellm.repositories.table_repositories import (
|
|||
AgentsRepository,
|
||||
MCPServerRepository,
|
||||
)
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -1086,7 +1085,7 @@ class MCPRequestHandler:
|
|||
assert_never(identity.subject_type)
|
||||
|
||||
@staticmethod
|
||||
async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
|
||||
async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
|
||||
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
|
||||
|
||||
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
|
||||
|
|
@ -1111,6 +1110,7 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=requires_fresh_policy,
|
||||
)
|
||||
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
|
||||
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
|
||||
|
|
@ -1119,6 +1119,7 @@ class MCPRequestHandler:
|
|||
if user_object is not None and object_permission is None and user_object.object_permission_id:
|
||||
object_permission = await get_object_permission(
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
check_db_only=requires_fresh_policy,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
|
@ -1147,6 +1148,7 @@ class MCPRequestHandler:
|
|||
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
|
||||
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
|
||||
admitted.mcp_admitted_user_subject = True
|
||||
admitted.requires_fresh_policy = requires_fresh_policy
|
||||
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
|
||||
# several teams under its own identity, so without this a cross-team user outruns every team's
|
||||
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
|
||||
|
|
@ -1597,6 +1599,11 @@ class MCPRequestHandler:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
|
||||
|
||||
return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
|
||||
|
||||
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
||||
try:
|
||||
|
|
@ -1606,7 +1613,7 @@ class MCPRequestHandler:
|
|||
# independent; an opt-out silences only its own source, inside the recursive call).
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
)
|
||||
|
||||
# Get allowed servers from key and team
|
||||
|
|
@ -1703,7 +1710,7 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth and user_api_key_auth.agent_id:
|
||||
agent_capped: Final = _agent_capped_servers(
|
||||
allowed_mcp_servers,
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
|
||||
)
|
||||
if agent_capped is not None:
|
||||
|
|
@ -1829,10 +1836,12 @@ class MCPRequestHandler:
|
|||
scoped.object_permission = auth.object_permission
|
||||
scoped.object_permission_id = auth.object_permission_id
|
||||
scoped.access_group_ids = auth.access_group_ids
|
||||
scoped.requires_fresh_policy = auth.requires_fresh_policy
|
||||
scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
|
||||
return scoped
|
||||
|
||||
@staticmethod
|
||||
async def _admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
|
||||
direct grants, plus every team they are a live roster member of.
|
||||
|
||||
|
|
@ -1886,6 +1895,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(auth and auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
|
||||
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
|
||||
|
|
@ -1941,11 +1951,11 @@ class MCPRequestHandler:
|
|||
roster instead of by grant charged unrelated teams' buckets)."""
|
||||
return [
|
||||
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
|
||||
for source in await MCPRequestHandler._admitted_subject_sources(auth)
|
||||
for source in await MCPRequestHandler.admitted_subject_sources(auth)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
|
||||
async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
|
||||
"""Union of what each of the admitted subject's sources reaches, each answered by the
|
||||
canonical resolver so no rule is reimplemented for this caller shape."""
|
||||
reachable: Final[set[str]] = set()
|
||||
|
|
@ -2007,7 +2017,7 @@ class MCPRequestHandler:
|
|||
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
|
||||
sources that actually grant that server.
|
||||
|
||||
|
|
@ -2088,6 +2098,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
if not team_obj:
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,13 +1,16 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, TypeAlias
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
AccessGroupIds: TypeAlias = tuple[str, ...]
|
||||
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
|
||||
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
|
||||
|
|
@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
|
|||
return tuple(agent.access_group_ids or ()) if agent is not None else ()
|
||||
|
||||
|
||||
async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
||||
async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
|
|
@ -47,6 +50,7 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
except HTTPException as e:
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
@ -73,3 +77,16 @@ async def resolve_agent_access_group_ceiling(
|
|||
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
|
||||
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
|
||||
)
|
||||
|
||||
|
||||
async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
|
||||
async def authoritative_group(group_id: str) -> LoadedAccessGroup:
|
||||
return await _load_access_group(group_id, check_db_only=True)
|
||||
|
||||
async def manual_ids(_agent_id: str) -> AccessGroupIds:
|
||||
return tuple(agent.access_group_ids or ())
|
||||
|
||||
manual: Final = await resolve_agent_access_group_ceiling(
|
||||
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
|
||||
)
|
||||
return (manual,) if manual is not None else ()
|
||||
|
|
|
|||
|
|
@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -86,7 +89,9 @@ class AgentRequestHandler:
|
|||
) -> AgentAccess:
|
||||
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
|
||||
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
|
||||
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
|
||||
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
|
||||
return await _managed_actor_agent_access(user_api_key_auth)
|
||||
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(user_api_key_auth)
|
||||
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
|
||||
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
|
||||
|
|
@ -104,13 +109,19 @@ class AgentRequestHandler:
|
|||
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
|
||||
|
||||
@staticmethod
|
||||
async def _resolve_key_team_agent_access(
|
||||
async def resolve_key_team_agent_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
try:
|
||||
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
|
||||
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
user_api_key_auth, strict=strict
|
||||
)
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
verbose_logger.warning("Failed to get allowed agents: %s", e)
|
||||
return UnrestrictedAgentAccess()
|
||||
return _intersect_agent_access(key_access, team_access)
|
||||
|
|
@ -144,6 +155,33 @@ class AgentRequestHandler:
|
|||
bool: True if agent is allowed, False otherwise
|
||||
"""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
|
||||
if prisma_client is not None or (registered is not None and registered.identity_managed):
|
||||
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
|
||||
if isinstance(target, AgentIdentityFailure):
|
||||
raise_identity_failure(target)
|
||||
if target is None and registered is not None and registered.identity_managed:
|
||||
return False
|
||||
if target is not None and target.identity_managed:
|
||||
if (
|
||||
not target.enabled
|
||||
or target.identity is None
|
||||
or not target.identity.active
|
||||
or user_api_key_auth is None
|
||||
):
|
||||
return False
|
||||
fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True})
|
||||
explicit: Final = await _granted_agent_ids(
|
||||
fresh_auth,
|
||||
_strict_agent_access,
|
||||
build_effective_auth_contexts,
|
||||
)
|
||||
return target.agent_id in explicit
|
||||
|
||||
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
|
||||
case UnrestrictedAgentAccess():
|
||||
|
|
@ -202,8 +240,10 @@ class AgentRequestHandler:
|
|||
return team_obj.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_agents_for_key(
|
||||
async def get_allowed_agents_for_key(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
"""
|
||||
Get allowed agents for a key.
|
||||
|
|
@ -237,24 +277,36 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
access_group_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
list(declared_access_groups), check_db_only=strict
|
||||
)
|
||||
)
|
||||
if declared_access_groups
|
||||
else ()
|
||||
)
|
||||
unified_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_unified_access_group_agents(
|
||||
list(key_access_group_ids), check_db_only=strict
|
||||
)
|
||||
)
|
||||
if key_access_group_ids
|
||||
else ()
|
||||
)
|
||||
|
||||
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
|
||||
return UnrestrictedAgentAccess()
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_agents_for_team(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
"""
|
||||
Get allowed agents for a team.
|
||||
|
|
@ -263,7 +315,7 @@ class AgentRequestHandler:
|
|||
2. Also includes agents from team's access_group_ids (unified access groups)
|
||||
|
||||
Fetches the team object once and reuses it for both permission sources.
|
||||
Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`.
|
||||
Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`.
|
||||
"""
|
||||
if user_api_key_auth is None:
|
||||
return UnrestrictedAgentAccess()
|
||||
|
|
@ -280,7 +332,7 @@ class AgentRequestHandler:
|
|||
)
|
||||
|
||||
if not prisma_client:
|
||||
return UnrestrictedAgentAccess()
|
||||
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
|
||||
|
||||
# Fetch the team object once for both permission sources
|
||||
team_obj: Final = await get_team_object(
|
||||
|
|
@ -289,10 +341,11 @@ class AgentRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=strict,
|
||||
)
|
||||
|
||||
if team_obj is None:
|
||||
return UnrestrictedAgentAccess()
|
||||
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
|
||||
|
||||
# 1. Get agents from object_permission (native permissions)
|
||||
object_permissions: Final = team_obj.object_permission
|
||||
|
|
@ -307,18 +360,28 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
access_group_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
list(declared_access_groups), check_db_only=strict
|
||||
)
|
||||
)
|
||||
if declared_access_groups
|
||||
else ()
|
||||
)
|
||||
unified_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_unified_access_group_agents(
|
||||
list(team_access_group_ids), check_db_only=strict
|
||||
)
|
||||
)
|
||||
if team_access_group_ids
|
||||
else ()
|
||||
)
|
||||
|
||||
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
# litellm-dashboard is the default UI team and will never have agents;
|
||||
# skip noisy warnings for it.
|
||||
if user_api_key_auth.team_id != UI_TEAM_ID:
|
||||
|
|
@ -326,7 +389,9 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
@staticmethod
|
||||
def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
|
||||
def _get_config_agent_ids_for_access_groups(
|
||||
config_agents: Sequence[AgentResponse], access_groups: list[str]
|
||||
) -> set[str]:
|
||||
"""
|
||||
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
|
||||
"""
|
||||
|
|
@ -339,7 +404,9 @@ class AgentRequestHandler:
|
|||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
|
||||
async def _get_db_agent_ids_for_access_groups(
|
||||
prisma_client, access_groups: list[str], *, check_db_only: bool = False
|
||||
) -> set[str]:
|
||||
"""
|
||||
Helper to get agent_ids from DB agents that match any of the given access groups.
|
||||
|
||||
|
|
@ -349,23 +416,27 @@ class AgentRequestHandler:
|
|||
if not access_groups or prisma_client is None:
|
||||
return set()
|
||||
|
||||
agents: Final = await AgentsRepository(prisma_client).table.find_many(
|
||||
agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
|
||||
where={"agent_access_groups": {"hasSome": access_groups}}
|
||||
)
|
||||
return {agent.agent_id for agent in agents}
|
||||
|
||||
@staticmethod
|
||||
async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
|
||||
async def _get_unified_access_group_agents(
|
||||
access_group_ids: list[str], *, check_db_only: bool = False
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve unified access group ids to agent IDs.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
|
||||
|
||||
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
|
||||
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
|
||||
|
||||
@staticmethod
|
||||
async def _get_agents_from_access_groups(
|
||||
access_groups: list[str],
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
|
||||
|
|
@ -373,14 +444,13 @@ class AgentRequestHandler:
|
|||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Use the helper for config-loaded agents
|
||||
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
|
||||
global_agent_registry.agent_list, access_groups
|
||||
)
|
||||
|
||||
# Use the helper for DB agents
|
||||
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
prisma_client, access_groups, check_db_only=check_db_only
|
||||
)
|
||||
|
||||
return list(config_agent_ids | db_agent_ids)
|
||||
|
|
@ -531,4 +601,58 @@ async def accessible_agents(
|
|||
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
|
||||
effective_contexts,
|
||||
)
|
||||
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
|
||||
allowed: Final = await asyncio.gather(
|
||||
*(
|
||||
AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
|
||||
for agent in agents
|
||||
if agent.identity_managed
|
||||
)
|
||||
)
|
||||
managed_ids: Final = frozenset(
|
||||
agent.agent_id
|
||||
for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
|
||||
if permitted
|
||||
)
|
||||
return tuple(
|
||||
agent
|
||||
for agent in agents
|
||||
if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
|
||||
)
|
||||
|
||||
|
||||
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
if auth.managed_agent_policy is not None:
|
||||
return await _managed_actor_agent_access(auth)
|
||||
return await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
|
||||
|
||||
|
||||
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
agent: Final = auth.managed_agent_policy
|
||||
if agent is None or not agent.object_permission:
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
|
||||
own_auth: Final = UserAPIKeyAuth(object_permission=permission)
|
||||
own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True))
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
|
||||
ceilings: Final = await resolve_managed_agent_ceilings(agent)
|
||||
capped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return RestrictedAgentAccess(capped)
|
||||
if context.user_id is None:
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
human_ids: Final = await verified_human_agent_grants(context.user_id)
|
||||
return RestrictedAgentAccess(capped.intersection(human_ids))
|
||||
|
||||
|
||||
async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if user_id is None:
|
||||
return frozenset()
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
sources: Final = await MCPRequestHandler.admitted_subject_sources(human)
|
||||
human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
|
||||
return frozenset().union(*(_granted_ids(access) for access in human_access))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -70,6 +70,9 @@ class _PrismaClientView(Protocol):
|
|||
@property
|
||||
def db(self) -> _PrismaTeamDb: ...
|
||||
|
||||
@property
|
||||
def writer_db(self) -> _PrismaTeamDb: ...
|
||||
|
||||
|
||||
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
|
||||
_JSON_ENCODED_TEAM_FIELDS: Final = (
|
||||
|
|
@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = (
|
|||
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
||||
"""Repository for team database operations."""
|
||||
|
||||
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
|
||||
super().__init__(prisma_client)
|
||||
self._use_writer = use_writer
|
||||
|
||||
@property
|
||||
def _db(self) -> _PrismaTeamDb:
|
||||
client: Final[_PrismaClientView] = self.prisma_client
|
||||
return client.db
|
||||
return client.writer_db if self._use_writer else client.db
|
||||
|
||||
@property
|
||||
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:
|
||||
|
|
|
|||
|
|
@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
|
|||
class UserRepository(BaseRepository[LiteLLM_UserTable]):
|
||||
"""Repository for user database operations."""
|
||||
|
||||
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
|
||||
super().__init__(prisma_client)
|
||||
self._use_writer = use_writer
|
||||
|
||||
@property
|
||||
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
|
||||
return self.prisma_client.db.litellm_usertable
|
||||
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
|
||||
return database.litellm_usertable
|
||||
|
||||
@property
|
||||
def model_class(self) -> type[LiteLLM_UserTable]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -145,7 +145,7 @@ async def test_build_effective_auth_contexts_appends_admitted_user_context(monke
|
|||
|
||||
assert contexts[-1].user_id == "user-42" and contexts[-1].team_id is None
|
||||
assert [ctx.team_id for ctx in contexts[:-1]] == ["team-one"]
|
||||
reload_mock.assert_awaited_once_with("user-42")
|
||||
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -198,7 +198,7 @@ async def test_acting_user_auth_returns_admitted_subject_for_non_admin_sessions(
|
|||
result = await acting_user_auth(user_auth)
|
||||
|
||||
assert result.user_id == "user-42" and result.team_id is None
|
||||
reload_mock.assert_awaited_once_with("user-42")
|
||||
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -67,7 +67,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 1: Both key and team have agents - intersection
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -86,7 +86,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 2: Team has agents, key has none - inherit from team
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -105,7 +105,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 3: Key has agents, team has none - key restrictions stand
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -120,7 +120,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 4: No grant anywhere - unrestricted (documented open-by-default)
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -141,7 +141,7 @@ class TestAgentRequestHandler:
|
|||
api_key="test-key", user_id="test-user", team_id="test-team"
|
||||
)
|
||||
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
|
||||
mock_key.return_value = RestrictedAgentAccess(frozenset({"agent-alpha"}))
|
||||
mock_team.return_value = RestrictedAgentAccess(frozenset({"agent-beta"}))
|
||||
|
|
@ -198,7 +198,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
@staticmethod
|
||||
def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock:
|
||||
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess:
|
||||
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess:
|
||||
assert user_api_key_auth is not None
|
||||
return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess())
|
||||
|
||||
|
|
@ -249,7 +249,6 @@ class TestAgentRequestHandler:
|
|||
frozenset({"agent-alpha"})
|
||||
)
|
||||
|
||||
|
||||
async def test_agent_access_groups_intersect_with_key_grants(self):
|
||||
agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent")
|
||||
resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"}))
|
||||
|
|
@ -299,7 +298,7 @@ class TestAgentRequestHandler:
|
|||
) as mock_groups:
|
||||
mock_groups.return_value = []
|
||||
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == RestrictedAgentAccess(frozenset())
|
||||
|
||||
|
|
@ -315,7 +314,7 @@ class TestAgentRequestHandler:
|
|||
) as mock_groups:
|
||||
mock_groups.side_effect = Exception("DB Error")
|
||||
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == UnrestrictedAgentAccess()
|
||||
|
||||
|
|
@ -404,7 +403,7 @@ class TestAgentRequestHandler:
|
|||
)
|
||||
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -489,9 +488,9 @@ class TestAgentRequestHandler:
|
|||
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
|
||||
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
|
||||
|
||||
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
|
||||
async def testget_allowed_agents_for_key_via_access_group_ids(self):
|
||||
"""
|
||||
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
|
||||
Test that get_allowed_agents_for_key includes agents from key's access_group_ids
|
||||
(unified access groups) when key has no native object_permission.
|
||||
"""
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
|
|
@ -508,16 +507,16 @@ class TestAgentRequestHandler:
|
|||
new_callable=AsyncMock,
|
||||
return_value=["agent-from-ag-1", "agent-from-ag-2"],
|
||||
):
|
||||
result = await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
result = await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
assert result == RestrictedAgentAccess(
|
||||
frozenset({"agent-from-ag-1", "agent-from-ag-2"})
|
||||
)
|
||||
|
||||
async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self):
|
||||
async def testget_allowed_agents_for_key_combines_native_and_access_groups(self):
|
||||
"""
|
||||
Test that _get_allowed_agents_for_key combines agents from native object_permission
|
||||
Test that get_allowed_agents_for_key combines agents from native object_permission
|
||||
and key's access_group_ids (unified access groups).
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
|
@ -540,7 +539,7 @@ class TestAgentRequestHandler:
|
|||
new_callable=AsyncMock,
|
||||
return_value=["agent-from-ag"],
|
||||
):
|
||||
result = await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
result = await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
assert result == RestrictedAgentAccess(
|
||||
|
|
@ -611,7 +610,7 @@ class TestAgentRequestHandler:
|
|||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
|
||||
registry,
|
||||
):
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
|
||||
for key_grant, team_grant in (
|
||||
(
|
||||
|
|
@ -632,3 +631,257 @@ class TestAgentRequestHandler:
|
|||
assert await AgentRequestHandler.resolve_agent_access(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"state,allowed",
|
||||
[
|
||||
({}, True),
|
||||
({"enabled": False}, False),
|
||||
],
|
||||
)
|
||||
async def test_managed_invocation_requires_local_and_directory_admission(
|
||||
monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="target",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
)
|
||||
target: Final = AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True
|
||||
).model_copy(update=state)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"])
|
||||
auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("delegated", [True, False])
|
||||
async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
|
||||
monkeypatch: pytest.MonkeyPatch, delegated: bool
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
database: Final = MagicMock()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"])
|
||||
human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"])
|
||||
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
|
||||
)
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None
|
||||
)
|
||||
access: Final = await AgentRequestHandler.resolve_agent_access(auth)
|
||||
assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"}))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"])
|
||||
async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches(
|
||||
monkeypatch: pytest.MonkeyPatch, revoked: str
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
|
||||
direct: Final = revoked == "direct-grant"
|
||||
grouped: Final = revoked == "access-group"
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
|
||||
human: Final = LiteLLM_UserTable(
|
||||
user_id="human",
|
||||
teams=[] if direct else ["team"],
|
||||
organization_memberships=[],
|
||||
object_permission_id="grant" if direct else None,
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="team",
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission_id=None if grouped else "grant",
|
||||
access_group_ids=["group"] if grouped else [],
|
||||
)
|
||||
group: Final = LiteLLM_AccessGroupTable(
|
||||
access_group_id="group", access_group_name="Group", access_agent_ids=["target"]
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("human", human)
|
||||
cache.set_cache("team_id:team", team)
|
||||
cache.set_cache(object_permission_cache_key("grant"), permission)
|
||||
cache.set_cache("access_group_id:group", group)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human)
|
||||
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
assert await verified_human_agent_grants("human") == frozenset({"target"})
|
||||
client.writer_db.litellm_usertable.find_unique.return_value = (
|
||||
human.model_copy(update={"teams": []}) if revoked == "user" else human
|
||||
)
|
||||
client.writer_db.litellm_teamtable.find_unique.return_value = (
|
||||
team.model_copy(update={"members_with_roles": []})
|
||||
if revoked == "team-member"
|
||||
else team.model_copy(update={"object_permission_id": None})
|
||||
if revoked == "team-grant"
|
||||
else team
|
||||
)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = (
|
||||
permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission
|
||||
)
|
||||
client.writer_db.litellm_accessgrouptable.find_unique.return_value = (
|
||||
group.model_copy(update={"access_agent_ids": []}) if grouped else group
|
||||
)
|
||||
assert await verified_human_agent_grants("human") == frozenset()
|
||||
client.db.litellm_usertable.find_unique.assert_not_called()
|
||||
client.db.litellm_teamtable.find_unique.assert_not_called()
|
||||
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
|
||||
client.db.litellm_accessgrouptable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={})
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(stale)
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
|
||||
database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
object_permission=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="permission", agent_access_groups=["group"]
|
||||
)
|
||||
)
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
|
||||
frozenset({"revoked"})
|
||||
)
|
||||
database.writer_db.litellm_agentstable.find_many.return_value = []
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
|
||||
frozenset()
|
||||
)
|
||||
database.db.litellm_agentstable.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("groups", [[], ["group"]])
|
||||
async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None:
|
||||
assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("team", [False, True])
|
||||
async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all(
|
||||
monkeypatch: pytest.MonkeyPatch, team: bool
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable"))
|
||||
database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
team_id="team" if team else None,
|
||||
object_permission=None if team else LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="grant", agent_access_groups=["group"]
|
||||
),
|
||||
)
|
||||
with pytest.raises(HTTPException, match="policy is unavailable") as denied:
|
||||
await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
|
||||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("available", [False, True])
|
||||
async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None)
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None))
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
UserAPIKeyAuth(team_id="missing"), strict=True
|
||||
) == RestrictedAgentAccess(frozenset())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outage", [False, True])
|
||||
async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, outage: bool
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True
|
||||
))
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None, side_effect=ConnectionError("unavailable") if outage else None
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
if outage:
|
||||
with pytest.raises(HTTPException, match="could not be loaded") as denied:
|
||||
await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth())
|
||||
assert denied.value.status_code == 503
|
||||
else:
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("grant", [False, True])
|
||||
async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None:
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={},
|
||||
object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None,
|
||||
)
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated")
|
||||
assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset())
|
||||
assert await verified_human_agent_grants(None) == frozenset()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue