mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
tests: add test
This commit is contained in:
parent
8b90e5f4dd
commit
5927a557fb
1 changed files with 103 additions and 1 deletions
|
|
@ -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