mcp-forge/tests/mcp/test_bridge.py

360 lines
10 KiB
Python
Raw Permalink Normal View History

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