From 831eeaf3ab8fc8d52816109c0871251770fec699 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 2 Aug 2025 09:48:05 -0700 Subject: [PATCH] [QA Fixes for MCP] - Ensure MCPs load + don't run a health check everytime we load MCPs on UI (#13228) * qa - mcps should load even if they don't have required fields * fix loading MCPs --- litellm/proxy/_experimental/mcp_server/db.py | 10 ++++-- .../mcp_server/mcp_server_manager.py | 35 +++++++------------ 2 files changed, 21 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 414f8094c32..3d90c99eee2 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -82,14 +82,20 @@ async def get_mcp_servers( """ Returns the matching mcp servers from the db with the server_ids """ - mcp_servers: List[LiteLLM_MCPServerTable] = ( + _mcp_servers: List[LiteLLM_MCPServerTable] = ( await prisma_client.db.litellm_mcpservertable.find_many( where={ "server_id": {"in": server_ids}, } ) ) - return mcp_servers + final_mcp_servers: List[LiteLLM_MCPServerTable] = [] + for _mcp_server in _mcp_servers: + final_mcp_servers.append( + LiteLLM_MCPServerTable(**_mcp_server.model_dump()) + ) + + return final_mcp_servers async def get_mcp_servers_by_verificationtoken( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0abd0232d21..c7ee6ca0612 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -12,13 +12,13 @@ import hashlib import json from typing import Any, Dict, List, Optional, cast +from fastapi import HTTPException from mcp.types import CallToolRequestParams as MCPCallToolRequestParams from mcp.types import CallToolResult from mcp.types import Tool as MCPTool from litellm._logging import verbose_logger from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException -from fastapi import HTTPException from litellm.experimental_mcp_client.client import MCPClient from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, @@ -26,10 +26,10 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( from litellm.proxy._experimental.mcp_server.utils import ( add_server_prefix_to_tool_name, get_server_name_prefix_tool_mcp, + get_server_prefix, is_tool_name_prefixed, normalize_server_name, validate_mcp_server_name, - get_server_prefix, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -81,21 +81,21 @@ def _convert_protocol_version_to_enum(protocol_version: Optional[str | MCPSpecVe MCPSpecVersionType: The enum value """ if not protocol_version: - return MCPSpecVersion.jun_2025 + return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025) # If it's already an MCPSpecVersion enum, return it if isinstance(protocol_version, MCPSpecVersion): - return protocol_version + return cast(MCPSpecVersionType, protocol_version) # If it's a string, try to match it to enum values if isinstance(protocol_version, str): for version in MCPSpecVersion: if version.value == protocol_version: - return version + return cast(MCPSpecVersionType, version) # If no match found, return default verbose_logger.warning(f"Unknown protocol version '{protocol_version}', using default") - return MCPSpecVersion.jun_2025 + return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025) class MCPServerManager: @@ -903,15 +903,18 @@ class MCPServerManager: Returns: List of MCP server objects with health and team data """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy._experimental.mcp_server.db import ( + get_all_mcp_servers, + get_mcp_servers, + ) from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view - from litellm.proxy._experimental.mcp_server.db import get_mcp_servers, get_all_mcp_servers - + from litellm.proxy.proxy_server import prisma_client + # Get allowed server IDs allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth) # Get servers from database - list_mcp_servers = [] + list_mcp_servers: List[LiteLLM_MCPServerTable] = [] if prisma_client is not None: list_mcp_servers = await get_mcp_servers(prisma_client, allowed_server_ids) @@ -974,14 +977,6 @@ class MCPServerManager: "organization_id": team.organization_id }) - # Get health check results if requested - all_health_results = {} - if include_health: - try: - all_health_results = await self.health_check_allowed_servers(user_api_key_auth) - except Exception as e: - verbose_logger.debug(f"Error performing health checks: {e}") - # Map servers to their teams and return with health data from typing import cast return [ @@ -1001,10 +996,6 @@ class MCPServerManager: mcp_access_groups=server.mcp_access_groups if server.mcp_access_groups is not None else [], mcp_info=server.mcp_info, teams=cast(List[Dict[str, str | None]], server_to_teams_map.get(server.server_id, [])), - # Health check status - status=all_health_results.get(server.server_id, {}).get("status", "unknown"), - last_health_check=datetime.datetime.fromisoformat(all_health_results.get(server.server_id, {}).get("last_health_check", datetime.datetime.now().isoformat())) if all_health_results.get(server.server_id, {}).get("last_health_check") else None, - health_check_error=all_health_results.get(server.server_id, {}).get("error"), # Stdio-specific fields command=getattr(server, 'command', None), args=getattr(server, 'args', None) or [],