mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(agents): enforce fresh MCP grants and repair auth fixtures
This commit is contained in:
parent
320a688bae
commit
e53e15c5e2
10 changed files with 311 additions and 64 deletions
|
|
@ -2109,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
|
||||
|
|
@ -2125,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(
|
||||
|
|
@ -2137,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(
|
||||
|
|
@ -2266,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
|
||||
)
|
||||
|
|
@ -2282,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
|
||||
|
|
@ -2382,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 = (
|
||||
|
|
@ -2537,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
|
||||
|
|
@ -2550,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 []
|
||||
)
|
||||
|
|
@ -2625,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
|
||||
|
|
@ -2640,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)
|
||||
)
|
||||
|
||||
|
|
@ -2704,7 +2736,11 @@ class MCPRequestHandler:
|
|||
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):
|
||||
|
|
@ -2834,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(
|
||||
|
|
@ -2843,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}
|
||||
|
|
@ -2935,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
|
||||
|
|
@ -3068,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)
|
||||
|
|
@ -3208,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
|
||||
|
|
@ -3237,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
|
||||
|
|
@ -3282,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
|
||||
|
|
@ -3403,9 +3458,12 @@ 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 user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
|
||||
|
|
@ -3475,7 +3533,9 @@ 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 is not None else None
|
||||
except Exception as e:
|
||||
|
|
@ -3497,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
|
||||
|
||||
|
|
@ -3534,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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -3583,6 +3583,8 @@ class MCPServerManager:
|
|||
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.
|
||||
|
|
@ -3592,6 +3594,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
|
||||
|
|
@ -3606,12 +3612,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:
|
||||
|
|
@ -3625,6 +3635,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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore |
|
|||
raise_identity_failure(failure)
|
||||
auth.managed_agent_policy = agent
|
||||
auth.billing_agent_policy = agent
|
||||
auth.requires_fresh_policy = True
|
||||
if auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated":
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
|
||||
|
|
|
|||
|
|
@ -162,6 +162,68 @@ async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
|
|||
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"])
|
||||
|
|
@ -255,7 +317,9 @@ async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_server
|
|||
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")]))
|
||||
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
|
||||
|
|
@ -284,7 +348,9 @@ async def test_manager_preserves_managed_server_grants_across_open_channels(
|
|||
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"})
|
||||
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == (
|
||||
{"slack"} if scoped else {"slack", "linear"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
@ -4721,7 +4723,9 @@ class TestAgentMCPPermissions:
|
|||
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
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
import asyncio
|
||||
import contextlib
|
||||
import contextvars
|
||||
|
|
@ -2244,7 +2245,7 @@ async def test_mcp_routing_chunked_initialize_to_stateful():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -2356,7 +2357,7 @@ async def test_mcp_routing_caps_body_peek_for_oversized_chunked_body():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
|
||||
patch(
|
||||
|
|
@ -2567,7 +2568,7 @@ async def test_mcp_routing_initialize_rejected_when_owner_at_session_cap():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
|
||||
patch(
|
||||
|
|
|
|||
|
|
@ -11152,6 +11152,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
|
||||
|
|
|
|||
|
|
@ -9,9 +9,12 @@ they may send a stale `mcp-session-id` header. This test verifies that:
|
|||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
|
||||
class TestHandleStaleMcpSession:
|
||||
"""Unit tests for the _handle_stale_mcp_session helper."""
|
||||
|
|
@ -260,7 +263,7 @@ async def test_stale_mcp_session_id_is_stripped():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -337,7 +340,7 @@ async def test_delete_stale_mcp_session_returns_success():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -386,7 +389,7 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
pytest.skip("MCP server not available")
|
||||
|
||||
session_id = "delete-failure-session"
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.api_key = "sk-test"
|
||||
user_auth.user_id = "test-user"
|
||||
auth_context = MagicMock()
|
||||
|
|
@ -491,7 +494,7 @@ async def test_valid_mcp_session_id_is_preserved():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -554,7 +557,7 @@ async def test_no_mcp_session_id_header_works_normally():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -613,7 +616,7 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
oauth_server = MagicMock()
|
||||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -700,7 +703,7 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me
|
|||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "sso-user-42"
|
||||
user_auth.mcp_admitted_user_subject = True
|
||||
oauth_server = MagicMock()
|
||||
|
|
@ -806,7 +809,7 @@ async def test_client_credentials_server_is_not_preemptively_challenged(m2m_fiel
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
m2m_server = MCPServer(
|
||||
server_id="m2m-server-id",
|
||||
|
|
@ -892,7 +895,7 @@ async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_cha
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
delegated_server = MagicMock()
|
||||
delegated_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -996,7 +999,7 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
oauth_server = MagicMock()
|
||||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -1092,7 +1095,7 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
delegated_server = MagicMock()
|
||||
delegated_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -1192,7 +1195,7 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
obo_server = MagicMock()
|
||||
obo_server.auth_type = MCPAuth.oauth2_token_exchange
|
||||
|
|
@ -1301,7 +1304,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
od_server = _build_passthrough_mode_server("od_server", MCPAuth.oauth_delegate)
|
||||
|
||||
|
|
@ -1366,7 +1369,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_with_forwarded_token_sk
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
od_server = _build_passthrough_mode_server("od_server", MCPAuth.oauth_delegate)
|
||||
|
||||
|
|
@ -1431,7 +1434,7 @@ async def _run_passthrough_connect(
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
server = _build_passthrough_mode_server(server_names[0], auth_type)
|
||||
|
||||
|
|
@ -1554,7 +1557,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
tp_server = _build_passthrough_mode_server("tp_server", MCPAuth.true_passthrough)
|
||||
|
||||
|
|
@ -1620,7 +1623,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
bridge_server = _build_passthrough_mode_server("tp_bridge_server", MCPAuth.true_passthrough).model_copy(
|
||||
update={"dcr_bridge": True}
|
||||
|
|
@ -1691,7 +1694,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_with_token_skips_prob
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
tp_server = _build_passthrough_mode_server("tp_server", MCPAuth.true_passthrough)
|
||||
|
||||
|
|
|
|||
|
|
@ -437,9 +437,12 @@ def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route
|
|||
|
||||
with pytest.raises(HTTPException, match="explicit or configured model"):
|
||||
managed_inference_request(route, {}, {"completion_model": "allowed-default"}, "cli")
|
||||
assert managed_inference_request(
|
||||
route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli"
|
||||
)["model"] == "requested"
|
||||
assert (
|
||||
managed_inference_request(route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli")[
|
||||
"model"
|
||||
]
|
||||
== "requested"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")])
|
||||
|
|
@ -482,3 +485,16 @@ async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None:
|
|||
assert auth.managed_agent_policy == agent()
|
||||
assert auth.billing_agent_policy == agent()
|
||||
assert auth.user_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_next_request() -> None:
|
||||
"""Managed MCP grants (toolsets, access groups) are read through the shared resolvers, which only
|
||||
bypass the warm cache and the replica when the subject carries requires_fresh_policy"""
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent())
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
assert auth.requires_fresh_policy is False
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.requires_fresh_policy is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue