mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat: expose platform mcp server
This commit is contained in:
parent
0463a74905
commit
6ad764423f
8 changed files with 463 additions and 434 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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'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>
|
||||
);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue