feat: optimize MCP server listing by separating health checks

This commit is contained in:
Yuta Saito
2025-12-31 06:55:47 +09:00
parent dce3d0db09
commit 7fdea85b5c
9 changed files with 653 additions and 36 deletions
@@ -2186,8 +2186,12 @@ class MCPServerManager:
async def _noop(session):
return "ok"
await client.run_with_session(_noop)
# Add timeout wrapper to prevent hanging
await asyncio.wait_for(client.run_with_session(_noop), timeout=10.0)
status = "healthy"
except asyncio.TimeoutError:
health_check_error = "Health check timed out after 10 seconds"
status = "unhealthy"
except Exception as e:
health_check_error = str(e)
status = "unhealthy"
@@ -2221,14 +2225,15 @@ class MCPServerManager:
async def get_all_mcp_servers_with_health_and_teams(
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
include_health: bool = True,
server_ids: Optional[List[str]] = None,
) -> List[LiteLLM_MCPServerTable]:
"""
Get all MCP servers that the user has access to, with health status and team information.
Args:
user_api_key_auth: User authentication info for access control
include_health: Whether to include health check information
server_ids: Optional list of server IDs to filter. If provided, only these servers
will be checked (subject to access control). If None, all accessible servers are checked.
Returns:
List of MCP server objects with health and team data
@@ -2237,10 +2242,16 @@ class MCPServerManager:
# Get allowed server IDs
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
# Filter by requested server_ids if provided
if server_ids:
# Only check servers that are both requested AND accessible
target_server_ids = [sid for sid in server_ids if sid in allowed_server_ids]
else:
# Check all accessible servers
target_server_ids = allowed_server_ids
# Run health checks concurrently
tasks = [
self.health_check_server(server_id) for server_id in allowed_server_ids
]
tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
results = await asyncio.gather(*tasks)
# Filter out None results (servers that were not found)
@@ -2248,6 +2259,62 @@ class MCPServerManager:
return list_mcp_servers
async def get_all_allowed_mcp_servers(
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[LiteLLM_MCPServerTable]:
"""
Get all MCP servers that the user has access to.
Args:
user_api_key_auth: User authentication info for access control
Returns:
List of MCP server objects without health status
"""
from datetime import datetime
# Get allowed server IDs
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
list_mcp_servers: List[LiteLLM_MCPServerTable] = []
for server_id in allowed_server_ids:
server = self.get_mcp_server_by_id(server_id)
if not server:
verbose_logger.warning(f"MCP Server {server_id} not found in registry")
continue
# Build LiteLLM_MCPServerTable without health check
mcp_server_table = LiteLLM_MCPServerTable(
server_id=server.server_id,
server_name=server.server_name,
alias=server.alias,
description=(
server.mcp_info.get("description") if server.mcp_info else None
),
url=server.url,
transport=server.transport,
auth_type=server.auth_type,
created_at=datetime.now(),
updated_at=datetime.now(),
teams=[],
mcp_access_groups=server.access_groups or [],
allowed_tools=server.allowed_tools or [],
extra_headers=server.extra_headers or [],
mcp_info=server.mcp_info,
static_headers=server.static_headers,
status=None, # No health check performed
last_health_check=None, # No health check performed
health_check_error=None,
command=getattr(server, "command", None),
args=getattr(server, "args", None) or [],
env=getattr(server, "env", None) or {},
)
list_mcp_servers.append(mcp_server_table)
return list_mcp_servers
async def reload_servers_from_database(self):
"""
Public method to reload all MCP servers from database into registry.
@@ -16,7 +16,7 @@ Endpoints here:
import importlib
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Any, Dict, Iterable, List, Optional
from typing import Any, Dict, Iterable, List, Literal, Optional
from fastapi import (
APIRouter,
@@ -24,6 +24,7 @@ from fastapi import (
Form,
Header,
HTTPException,
Query,
Request,
Response,
status,
@@ -318,7 +319,7 @@ if MCP_AVAILABLE:
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
for auth_context in auth_contexts:
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
user_api_key_auth=auth_context
)
for server in servers:
@@ -336,6 +337,56 @@ if MCP_AVAILABLE:
server.mcp_info["is_public"] = True
return redacted_mcp_servers
@router.get(
"/server/health",
description="Health check for MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
async def health_check_servers(
server_ids: Optional[List[str]] = Query(
None,
description="Server IDs to check. If not provided, checks all accessible servers.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Perform health checks on one or more MCP servers.
Parameters:
- server_ids: Optional list of server IDs. If not provided, checks all accessible servers.
Returns:
- Health check results for requested servers
```
# Check all accessible servers
curl --location 'http://localhost:4000/v1/mcp/server/health' \
--header 'Authorization: Bearer your_api_key_here'
# Check specific servers
curl --location 'http://localhost:4000/v1/mcp/server/health?server_ids=server-1&server_ids=server-2' \
--header 'Authorization: Bearer your_api_key_here'
```
"""
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
server_status_map: Dict[
str, Optional[Literal["healthy", "unhealthy", "unknown"]]
] = {}
for auth_context in auth_contexts:
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams(
user_api_key_auth=auth_context,
server_ids=server_ids,
)
for server in servers:
if server.server_id not in server_status_map:
server_status_map[server.server_id] = server.status
return [
{"server_id": server_id, "status": status}
for server_id, status in server_status_map.items()
]
@router.get(
"/server/{server_id}",
description="Returns the mcp server info",
@@ -169,8 +169,8 @@ class TestListMCPServers:
return_value=["config_server_1", "config_server_2"]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
# Mock the new method that returns servers without health check
mock_servers = [
generate_mock_mcp_server_db_record(
server_id="config_server_1",
alias="Zapier MCP",
@@ -184,11 +184,11 @@ class TestListMCPServers:
transport="http",
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
mock_manager.get_all_allowed_mcp_servers = AsyncMock(
return_value=mock_servers
)
for idx, server in enumerate(mock_servers_with_health):
for idx, server in enumerate(mock_servers):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
@@ -200,6 +200,9 @@ class TestListMCPServers:
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@@ -300,8 +303,8 @@ class TestListMCPServers:
]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
# Mock the new method that returns servers without health check
mock_servers = [
db_server_1,
db_server_2,
generate_mock_mcp_server_db_record(
@@ -317,11 +320,11 @@ class TestListMCPServers:
transport="http",
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
mock_manager.get_all_allowed_mcp_servers = AsyncMock(
return_value=mock_servers
)
for idx, server in enumerate(mock_servers_with_health):
for idx, server in enumerate(mock_servers):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
@@ -333,6 +336,9 @@ class TestListMCPServers:
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@@ -425,8 +431,8 @@ class TestListMCPServers:
return_value=["db_server_allowed", "config_server_allowed"]
)
# Mock the new method that returns servers with health and team data
mock_servers_with_health = [
# Mock the new method that returns servers without health check
mock_servers = [
db_server_allowed,
generate_mock_mcp_server_db_record(
server_id="config_server_allowed",
@@ -434,11 +440,11 @@ class TestListMCPServers:
url="https://actions.zapier.com/mcp/sse",
),
]
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=mock_servers_with_health
mock_manager.get_all_allowed_mcp_servers = AsyncMock(
return_value=mock_servers
)
for idx, server in enumerate(mock_servers_with_health):
for idx, server in enumerate(mock_servers):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
@@ -450,6 +456,9 @@ class TestListMCPServers:
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@@ -975,3 +984,163 @@ class TestUpdateMCPServer:
# Verify the result includes extra_headers
assert result.extra_headers == ["X-Custom-Header", "X-Another-Header"]
assert result.alias == "Updated Test Server"
class TestHealthCheckServers:
"""Test suite for health check servers endpoint"""
@pytest.mark.asyncio
async def test_health_check_all_servers(self):
"""
Test health check for all accessible servers
Scenario: User has access to 2 servers, checks all
Expected: Returns health status for both servers
"""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_servers,
)
# Mock user auth
mock_user_auth = generate_mock_user_api_key_auth()
# Mock health check results
mock_health_result_1 = generate_mock_mcp_server_db_record(
server_id="server-1",
alias="Server 1",
url="https://server1.example.com",
)
mock_health_result_1.status = "healthy"
mock_health_result_1.last_health_check = datetime.now()
mock_health_result_1.health_check_error = None
mock_health_result_2 = generate_mock_mcp_server_db_record(
server_id="server-2",
alias="Server 2",
url="https://server2.example.com",
)
mock_health_result_2.status = "unhealthy"
mock_health_result_2.last_health_check = datetime.now()
mock_health_result_2.health_check_error = "Connection timeout"
# Mock manager
mock_manager = MagicMock()
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=[mock_health_result_1, mock_health_result_2]
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
result = await health_check_servers(
server_ids=None,
user_api_key_dict=mock_user_auth,
)
# Verify results
assert len(result) == 2
assert result[0]["server_id"] == "server-1"
assert result[0]["status"] == "healthy"
assert result[1]["server_id"] == "server-2"
assert result[1]["status"] == "unhealthy"
@pytest.mark.asyncio
async def test_health_check_specific_servers(self):
"""
Test health check for specific servers
Scenario: User requests health check for specific server IDs
Expected: Returns health status only for requested servers
"""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_servers,
)
# Mock user auth
mock_user_auth = generate_mock_user_api_key_auth()
# Mock health check result
mock_health_result = generate_mock_mcp_server_db_record(
server_id="server-1",
alias="Server 1",
url="https://server1.example.com",
)
mock_health_result.status = "healthy"
mock_health_result.last_health_check = datetime.now()
mock_health_result.health_check_error = None
# Mock manager
mock_manager = MagicMock()
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=[mock_health_result]
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
result = await health_check_servers(
server_ids=["server-1"],
user_api_key_dict=mock_user_auth,
)
# Verify results
assert len(result) == 1
assert result[0]["server_id"] == "server-1"
assert result[0]["status"] == "healthy"
@pytest.mark.asyncio
async def test_health_check_unauthorized_servers(self):
"""
Test health check with unauthorized servers
Scenario: User requests health check for servers they don't have access to
Expected: Only checks accessible servers, unauthorized servers are filtered out
"""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_servers,
)
# Mock user auth
mock_user_auth = generate_mock_user_api_key_auth()
# Mock health check result for authorized server
mock_health_result = generate_mock_mcp_server_db_record(
server_id="server-1",
alias="Server 1",
url="https://server1.example.com",
)
mock_health_result.status = "healthy"
mock_health_result.last_health_check = datetime.now()
mock_health_result.health_check_error = None
# Mock manager - server_ids filter is applied inside get_all_mcp_servers_with_health_and_teams
# So it only returns servers the user has access to
mock_manager = MagicMock()
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=[mock_health_result] # Only server-1 is returned (accessible)
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
):
result = await health_check_servers(
server_ids=["server-1", "server-unauthorized"],
user_api_key_dict=mock_user_auth,
)
# Verify results - only accessible server is returned
assert len(result) == 1
assert result[0]["server_id"] == "server-1"
assert result[0]["status"] == "healthy"
@@ -0,0 +1,127 @@
/* @vitest-environment jsdom */
import React from "react";
import { renderHook, waitFor } from "@testing-library/react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { useMCPServerHealth } from "./useMCPServerHealth";
import * as networking from "@/components/networking";
// Mock the networking module
vi.mock("@/components/networking", () => ({
fetchMCPServerHealth: vi.fn(),
}));
// Mock useAuthorized hook
vi.mock("../useAuthorized", () => ({
default: vi.fn(() => ({
accessToken: "test-token-123",
})),
}));
const createQueryClient = () =>
new QueryClient({
defaultOptions: {
queries: {
retry: false,
gcTime: 0,
},
},
});
const wrapper = ({ children }: { children: React.ReactNode }) => {
const queryClient = createQueryClient();
return React.createElement(QueryClientProvider, { client: queryClient }, children);
};
describe("useMCPServerHealth", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should fetch health status for given server IDs", async () => {
const mockHealthStatuses = [
{ server_id: "server-1", status: "healthy" },
{ server_id: "server-2", status: "unhealthy" },
];
vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue(mockHealthStatuses);
const { result } = renderHook(() => useMCPServerHealth(["server-1", "server-2"]), {
wrapper,
});
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(networking.fetchMCPServerHealth).toHaveBeenCalledWith("test-token-123", ["server-1", "server-2"]);
expect(result.current.data).toEqual(mockHealthStatuses);
});
it("should fetch health status for all servers when no server IDs provided", async () => {
const mockHealthStatuses = [
{ server_id: "server-1", status: "healthy" },
{ server_id: "server-2", status: "healthy" },
{ server_id: "server-3", status: "unhealthy" },
];
vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue(mockHealthStatuses);
const { result } = renderHook(() => useMCPServerHealth(), {
wrapper,
});
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(networking.fetchMCPServerHealth).toHaveBeenCalledWith("test-token-123", undefined);
expect(result.current.data).toEqual(mockHealthStatuses);
});
it("should handle empty server list", async () => {
vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]);
const { result } = renderHook(() => useMCPServerHealth([]), {
wrapper,
});
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(networking.fetchMCPServerHealth).toHaveBeenCalledWith("test-token-123", []);
expect(result.current.data).toEqual([]);
});
it("should handle errors when fetching health status", async () => {
const mockError = new Error("Failed to fetch health status");
vi.mocked(networking.fetchMCPServerHealth).mockRejectedValue(mockError);
const { result } = renderHook(() => useMCPServerHealth(["server-1"]), {
wrapper,
});
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toEqual(mockError);
});
it("should not fetch when accessToken is not available", async () => {
// Mock useAuthorized to return no token
const useAuthorizedModule = await import("../useAuthorized");
vi.mocked(useAuthorizedModule.default).mockReturnValue({
accessToken: null,
} as any);
const { result } = renderHook(() => useMCPServerHealth(["server-1"]), {
wrapper,
});
// Should remain in idle state since query is not enabled
expect(result.current.status).toBe("pending");
expect(networking.fetchMCPServerHealth).not.toHaveBeenCalled();
});
});
@@ -0,0 +1,22 @@
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import { fetchMCPServerHealth } from "@/components/networking";
import useAuthorized from "../useAuthorized";
const mcpServerHealthKeys = createQueryKeys("mcpServerHealth");
interface MCPServerHealth {
server_id: string;
status: string;
}
export const useMCPServerHealth = (serverIds?: string[]) => {
const { accessToken } = useAuthorized();
return useQuery<MCPServerHealth[]>({
queryKey: mcpServerHealthKeys.list({ serverIds }),
queryFn: async () => await fetchMCPServerHealth(accessToken!, serverIds),
enabled: !!accessToken,
// Refetch health status every 30 seconds to keep it up to date
refetchInterval: 30000,
});
};
@@ -10,6 +10,7 @@ export const mcpServerColumns = (
onView: (serverId: string) => void,
onEdit: (serverId: string) => void,
onDelete: (serverId: string) => void,
isLoadingHealth?: boolean,
): ColumnDef<MCPServer>[] => [
{
accessorKey: "server_id",
@@ -58,6 +59,19 @@ export const mcpServerColumns = (
const lastCheck = server.last_health_check;
const error = server.health_check_error;
// Show loading spinner if health check is in progress
if (isLoadingHealth) {
return (
<div className="flex items-center text-gray-500">
<svg className="animate-spin h-4 w-4 mr-1" xmlns="http://www.w3.org/2000/svg" fill="none" viewBox="0 0 24 24">
<circle className="opacity-25" cx="12" cy="12" r="10" stroke="currentColor" strokeWidth="4"></circle>
<path className="opacity-75" fill="currentColor" d="M4 12a8 8 0 018-8V0C5.373 0 0 5.373 0 12h4zm2 5.291A7.962 7.962 0 014 12H0c0 3.042 1.135 5.824 3 7.938l3-2.647z"></path>
</svg>
<span className="text-xs">Loading...</span>
</div>
);
}
const getStatusColor = (status: string) => {
switch (status) {
case "healthy":
@@ -8,6 +8,7 @@ import * as networking from "../networking";
// Mock the networking module
vi.mock("../networking", () => ({
fetchMCPServers: vi.fn(),
fetchMCPServerHealth: vi.fn(),
deleteMCPServer: vi.fn(),
getProxyBaseUrl: vi.fn().mockReturnValue("http://localhost:4000"),
}));
@@ -123,4 +124,108 @@ describe("MCPServers", () => {
// Note: useMCPServers uses useAuthorized() internally, which returns "123" from global mock
expect(networking.fetchMCPServers).toHaveBeenCalledWith("123");
});
it("should fetch and merge health status for servers", async () => {
// Mock MCP servers data without health status
const mockServers = [
{
server_id: "server-1",
server_name: "Test Server 1",
alias: "test-server-1",
url: "https://example.com/mcp",
transport: "http",
auth_type: "none",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
teams: [],
mcp_access_groups: [],
status: undefined,
},
{
server_id: "server-2",
server_name: "Test Server 2",
alias: "test-server-2",
url: "https://example2.com/mcp",
transport: "sse",
auth_type: "api_key",
created_at: "2024-01-02T00:00:00Z",
created_by: "user-2",
updated_at: "2024-01-02T00:00:00Z",
updated_by: "user-2",
teams: [],
mcp_access_groups: ["group-1"],
status: undefined,
},
];
// Mock health status data
const mockHealthStatuses = [
{ server_id: "server-1", status: "healthy" },
{ server_id: "server-2", status: "unhealthy" },
];
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue(mockHealthStatuses);
const queryClient = createQueryClient();
const { getByText } = render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
);
// Wait for the component to load
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
});
// Verify the health check API was called with server IDs
await waitFor(() => {
expect(networking.fetchMCPServerHealth).toHaveBeenCalledWith("123", ["server-1", "server-2"]);
});
});
it("should display loading state while health check is in progress", async () => {
const mockServers = [
{
server_id: "server-1",
server_name: "Test Server 1",
alias: "test-server-1",
url: "https://example.com/mcp",
transport: "http",
auth_type: "none",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
teams: [],
mcp_access_groups: [],
},
];
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
// Mock health check to never resolve (to test loading state)
vi.mocked(networking.fetchMCPServerHealth).mockImplementation(
() => new Promise(() => {}), // Never resolves
);
const queryClient = createQueryClient();
const { getByText } = render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
);
// Wait for the component to load
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
});
// Verify that health check was initiated
await waitFor(() => {
expect(networking.fetchMCPServerHealth).toHaveBeenCalled();
});
});
});
@@ -2,8 +2,9 @@ import { isAdminRole } from "@/utils/roles";
import { QuestionCircleOutlined } from "@ant-design/icons";
import { Button, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react";
import { Descriptions, Modal, Select, Tooltip, Typography } from "antd";
import React, { useEffect, useState } from "react";
import React, { useEffect, useState, useMemo } from "react";
import { useMCPServers } from "../../app/(dashboard)/hooks/mcpServers/useMCPServers";
import { useMCPServerHealth } from "../../app/(dashboard)/hooks/mcpServers/useMCPServerHealth";
import NotificationsManager from "../molecules/notifications_manager";
import { deleteMCPServer } from "../networking";
import { DataTable } from "../view_logs/table";
@@ -19,7 +20,29 @@ const EDIT_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-edit-state";
const { Option } = Select;
const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID }) => {
const { data: mcpServers, isLoading: isLoadingServers, refetch, dataUpdatedAt } = useMCPServers();
const { data: mcpServers, isLoading: isLoadingServers, refetch } = useMCPServers();
// Fetch health status for all servers
const serverIds = useMemo(() => mcpServers?.map((server) => server.server_id), [mcpServers]);
const { data: healthStatuses, isLoading: isLoadingHealth } = useMCPServerHealth(serverIds);
// Merge health status data into servers
const serversWithHealth = useMemo(() => {
if (!mcpServers) return [];
if (!healthStatuses) return mcpServers;
const healthMap = new Map(healthStatuses.map((h) => [h.server_id, h.status]));
return mcpServers.map((server) => {
const healthStatus = healthMap.get(server.server_id);
return {
...server,
status: healthStatus
? (healthStatus as "healthy" | "unhealthy" | "unknown")
: server.status,
};
});
}, [mcpServers, healthStatuses]);
// Log allowed_tools from fetched servers
React.useEffect(() => {
@@ -65,10 +88,10 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
// Get unique teams from all servers
const uniqueTeams = React.useMemo(() => {
if (!mcpServers) return [];
if (!serversWithHealth) return [];
const teamsSet = new Set<string>();
const uniqueTeamsArray: Team[] = [];
mcpServers.forEach((server: MCPServer) => {
serversWithHealth.forEach((server: MCPServer) => {
if (server.teams) {
server.teams.forEach((team: Team) => {
const teamKey = team.team_id;
@@ -80,17 +103,17 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
}
});
return uniqueTeamsArray;
}, [mcpServers]);
}, [serversWithHealth]);
// Get unique MCP access groups from all servers
const uniqueMcpAccessGroups = React.useMemo(() => {
if (!mcpServers) return [];
if (!serversWithHealth) return [];
return Array.from(
new Set(
mcpServers.flatMap((server) => server.mcp_access_groups).filter((group): group is string => group != null),
serversWithHealth.flatMap((server) => server.mcp_access_groups).filter((group): group is string => group != null),
),
);
}, [mcpServers]);
}, [serversWithHealth]);
// Handle team filter change
const handleTeamChange = (teamId: string) => {
@@ -106,8 +129,8 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
// Filtering logic for both team and access group
const filterServers = (teamId: string, group: string) => {
if (!mcpServers) return setFilteredServers([]);
let filtered = mcpServers;
if (!serversWithHealth) return setFilteredServers([]);
let filtered = serversWithHealth;
if (teamId === "personal") {
setFilteredServers([]);
return;
@@ -123,10 +146,10 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
setFilteredServers(filtered);
};
// Initial and effect-based filtering (trigger on query data updates)
// Initial and effect-based filtering (trigger on query data updates and health data updates)
useEffect(() => {
filterServers(selectedTeam, selectedMcpAccessGroup);
}, [dataUpdatedAt]);
}, [serversWithHealth, selectedTeam, selectedMcpAccessGroup]);
const columns = React.useMemo(
() =>
@@ -141,8 +164,9 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
setEditServer(true);
},
handleDelete,
isLoadingHealth,
),
[userRole],
[userRole, isLoadingHealth],
);
function handleDelete(server_id: string) {
@@ -5687,6 +5687,44 @@ export const fetchMCPServers = async (accessToken: string) => {
}
};
export const fetchMCPServerHealth = async (accessToken: string, serverIds?: string[]) => {
try {
// Construct base URL
let url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server/health` : `/v1/mcp/server/health`;
// Add server_ids query parameters if provided
if (serverIds && serverIds.length > 0) {
const params = new URLSearchParams();
serverIds.forEach((id) => params.append("server_ids", id));
url = `${url}?${params.toString()}`;
}
console.log("Fetching MCP server health from:", url);
const response = await fetch(url, {
method: HTTP_REQUEST.GET,
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
console.log("Fetched MCP server health:", data);
return data;
} catch (error) {
console.error("Failed to fetch MCP server health:", error);
throw error;
}
};
export const fetchMCPAccessGroups = async (accessToken: string) => {
try {
// Construct base URL