feat(mcp): add Bearer token authentication and tenant auth propagation (#241)

* feat(mcp): add Bearer token authentication support

Add HINDSIGHT_API_MCP_AUTH_TOKEN environment variable to enable
authentication for MCP endpoint. When set, all requests must include
a valid Authorization header (Bearer token or direct token).

If not set, MCP endpoint remains open for backwards compatibility
with local development environments.

* fix: propagate Bearer token from MCP middleware to tools for tenant auth

MCP tools were creating RequestContext() without api_key, causing
"Invalid API key" errors when tenant extension validates requests.
Now the Bearer token is extracted in middleware, stored in a context
variable, and passed through to all MCP tool RequestContext instances.
This commit is contained in:
Anton Evseev 2026-01-30 18:08:18 +10:00 committed by GitHub
parent d57e8639c5
commit 0da77ce2c9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 120 additions and 17 deletions

View file

@ -29,15 +29,26 @@ logger = logging.getLogger(__name__)
# Default bank_id from environment variable # Default bank_id from environment variable
DEFAULT_BANK_ID = os.environ.get("HINDSIGHT_MCP_BANK_ID", "default") 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 # Context variable to hold the current bank_id
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None) _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: def get_current_bank_id() -> str | None:
"""Get the current bank_id from context.""" """Get the current bank_id from context."""
return _current_bank_id.get() 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: def create_mcp_server(memory: MemoryEngine) -> FastMCP:
""" """
Create and configure the Hindsight MCP server. 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 # Configure and register tools using shared module
config = MCPToolsConfig( config = MCPToolsConfig(
bank_id_resolver=get_current_bank_id, 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 include_bank_id_param=True, # HTTP MCP supports multi-bank via parameter
tools=None, # All tools tools=None, # All tools
retain_fire_and_forget=False, # HTTP MCP supports sync/async modes retain_fire_and_forget=False, # HTTP MCP supports sync/async modes
@ -65,7 +77,11 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
class MCPMiddleware: 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: Bank ID can be provided via:
1. X-Bank-Id header (recommended for Claude Code) 1. X-Bank-Id header (recommended for Claude Code)
@ -74,7 +90,7 @@ class MCPMiddleware:
For Claude Code, configure with: For Claude Code, configure with:
claude mcp add --transport http hindsight http://localhost:8888/mcp \\ 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 <token>"
""" """
def __init__(self, app, memory: MemoryEngine): def __init__(self, app, memory: MemoryEngine):
@ -98,6 +114,22 @@ class MCPMiddleware:
await self.mcp_app(scope, receive, send) await self.mcp_app(scope, receive, send)
return 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 <token>" 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", "") path = scope.get("path", "")
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped # Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
@ -132,8 +164,10 @@ class MCPMiddleware:
bank_id = DEFAULT_BANK_ID bank_id = DEFAULT_BANK_ID
logger.debug(f"Using default bank_id: {bank_id}") logger.debug(f"Using default bank_id: {bank_id}")
# Set bank_id context # Set bank_id and api_key context
token = _current_bank_id.set(bank_id) 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: try:
new_scope = scope.copy() new_scope = scope.copy()
new_scope["path"] = new_path new_scope["path"] = new_path
@ -152,7 +186,9 @@ class MCPMiddleware:
await self.mcp_app(new_scope, receive, send_wrapper) await self.mcp_app(new_scope, receive, send_wrapper)
finally: 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): async def _send_error(self, send, status: int, message: str):
"""Send an error response.""" """Send an error response."""
@ -176,6 +212,10 @@ def create_mcp_app(memory: MemoryEngine):
""" """
Create an ASGI app that handles MCP requests. 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: 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" 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}/ 2. URL path: /mcp/{bank_id}/

View file

