fleet-memory/hindsight-api/hindsight_api/api/mcp.py
DK09876 e798979733
Harden MCP server: fix routing, validation, and usage metering (#341)
* fix: move mental model usage metering into engine for MCP support

Mental model validation hooks (validate_mental_model_get, validate_mental_model_refresh)
were only called in REST HTTP handlers, not in the engine. MCP tools call engine methods
directly, so usage metering was skipped entirely for MCP mental model operations.

Moved pre-validation and post-completion hooks into memory_engine.py (matching the
retain/recall/reflect pattern) and removed the duplicate code from http.py.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* fix: remove double validation from create_mental_model and add internal checks

- Remove pre-validation from create_mental_model since callers always call
  submit_async_refresh_mental_model next (which validates), preventing
  double credit checks
- Add is_internal checks to mental model metering validators (matching
  the existing pattern for recall/reflect) so background worker tasks
  skip billing

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* fix: prevent 307 redirect on /mcp that breaks MCP tool discovery

Starlette's Mount class redirects /mcp to /mcp/ with a 307 Temporary
Redirect. Many MCP clients don't follow POST redirects, which causes
tool discovery to fail (0 tools discovered despite successful auth).

Add _MCPPathRewriteMiddleware that rewrites /mcp to /mcp/ at the ASGI
level before routing, preventing the redirect entirely. Both /mcp and
/mcp/ now work identically.

Add regression test test_mcp_no_trailing_slash_works to verify URLs
with and without trailing slashes discover tools correctly.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* harden MCP server for real-world usage

- Remove MCP_ENDPOINTS blocklist so banks named "sse"/"messages" route correctly
- Scope SSE body rewriting to text/event-stream responses only to prevent data corruption
- Add _validate_mental_model_inputs for name, source_query, max_tokens validation in MCP tools
- Improve "not found" error messages to include bank_id context
- Fix fragile tool count assertions (exact → minimum bounds)
- Add integration tests: tool execution, input validation, edge-case bank names
- Add unit tests for validation helper and tool-level validation

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* refactor: replace Mount + rewrite middleware with wrapping middleware

Starlette's Mount class redirects /mcp -> /mcp/ with 307, which MCP clients
don't follow. Previously we patched this with _MCPPathRewriteMiddleware.

Now MCPMiddleware wraps the FastAPI app directly via add_middleware, intercepting
/mcp* requests before they reach Starlette's router. No Mount means no redirect.

- Remove _MCPPathRewriteMiddleware (no longer needed)
- Remove app.mount() call
- Add prefix parameter to MCPMiddleware
- Use app.add_middleware() for proper Starlette integration
- Simplify path stripping (just remove prefix, no mount/root_path handling)
- Update routing test to match current behavior (no MCP_ENDPOINTS blocklist)

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* fix: update stale docstring referencing removed _MCPPathRewriteMiddleware

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-11 10:41:20 +01:00

347 lines
14 KiB
Python

"""Hindsight MCP Server implementation using FastMCP (HTTP transport)."""
import json
import logging
import os
from contextvars import ContextVar
from fastmcp import FastMCP
from hindsight_api import MemoryEngine
from hindsight_api.engine.memory_engine import _current_schema
from hindsight_api.extensions import MCPExtension, load_extension
from hindsight_api.extensions.tenant import AuthenticationError
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
from hindsight_api.models import RequestContext
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
_log_level_map = {
"critical": logging.CRITICAL,
"error": logging.ERROR,
"warning": logging.WARNING,
"info": logging.INFO,
"debug": logging.DEBUG,
"trace": logging.DEBUG,
}
logging.basicConfig(
level=_log_level_map.get(_log_level_str, logging.INFO),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
)
logger = logging.getLogger(__name__)
# Default bank_id from environment variable
DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default")
# Legacy MCP authentication token (for backwards compatibility)
# If set, this token is checked first before TenantExtension auth
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)
# Context variables for tenant_id and api_key_id (set by authenticate, used by usage metering)
_current_tenant_id: ContextVar[str | None] = ContextVar("current_tenant_id", default=None)
_current_api_key_id: ContextVar[str | None] = ContextVar("current_api_key_id", 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 get_current_tenant_id() -> str | None:
"""Get the current tenant_id from context."""
return _current_tenant_id.get()
def get_current_api_key_id() -> str | None:
"""Get the current api_key_id from context."""
return _current_api_key_id.get()
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
"""
Create and configure the Hindsight MCP server.
Args:
memory: MemoryEngine instance (required)
multi_bank: If True, expose all tools with bank_id parameters (default).
If False, only expose bank-scoped tools without bank_id parameters.
Returns:
Configured FastMCP server instance with stateless_http enabled
"""
# Use stateless_http=True for Claude Code compatibility
mcp = FastMCP("hindsight-mcp-server", stateless_http=True)
# 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
tenant_id_resolver=get_current_tenant_id, # Propagate tenant_id for usage metering
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
include_bank_id_param=multi_bank,
tools=None
if multi_bank
else {
"retain",
"recall",
"reflect",
"list_mental_models",
"get_mental_model",
"create_mental_model",
"update_mental_model",
"delete_mental_model",
"refresh_mental_model",
}, # Scoped tools for single-bank mode (excludes bank management: list_banks, create_bank)
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
)
register_mcp_tools(mcp, memory, config)
# Load and register additional tools from MCP extension if configured
mcp_extension = load_extension("MCP", MCPExtension)
if mcp_extension:
logger.info(f"Loading MCP extension: {mcp_extension.__class__.__name__}")
mcp_extension.register_tools(mcp, memory)
return mcp
class MCPMiddleware:
"""ASGI middleware that intercepts MCP requests and routes to appropriate MCP server.
This middleware wraps the main FastAPI app and intercepts requests matching the
configured prefix (default: /mcp). Non-MCP requests pass through to the inner app.
Authentication:
1. If HINDSIGHT_API_MCP_AUTH_TOKEN is set (legacy), validates against that token
2. Otherwise, uses TenantExtension.authenticate_mcp() from the MemoryEngine
- DefaultTenantExtension: no auth required (local dev)
- ApiKeyTenantExtension: validates against env var
Two modes based on URL structure:
1. Multi-bank mode (for /mcp/ root endpoint):
- Exposes all tools: retain, recall, reflect, list_banks, create_bank
- All tools include optional bank_id parameter for cross-bank operations
- Bank ID from: X-Bank-Id header or HINDSIGHT_MCP_BANK_ID env var
2. Single-bank mode (for /mcp/{bank_id}/ endpoints):
- Exposes bank-scoped tools only: retain, recall, reflect
- No bank_id parameter (comes from URL)
- No bank management tools (list_banks, create_bank)
- Recommended for agent isolation
Examples:
# Single-bank mode (recommended for agent isolation)
claude mcp add --transport http my-agent http://localhost:8888/mcp/my-agent-bank/ \\
--header "Authorization: Bearer <token>"
# Multi-bank mode (for cross-bank operations)
claude mcp add --transport http hindsight http://localhost:8888/mcp \\
--header "X-Bank-Id: my-bank" --header "Authorization: Bearer <token>"
"""
def __init__(
self,
app,
memory: MemoryEngine,
prefix: str = "/mcp",
multi_bank_app=None,
single_bank_app=None,
multi_bank_server=None,
single_bank_server=None,
):
self.app = app
self.prefix = prefix
self.memory = memory
self.tenant_extension = memory._tenant_extension
if multi_bank_app and single_bank_app:
# Pre-created servers (used when called via add_middleware from create_app)
self.multi_bank_app = multi_bank_app
self.single_bank_app = single_bank_app
self.multi_bank_server = multi_bank_server
self.single_bank_server = single_bank_server
else:
# Create servers internally (for direct construction / tests)
self.multi_bank_server = create_mcp_server(memory, multi_bank=True)
self.multi_bank_app = self.multi_bank_server.http_app(path="/")
self.single_bank_server = create_mcp_server(memory, multi_bank=False)
self.single_bank_app = self.single_bank_server.http_app(path="/")
def _get_header(self, scope: dict, name: str) -> str | None:
"""Extract a header value from ASGI scope."""
name_lower = name.lower().encode()
for header_name, header_value in scope.get("headers", []):
if header_name.lower() == name_lower:
return header_value.decode()
return None
async def __call__(self, scope, receive, send):
if scope["type"] != "http":
await self.app(scope, receive, send)
return
path = scope.get("path", "")
# Check if this is an MCP request (matches prefix)
if not (path == self.prefix or path.startswith(self.prefix + "/")):
# Not an MCP request — pass through to the inner app
await self.app(scope, receive, send)
return
# Strip prefix from path
path = path[len(self.prefix) :] or "/"
# 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 <token>" and direct token
auth_token = auth_header[7:].strip() if auth_header.startswith("Bearer ") else auth_header.strip()
# Authenticate: check legacy MCP_AUTH_TOKEN first, then TenantExtension
tenant_context = None
auth_tenant_id: str | None = None
auth_api_key_id: str | None = None
if MCP_AUTH_TOKEN:
# Legacy authentication mode - validate against static 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
# Legacy mode doesn't use tenant schemas
tenant_context = None
else:
# Use TenantExtension.authenticate_mcp() for auth
try:
auth_context = RequestContext(api_key=auth_token)
tenant_context = await self.tenant_extension.authenticate_mcp(auth_context)
# Capture tenant_id and api_key_id set by authenticate() for usage metering
auth_tenant_id = auth_context.tenant_id
auth_api_key_id = auth_context.api_key_id
except AuthenticationError as e:
await self._send_error(send, 401, str(e))
return
# Set schema from tenant context so downstream DB queries use the correct schema
schema_token = (
_current_schema.set(tenant_context.schema_name) if tenant_context and tenant_context.schema_name else None
)
# Try to get bank_id from header first (for Claude Code compatibility)
bank_id = self._get_header(scope, "X-Bank-Id")
bank_id_from_path = False
# If no header, try to extract from path: /{bank_id}/...
new_path = path
if not bank_id and path.startswith("/") and len(path) > 1:
parts = path[1:].split("/", 1)
if parts[0]:
# First segment looks like a bank_id
bank_id = parts[0]
bank_id_from_path = True
new_path = "/" + parts[1] if len(parts) > 1 else "/"
# Fall back to default bank_id
if not bank_id:
bank_id = DEFAULT_BANK_ID
logger.debug(f"Using default bank_id: {bank_id}")
# Select the appropriate MCP app based on how bank_id was provided:
# - Path-based bank_id → single-bank app (no bank_id param, scoped tools)
# - Header/env bank_id → multi-bank app (bank_id param, all tools)
target_app = self.single_bank_app if bank_id_from_path else self.multi_bank_app
# Set bank_id, api_key, tenant_id, and api_key_id 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
# Store tenant_id and api_key_id from authentication for usage metering
tenant_id_token = _current_tenant_id.set(auth_tenant_id) if auth_tenant_id else None
api_key_id_token = _current_api_key_id.set(auth_api_key_id) if auth_api_key_id else None
try:
new_scope = scope.copy()
new_scope["path"] = new_path
# Clear root_path since we're passing directly to the app
new_scope["root_path"] = ""
# Wrap send to rewrite the SSE endpoint URL to include bank_id if using path-based routing.
# Only rewrite SSE (text/event-stream) responses to avoid corrupting tool results
# that might contain the literal string "data: /messages".
is_sse_response = False
async def send_wrapper(message):
nonlocal is_sse_response
if message["type"] == "http.response.start":
for header_name, header_value in message.get("headers", []):
if header_name == b"content-type" and b"text/event-stream" in header_value:
is_sse_response = True
break
if message["type"] == "http.response.body" and bank_id_from_path and is_sse_response:
body = message.get("body", b"")
if body and b"/messages" in body:
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
body = body.replace(b"data: /messages", f"data: /{bank_id}/messages".encode())
message = {**message, "body": body}
await send(message)
await target_app(new_scope, receive, send_wrapper)
finally:
_current_bank_id.reset(bank_id_token)
if api_key_token is not None:
_current_api_key.reset(api_key_token)
if tenant_id_token is not None:
_current_tenant_id.reset(tenant_id_token)
if api_key_id_token is not None:
_current_api_key_id.reset(api_key_id_token)
if schema_token is not None:
_current_schema.reset(schema_token)
async def _send_error(self, send, status: int, message: str):
"""Send an error response."""
body = json.dumps({"error": message}).encode()
await send(
{
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
}
)
await send(
{
"type": "http.response.body",
"body": body,
}
)
def create_mcp_servers(memory: MemoryEngine):
"""Create multi-bank and single-bank MCP servers and their Starlette apps.
Returns the servers and apps separately so lifespans can be chained before
the middleware wraps the main app.
Returns:
Tuple of (multi_bank_server, single_bank_server, multi_bank_app, single_bank_app)
"""
multi_bank_server = create_mcp_server(memory, multi_bank=True)
multi_bank_app = multi_bank_server.http_app(path="/")
single_bank_server = create_mcp_server(memory, multi_bank=False)
single_bank_app = single_bank_server.http_app(path="/")
return multi_bank_server, single_bank_server, multi_bank_app, single_bank_app