fix: MCP tool calls fail when MCP_AUTH_TOKEN and TENANT_API_KEY differ (#635)
* fix: MCP tool calls fail when MCP_AUTH_TOKEN and TENANT_API_KEY differ When both HINDSIGHT_API_MCP_AUTH_TOKEN and ApiKeyTenantExtension are configured with different values, MCP transport auth passes but tool execution fails because the MCP token gets re-validated against the tenant API key in the engine layer. Add mcp_authenticated flag to RequestContext so the engine skips tenant re-validation when MCP transport auth already succeeded. Fixes #627 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * test: strengthen assertion to verify no auth error in tool response The original test only checked that "banks" key existed in the response, which was true even for error responses like {"error": "...", "banks": []}. Now asserts "error" not in parsed to properly catch auth failures. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
5a486883e8
commit
8364b9c5d5
5 changed files with 123 additions and 3 deletions
|
|
@ -83,6 +83,9 @@ _current_api_key: ContextVar[str | None] = ContextVar("current_api_key", default
|
||||||
_current_tenant_id: ContextVar[str | None] = ContextVar("current_tenant_id", default=None)
|
_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)
|
_current_api_key_id: ContextVar[str | None] = ContextVar("current_api_key_id", default=None)
|
||||||
|
|
||||||
|
# Context variable for MCP pre-authentication flag (set when MCP_AUTH_TOKEN validates)
|
||||||
|
_current_mcp_authenticated: ContextVar[bool] = ContextVar("current_mcp_authenticated", default=False)
|
||||||
|
|
||||||
|
|
||||||
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."""
|
||||||
|
|
@ -104,6 +107,11 @@ def get_current_api_key_id() -> str | None:
|
||||||
return _current_api_key_id.get()
|
return _current_api_key_id.get()
|
||||||
|
|
||||||
|
|
||||||
|
def get_current_mcp_authenticated() -> bool:
|
||||||
|
"""Get whether the request was pre-authenticated by MCP transport auth."""
|
||||||
|
return _current_mcp_authenticated.get()
|
||||||
|
|
||||||
|
|
||||||
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||||
"""
|
"""
|
||||||
Create and configure the Hindsight MCP server.
|
Create and configure the Hindsight MCP server.
|
||||||
|
|
@ -164,6 +172,7 @@ def create_mcp_server(memory: MemoryEngine, multi_bank: bool = True) -> FastMCP:
|
||||||
api_key_resolver=get_current_api_key, # Propagate API key for tenant auth
|
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
|
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
|
api_key_id_resolver=get_current_api_key_id, # Propagate api_key_id for usage metering
|
||||||
|
mcp_authenticated_resolver=get_current_mcp_authenticated, # Propagate MCP pre-auth flag
|
||||||
include_bank_id_param=multi_bank,
|
include_bank_id_param=multi_bank,
|
||||||
tools=base_tools,
|
tools=base_tools,
|
||||||
)
|
)
|
||||||
|
|
@ -312,6 +321,7 @@ class MCPMiddleware:
|
||||||
tenant_context = None
|
tenant_context = None
|
||||||
auth_tenant_id: str | None = None
|
auth_tenant_id: str | None = None
|
||||||
auth_api_key_id: str | None = None
|
auth_api_key_id: str | None = None
|
||||||
|
mcp_pre_authenticated = False
|
||||||
if MCP_AUTH_TOKEN:
|
if MCP_AUTH_TOKEN:
|
||||||
# Legacy authentication mode - validate against static token
|
# Legacy authentication mode - validate against static token
|
||||||
if not auth_token:
|
if not auth_token:
|
||||||
|
|
@ -320,8 +330,9 @@ class MCPMiddleware:
|
||||||
if auth_token != MCP_AUTH_TOKEN:
|
if auth_token != MCP_AUTH_TOKEN:
|
||||||
await self._send_error(send, 401, "Invalid authentication token")
|
await self._send_error(send, 401, "Invalid authentication token")
|
||||||
return
|
return
|
||||||
# Legacy mode doesn't use tenant schemas
|
# Legacy mode: mark as pre-authenticated so tenant extension won't re-validate
|
||||||
tenant_context = None
|
tenant_context = None
|
||||||
|
mcp_pre_authenticated = True
|
||||||
else:
|
else:
|
||||||
# Use TenantExtension.authenticate_mcp() for auth
|
# Use TenantExtension.authenticate_mcp() for auth
|
||||||
try:
|
try:
|
||||||
|
|
@ -368,13 +379,15 @@ class MCPMiddleware:
|
||||||
# - Header/env bank_id → multi-bank app (bank_id param, all 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
|
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
|
# Set bank_id, api_key, tenant_id, api_key_id, and mcp_authenticated context
|
||||||
bank_id_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
|
# Store the auth token for tenant extension to validate
|
||||||
api_key_token = _current_api_key.set(auth_token) if auth_token else None
|
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
|
# 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
|
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
|
api_key_id_token = _current_api_key_id.set(auth_api_key_id) if auth_api_key_id else None
|
||||||
|
# Store MCP pre-authentication flag to skip tenant re-validation
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(mcp_pre_authenticated)
|
||||||
try:
|
try:
|
||||||
new_scope = scope.copy()
|
new_scope = scope.copy()
|
||||||
new_scope["path"] = new_path
|
new_scope["path"] = new_path
|
||||||
|
|
@ -419,6 +432,7 @@ class MCPMiddleware:
|
||||||
_current_tenant_id.reset(tenant_id_token)
|
_current_tenant_id.reset(tenant_id_token)
|
||||||
if api_key_id_token is not None:
|
if api_key_id_token is not None:
|
||||||
_current_api_key_id.reset(api_key_id_token)
|
_current_api_key_id.reset(api_key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
if schema_token is not None:
|
if schema_token is not None:
|
||||||
_current_schema.reset(schema_token)
|
_current_schema.reset(schema_token)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -545,6 +545,12 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
if request_context.internal:
|
if request_context.internal:
|
||||||
return _current_schema.get()
|
return _current_schema.get()
|
||||||
|
|
||||||
|
# For MCP requests already authenticated via MCP_AUTH_TOKEN, skip tenant re-validation.
|
||||||
|
# The MCP transport layer already verified the token; re-validating against the tenant
|
||||||
|
# extension would fail when MCP_AUTH_TOKEN and TENANT_API_KEY differ.
|
||||||
|
if request_context.mcp_authenticated:
|
||||||
|
return _current_schema.get()
|
||||||
|
|
||||||
# Authenticate through tenant extension (always set, may be default no-auth extension)
|
# Authenticate through tenant extension (always set, may be default no-auth extension)
|
||||||
tenant_context = await self._tenant_extension.authenticate(request_context)
|
tenant_context = await self._tenant_extension.authenticate(request_context)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -42,6 +42,9 @@ class MCPToolsConfig:
|
||||||
# How to resolve api_key_id for usage metering (set by MCP middleware after auth)
|
# How to resolve api_key_id for usage metering (set by MCP middleware after auth)
|
||||||
api_key_id_resolver: Callable[[], str | None] | None = None
|
api_key_id_resolver: Callable[[], str | None] | None = None
|
||||||
|
|
||||||
|
# How to resolve mcp_authenticated flag (set when MCP_AUTH_TOKEN validates)
|
||||||
|
mcp_authenticated_resolver: Callable[[], bool] | 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
|
||||||
|
|
||||||
|
|
@ -64,7 +67,10 @@ def _get_request_context(config: MCPToolsConfig) -> RequestContext:
|
||||||
api_key = config.api_key_resolver() if config.api_key_resolver else None
|
api_key = config.api_key_resolver() if config.api_key_resolver else None
|
||||||
tenant_id = config.tenant_id_resolver() if config.tenant_id_resolver else None
|
tenant_id = config.tenant_id_resolver() if config.tenant_id_resolver else None
|
||||||
api_key_id = config.api_key_id_resolver() if config.api_key_id_resolver else None
|
api_key_id = config.api_key_id_resolver() if config.api_key_id_resolver else None
|
||||||
return RequestContext(api_key=api_key, tenant_id=tenant_id, api_key_id=api_key_id)
|
mcp_authenticated = config.mcp_authenticated_resolver() if config.mcp_authenticated_resolver else False
|
||||||
|
return RequestContext(
|
||||||
|
api_key=api_key, tenant_id=tenant_id, api_key_id=api_key_id, mcp_authenticated=mcp_authenticated
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def parse_timestamp(timestamp: str) -> datetime | None:
|
def parse_timestamp(timestamp: str) -> datetime | None:
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,7 @@ class RequestContext:
|
||||||
api_key_id: str | None = None # UUID of the API key used for authentication
|
api_key_id: str | None = None # UUID of the API key used for authentication
|
||||||
tenant_id: str | None = None # Tenant identifier (set by extension after auth)
|
tenant_id: str | None = None # Tenant identifier (set by extension after auth)
|
||||||
internal: bool = False # True for background/internal operations (skips extension auth)
|
internal: bool = False # True for background/internal operations (skips extension auth)
|
||||||
|
mcp_authenticated: bool = False # True when MCP transport auth already validated (skips tenant re-auth)
|
||||||
user_initiated: bool = False # True for async operations that originated from a user request
|
user_initiated: bool = False # True for async operations that originated from a user request
|
||||||
allowed_bank_ids: list[str] | None = None # None = unrestricted (all banks)
|
allowed_bank_ids: list[str] | None = None # None = unrestricted (all banks)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,9 @@ This test verifies that /mcp/ and /mcp/{bank_id}/ expose different tool sets,
|
||||||
and that URLs with or without trailing slashes both work (no 307 redirect).
|
and that URLs with or without trailing slashes both work (no 307 redirect).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
from mcp.client.session import ClientSession
|
from mcp.client.session import ClientSession
|
||||||
|
|
@ -278,3 +281,93 @@ async def test_mcp_bank_named_messages_routes_to_single_bank(memory):
|
||||||
|
|
||||||
assert "retain" in tools
|
assert "retain" in tools
|
||||||
assert "list_banks" not in tools, "Bank 'messages' should route to single-bank mode"
|
assert "list_banks" not in tools, "Bank 'messages' should route to single-bank mode"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mcp_tool_execution_with_different_mcp_and_tenant_tokens(memory):
|
||||||
|
"""Test that MCP tool calls work when MCP_AUTH_TOKEN and TENANT_API_KEY differ.
|
||||||
|
|
||||||
|
Regression test for https://github.com/vectorize-io/hindsight/issues/627
|
||||||
|
When both HINDSIGHT_API_MCP_AUTH_TOKEN and ApiKeyTenantExtension are configured
|
||||||
|
with different values, tool calls should succeed because MCP transport auth
|
||||||
|
already validated the token — the tenant extension should not re-validate.
|
||||||
|
"""
|
||||||
|
from httpx import ASGITransport
|
||||||
|
|
||||||
|
from hindsight_api.api import create_app
|
||||||
|
from hindsight_api.extensions import ApiKeyTenantExtension
|
||||||
|
|
||||||
|
mcp_token = "mcp-secret-token"
|
||||||
|
tenant_key = "tenant-secret-key"
|
||||||
|
|
||||||
|
# Configure ApiKeyTenantExtension with a different key than the MCP token
|
||||||
|
tenant_ext = ApiKeyTenantExtension({"api_key": tenant_key})
|
||||||
|
memory._tenant_extension = tenant_ext
|
||||||
|
|
||||||
|
# Patch MCP_AUTH_TOKEN so the MCP middleware uses legacy auth
|
||||||
|
with patch("hindsight_api.api.mcp.MCP_AUTH_TOKEN", mcp_token):
|
||||||
|
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||||
|
|
||||||
|
async with app.router.lifespan_context(app):
|
||||||
|
# Pass auth header via the httpx client (streamable_http_client doesn't accept headers)
|
||||||
|
async with httpx.AsyncClient(
|
||||||
|
transport=ASGITransport(app=app),
|
||||||
|
base_url="http://test",
|
||||||
|
headers={"Authorization": f"Bearer {mcp_token}"},
|
||||||
|
) as http_client:
|
||||||
|
async with streamable_http_client("http://test/mcp/", http_client=http_client) as (
|
||||||
|
read_stream,
|
||||||
|
write_stream,
|
||||||
|
_,
|
||||||
|
):
|
||||||
|
async with ClientSession(read_stream, write_stream) as session:
|
||||||
|
await session.initialize()
|
||||||
|
|
||||||
|
# list_tools should work
|
||||||
|
tools_result = await session.list_tools()
|
||||||
|
tool_names = {t.name for t in tools_result.tools}
|
||||||
|
assert "get_bank" in tool_names
|
||||||
|
|
||||||
|
# Tool execution should work (this was failing before the fix)
|
||||||
|
result = await session.call_tool("list_banks", arguments={})
|
||||||
|
assert result is not None
|
||||||
|
assert len(result.content) > 0
|
||||||
|
parsed = json.loads(result.content[0].text)
|
||||||
|
assert "banks" in parsed
|
||||||
|
assert "error" not in parsed, f"Tool call failed with: {parsed.get('error')}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mcp_rejects_wrong_mcp_token_even_if_matches_tenant_key(memory):
|
||||||
|
"""Test that an invalid MCP token is rejected even if it matches the tenant key.
|
||||||
|
|
||||||
|
When MCP_AUTH_TOKEN is set, the MCP middleware should validate against that token,
|
||||||
|
not the tenant API key.
|
||||||
|
"""
|
||||||
|
from httpx import ASGITransport
|
||||||
|
|
||||||
|
from hindsight_api.api import create_app
|
||||||
|
from hindsight_api.extensions import ApiKeyTenantExtension
|
||||||
|
|
||||||
|
mcp_token = "mcp-secret-token"
|
||||||
|
tenant_key = "tenant-secret-key"
|
||||||
|
|
||||||
|
tenant_ext = ApiKeyTenantExtension({"api_key": tenant_key})
|
||||||
|
memory._tenant_extension = tenant_ext
|
||||||
|
|
||||||
|
with patch("hindsight_api.api.mcp.MCP_AUTH_TOKEN", mcp_token):
|
||||||
|
app = create_app(memory, mcp_api_enabled=True, initialize_memory=False)
|
||||||
|
|
||||||
|
async with app.router.lifespan_context(app):
|
||||||
|
async with httpx.AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as http_client:
|
||||||
|
# Try connecting with the tenant key (wrong for MCP auth)
|
||||||
|
response = await http_client.post(
|
||||||
|
"http://test/mcp/",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {tenant_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"Accept": "application/json, text/event-stream",
|
||||||
|
},
|
||||||
|
json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}},
|
||||||
|
)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue