Merge pull request #18443 from BerriAI/litellm_fix_health_status

fix health status
This commit is contained in:
YutaSaito 2025-12-29 14:16:38 +09:00 committed by GitHub
commit a569369195
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 237 additions and 718 deletions

View file

@ -11,7 +11,7 @@ import datetime
import hashlib
import json
import re
from typing import Any, Dict, List, Optional, Set, Tuple, Union, cast
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from urllib.parse import urlparse
from fastapi import HTTPException
@ -2127,7 +2127,7 @@ class MCPServerManager:
async def health_check_server(
self, server_id: str, mcp_auth_header: Optional[str] = None
) -> Dict[str, Any]:
) -> LiteLLM_MCPServerTable:
"""
Perform a health check on a specific MCP server.
@ -2138,96 +2138,85 @@ class MCPServerManager:
Returns:
Dict containing health check results
"""
import time
from datetime import datetime
server = self.get_mcp_server_by_id(server_id)
if not server:
return {
"server_id": server_id,
"server_name": None,
"status": "unknown",
"error": "Server not found",
"last_health_check": datetime.now().isoformat(),
"response_time_ms": None,
}
start_time = time.time()
try:
# Try to get tools from the server as a health check
tools = await self._get_tools_from_server(server, mcp_auth_header)
response_time = (time.time() - start_time) * 1000
return {
"server_id": server_id,
"server_name": server.name,
"status": "healthy",
"tools_count": len(tools),
"last_health_check": datetime.now().isoformat(),
"response_time_ms": round(response_time, 2),
"error": None,
}
except Exception as e:
response_time = (time.time() - start_time) * 1000
error_message = str(e)
return {
"server_id": server_id,
"server_name": server.name,
"status": "unhealthy",
"last_health_check": datetime.now().isoformat(),
"response_time_ms": round(response_time, 2),
"error": error_message,
}
async def health_check_all_servers(
self, mcp_auth_header: Optional[str] = None
) -> Dict[str, Any]:
"""
Perform health checks on all MCP servers.
Args:
mcp_auth_header: Optional authentication header for the MCP servers
Returns:
Dict containing health check results for all servers
"""
all_servers = self.get_registry()
results = {}
for server_id, server in all_servers.items():
results[server_id] = await self.health_check_server(
server_id, mcp_auth_header
verbose_logger.warning(f"MCP Server {server_id} not found")
return LiteLLM_MCPServerTable(
server_id=server_id,
server_name=None,
transport=MCPTransport.http, # Default transport for not found servers
status="unknown",
health_check_error="Server not found",
last_health_check=datetime.now(),
)
return results
status: Literal["healthy", "unhealthy", "unknown"] = "unknown"
health_check_error = None
async def health_check_allowed_servers(
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
) -> Dict[str, Any]:
"""
Perform health checks on all MCP servers that the user has access to.
# Check if we should skip health check based on auth configuration
should_skip_health_check = False
Args:
user_api_key_auth: User authentication info for access control
mcp_auth_header: Optional authentication header for the MCP servers
# 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
Returns:
Dict containing health check results for accessible servers
"""
# Get allowed servers for the user
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
if not should_skip_health_check:
extra_headers = {}
if server.static_headers:
extra_headers.update(server.static_headers)
# Perform health checks on allowed servers
results = {}
for server_id in allowed_server_ids:
results[server_id] = await self.health_check_server(
server_id, mcp_auth_header
client = self._create_mcp_client(
server=server,
mcp_auth_header=None,
extra_headers=extra_headers,
stdio_env=None,
)
return results
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_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=status,
last_health_check=datetime.now(),
health_check_error=health_check_error,
command=getattr(server, "command", None),
args=getattr(server, "args", None) or [],
env=getattr(server, "env", None) or {},
)
async def get_all_mcp_servers_with_health_and_teams(
self,
@ -2244,100 +2233,18 @@ 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)
# Run health checks concurrently
tasks = [
self.health_check_server(server_id) for server_id in allowed_server_ids
]
results = await asyncio.gather(*tasks)
# 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)
# 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 {},
}
)
)
# 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}
)
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)
# 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

View file

@ -296,117 +296,6 @@ if MCP_AVAILABLE:
access_groups_list = sorted(list(access_groups))
return {"access_groups": access_groups_list}
@router.get(
"/server/{server_id}/health",
description="Perform health check on a specific MCP server",
dependencies=[Depends(user_api_key_auth)],
)
async def health_check_mcp_server(
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Perform a health check on the MCP server specified by the `server_id`
Parameters:
- server_id: str - Required. The unique identifier of the mcp server to health check.
```
curl --location 'http://localhost:4000/v1/mcp/server/{server_id}/health' \
--header 'Authorization: Bearer your_api_key_here'
```
"""
# Check if server exists and user has access
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to your proxy"
)
# check to see if server exists for all users
mcp_server = await get_mcp_server(prisma_client, server_id)
if mcp_server is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={"error": f"MCP Server with id {server_id} not found"},
)
# Implement authz restriction from requested user
if not _user_has_admin_view(user_api_key_dict):
# Perform authz check to filter the mcp servers user has access to
mcp_server_records = await get_all_mcp_servers_for_user(
prisma_client, user_api_key_dict
)
exists = does_mcp_server_exist(mcp_server_records, server_id)
if not exists:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"error": f"User does not have permission to access mcp server with id {server_id}. You can only access mcp servers that you have access to."
},
)
# Perform health check using server manager
try:
health_result = await global_mcp_server_manager.health_check_server(
server_id
)
return health_result
except Exception as e:
verbose_proxy_logger.exception(
f"Error performing health check on MCP server {server_id}: {str(e)}"
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Error performing health check: {str(e)}"},
)
@router.get(
"/server/health",
description="Perform health check on all accessible MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
async def health_check_all_mcp_servers(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Perform health checks on all MCP servers accessible to the user
```
curl --location 'http://localhost:4000/v1/mcp/server/health' \
--header 'Authorization: Bearer your_api_key_here'
```
"""
# Use server manager to get health checks for allowed servers
try:
all_health_results = (
await global_mcp_server_manager.health_check_allowed_servers(
user_api_key_auth=user_api_key_dict
)
)
return {
"total_servers": len(all_health_results),
"healthy_count": len(
[r for r in all_health_results.values() if r["status"] == "healthy"]
),
"unhealthy_count": len(
[
r
for r in all_health_results.values()
if r["status"] == "unhealthy"
]
),
"unknown_count": len(
[r for r in all_health_results.values() if r["status"] == "unknown"]
),
"servers": all_health_results,
}
except Exception as e:
verbose_proxy_logger.exception(
f"Error performing health checks on MCP servers: {str(e)}"
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Error performing health checks: {str(e)}"},
)
## FastAPI Routes
@router.get(
"/server",
@ -484,15 +373,11 @@ if MCP_AVAILABLE:
server_id
)
# Update the server object with health check results
mcp_server.status = health_result.get("status", "unknown")
mcp_server.last_health_check = (
datetime.fromisoformat(
health_result.get("last_health_check", datetime.now().isoformat())
)
if health_result.get("last_health_check")
else None
mcp_server.status = (
health_result.status if health_result.status else "unknown"
)
mcp_server.health_check_error = health_result.get("error")
mcp_server.last_health_check = health_result.last_health_check
mcp_server.health_check_error = health_result.health_check_error
except Exception as e:
verbose_proxy_logger.debug(
f"Error performing health check on server {server_id}: {e}"

View file

@ -641,37 +641,31 @@ class TestMCPServerManager:
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.http,
auth_type=None,
authentication_token="test-token",
url="http://test-server.com",
)
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock successful _get_tools_from_server
async def mock_get_tools_from_server(
server,
mcp_auth_header=None,
raw_headers=None,
):
tool1 = MagicMock()
tool1.name = "tool1"
tool2 = MagicMock()
tool2.name = "tool2"
return [tool1, tool2]
manager._get_tools_from_server = mock_get_tools_from_server
# Mock successful client.run_with_session
mock_client = AsyncMock()
mock_client.run_with_session = AsyncMock(return_value="ok")
manager._create_mcp_client = MagicMock(return_value=mock_client)
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
assert result["tools_count"] == 2
assert result["error"] is None
assert "last_health_check" in result
assert "response_time_ms" in result
assert result["response_time_ms"] >= 0 # Allow 0 for very fast mocks
# Verify results - result is now LiteLLM_MCPServerTable
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "test-server"
assert result.status == "healthy"
assert result.health_check_error is None
assert result.last_health_check is not None
@pytest.mark.asyncio
async def test_health_check_server_unhealthy(self):
@ -679,32 +673,33 @@ class TestMCPServerManager:
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.http,
auth_type=None,
authentication_token="test-token",
url="http://test-server.com",
)
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock failed _get_tools_from_server
async def mock_get_tools_from_server(
server,
mcp_auth_header=None,
raw_headers=None,
):
raise Exception("Connection timeout")
manager._get_tools_from_server = mock_get_tools_from_server
# Mock failed client.run_with_session
mock_client = AsyncMock()
mock_client.run_with_session = AsyncMock(
side_effect=Exception("Connection timeout")
)
manager._create_mcp_client = MagicMock(return_value=mock_client)
# Perform health check
result = await manager.health_check_server("test-server")
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "unhealthy"
assert result["error"] == "Connection timeout"
assert "last_health_check" in result
assert "response_time_ms" in result
assert result["response_time_ms"] >= 0 # Allow 0 for very fast mocks
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "test-server"
assert result.status == "unhealthy"
assert result.health_check_error == "Connection timeout"
assert result.last_health_check is not None
@pytest.mark.asyncio
async def test_health_check_server_not_found(self):
@ -718,104 +713,121 @@ class TestMCPServerManager:
result = await manager.health_check_server("non-existent-server")
# Verify results
assert result["server_id"] == "non-existent-server"
assert result["status"] == "unknown"
assert result["error"] == "Server not found"
assert result["response_time_ms"] is None
assert "last_health_check" in result
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "non-existent-server"
assert result.server_name is None
assert result.status == "unknown"
assert result.health_check_error == "Server not found"
assert result.last_health_check is not None
@pytest.mark.asyncio
async def test_health_check_all_servers(self):
"""Test health check for all servers"""
async def test_health_check_server_oauth2_skips_check(self):
"""Test that health check is skipped for OAuth2 servers and returns unknown status"""
manager = MCPServerManager()
# Mock servers
server1 = MagicMock()
server1.server_id = "server1"
server1.name = "server1"
server2 = MagicMock()
server2.server_id = "server2"
server2.name = "server2"
# Mock registry
manager.registry = {"server1": server1, "server2": server2}
# Mock get_mcp_server_by_id
def mock_get_server_by_id(server_id):
if server_id == "server1":
return server1
elif server_id == "server2":
return server2
return None
manager.get_mcp_server_by_id = mock_get_server_by_id
# Mock _get_tools_from_server with different results
async def mock_get_tools_from_server(
server,
mcp_auth_header=None,
raw_headers=None,
):
if server.server_id == "server1":
tool = MagicMock()
tool.name = "tool1"
return [tool]
elif server.server_id == "server2":
raise Exception("Connection failed")
return []
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check for all servers
result = await manager.health_check_all_servers()
# Verify results
assert len(result) == 2
assert "server1" in result
assert "server2" in result
# Check server1 (healthy)
assert result["server1"]["status"] == "healthy"
assert result["server1"]["tools_count"] == 1
assert result["server1"]["error"] is None
# Check server2 (unhealthy)
assert result["server2"]["status"] == "unhealthy"
assert result["server2"]["error"] == "Connection failed"
@pytest.mark.asyncio
async def test_health_check_server_with_auth_header(self):
"""Test health check with authentication header"""
manager = MCPServerManager()
# Mock server
server = MagicMock()
server.server_id = "test-server"
server.name = "test-server"
# Mock OAuth2 server
server = MCPServer(
server_id="oauth2-server",
name="oauth2-server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
url="http://oauth2-server.com",
)
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock _get_tools_from_server to verify auth header is passed
async def mock_get_tools_from_server(
server,
mcp_auth_header=None,
raw_headers=None,
):
assert mcp_auth_header == "test-token"
tool = MagicMock()
tool.name = "tool1"
return [tool]
# _create_mcp_client should not be called for OAuth2 servers
manager._create_mcp_client = MagicMock()
manager._get_tools_from_server = mock_get_tools_from_server
# Perform health check
result = await manager.health_check_server("oauth2-server")
# Perform health check with auth header
result = await manager.health_check_server("test-server", "test-token")
# Verify that client was not created (health check was skipped)
manager._create_mcp_client.assert_not_called()
# Verify results
assert result["server_id"] == "test-server"
assert result["status"] == "healthy"
assert result["tools_count"] == 1
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "oauth2-server"
assert result.status == "unknown"
assert result.health_check_error is None
assert result.last_health_check is not None
@pytest.mark.asyncio
async def test_health_check_server_no_token_skips_check(self):
"""Test that health check is skipped when auth_type is set but authentication_token is missing"""
manager = MCPServerManager()
# Mock server with auth_type but no authentication_token
server = MCPServer(
server_id="no-token-server",
name="no-token-server",
transport=MCPTransport.http,
auth_type=MCPAuth.bearer_token,
authentication_token=None, # No token
url="http://no-token-server.com",
)
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# _create_mcp_client should not be called
manager._create_mcp_client = MagicMock()
# Perform health check
result = await manager.health_check_server("no-token-server")
# Verify that client was not created (health check was skipped)
manager._create_mcp_client.assert_not_called()
# Verify results
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "no-token-server"
assert result.status == "unknown"
assert result.health_check_error is None
assert result.last_health_check is not None
@pytest.mark.asyncio
async def test_health_check_server_with_static_headers(self):
"""Test health check with static headers configured"""
manager = MCPServerManager()
# Mock server with static_headers
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.http,
auth_type=None,
authentication_token="test-token",
url="http://test-server.com",
static_headers={"X-Custom-Header": "custom-value"},
)
manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock successful client
mock_client = AsyncMock()
mock_client.run_with_session = AsyncMock(return_value="ok")
# Capture the extra_headers passed to _create_mcp_client
captured_extra_headers = None
def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env):
nonlocal captured_extra_headers
captured_extra_headers = extra_headers
return mock_client
manager._create_mcp_client = MagicMock(side_effect=capture_create_mcp_client)
# Perform health check
result = await manager.health_check_server("test-server")
# Verify static headers were passed
assert captured_extra_headers == {"X-Custom-Header": "custom-value"}
# Verify results
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "test-server"
assert result.status == "healthy"
assert result.health_check_error is None
@pytest.mark.asyncio
async def test_pre_call_tool_check_allowed_tools_list_allows_tool(self):

