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 = 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 <token>"
|
||||
"""
|
||||
|
||||
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 <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", "")
|
||||
|
||||
# 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}/
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue