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:
parent
d57e8639c5
commit
0da77ce2c9
3 changed files with 120 additions and 17 deletions
|
|
@ -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}/
|
||||||
|
|
|
||||||
|
|
@ -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"):
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue