refactor(session): migrate to server-based session management and add tests (#9077)
* update MCP Tests * [autofix.ci] apply automated fixes * Update util.py * [autofix.ci] apply automated fixes * Refactor MCP session manager for better configurability and cleanup (#9176) * Add log rotation and header validation features Introduces support for log rotation via the LANGFLOW_LOG_ROTATION environment variable and CLI/config options, with documentation updates. Adds header validation and sanitization for MCP connections, ensuring RFC 7230 compliance and security. Frontend and backend now support passing custom headers for MCP servers. Includes extensive new and updated unit tests for header handling, MCP utilities, and log rotation. * Add unit tests for MCP utilities and update disconnect logic Added comprehensive unit tests for MCP utility functions, session management, header validation, and client classes in test_mcp_util.py. Updated MCPStdioClient and MCPSseClient disconnect methods for clearer session cleanup logic. Refactored test_mcp_component.py to remove redundant and duplicated tests, consolidating coverage in the new test suite. * [autofix.ci] apply automated fixes * Update test_mcp_memory_leak.py * Update util.py --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Gabriel Luiz Freitas Almeida <gabriel@langflow.org>
This commit is contained in:
parent
80ebe03d94
commit
b093c1fadb
6 changed files with 1586 additions and 664 deletions
|
|
@ -0,0 +1,365 @@
|
|||
"""Integration tests for MCP memory leak fix.
|
||||
|
||||
These tests verify that the MCP session manager properly handles session reuse
|
||||
and cleanup to prevent subprocess leaks.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import platform
|
||||
import shutil
|
||||
|
||||
import psutil
|
||||
import pytest
|
||||
from langflow.base.mcp.util import MCPSessionManager
|
||||
from loguru import logger
|
||||
from mcp import StdioServerParameters
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mcp_server_params():
|
||||
"""Create MCP server parameters for testing."""
|
||||
command = ["npx", "-y", "@modelcontextprotocol/server-everything"]
|
||||
env_data = {"DEBUG": "true", "PATH": os.environ["PATH"]}
|
||||
|
||||
if platform.system() == "Windows":
|
||||
return StdioServerParameters(
|
||||
command="cmd",
|
||||
args=["/c", f"{command[0]} {' '.join(command[1:])}"],
|
||||
env=env_data,
|
||||
)
|
||||
return StdioServerParameters(
|
||||
command="bash",
|
||||
args=["-c", f"exec {' '.join(command)}"],
|
||||
env=env_data,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def process_tracker():
|
||||
"""Track subprocess count for memory leak detection."""
|
||||
process = psutil.Process()
|
||||
initial_count = len(process.children(recursive=True))
|
||||
|
||||
yield process, initial_count
|
||||
|
||||
# Cleanup any remaining child processes
|
||||
try:
|
||||
for child in process.children(recursive=True):
|
||||
try:
|
||||
child.terminate()
|
||||
child.wait(timeout=3)
|
||||
except (psutil.NoSuchProcess, psutil.TimeoutExpired):
|
||||
with contextlib.suppress(psutil.NoSuchProcess):
|
||||
child.kill()
|
||||
except Exception as e:
|
||||
logger.exception("Error cleaning up child processes: %s", e)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_session_reuse_prevents_subprocess_leak(mcp_server_params, process_tracker):
|
||||
"""Test that session reuse prevents subprocess proliferation."""
|
||||
process, initial_count = process_tracker
|
||||
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
try:
|
||||
# Create multiple sessions with different context IDs but same server
|
||||
sessions = []
|
||||
for i in range(3):
|
||||
context_id = f"test_context_{i}"
|
||||
session = await session_manager.get_session(context_id, mcp_server_params, "stdio")
|
||||
sessions.append(session)
|
||||
|
||||
# Verify session is working
|
||||
tools_response = await session.list_tools()
|
||||
assert len(tools_response.tools) > 0
|
||||
|
||||
# Check subprocess count after creating sessions
|
||||
current_count = len(process.children(recursive=True))
|
||||
subprocess_increase = current_count - initial_count
|
||||
|
||||
# With the fix, we should have minimal subprocess increase
|
||||
# (ideally 2 subprocesses max for the MCP server)
|
||||
assert subprocess_increase <= 4, f"Too many subprocesses created: {subprocess_increase}"
|
||||
|
||||
# Verify all sessions are functional
|
||||
for session in sessions:
|
||||
tools_response = await session.list_tools()
|
||||
assert len(tools_response.tools) > 0
|
||||
|
||||
finally:
|
||||
await session_manager.cleanup_all()
|
||||
await asyncio.sleep(2) # Allow cleanup to complete
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_session_cleanup_removes_subprocesses(mcp_server_params, process_tracker):
|
||||
"""Test that session cleanup properly removes subprocesses."""
|
||||
process, initial_count = process_tracker
|
||||
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
try:
|
||||
# Create a session
|
||||
session = await session_manager.get_session("cleanup_test", mcp_server_params, "stdio")
|
||||
tools_response = await session.list_tools()
|
||||
assert len(tools_response.tools) > 0
|
||||
|
||||
# Verify subprocess was created
|
||||
after_creation_count = len(process.children(recursive=True))
|
||||
assert after_creation_count > initial_count
|
||||
|
||||
finally:
|
||||
# Clean up session
|
||||
await session_manager.cleanup_all()
|
||||
await asyncio.sleep(2) # Allow cleanup to complete
|
||||
|
||||
# Verify subprocess was cleaned up
|
||||
after_cleanup_count = len(process.children(recursive=True))
|
||||
# Allow some tolerance for cleanup timing and system processes
|
||||
assert after_cleanup_count <= initial_count + 1, (
|
||||
f"Subprocesses not cleaned up properly: {after_cleanup_count} vs {initial_count}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_session_health_check_and_recovery(mcp_server_params, process_tracker):
|
||||
"""Test that unhealthy sessions are properly detected and recreated."""
|
||||
process, initial_count = process_tracker
|
||||
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
try:
|
||||
# Create a session
|
||||
session1 = await session_manager.get_session("health_test", mcp_server_params, "stdio")
|
||||
tools_response = await session1.list_tools()
|
||||
assert len(tools_response.tools) > 0
|
||||
|
||||
# Simulate session becoming unhealthy by accessing internal state
|
||||
# This is a bit of a hack but necessary for testing
|
||||
server_key = session_manager._get_server_key(mcp_server_params, "stdio")
|
||||
if hasattr(session_manager, "sessions_by_server"):
|
||||
# For the fixed version
|
||||
sessions = session_manager.sessions_by_server.get(server_key, {})
|
||||
if sessions:
|
||||
session_id = next(iter(sessions.keys()))
|
||||
session_info = sessions[session_id]
|
||||
if "task" in session_info:
|
||||
task = session_info["task"]
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await task
|
||||
elif hasattr(session_manager, "sessions"):
|
||||
# For the original version
|
||||
for session_info in session_manager.sessions.values():
|
||||
if "task" in session_info:
|
||||
task = session_info["task"]
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
# Wait a bit for the task to be cancelled
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Try to get a session again - should create a new healthy one
|
||||
session2 = await session_manager.get_session("health_test_2", mcp_server_params, "stdio")
|
||||
tools_response = await session2.list_tools()
|
||||
assert len(tools_response.tools) > 0
|
||||
|
||||
finally:
|
||||
await session_manager.cleanup_all()
|
||||
await asyncio.sleep(2)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_multiple_servers_isolation(process_tracker):
|
||||
"""Test that different servers get separate sessions."""
|
||||
process, initial_count = process_tracker
|
||||
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
# Create parameters for different servers
|
||||
server1_params = StdioServerParameters(
|
||||
command="bash",
|
||||
args=["-c", "exec npx -y @modelcontextprotocol/server-everything"],
|
||||
env={"DEBUG": "true", "PATH": os.environ["PATH"]},
|
||||
)
|
||||
|
||||
server2_params = StdioServerParameters(
|
||||
command="bash",
|
||||
args=["-c", "exec npx -y @modelcontextprotocol/server-everything"],
|
||||
env={"DEBUG": "false", "PATH": os.environ["PATH"]}, # Different env
|
||||
)
|
||||
|
||||
try:
|
||||
# Create sessions for different servers
|
||||
session1 = await session_manager.get_session("server1_test", server1_params, "stdio")
|
||||
session2 = await session_manager.get_session("server2_test", server2_params, "stdio")
|
||||
|
||||
# Verify both sessions work
|
||||
tools1 = await session1.list_tools()
|
||||
tools2 = await session2.list_tools()
|
||||
|
||||
assert len(tools1.tools) > 0
|
||||
assert len(tools2.tools) > 0
|
||||
|
||||
# Sessions should be different objects for different servers (different environments)
|
||||
# Since the servers have different environments, they should get different server keys
|
||||
server_key1 = session_manager._get_server_key(server1_params, "stdio")
|
||||
server_key2 = session_manager._get_server_key(server2_params, "stdio")
|
||||
assert server_key1 != server_key2, "Different server environments should generate different keys"
|
||||
assert session1 is not session2
|
||||
|
||||
finally:
|
||||
await session_manager.cleanup_all()
|
||||
await asyncio.sleep(2)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_manager_server_key_generation():
|
||||
"""Test that server key generation works correctly."""
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
# Test stdio server key
|
||||
stdio_params = StdioServerParameters(
|
||||
command="test_command",
|
||||
args=["arg1", "arg2"],
|
||||
env={"TEST": "value"},
|
||||
)
|
||||
|
||||
key1 = session_manager._get_server_key(stdio_params, "stdio")
|
||||
key2 = session_manager._get_server_key(stdio_params, "stdio")
|
||||
|
||||
# Same parameters should generate same key
|
||||
assert key1 == key2
|
||||
assert key1.startswith("stdio_")
|
||||
|
||||
# Different parameters should generate different keys
|
||||
stdio_params2 = StdioServerParameters(
|
||||
command="different_command",
|
||||
args=["arg1", "arg2"],
|
||||
env={"TEST": "value"},
|
||||
)
|
||||
|
||||
key3 = session_manager._get_server_key(stdio_params2, "stdio")
|
||||
assert key1 != key3
|
||||
|
||||
# Test SSE server key
|
||||
sse_params = {
|
||||
"url": "http://example.com/sse",
|
||||
"headers": {"Authorization": "Bearer token"},
|
||||
"timeout_seconds": 30,
|
||||
"sse_read_timeout_seconds": 30,
|
||||
}
|
||||
|
||||
sse_key1 = session_manager._get_server_key(sse_params, "sse")
|
||||
sse_key2 = session_manager._get_server_key(sse_params, "sse")
|
||||
|
||||
assert sse_key1 == sse_key2
|
||||
assert sse_key1.startswith("sse_")
|
||||
|
||||
# Different URL should generate different key
|
||||
sse_params2 = sse_params.copy()
|
||||
sse_params2["url"] = "http://different.com/sse"
|
||||
|
||||
sse_key3 = session_manager._get_server_key(sse_params2, "sse")
|
||||
assert sse_key1 != sse_key3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_manager_connectivity_validation():
|
||||
"""Test session connectivity validation."""
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
# Mock a session that responds to list_tools
|
||||
class MockSession:
|
||||
def __init__(self, should_fail=False): # noqa: FBT002
|
||||
self.should_fail = should_fail
|
||||
|
||||
async def list_tools(self):
|
||||
if self.should_fail:
|
||||
msg = "Connection failed"
|
||||
raise Exception(msg) # noqa: TRY002
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self):
|
||||
self.tools = ["tool1", "tool2"]
|
||||
|
||||
return MockResponse()
|
||||
|
||||
# Test healthy session
|
||||
healthy_session = MockSession(should_fail=False)
|
||||
is_healthy = await session_manager._validate_session_connectivity(healthy_session)
|
||||
assert is_healthy is True
|
||||
|
||||
# Test unhealthy session
|
||||
unhealthy_session = MockSession(should_fail=True)
|
||||
is_healthy = await session_manager._validate_session_connectivity(unhealthy_session)
|
||||
assert is_healthy is False
|
||||
|
||||
# Test session that returns None
|
||||
class MockNoneSession:
|
||||
async def list_tools(self):
|
||||
return None
|
||||
|
||||
none_session = MockNoneSession()
|
||||
is_healthy = await session_manager._validate_session_connectivity(none_session)
|
||||
assert is_healthy is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_manager_cleanup_all():
|
||||
"""Test that cleanup_all properly cleans up all sessions."""
|
||||
session_manager = MCPSessionManager()
|
||||
|
||||
# Mock some sessions using the correct structure
|
||||
session_manager.sessions_by_server = {
|
||||
"server1": {
|
||||
"sessions": {
|
||||
"session1": {
|
||||
"session": "mock_session",
|
||||
"task": asyncio.create_task(asyncio.sleep(10)),
|
||||
"type": "stdio",
|
||||
"last_used": asyncio.get_event_loop().time(),
|
||||
}
|
||||
}
|
||||
},
|
||||
"server2": {
|
||||
"sessions": {
|
||||
"session2": {
|
||||
"session": "mock_session",
|
||||
"task": asyncio.create_task(asyncio.sleep(10)),
|
||||
"type": "sse",
|
||||
"last_used": asyncio.get_event_loop().time(),
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
# Add some background tasks
|
||||
task1 = asyncio.create_task(asyncio.sleep(10))
|
||||
task2 = asyncio.create_task(asyncio.sleep(10))
|
||||
session_manager._background_tasks = {task1, task2}
|
||||
|
||||
# Cleanup all
|
||||
await session_manager.cleanup_all()
|
||||
|
||||
# Verify cleanup
|
||||
if hasattr(session_manager, "sessions_by_server"):
|
||||
# For fixed version
|
||||
assert len(session_manager.sessions_by_server) == 0
|
||||
elif hasattr(session_manager, "sessions"):
|
||||
# For original version
|
||||
assert len(session_manager.sessions) == 0
|
||||
|
||||
# Verify background tasks were cancelled
|
||||
assert task1.done()
|
||||
assert task2.done()
|
||||
0
src/backend/tests/unit/base/mcp/__init__.py
Normal file
0
src/backend/tests/unit/base/mcp/__init__.py
Normal file
806
src/backend/tests/unit/base/mcp/test_mcp_util.py
Normal file
806
src/backend/tests/unit/base/mcp/test_mcp_util.py
Normal file
|
|
@ -0,0 +1,806 @@
|
|||
"""Unit tests for MCP utility functions.
|
||||
|
||||
This test suite validates the MCP utility functions including:
|
||||
- Session management
|
||||
- Header validation and processing
|
||||
- Utility functions for name sanitization and schema conversion
|
||||
"""
|
||||
|
||||
import shutil
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from langflow.base.mcp import util
|
||||
from langflow.base.mcp.util import MCPSessionManager, MCPSseClient, MCPStdioClient, _process_headers, validate_headers
|
||||
|
||||
|
||||
class TestMCPSessionManager:
|
||||
@pytest.fixture
|
||||
async def session_manager(self):
|
||||
"""Create a session manager and clean it up after the test."""
|
||||
manager = MCPSessionManager()
|
||||
yield manager
|
||||
# Clean up after test
|
||||
await manager.cleanup_all()
|
||||
|
||||
async def test_session_caching(self, session_manager):
|
||||
"""Test that sessions are properly cached and reused."""
|
||||
context_id = "test_context"
|
||||
connection_params = MagicMock()
|
||||
transport_type = "stdio"
|
||||
|
||||
# Create a mock session that will appear healthy
|
||||
mock_session = AsyncMock()
|
||||
mock_session._write_stream = MagicMock()
|
||||
mock_session._write_stream._closed = False
|
||||
|
||||
# Create a mock task that appears to be running
|
||||
mock_task = AsyncMock()
|
||||
mock_task.done = MagicMock(return_value=False)
|
||||
|
||||
with (
|
||||
patch.object(session_manager, "_create_stdio_session") as mock_create,
|
||||
patch.object(session_manager, "_validate_session_connectivity", return_value=True),
|
||||
):
|
||||
mock_create.return_value = (mock_session, mock_task)
|
||||
|
||||
# First call should create session
|
||||
session1 = await session_manager.get_session(context_id, connection_params, transport_type)
|
||||
|
||||
# Second call should return cached session without creating new one
|
||||
session2 = await session_manager.get_session(context_id, connection_params, transport_type)
|
||||
|
||||
assert session1 == session2
|
||||
assert session1 == mock_session
|
||||
# Should only create once since the second call should use the cached session
|
||||
mock_create.assert_called_once()
|
||||
|
||||
async def test_session_cleanup(self, session_manager):
|
||||
"""Test session cleanup functionality."""
|
||||
context_id = "test_context"
|
||||
server_key = "test_server"
|
||||
session_id = "test_session"
|
||||
|
||||
# Add a session to the manager with proper mock setup using new structure
|
||||
mock_task = AsyncMock()
|
||||
mock_task.done = MagicMock(return_value=False) # Use MagicMock for sync method
|
||||
mock_task.cancel = MagicMock() # Use MagicMock for sync method
|
||||
|
||||
# Set up the new session structure
|
||||
session_manager.sessions_by_server[server_key] = {
|
||||
"sessions": {session_id: {"session": AsyncMock(), "task": mock_task, "type": "stdio", "last_used": 0}},
|
||||
"last_cleanup": 0,
|
||||
}
|
||||
|
||||
# Set up mapping for backwards compatibility
|
||||
session_manager._context_to_session[context_id] = (server_key, session_id)
|
||||
|
||||
await session_manager._cleanup_session(context_id)
|
||||
|
||||
# Should cancel the task and remove from sessions
|
||||
mock_task.cancel.assert_called_once()
|
||||
assert session_id not in session_manager.sessions_by_server[server_key]["sessions"]
|
||||
|
||||
async def test_server_switch_detection(self, session_manager):
|
||||
"""Test that server switches are properly detected and handled."""
|
||||
context_id = "test_context"
|
||||
|
||||
# First server
|
||||
server1_params = MagicMock()
|
||||
server1_params.command = "server1"
|
||||
|
||||
# Second server
|
||||
server2_params = MagicMock()
|
||||
server2_params.command = "server2"
|
||||
|
||||
with (
|
||||
patch.object(session_manager, "_create_stdio_session") as mock_create,
|
||||
patch.object(session_manager, "_validate_session_connectivity", return_value=True),
|
||||
):
|
||||
mock_session1 = AsyncMock()
|
||||
mock_session2 = AsyncMock()
|
||||
mock_task1 = AsyncMock()
|
||||
mock_task2 = AsyncMock()
|
||||
mock_create.side_effect = [(mock_session1, mock_task1), (mock_session2, mock_task2)]
|
||||
|
||||
# First connection
|
||||
session1 = await session_manager.get_session(context_id, server1_params, "stdio")
|
||||
|
||||
# Switch to different server should create new session
|
||||
session2 = await session_manager.get_session(context_id, server2_params, "stdio")
|
||||
|
||||
assert session1 != session2
|
||||
assert mock_create.call_count == 2
|
||||
|
||||
|
||||
class TestHeaderValidation:
|
||||
"""Test the header validation functionality."""
|
||||
|
||||
def test_validate_headers_valid_input(self):
|
||||
"""Test header validation with valid headers."""
|
||||
headers = {"Authorization": "Bearer token123", "Content-Type": "application/json", "X-API-Key": "secret-key"}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Headers should be normalized to lowercase
|
||||
expected = {"authorization": "Bearer token123", "content-type": "application/json", "x-api-key": "secret-key"}
|
||||
assert result == expected
|
||||
|
||||
def test_validate_headers_empty_input(self):
|
||||
"""Test header validation with empty/None input."""
|
||||
assert validate_headers({}) == {}
|
||||
assert validate_headers(None) == {}
|
||||
|
||||
def test_validate_headers_invalid_names(self):
|
||||
"""Test header validation with invalid header names."""
|
||||
headers = {
|
||||
"Invalid Header": "value", # spaces not allowed
|
||||
"Header@Name": "value", # @ not allowed
|
||||
"Header Name": "value", # spaces not allowed
|
||||
"Valid-Header": "value", # this should pass
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Only the valid header should remain
|
||||
assert result == {"valid-header": "value"}
|
||||
|
||||
def test_validate_headers_sanitize_values(self):
|
||||
"""Test header value sanitization."""
|
||||
headers = {
|
||||
"Authorization": "Bearer \x00token\x1f with\r\ninjection",
|
||||
"Clean-Header": " clean value ",
|
||||
"Empty-After-Clean": "\x00\x01\x02",
|
||||
"Tab-Header": "value\twith\ttabs", # tabs should be preserved
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Control characters should be removed, whitespace trimmed
|
||||
# Header with injection attempts should be skipped
|
||||
expected = {"clean-header": "clean value", "tab-header": "value\twith\ttabs"}
|
||||
assert result == expected
|
||||
|
||||
def test_validate_headers_non_string_values(self):
|
||||
"""Test header validation with non-string values."""
|
||||
headers = {"String-Header": "valid", "Number-Header": 123, "None-Header": None, "List-Header": ["value"]}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Only string headers should remain
|
||||
assert result == {"string-header": "valid"}
|
||||
|
||||
def test_validate_headers_injection_attempts(self):
|
||||
"""Test header validation against injection attempts."""
|
||||
headers = {
|
||||
"Injection1": "value\r\nInjected-Header: malicious",
|
||||
"Injection2": "value\nX-Evil: attack",
|
||||
"Safe-Header": "safe-value",
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Injection attempts should be filtered out
|
||||
assert result == {"safe-header": "safe-value"}
|
||||
|
||||
|
||||
class TestSSEHeaderIntegration:
|
||||
"""Integration test to verify headers are properly passed through the entire SSE flow."""
|
||||
|
||||
async def test_headers_processing(self):
|
||||
"""Test that headers flow properly from server config through to SSE client connection."""
|
||||
# Test the header processing function directly
|
||||
headers_input = [
|
||||
{"key": "Authorization", "value": "Bearer test-token"},
|
||||
{"key": "X-API-Key", "value": "secret-key"},
|
||||
]
|
||||
|
||||
expected_headers = {
|
||||
"authorization": "Bearer test-token", # normalized to lowercase
|
||||
"x-api-key": "secret-key",
|
||||
}
|
||||
|
||||
# Test _process_headers function with validation
|
||||
processed_headers = _process_headers(headers_input)
|
||||
assert processed_headers == expected_headers
|
||||
|
||||
# Test different input formats
|
||||
# Test dict input with validation
|
||||
dict_headers = {"Authorization": "Bearer dict-token", "Invalid Header": "bad"}
|
||||
result = _process_headers(dict_headers)
|
||||
# Invalid header should be filtered out, valid header normalized
|
||||
assert result == {"authorization": "Bearer dict-token"}
|
||||
|
||||
# Test None input
|
||||
assert _process_headers(None) == {}
|
||||
|
||||
# Test empty list
|
||||
assert _process_headers([]) == {}
|
||||
|
||||
# Test malformed list
|
||||
malformed_headers = [{"key": "Auth"}, {"value": "token"}] # Missing value/key
|
||||
assert _process_headers(malformed_headers) == {}
|
||||
|
||||
# Test list with invalid header names
|
||||
invalid_headers = [
|
||||
{"key": "Valid-Header", "value": "good"},
|
||||
{"key": "Invalid Header", "value": "bad"}, # spaces not allowed
|
||||
]
|
||||
result = _process_headers(invalid_headers)
|
||||
assert result == {"valid-header": "good"}
|
||||
|
||||
async def test_sse_client_header_storage(self):
|
||||
"""Test that SSE client properly stores headers in connection params."""
|
||||
sse_client = MCPSseClient()
|
||||
test_url = "http://test.url"
|
||||
test_headers = {"Authorization": "Bearer test123", "Custom": "value"}
|
||||
|
||||
# Test that headers are properly stored in connection params
|
||||
# Set connection params as a dict like the implementation expects
|
||||
sse_client._connection_params = {
|
||||
"url": test_url,
|
||||
"headers": test_headers,
|
||||
"timeout_seconds": 30,
|
||||
"sse_read_timeout_seconds": 30,
|
||||
}
|
||||
|
||||
# Verify headers are stored
|
||||
assert sse_client._connection_params["url"] == test_url
|
||||
assert sse_client._connection_params["headers"] == test_headers
|
||||
|
||||
|
||||
class TestMCPUtilityFunctions:
|
||||
"""Test utility functions from util.py that don't have dedicated test classes."""
|
||||
|
||||
def test_sanitize_mcp_name(self):
|
||||
"""Test MCP name sanitization."""
|
||||
assert util.sanitize_mcp_name("Test Name 123") == "test_name_123"
|
||||
assert util.sanitize_mcp_name(" ") == ""
|
||||
assert util.sanitize_mcp_name("123abc") == "_123abc"
|
||||
assert util.sanitize_mcp_name("Tést-😀-Námé") == "test_name"
|
||||
assert util.sanitize_mcp_name("a" * 100) == "a" * 46
|
||||
|
||||
def test_get_unique_name(self):
|
||||
"""Test unique name generation."""
|
||||
names = {"foo", "foo_1"}
|
||||
assert util.get_unique_name("foo", 10, names) == "foo_2"
|
||||
assert util.get_unique_name("bar", 10, names) == "bar"
|
||||
assert util.get_unique_name("longname", 4, {"long"}) == "lo_1"
|
||||
|
||||
def test_is_valid_key_value_item(self):
|
||||
"""Test key-value item validation."""
|
||||
assert util._is_valid_key_value_item({"key": "a", "value": "b"}) is True
|
||||
assert util._is_valid_key_value_item({"key": "a"}) is False
|
||||
assert util._is_valid_key_value_item(["key", "value"]) is False
|
||||
assert util._is_valid_key_value_item(None) is False
|
||||
|
||||
def test_validate_node_installation(self):
|
||||
"""Test Node.js installation validation."""
|
||||
if shutil.which("node"):
|
||||
assert util._validate_node_installation("npx something") == "npx something"
|
||||
else:
|
||||
with pytest.raises(ValueError, match="Node.js is not installed"):
|
||||
util._validate_node_installation("npx something")
|
||||
assert util._validate_node_installation("echo test") == "echo test"
|
||||
|
||||
def test_create_input_schema_from_json_schema(self):
|
||||
"""Test JSON schema to Pydantic model conversion."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"foo": {"type": "string", "description": "desc"},
|
||||
"bar": {"type": "integer"},
|
||||
},
|
||||
"required": ["foo"],
|
||||
}
|
||||
model_class = util.create_input_schema_from_json_schema(schema)
|
||||
instance = model_class(foo="abc", bar=1)
|
||||
assert instance.foo == "abc"
|
||||
assert instance.bar == 1
|
||||
|
||||
with pytest.raises(Exception): # noqa: B017, PT011
|
||||
model_class(bar=1) # missing required field
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_connection_params(self):
|
||||
"""Test connection parameter validation."""
|
||||
# Valid parameters
|
||||
await util._validate_connection_params("Stdio", command="echo test")
|
||||
await util._validate_connection_params("SSE", url="http://test")
|
||||
|
||||
# Invalid parameters
|
||||
with pytest.raises(ValueError, match="Command is required for Stdio mode"):
|
||||
await util._validate_connection_params("Stdio", command=None)
|
||||
with pytest.raises(ValueError, match="URL is required for SSE mode"):
|
||||
await util._validate_connection_params("SSE", url=None)
|
||||
with pytest.raises(ValueError, match="Invalid mode"):
|
||||
await util._validate_connection_params("InvalidMode")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_flow_snake_case_mocked(self):
|
||||
"""Test flow lookup by snake case name with mocked session."""
|
||||
|
||||
class DummyFlow:
|
||||
def __init__(self, name: str, user_id: str, *, is_component: bool = False, action_name: str | None = None):
|
||||
self.name = name
|
||||
self.user_id = user_id
|
||||
self.is_component = is_component
|
||||
self.action_name = action_name
|
||||
|
||||
class DummyExec:
|
||||
def __init__(self, flows: list[DummyFlow]):
|
||||
self._flows = flows
|
||||
|
||||
def all(self):
|
||||
return self._flows
|
||||
|
||||
class DummySession:
|
||||
def __init__(self, flows: list[DummyFlow]):
|
||||
self._flows = flows
|
||||
|
||||
async def exec(self, stmt): # noqa: ARG002
|
||||
return DummyExec(self._flows)
|
||||
|
||||
user_id = "123e4567-e89b-12d3-a456-426614174000"
|
||||
flows = [DummyFlow("Test Flow", user_id), DummyFlow("Other", user_id)]
|
||||
|
||||
# Should match sanitized name
|
||||
result = await util.get_flow_snake_case(util.sanitize_mcp_name("Test Flow"), user_id, DummySession(flows))
|
||||
assert result is flows[0]
|
||||
|
||||
# Should return None if not found
|
||||
result = await util.get_flow_snake_case("notfound", user_id, DummySession(flows))
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestMCPStdioClientWithEverythingServer:
|
||||
"""Test MCPStdioClient with the Everything MCP server."""
|
||||
|
||||
@pytest.fixture
|
||||
def stdio_client(self):
|
||||
"""Create a stdio client for testing."""
|
||||
return MCPStdioClient()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_connect_to_everything_server(self, stdio_client):
|
||||
"""Test connecting to the Everything MCP server."""
|
||||
command = "npx -y @modelcontextprotocol/server-everything"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await stdio_client.connect_to_server(command)
|
||||
|
||||
# Verify tools were returned
|
||||
assert len(tools) > 0
|
||||
|
||||
# Find the echo tool
|
||||
echo_tool = None
|
||||
for tool in tools:
|
||||
if hasattr(tool, "name") and tool.name == "echo":
|
||||
echo_tool = tool
|
||||
break
|
||||
|
||||
assert echo_tool is not None, "Echo tool not found in server tools"
|
||||
assert echo_tool.description is not None
|
||||
|
||||
# Verify the echo tool has the expected input schema
|
||||
assert hasattr(echo_tool, "inputSchema")
|
||||
assert echo_tool.inputSchema is not None
|
||||
|
||||
finally:
|
||||
# Clean up the connection
|
||||
await stdio_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_run_echo_tool(self, stdio_client):
|
||||
"""Test running the echo tool from the Everything server."""
|
||||
command = "npx -y @modelcontextprotocol/server-everything"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await stdio_client.connect_to_server(command)
|
||||
|
||||
# Find the echo tool
|
||||
echo_tool = None
|
||||
for tool in tools:
|
||||
if hasattr(tool, "name") and tool.name == "echo":
|
||||
echo_tool = tool
|
||||
break
|
||||
|
||||
assert echo_tool is not None, "Echo tool not found"
|
||||
|
||||
# Run the echo tool
|
||||
test_message = "Hello, MCP!"
|
||||
result = await stdio_client.run_tool("echo", {"message": test_message})
|
||||
|
||||
# Verify the result
|
||||
assert result is not None
|
||||
assert hasattr(result, "content")
|
||||
assert len(result.content) > 0
|
||||
|
||||
# Check that the echo worked - content should contain our message
|
||||
content_text = str(result.content[0])
|
||||
assert test_message in content_text or "Echo:" in content_text
|
||||
|
||||
finally:
|
||||
await stdio_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_list_all_tools(self, stdio_client):
|
||||
"""Test listing all available tools from the Everything server."""
|
||||
command = "npx -y @modelcontextprotocol/server-everything"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await stdio_client.connect_to_server(command)
|
||||
|
||||
# Verify we have multiple tools
|
||||
assert len(tools) >= 3 # Everything server typically has several tools
|
||||
|
||||
# Check that tools have the expected attributes
|
||||
for tool in tools:
|
||||
assert hasattr(tool, "name")
|
||||
assert hasattr(tool, "description")
|
||||
assert hasattr(tool, "inputSchema")
|
||||
assert tool.name is not None
|
||||
assert len(tool.name) > 0
|
||||
|
||||
# Common tools that should be available
|
||||
expected_tools = ["echo"] # Echo is typically available
|
||||
for expected_tool in expected_tools:
|
||||
assert any(tool.name == expected_tool for tool in tools), f"Expected tool '{expected_tool}' not found"
|
||||
|
||||
finally:
|
||||
await stdio_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_session_reuse(self, stdio_client):
|
||||
"""Test that sessions are properly reused."""
|
||||
command = "npx -y @modelcontextprotocol/server-everything"
|
||||
|
||||
try:
|
||||
# Set session context
|
||||
stdio_client.set_session_context("test_session_reuse")
|
||||
|
||||
# Connect to the server
|
||||
tools1 = await stdio_client.connect_to_server(command)
|
||||
|
||||
# Connect again - should reuse the session
|
||||
tools2 = await stdio_client.connect_to_server(command)
|
||||
|
||||
# Should have the same tools
|
||||
assert len(tools1) == len(tools2)
|
||||
|
||||
# Run a tool to verify the session is working
|
||||
result = await stdio_client.run_tool("echo", {"message": "Session reuse test"})
|
||||
assert result is not None
|
||||
|
||||
finally:
|
||||
await stdio_client.disconnect()
|
||||
|
||||
|
||||
class TestMCPSseClientWithDeepWikiServer:
|
||||
"""Test MCPSseClient with the DeepWiki MCP server."""
|
||||
|
||||
@pytest.fixture
|
||||
def sse_client(self):
|
||||
"""Create an SSE client for testing."""
|
||||
return MCPSseClient()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_to_deepwiki_server(self, sse_client):
|
||||
"""Test connecting to the DeepWiki MCP server."""
|
||||
url = "https://mcp.deepwiki.com/sse"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await sse_client.connect_to_server(url)
|
||||
|
||||
# Verify tools were returned
|
||||
assert len(tools) > 0
|
||||
|
||||
# Check for expected DeepWiki tools
|
||||
expected_tools = ["read_wiki_structure", "read_wiki_contents", "ask_question"]
|
||||
|
||||
# Verify we have the expected tools
|
||||
for expected_tool in expected_tools:
|
||||
assert any(tool.name == expected_tool for tool in tools), f"Expected tool '{expected_tool}' not found"
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible, skip the test
|
||||
pytest.skip(f"DeepWiki server not accessible: {e}")
|
||||
finally:
|
||||
await sse_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_wiki_structure_tool(self, sse_client):
|
||||
"""Test running the read_wiki_structure tool."""
|
||||
url = "https://mcp.deepwiki.com/sse"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await sse_client.connect_to_server(url)
|
||||
|
||||
# Find the read_wiki_structure tool
|
||||
wiki_tool = None
|
||||
for tool in tools:
|
||||
if hasattr(tool, "name") and tool.name == "read_wiki_structure":
|
||||
wiki_tool = tool
|
||||
break
|
||||
|
||||
assert wiki_tool is not None, "read_wiki_structure tool not found"
|
||||
|
||||
# Run the tool with a test repository (use repoName as expected by the API)
|
||||
result = await sse_client.run_tool("read_wiki_structure", {"repoName": "microsoft/vscode"})
|
||||
|
||||
# Verify the result
|
||||
assert result is not None
|
||||
assert hasattr(result, "content")
|
||||
assert len(result.content) > 0
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible or the tool fails, skip the test
|
||||
pytest.skip(f"DeepWiki server test failed: {e}")
|
||||
finally:
|
||||
await sse_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ask_question_tool(self, sse_client):
|
||||
"""Test running the ask_question tool."""
|
||||
url = "https://mcp.deepwiki.com/sse"
|
||||
|
||||
try:
|
||||
# Connect to the server
|
||||
tools = await sse_client.connect_to_server(url)
|
||||
|
||||
# Find the ask_question tool
|
||||
ask_tool = None
|
||||
for tool in tools:
|
||||
if hasattr(tool, "name") and tool.name == "ask_question":
|
||||
ask_tool = tool
|
||||
break
|
||||
|
||||
assert ask_tool is not None, "ask_question tool not found"
|
||||
|
||||
# Run the tool with a test question (use repoName as expected by the API)
|
||||
result = await sse_client.run_tool(
|
||||
"ask_question", {"repoName": "microsoft/vscode", "question": "What is VS Code?"}
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
assert result is not None
|
||||
assert hasattr(result, "content")
|
||||
assert len(result.content) > 0
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible or the tool fails, skip the test
|
||||
pytest.skip(f"DeepWiki server test failed: {e}")
|
||||
finally:
|
||||
await sse_client.disconnect()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_url_validation(self, sse_client):
|
||||
"""Test URL validation for SSE connections."""
|
||||
# Test valid URL
|
||||
valid_url = "https://mcp.deepwiki.com/sse"
|
||||
is_valid, error = await sse_client.validate_url(valid_url)
|
||||
assert is_valid or error == "" # Either valid or accessible
|
||||
|
||||
# Test invalid URL
|
||||
invalid_url = "not_a_url"
|
||||
is_valid, error = await sse_client.validate_url(invalid_url)
|
||||
assert not is_valid
|
||||
assert error != ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redirect_handling(self, sse_client):
|
||||
"""Test redirect handling for SSE connections."""
|
||||
# Test with the DeepWiki URL
|
||||
url = "https://mcp.deepwiki.com/sse"
|
||||
|
||||
try:
|
||||
# Check for redirects
|
||||
final_url = await sse_client.pre_check_redirect(url)
|
||||
|
||||
# Should return a URL (either original or redirected)
|
||||
assert final_url is not None
|
||||
assert isinstance(final_url, str)
|
||||
assert final_url.startswith("http")
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible, skip the test
|
||||
pytest.skip(f"DeepWiki server not accessible for redirect test: {e}")
|
||||
|
||||
@pytest.fixture
|
||||
def mock_tool(self):
|
||||
"""Create a mock MCP tool."""
|
||||
tool = MagicMock()
|
||||
tool.name = "test_tool"
|
||||
tool.description = "Test tool description"
|
||||
tool.inputSchema = {
|
||||
"type": "object",
|
||||
"properties": {"test_param": {"type": "string", "description": "Test parameter"}},
|
||||
"required": ["test_param"],
|
||||
}
|
||||
return tool
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session(self, mock_tool):
|
||||
"""Create a mock ClientSession."""
|
||||
session = AsyncMock()
|
||||
session.initialize = AsyncMock()
|
||||
list_tools_result = MagicMock()
|
||||
list_tools_result.tools = [mock_tool]
|
||||
session.list_tools = AsyncMock(return_value=list_tools_result)
|
||||
session.call_tool = AsyncMock(
|
||||
return_value=MagicMock(content=[MagicMock(model_dump=lambda: {"result": "success"})])
|
||||
)
|
||||
return session
|
||||
|
||||
|
||||
class TestMCPSseClientUnit:
|
||||
"""Unit tests for MCPSseClient functionality."""
|
||||
|
||||
@pytest.fixture
|
||||
def sse_client(self):
|
||||
return MCPSseClient()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_initialization(self, sse_client):
|
||||
"""Test that SSE client initializes correctly."""
|
||||
# Client should initialize with default values
|
||||
assert sse_client.session is None
|
||||
assert sse_client._connection_params is None
|
||||
assert sse_client._connected is False
|
||||
assert sse_client._session_context is None
|
||||
|
||||
async def test_validate_url_valid(self, sse_client):
|
||||
"""Test URL validation with valid URL."""
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
is_valid, error_msg = await sse_client.validate_url("http://test.url", {})
|
||||
|
||||
assert is_valid is True
|
||||
assert error_msg == ""
|
||||
|
||||
async def test_validate_url_invalid_format(self, sse_client):
|
||||
"""Test URL validation with invalid format."""
|
||||
is_valid, error_msg = await sse_client.validate_url("invalid-url", {})
|
||||
|
||||
assert is_valid is False
|
||||
assert "Invalid URL format" in error_msg
|
||||
|
||||
async def test_validate_url_with_404_response(self, sse_client):
|
||||
"""Test URL validation with 404 response (should be valid for SSE)."""
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 404
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
is_valid, error_msg = await sse_client.validate_url("http://test.url", {})
|
||||
|
||||
assert is_valid is True
|
||||
assert error_msg == ""
|
||||
|
||||
async def test_connect_to_server_with_headers(self, sse_client):
|
||||
"""Test connecting to server via SSE with custom headers."""
|
||||
test_url = "http://test.url"
|
||||
test_headers = {"Authorization": "Bearer token123", "Custom-Header": "value"}
|
||||
expected_headers = {"authorization": "Bearer token123", "custom-header": "value"} # normalized
|
||||
|
||||
with (
|
||||
patch.object(sse_client, "validate_url", return_value=(True, "")),
|
||||
patch.object(sse_client, "pre_check_redirect", return_value=test_url),
|
||||
patch.object(sse_client, "_get_or_create_session") as mock_get_session,
|
||||
):
|
||||
# Mock session
|
||||
mock_session = AsyncMock()
|
||||
mock_tool = MagicMock()
|
||||
mock_tool.name = "test_tool"
|
||||
list_tools_result = MagicMock()
|
||||
list_tools_result.tools = [mock_tool]
|
||||
mock_session.list_tools = AsyncMock(return_value=list_tools_result)
|
||||
mock_get_session.return_value = mock_session
|
||||
|
||||
tools = await sse_client.connect_to_server(test_url, test_headers)
|
||||
|
||||
assert len(tools) == 1
|
||||
assert tools[0].name == "test_tool"
|
||||
assert sse_client._connected is True
|
||||
|
||||
# Verify headers are stored in connection params (normalized)
|
||||
assert sse_client._connection_params is not None
|
||||
assert sse_client._connection_params["headers"] == expected_headers
|
||||
assert sse_client._connection_params["url"] == test_url
|
||||
|
||||
async def test_headers_passed_to_session_manager(self, sse_client):
|
||||
"""Test that headers are properly passed to the session manager."""
|
||||
test_url = "http://test.url"
|
||||
expected_headers = {"authorization": "Bearer token123", "x-api-key": "secret"} # normalized
|
||||
|
||||
sse_client._session_context = "test_context"
|
||||
sse_client._connection_params = {
|
||||
"url": test_url,
|
||||
"headers": expected_headers, # Use normalized headers
|
||||
"timeout_seconds": 30,
|
||||
"sse_read_timeout_seconds": 30,
|
||||
}
|
||||
|
||||
with patch.object(sse_client, "_get_session_manager") as mock_get_manager:
|
||||
mock_manager = AsyncMock()
|
||||
mock_session = AsyncMock()
|
||||
mock_manager.get_session = AsyncMock(return_value=mock_session)
|
||||
mock_get_manager.return_value = mock_manager
|
||||
|
||||
result_session = await sse_client._get_or_create_session()
|
||||
|
||||
# Verify session manager was called with correct parameters including normalized headers
|
||||
mock_manager.get_session.assert_called_once_with("test_context", sse_client._connection_params, "sse")
|
||||
assert result_session == mock_session
|
||||
|
||||
async def test_pre_check_redirect_with_headers(self, sse_client):
|
||||
"""Test pre-check redirect functionality with custom headers."""
|
||||
test_url = "http://test.url"
|
||||
redirect_url = "http://redirect.url"
|
||||
# Use pre-validated headers since pre_check_redirect expects already validated headers
|
||||
test_headers = {"authorization": "Bearer token123"} # already normalized
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 307
|
||||
mock_response.headers.get.return_value = redirect_url
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
result = await sse_client.pre_check_redirect(test_url, test_headers)
|
||||
|
||||
assert result == redirect_url
|
||||
# Verify validated headers were passed to the request
|
||||
mock_client.return_value.__aenter__.return_value.get.assert_called_with(
|
||||
test_url, timeout=2.0, headers={"Accept": "text/event-stream", **test_headers}
|
||||
)
|
||||
|
||||
async def test_run_tool_with_retry_on_connection_error(self, sse_client):
|
||||
"""Test that run_tool retries on connection errors."""
|
||||
# Setup connection state
|
||||
sse_client._connected = True
|
||||
sse_client._connection_params = {"url": "http://test.url", "headers": {}}
|
||||
sse_client._session_context = "test_context"
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def mock_get_session_side_effect():
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
session = AsyncMock()
|
||||
if call_count == 1:
|
||||
# First call fails with connection error
|
||||
from anyio import ClosedResourceError
|
||||
|
||||
session.call_tool = AsyncMock(side_effect=ClosedResourceError())
|
||||
else:
|
||||
# Second call succeeds
|
||||
mock_result = MagicMock()
|
||||
session.call_tool = AsyncMock(return_value=mock_result)
|
||||
return session
|
||||
|
||||
with (
|
||||
patch.object(sse_client, "_get_or_create_session", side_effect=mock_get_session_side_effect),
|
||||
patch.object(sse_client, "_get_session_manager") as mock_get_manager,
|
||||
):
|
||||
mock_manager = AsyncMock()
|
||||
mock_get_manager.return_value = mock_manager
|
||||
|
||||
result = await sse_client.run_tool("test_tool", {"param": "value"})
|
||||
|
||||
# Should have retried and succeeded on second attempt
|
||||
assert call_count == 2
|
||||
assert result is not None
|
||||
# Should have cleaned up the failed session
|
||||
mock_manager._cleanup_session.assert_called_once_with("test_context")
|
||||
|
|
@ -1,8 +1,15 @@
|
|||
"""Unit tests for MCP component with actual MCP servers.
|
||||
|
||||
This test suite validates the MCP component functionality using real MCP servers:
|
||||
- Everything server (stdio mode) - provides echo and other tools
|
||||
- DeepWiki server (SSE mode) - provides wiki-related tools
|
||||
"""
|
||||
|
||||
import shutil
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from langflow.base.mcp import util
|
||||
from langflow.base.mcp.util import MCPSessionManager, MCPSseClient, MCPStdioClient, _process_headers, validate_headers
|
||||
from langflow.base.mcp.util import MCPSessionManager, MCPSseClient, MCPStdioClient
|
||||
from langflow.components.agents.mcp_component import MCPToolsComponent
|
||||
|
||||
from tests.base import ComponentTestBaseWithoutClient, VersionComponentMapping
|
||||
|
|
@ -18,8 +25,11 @@ class TestMCPToolsComponent(ComponentTestBaseWithoutClient):
|
|||
def default_kwargs(self):
|
||||
"""Return the default kwargs for the component."""
|
||||
return {
|
||||
"mode": "Stdio",
|
||||
"command": "npx -y @modelcontextprotocol/server-everything",
|
||||
"sse_url": "https://mcp.deepwiki.com/sse",
|
||||
"tool": "echo",
|
||||
"mcp_server": {"name": "test_server", "config": {"command": "uvx mcp-server-fetch"}},
|
||||
"tool": "",
|
||||
}
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -27,34 +37,106 @@ class TestMCPToolsComponent(ComponentTestBaseWithoutClient):
|
|||
"""Return the file names mapping for different versions."""
|
||||
return []
|
||||
|
||||
@pytest.fixture
|
||||
def mock_tool(self):
|
||||
"""Create a mock MCP tool."""
|
||||
tool = MagicMock()
|
||||
tool.name = "test_tool"
|
||||
tool.description = "Test tool description"
|
||||
tool.inputSchema = {
|
||||
"type": "object",
|
||||
"properties": {"test_param": {"type": "string", "description": "Test parameter"}},
|
||||
"required": ["test_param"],
|
||||
}
|
||||
return tool
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_component_initialization(self, component_class, default_kwargs):
|
||||
"""Test that the component initializes correctly."""
|
||||
component = component_class(**default_kwargs)
|
||||
|
||||
# Check that the component has the expected attributes
|
||||
assert hasattr(component, "stdio_client")
|
||||
assert hasattr(component, "sse_client")
|
||||
assert isinstance(component.stdio_client, MCPStdioClient)
|
||||
assert isinstance(component.sse_client, MCPSseClient)
|
||||
|
||||
# Check that the component has a session manager
|
||||
session_manager = component.stdio_client._get_session_manager()
|
||||
assert isinstance(session_manager, MCPSessionManager)
|
||||
|
||||
|
||||
class TestMCPToolsComponentIntegration:
|
||||
"""Integration tests for the MCPToolsComponent."""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session(self, mock_tool):
|
||||
"""Create a mock ClientSession."""
|
||||
session = AsyncMock()
|
||||
session.initialize = AsyncMock()
|
||||
list_tools_result = MagicMock()
|
||||
list_tools_result.tools = [mock_tool]
|
||||
session.list_tools = AsyncMock(return_value=list_tools_result)
|
||||
session.call_tool = AsyncMock(
|
||||
return_value=MagicMock(content=[MagicMock(model_dump=lambda: {"result": "success"})])
|
||||
)
|
||||
return session
|
||||
def component(self):
|
||||
"""Create a component for testing."""
|
||||
return MCPToolsComponent()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not shutil.which("npx"), reason="Node.js not available")
|
||||
async def test_stdio_mode_integration(self, component):
|
||||
"""Test the component in stdio mode with Everything server."""
|
||||
# Configure for stdio mode
|
||||
component.mode = "Stdio"
|
||||
component.command = "npx -y @modelcontextprotocol/server-everything"
|
||||
component.tool = "echo"
|
||||
|
||||
try:
|
||||
# Mock the update_tool_list method to simulate server connection
|
||||
tools, server_info = await component.update_tool_list()
|
||||
|
||||
# Should have tools
|
||||
assert len(tools) > 0
|
||||
|
||||
# Should have server info
|
||||
assert server_info is not None
|
||||
assert isinstance(server_info, dict)
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible, skip the test
|
||||
pytest.skip(f"Everything server not accessible: {e}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_mode_integration(self, component):
|
||||
"""Test the component in SSE mode with DeepWiki server."""
|
||||
# Configure for SSE mode
|
||||
component.mode = "SSE"
|
||||
component.sse_url = "https://mcp.deepwiki.com/sse"
|
||||
|
||||
try:
|
||||
# Mock the update_tool_list method to simulate server connection
|
||||
tools, server_info = await component.update_tool_list()
|
||||
|
||||
# Should have tools
|
||||
assert len(tools) > 0
|
||||
|
||||
# Should have server info
|
||||
assert server_info is not None
|
||||
assert isinstance(server_info, dict)
|
||||
|
||||
except Exception as e:
|
||||
# If the server is not accessible, skip the test
|
||||
pytest.skip(f"DeepWiki server not accessible: {e}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_context_setting(self, component):
|
||||
"""Test that session context is properly set."""
|
||||
# Set session context
|
||||
component.stdio_client.set_session_context("test_context")
|
||||
component.sse_client.set_session_context("test_context")
|
||||
|
||||
# Verify context was set
|
||||
assert component.stdio_client._session_context == "test_context"
|
||||
assert component.sse_client._session_context == "test_context"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_manager_sharing(self, component):
|
||||
"""Test that session managers are shared through component cache."""
|
||||
# Get session managers
|
||||
stdio_manager = component.stdio_client._get_session_manager()
|
||||
sse_manager = component.sse_client._get_session_manager()
|
||||
|
||||
# Both should be MCPSessionManager instances
|
||||
assert isinstance(stdio_manager, MCPSessionManager)
|
||||
assert isinstance(sse_manager, MCPSessionManager)
|
||||
|
||||
# They should be the same instance (shared through cache)
|
||||
assert stdio_manager is sse_manager
|
||||
|
||||
|
||||
class TestMCPStdioClient:
|
||||
class TestMCPComponentErrorHandling:
|
||||
"""Test error handling in MCP components."""
|
||||
|
||||
@pytest.fixture
|
||||
def stdio_client(self):
|
||||
return MCPStdioClient()
|
||||
|
|
@ -122,494 +204,3 @@ class TestMCPStdioClient:
|
|||
mock_manager._cleanup_session.assert_called_once_with("test_context")
|
||||
assert stdio_client.session is None
|
||||
assert stdio_client._connected is False
|
||||
|
||||
|
||||
class TestMCPSseClient:
|
||||
@pytest.fixture
|
||||
def sse_client(self):
|
||||
return MCPSseClient()
|
||||
|
||||
async def test_validate_url_valid(self, sse_client):
|
||||
"""Test URL validation with valid URL."""
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
is_valid, error_msg = await sse_client.validate_url("http://test.url", {})
|
||||
|
||||
assert is_valid is True
|
||||
assert error_msg == ""
|
||||
|
||||
async def test_validate_url_invalid_format(self, sse_client):
|
||||
"""Test URL validation with invalid format."""
|
||||
is_valid, error_msg = await sse_client.validate_url("invalid-url", {})
|
||||
|
||||
assert is_valid is False
|
||||
assert "Invalid URL format" in error_msg
|
||||
|
||||
async def test_validate_url_with_404_response(self, sse_client):
|
||||
"""Test URL validation with 404 response (should be valid for SSE)."""
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 404
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
is_valid, error_msg = await sse_client.validate_url("http://test.url", {})
|
||||
|
||||
assert is_valid is True
|
||||
assert error_msg == ""
|
||||
|
||||
async def test_connect_to_server_with_headers(self, sse_client):
|
||||
"""Test connecting to server via SSE with custom headers."""
|
||||
test_url = "http://test.url"
|
||||
test_headers = {"Authorization": "Bearer token123", "Custom-Header": "value"}
|
||||
expected_headers = {"authorization": "Bearer token123", "custom-header": "value"} # normalized
|
||||
|
||||
with (
|
||||
patch.object(sse_client, "validate_url", return_value=(True, "")),
|
||||
patch.object(sse_client, "pre_check_redirect", return_value=test_url),
|
||||
patch.object(sse_client, "_get_or_create_session") as mock_get_session,
|
||||
):
|
||||
# Mock session
|
||||
mock_session = AsyncMock()
|
||||
mock_tool = MagicMock()
|
||||
mock_tool.name = "test_tool"
|
||||
list_tools_result = MagicMock()
|
||||
list_tools_result.tools = [mock_tool]
|
||||
mock_session.list_tools = AsyncMock(return_value=list_tools_result)
|
||||
mock_get_session.return_value = mock_session
|
||||
|
||||
tools = await sse_client.connect_to_server(test_url, test_headers)
|
||||
|
||||
assert len(tools) == 1
|
||||
assert tools[0].name == "test_tool"
|
||||
assert sse_client._connected is True
|
||||
|
||||
# Verify headers are stored in connection params (normalized)
|
||||
assert sse_client._connection_params is not None
|
||||
assert sse_client._connection_params["headers"] == expected_headers
|
||||
assert sse_client._connection_params["url"] == test_url
|
||||
|
||||
async def test_headers_passed_to_session_manager(self, sse_client):
|
||||
"""Test that headers are properly passed to the session manager."""
|
||||
test_url = "http://test.url"
|
||||
expected_headers = {"authorization": "Bearer token123", "x-api-key": "secret"} # normalized
|
||||
|
||||
sse_client._session_context = "test_context"
|
||||
sse_client._connection_params = {
|
||||
"url": test_url,
|
||||
"headers": expected_headers, # Use normalized headers
|
||||
"timeout_seconds": 30,
|
||||
"sse_read_timeout_seconds": 30,
|
||||
}
|
||||
|
||||
with patch.object(sse_client, "_get_session_manager") as mock_get_manager:
|
||||
mock_manager = AsyncMock()
|
||||
mock_session = AsyncMock()
|
||||
mock_manager.get_session = AsyncMock(return_value=mock_session)
|
||||
mock_get_manager.return_value = mock_manager
|
||||
|
||||
result_session = await sse_client._get_or_create_session()
|
||||
|
||||
# Verify session manager was called with correct parameters including normalized headers
|
||||
mock_manager.get_session.assert_called_once_with("test_context", sse_client._connection_params, "sse")
|
||||
assert result_session == mock_session
|
||||
|
||||
async def test_pre_check_redirect_with_headers(self, sse_client):
|
||||
"""Test pre-check redirect functionality with custom headers."""
|
||||
test_url = "http://test.url"
|
||||
redirect_url = "http://redirect.url"
|
||||
# Use pre-validated headers since pre_check_redirect expects already validated headers
|
||||
test_headers = {"authorization": "Bearer token123"} # already normalized
|
||||
|
||||
with patch("httpx.AsyncClient") as mock_client:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 307
|
||||
mock_response.headers.get.return_value = redirect_url
|
||||
mock_client.return_value.__aenter__.return_value.get.return_value = mock_response
|
||||
|
||||
result = await sse_client.pre_check_redirect(test_url, test_headers)
|
||||
|
||||
assert result == redirect_url
|
||||
# Verify validated headers were passed to the request
|
||||
mock_client.return_value.__aenter__.return_value.get.assert_called_with(
|
||||
test_url, timeout=2.0, headers={"Accept": "text/event-stream", **test_headers}
|
||||
)
|
||||
|
||||
async def test_run_tool_with_retry_on_connection_error(self, sse_client):
|
||||
"""Test that run_tool retries on connection errors."""
|
||||
# Setup connection state
|
||||
sse_client._connected = True
|
||||
sse_client._connection_params = {"url": "http://test.url", "headers": {}}
|
||||
sse_client._session_context = "test_context"
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def mock_get_session_side_effect():
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
session = AsyncMock()
|
||||
if call_count == 1:
|
||||
# First call fails with connection error
|
||||
from anyio import ClosedResourceError
|
||||
|
||||
session.call_tool = AsyncMock(side_effect=ClosedResourceError())
|
||||
else:
|
||||
# Second call succeeds
|
||||
mock_result = MagicMock()
|
||||
session.call_tool = AsyncMock(return_value=mock_result)
|
||||
return session
|
||||
|
||||
with (
|
||||
patch.object(sse_client, "_get_or_create_session", side_effect=mock_get_session_side_effect),
|
||||
patch.object(sse_client, "_get_session_manager") as mock_get_manager,
|
||||
):
|
||||
mock_manager = AsyncMock()
|
||||
mock_get_manager.return_value = mock_manager
|
||||
|
||||
result = await sse_client.run_tool("test_tool", {"param": "value"})
|
||||
|
||||
# Should have retried and succeeded on second attempt
|
||||
assert call_count == 2
|
||||
assert result is not None
|
||||
# Should have cleaned up the failed session
|
||||
mock_manager._cleanup_session.assert_called_once_with("test_context")
|
||||
|
||||
|
||||
class TestMCPSessionManager:
|
||||
@pytest.fixture
|
||||
def session_manager(self):
|
||||
return MCPSessionManager()
|
||||
|
||||
async def test_session_caching(self, session_manager):
|
||||
"""Test that sessions are properly cached and reused."""
|
||||
context_id = "test_context"
|
||||
connection_params = MagicMock()
|
||||
transport_type = "stdio"
|
||||
|
||||
# Create a mock session that will appear healthy
|
||||
mock_session = AsyncMock()
|
||||
mock_session._write_stream = MagicMock()
|
||||
mock_session._write_stream._closed = False
|
||||
|
||||
# Create a mock task that appears to be running
|
||||
mock_task = AsyncMock()
|
||||
mock_task.done = MagicMock(return_value=False)
|
||||
|
||||
with (
|
||||
patch.object(session_manager, "_create_stdio_session") as mock_create,
|
||||
patch.object(session_manager, "_validate_session_connectivity", return_value=True),
|
||||
):
|
||||
mock_create.return_value = mock_session
|
||||
|
||||
# First call should create session
|
||||
session1 = await session_manager.get_session(context_id, connection_params, transport_type)
|
||||
|
||||
# Manually populate the sessions cache as if the session was created properly
|
||||
session_manager.sessions[context_id] = {"session": mock_session, "task": mock_task, "type": transport_type}
|
||||
|
||||
# Second call should return cached session without creating new one
|
||||
session2 = await session_manager.get_session(context_id, connection_params, transport_type)
|
||||
|
||||
assert session1 == session2
|
||||
assert session1 == mock_session
|
||||
# Should only create once since the second call should use the cached session
|
||||
mock_create.assert_called_once()
|
||||
|
||||
async def test_session_cleanup(self, session_manager):
|
||||
"""Test session cleanup functionality."""
|
||||
context_id = "test_context"
|
||||
|
||||
# Add a session to the manager with proper mock setup
|
||||
mock_task = AsyncMock()
|
||||
mock_task.done = MagicMock(return_value=False) # Use MagicMock for sync method
|
||||
mock_task.cancel = MagicMock() # Use MagicMock for sync method
|
||||
|
||||
session_manager.sessions[context_id] = {"session": AsyncMock(), "task": mock_task, "type": "stdio"}
|
||||
|
||||
await session_manager._cleanup_session(context_id)
|
||||
|
||||
# Should cancel the task and remove from sessions
|
||||
mock_task.cancel.assert_called_once()
|
||||
assert context_id not in session_manager.sessions
|
||||
|
||||
async def test_server_switch_detection(self, session_manager):
|
||||
"""Test that server switches are properly detected and handled."""
|
||||
context_id = "test_context"
|
||||
|
||||
# First server
|
||||
server1_params = MagicMock()
|
||||
server1_params.command = "server1"
|
||||
|
||||
# Second server
|
||||
server2_params = MagicMock()
|
||||
server2_params.command = "server2"
|
||||
|
||||
with (
|
||||
patch.object(session_manager, "_create_stdio_session") as mock_create,
|
||||
patch.object(session_manager, "_validate_session_connectivity", return_value=True),
|
||||
):
|
||||
mock_session1 = AsyncMock()
|
||||
mock_session2 = AsyncMock()
|
||||
mock_create.side_effect = [mock_session1, mock_session2]
|
||||
|
||||
# First connection
|
||||
session1 = await session_manager.get_session(context_id, server1_params, "stdio")
|
||||
|
||||
# Switch to different server should create new session
|
||||
session2 = await session_manager.get_session(context_id, server2_params, "stdio")
|
||||
|
||||
assert session1 != session2
|
||||
assert mock_create.call_count == 2
|
||||
|
||||
|
||||
# Integration test for header functionality
|
||||
class TestHeaderValidation:
|
||||
"""Test the header validation functionality."""
|
||||
|
||||
def test_validate_headers_valid_input(self):
|
||||
"""Test header validation with valid headers."""
|
||||
headers = {"Authorization": "Bearer token123", "Content-Type": "application/json", "X-API-Key": "secret-key"}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Headers should be normalized to lowercase
|
||||
expected = {"authorization": "Bearer token123", "content-type": "application/json", "x-api-key": "secret-key"}
|
||||
assert result == expected
|
||||
|
||||
def test_validate_headers_empty_input(self):
|
||||
"""Test header validation with empty/None input."""
|
||||
assert validate_headers({}) == {}
|
||||
assert validate_headers(None) == {}
|
||||
|
||||
def test_validate_headers_invalid_names(self):
|
||||
"""Test header validation with invalid header names."""
|
||||
headers = {
|
||||
"Invalid Header": "value", # spaces not allowed
|
||||
"Header@Name": "value", # @ not allowed
|
||||
"Header Name": "value", # spaces not allowed
|
||||
"Valid-Header": "value", # this should pass
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Only the valid header should remain
|
||||
assert result == {"valid-header": "value"}
|
||||
|
||||
def test_validate_headers_sanitize_values(self):
|
||||
"""Test header value sanitization."""
|
||||
headers = {
|
||||
"Authorization": "Bearer \x00token\x1f with\r\ninjection",
|
||||
"Clean-Header": " clean value ",
|
||||
"Empty-After-Clean": "\x00\x01\x02",
|
||||
"Tab-Header": "value\twith\ttabs", # tabs should be preserved
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Control characters should be removed, whitespace trimmed
|
||||
# Header with injection attempts should be skipped
|
||||
expected = {"clean-header": "clean value", "tab-header": "value\twith\ttabs"}
|
||||
assert result == expected
|
||||
|
||||
def test_validate_headers_non_string_values(self):
|
||||
"""Test header validation with non-string values."""
|
||||
headers = {"String-Header": "valid", "Number-Header": 123, "None-Header": None, "List-Header": ["value"]}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Only string headers should remain
|
||||
assert result == {"string-header": "valid"}
|
||||
|
||||
def test_validate_headers_injection_attempts(self):
|
||||
"""Test header validation against injection attempts."""
|
||||
headers = {
|
||||
"Injection1": "value\r\nInjected-Header: malicious",
|
||||
"Injection2": "value\nX-Evil: attack",
|
||||
"Safe-Header": "safe-value",
|
||||
}
|
||||
|
||||
result = validate_headers(headers)
|
||||
|
||||
# Injection attempts should be filtered out
|
||||
assert result == {"safe-header": "safe-value"}
|
||||
|
||||
|
||||
class TestSSEHeaderIntegration:
|
||||
"""Integration test to verify headers are properly passed through the entire SSE flow."""
|
||||
|
||||
async def test_headers_processing(self):
|
||||
"""Test that headers flow properly from server config through to SSE client connection."""
|
||||
# Test the header processing function directly
|
||||
headers_input = [
|
||||
{"key": "Authorization", "value": "Bearer test-token"},
|
||||
{"key": "X-API-Key", "value": "secret-key"},
|
||||
]
|
||||
|
||||
expected_headers = {
|
||||
"authorization": "Bearer test-token", # normalized to lowercase
|
||||
"x-api-key": "secret-key",
|
||||
}
|
||||
|
||||
# Test _process_headers function with validation
|
||||
processed_headers = _process_headers(headers_input)
|
||||
assert processed_headers == expected_headers
|
||||
|
||||
# Test different input formats
|
||||
# Test dict input with validation
|
||||
dict_headers = {"Authorization": "Bearer dict-token", "Invalid Header": "bad"}
|
||||
result = _process_headers(dict_headers)
|
||||
# Invalid header should be filtered out, valid header normalized
|
||||
assert result == {"authorization": "Bearer dict-token"}
|
||||
|
||||
# Test None input
|
||||
assert _process_headers(None) == {}
|
||||
|
||||
# Test empty list
|
||||
assert _process_headers([]) == {}
|
||||
|
||||
# Test malformed list
|
||||
malformed_headers = [{"key": "Auth"}, {"value": "token"}] # Missing value/key
|
||||
assert _process_headers(malformed_headers) == {}
|
||||
|
||||
# Test list with invalid header names
|
||||
invalid_headers = [
|
||||
{"key": "Valid-Header", "value": "good"},
|
||||
{"key": "Invalid Header", "value": "bad"}, # spaces not allowed
|
||||
]
|
||||
result = _process_headers(invalid_headers)
|
||||
assert result == {"valid-header": "good"}
|
||||
|
||||
async def test_sse_client_header_storage(self):
|
||||
"""Test that SSE client properly stores headers in connection params."""
|
||||
sse_client = MCPSseClient()
|
||||
test_url = "http://test.url"
|
||||
test_headers = {"Authorization": "Bearer test123", "Custom": "value"}
|
||||
expected_headers = {"authorization": "Bearer test123", "custom": "value"} # normalized
|
||||
|
||||
with (
|
||||
patch.object(sse_client, "validate_url", return_value=(True, "")),
|
||||
patch.object(sse_client, "pre_check_redirect", return_value=test_url),
|
||||
patch.object(sse_client, "_get_or_create_session") as mock_get_session,
|
||||
):
|
||||
mock_session = AsyncMock()
|
||||
mock_tool = MagicMock()
|
||||
mock_tool.name = "test_tool"
|
||||
list_tools_result = MagicMock()
|
||||
list_tools_result.tools = [mock_tool]
|
||||
mock_session.list_tools = AsyncMock(return_value=list_tools_result)
|
||||
mock_get_session.return_value = mock_session
|
||||
|
||||
await sse_client.connect_to_server(test_url, test_headers)
|
||||
|
||||
# Verify headers are stored correctly in connection params (normalized)
|
||||
assert sse_client._connection_params is not None
|
||||
assert sse_client._connection_params["headers"] == expected_headers
|
||||
assert sse_client._connection_params["url"] == test_url
|
||||
|
||||
|
||||
class TestMCPUtilityFunctions:
|
||||
"""Test utility functions from util.py that don't have dedicated test classes."""
|
||||
|
||||
def test_sanitize_mcp_name(self):
|
||||
"""Test MCP name sanitization."""
|
||||
assert util.sanitize_mcp_name("Test Name 123") == "test_name_123"
|
||||
assert util.sanitize_mcp_name(" ") == ""
|
||||
assert util.sanitize_mcp_name("123abc") == "_123abc"
|
||||
assert util.sanitize_mcp_name("Tést-😀-Námé") == "test_name"
|
||||
assert util.sanitize_mcp_name("a" * 100) == "a" * 46
|
||||
|
||||
def test_get_unique_name(self):
|
||||
"""Test unique name generation."""
|
||||
names = {"foo", "foo_1"}
|
||||
assert util.get_unique_name("foo", 10, names) == "foo_2"
|
||||
assert util.get_unique_name("bar", 10, names) == "bar"
|
||||
assert util.get_unique_name("longname", 4, {"long"}) == "lo_1"
|
||||
|
||||
def test_is_valid_key_value_item(self):
|
||||
"""Test key-value item validation."""
|
||||
assert util._is_valid_key_value_item({"key": "a", "value": "b"}) is True
|
||||
assert util._is_valid_key_value_item({"key": "a"}) is False
|
||||
assert util._is_valid_key_value_item(["key", "value"]) is False
|
||||
assert util._is_valid_key_value_item(None) is False
|
||||
|
||||
def test_validate_node_installation(self):
|
||||
"""Test Node.js installation validation."""
|
||||
import shutil
|
||||
|
||||
if shutil.which("node"):
|
||||
assert util._validate_node_installation("npx something") == "npx something"
|
||||
else:
|
||||
with pytest.raises(ValueError, match="Node.js is not installed"):
|
||||
util._validate_node_installation("npx something")
|
||||
assert util._validate_node_installation("echo test") == "echo test"
|
||||
|
||||
def test_create_input_schema_from_json_schema(self):
|
||||
"""Test JSON schema to Pydantic model conversion."""
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"foo": {"type": "string", "description": "desc"},
|
||||
"bar": {"type": "integer"},
|
||||
},
|
||||
"required": ["foo"],
|
||||
}
|
||||
model_class = util.create_input_schema_from_json_schema(schema)
|
||||
instance = model_class(foo="abc", bar=1)
|
||||
assert instance.foo == "abc"
|
||||
assert instance.bar == 1
|
||||
|
||||
with pytest.raises(Exception): # noqa: B017, PT011
|
||||
model_class(bar=1) # missing required field
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_connection_params(self):
|
||||
"""Test connection parameter validation."""
|
||||
# Valid parameters
|
||||
await util._validate_connection_params("Stdio", command="echo test")
|
||||
await util._validate_connection_params("SSE", url="http://test")
|
||||
|
||||
# Invalid parameters
|
||||
with pytest.raises(ValueError, match="Command is required for Stdio mode"):
|
||||
await util._validate_connection_params("Stdio", command=None)
|
||||
with pytest.raises(ValueError, match="URL is required for SSE mode"):
|
||||
await util._validate_connection_params("SSE", url=None)
|
||||
with pytest.raises(ValueError, match="Invalid mode"):
|
||||
await util._validate_connection_params("InvalidMode")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_flow_snake_case_mocked(self):
|
||||
"""Test flow lookup by snake case name with mocked session."""
|
||||
|
||||
class DummyFlow:
|
||||
def __init__(self, name: str, user_id: str, *, is_component: bool = False, action_name: str | None = None):
|
||||
self.name = name
|
||||
self.user_id = user_id
|
||||
self.is_component = is_component
|
||||
self.action_name = action_name
|
||||
|
||||
class DummyExec:
|
||||
def __init__(self, flows: list[DummyFlow]):
|
||||
self._flows = flows
|
||||
|
||||
def all(self):
|
||||
return self._flows
|
||||
|
||||
class DummySession:
|
||||
def __init__(self, flows: list[DummyFlow]):
|
||||
self._flows = flows
|
||||
|
||||
async def exec(self, stmt): # noqa: ARG002
|
||||
return DummyExec(self._flows)
|
||||
|
||||
user_id = "123e4567-e89b-12d3-a456-426614174000"
|
||||
flows = [DummyFlow("Test Flow", user_id), DummyFlow("Other", user_id)]
|
||||
|
||||
# Should match sanitized name
|
||||
result = await util.get_flow_snake_case(util.sanitize_mcp_name("Test Flow"), user_id, DummySession(flows))
|
||||
assert result is flows[0]
|
||||
|
||||
# Should return None if not found
|
||||
result = await util.get_flow_snake_case("notfound", user_id, DummySession(flows))
|
||||
assert result is None
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue