initial commit after one day coding agent session
This commit is contained in:
commit
372af75b90
88 changed files with 22694 additions and 0 deletions
359
tests/mcp/test_bridge.py
Normal file
359
tests/mcp/test_bridge.py
Normal file
|
|
@ -0,0 +1,359 @@
|
|||
"""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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue