[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
This commit is contained in:
Ishaan Jaff 2025-08-02 09:48:05 -07:00 • committed by GitHub
parent a107a4bdba
commit 831eeaf3ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 21 additions and 24 deletions

View file

@ -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(

View file

@ -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 [],