* feat: improve mcp tools based on endpoint
* feat: improve mcp tools based on endpoint
* test: add integration test for MCP endpoint routing
- Add test_mcp_endpoint_routing.py to verify single-bank vs multi-bank tool exposure
- Verifies /mcp/ exposes all tools with bank_id parameters
- Verifies /mcp/{bank_id}/ only exposes scoped tools without bank_id parameters
- Regression test for issue #317
Related: #317, #318
* test: use StreamableHTTP client for MCP endpoint routing test
Replace httpx AsyncClient SSE parsing with proper MCP StreamableHTTP
client. This correctly tests the MCP server using the actual protocol
that clients will use.
Fixes #317
292 lines
10 KiB
Python
292 lines
10 KiB
Python
"""Test MCP server routing with dynamic bank_id."""
|
|
|
|
import pytest
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_memory():
|
|
"""Create a mock MemoryEngine."""
|
|
memory = MagicMock()
|
|
memory.retain_batch_async = AsyncMock()
|
|
memory.submit_async_retain = AsyncMock(return_value={"operation_id": "test-op-123"})
|
|
memory.recall_async = AsyncMock(return_value=MagicMock(results=[]))
|
|
return memory
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mcp_context_variable():
|
|
"""Test that context variable works correctly."""
|
|
from hindsight_api.api.mcp import get_current_bank_id, _current_bank_id
|
|
|
|
# Initially None
|
|
assert get_current_bank_id() is None
|
|
|
|
# Set and verify
|
|
token = _current_bank_id.set("test-bank-123")
|
|
try:
|
|
assert get_current_bank_id() == "test-bank-123"
|
|
finally:
|
|
_current_bank_id.reset(token)
|
|
|
|
# Back to None after reset
|
|
assert get_current_bank_id() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mcp_tools_use_context_bank_id(mock_memory):
|
|
"""Test that MCP tools use bank_id from context."""
|
|
from hindsight_api.api.mcp import create_mcp_server, _current_bank_id
|
|
|
|
mcp_server = create_mcp_server(mock_memory)
|
|
|
|
# Get the tools
|
|
tools = mcp_server._tool_manager._tools
|
|
assert "retain" in tools
|
|
assert "recall" in tools
|
|
|
|
# Test retain with bank_id from context (use async_processing=False for synchronous test)
|
|
token = _current_bank_id.set("context-bank-id")
|
|
try:
|
|
retain_tool = tools["retain"]
|
|
result = await retain_tool.fn(content="test content", context="test_context", async_processing=False)
|
|
assert "successfully" in result.lower()
|
|
|
|
# Verify the memory was called with the context bank_id
|
|
mock_memory.retain_batch_async.assert_called_once()
|
|
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
|
|
assert call_kwargs["bank_id"] == "context-bank-id"
|
|
finally:
|
|
_current_bank_id.reset(token)
|
|
|
|
|
|
def test_path_parsing_logic():
|
|
"""Test the path parsing logic for bank_id extraction."""
|
|
def parse_path(path):
|
|
"""Simulate the path parsing logic from MCPMiddleware."""
|
|
if not path.startswith("/") or len(path) <= 1:
|
|
return None, None # Error case
|
|
|
|
parts = path[1:].split("/", 1)
|
|
if not parts[0]:
|
|
return None, None # Error case
|
|
|
|
bank_id = parts[0]
|
|
new_path = "/" + parts[1] if len(parts) > 1 else "/"
|
|
return bank_id, new_path
|
|
|
|
# Test bank-specific paths
|
|
bank_id, remaining = parse_path("/my-bank/")
|
|
assert bank_id == "my-bank"
|
|
assert remaining == "/"
|
|
|
|
bank_id, remaining = parse_path("/my-bank")
|
|
assert bank_id == "my-bank"
|
|
assert remaining == "/"
|
|
|
|
# Test error case - no bank_id
|
|
bank_id, remaining = parse_path("/")
|
|
assert bank_id is None
|
|
|
|
# Test with complex bank_id
|
|
bank_id, remaining = parse_path("/user_12345/")
|
|
assert bank_id == "user_12345"
|
|
assert remaining == "/"
|
|
|
|
# Test with additional path after bank_id
|
|
bank_id, remaining = parse_path("/my-bank/some/path")
|
|
assert bank_id == "my-bank"
|
|
assert remaining == "/some/path"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_api_key_context_variable():
|
|
"""Test that API key context variable works correctly."""
|
|
from hindsight_api.api.mcp import get_current_api_key, _current_api_key
|
|
|
|
# Initially None
|
|
assert get_current_api_key() is None
|
|
|
|
# Set and verify
|
|
token = _current_api_key.set("test-api-key-123")
|
|
try:
|
|
assert get_current_api_key() == "test-api-key-123"
|
|
finally:
|
|
_current_api_key.reset(token)
|
|
|
|
# Back to None after reset
|
|
assert get_current_api_key() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mcp_tools_propagate_api_key(mock_memory):
|
|
"""Test that MCP tools propagate API key to RequestContext."""
|
|
from hindsight_api.api.mcp import create_mcp_server, _current_bank_id, _current_api_key
|
|
|
|
mcp_server = create_mcp_server(mock_memory)
|
|
tools = mcp_server._tool_manager._tools
|
|
|
|
# Set both bank_id and api_key context
|
|
bank_token = _current_bank_id.set("test-bank")
|
|
api_key_token = _current_api_key.set("test-bearer-token")
|
|
try:
|
|
retain_tool = tools["retain"]
|
|
result = await retain_tool.fn(content="test content", context="test_context", async_processing=False)
|
|
assert "successfully" in result.lower()
|
|
|
|
# Verify the memory was called with request_context containing api_key
|
|
mock_memory.retain_batch_async.assert_called_once()
|
|
call_kwargs = mock_memory.retain_batch_async.call_args.kwargs
|
|
assert call_kwargs["request_context"].api_key == "test-bearer-token"
|
|
finally:
|
|
_current_bank_id.reset(bank_token)
|
|
_current_api_key.reset(api_key_token)
|
|
|
|
|
|
def test_multi_bank_mode_exposes_all_tools(mock_memory):
|
|
"""Test that multi-bank mode exposes all tools including bank management."""
|
|
from hindsight_api.api.mcp import create_mcp_server
|
|
|
|
# Create server in multi-bank mode (default)
|
|
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
|
|
tools = mcp_server._tool_manager._tools
|
|
|
|
# Should have all tools
|
|
assert "retain" in tools
|
|
assert "recall" in tools
|
|
assert "reflect" in tools
|
|
assert "list_banks" in tools
|
|
assert "create_bank" in tools
|
|
|
|
|
|
def test_single_bank_mode_excludes_bank_management_tools(mock_memory):
|
|
"""Test that single-bank mode only exposes bank-scoped tools."""
|
|
from hindsight_api.api.mcp import create_mcp_server
|
|
|
|
# Create server in single-bank mode
|
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
|
tools = mcp_server._tool_manager._tools
|
|
|
|
# Should only have bank-scoped tools
|
|
assert "retain" in tools
|
|
assert "recall" in tools
|
|
assert "reflect" in tools
|
|
|
|
# Should NOT have bank management tools
|
|
assert "list_banks" not in tools
|
|
assert "create_bank" not in tools
|
|
|
|
|
|
def test_multi_bank_mode_tools_have_bank_id_param(mock_memory):
|
|
"""Test that multi-bank mode tools include bank_id parameter."""
|
|
from hindsight_api.api.mcp import create_mcp_server
|
|
import inspect
|
|
|
|
mcp_server = create_mcp_server(mock_memory, multi_bank=True)
|
|
tools = mcp_server._tool_manager._tools
|
|
|
|
# Check that tools have bank_id parameter
|
|
retain_tool = tools["retain"]
|
|
retain_sig = inspect.signature(retain_tool.fn)
|
|
assert "bank_id" in retain_sig.parameters
|
|
|
|
recall_tool = tools["recall"]
|
|
recall_sig = inspect.signature(recall_tool.fn)
|
|
assert "bank_id" in recall_sig.parameters
|
|
|
|
reflect_tool = tools["reflect"]
|
|
reflect_sig = inspect.signature(reflect_tool.fn)
|
|
assert "bank_id" in reflect_sig.parameters
|
|
|
|
|
|
def test_single_bank_mode_tools_no_bank_id_param(mock_memory):
|
|
"""Test that single-bank mode tools do NOT include bank_id parameter."""
|
|
from hindsight_api.api.mcp import create_mcp_server
|
|
import inspect
|
|
|
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
|
tools = mcp_server._tool_manager._tools
|
|
|
|
# Check that tools do NOT have bank_id parameter
|
|
retain_tool = tools["retain"]
|
|
retain_sig = inspect.signature(retain_tool.fn)
|
|
assert "bank_id" not in retain_sig.parameters
|
|
|
|
recall_tool = tools["recall"]
|
|
recall_sig = inspect.signature(recall_tool.fn)
|
|
assert "bank_id" not in recall_sig.parameters
|
|
|
|
reflect_tool = tools["reflect"]
|
|
reflect_sig = inspect.signature(reflect_tool.fn)
|
|
assert "bank_id" not in reflect_sig.parameters
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_middleware_handles_both_endpoints(mock_memory):
|
|
"""Test that MCPMiddleware routes to correct server based on URL path."""
|
|
from hindsight_api.api.mcp import MCPMiddleware
|
|
|
|
# Create middleware (single instance)
|
|
middleware = MCPMiddleware(None, mock_memory)
|
|
|
|
# Verify both server instances exist
|
|
assert middleware.multi_bank_app is not None
|
|
assert middleware.single_bank_app is not None
|
|
|
|
# Verify they expose different tools
|
|
multi_bank_tools = middleware.multi_bank_server._tool_manager._tools
|
|
single_bank_tools = middleware.single_bank_server._tool_manager._tools
|
|
|
|
# Multi-bank should have all tools
|
|
assert "retain" in multi_bank_tools
|
|
assert "recall" in multi_bank_tools
|
|
assert "list_banks" in multi_bank_tools
|
|
assert "create_bank" in multi_bank_tools
|
|
|
|
# Single-bank should only have scoped tools
|
|
assert "retain" in single_bank_tools
|
|
assert "recall" in single_bank_tools
|
|
assert "list_banks" not in single_bank_tools
|
|
assert "create_bank" not in single_bank_tools
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_routing_logic_from_url_path():
|
|
"""Test that routing correctly selects server based on URL structure."""
|
|
from hindsight_api.api.mcp import MCPMiddleware
|
|
from unittest.mock import AsyncMock
|
|
|
|
# Mock memory
|
|
mock_memory = MagicMock()
|
|
|
|
# Create middleware
|
|
middleware = MCPMiddleware(None, mock_memory)
|
|
|
|
# Simulate different URL patterns and verify routing
|
|
test_cases = [
|
|
# (path_after_stripping_mcp, expected_bank_id_from_path, expected_bank_id, description)
|
|
("/alice/messages", True, "alice", "Bank ID in path with endpoint"),
|
|
("/my-agent-123/", True, "my-agent-123", "Bank ID in path with trailing slash"),
|
|
("ciccio/messages", True, "ciccio", "Bank ID without leading slash (after mount strip)"),
|
|
("bob", True, "bob", "Bank ID only, no leading slash"),
|
|
("/messages", False, None, "MCP endpoint, no bank ID"),
|
|
("/", False, None, "Root path, no bank ID"),
|
|
]
|
|
|
|
for path, expected_bank_from_path, expected_bank_id, description in test_cases:
|
|
# Simulate the path parsing logic with leading slash normalization
|
|
if path and not path.startswith("/"):
|
|
path = "/" + path
|
|
|
|
bank_id = None
|
|
bank_id_from_path = False
|
|
MCP_ENDPOINTS = {"sse", "messages"}
|
|
|
|
if path.startswith("/") and len(path) > 1:
|
|
parts = path[1:].split("/", 1)
|
|
if parts[0] and parts[0] not in MCP_ENDPOINTS:
|
|
bank_id = parts[0]
|
|
bank_id_from_path = True
|
|
|
|
assert bank_id_from_path == expected_bank_from_path, f"Failed for: {description} (path={path})"
|
|
assert bank_id == expected_bank_id, f"Failed bank_id for: {description} (path={path}, got={bank_id})"
|