View file

@ -486,11 +486,14 @@ class TestListMCPServers:
mock_server.credentials = {"auth_value": "top-secret"}
mock_prisma_client = MagicMock()
mock_health_result = {
"status": "healthy",
"last_health_check": datetime.now().isoformat(),
"error": None,
}
# Mock health check result as LiteLLM_MCPServerTable
mock_health_result = generate_mock_mcp_server_db_record(
server_id="server-1", alias="Server 1"
)
mock_health_result.status = "healthy"
mock_health_result.last_health_check = datetime.now()
mock_health_result.health_check_error = None
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
@ -531,11 +534,14 @@ class TestListMCPServers:
delattr(mock_server, "credentials")
mock_prisma_client = MagicMock()
mock_health_result = {
"status": "healthy",
"last_health_check": datetime.now().isoformat(),
"error": None,
}
# Mock health check result as LiteLLM_MCPServerTable
mock_health_result = generate_mock_mcp_server_db_record(
server_id="server-2", alias="Server 2"
)
mock_health_result.status = "healthy"
mock_health_result.last_health_check = datetime.now()
mock_health_result.health_check_error = None
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
@ -568,296 +574,6 @@ class TestListMCPServers:
assert result.status == "healthy"
class TestMCPHealthCheckEndpoints:
"""Test MCP health check endpoints"""
@pytest.mark.asyncio
async def test_health_check_mcp_server_success(self):
"""Test successful health check for a specific MCP server"""
# Mock server
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server", alias="Test Server"
)
# Mock dependencies
mock_prisma_client = MagicMock()
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.health_check_server = AsyncMock(
return_value={
"server_id": "test-server",
"server_name": "Test Server",
"status": "healthy",
"tools_count": 3,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 150.5,
"error": None,
}
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with 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._user_has_admin_view",
return_value=True,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
result = await health_check_mcp_server(
server_id="test-server", user_api_key_dict=mock_user_auth
)
# Verify results
assert result["server_id"] == "test-server"
assert result["server_name"] == "Test Server"
assert result["status"] == "healthy"
assert result["tools_count"] == 3
assert result["response_time_ms"] == 150.5
assert result["error"] is None
@pytest.mark.asyncio
async def test_health_check_mcp_server_not_found(self):
"""Test health check for a server that doesn't exist"""
# Mock dependencies
mock_prisma_client = MagicMock()
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with 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.get_mcp_server",
AsyncMock(return_value=None),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
# Should raise HTTPException
with pytest.raises(Exception) as exc_info:
await health_check_mcp_server(
server_id="non-existent-server", user_api_key_dict=mock_user_auth
)
assert "not found" in str(exc_info.value)
@pytest.mark.asyncio
async def test_health_check_mcp_server_unauthorized(self):
"""Test health check for a server user doesn't have access to"""
# Mock server
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server", alias="Test Server"
)
# Mock dependencies
mock_prisma_client = MagicMock()
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER # Non-admin user
)
# Mock user doesn't have access to this server
mock_user_servers = []
with 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._user_has_admin_view",
return_value=False,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_all_mcp_servers_for_user",
return_value=mock_user_servers,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_mcp_server,
)
# Should raise HTTPException
with pytest.raises(Exception) as exc_info:
await health_check_mcp_server(
server_id="test-server", user_api_key_dict=mock_user_auth
)
assert "permission" in str(exc_info.value)
@pytest.mark.asyncio
async def test_health_check_all_mcp_servers(self):
"""Test health check for all accessible MCP servers"""
# Mock team records
team_records = [
generate_mock_team_record(
team_id="team1",
team_alias="Team 1",
organization_id="org1",
mcp_servers=["server1", "server2"],
)
]
# Mock DB servers
db_servers = [
generate_mock_mcp_server_db_record(server_id="server1"),
generate_mock_mcp_server_db_record(server_id="server2"),
]
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
mock_prisma_client=mock_prisma_client,
team_records=team_records,
mcp_servers=db_servers,
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.health_check_allowed_servers = AsyncMock(
return_value={
"server1": {
"server_id": "server1",
"server_name": "Test DB Server",
"status": "healthy",
"tools_count": 2,
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 100.0,
"error": None,
},
"server2": {
"server_id": "server2",
"server_name": "Test DB Server",
"status": "unhealthy",
"last_health_check": "2024-01-01T12:00:00",
"response_time_ms": 5000.0,
"error": "Connection timeout",
},
}
)
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1", "server2"]
)
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER
)
with 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._user_has_admin_view",
return_value=False,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
health_check_all_mcp_servers,
)
result = await health_check_all_mcp_servers(
user_api_key_dict=mock_user_auth
)
# Verify results
assert result["total_servers"] == 2
assert result["healthy_count"] == 1
assert result["unhealthy_count"] == 1
assert result["unknown_count"] == 0
assert "server1" in result["servers"]
assert "server2" in result["servers"]
# Check individual server results
assert result["servers"]["server1"]["status"] == "healthy"
assert result["servers"]["server1"]["tools_count"] == 2
assert result["servers"]["server1"]["server_name"] == "Test DB Server"
assert result["servers"]["server2"]["status"] == "unhealthy"
assert result["servers"]["server2"]["error"] == "Connection timeout"
assert result["servers"]["server2"]["server_name"] == "Test DB Server"
@pytest.mark.asyncio
async def test_fetch_all_mcp_servers_with_health_status(self):
"""Test that fetch_all_mcp_servers includes health check status"""
# Mock server with health status
mock_server = generate_mock_mcp_server_db_record(
server_id="test-server", alias="Test Server"
)
# Add health status to the mock server
mock_server.status = "healthy"
mock_server.last_health_check = datetime.now()
mock_server.health_check_error = None
# Mock dependencies
mock_prisma_client = MagicMock()
mock_prisma_client = setup_mock_prisma_client(
mock_prisma_client=mock_prisma_client,
team_records=[],
mcp_servers=[], # Don't add servers here since we're mocking get_all_mcp_servers
)
# Mock global MCP server manager
mock_manager = MagicMock()
mock_manager.config_mcp_servers = {}
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[])
mock_manager.get_all_mcp_servers_with_health_and_teams = AsyncMock(
return_value=[mock_server]
)
mock_server.credentials = {"auth_value": "secret"}
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with 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._user_has_admin_view",
return_value=True,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
):
# Import and call the function
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
# Verify health check status is included
assert len(result) == 1
server = result[0]
assert server.server_id == "test-server"
assert server.status == "healthy"
assert server.last_health_check is not None
assert server.health_check_error is None
assert server.credentials is None
class TestTemporaryMCPSessionEndpoints:
def test_inherit_credentials_from_existing_server(self):
payload = NewMCPServerRequest(
@ -1170,7 +886,6 @@ class TestTemporaryMCPSessionEndpoints:
fallback_client_id="server-1",
)
class TestUpdateMCPServer:
"""Test suite for update MCP server functionality"""