mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-24 02:25:29 +00:00
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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user