feat: expose platform mcp server

This commit is contained in:
Krrish Dholakia 2026-06-22 19:24:31 -07:00
parent 0463a74905
commit 6ad764423f
8 changed files with 463 additions and 434 deletions

View file

@ -569,6 +569,12 @@ class MCPServerManager:
as the cooldown avoids reconnecting on every gateway initialize when
upstream returns empty or fails).
"""
from litellm.proxy._experimental.mcp_server.platform_mcp import (
is_platform_mcp_server,
)
if is_platform_mcp_server(server):
return
if server.spec_path:
return
if server.instructions and server.instructions.strip():
@ -636,7 +642,16 @@ class MCPServerManager:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
"""
return self.config_mcp_servers | self.registry
from litellm.proxy._experimental.mcp_server.platform_mcp import (
PLATFORM_MCP_SERVER_ID,
build_platform_mcp_server,
)
return (
self.config_mcp_servers
| self.registry
| {PLATFORM_MCP_SERVER_ID: build_platform_mcp_server()}
)
async def load_servers_from_config(
self,
@ -1391,6 +1406,18 @@ class MCPServerManager:
if not in_toolset_scope:
combined_servers.update(allow_all_server_ids)
from litellm.proxy._experimental.mcp_server.platform_mcp import (
get_platform_mcp_enabled,
is_platform_mcp_server_identifier,
)
if not await get_platform_mcp_enabled():
combined_servers = {
server_id
for server_id in combined_servers
if not is_platform_mcp_server_identifier(server_id)
}
# For anonymous callers (no user_id, no role), also surface any
# servers the operator has opted into upstream-delegated auth.
# These servers handle their own auth at the upstream level, so
@ -2054,6 +2081,10 @@ class MCPServerManager:
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
from litellm.proxy._experimental.mcp_server.platform_mcp import (
build_platform_mcp_tools,
is_platform_mcp_server,
)
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"_get_tools_from_server for {server.name}...")
@ -2061,6 +2092,13 @@ class MCPServerManager:
client = None
try:
if is_platform_mcp_server(server):
return self._create_prefixed_tools(
build_platform_mcp_tools(),
server,
add_prefix=add_prefix,
)
# Tool *listing* must not be blocked by missing per-user env vars —
# the server's tools should still appear so the client connects. The
# friendly "missing vars" error is raised only on the tool-*call*
@ -3689,6 +3727,9 @@ class MCPServerManager:
"""
start_time = datetime.datetime.now()
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)
from litellm.proxy._experimental.mcp_server.platform_mcp import (
is_platform_mcp_server,
)
#########################################################
# Pre MCP Tool Call Hook
@ -3726,6 +3767,25 @@ class MCPServerManager:
mcp_server, oauth2_headers, user_api_key_auth
)
if is_platform_mcp_server(mcp_server):
from litellm.proxy._experimental.mcp_server.platform_mcp import (
normalize_platform_mcp_tool_name,
)
from litellm.proxy._experimental.mcp_server.server import (
_handle_platform_mcp_tool_call,
)
return await _handle_platform_mcp_tool_call(
name=normalize_platform_mcp_tool_name(name),
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=[mcp_server.server_id],
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
# For OpenAPI servers, call the tool handler directly instead of via MCP client
if mcp_server.spec_path:
verbose_logger.debug(
@ -4251,6 +4311,19 @@ class MCPServerManager:
last_health_check=datetime.now(),
)
from litellm.proxy._experimental.mcp_server.platform_mcp import (
is_platform_mcp_server,
)
if is_platform_mcp_server(server):
return self._build_mcp_server_table(server).model_copy(
update={
"status": "healthy",
"health_check_error": None,
"last_health_check": datetime.now(),
}
)
status: Literal["healthy", "unhealthy", "unknown"] = "unknown"
health_check_error = None

View file

@ -1,8 +1,6 @@
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:
@ -11,20 +9,25 @@ except ImportError:
MCPTool = None # type: ignore
DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD = 10
PLATFORM_MCP_SERVER_ID = "platform_mcp"
PLATFORM_MCP_SERVER_NAME = "platform_mcp"
PLATFORM_MCP_SERVER_DESCRIPTION = (
"Built-in LiteLLM Platform MCP server for discovering and invoking accessible "
"downstream MCP servers."
)
PLATFORM_MCP_LIST_SERVERS_TOOL_NAME = "list_servers"
PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME = "enable_server"
PLATFORM_MCP_GET_SERVER_TOOLS_TOOL_NAME = "get_server_tools"
PLATFORM_MCP_CALL_TOOL_NAME = "call_tool"
PLATFORM_MCP_TOOL_NAMES = frozenset(
{
PLATFORM_MCP_LIST_SERVERS_TOOL_NAME,
PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME,
PLATFORM_MCP_GET_SERVER_TOOLS_TOOL_NAME,
PLATFORM_MCP_CALL_TOOL_NAME,
}
)
_enabled_servers_by_session: "weakref.WeakKeyDictionary[Any, frozenset[str]]" = weakref.WeakKeyDictionary()
async def get_platform_mcp_settings() -> tuple[bool, int]:
async def get_platform_mcp_enabled() -> bool:
from litellm.proxy.proxy_server import general_settings, prisma_client
settings = dict(general_settings or {})
@ -36,9 +39,49 @@ async def get_platform_mcp_settings() -> tuple[bool, int]:
if isinstance(param_value, dict):
settings.update(param_value)
enabled = _coerce_enabled(settings.get("platform_mcp_enabled"))
threshold = _coerce_positive_threshold(settings.get("platform_mcp_tool_threshold"))
return enabled, threshold
return _coerce_enabled(settings.get("platform_mcp_enabled"))
def build_platform_mcp_server() -> MCPServer:
from litellm.types.mcp import MCPTransport
return MCPServer(
server_id=PLATFORM_MCP_SERVER_ID,
name=PLATFORM_MCP_SERVER_NAME,
alias=PLATFORM_MCP_SERVER_NAME,
server_name=PLATFORM_MCP_SERVER_NAME,
url=None,
transport=MCPTransport.http,
auth_type=None,
mcp_info={
"description": PLATFORM_MCP_SERVER_DESCRIPTION,
"server_name": PLATFORM_MCP_SERVER_NAME,
"is_platform_mcp": True,
},
allow_all_keys=False,
available_on_public_internet=True,
)
def is_platform_mcp_server_identifier(value: Optional[str]) -> bool:
if not value:
return False
return value in {
PLATFORM_MCP_SERVER_ID,
PLATFORM_MCP_SERVER_NAME,
}
def is_platform_mcp_server(server: Optional[MCPServer]) -> bool:
if server is None:
return False
return is_platform_mcp_server_identifier(server.server_id) or bool(
server.mcp_info and server.mcp_info.get("is_platform_mcp") is True
)
def without_platform_mcp_servers(servers: Iterable[MCPServer]) -> list[MCPServer]:
return [server for server in servers if not is_platform_mcp_server(server)]
def build_platform_mcp_tools() -> list[Any]:
@ -50,7 +93,7 @@ def build_platform_mcp_tools() -> list[Any]:
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."
"name and description so you can choose which server to inspect."
),
inputSchema={
"type": "object",
@ -59,7 +102,7 @@ def build_platform_mcp_tools() -> list[Any]:
},
),
MCPTool(
name=PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME,
name=PLATFORM_MCP_GET_SERVER_TOOLS_TOOL_NAME,
description=(
"Return full tool definitions for one accessible MCP server. "
"Use a server name returned by list_servers."
@ -76,63 +119,51 @@ def build_platform_mcp_tools() -> list[Any]:
"additionalProperties": False,
},
),
MCPTool(
name=PLATFORM_MCP_CALL_TOOL_NAME,
description=(
"Call a tool on one accessible MCP server. Use tool names and input "
"schemas returned by get_server_tools."
),
inputSchema={
"type": "object",
"properties": {
"server_name": {
"type": "string",
"description": "The MCP server name returned by list_servers.",
},
"tool_name": {
"type": "string",
"description": "The tool name returned by get_server_tools.",
},
"arguments": {
"type": "object",
"description": "Arguments to pass to the downstream MCP tool.",
"additionalProperties": True,
},
},
"required": ["server_name", "tool_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
if name in PLATFORM_MCP_TOOL_NAMES:
return True
prefix = f"{PLATFORM_MCP_SERVER_NAME}-"
return name.startswith(prefix) and name[len(prefix) :] in PLATFORM_MCP_TOOL_NAMES
def extract_enable_server_name(arguments: Optional[dict[str, Any]]) -> Optional[str]:
def normalize_platform_mcp_tool_name(name: str) -> str:
prefix = f"{PLATFORM_MCP_SERVER_NAME}-"
if name.startswith(prefix):
return name[len(prefix) :]
return name
def extract_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")
@ -200,20 +231,6 @@ def serialize_tool(tool: Any) -> dict[str, Any]:
return serialized_tool
def _coerce_positive_threshold(value: Any) -> int:
if isinstance(value, bool):
return DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD
if isinstance(value, int) and value > 0:
return value
if isinstance(value, str):
value = value.strip()
if value.isdigit():
threshold = int(value)
if threshold > 0:
return threshold
return DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD
def _coerce_enabled(value: Any) -> bool:
if isinstance(value, bool):
return value
@ -222,10 +239,6 @@ def _coerce_enabled(value: Any) -> bool:
return False
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

View file

@ -673,25 +673,6 @@ 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
@ -2182,7 +2163,6 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]] = None,
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
enable_platform_mcp_compression: bool = True,
) -> List[MCPTool]:
"""
List all available MCP tools.
@ -2202,28 +2182,6 @@ 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()
)
platform_mcp_enabled = (
platform_mcp_enabled and enable_platform_mcp_compression
)
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 = []
@ -2231,7 +2189,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=effective_mcp_servers,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
@ -2247,22 +2205,6 @@ 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(
@ -2276,26 +2218,44 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]],
) -> CallToolResult:
from litellm.proxy._experimental.mcp_server.platform_mcp import (
PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME,
PLATFORM_MCP_CALL_TOOL_NAME,
PLATFORM_MCP_GET_SERVER_TOOLS_TOOL_NAME,
PLATFORM_MCP_LIST_SERVERS_TOOL_NAME,
enable_server_for_session,
extract_enable_server_name,
get_platform_mcp_settings,
extract_server_name,
get_platform_mcp_enabled,
is_platform_mcp_server,
serialize_server_tool_response,
serialize_servers_response,
without_platform_mcp_servers,
)
platform_mcp_enabled, _ = await get_platform_mcp_settings()
if not platform_mcp_enabled:
if not await get_platform_mcp_enabled():
return CallToolResult(
content=[TextContent(text="Platform MCP is disabled.", type="text")],
isError=True,
)
allowed_servers = await _get_allowed_mcp_servers(
platform_scope_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
)
if not any(is_platform_mcp_server(server) for server in platform_scope_servers):
return CallToolResult(
content=[
TextContent(
text="Platform MCP is not available to this key.",
type="text",
)
],
isError=True,
)
allowed_servers = without_platform_mcp_servers(
await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=None,
)
)
if name == PLATFORM_MCP_LIST_SERVERS_TOOL_NAME:
return CallToolResult(
@ -2305,13 +2265,16 @@ if MCP_AVAILABLE:
isError=False,
)
if name != PLATFORM_MCP_ENABLE_SERVER_TOOL_NAME:
if name not in {
PLATFORM_MCP_GET_SERVER_TOOLS_TOOL_NAME,
PLATFORM_MCP_CALL_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)
requested_server_name = extract_server_name(arguments)
if requested_server_name is None:
return CallToolResult(
content=[
@ -2351,6 +2314,46 @@ if MCP_AVAILABLE:
selected_server_name = (
selected_server.alias or selected_server.server_name or selected_server.name
)
if name == PLATFORM_MCP_CALL_TOOL_NAME:
tool_name = arguments.get("tool_name") if arguments else None
if not isinstance(tool_name, str) or not tool_name.strip():
return CallToolResult(
content=[
TextContent(
text="Missing required argument: tool_name",
type="text",
)
],
isError=True,
)
tool_arguments = arguments.get("arguments") if arguments else None
if tool_arguments is None:
tool_arguments = {}
if not isinstance(tool_arguments, dict):
return CallToolResult(
content=[
TextContent(
text="Argument 'arguments' must be an object.",
type="text",
)
],
isError=True,
)
return await execute_mcp_tool(
name=tool_name,
arguments=tool_arguments,
allowed_mcp_servers=[selected_server],
start_time=datetime.now(),
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
requested_server_id=selected_server.server_id,
)
tools = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
@ -2359,9 +2362,6 @@ if MCP_AVAILABLE:
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=[
@ -2376,18 +2376,6 @@ if MCP_AVAILABLE:
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

@ -2341,11 +2341,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
)
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.",
description="If True, enables the built-in Platform MCP server.",
)
mcp_trusted_proxy_ranges: Optional[List[str]] = Field(
None,

View file

@ -734,7 +734,6 @@ if MCP_AVAILABLE:
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
enable_platform_mcp_compression=False,
)
dumped_tools = [dict(tool) for tool in tools]

View file

@ -4,7 +4,7 @@ import sys
from typing import Optional
import pytest
from mcp.types import Tool
from mcp.types import CallToolResult, TextContent, Tool
sys.path.insert(0, os.path.abspath("../../.."))
@ -30,7 +30,7 @@ def _rich_tool(name: str) -> Tool:
description=f"{name} description",
inputSchema={"type": "object", "properties": {}},
outputSchema={"type": "object", "properties": {"id": {"type": "string"}}},
_meta={"source": "platform-mcp-test"},
_meta={"source": "platform_mcp-test"},
)
@ -56,12 +56,12 @@ 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 _enabled_platform_mcp() -> bool:
return True
async def _disabled_platform_settings() -> tuple[bool, int]:
return False, 10
async def _disabled_platform_mcp() -> bool:
return False
@pytest.mark.asyncio
@ -83,8 +83,8 @@ async def test_platform_mcp_disabled_returns_normal_tools(monkeypatch):
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_disabled_platform_settings,
"get_platform_mcp_enabled",
_disabled_platform_mcp,
)
tools = await mcp_server_module._list_mcp_tools()
@ -100,44 +100,27 @@ def test_platform_mcp_advertises_tool_list_changed_capability():
@pytest.mark.asyncio
async def test_platform_mcp_settings_accept_config_field_string_values(monkeypatch):
async def test_platform_mcp_enabled_accepts_config_field_string_values(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(
proxy_server,
"general_settings",
{
"platform_mcp_enabled": "true",
"platform_mcp_tool_threshold": "12",
},
)
monkeypatch.setattr(proxy_server, "general_settings", {"platform_mcp_enabled": "true"})
monkeypatch.setattr(proxy_server, "prisma_client", None)
assert await platform_mcp.get_platform_mcp_settings() == (True, 12)
assert await platform_mcp.get_platform_mcp_enabled() is True
@pytest.mark.asyncio
async def test_platform_mcp_threshold_ignores_bool_values(monkeypatch):
async def test_platform_mcp_enabled_defaults_to_false(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(
proxy_server,
"general_settings",
{
"platform_mcp_enabled": True,
"platform_mcp_tool_threshold": True,
},
)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "prisma_client", None)
assert await platform_mcp.get_platform_mcp_settings() == (
True,
platform_mcp.DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD,
)
assert await platform_mcp.get_platform_mcp_enabled() is False
@pytest.mark.asyncio
async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatch):
async def test_platform_mcp_enabled_does_not_change_aggregate_tool_list(monkeypatch):
normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]
async def fake_get_tools(**kwargs):
@ -155,47 +138,17 @@ async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatc
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_enabled_platform_settings,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
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_catalog_call_can_bypass_compression(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(
enable_platform_mcp_compression=False
)
assert tools == normal_tools
@pytest.mark.asyncio
async def test_platform_mcp_does_not_compress_scoped_server_tools(monkeypatch):
async def test_platform_mcp_enabled_does_not_change_scoped_server_tools(monkeypatch):
normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]
async def fake_get_tools(**kwargs):
@ -213,8 +166,8 @@ async def test_platform_mcp_does_not_compress_scoped_server_tools(monkeypatch):
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_enabled_platform_settings,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
tools = await mcp_server_module._list_mcp_tools(mcp_servers=["servicenow"])
@ -223,14 +176,11 @@ async def test_platform_mcp_does_not_compress_scoped_server_tools(monkeypatch):
@pytest.mark.asyncio
async def test_platform_mcp_enabled_session_returns_meta_tools_and_enabled_server_tools(
async def test_platform_mcp_virtual_server_returns_only_platform_tools(
monkeypatch,
):
selected_tools = [_tool("servicenow_get_ticket")]
async def fake_get_tools(**kwargs):
assert kwargs["mcp_servers"] == ["servicenow"]
return selected_tools
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return [platform_mcp.build_platform_mcp_server()]
monkeypatch.setattr(
mcp_server_module,
@ -239,33 +189,32 @@ async def test_platform_mcp_enabled_session_returns_meta_tools_and_enabled_serve
)
monkeypatch.setattr(
mcp_server_module,
"_get_tools_from_mcp_servers",
fake_get_tools,
"_get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_enabled_platform_settings,
)
monkeypatch.setattr(
platform_mcp,
"get_enabled_server_names_for_session",
lambda _session: ("servicenow",),
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
tools = await mcp_server_module._list_mcp_tools()
tools = await mcp_server_module._list_mcp_tools(mcp_servers=["platform_mcp"])
assert [tool.name for tool in tools] == [
"list_servers",
"enable_server",
"servicenow_get_ticket",
"platform_mcp-list_servers",
"platform_mcp-get_server_tools",
"platform_mcp-call_tool",
]
@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")]
return [
platform_mcp.build_platform_mcp_server(),
_server(name="servicenow"),
_server(name="github", description="Code"),
]
monkeypatch.setattr(
mcp_server_module,
@ -274,8 +223,8 @@ async def test_platform_mcp_list_servers_returns_names_and_descriptions(monkeypa
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_enabled_platform_settings,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
result = await mcp_server_module._handle_platform_mcp_tool_call(
@ -300,82 +249,9 @@ async def test_platform_mcp_list_servers_returns_names_and_descriptions(monkeypa
@pytest.mark.asyncio
async def test_platform_mcp_enable_server_returns_selected_server_tool_definitions(
async def test_platform_mcp_get_server_tools_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
@ -383,18 +259,20 @@ async def test_platform_mcp_enable_server_updates_same_session_list_tools(monkey
async def send_tool_list_changed(self):
self.tool_list_changed_count += 1
platform_server = platform_mcp.build_platform_mcp_server()
selected_server = _server(name="servicenow")
selected_tools = [_tool("servicenow_get_ticket")]
selected_tools = [_rich_tool("servicenow_get_ticket")]
normal_tools = [_tool("github_list_issues")]
requested_servers = []
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return [selected_server]
return [platform_server, selected_server]
async def fake_get_tools(**kwargs):
requested_servers.append(kwargs["mcp_servers"])
return selected_tools
session = FakeSession()
if kwargs["mcp_servers"] == ["servicenow"]:
return selected_tools
return normal_tools
monkeypatch.setattr(
mcp_server_module,
@ -413,14 +291,14 @@ async def test_platform_mcp_enable_server_updates_same_session_list_tools(monkey
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_settings",
_enabled_platform_settings,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
platform_mcp._enabled_servers_by_session.clear()
session = FakeSession()
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",
result = await mcp_server_module._handle_platform_mcp_tool_call(
name="get_server_tools",
arguments={"server_name": "servicenow"},
user_api_key_auth=UserAPIKeyAuth(api_key="test"),
mcp_auth_header=None,
@ -432,13 +310,143 @@ async def test_platform_mcp_enable_server_updates_same_session_list_tools(monkey
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",
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"},
}
]
assert session.tool_list_changed_count == 0
assert requested_servers == [["servicenow"], None]
assert tools == normal_tools
@pytest.mark.asyncio
async def test_platform_mcp_call_tool_dispatches_to_selected_server(monkeypatch):
platform_server = platform_mcp.build_platform_mcp_server()
selected_server = _server(name="servicenow")
auth = UserAPIKeyAuth(api_key="test")
calls = []
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return [platform_server, selected_server]
async def fake_execute_mcp_tool(**kwargs):
calls.append(kwargs)
return CallToolResult(
content=[TextContent(type="text", text="ticket result")],
isError=False,
)
monkeypatch.setattr(
mcp_server_module,
"_get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
)
monkeypatch.setattr(
mcp_server_module,
"execute_mcp_tool",
fake_execute_mcp_tool,
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
result = await mcp_server_module._handle_platform_mcp_tool_call(
name="call_tool",
arguments={
"server_name": "servicenow",
"tool_name": "servicenow_get_ticket",
"arguments": {"ticket_id": "INC-1"},
},
user_api_key_auth=auth,
mcp_auth_header=None,
mcp_servers=["platform_mcp"],
mcp_server_auth_headers=None,
oauth2_headers=None,
raw_headers=None,
)
assert result.isError is False
assert result.content[0].text == "ticket result"
assert calls == [
{
"name": "servicenow_get_ticket",
"arguments": {"ticket_id": "INC-1"},
"allowed_mcp_servers": [selected_server],
"start_time": calls[0]["start_time"],
"user_api_key_auth": auth,
"mcp_auth_header": None,
"mcp_server_auth_headers": None,
"oauth2_headers": None,
"raw_headers": None,
"requested_server_id": selected_server.server_id,
}
]
@pytest.mark.asyncio
async def test_platform_mcp_call_tool_rejects_unavailable_downstream_server(
monkeypatch,
):
platform_server = platform_mcp.build_platform_mcp_server()
selected_server = _server(name="servicenow")
calls = []
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return [platform_server, selected_server]
async def fake_execute_mcp_tool(**kwargs):
calls.append(kwargs)
return CallToolResult(content=[], isError=False)
monkeypatch.setattr(
mcp_server_module,
"_get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
)
monkeypatch.setattr(
mcp_server_module,
"execute_mcp_tool",
fake_execute_mcp_tool,
)
monkeypatch.setattr(
platform_mcp,
"get_platform_mcp_enabled",
_enabled_platform_mcp,
)
result = await mcp_server_module._handle_platform_mcp_tool_call(
name="call_tool",
arguments={
"server_name": "github",
"tool_name": "list_issues",
"arguments": {},
},
user_api_key_auth=UserAPIKeyAuth(api_key="test"),
mcp_auth_header=None,
mcp_servers=["platform_mcp"],
mcp_server_auth_headers=None,
oauth2_headers=None,
raw_headers=None,
)
assert result.isError is True
assert "not available to this key" in result.content[0].text
assert calls == []

View file

@ -16,53 +16,33 @@ describe("PlatformMCPTab", () => {
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 () => {
it("shows the pre-v0 warning and platform MCP tools", async () => {
render(<PlatformMCPTab accessToken="token" />);
expect(await screen.findByText(/This can change unexpectedly/i)).toBeInTheDocument();
expect(screen.getByText(/product@berri.ai/i)).toBeInTheDocument();
expect(screen.getByText(/tools\/list responses/i)).toBeInTheDocument();
expect(screen.getByText(/platform-managed MCP discovery and tool calling/i)).toBeInTheDocument();
expect(screen.getByText(/list_servers, get_server_tools, and/i)).toBeInTheDocument();
expect(screen.getByText("list_servers")).toBeInTheDocument();
expect(screen.getByText("enable_server")).toBeInTheDocument();
expect(screen.getByText("get_server_tools")).toBeInTheDocument();
expect(screen.getByText("call_tool")).toBeInTheDocument();
expect(screen.queryByText("search_tools")).not.toBeInTheDocument();
expect(screen.queryByText("sandbox_execute")).not.toBeInTheDocument();
expect(networking.getConfigFieldSetting).toHaveBeenCalledWith("token", "platform_mcp_enabled");
});
it("updates the enabled setting through the existing config field endpoint", async () => {
render(<PlatformMCPTab accessToken="token" />);
const toggle = await screen.findByRole("switch");
const toggle = await screen.findByRole("switch", { name: /enable platform mcp/i });
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,
);
expect(networking.updateConfigFieldSetting).toHaveBeenCalledWith("token", "platform_mcp_enabled", false);
});
});
});

View file

@ -1,26 +1,24 @@
import React, { useEffect, useState } from "react";
import { ExperimentOutlined, SaveOutlined, ToolOutlined } from "@ant-design/icons";
import { Button, Card, InputNumber, Spin, Switch, Typography } from "antd";
import { ExperimentOutlined, ToolOutlined } from "@ant-design/icons";
import { Card, 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 => {
const getFieldValue = (response: unknown, fallback: boolean): boolean => {
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;
return Boolean(field.field_value);
}
if (field.field_default_value !== undefined && field.field_default_value !== null) {
return field.field_default_value as boolean | number;
return Boolean(field.field_default_value);
}
}
return fallback;
@ -29,9 +27,7 @@ const getFieldValue = (response: unknown, fallback: boolean | number): boolean |
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 () => {
@ -42,13 +38,8 @@ const PlatformMCPTab: React.FC<PlatformMCPTabProps> = ({ accessToken }) => {
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);
const enabledResponse = await getConfigFieldSetting(accessToken, PLATFORM_MCP_ENABLED_FIELD);
setEnabled(getFieldValue(enabledResponse, false));
} catch (error) {
console.error("Failed to load Platform MCP settings:", error);
} finally {
@ -72,18 +63,6 @@ const PlatformMCPTab: React.FC<PlatformMCPTabProps> = ({ accessToken }) => {
}
};
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">
@ -105,42 +84,24 @@ const PlatformMCPTab: React.FC<PlatformMCPTabProps> = ({ accessToken }) => {
<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.
When enabled, LiteLLM exposes platform-managed MCP discovery and tool calling through the proxy.
</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>
<Text className="font-medium">Enable Platform MCP</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.
Disabled leaves existing MCP behavior unchanged. Enabled makes list_servers, get_server_tools, and
call_tool available for keys with platform_mcp access.
</p>
</div>
<Switch checked={enabled} loading={savingEnabled} onChange={handleToggle} />
<Switch aria-label="Enable Platform MCP" 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">
<div className="grid grid-cols-1 gap-4 md:grid-cols-3">
<Card>
<div className="flex items-start gap-3">
<ToolOutlined className="mt-1 text-gray-500" />
@ -156,13 +117,24 @@ const PlatformMCPTab: React.FC<PlatformMCPTabProps> = ({ accessToken }) => {
<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>
<Text className="font-mono font-semibold text-blue-600">get_server_tools</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>
<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">call_tool</Text>
<p className="mb-0 mt-1 text-sm text-gray-500">
Calls a tool on an accessible MCP server through the platform.
</p>
</div>
</div>
</Card>
</div>
</div>
);