feat: add platform mcp compression

This commit is contained in:
Krrish Dholakia 2026-06-22 18:02:59 -07:00
parent dcf1b445e6
commit 7a4bc18271
12 changed files with 1205 additions and 7 deletions

View file

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

View file

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

View file

@ -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.",

View file

@ -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",
]

View file

@ -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(<PlatformMCPTab accessToken="token" />);
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(<PlatformMCPTab accessToken="token" />);
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(<PlatformMCPTab accessToken="token" />);
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,
);
});
});
});

View file

@ -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<PlatformMCPTabProps> = ({ 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 (
<div className="flex justify-center py-12">
<Spin />
</div>
);
}
return (
<div className="space-y-6 p-4">
<div className="flex items-center gap-2 rounded-lg border border-amber-200 bg-amber-50 px-4 py-3 text-sm text-amber-800">
<ExperimentOutlined className="flex-shrink-0 text-amber-600" />
<span>
Platform MCP is a pre-v0 feature. The dashboard control is only available to proxy admins and should not be
used in production.
</span>
</div>
<div>
<Text className="text-lg font-semibold">Platform MCP</Text>
<p className="mt-1 text-sm text-gray-500">
When enabled, LiteLLM compresses aggregate MCP tools/list responses only after the caller&apos;s filtered tool
count is over the configured threshold.
</p>
</div>
<Card>
<div className="flex flex-col gap-4 lg:flex-row lg:items-center lg:justify-between">
<div>
<Text className="font-medium">Enable Platform MCP compression</Text>
<p className="mb-0 mt-1 text-sm text-gray-500">
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.
</p>
</div>
<Switch checked={enabled} loading={savingEnabled} onChange={handleToggle} />
</div>
</Card>
<Card>
<div className="flex flex-col gap-4 lg:flex-row lg:items-center lg:justify-between">
<div>
<Text className="font-medium">Compression threshold</Text>
<p className="mb-0 mt-1 text-sm text-gray-500">
Default is 10 tools. Compression starts when the final accessible tool count is greater than this value.
</p>
</div>
<div className="flex items-center gap-2">
<InputNumber min={1} value={threshold} onChange={(value) => setThreshold(Number(value || 1))} />
<Button type="primary" icon={<SaveOutlined />} loading={savingThreshold} onClick={handleSaveThreshold}>
Save
</Button>
</div>
</div>
</Card>
<div className="grid grid-cols-1 gap-4 md:grid-cols-2">
<Card>
<div className="flex items-start gap-3">
<ToolOutlined className="mt-1 text-gray-500" />
<div>
<Text className="font-mono font-semibold text-blue-600">list_servers</Text>
<p className="mb-0 mt-1 text-sm text-gray-500">
Returns MCP server names and descriptions for servers the key can access.
</p>
</div>
</div>
</Card>
<Card>
<div className="flex items-start gap-3">
<ToolOutlined className="mt-1 text-gray-500" />
<div>
<Text className="font-mono font-semibold text-blue-600">enable_server</Text>
<p className="mb-0 mt-1 text-sm text-gray-500">
Returns full tool definitions for one accessible MCP server.
</p>
</div>
</div>
</Card>
</div>
</div>
);
};
export default PlatformMCPTab;

View file

@ -62,12 +62,21 @@ vi.mock("./mcp_tool_configuration", () => ({
>
Disable all tools
</button>
<button
type="button"
onClick={() => {
onToolAllowlistInteraction?.();
onAllowedToolsChange?.(["read_user"]);
}}
>
Enable read_user only
</button>
</div>
),
}));
vi.mock("./mcp_connection_status", () => ({
default: ({ tools }: { tools?: any[] }) => (
default: ({ tools }: { tools?: unknown[] }) => (
<div data-testid="mcp-connection-status" data-tool-count={tools?.length ?? 0} />
),
}));
@ -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", () => {

View file

@ -1112,6 +1112,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
externalIsLoading={isLoadingTools}
externalError={toolsError}
externalCanFetch={canFetchTools}
defaultViewMode="flat"
/>
</div>

View file

@ -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(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} userRole="proxy_admin" />
</QueryClientProvider>,
);
await waitFor(() => {
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
expect(screen.getByRole("tab", { name: /Platform MCP/ })).toBeInTheDocument();
rerender(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} userRole="proxy_admin_viewer" />
</QueryClientProvider>,
);
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 = [

View file

@ -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<MCPServerProps> = ({ accessToken, userRole, userID })
</span>
</Tab>
)}
{isProxyAdminRole(userRole) && (
<Tab>
<span className="flex items-center gap-2">
Platform MCP <NewBadge />
</span>
</Tab>
)}
</div>
</TabList>
<TabPanels>
@ -663,6 +671,11 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
<MCPSubmissionsTab accessToken={accessToken} />
</TabPanel>
)}
{isProxyAdminRole(userRole) && (
<TabPanel>
<PlatformMCPTab accessToken={accessToken} />
</TabPanel>
)}
</TabPanels>
</TabGroup>

View file

@ -30,6 +30,69 @@ const renderToolConfiguration = (onAllowedToolsChange = vi.fn()) => {
};
describe("MCPToolConfiguration", () => {
it("can start new-server onboarding in the flat checklist view", async () => {
render(
<MCPToolConfiguration
accessToken="token"
formValues={{ url: "https://example.com/mcp", transport: "http", auth_type: "none" }}
allowedTools={[]}
existingAllowedTools={null}
onAllowedToolsChange={vi.fn()}
toolNameToDisplayName={{}}
toolNameToDescription={{}}
onToolNameToDisplayNameChange={vi.fn()}
onToolNameToDescriptionChange={vi.fn()}
externalTools={tools}
externalCanFetch
defaultViewMode="flat"
/>,
);
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<string[]>([]);
return (
<MCPToolConfiguration
accessToken="token"
formValues={{ url: "https://example.com/mcp", transport: "http", auth_type: "none" }}
allowedTools={allowedTools}
existingAllowedTools={null}
onAllowedToolsChange={(nextAllowedTools) => {
onAllowedToolsChange(nextAllowedTools);
setAllowedTools(nextAllowedTools);
}}
toolNameToDisplayName={{}}
toolNameToDescription={{}}
onToolNameToDisplayNameChange={vi.fn()}
onToolNameToDescriptionChange={vi.fn()}
externalTools={tools}
externalCanFetch
defaultViewMode="flat"
/>
);
};
render(<Wrapper />);
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();

View file

@ -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<ToolRowProps> = ({
>
<div className="p-4 cursor-pointer" onClick={() => onToggle(tool.name)}>
<div className="flex items-start gap-3">
<Checkbox checked={isEnabled} onChange={() => onToggle(tool.name)} />
<Checkbox
checked={isEnabled}
onClick={(e) => e.stopPropagation()}
onChange={() => onToggle(tool.name)}
/>
<div className="flex-1">
<div className="flex items-center gap-2">
<Text className="font-medium text-gray-900">{toolNameToDisplayName[tool.name] || tool.name}</Text>
@ -158,10 +163,11 @@ const MCPToolConfiguration: React.FC<MCPToolConfigurationProps> = ({
externalError,
externalCanFetch,
isEditMode = false,
defaultViewMode = "crud",
}) => {
const previousToolsRef = useRef<ToolEntry[]>([]);
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<string>("");
const [expandedTools, setExpandedTools] = useState<Set<string>>(new Set());