diff --git a/litellm/proxy/_experimental/mcp_server/platform_mcp.py b/litellm/proxy/_experimental/mcp_server/platform_mcp.py new file mode 100644 index 00000000000..66ea264e577 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/platform_mcp.py @@ -0,0 +1,220 @@ +import json +import weakref +from typing import Any, Iterable, Optional, Sequence + +from litellm._logging import verbose_logger +from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer + +try: + from mcp.types import Tool as MCPTool +except ImportError: + MCPTool = None # type: ignore + + +DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD = 10 +PLATFORM_MCP_LIST_SERVERS_TOOL_NAME = "list_servers" +PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME = "enable_server" +PLATFORM_MCP_TOOL_NAMES = frozenset( + { + PLATFORM_MCP_LIST_SERVERS_TOOL_NAME, + PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME, + } +) + +_enabled_servers_by_session: "weakref.WeakKeyDictionary[Any, frozenset[str]]" = weakref.WeakKeyDictionary() + + +async def get_platform_mcp_settings() -> tuple[bool, int]: + from litellm.proxy.proxy_server import general_settings, prisma_client + + settings = dict(general_settings or {}) + if prisma_client is not None: + from litellm.proxy.utils import get_config_param + + row = await get_config_param(prisma_client, "general_settings") + param_value = getattr(row, "param_value", None) if row is not None else None + if isinstance(param_value, dict): + settings.update(param_value) + + enabled = settings.get("platform_mcp_enabled") is True + threshold = _coerce_positive_threshold(settings.get("platform_mcp_tool_threshold")) + return enabled, threshold + + +def build_platform_mcp_tools() -> list[Any]: + if MCPTool is None: + return [] + + return [ + MCPTool( + name=PLATFORM_MCP_LIST_SERVERS_TOOL_NAME, + description=( + "List the MCP servers this key can access, including the server " + "name and description so you can choose which server to enable." + ), + inputSchema={ + "type": "object", + "properties": {}, + "additionalProperties": False, + }, + ), + MCPTool( + name=PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME, + description=( + "Return full tool definitions for one accessible MCP server. " + "Use a server name returned by list_servers." + ), + inputSchema={ + "type": "object", + "properties": { + "server_name": { + "type": "string", + "description": "The MCP server name returned by list_servers.", + } + }, + "required": ["server_name"], + "additionalProperties": False, + }, + ), + ] + + +def should_compress_tools( + *, + platform_mcp_enabled: bool, + threshold: int, + tool_count: int, + requested_mcp_servers: Optional[Sequence[str]], + enabled_server_names: Sequence[str], +) -> bool: + return ( + platform_mcp_enabled and requested_mcp_servers is None and not enabled_server_names and tool_count > threshold + ) + + +def should_include_platform_meta_tools( + *, + platform_mcp_enabled: bool, + requested_mcp_servers: Optional[Sequence[str]], + enabled_server_names: Sequence[str], +) -> bool: + return platform_mcp_enabled and requested_mcp_servers is None and len(enabled_server_names) > 0 + + +def get_enabled_server_names_for_session(session: Optional[Any]) -> tuple[str, ...]: + if session is None: + return () + try: + return tuple(sorted(_enabled_servers_by_session.get(session, frozenset()))) + except TypeError: + verbose_logger.debug( + "Platform MCP session object cannot be used for enabled-server storage: %s", + type(session).__name__, + ) + return () + + +def enable_server_for_session(session: Optional[Any], server: MCPServer) -> None: + if session is None: + return + current = frozenset(get_enabled_server_names_for_session(session)) + next_value = current | frozenset([_server_match_name(server)]) + try: + _enabled_servers_by_session[session] = next_value + except TypeError: + verbose_logger.debug( + "Platform MCP could not store enabled server for session type: %s", + type(session).__name__, + ) + + +def is_platform_mcp_tool(name: str) -> bool: + return name in PLATFORM_MCP_TOOL_NAMES + + +def extract_enable_server_name(arguments: Optional[dict[str, Any]]) -> Optional[str]: + if not arguments: + return None + value = arguments.get("server_name") or arguments.get("mcp_name") or arguments.get("name") + return value if isinstance(value, str) and value.strip() else None + + +def serialize_server_summary(server: MCPServer) -> dict[str, str]: + return { + "name": _server_display_name(server), + "description": _server_description(server), + } + + +def serialize_server_tool_response( + *, + server: MCPServer, + tools: Sequence[Any], +) -> str: + return json.dumps( + { + "server": serialize_server_summary(server), + "tools": [serialize_tool(tool) for tool in tools], + } + ) + + +def serialize_servers_response(servers: Iterable[MCPServer]) -> str: + return json.dumps( + {"servers": [serialize_server_summary(server) for server in sorted(servers, key=_server_display_name)]} + ) + + +def serialize_tool(tool: Any) -> dict[str, Any]: + if hasattr(tool, "model_dump"): + try: + dumped = tool.model_dump( + mode="json", + by_alias=True, + exclude_none=True, + ) + except TypeError: + dumped = tool.model_dump() + if isinstance(dumped, dict): + return dumped + + input_schema = getattr(tool, "inputSchema", None) + if input_schema is None: + input_schema = getattr(tool, "input_schema", {}) + serialized_tool = { + "name": getattr(tool, "name", ""), + "description": getattr(tool, "description", "") or "", + "inputSchema": input_schema or {}, + } + for attr_name, output_name in [ + ("title", "title"), + ("outputSchema", "outputSchema"), + ("icons", "icons"), + ("annotations", "annotations"), + ("meta", "_meta"), + ("execution", "execution"), + ]: + value = getattr(tool, attr_name, None) + if value is not None: + serialized_tool[output_name] = value + return serialized_tool + + +def _coerce_positive_threshold(value: Any) -> int: + if isinstance(value, int) and value > 0: + return value + return DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD + + +def _server_match_name(server: MCPServer) -> str: + return server.alias or server.server_name or server.name + + +def _server_display_name(server: MCPServer) -> str: + return server.alias or server.server_name or server.name + + +def _server_description(server: MCPServer) -> str: + mcp_info: Optional[MCPInfo] = server.mcp_info + description = mcp_info.get("description") if mcp_info else None + return description if isinstance(description, str) else "" diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 08e42e918e9..ffae2e169a1 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -322,6 +322,19 @@ if MCP_AVAILABLE: notification_options: Optional[NotificationOptions] = None, experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None, ) -> InitializationOptions: + notification_options = NotificationOptions( + prompts_changed=( + notification_options.prompts_changed + if notification_options is not None + else False + ), + resources_changed=( + notification_options.resources_changed + if notification_options is not None + else False + ), + tools_changed=True, + ) opts = Server.create_initialization_options( self, notification_options=notification_options, @@ -660,6 +673,25 @@ if MCP_AVAILABLE: verbose_logger.debug( f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}" ) + from litellm.proxy._experimental.mcp_server.platform_mcp import ( + get_platform_mcp_settings, + is_platform_mcp_tool, + ) + + if is_platform_mcp_tool(name): + platform_mcp_enabled, _ = await get_platform_mcp_settings() + if platform_mcp_enabled: + return await _handle_platform_mcp_tool_call( + name=name, + arguments=arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + host_progress_callback = None try: host_ctx = server.request_context @@ -2169,6 +2201,25 @@ if MCP_AVAILABLE: # Resolve toolset permissions and merge into the key's object_permission # so that the existing filter_tools_by_key_team_permissions logic picks them up. user_api_key_auth = await _merge_toolset_permissions(user_api_key_auth) + from litellm.proxy._experimental.mcp_server.platform_mcp import ( + build_platform_mcp_tools, + get_enabled_server_names_for_session, + get_platform_mcp_settings, + should_compress_tools, + should_include_platform_meta_tools, + ) + + platform_mcp_enabled, platform_mcp_threshold = ( + await get_platform_mcp_settings() + ) + enabled_server_names = get_enabled_server_names_for_session( + get_active_mcp_session() + ) + effective_mcp_servers = ( + list(enabled_server_names) + if platform_mcp_enabled and mcp_servers is None and enabled_server_names + else mcp_servers + ) # Get tools from managed MCP servers with error handling managed_tools = [] @@ -2176,7 +2227,7 @@ if MCP_AVAILABLE: managed_tools = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, + mcp_servers=effective_mcp_servers, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, @@ -2192,8 +2243,147 @@ if MCP_AVAILABLE: ) # Continue with empty managed tools list instead of failing completely + if should_compress_tools( + platform_mcp_enabled=platform_mcp_enabled, + threshold=platform_mcp_threshold, + tool_count=len(managed_tools), + requested_mcp_servers=mcp_servers, + enabled_server_names=enabled_server_names, + ): + return build_platform_mcp_tools() + + if should_include_platform_meta_tools( + platform_mcp_enabled=platform_mcp_enabled, + requested_mcp_servers=mcp_servers, + enabled_server_names=enabled_server_names, + ): + return build_platform_mcp_tools() + managed_tools + return managed_tools + async def _handle_platform_mcp_tool_call( + name: str, + arguments: Optional[Dict[str, Any]], + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], + mcp_servers: Optional[List[str]], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> CallToolResult: + from litellm.proxy._experimental.mcp_server.platform_mcp import ( + PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME, + PLATFORM_MCP_LIST_SERVERS_TOOL_NAME, + enable_server_for_session, + extract_enable_server_name, + get_platform_mcp_settings, + serialize_server_tool_response, + serialize_servers_response, + ) + + platform_mcp_enabled, _ = await get_platform_mcp_settings() + if not platform_mcp_enabled: + return CallToolResult( + content=[TextContent(text="Platform MCP is disabled.", type="text")], + isError=True, + ) + + allowed_servers = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + ) + + if name == PLATFORM_MCP_LIST_SERVERS_TOOL_NAME: + return CallToolResult( + content=[ + TextContent(text=serialize_servers_response(allowed_servers), type="text") + ], + isError=False, + ) + + if name != PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME: + return CallToolResult( + content=[TextContent(text=f"Unknown Platform MCP tool: {name}", type="text")], + isError=True, + ) + + requested_server_name = extract_enable_server_name(arguments) + if requested_server_name is None: + return CallToolResult( + content=[ + TextContent( + text="Missing required argument: server_name", + type="text", + ) + ], + isError=True, + ) + + requested_server_name_lower = requested_server_name.lower() + selected_server = next( + ( + server + for server in allowed_servers + if requested_server_name_lower + in [ + known_name.lower() + for known_name in iter_known_server_prefixes(server) + if known_name + ] + ), + None, + ) + if selected_server is None: + return CallToolResult( + content=[ + TextContent( + text=f"MCP server '{requested_server_name}' is not available to this key.", + type="text", + ) + ], + isError=True, + ) + + selected_server_name = ( + selected_server.alias or selected_server.server_name or selected_server.name + ) + tools = await _get_tools_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=[selected_server_name], + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + active_session = get_active_mcp_session() + enable_server_for_session(active_session, selected_server) + await _send_platform_mcp_tool_list_changed(active_session) + + return CallToolResult( + content=[ + TextContent( + text=serialize_server_tool_response( + server=selected_server, + tools=tools, + ), + type="text", + ) + ], + isError=False, + ) + + async def _send_platform_mcp_tool_list_changed(session: Optional[Any]) -> None: + if session is None or not hasattr(session, "send_tool_list_changed"): + return + try: + import inspect + + result = session.send_tool_list_changed() + if inspect.isawaitable(result): + await result + except Exception as e: + verbose_logger.debug("Platform MCP tool-list notification failed: %s", e) + async def _list_mcp_prompts( user_api_key_auth: Optional[UserAPIKeyAuth] = None, mcp_auth_header: Optional[str] = None, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8e2ec423cde..0bc373cc281 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2339,6 +2339,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).", ) + platform_mcp_enabled: Optional[bool] = Field( + False, + description="If True, enables Platform MCP staged tool loading for aggregate MCP tools/list responses.", + ) + platform_mcp_tool_threshold: Optional[int] = Field( + 10, + description="When Platform MCP is enabled, aggregate MCP tools/list responses with more than this many tools are compressed to Platform MCP meta-tools.", + ) mcp_trusted_proxy_ranges: Optional[List[str]] = Field( None, description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.", diff --git a/tests/mcp_tests/test_platform_mcp.py b/tests/mcp_tests/test_platform_mcp.py new file mode 100644 index 00000000000..b5b46325f22 --- /dev/null +++ b/tests/mcp_tests/test_platform_mcp.py @@ -0,0 +1,377 @@ +import json +import os +import sys +from typing import Optional + +import pytest +from mcp.types import Tool + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy._experimental.mcp_server import platform_mcp +from litellm.proxy._experimental.mcp_server import server as mcp_server_module +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _tool(name: str) -> Tool: + return Tool( + name=name, + description=f"{name} description", + inputSchema={"type": "object", "properties": {}}, + ) + + +def _rich_tool(name: str) -> Tool: + return Tool( + name=name, + title="Get Ticket", + description=f"{name} description", + inputSchema={"type": "object", "properties": {}}, + outputSchema={"type": "object", "properties": {"id": {"type": "string"}}}, + _meta={"source": "platform-mcp-test"}, + ) + + +def _server( + *, + server_id: str = "server-1", + name: str = "servicenow", + alias: Optional[str] = None, + description: str = "Service management", +) -> MCPServer: + return MCPServer( + server_id=server_id, + name=name, + alias=alias, + server_name=name, + url="https://example.com/mcp", + transport=MCPTransport.http, + mcp_info={"description": description}, + ) + + +async def _merge_toolset_permissions(user_api_key_auth): + return user_api_key_auth + + +async def _enabled_platform_settings() -> tuple[bool, int]: + return True, 10 + + +async def _disabled_platform_settings() -> tuple[bool, int]: + return False, 10 + + +@pytest.mark.asyncio +async def test_platform_mcp_disabled_returns_normal_tools(monkeypatch): + normal_tools = [_tool(f"tool_{idx}") for idx in range(11)] + + async def fake_get_tools(**kwargs): + return normal_tools + + monkeypatch.setattr( + mcp_server_module, + "_merge_toolset_permissions", + _merge_toolset_permissions, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _disabled_platform_settings, + ) + + tools = await mcp_server_module._list_mcp_tools() + + assert tools == normal_tools + + +def test_platform_mcp_advertises_tool_list_changed_capability(): + options = mcp_server_module.server.create_initialization_options() + + assert options.capabilities.tools is not None + assert options.capabilities.tools.listChanged is True + + +@pytest.mark.asyncio +async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatch): + normal_tools = [_tool(f"tool_{idx}") for idx in range(11)] + + async def fake_get_tools(**kwargs): + return normal_tools + + monkeypatch.setattr( + mcp_server_module, + "_merge_toolset_permissions", + _merge_toolset_permissions, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + + tools = await mcp_server_module._list_mcp_tools() + + assert [tool.name for tool in tools] == ["list_servers", "enable_server"] + + +@pytest.mark.asyncio +async def test_platform_mcp_does_not_compress_scoped_server_tools(monkeypatch): + normal_tools = [_tool(f"tool_{idx}") for idx in range(11)] + + async def fake_get_tools(**kwargs): + return normal_tools + + monkeypatch.setattr( + mcp_server_module, + "_merge_toolset_permissions", + _merge_toolset_permissions, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + + tools = await mcp_server_module._list_mcp_tools(mcp_servers=["servicenow"]) + + assert tools == normal_tools + + +@pytest.mark.asyncio +async def test_platform_mcp_enabled_session_returns_meta_tools_and_enabled_server_tools( + monkeypatch, +): + selected_tools = [_tool("servicenow_get_ticket")] + + async def fake_get_tools(**kwargs): + assert kwargs["mcp_servers"] == ["servicenow"] + return selected_tools + + monkeypatch.setattr( + mcp_server_module, + "_merge_toolset_permissions", + _merge_toolset_permissions, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + monkeypatch.setattr( + platform_mcp, + "get_enabled_server_names_for_session", + lambda _session: ("servicenow",), + ) + + tools = await mcp_server_module._list_mcp_tools() + + assert [tool.name for tool in tools] == [ + "list_servers", + "enable_server", + "servicenow_get_ticket", + ] + + +@pytest.mark.asyncio +async def test_platform_mcp_list_servers_returns_names_and_descriptions(monkeypatch): + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [_server(name="servicenow"), _server(name="github", description="Code")] + + monkeypatch.setattr( + mcp_server_module, + "_get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + + result = await mcp_server_module._handle_platform_mcp_tool_call( + name="list_servers", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(api_key="test"), + mcp_auth_header=None, + mcp_servers=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + ) + + assert result.isError is False + payload = json.loads(result.content[0].text) + assert payload == { + "servers": [ + {"name": "github", "description": "Code"}, + {"name": "servicenow", "description": "Service management"}, + ] + } + + +@pytest.mark.asyncio +async def test_platform_mcp_enable_server_returns_selected_server_tool_definitions( + monkeypatch, +): + selected_server = _server(name="servicenow") + selected_tools = [_rich_tool("servicenow_get_ticket")] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [selected_server] + + async def fake_get_tools(**kwargs): + assert kwargs["mcp_servers"] == ["servicenow"] + return selected_tools + + async def fake_send_tool_list_changed(_session): + return None + + monkeypatch.setattr( + mcp_server_module, + "_get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + monkeypatch.setattr( + platform_mcp, + "enable_server_for_session", + lambda _session, _server: None, + ) + monkeypatch.setattr( + mcp_server_module, + "_send_platform_mcp_tool_list_changed", + fake_send_tool_list_changed, + ) + + result = await mcp_server_module._handle_platform_mcp_tool_call( + name="enable_server", + arguments={"server_name": "servicenow"}, + user_api_key_auth=UserAPIKeyAuth(api_key="test"), + mcp_auth_header=None, + mcp_servers=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + ) + + assert result.isError is False + payload = json.loads(result.content[0].text) + assert payload["server"] == { + "name": "servicenow", + "description": "Service management", + } + assert payload["tools"] == [ + { + "name": "servicenow_get_ticket", + "title": "Get Ticket", + "description": "servicenow_get_ticket description", + "inputSchema": {"type": "object", "properties": {}}, + "outputSchema": { + "type": "object", + "properties": {"id": {"type": "string"}}, + }, + "_meta": {"source": "platform-mcp-test"}, + } + ] + + +@pytest.mark.asyncio +async def test_platform_mcp_enable_server_updates_same_session_list_tools(monkeypatch): + class FakeSession: + def __init__(self): + self.tool_list_changed_count = 0 + + async def send_tool_list_changed(self): + self.tool_list_changed_count += 1 + + selected_server = _server(name="servicenow") + selected_tools = [_tool("servicenow_get_ticket")] + requested_servers = [] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [selected_server] + + async def fake_get_tools(**kwargs): + requested_servers.append(kwargs["mcp_servers"]) + return selected_tools + + session = FakeSession() + + monkeypatch.setattr( + mcp_server_module, + "_merge_toolset_permissions", + _merge_toolset_permissions, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + ) + monkeypatch.setattr( + mcp_server_module, + "_get_tools_from_mcp_servers", + fake_get_tools, + ) + monkeypatch.setattr( + platform_mcp, + "get_platform_mcp_settings", + _enabled_platform_settings, + ) + platform_mcp._enabled_servers_by_session.clear() + token = mcp_server_module.active_mcp_session_var.set(session) + try: + enable_result = await mcp_server_module._handle_platform_mcp_tool_call( + name="enable_server", + arguments={"server_name": "servicenow"}, + user_api_key_auth=UserAPIKeyAuth(api_key="test"), + mcp_auth_header=None, + mcp_servers=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + ) + tools = await mcp_server_module._list_mcp_tools() + finally: + mcp_server_module.active_mcp_session_var.reset(token) + platform_mcp._enabled_servers_by_session.clear() + + assert enable_result.isError is False + assert session.tool_list_changed_count == 1 + assert requested_servers == [["servicenow"], ["servicenow"]] + assert [tool.name for tool in tools] == [ + "list_servers", + "enable_server", + "servicenow_get_ticket", + ] diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx new file mode 100644 index 00000000000..73e8292a07c --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx @@ -0,0 +1,67 @@ +import React from "react"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import PlatformMCPTab from "./PlatformMCPTab"; +import * as networking from "../networking"; + +vi.mock("../networking", () => ({ + getConfigFieldSetting: vi.fn(), + updateConfigFieldSetting: vi.fn().mockResolvedValue(undefined), +})); + +describe("PlatformMCPTab", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(networking.getConfigFieldSetting).mockImplementation(async (_accessToken, fieldName) => { + if (fieldName === "platform_mcp_enabled") { + return { field_value: true }; + } + if (fieldName === "platform_mcp_tool_threshold") { + return { field_value: 10 }; + } + return { field_value: null }; + }); + }); + + it("shows the pre-v0 warning and v0 meta-tools", async () => { + render(); + + expect(await screen.findByText(/Platform MCP is a pre-v0 feature/i)).toBeInTheDocument(); + expect(screen.getByText(/tools\/list responses/i)).toBeInTheDocument(); + expect(screen.getByText("list_servers")).toBeInTheDocument(); + expect(screen.getByText("enable_server")).toBeInTheDocument(); + expect(screen.queryByText("search_tools")).not.toBeInTheDocument(); + expect(screen.queryByText("sandbox_execute")).not.toBeInTheDocument(); + }); + + it("updates the enabled setting through the existing config field endpoint", async () => { + render(); + + const toggle = await screen.findByRole("switch"); + fireEvent.click(toggle); + + await waitFor(() => { + expect(networking.updateConfigFieldSetting).toHaveBeenCalledWith( + "token", + "platform_mcp_enabled", + false, + ); + }); + }); + + it("updates the threshold through the existing config field endpoint", async () => { + render(); + + const thresholdInput = await screen.findByRole("spinbutton"); + fireEvent.change(thresholdInput, { target: { value: "12" } }); + fireEvent.click(screen.getByRole("button", { name: /save/i })); + + await waitFor(() => { + expect(networking.updateConfigFieldSetting).toHaveBeenCalledWith( + "token", + "platform_mcp_tool_threshold", + 12, + ); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx new file mode 100644 index 00000000000..7bbb0662f6e --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx @@ -0,0 +1,171 @@ +import React, { useEffect, useState } from "react"; +import { ExperimentOutlined, SaveOutlined, ToolOutlined } from "@ant-design/icons"; +import { Button, Card, InputNumber, Spin, Switch, Typography } from "antd"; +import { getConfigFieldSetting, updateConfigFieldSetting } from "../networking"; + +const { Text } = Typography; + +const PLATFORM_MCP_ENABLED_FIELD = "platform_mcp_enabled"; +const PLATFORM_MCP_THRESHOLD_FIELD = "platform_mcp_tool_threshold"; +const DEFAULT_THRESHOLD = 10; + +interface PlatformMCPTabProps { + accessToken: string | null; +} + +const getFieldValue = (response: unknown, fallback: boolean | number): boolean | number => { + if (response && typeof response === "object") { + const field = response as { field_value?: unknown; field_default_value?: unknown }; + if (field.field_value !== undefined && field.field_value !== null) { + return field.field_value as boolean | number; + } + if (field.field_default_value !== undefined && field.field_default_value !== null) { + return field.field_default_value as boolean | number; + } + } + return fallback; +}; + +const PlatformMCPTab: React.FC = ({ accessToken }) => { + const [loading, setLoading] = useState(true); + const [savingEnabled, setSavingEnabled] = useState(false); + const [savingThreshold, setSavingThreshold] = useState(false); + const [enabled, setEnabled] = useState(false); + const [threshold, setThreshold] = useState(DEFAULT_THRESHOLD); + + useEffect(() => { + const loadSettings = async () => { + if (!accessToken) { + setLoading(false); + return; + } + + setLoading(true); + try { + const [enabledResponse, thresholdResponse] = await Promise.all([ + getConfigFieldSetting(accessToken, PLATFORM_MCP_ENABLED_FIELD), + getConfigFieldSetting(accessToken, PLATFORM_MCP_THRESHOLD_FIELD), + ]); + setEnabled(Boolean(getFieldValue(enabledResponse, false))); + const nextThreshold = Number(getFieldValue(thresholdResponse, DEFAULT_THRESHOLD)); + setThreshold(Number.isFinite(nextThreshold) && nextThreshold > 0 ? nextThreshold : DEFAULT_THRESHOLD); + } catch (error) { + console.error("Failed to load Platform MCP settings:", error); + } finally { + setLoading(false); + } + }; + + loadSettings(); + }, [accessToken]); + + const handleToggle = async (checked: boolean) => { + if (!accessToken) return; + setSavingEnabled(true); + try { + await updateConfigFieldSetting(accessToken, PLATFORM_MCP_ENABLED_FIELD, checked); + setEnabled(checked); + } catch (error) { + console.error("Failed to update Platform MCP enabled setting:", error); + } finally { + setSavingEnabled(false); + } + }; + + const handleSaveThreshold = async () => { + if (!accessToken) return; + setSavingThreshold(true); + try { + await updateConfigFieldSetting(accessToken, PLATFORM_MCP_THRESHOLD_FIELD, threshold); + } catch (error) { + console.error("Failed to update Platform MCP threshold:", error); + } finally { + setSavingThreshold(false); + } + }; + + if (loading) { + return ( +
+ +
+ ); + } + + return ( +
+
+ + + Platform MCP is a pre-v0 feature. The dashboard control is only available to proxy admins and should not be + used in production. + +
+ +
+ Platform MCP +

+ When enabled, LiteLLM compresses aggregate MCP tools/list responses only after the caller's filtered tool + count is over the configured threshold. +

+
+ + +
+
+ Enable Platform MCP compression +

+ Disabled returns the current full tool list. Enabled keeps the full tool list at or below the threshold, + then returns only list_servers and enable_server above it. +

+
+ +
+
+ + +
+
+ Compression threshold +

+ Default is 10 tools. Compression starts when the final accessible tool count is greater than this value. +

+
+
+ setThreshold(Number(value || 1))} /> + +
+
+
+ +
+ +
+ +
+ list_servers +

