fix(agents): enforce fresh MCP grants and repair auth fixtures

This commit is contained in:
Joshua Valluru 2026-09-27 09:31:55 -07:00
parent 320a688bae
commit e53e15c5e2
10 changed files with 311 additions and 64 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -369,7 +369,9 @@ class TestMCPRequestHandler:
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
assert result == ["server-a"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(
toolset_ids=["toolset-1"], requires_fresh_policy=False
)
async def test_get_allowed_mcp_servers_for_key_skips_toolset_resolution_when_none_granted(self):
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
@ -4147,7 +4149,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
"group-server2",
}
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
mock_get_access_group_servers.assert_called_once_with(["dev-group"], requires_fresh_policy=False)
finally:
for sid in ("direct-server1", "direct-server2"):
global_mcp_server_manager.registry.pop(sid, None)
@ -4316,7 +4318,7 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
assert set(result) == {"direct-server", "group-server"}
mock_get_perm.assert_not_called()
mock_access_groups.assert_called_once_with(["grp-alpha"])
mock_access_groups.assert_called_once_with(["grp-alpha"], requires_fresh_policy=False)
finally:
global_mcp_server_manager.registry.pop("direct-server", None)
@ -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

View file

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

View file

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

View file

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

View file

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