diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 8062243dfdd..f1558ac5791 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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