feat(mcp): add filter_mcp_tools hook for per-user tool visibility (#737)
Add optional filter_mcp_tools() method to OperationValidatorExtension. Called during tools/list after bank-level mcp_enabled_tools filtering. Extensions can override to hide MCP tools per-user-per-bank based on access policies. Default returns all tools unchanged. - Add filter_mcp_tools to OperationValidatorExtension with default pass-through - Wire into _get_enabled_tools in _apply_bank_tool_filtering - Move _ALL_TOOLS to mcp_tools.py to avoid circular import (re-exported from mcp.py) - Fail-open: if filter raises, log warning and return unfiltered tools - Enforce ceiling: validator can narrow but never expand beyond bank config - Add 8 tests: default, filtering, empty set, integration, composition, can't-add-tools, exception fail-open, no-validator passthrough
This commit is contained in:
parent
a209ef1ae2
commit
f8285b7b90
4 changed files with 422 additions and 40 deletions
|
|
@ -12,44 +12,9 @@ from hindsight_api.config import _get_raw_config
|
||||||
from hindsight_api.engine.memory_engine import _current_schema
|
from hindsight_api.engine.memory_engine import _current_schema
|
||||||
from hindsight_api.extensions import MCPExtension, load_extension
|
from hindsight_api.extensions import MCPExtension, load_extension
|
||||||
from hindsight_api.extensions.tenant import AuthenticationError
|
from hindsight_api.extensions.tenant import AuthenticationError
|
||||||
from hindsight_api.mcp_tools import MCPToolsConfig, register_mcp_tools
|
from hindsight_api.mcp_tools import _ALL_TOOLS, MCPToolsConfig, register_mcp_tools
|
||||||
from hindsight_api.models import RequestContext
|
from hindsight_api.models import RequestContext
|
||||||
|
|
||||||
# All tools available in the system (explicit list — no wildcards)
|
|
||||||
_ALL_TOOLS: frozenset[str] = frozenset(
|
|
||||||
{
|
|
||||||
"retain",
|
|
||||||
"recall",
|
|
||||||
"reflect",
|
|
||||||
"list_banks",
|
|
||||||
"create_bank",
|
|
||||||
"list_mental_models",
|
|
||||||
"get_mental_model",
|
|
||||||
"create_mental_model",
|
|
||||||
"update_mental_model",
|
|
||||||
"delete_mental_model",
|
|
||||||
"refresh_mental_model",
|
|
||||||
"list_directives",
|
|
||||||
"create_directive",
|
|
||||||
"delete_directive",
|
|
||||||
"list_memories",
|
|
||||||
"get_memory",
|
|
||||||
"delete_memory",
|
|
||||||
"list_documents",
|
|
||||||
"get_document",
|
|
||||||
"delete_document",
|
|
||||||
"list_operations",
|
|
||||||
"get_operation",
|
|
||||||
"cancel_operation",
|
|
||||||
"list_tags",
|
|
||||||
"get_bank",
|
|
||||||
"get_bank_stats",
|
|
||||||
"update_bank",
|
|
||||||
"delete_bank",
|
|
||||||
"clear_memories",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
|
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
|
||||||
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
|
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
|
||||||
_log_level_map = {
|
_log_level_map = {
|
||||||
|
|
|
||||||
|
|
@ -722,3 +722,28 @@ class OperationValidatorExtension(Extension, ABC):
|
||||||
BankListResult with the filtered list of banks.
|
BankListResult with the filtered list of banks.
|
||||||
"""
|
"""
|
||||||
return BankListResult(banks=ctx.banks)
|
return BankListResult(banks=ctx.banks)
|
||||||
|
|
||||||
|
async def filter_mcp_tools(
|
||||||
|
self,
|
||||||
|
bank_id: str,
|
||||||
|
request_context: "RequestContext",
|
||||||
|
tools: frozenset[str],
|
||||||
|
) -> frozenset[str]:
|
||||||
|
"""
|
||||||
|
Filter MCP tools visible to this user on this bank.
|
||||||
|
|
||||||
|
Called during tools/list after bank-level mcp_enabled_tools filtering.
|
||||||
|
The input set is already narrowed by bank config — this method can only
|
||||||
|
remove tools, never add ones the bank config excluded.
|
||||||
|
|
||||||
|
Default: return all tools unchanged (no per-user filtering).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
bank_id: Target bank ID (from URL path or header).
|
||||||
|
request_context: Authenticated context with tenant_id set.
|
||||||
|
tools: Tools remaining after bank config filtering.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Subset of tools this user should see.
|
||||||
|
"""
|
||||||
|
return tools
|
||||||
|
|
|
||||||
|
|
@ -24,6 +24,42 @@ from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
|
||||||
from hindsight_api.extensions import OperationValidationError
|
from hindsight_api.extensions import OperationValidationError
|
||||||
from hindsight_api.models import RequestContext
|
from hindsight_api.models import RequestContext
|
||||||
|
|
||||||
|
# All tools available in the system (explicit list — no wildcards).
|
||||||
|
# Defined here (shared module) to avoid circular imports with api/mcp.py.
|
||||||
|
_ALL_TOOLS: frozenset[str] = frozenset(
|
||||||
|
{
|
||||||
|
"retain",
|
||||||
|
"recall",
|
||||||
|
"reflect",
|
||||||
|
"list_banks",
|
||||||
|
"create_bank",
|
||||||
|
"list_mental_models",
|
||||||
|
"get_mental_model",
|
||||||
|
"create_mental_model",
|
||||||
|
"update_mental_model",
|
||||||
|
"delete_mental_model",
|
||||||
|
"refresh_mental_model",
|
||||||
|
"list_directives",
|
||||||
|
"create_directive",
|
||||||
|
"delete_directive",
|
||||||
|
"list_memories",
|
||||||
|
"get_memory",
|
||||||
|
"delete_memory",
|
||||||
|
"list_documents",
|
||||||
|
"get_document",
|
||||||
|
"delete_document",
|
||||||
|
"list_operations",
|
||||||
|
"get_operation",
|
||||||
|
"cancel_operation",
|
||||||
|
"list_tags",
|
||||||
|
"get_bank",
|
||||||
|
"get_bank_stats",
|
||||||
|
"update_bank",
|
||||||
|
"delete_bank",
|
||||||
|
"clear_memories",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -305,11 +341,29 @@ def _apply_bank_tool_filtering(mcp: FastMCP, memory: MemoryEngine, config: MCPTo
|
||||||
if not bank_id:
|
if not bank_id:
|
||||||
return None
|
return None
|
||||||
request_context = _get_request_context(config)
|
request_context = _get_request_context(config)
|
||||||
|
|
||||||
|
# Layer 1: bank config filter (existing)
|
||||||
bank_cfg = await memory._config_resolver.get_bank_config(bank_id, request_context)
|
bank_cfg = await memory._config_resolver.get_bank_config(bank_id, request_context)
|
||||||
enabled: list[str] | None = bank_cfg.get("mcp_enabled_tools")
|
bank_tools: list[str] | None = bank_cfg.get("mcp_enabled_tools")
|
||||||
if enabled is None:
|
enabled: set[str] | None = set(bank_tools) if bank_tools is not None else None
|
||||||
return None
|
|
||||||
return set(enabled)
|
# Layer 2: operation validator filter
|
||||||
|
validator = memory._operation_validator
|
||||||
|
if validator is not None:
|
||||||
|
candidate = frozenset(enabled) if enabled is not None else _ALL_TOOLS
|
||||||
|
try:
|
||||||
|
filtered = await validator.filter_mcp_tools(bank_id, request_context, candidate)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("filter_mcp_tools raised, returning unfiltered tools", exc_info=True)
|
||||||
|
return enabled
|
||||||
|
if filtered != candidate:
|
||||||
|
# Validator can only narrow, never expand beyond the bank config ceiling.
|
||||||
|
if bank_tools is not None:
|
||||||
|
enabled = set(filtered) & set(bank_tools)
|
||||||
|
else:
|
||||||
|
enabled = set(filtered)
|
||||||
|
|
||||||
|
return enabled
|
||||||
|
|
||||||
if hasattr(mcp, "list_tools"):
|
if hasattr(mcp, "list_tools"):
|
||||||
# FastMCP 3.x: wrap list_tools() and get_tool() on the instance
|
# FastMCP 3.x: wrap list_tools() and get_tool() on the instance
|
||||||
|
|
|
||||||
338
hindsight-api-slim/tests/test_mcp_tool_filtering.py
Normal file
338
hindsight-api-slim/tests/test_mcp_tool_filtering.py
Normal file
|
|
@ -0,0 +1,338 @@
|
||||||
|
"""Tests for filter_mcp_tools on OperationValidatorExtension."""
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from hindsight_api.api.mcp import (
|
||||||
|
_current_api_key,
|
||||||
|
_current_api_key_id,
|
||||||
|
_current_bank_id,
|
||||||
|
_current_mcp_authenticated,
|
||||||
|
_current_tenant_id,
|
||||||
|
create_mcp_server,
|
||||||
|
)
|
||||||
|
from hindsight_api.extensions.operation_validator import OperationValidatorExtension, ValidationResult
|
||||||
|
from hindsight_api.models import RequestContext
|
||||||
|
|
||||||
|
|
||||||
|
class MinimalValidator(OperationValidatorExtension):
|
||||||
|
"""Minimal concrete subclass — only implements abstract methods."""
|
||||||
|
|
||||||
|
async def validate_retain(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
async def validate_recall(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
async def validate_reflect(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
|
||||||
|
class FilteringValidator(OperationValidatorExtension):
|
||||||
|
"""Validator that removes retain from the tool set."""
|
||||||
|
|
||||||
|
async def validate_retain(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
async def validate_recall(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
async def validate_reflect(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
|
||||||
|
async def filter_mcp_tools(self, bank_id, request_context, tools):
|
||||||
|
return tools - {"retain"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_filter_mcp_tools_default_returns_all():
|
||||||
|
"""Default implementation returns the input unchanged."""
|
||||||
|
validator = MinimalValidator({})
|
||||||
|
tools = frozenset({"retain", "recall", "reflect", "list_memories"})
|
||||||
|
ctx = RequestContext()
|
||||||
|
|
||||||
|
result = await validator.filter_mcp_tools("test-bank", ctx, tools)
|
||||||
|
|
||||||
|
assert result == tools
|
||||||
|
assert isinstance(result, frozenset)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_filter_mcp_tools_subclass_removes_tools():
|
||||||
|
"""Subclass can remove tools from the set."""
|
||||||
|
validator = FilteringValidator({})
|
||||||
|
tools = frozenset({"retain", "recall", "reflect"})
|
||||||
|
ctx = RequestContext()
|
||||||
|
|
||||||
|
result = await validator.filter_mcp_tools("test-bank", ctx, tools)
|
||||||
|
|
||||||
|
assert result == frozenset({"recall", "reflect"})
|
||||||
|
assert "retain" not in result
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_filter_mcp_tools_returns_empty_set():
|
||||||
|
"""Validator can return empty set — no tools visible."""
|
||||||
|
|
||||||
|
class DenyAllValidator(OperationValidatorExtension):
|
||||||
|
async def validate_retain(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_recall(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_reflect(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def filter_mcp_tools(self, bank_id, request_context, tools):
|
||||||
|
return frozenset()
|
||||||
|
|
||||||
|
validator = DenyAllValidator({})
|
||||||
|
tools = frozenset({"retain", "recall", "reflect"})
|
||||||
|
ctx = RequestContext()
|
||||||
|
|
||||||
|
result = await validator.filter_mcp_tools("test-bank", ctx, tools)
|
||||||
|
|
||||||
|
assert result == frozenset()
|
||||||
|
assert len(result) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_validator_filters_tools_list():
|
||||||
|
"""Validator filter is applied during tools/list via _get_enabled_tools."""
|
||||||
|
mock_memory = MagicMock()
|
||||||
|
mock_memory._tenant_extension = MagicMock()
|
||||||
|
mock_memory._tenant_extension.authenticate_mcp = AsyncMock()
|
||||||
|
mock_memory.retain_batch_async = AsyncMock()
|
||||||
|
mock_memory.submit_async_retain = AsyncMock()
|
||||||
|
mock_memory.recall_async = AsyncMock()
|
||||||
|
mock_memory.reflect_async = AsyncMock()
|
||||||
|
mock_memory.list_banks = AsyncMock(return_value=[])
|
||||||
|
|
||||||
|
validator = FilteringValidator({})
|
||||||
|
mock_memory._operation_validator = validator
|
||||||
|
|
||||||
|
mock_config = {"mcp_enabled_tools": None}
|
||||||
|
mock_memory._config_resolver = MagicMock()
|
||||||
|
mock_memory._config_resolver.get_bank_config = AsyncMock(return_value=mock_config)
|
||||||
|
|
||||||
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||||
|
|
||||||
|
bank_token = _current_bank_id.set("test-bank")
|
||||||
|
api_key_token = _current_api_key.set("hsk_test_key")
|
||||||
|
tenant_token = _current_tenant_id.set("alice")
|
||||||
|
key_id_token = _current_api_key_id.set("key-uuid")
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(False)
|
||||||
|
try:
|
||||||
|
if hasattr(mcp_server, "list_tools"):
|
||||||
|
tools = await mcp_server.list_tools()
|
||||||
|
tool_names = {t.name for t in tools}
|
||||||
|
else:
|
||||||
|
tools = await mcp_server._tool_manager.get_tools()
|
||||||
|
tool_names = set(tools.keys())
|
||||||
|
|
||||||
|
assert "recall" in tool_names
|
||||||
|
assert "reflect" in tool_names
|
||||||
|
assert "retain" not in tool_names
|
||||||
|
finally:
|
||||||
|
_current_bank_id.reset(bank_token)
|
||||||
|
_current_api_key.reset(api_key_token)
|
||||||
|
_current_tenant_id.reset(tenant_token)
|
||||||
|
_current_api_key_id.reset(key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_bank_config_and_validator_compose():
|
||||||
|
"""Bank config sets ceiling, validator narrows further."""
|
||||||
|
mock_memory = MagicMock()
|
||||||
|
mock_memory._tenant_extension = MagicMock()
|
||||||
|
mock_memory._tenant_extension.authenticate_mcp = AsyncMock()
|
||||||
|
mock_memory.retain_batch_async = AsyncMock()
|
||||||
|
mock_memory.submit_async_retain = AsyncMock()
|
||||||
|
mock_memory.recall_async = AsyncMock()
|
||||||
|
mock_memory.reflect_async = AsyncMock()
|
||||||
|
mock_memory.list_banks = AsyncMock(return_value=[])
|
||||||
|
|
||||||
|
mock_memory._operation_validator = FilteringValidator({})
|
||||||
|
|
||||||
|
mock_config = {"mcp_enabled_tools": ["recall", "retain", "reflect"]}
|
||||||
|
mock_memory._config_resolver = MagicMock()
|
||||||
|
mock_memory._config_resolver.get_bank_config = AsyncMock(return_value=mock_config)
|
||||||
|
|
||||||
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||||
|
|
||||||
|
bank_token = _current_bank_id.set("test-bank")
|
||||||
|
api_key_token = _current_api_key.set("hsk_test")
|
||||||
|
tenant_token = _current_tenant_id.set("alice")
|
||||||
|
key_id_token = _current_api_key_id.set("key-1")
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(False)
|
||||||
|
try:
|
||||||
|
if hasattr(mcp_server, "list_tools"):
|
||||||
|
tools = await mcp_server.list_tools()
|
||||||
|
tool_names = {t.name for t in tools}
|
||||||
|
else:
|
||||||
|
tools = await mcp_server._tool_manager.get_tools()
|
||||||
|
tool_names = set(tools.keys())
|
||||||
|
|
||||||
|
assert tool_names == {"recall", "reflect"}
|
||||||
|
finally:
|
||||||
|
_current_bank_id.reset(bank_token)
|
||||||
|
_current_api_key.reset(api_key_token)
|
||||||
|
_current_tenant_id.reset(tenant_token)
|
||||||
|
_current_api_key_id.reset(key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_validator_cannot_add_tools_beyond_bank_config():
|
||||||
|
"""Validator returning tools not in bank config doesn't expand the set."""
|
||||||
|
|
||||||
|
class PermissiveValidator(OperationValidatorExtension):
|
||||||
|
async def validate_retain(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_recall(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_reflect(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def filter_mcp_tools(self, bank_id, request_context, tools):
|
||||||
|
return tools | {"retain", "delete_bank"}
|
||||||
|
|
||||||
|
mock_memory = MagicMock()
|
||||||
|
mock_memory._tenant_extension = MagicMock()
|
||||||
|
mock_memory._tenant_extension.authenticate_mcp = AsyncMock()
|
||||||
|
mock_memory.retain_batch_async = AsyncMock()
|
||||||
|
mock_memory.submit_async_retain = AsyncMock()
|
||||||
|
mock_memory.recall_async = AsyncMock()
|
||||||
|
mock_memory.reflect_async = AsyncMock()
|
||||||
|
mock_memory.list_banks = AsyncMock(return_value=[])
|
||||||
|
|
||||||
|
mock_memory._operation_validator = PermissiveValidator({})
|
||||||
|
|
||||||
|
mock_config = {"mcp_enabled_tools": ["recall"]}
|
||||||
|
mock_memory._config_resolver = MagicMock()
|
||||||
|
mock_memory._config_resolver.get_bank_config = AsyncMock(return_value=mock_config)
|
||||||
|
|
||||||
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||||
|
|
||||||
|
bank_token = _current_bank_id.set("test-bank")
|
||||||
|
api_key_token = _current_api_key.set("hsk_test")
|
||||||
|
tenant_token = _current_tenant_id.set("alice")
|
||||||
|
key_id_token = _current_api_key_id.set("key-1")
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(False)
|
||||||
|
try:
|
||||||
|
if hasattr(mcp_server, "list_tools"):
|
||||||
|
tools = await mcp_server.list_tools()
|
||||||
|
tool_names = {t.name for t in tools}
|
||||||
|
else:
|
||||||
|
tools = await mcp_server._tool_manager.get_tools()
|
||||||
|
tool_names = set(tools.keys())
|
||||||
|
|
||||||
|
assert "recall" in tool_names
|
||||||
|
assert "retain" not in tool_names
|
||||||
|
assert "delete_bank" not in tool_names
|
||||||
|
finally:
|
||||||
|
_current_bank_id.reset(bank_token)
|
||||||
|
_current_api_key.reset(api_key_token)
|
||||||
|
_current_tenant_id.reset(tenant_token)
|
||||||
|
_current_api_key_id.reset(key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_validator_exception_fails_open(caplog):
|
||||||
|
"""If filter_mcp_tools raises, all tools remain visible and warning is logged."""
|
||||||
|
import logging
|
||||||
|
caplog.set_level(logging.WARNING)
|
||||||
|
|
||||||
|
class BrokenValidator(OperationValidatorExtension):
|
||||||
|
async def validate_retain(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_recall(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def validate_reflect(self, ctx):
|
||||||
|
return ValidationResult.accept()
|
||||||
|
async def filter_mcp_tools(self, bank_id, request_context, tools):
|
||||||
|
raise RuntimeError("Policy backend unreachable")
|
||||||
|
|
||||||
|
mock_memory = MagicMock()
|
||||||
|
mock_memory._tenant_extension = MagicMock()
|
||||||
|
mock_memory._tenant_extension.authenticate_mcp = AsyncMock()
|
||||||
|
mock_memory.retain_batch_async = AsyncMock()
|
||||||
|
mock_memory.submit_async_retain = AsyncMock()
|
||||||
|
mock_memory.recall_async = AsyncMock()
|
||||||
|
mock_memory.reflect_async = AsyncMock()
|
||||||
|
mock_memory.list_banks = AsyncMock(return_value=[])
|
||||||
|
|
||||||
|
mock_memory._operation_validator = BrokenValidator({})
|
||||||
|
mock_config = {"mcp_enabled_tools": None}
|
||||||
|
mock_memory._config_resolver = MagicMock()
|
||||||
|
mock_memory._config_resolver.get_bank_config = AsyncMock(return_value=mock_config)
|
||||||
|
|
||||||
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||||
|
|
||||||
|
bank_token = _current_bank_id.set("test-bank")
|
||||||
|
api_key_token = _current_api_key.set("hsk_test")
|
||||||
|
tenant_token = _current_tenant_id.set("alice")
|
||||||
|
key_id_token = _current_api_key_id.set("key-1")
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(False)
|
||||||
|
try:
|
||||||
|
if hasattr(mcp_server, "list_tools"):
|
||||||
|
tools = await mcp_server.list_tools()
|
||||||
|
tool_names = {t.name for t in tools}
|
||||||
|
else:
|
||||||
|
tools = await mcp_server._tool_manager.get_tools()
|
||||||
|
tool_names = set(tools.keys())
|
||||||
|
|
||||||
|
assert "retain" in tool_names
|
||||||
|
assert "recall" in tool_names
|
||||||
|
assert "reflect" in tool_names
|
||||||
|
|
||||||
|
assert any("filter_mcp_tools raised" in r.message for r in caplog.records)
|
||||||
|
finally:
|
||||||
|
_current_bank_id.reset(bank_token)
|
||||||
|
_current_api_key.reset(api_key_token)
|
||||||
|
_current_tenant_id.reset(tenant_token)
|
||||||
|
_current_api_key_id.reset(key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_no_validator_returns_unfiltered():
|
||||||
|
"""Without an operation validator, tools/list returns all tools."""
|
||||||
|
mock_memory = MagicMock()
|
||||||
|
mock_memory._tenant_extension = MagicMock()
|
||||||
|
mock_memory._tenant_extension.authenticate_mcp = AsyncMock()
|
||||||
|
mock_memory.retain_batch_async = AsyncMock()
|
||||||
|
mock_memory.submit_async_retain = AsyncMock()
|
||||||
|
mock_memory.recall_async = AsyncMock()
|
||||||
|
mock_memory.reflect_async = AsyncMock()
|
||||||
|
mock_memory.list_banks = AsyncMock(return_value=[])
|
||||||
|
|
||||||
|
mock_memory._operation_validator = None
|
||||||
|
mock_config = {"mcp_enabled_tools": None}
|
||||||
|
mock_memory._config_resolver = MagicMock()
|
||||||
|
mock_memory._config_resolver.get_bank_config = AsyncMock(return_value=mock_config)
|
||||||
|
|
||||||
|
mcp_server = create_mcp_server(mock_memory, multi_bank=False)
|
||||||
|
|
||||||
|
bank_token = _current_bank_id.set("test-bank")
|
||||||
|
api_key_token = _current_api_key.set("hsk_test")
|
||||||
|
tenant_token = _current_tenant_id.set("alice")
|
||||||
|
key_id_token = _current_api_key_id.set("key-1")
|
||||||
|
mcp_auth_token = _current_mcp_authenticated.set(False)
|
||||||
|
try:
|
||||||
|
if hasattr(mcp_server, "list_tools"):
|
||||||
|
tools = await mcp_server.list_tools()
|
||||||
|
tool_names = {t.name for t in tools}
|
||||||
|
else:
|
||||||
|
tools = await mcp_server._tool_manager.get_tools()
|
||||||
|
tool_names = set(tools.keys())
|
||||||
|
|
||||||
|
assert "retain" in tool_names
|
||||||
|
assert "recall" in tool_names
|
||||||
|
assert "reflect" in tool_names
|
||||||
|
finally:
|
||||||
|
_current_bank_id.reset(bank_token)
|
||||||
|
_current_api_key.reset(api_key_token)
|
||||||
|
_current_tenant_id.reset(tenant_token)
|
||||||
|
_current_api_key_id.reset(key_id_token)
|
||||||
|
_current_mcp_authenticated.reset(mcp_auth_token)
|
||||||
Loading…
Reference in a new issue