@ -32,6 +32,9 @@ class MCPToolsConfig:
# How to resolve bank_id for operations # How to resolve bank_id for operations
bank_id_resolver: Callable[[], str | None] 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) # Whether to include bank_id as a parameter on tools (for multi-bank support)
include_bank_id_param: bool = False 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 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: def parse_timestamp(timestamp: str) -> datetime | None:
"""Parse an ISO format timestamp string. """Parse an ISO format timestamp string.
@ -155,12 +168,14 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
if error: if error:
return {"status": "error", "message": error} return {"status": "error", "message": error}
request_context = _get_request_context(config)
async def _retain(): async def _retain():
try: try:
await memory.retain_batch_async( await memory.retain_batch_async(
bank_id=target_bank, bank_id=target_bank,
contents=[content_dict], contents=[content_dict],
request_context=RequestContext(), request_context=request_context,
) )
except Exception as e: except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True) 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}" return f"Error: {error}"
contents = [content_dict] contents = [content_dict]
request_context = _get_request_context(config)
if async_processing: if async_processing:
result = await memory.submit_async_retain( 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')})" return f"Memory queued for background processing (operation_id: {result.get('operation_id', 'N/A')})"
else: else:
await memory.retain_batch_async( await memory.retain_batch_async(
bank_id=target_bank, bank_id=target_bank,
contents=contents, contents=contents,
request_context=RequestContext(), request_context=request_context,
) )
return f"Memory stored successfully in bank '{target_bank}'" return f"Memory stored successfully in bank '{target_bank}'"
except Exception as e: except Exception as e:
@ -237,12 +253,14 @@ def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig)
if error: if error:
return {"status": "error", "message": error} return {"status": "error", "message": error}
request_context = _get_request_context(config)
async def _retain(): async def _retain():
try: try:
await memory.retain_batch_async( await memory.retain_batch_async(
bank_id=target_bank, bank_id=target_bank,
contents=[content_dict], contents=[content_dict],
request_context=RequestContext(), request_context=request_context,
) )
except Exception as e: except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True) 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), fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.HIGH, budget=Budget.HIGH,
max_tokens=max_tokens, max_tokens=max_tokens,
request_context=RequestContext(), request_context=_get_request_context(config),
) )
return recall_result.model_dump_json(indent=2) 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), fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.HIGH, budget=Budget.HIGH,
max_tokens=max_tokens, max_tokens=max_tokens,
request_context=RequestContext(), request_context=_get_request_context(config),
) )
return recall_result.model_dump() return recall_result.model_dump()
@ -370,7 +388,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
query=query, query=query,
budget=budget_enum, budget=budget_enum,
context=context, context=context,
request_context=RequestContext(), request_context=_get_request_context(config),
) )
return reflect_result.model_dump_json(indent=2) return reflect_result.model_dump_json(indent=2)
@ -423,7 +441,7 @@ def _register_reflect(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig
query=query, query=query,
budget=budget_enum, budget=budget_enum,
context=context, context=context,
request_context=RequestContext(), request_context=_get_request_context(config),
) )
return reflect_result.model_dump() 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. JSON list of banks with their IDs, names, dispositions, and missions.
""" """
try: 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) return json.dumps({"banks": banks}, indent=2)
except Exception as e: except Exception as e:
logger.error(f"Error listing banks: {e}", exc_info=True) 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 mission: Optional mission describing who the agent is and what they're trying to accomplish
""" """
try: try:
request_context = _get_request_context(config)
# get_bank_profile auto-creates bank if it doesn't exist # 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 # Update name/mission if provided
if name is not None or mission is not None: 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, bank_id,
name=name, name=name,
mission=mission, mission=mission,
request_context=RequestContext(), request_context=request_context,
) )
# Fetch updated profile # 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 # Serialize disposition if it's a Pydantic model
if "disposition" in profile and hasattr(profile["disposition"], "model_dump"): if "disposition" in profile and hasattr(profile["disposition"], "model_dump"):

View file

@ -97,3 +97,47 @@ def test_path_parsing_logic():
bank_id, remaining = parse_path("/my-bank/some/path") bank_id, remaining = parse_path("/my-bank/some/path")
assert bank_id == "my-bank" assert bank_id == "my-bank"
assert remaining == "/some/path" 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)