diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ceb2b612568..22b12b2f229 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3434,6 +3434,10 @@ 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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py index a7df7cf3b25..b6ddd9cb7ac 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py @@ -259,3 +259,43 @@ async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatc 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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index b467960a3a0..ede304742e0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -4708,7 +4708,7 @@ class TestAgentMCPPermissions: ), ) - async def testget_allowed_mcp_servers_for_agent_includes_toolset_servers(self): + async def test_get_allowed_mcp_servers_for_agent_includes_toolset_servers(self): """An agent granted only mcp_toolsets reaches the toolset's servers, exactly as a key, team, or org granted only toolsets does""" user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets") @@ -4777,7 +4777,7 @@ class TestAgentMCPPermissions: assert result == [] - async def testget_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self): + async def test_get_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self): """The agent's tool ceiling on a server is its direct tool grants plus the tools its toolsets grant there, and None only when neither names the server""" user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")