mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: get_all_mcp_servers_with_health_and_teams
This commit is contained in:
parent
78b5d40664
commit
3a141e642a
1 changed files with 63 additions and 83 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue