diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3eaac0fafe4..65f772d1664 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/platform_mcp.py b/litellm/proxy/_experimental/mcp_server/platform_mcp.py index 6eae8b02d35..4e7ca41a99f 100644 --- a/litellm/proxy/_experimental/mcp_server/platform_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/platform_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f263c851723..0a0a29dcbec 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0bc373cc281..585bb573b9d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 416d275ad24..e86982307e7 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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] diff --git a/tests/mcp_tests/test_platform_mcp.py b/tests/mcp_tests/test_platform_mcp.py index 96359ac1194..7d5c4d0083a 100644 --- a/tests/mcp_tests/test_platform_mcp.py +++ b/tests/mcp_tests/test_platform_mcp.py @@ -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 == [] diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx index c1798da8b9d..ff1f1d5bd68 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.test.tsx @@ -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(); 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(); - 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(); - - 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); }); }); }); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx index 2e5e5e83770..d33f44caa6a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/PlatformMCPTab.tsx @@ -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 = ({ 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 = ({ 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 = ({ 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 (
@@ -105,42 +84,24 @@ const PlatformMCPTab: React.FC = ({ accessToken }) => {
Platform MCP

- 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.

- Enable Platform MCP compression + Enable Platform MCP

- 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.

- +
- -
-
- Compression threshold -

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

-
-
- setThreshold(Number(value || 1))} /> - -
-
-
- -
+
@@ -156,13 +117,24 @@ const PlatformMCPTab: React.FC = ({ accessToken }) => {
- enable_server + get_server_tools

Returns full tool definitions for one accessible MCP server.

+ +
+ +
+ call_tool +

+ Calls a tool on an accessible MCP server through the platform. +

+
+
+
);