mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
feat: add user_mcp_management_mode for view_all visibility
This commit is contained in:
parent
7937c8674b
commit
694bcb6186
4 changed files with 205 additions and 50 deletions
|
|
@ -2284,14 +2284,7 @@ class MCPServerManager:
|
|||
# Check all accessible servers
|
||||
target_server_ids = allowed_server_ids
|
||||
|
||||
# Run health checks concurrently
|
||||
tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
# Filter out None results (servers that were not found)
|
||||
list_mcp_servers = [server for server in results if server is not None]
|
||||
|
||||
return list_mcp_servers
|
||||
return await self._run_health_checks(target_server_ids)
|
||||
|
||||
async def get_all_allowed_mcp_servers(
|
||||
self,
|
||||
|
|
@ -2306,8 +2299,6 @@ class MCPServerManager:
|
|||
Returns:
|
||||
List of MCP server objects without health status
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
# Get allowed server IDs
|
||||
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
|
||||
|
||||
|
|
@ -2319,40 +2310,56 @@ class MCPServerManager:
|
|||
verbose_logger.warning(f"MCP Server {server_id} not found in registry")
|
||||
continue
|
||||
|
||||
# Build LiteLLM_MCPServerTable without health check
|
||||
mcp_server_table = LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
alias=server.alias,
|
||||
description=(
|
||||
server.mcp_info.get("description") if server.mcp_info else None
|
||||
),
|
||||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
extra_headers=server.extra_headers or [],
|
||||
mcp_info=server.mcp_info,
|
||||
static_headers=server.static_headers,
|
||||
status=None, # No health check performed
|
||||
last_health_check=None, # No health check performed
|
||||
health_check_error=None,
|
||||
command=getattr(server, "command", None),
|
||||
args=getattr(server, "args", None) or [],
|
||||
env=getattr(server, "env", None) or {},
|
||||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
)
|
||||
mcp_server_table = self._build_mcp_server_table(server)
|
||||
list_mcp_servers.append(mcp_server_table)
|
||||
|
||||
return list_mcp_servers
|
||||
|
||||
def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
|
||||
from datetime import datetime
|
||||
|
||||
return LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.server_name,
|
||||
alias=server.alias,
|
||||
description=(
|
||||
server.mcp_info.get("description") if server.mcp_info else None
|
||||
),
|
||||
url=server.url,
|
||||
transport=server.transport,
|
||||
auth_type=server.auth_type,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
extra_headers=server.extra_headers or [],
|
||||
mcp_info=server.mcp_info,
|
||||
static_headers=server.static_headers,
|
||||
status=None, # No health check performed
|
||||
last_health_check=None, # No health check performed
|
||||
health_check_error=None,
|
||||
command=getattr(server, "command", None),
|
||||
args=getattr(server, "args", None) or [],
|
||||
env=getattr(server, "env", None) or {},
|
||||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
|
||||
"""Return all MCP servers from registry without applying access controls."""
|
||||
|
||||
registry = self.get_registry()
|
||||
if not registry:
|
||||
return []
|
||||
|
||||
servers: List[LiteLLM_MCPServerTable] = []
|
||||
for server in registry.values():
|
||||
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.
|
||||
|
|
@ -2360,5 +2367,34 @@ class MCPServerManager:
|
|||
"""
|
||||
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]:
|
||||
"""Return health info for all servers in registry regardless of user access."""
|
||||
|
||||
registry = self.get_registry()
|
||||
if not registry:
|
||||
return []
|
||||
|
||||
if server_ids:
|
||||
target_server_ids = [sid for sid in server_ids if sid in registry]
|
||||
else:
|
||||
target_server_ids = list(registry.keys())
|
||||
|
||||
if not target_server_ids:
|
||||
return []
|
||||
|
||||
return await self._run_health_checks(target_server_ids)
|
||||
|
||||
async def _run_health_checks(
|
||||
self, target_server_ids: List[str]
|
||||
) -> List[LiteLLM_MCPServerTable]:
|
||||
if not target_server_ids:
|
||||
return []
|
||||
|
||||
tasks = [self.health_check_server(server_id) for server_id in target_server_ids]
|
||||
results = await asyncio.gather(*tasks)
|
||||
return [server for server in results if server is not None]
|
||||
|
||||
|
||||
global_mcp_server_manager: MCPServerManager = MCPServerManager()
|
||||
|
|
|
|||
|
|
@ -1908,6 +1908,9 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase):
|
|||
}
|
||||
|
||||
|
||||
UserMCPManagementMode = Literal["restricted", "view_all"]
|
||||
|
||||
|
||||
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
||||
"""
|
||||
Documents all the fields supported by `general_settings` in config.yaml
|
||||
|
|
@ -2025,6 +2028,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).",
|
||||
)
|
||||
user_mcp_management_mode: Optional[UserMCPManagementMode] = Field(
|
||||
None,
|
||||
description="Controls how non-admin users interact with MCP servers in the dashboard. 'restricted' shows only accessible servers, 'view_all' lists every server in read-only mode.",
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -76,6 +76,7 @@ if MCP_AVAILABLE:
|
|||
SpecialMCPServerName,
|
||||
UpdateMCPServerRequest,
|
||||
UserAPIKeyAuth,
|
||||
UserMCPManagementMode,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
|
|
@ -302,6 +303,19 @@ if MCP_AVAILABLE:
|
|||
return {"access_groups": access_groups_list}
|
||||
|
||||
## FastAPI Routes
|
||||
def _get_user_mcp_management_mode() -> UserMCPManagementMode:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings as proxy_general_settings,
|
||||
)
|
||||
except Exception:
|
||||
proxy_general_settings = None
|
||||
|
||||
mode = (proxy_general_settings or {}).get("user_mcp_management_mode")
|
||||
if mode == "view_all":
|
||||
return "view_all"
|
||||
return "restricted"
|
||||
|
||||
@router.get(
|
||||
"/server",
|
||||
description="Returns the mcp server list with associated teams",
|
||||
|
|
@ -319,18 +333,26 @@ if MCP_AVAILABLE:
|
|||
```
|
||||
"""
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
user_mcp_management_mode = _get_user_mcp_management_mode()
|
||||
|
||||
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
|
||||
user_api_key_auth=auth_context
|
||||
if user_mcp_management_mode == "view_all":
|
||||
servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(servers)
|
||||
else:
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
|
||||
for auth_context in auth_contexts:
|
||||
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
|
||||
user_api_key_auth=auth_context
|
||||
)
|
||||
for server in servers:
|
||||
if server.server_id not in aggregated_servers:
|
||||
aggregated_servers[server.server_id] = server
|
||||
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(
|
||||
aggregated_servers.values()
|
||||
)
|
||||
for server in servers:
|
||||
if server.server_id not in aggregated_servers:
|
||||
aggregated_servers[server.server_id] = server
|
||||
|
||||
redacted_mcp_servers = _redact_mcp_credentials_list(aggregated_servers.values())
|
||||
|
||||
# augment the mcp servers with public status
|
||||
if litellm.public_mcp_servers is not None:
|
||||
|
|
@ -372,6 +394,17 @@ if MCP_AVAILABLE:
|
|||
--header 'Authorization: Bearer your_api_key_here'
|
||||
```
|
||||
"""
|
||||
user_mcp_management_mode = _get_user_mcp_management_mode()
|
||||
|
||||
if user_mcp_management_mode == "view_all":
|
||||
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(
|
||||
server_ids=server_ids
|
||||
)
|
||||
return [
|
||||
{"server_id": server.server_id, "status": server.status}
|
||||
for server in servers
|
||||
]
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
server_status_map: Dict[
|
||||
|
|
|
|||
|
|
@ -231,6 +231,40 @@ class TestListMCPServers:
|
|||
assert server.url == "https://mcp.deepwiki.com/mcp"
|
||||
assert server.transport == "http"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_mcp_servers_view_all_mode(self):
|
||||
"""Users should see all MCP servers when view_all mode is enabled."""
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
mock_servers = [
|
||||
generate_mock_mcp_server_db_record(server_id="server-1", alias="One"),
|
||||
generate_mock_mcp_server_db_record(server_id="server-2", alias="Two"),
|
||||
]
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_all_mcp_servers_unfiltered = AsyncMock(
|
||||
return_value=mock_servers
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
|
||||
return_value="view_all",
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_all_mcp_servers,
|
||||
)
|
||||
|
||||
result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth)
|
||||
|
||||
assert len(result) == 2
|
||||
assert {server.server_id for server in result} == {"server-1", "server-2"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_mcp_servers_combined_config_and_db(self):
|
||||
"""
|
||||
|
|
@ -1096,6 +1130,51 @@ class TestHealthCheckServers:
|
|||
assert result[0]["server_id"] == "server-1"
|
||||
assert result[0]["status"] == "healthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_view_all_mode(self):
|
||||
"""view_all mode should return health info for all MCP servers."""
|
||||
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
health_check_servers,
|
||||
)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
health_result_one = generate_mock_mcp_server_db_record(
|
||||
server_id="server-1", alias="One"
|
||||
)
|
||||
health_result_one.status = "healthy"
|
||||
|
||||
health_result_two = generate_mock_mcp_server_db_record(
|
||||
server_id="server-2", alias="Two"
|
||||
)
|
||||
health_result_two.status = "unhealthy"
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_all_mcp_servers_with_health_unfiltered = AsyncMock(
|
||||
return_value=[health_result_one, health_result_two]
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
|
||||
return_value="view_all",
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
):
|
||||
result = await health_check_servers(
|
||||
server_ids=None,
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0]["server_id"] == "server-1"
|
||||
assert result[0]["status"] == "healthy"
|
||||
assert result[1]["server_id"] == "server-2"
|
||||
assert result[1]["status"] == "unhealthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_unauthorized_servers(self):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue