diff --git a/hindsight-api/hindsight_api/api/mcp.py b/hindsight-api/hindsight_api/api/mcp.py index 1821dd00..c2d0a71a 100644 --- a/hindsight-api/hindsight_api/api/mcp.py +++ b/hindsight-api/hindsight_api/api/mcp.py @@ -29,15 +29,26 @@ logger = logging.getLogger(__name__) # Default bank_id from environment variable DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default") +# MCP authentication token (optional - if set, Bearer token auth is required) +MCP_AUTH_TOKEN = os.environ.get("HINDSIGHT_API_MCP_AUTH_TOKEN") + # Context variable to hold the current bank_id _current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None) +# Context variable to hold the current API key (for tenant auth propagation) +_current_api_key: ContextVar[str | None] = ContextVar("current_api_key", default=None) + def get_current_bank_id() -> str | None: """Get the current bank_id from context.""" return _current_bank_id.get() +def get_current_api_key() -> str | None: + """Get the current API key from context.""" + return _current_api_key.get() + + def create_mcp_server(memory: MemoryEngine) -> FastMCP: """ Create and configure the Hindsight MCP server. @@ -54,6 +65,7 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP: # Configure and register tools using shared module config = MCPToolsConfig( bank_id_resolver=get_current_bank_id, + api_key_resolver=get_current_api_key, # Propagate API key for tenant auth include_bank_id_param=True, # HTTP MCP supports multi-bank via parameter tools=None, # All tools retain_fire_and_forget=False, # HTTP MCP supports sync/async modes @@ -65,7 +77,11 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP: class MCPMiddleware: - """ASGI middleware that extracts bank_id from header or path and sets context. + """ASGI middleware that handles authentication and extracts bank_id from header or path. + + Authentication: + If HINDSIGHT_API_MCP_AUTH_TOKEN is set, all requests must include a valid + Authorization header with Bearer token or direct token matching the configured value. Bank ID can be provided via: 1. X-Bank-Id header (recommended for Claude Code) @@ -74,7 +90,7 @@ class MCPMiddleware: For Claude Code, configure with: claude mcp add --transport http hindsight http://localhost:8888/mcp \\ - --header "X-Bank-Id: my-bank" + --header "X-Bank-Id: my-bank" --header "Authorization: Bearer " """ def __init__(self, app, memory: MemoryEngine): @@ -98,6 +114,22 @@ class MCPMiddleware: await self.mcp_app(scope, receive, send) return + # Extract auth token from header (for tenant auth propagation) + auth_header = self._get_header(scope, "Authorization") + auth_token: str | None = None + if auth_header: + # Support both "Bearer " and direct token + auth_token = auth_header[7:].strip() if auth_header.startswith("Bearer ") else auth_header.strip() + + # Authenticate if MCP_AUTH_TOKEN is configured + if MCP_AUTH_TOKEN: + if not auth_token: + await self._send_error(send, 401, "Authorization header required") + return + if auth_token != MCP_AUTH_TOKEN: + await self._send_error(send, 401, "Invalid authentication token") + return + path = scope.get("path", "") # Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped @@ -132,8 +164,10 @@ class MCPMiddleware: bank_id = DEFAULT_BANK_ID logger.debug(f"Using default bank_id: {bank_id}") - # Set bank_id context - token = _current_bank_id.set(bank_id) + # Set bank_id and api_key context + bank_id_token = _current_bank_id.set(bank_id) + # Store the auth token for tenant extension to validate + api_key_token = _current_api_key.set(auth_token) if auth_token else None try: new_scope = scope.copy() new_scope["path"] = new_path @@ -152,7 +186,9 @@ class MCPMiddleware: await self.mcp_app(new_scope, receive, send_wrapper) finally: - _current_bank_id.reset(token) + _current_bank_id.reset(bank_id_token) + if api_key_token is not None: + _current_api_key.reset(api_key_token) async def _send_error(self, send, status: int, message: str): """Send an error response.""" @@ -176,6 +212,10 @@ def create_mcp_app(memory: MemoryEngine): """ Create an ASGI app that handles MCP requests. + Authentication: + Set HINDSIGHT_API_MCP_AUTH_TOKEN to require Bearer token authentication. + If not set, MCP endpoint is open (for local development). + Bank ID can be provided via: 1. X-Bank-Id header: claude mcp add --transport http hindsight http://localhost:8888/mcp --header "X-Bank-Id: my-bank" 2. URL path: /mcp/{bank_id}/ diff --git a/hindsight-api/hindsight_api/mcp_tools.py b/hindsight-api/hindsight_api/mcp_tools.py index 0cd4a31b..db4dfdd9 100644 --- a/hindsight-api/hindsight_api/mcp_tools.py +++ b/hindsight-api/hindsight_api/mcp_tools.py @@ -32,6 +32,9 @@ class MCPToolsConfig: # How to resolve bank_id for operations bank_id_resolver: Callable[[], str | None] + # How to resolve API key for tenant auth (optional) + api_key_resolver: Callable[[], str | None] | None = None + # Whether to include bank_id as a parameter on tools (for multi-bank support) include_bank_id_param: bool = False @@ -46,6 +49,16 @@ class MCPToolsConfig: retain_fire_and_forget: bool = False # If True, use asyncio.create_task pattern +def _get_request_context(config: MCPToolsConfig) -> RequestContext: + """Create RequestContext with API key from resolver if available. + + This enables tenant auth to work with MCP tools by propagating + the Bearer token from the MCP middleware to the memory engine. + """ + api_key = config.api_key_resolver() if config.api_key_resolver else None + return RequestContext(api_key=api_key) + + def parse_timestamp(timestamp: str) -> datetime | None: """Parse an ISO format timestamp string. @@ -155,12 +168,14 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) if error: return {"status": "error", "message": error} + request_context = _get_request_context(config) + async def _retain(): try: await memory.retain_batch_async( bank_id=target_bank, contents=[content_dict], - request_context=RequestContext(), + request_context=request_context, ) except Exception as e: logger.error(f"Error storing memory: {e}", exc_info=True) @@ -196,16 +211,17 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) return f"Error: {error}" contents = [content_dict] + request_context = _get_request_context(config) if async_processing: result = await memory.submit_async_retain( - bank_id=target_bank, contents=contents, request_context=RequestContext() + bank_id=target_bank, contents=contents, request_context=request_context ) return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})" else: await memory.retain_batch_async( bank_id=target_bank, contents=contents, - request_context=RequestContext(), + request_context=request_context, ) return f"Memory stored successfully in bank '{target_bank}'" except Exception as e: @@ -237,12 +253,14 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) if error: return {"status": "error", "message": error} + request_context = _get_request_context(config) + async def _retain(): try: await memory.retain_batch_async( bank_id=target_bank, contents=[content_dict], - request_context=RequestContext(), + request_context=request_context, ) except Exception as e: logger.error(f"Error storing memory: {e}", exc_info=True) @@ -280,7 +298,7 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) fact_type=list(VALID_RECALL_FACT_TYPES), budget=Budget.HIGH, max_tokens=max_tokens, - request_context=RequestContext(), + request_context=_get_request_context(config), ) return recall_result.model_dump_json(indent=2) @@ -311,7 +329,7 @@ def _register_recall(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) fact_type=list(VALID_RECALL_FACT_TYPES), budget=Budget.HIGH, max_tokens=max_tokens, - request_context=RequestContext(), + request_context=_get_request_context(config), ) return recall_result.model_dump() @@ -370,7 +388,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig query=query, budget=budget_enum, context=context, - request_context=RequestContext(), + request_context=_get_request_context(config), ) return reflect_result.model_dump_json(indent=2) @@ -423,7 +441,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig query=query, budget=budget_enum, context=context, - request_context=RequestContext(), + request_context=_get_request_context(config), ) return reflect_result.model_dump() @@ -447,7 +465,7 @@ def _register_list_banks(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCon JSON list of banks with their IDs, names, dispositions, and missions. """ try: - banks = await memory.list_banks(request_context=RequestContext()) + banks = await memory.list_banks(request_context=_get_request_context(config)) return json.dumps({"banks": banks}, indent=2) except Exception as e: logger.error(f"Error listing banks: {e}", exc_info=True) @@ -471,8 +489,9 @@ def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo mission: Optional mission describing who the agent is and what they're trying to accomplish """ try: + request_context = _get_request_context(config) # get_bank_profile auto-creates bank if it doesn't exist - profile = await memory.get_bank_profile(bank_id, request_context=RequestContext()) + profile = await memory.get_bank_profile(bank_id, request_context=request_context) # Update name/mission if provided if name is not None or mission is not None: @@ -480,10 +499,10 @@ def _register_create_bank(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsCo bank_id, name=name, mission=mission, - request_context=RequestContext(), + request_context=request_context, ) # Fetch updated profile - profile = await memory.get_bank_profile(bank_id, request_context=RequestContext()) + profile = await memory.get_bank_profile(bank_id, request_context=request_context) # Serialize disposition if it's a Pydantic model if "disposition" in profile and hasattr(profile["disposition"], "model_dump"): diff --git a/hindsight-api/tests/test_mcp_routing.py b/hindsight-api/tests/test_mcp_routing.py index d54e9f4d..fb5676cd 100644 --- a/hindsight-api/tests/test_mcp_routing.py +++ b/hindsight-api/tests/test_mcp_routing.py @@ -97,3 +97,47 @@ def test_path_parsing_logic(): 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)