359 lines
10 KiB
Python
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()
|