fix: get_all_mcp_servers_with_health_and_teams

This commit is contained in:
Yuta Saito 2025-12-26 15:37:41 +09:00
parent 78b5d40664
commit 3a141e642a

View file

@ -2244,100 +2244,80 @@ class MCPServerManager:
Returns:
List of MCP server objects with health and team data
"""
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.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[LiteLLM_MCPServerTable] = []
if prisma_client is not None:
list_mcp_servers = await get_mcp_servers(prisma_client, allowed_server_ids)
async def _check_server_health(server_id: str) -> Optional[LiteLLM_MCPServerTable]:
"""Helper function to check health of a single server"""
server = self.get_mcp_server_by_id(server_id)
if server is None:
verbose_logger.warning(f"MCP Server {server_id} not found")
return None
# If admin, also get all servers from database
if user_api_key_auth and _user_has_admin_view(user_api_key_auth):
all_mcp_servers = await get_all_mcp_servers(prisma_client)
for server in all_mcp_servers:
if server.server_id not in allowed_server_ids:
list_mcp_servers.append(server)
status = "unknown"
health_check_error = None
# Add config.yaml servers
for _server_id, _server_config in self.config_mcp_servers.items():
if _server_id in allowed_server_ids:
list_mcp_servers.append(
LiteLLM_MCPServerTable(
**{
**_server_config.model_dump(),
"created_at": datetime.datetime.now(),
"updated_at": datetime.datetime.now(),
"description": (
_server_config.mcp_info.get("description")
if _server_config.mcp_info
else None
),
"allowed_tools": _server_config.allowed_tools or [],
"mcp_info": _server_config.mcp_info,
"mcp_access_groups": _server_config.access_groups or [],
"extra_headers": _server_config.extra_headers or [],
"command": getattr(_server_config, "command", None),
"args": getattr(_server_config, "args", None) or [],
"env": getattr(_server_config, "env", None) or {},
}
)
# Check if we should skip health check based on auth configuration
should_skip_health_check = False
# Skip if auth_type is oauth2
if server.auth_type == MCPAuth.oauth2:
should_skip_health_check = True
# Skip if auth_type is not none and authentication_token is missing
elif server.auth_type and server.auth_type != MCPAuth.none and not server.authentication_token:
should_skip_health_check = True
if not should_skip_health_check:
extra_headers = {}
if server.static_headers:
extra_headers.update(server.static_headers)
client = self._create_mcp_client(
server=server,
mcp_auth_header=None,
extra_headers=extra_headers,
stdio_env=None,
)
# Get team information for non-admin users
server_to_teams_map: Dict[str, List[Dict[str, str]]] = {}
if (
user_api_key_auth
and not _user_has_admin_view(user_api_key_auth)
and prisma_client is not None
):
teams = await prisma_client.db.litellm_teamtable.find_many(
include={"object_permission": True}
try:
async def _noop(session):
return "ok"
await client.run_with_session(_noop)
status = "healthy"
except Exception as e:
health_check_error = str(e)
status = "unhealthy"
return LiteLLM_MCPServerTable(
**{
**server.model_dump(),
"created_at": datetime.datetime.now(),
"updated_at": datetime.datetime.now(),
"description": (
server.mcp_info.get("description")
if server.mcp_info
else None
),
"allowed_tools": server.allowed_tools or [],
"mcp_info": server.mcp_info,
"mcp_access_groups": server.access_groups or [],
"extra_headers": server.extra_headers or [],
"command": getattr(server, "command", None),
"args": getattr(server, "args", None) or [],
"env": getattr(server, "env", None) or {},
"status": status,
"health_check_error": health_check_error,
}
)
user_teams = []
for team in teams:
if team.members_with_roles:
for member in team.members_with_roles:
if (
"user_id" in member
and member["user_id"] is not None
and member["user_id"] == user_api_key_auth.user_id
):
user_teams.append(team)
# Run health checks concurrently
tasks = [_check_server_health(server_id) for server_id in allowed_server_ids]
results = await asyncio.gather(*tasks)
# Create a mapping of server_id to teams that have access to it
for team in user_teams:
if team.object_permission and team.object_permission.mcp_servers:
for server_id in team.object_permission.mcp_servers:
if server_id not in server_to_teams_map:
server_to_teams_map[server_id] = []
server_to_teams_map[server_id].append(
{
"team_id": team.team_id,
"team_alias": team.team_alias,
"organization_id": team.organization_id,
}
)
## mark invalid servers w/ reason for being invalid
valid_server_ids = self.get_all_mcp_server_ids()
for server in list_mcp_servers:
if server.server_id not in valid_server_ids:
server.status = "unhealthy"
## try adding server to registry to get error
try:
await self.add_update_server(server)
except Exception as e:
server.health_check_error = str(e)
server.health_check_error = "Server is not in in memory registry yet. This could be a temporary sync issue."
# Filter out None results (servers that were not found)
list_mcp_servers = [server for server in results if server is not None]
return list_mcp_servers