+ Returns MCP server names and descriptions for servers the key can access. +

+
+
+
+ +
+ +
+ enable_server +

+ Returns full tool definitions for one accessible MCP server. +

+
+
+
+
+
+ ); +}; + +export default PlatformMCPTab; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx index e70548d6a96..df38cf46c76 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx @@ -62,12 +62,21 @@ vi.mock("./mcp_tool_configuration", () => ({ > Disable all tools + ), })); vi.mock("./mcp_connection_status", () => ({ - default: ({ tools }: { tools?: any[] }) => ( + default: ({ tools }: { tools?: unknown[] }) => (
), })); @@ -419,6 +428,50 @@ describe("CreateMCPServer", () => { expect(payload.mcp_info.tool_allowlist_enforced).toBe(true); expect(payload.allowed_tools).toEqual([]); }); + + it("creates the server with the selected tool allowlist", async () => { + await selectHttpTransport(); + + const user = userEvent.setup({ delay: null }); + + const nameInput = getServerNameInput(); + await user.type(nameInput, "Selected_Tools_Server"); + + const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com"); + await user.type(urlInput, "https://example.com/mcp"); + + await selectAntOption("Authentication", "None"); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Enable read_user only" })); + }); + + vi.mocked(networking.createMCPServer).mockResolvedValue({ + server_id: "new-server-1", + server_name: "Selected_Tools_Server", + alias: "Selected_Tools_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "none", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }); + + const submitButton = screen.getByRole("button", { name: "Add MCP Server" }); + await act(async () => { + fireEvent.click(submitButton); + }); + + await waitFor(() => { + expect(networking.createMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.mcp_info.tool_allowlist_enforced).toBe(true); + expect(payload.allowed_tools).toEqual(["read_user"]); + }); }); describe("when OAuth interactive auth is selected", () => { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index ddcc9f65d38..c690d5bca2d 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -1112,6 +1112,7 @@ const CreateMCPServer: React.FC = ({ externalIsLoading={isLoadingTools} externalError={toolsError} externalCanFetch={canFetchTools} + defaultViewMode="flat" />
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx index 446fdd8c22d..299bbcd4f2e 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx @@ -7,15 +7,17 @@ import * as networking from "../networking"; // Mock the networking module vi.mock("../networking", () => ({ - fetchMCPServers: vi.fn(), - fetchMCPServerHealth: vi.fn(), + fetchMCPServers: vi.fn().mockResolvedValue([]), + fetchMCPServerHealth: vi.fn().mockResolvedValue([]), deleteMCPServer: vi.fn(), getProxyBaseUrl: vi.fn().mockReturnValue("http://localhost:4000"), fetchMCPClientIp: vi.fn().mockResolvedValue(null), + getConfigFieldSetting: vi.fn().mockResolvedValue({ field_value: null }), getGeneralSettingsCall: vi.fn().mockResolvedValue([]), updateConfigFieldSetting: vi.fn().mockResolvedValue(undefined), deleteConfigFieldSetting: vi.fn().mockResolvedValue(undefined), listMCPUserEnvVarStatus: vi.fn().mockResolvedValue([]), + modelHubCall: vi.fn().mockResolvedValue({ data: [] }), })); // Mock NotificationsManager @@ -67,6 +69,33 @@ describe("MCPServers", () => { expect(getByText("MCP Servers")).toBeInTheDocument(); }); + it("should only show Platform MCP tab to proxy admins", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + + const queryClient = createQueryClient(); + const { rerender } = render( + + + , + ); + + await waitFor(() => { + expect(screen.getByText("MCP Servers")).toBeInTheDocument(); + }); + expect(screen.getByRole("tab", { name: /Platform MCP/ })).toBeInTheDocument(); + + rerender( + + + , + ); + + await waitFor(() => { + expect(screen.queryByRole("tab", { name: /Platform MCP/ })).not.toBeInTheDocument(); + }); + }); + it("should render mocked MCP servers data in the table", async () => { // Mock MCP servers data const mockServers = [ diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index 0a445265e2d..541658b42e8 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -1,4 +1,4 @@ -import { isAdminRole } from "@/utils/roles"; +import { isAdminRole, isProxyAdminRole } from "@/utils/roles"; import { QuestionCircleOutlined, SearchOutlined } from "@ant-design/icons"; import { Button, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react"; import NewBadge from "../common_components/NewBadge"; @@ -20,6 +20,7 @@ import MCPSemanticFilterSettings from "../Settings/AdminSettings/MCPSemanticFilt import MCPNetworkSettings from "./MCPNetworkSettings"; import MCPDiscovery from "./mcp_discovery"; import { ByokCredentialModal } from "./ByokCredentialModal"; +import PlatformMCPTab from "./PlatformMCPTab"; import { getSecureItem } from "@/utils/secureStorage"; import { TOOLS_OAUTH_UI_STATE_KEY } from "@/hooks/mcpOAuthUtils"; import UserEnvVarsModal from "./UserEnvVarsModal"; @@ -506,6 +507,13 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) )} + {isProxyAdminRole(userRole) && ( + + + Platform MCP + + + )} @@ -663,6 +671,11 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) )} + {isProxyAdminRole(userRole) && ( + + + + )} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx index 064e3dea614..512eed99e95 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx @@ -30,6 +30,69 @@ const renderToolConfiguration = (onAllowedToolsChange = vi.fn()) => { }; describe("MCPToolConfiguration", () => { + it("can start new-server onboarding in the flat checklist view", async () => { + render( + , + ); + + expect(screen.getByLabelText("Flat List")).toBeChecked(); + }); + + it("toggles a flat-list checkbox once without bubbling to the row", async () => { + const onAllowedToolsChange = vi.fn(); + + const Wrapper = () => { + const [allowedTools, setAllowedTools] = useState([]); + + return ( + { + onAllowedToolsChange(nextAllowedTools); + setAllowedTools(nextAllowedTools); + }} + toolNameToDisplayName={{}} + toolNameToDescription={{}} + onToolNameToDisplayNameChange={vi.fn()} + onToolNameToDescriptionChange={vi.fn()} + externalTools={tools} + externalCanFetch + defaultViewMode="flat" + /> + ); + }; + + render(); + + await waitFor(() => { + expect(screen.getByText("2 of 2 tools enabled for user access")).toBeInTheDocument(); + }); + onAllowedToolsChange.mockClear(); + + fireEvent.click(screen.getAllByRole("checkbox")[0]); + + await waitFor(() => { + expect(onAllowedToolsChange).toHaveBeenCalledTimes(1); + expect(onAllowedToolsChange).toHaveBeenCalledWith(["delete_user"]); + }); + }); + it("shows legacy unrestricted edit tools enabled in flat view", async () => { const onAllowedToolsChange = renderToolConfiguration(); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx index 870cdfb1e3b..4ca21f1b968 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx @@ -30,6 +30,7 @@ interface MCPToolConfigurationProps { externalCanFetch?: boolean; /** When true, do not auto-select all tools for servers with no stored allowlist. */ isEditMode?: boolean; + defaultViewMode?: "crud" | "flat"; } interface ToolEntry { @@ -69,7 +70,11 @@ const ToolRow: React.FC = ({ >
onToggle(tool.name)}>
- onToggle(tool.name)} /> + e.stopPropagation()} + onChange={() => onToggle(tool.name)} + />
{toolNameToDisplayName[tool.name] || tool.name} @@ -158,10 +163,11 @@ const MCPToolConfiguration: React.FC = ({ externalError, externalCanFetch, isEditMode = false, + defaultViewMode = "crud", }) => { const previousToolsRef = useRef([]); const [toolSearchTerm, setToolSearchTerm] = useState(""); - const [viewMode, setViewMode] = useState<"crud" | "flat">("crud"); + const [viewMode, setViewMode] = useState<"crud" | "flat">(defaultViewMode); const hasInitializedRef = useRef(false); const previousSuggestedToolNamesRef = useRef(""); const [expandedTools, setExpandedTools] = useState>(new Set());