diff --git a/litellm/proxy/_experimental/mcp_server/platform_mcp.py b/litellm/proxy/_experimental/mcp_server/platform_mcp.py index 66ea264e577..6eae8b02d35 100644 --- a/litellm/proxy/_experimental/mcp_server/platform_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/platform_mcp.py @@ -36,7 +36,7 @@ async def get_platform_mcp_settings() -> tuple[bool, int]: if isinstance(param_value, dict): settings.update(param_value) - enabled = settings.get("platform_mcp_enabled") is True + enabled = _coerce_enabled(settings.get("platform_mcp_enabled")) threshold = _coerce_positive_threshold(settings.get("platform_mcp_tool_threshold")) return enabled, threshold @@ -201,11 +201,27 @@ def serialize_tool(tool: Any) -> dict[str, Any]: 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 + if isinstance(value, str): + return value.strip().lower() in {"1", "true", "yes", "on"} + return False + + def _server_match_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 ffae2e169a1..5af1aac1463 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2182,6 +2182,7 @@ 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. @@ -2212,6 +2213,9 @@ if MCP_AVAILABLE: 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() ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e86982307e7..416d275ad24 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -734,6 +734,7 @@ 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 b5b46325f22..96359ac1194 100644 --- a/tests/mcp_tests/test_platform_mcp.py +++ b/tests/mcp_tests/test_platform_mcp.py @@ -99,6 +99,43 @@ def test_platform_mcp_advertises_tool_list_changed_capability(): assert options.capabilities.tools.listChanged is True +@pytest.mark.asyncio +async def test_platform_mcp_settings_accept_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, "prisma_client", None) + + assert await platform_mcp.get_platform_mcp_settings() == (True, 12) + + +@pytest.mark.asyncio +async def test_platform_mcp_threshold_ignores_bool_values(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, "prisma_client", None) + + assert await platform_mcp.get_platform_mcp_settings() == ( + True, + platform_mcp.DEFAULT_PLATFORM_MCP_TOOL_THRESHOLD, + ) + + @pytest.mark.asyncio async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatch): normal_tools = [_tool(f"tool_{idx}") for idx in range(11)] @@ -127,6 +164,36 @@ async def test_platform_mcp_compresses_aggregate_tools_over_threshold(monkeypatc 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): normal_tools = [_tool(f"tool_{idx}") for idx in range(11)]