mcp-forge/tests/mcp/test_bridge.py
2026-02-07 07:45:57 +01:00

359 lines
10 KiB
Python

"""Tests for Tool Bridge Server."""
import pytest
import socket
import json
import tempfile
from pathlib import Path
from unittest.mock import AsyncMock, Mock, patch
from mcp_forge.mcp.bridge import ToolBridgeServer
@pytest.fixture
def temp_socket_path():
"""Create temporary socket path."""
with tempfile.TemporaryDirectory() as tmpdir:
yield Path(tmpdir) / "test_bridge.sock"
@pytest.fixture
def mock_client_manager():
"""Create mock MCP client manager."""
manager = AsyncMock()
manager.call_tool = AsyncMock(return_value={"result": "success"})
return manager
@pytest.fixture
def mock_audit_logger():
"""Create mock audit logger."""
logger = Mock()
logger.log = Mock()
return logger
@pytest.mark.asyncio
async def test_bridge_server_starts_and_stops(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test that bridge server starts and stops cleanly."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
# Start server
bridge.start()
# Verify socket was created
assert temp_socket_path.exists()
# Stop server
bridge.stop()
# Verify socket was removed
assert not temp_socket_path.exists()
@pytest.mark.asyncio
async def test_receive_and_forward_tool_call(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test receiving tool call request and forwarding to client."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
# Connect as client and send tool call request
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
request = {
"tool": "test_tool",
"params": {"arg1": "value1", "arg2": 42}
}
client_sock.sendall(json.dumps(request).encode('utf-8'))
client_sock.shutdown(socket.SHUT_WR)
# Receive response
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
response = json.loads(response_data.decode('utf-8'))
# Verify response
assert response["success"] is True
assert response["result"] == {"result": "success"}
# Verify tool was called with correct arguments
mock_client_manager.call_tool.assert_called_once_with(
"test_tool",
{"arg1": "value1", "arg2": 42}
)
# Verify audit log was called
mock_audit_logger.log.assert_called()
client_sock.close()
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_handle_tool_call_error(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test handling tool call errors."""
# Make client manager raise error
mock_client_manager.call_tool = AsyncMock(side_effect=RuntimeError("Tool failed"))
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
# Connect and send request
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
request = {"tool": "failing_tool", "params": {}}
client_sock.sendall(json.dumps(request).encode('utf-8'))
client_sock.shutdown(socket.SHUT_WR)
# Receive response
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
response = json.loads(response_data.decode('utf-8'))
# Verify error response
assert response["success"] is False
assert "Tool failed" in response["error"]
client_sock.close()
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_handle_invalid_json(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test handling invalid JSON in request."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
# Connect and send invalid JSON
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
client_sock.sendall(b"not valid json")
client_sock.shutdown(socket.SHUT_WR)
# Receive response
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
response = json.loads(response_data.decode('utf-8'))
# Verify error response
assert response["success"] is False
assert "Invalid JSON" in response["error"]
client_sock.close()
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_handle_missing_tool_field(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test handling request missing 'tool' field."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
# Connect and send request without 'tool' field
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
request = {"params": {"arg1": "value1"}} # Missing 'tool'
client_sock.sendall(json.dumps(request).encode('utf-8'))
client_sock.shutdown(socket.SHUT_WR)
# Receive response
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
response = json.loads(response_data.decode('utf-8'))
# Verify error response
assert response["success"] is False
assert "tool" in response["error"].lower()
client_sock.close()
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_concurrent_requests(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test handling multiple concurrent requests."""
import threading
# Track call counts
call_count = {"count": 0}
async def mock_call_tool(tool_name, arguments):
call_count["count"] += 1
return {"result": f"success_{call_count['count']}"}
mock_client_manager.call_tool = mock_call_tool
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
results = []
def make_request(tool_name):
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
request = {"tool": tool_name, "params": {}}
client_sock.sendall(json.dumps(request).encode('utf-8'))
client_sock.shutdown(socket.SHUT_WR)
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
response = json.loads(response_data.decode('utf-8'))
results.append(response)
client_sock.close()
# Make 3 concurrent requests
threads = []
for i in range(3):
thread = threading.Thread(target=make_request, args=(f"tool_{i}",))
threads.append(thread)
thread.start()
# Wait for all threads
for thread in threads:
thread.join()
# Verify all requests succeeded
assert len(results) == 3
for result in results:
assert result["success"] is True
# Verify all were processed
assert call_count["count"] == 3
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_audit_logging_tool_name_only(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test that audit log only logs tool name, not parameters."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
try:
# Send request with sensitive parameters
client_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
client_sock.connect(str(temp_socket_path))
request = {
"tool": "sensitive_tool",
"params": {"password": "secret123", "token": "abc123"}
}
client_sock.sendall(json.dumps(request).encode('utf-8'))
client_sock.shutdown(socket.SHUT_WR)
# Receive response
response_data = b''
while True:
chunk = client_sock.recv(4096)
if not chunk:
break
response_data += chunk
client_sock.close()
# Verify audit log was called
mock_audit_logger.log.assert_called()
# Get the log call arguments
log_call = mock_audit_logger.log.call_args
# Verify tool name is in log
log_str = str(log_call)
assert "sensitive_tool" in log_str
# Verify sensitive parameters are NOT in log
assert "secret123" not in log_str
assert "abc123" not in log_str
finally:
bridge.stop()
@pytest.mark.asyncio
async def test_socket_cleanup_on_error(temp_socket_path, mock_client_manager, mock_audit_logger):
"""Test that socket is cleaned up even if server encounters error."""
bridge = ToolBridgeServer(
socket_path=temp_socket_path,
client_manager=mock_client_manager,
audit_logger=mock_audit_logger
)
bridge.start()
assert temp_socket_path.exists()
# Stop should cleanup
bridge.stop()
assert not temp_socket_path.exists()