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