Merge pull request #18855 from BerriAI/litellm_fix_mcp-error-in-multiple-server

[fix] mcp error in multiple servers
This commit is contained in:
YutaSaito 2026-01-10 07:26:16 +09:00 • committed by GitHub
commit 07db8fe656
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 162 additions and 17 deletions

View file

@ -551,6 +551,7 @@ class MCPServerManager:
allowed_tools=getattr(mcp_server, "allowed_tools", None),
disallowed_tools=getattr(mcp_server, "disallowed_tools", None),
allow_all_keys=mcp_server.allow_all_keys,
updated_at=getattr(mcp_server, "updated_at", None),
)
return new_server
@ -697,9 +698,7 @@ class MCPServerManager:
results = await asyncio.gather(*tasks)
# Flatten results into single list
list_tools_result: List[MCPTool] = [
tool for tools in results for tool in tools
]
list_tools_result: List[MCPTool] = [tool for tools in results for tool in tools]
verbose_logger.info(
f"Successfully fetched {len(list_tools_result)} tools total from all servers"
@ -2059,7 +2058,8 @@ class MCPServerManager:
return None
async def _add_mcp_servers_from_db_to_in_memory_registry(self):
async def reload_servers_from_database(self):
"""Re-synchronize the in-memory MCP server registry with the database."""
from litellm.proxy._experimental.mcp_server.db import get_all_mcp_servers
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
@ -2074,15 +2074,34 @@ class MCPServerManager:
db_mcp_servers = await get_all_mcp_servers(prisma_client)
verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database")
# ensure the global_mcp_server_manager is up to date with the db
previous_registry = self.registry
new_registry: Dict[str, MCPServer] = {}
for server in db_mcp_servers:
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
verbose_logger.debug(
f"Adding server to registry: {server.server_id} ({server.server_name})"
f"Building server from DB: {server.server_id} ({server.server_name})"
)
await self.add_server(server)
new_registry[server.server_id] = await self.build_mcp_server_from_table(
server
)
self.registry = new_registry
verbose_logger.debug(
f"Registry now contains {len(self.get_registry())} servers"
"MCP registry refreshed (%s servers in registry)", len(new_registry)
)
def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]:
@ -2369,13 +2388,6 @@ class MCPServerManager:
servers.append(self._build_mcp_server_table(server))
return servers
async def reload_servers_from_database(self):
"""
Public method to reload all MCP servers from database into registry.
This can be called from management endpoints to ensure registry is up to date.
"""
await self._add_mcp_servers_from_db_to_in_memory_registry()
async def get_all_mcp_servers_with_health_unfiltered(
self, server_ids: Optional[List[str]] = None
) -> List[LiteLLM_MCPServerTable]:

View file

@ -3983,7 +3983,7 @@ class ProxyConfig:
)
try:
await global_mcp_server_manager._add_mcp_servers_from_db_to_in_memory_registry()
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - {}".format(
@ -4120,6 +4120,23 @@ class ProxyConfig:
return []
async def _reload_mcp_servers_job():
"""Background job entrypoint for MCP registry refreshes."""
if proxy_config._should_load_db_object(object_type="mcp") is False:
return
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
await global_mcp_server_manager.reload_servers_from_database()
except Exception as e:
verbose_proxy_logger.exception(
"Failed to reload MCP servers from database: %s", str(e)
)
proxy_config = ProxyConfig()
@ -4657,6 +4674,18 @@ class ProxyStartupEvent:
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
await proxy_config.get_credentials(prisma_client=prisma_client)
from litellm.proxy._experimental.mcp_server.utils import is_mcp_available
if is_mcp_available():
scheduler.add_job(
_reload_mcp_servers_job,
"interval",
seconds=30,
id="reload_mcp_servers_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
await cls._initialize_slack_alerting_jobs(
scheduler=scheduler,
general_settings=general_settings,

View file

@ -1,3 +1,4 @@
from datetime import datetime
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, ConfigDict
@ -50,4 +51,5 @@ class MCPServer(BaseModel):
env: Optional[Dict[str, str]] = None
access_groups: Optional[List[str]] = None
allow_all_keys: bool = False
updated_at: Optional[datetime] = None
model_config = ConfigDict(arbitrary_types_allowed=True)

View file

@ -1,4 +1,5 @@
import asyncio
from datetime import datetime, timedelta
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
@ -7,7 +8,12 @@ from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
from mcp.types import Prompt, ResourceTemplate, TextResourceContents
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPTransport,
UserAPIKeyAuth,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@pytest.mark.asyncio
@ -1688,3 +1694,99 @@ def test_filter_tools_by_allowed_tools():
assert len(filtered_tools) == 2
assert filtered_tools[0].name == "my_api_mcp-getpetbyid"
assert filtered_tools[1].name == "my_api_mcp-findpetsbystatus"
def _make_db_mcp_server(server_id: str, updated_at: datetime) -> LiteLLM_MCPServerTable:
return LiteLLM_MCPServerTable(
server_id=server_id,
server_name="server",
alias="server",
url="https://example.com",
transport=MCPTransport.http,
created_at=updated_at,
updated_at=updated_at,
mcp_info={},
)
class TestMCPServerManagerReload:
@pytest.mark.asyncio
async def test_reuses_existing_server_when_updated_at_matches(self):
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
except ImportError:
pytest.skip("MCP server not available")
manager = MCPServerManager()
timestamp = datetime.utcnow()
existing_server = MCPServer(
server_id="server-1",
name="server",
transport=MCPTransport.http,
updated_at=timestamp,
)
manager.registry = {existing_server.server_id: existing_server}
db_row = _make_db_mcp_server("server-1", timestamp)
with patch(
"litellm.proxy._experimental.mcp_server.db.get_all_mcp_servers",
new=AsyncMock(return_value=[db_row]),
) as mock_get_all, patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=object(),
), patch.object(
manager, "build_mcp_server_from_table", AsyncMock()
) as mock_build:
await manager.reload_servers_from_database()
mock_get_all.assert_awaited_once()
mock_build.assert_not_awaited()
assert manager.registry["server-1"] is existing_server
@pytest.mark.asyncio
async def test_rebuilds_server_when_updated_at_changes(self):
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
except ImportError:
pytest.skip("MCP server not available")
manager = MCPServerManager()
timestamp = datetime.utcnow()
existing_server = MCPServer(
server_id="server-1",
name="server",
transport=MCPTransport.http,
updated_at=timestamp,
)
manager.registry = {existing_server.server_id: existing_server}
new_timestamp = timestamp + timedelta(minutes=5)
db_row = _make_db_mcp_server("server-1", new_timestamp)
rebuilt_server = MCPServer(
server_id="server-1",
name="server",
transport=MCPTransport.http,
updated_at=new_timestamp,
)
with patch(
"litellm.proxy._experimental.mcp_server.db.get_all_mcp_servers",
new=AsyncMock(return_value=[db_row]),
) as mock_get_all, patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=object(),
), patch.object(
manager,
"build_mcp_server_from_table",
AsyncMock(return_value=rebuilt_server),
) as mock_build:
await manager.reload_servers_from_database()
mock_get_all.assert_awaited_once()
mock_build.assert_awaited_once_with(db_row)
assert manager.registry["server-1"] is rebuilt_server