mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #18855 from BerriAI/litellm_fix_mcp-error-in-multiple-server
[fix] mcp error in multiple servers
This commit is contained in:
commit
07db8fe656
4 changed files with 162 additions and 17 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue