Feat/persist mcp credentials in db (#16308)

* feat: persist mcp credentials in db

* feat: remove Auth Value field from MCP Tool Testing Playground

* fix: test
This commit is contained in:
YutaSaito
2025-11-07 19:22:49 -08:00
committed by GitHub
parent b6f792f301
commit 6eb74bd62a
22 changed files with 529 additions and 352 deletions
+3
View File
@@ -1342,6 +1342,7 @@ def test_add_update_server_with_alias():
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.auth_type = None
mock_mcp_server.credentials = {}
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}
mock_mcp_server.static_headers = {}
@@ -1380,6 +1381,7 @@ def test_add_update_server_without_alias():
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.auth_type = None
mock_mcp_server.credentials = {}
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}
mock_mcp_server.static_headers = {}
@@ -1418,6 +1420,7 @@ def test_add_update_server_fallback_to_server_id():
mock_mcp_server.url = "https://test-server.com/mcp"
mock_mcp_server.transport = MCPTransport.http
mock_mcp_server.auth_type = None
mock_mcp_server.credentials = {}
mock_mcp_server.description = "Test server description"
mock_mcp_server.mcp_info = {}
mock_mcp_server.static_headers = {}
@@ -15,6 +15,7 @@ from litellm.proxy._types import (
MCPTransportType,
MCPTransport,
NewMCPServerRequest,
UpdateMCPServerRequest,
LiteLLM_MCPServerTable,
LitellmUserRoles,
UserAPIKeyAuth,
@@ -165,6 +166,7 @@ async def test_create_mcp_server_direct():
updated_by=LITELLM_PROXY_ADMIN_NAME,
teams=[],
)
expected_response.credentials = {"auth_value": "secret"}
# Mock the database calls
mock_get_server.return_value = None # Server doesn't exist yet
@@ -188,6 +190,7 @@ async def test_create_mcp_server_direct():
assert result.alias == expected_alias # Check against normalized alias
assert result.url == mcp_server_request.url
assert result.transport == mcp_server_request.transport
assert result.credentials is None
# Verify mocks were called
mock_get_server.assert_called_once_with(mock_prisma, server_id)
@@ -353,6 +356,69 @@ async def test_create_mcp_server_invalid_alias():
)
@pytest.mark.asyncio
async def test_edit_mcp_server_redacts_credentials():
with mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
True,
), mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
) as mock_get_prisma, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
new_callable=mock.AsyncMock,
) as mock_update, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
autospec=True,
) as mock_validate, mock.patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager"
) as mock_manager:
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
edit_mcp_server,
)
mock_prisma = mock.Mock()
mock_get_prisma.return_value = mock_prisma
mock_manager.add_update_server = mock.Mock()
mock_manager.reload_servers_from_database = mock.AsyncMock()
server_id = str(uuid.uuid4())
updated_server = LiteLLM_MCPServerTable(
server_id=server_id,
alias="Updated Server",
url="https://updated.example.com/mcp",
transport=MCPTransport.http,
created_at=datetime.now(),
updated_at=datetime.now(),
teams=[],
)
updated_server.credentials = {"auth_value": "secret"}
mock_update.return_value = updated_server
payload = UpdateMCPServerRequest(
server_id=server_id,
alias="Updated Server",
url="https://updated.example.com/mcp",
transport=MCPTransport.http,
)
user_auth = UserAPIKeyAuth(
api_key=TEST_MASTER_KEY,
user_id="test-user",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
result = await edit_mcp_server(payload=payload, user_api_key_dict=user_auth)
assert result.server_id == server_id
assert result.credentials is None
assert updated_server.credentials == {"auth_value": "secret"}
mock_validate.assert_called_once()
mock_update.assert_awaited_once()
mock_manager.add_update_server.assert_called_once_with(updated_server)
mock_manager.reload_servers_from_database.assert_awaited_once()
def test_validate_mcp_server_name_direct():
"""
Test the validation function directly to ensure it works.
@@ -186,6 +186,9 @@ class TestListMCPServers:
return_value=mock_servers_with_health
)
for idx, server in enumerate(mock_servers_with_health):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@@ -205,6 +208,7 @@ class TestListMCPServers:
# Verify results
assert len(result) == 2
assert all(server.credentials is None for server in result)
# Check that both config servers are returned
server_ids = [server.server_id for server in result]
@@ -315,6 +319,9 @@ class TestListMCPServers:
return_value=mock_servers_with_health
)
for idx, server in enumerate(mock_servers_with_health):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@@ -334,6 +341,7 @@ class TestListMCPServers:
# Verify results
assert len(result) == 4
assert all(server.credentials is None for server in result)
# Check that both DB and config servers are returned
server_ids = [server.server_id for server in result]
@@ -428,6 +436,9 @@ class TestListMCPServers:
return_value=mock_servers_with_health
)
for idx, server in enumerate(mock_servers_with_health):
server.credentials = {"auth_value": f"secret_{idx}"}
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
@@ -447,6 +458,7 @@ class TestListMCPServers:
# Verify results - should only return servers user has access to
assert len(result) == 2
assert all(server.credentials is None for server in result)
# Check that only allowed servers are returned
server_ids = [server.server_id for server in result]
@@ -464,6 +476,51 @@ class TestListMCPServers:
assert server.url == "https://actions.zapier.com/mcp/sse"
@pytest.mark.asyncio
async def test_fetch_single_mcp_server_redacts_credentials(self):
mock_server = generate_mock_mcp_server_db_record(
server_id="server-1", alias="Server 1"
)
mock_server.credentials = {"auth_value": "top-secret"}
mock_prisma_client = MagicMock()
mock_health_result = {
"status": "healthy",
"last_health_check": datetime.now().isoformat(),
"error": None,
}
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma_client,
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
AsyncMock(return_value=mock_server),
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager.health_check_server",
AsyncMock(return_value=mock_health_result),
), patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=True,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_mcp_server,
)
result = await fetch_mcp_server(
server_id="server-1", user_api_key_dict=mock_user_auth
)
assert result.server_id == "server-1"
assert result.credentials is None
assert mock_server.credentials == {"auth_value": "top-secret"}
assert result.status == "healthy"
class TestMCPHealthCheckEndpoints:
"""Test MCP health check endpoints"""
@@ -721,6 +778,8 @@ class TestMCPHealthCheckEndpoints:
return_value=[mock_server]
)
mock_server.credentials = {"auth_value": "secret"}
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
@@ -749,3 +808,4 @@ class TestMCPHealthCheckEndpoints:
assert server.status == "healthy"
assert server.last_health_check is not None
assert server.health_check_error is None
assert server.credentials is None