mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-18 06:26:16 +00:00
fix: replace list in for set in
faster
This commit is contained in:
@@ -10,7 +10,7 @@ import asyncio
|
||||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
from typing import Any, Dict, List, Optional, Set, Union, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
@@ -285,12 +285,12 @@ class MCPServerManager:
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
verbose_logger.debug(f"Added MCP Server: {name_for_prefix}")
|
||||
|
||||
def get_all_mcp_server_ids(self) -> List[str]:
|
||||
def get_all_mcp_server_ids(self) -> Set[str]:
|
||||
"""
|
||||
Get all MCP server IDs
|
||||
"""
|
||||
all_servers = self.get_registry().values()
|
||||
return [server.server_id for server in all_servers]
|
||||
all_servers = list(self.get_registry().values())
|
||||
return {server.server_id for server in all_servers}
|
||||
|
||||
async def get_allowed_mcp_servers(
|
||||
self, user_api_key_auth: Optional[UserAPIKeyAuth] = None
|
||||
@@ -1187,7 +1187,7 @@ class MCPServerManager:
|
||||
valid_server_ids = self.get_all_mcp_server_ids()
|
||||
filtered_list_mcp_servers = []
|
||||
for server in list_mcp_servers:
|
||||
if server.server_id not in valid_server_ids:
|
||||
if server.server_id in valid_server_ids:
|
||||
filtered_list_mcp_servers.append(server)
|
||||
|
||||
# Map servers to their teams and return with health data
|
||||
|
||||
@@ -654,6 +654,7 @@ class TestMCPServerManager:
|
||||
"Tool tool3 is not allowed for server test-server"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
async def test_get_tools_from_server_add_prefix(self):
|
||||
"""Verify _get_tools_from_server respects add_prefix True/False."""
|
||||
manager = MCPServerManager()
|
||||
@@ -913,16 +914,20 @@ class TestMCPServerManager:
|
||||
async def test_get_all_mcp_servers_filters_invalid_server_ids(self):
|
||||
"""Test that get_all_mcp_servers_with_health_and_teams filters out servers with invalid server_ids"""
|
||||
manager = MCPServerManager()
|
||||
|
||||
|
||||
# Mock get_allowed_mcp_servers to return some server IDs
|
||||
manager.get_allowed_mcp_servers = AsyncMock(return_value=["valid-server-1", "invalid-server-1"])
|
||||
|
||||
manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["valid-server-1", "invalid-server-1"]
|
||||
)
|
||||
|
||||
# Mock get_all_mcp_server_ids to return only valid server IDs
|
||||
manager.get_all_mcp_server_ids = MagicMock(return_value=["valid-server-1", "another-valid-server"])
|
||||
|
||||
manager.get_all_mcp_server_ids = MagicMock(
|
||||
return_value=["valid-server-1", "another-valid-server"]
|
||||
)
|
||||
|
||||
# Mock config_mcp_servers to be empty for this test
|
||||
manager.config_mcp_servers = {}
|
||||
|
||||
|
||||
# Create mock servers - one valid, one invalid
|
||||
valid_server = LiteLLM_MCPServerTable(
|
||||
server_id="valid-server-1",
|
||||
@@ -934,43 +939,53 @@ class TestMCPServerManager:
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
|
||||
invalid_server = LiteLLM_MCPServerTable(
|
||||
server_id="invalid-server-1",
|
||||
server_id="invalid-server-1",
|
||||
server_name="Invalid Server",
|
||||
alias="invalid",
|
||||
description="An invalid server",
|
||||
url="http://invalid.com",
|
||||
url="http://invalid.com",
|
||||
transport=MCPTransport.http,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
|
||||
# Mock the database call to return both servers
|
||||
with patch('litellm.proxy._experimental.mcp_server.db.get_mcp_servers') as mock_get_servers:
|
||||
mock_get_servers.return_value = [valid_server, invalid_server]
|
||||
|
||||
# Mock prisma_client to not be None
|
||||
with patch('litellm.proxy.proxy_server.prisma_client', new=MagicMock()):
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.db.get_mcp_servers",
|
||||
new=AsyncMock(return_value=[valid_server, invalid_server]),
|
||||
) as mock_get_servers:
|
||||
|
||||
# Mock prisma_client with proper async methods
|
||||
mock_prisma = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
mock_teamtable = MagicMock()
|
||||
mock_teamtable.find_many = AsyncMock(return_value=[])
|
||||
mock_db.litellm_teamtable = mock_teamtable
|
||||
mock_prisma.db = mock_db
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", new=mock_prisma):
|
||||
# Mock _user_has_admin_view to return False (non-admin user)
|
||||
with patch('litellm.proxy.management_endpoints.common_utils._user_has_admin_view') as mock_admin_view:
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view"
|
||||
) as mock_admin_view:
|
||||
mock_admin_view.return_value = False
|
||||
|
||||
|
||||
# Create mock user auth
|
||||
user_auth = MagicMock()
|
||||
user_auth.user_id = "test-user"
|
||||
|
||||
|
||||
# Call the method
|
||||
result = await manager.get_all_mcp_servers_with_health_and_teams(
|
||||
user_api_key_auth=user_auth,
|
||||
include_health=False
|
||||
user_api_key_auth=user_auth, include_health=False
|
||||
)
|
||||
|
||||
|
||||
# Verify that only the valid server is returned
|
||||
assert len(result) == 1
|
||||
assert result[0].server_id == "valid-server-1"
|
||||
assert result[0].server_name == "Valid Server"
|
||||
|
||||
|
||||
# Verify invalid server is not in results
|
||||
server_ids = [server.server_id for server in result]
|
||||
assert "invalid-server-1" not in server_ids
|
||||
|
||||
Reference in New Issue
Block a user