mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat: add platform mcp compression
This commit is contained in:
parent
dcf1b445e6
commit
7a4bc18271
12 changed files with 1205 additions and 7 deletions
220
litellm/proxy/_experimental/mcp_server/platform_mcp.py
Normal file
220
litellm/proxy/_experimental/mcp_server/platform_mcp.py
Normal 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 ""
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
377
tests/mcp_tests/test_platform_mcp.py
Normal file
377
tests/mcp_tests/test_platform_mcp.py
Normal 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",
|
||||
]
|
||||
|
|
@ -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,
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
171
ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx
Normal file
171
ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx
Normal 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'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;
|
||||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -1112,6 +1112,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
externalIsLoading={isLoadingTools}
|
||||
externalError={toolsError}
|
||||
externalCanFetch={canFetchTools}
|
||||
defaultViewMode="flat"
|
||||
/>
|
||||
</div>
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue