feat: add user_mcp_management_mode for view_all visibility

This commit is contained in:
Yuta Saito 2026-01-06 11:22:28 +09:00
parent 7937c8674b
commit 694bcb6186
4 changed files with 205 additions and 50 deletions

View file

@ -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()

View file

@ -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):

View file

@ -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[

View file

@ -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):
"""