Fix bank-level MCP tool filtering for FastMCP 3.x (#491)

FastMCP 3.x replaced _tool_manager.get_tools() with a provider pattern
(LocalProvider._list_tools via _components). The existing wrapper on
_tool_manager.get_tools() silently failed (caught AttributeError) since
_tool_manager no longer exists in v3.

Now wraps FastMCP.list_tools() and FastMCP.get_tool() for v3, while
preserving the _tool_manager approach for v2 compatibility.
This commit is contained in:
Chris Bartholomew 2026-03-04 10:29:29 -05:00 committed by GitHub
parent 66423b85f5
commit f17406fdf0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -271,48 +271,71 @@ def register_mcp_tools(
def _apply_bank_tool_filtering(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None:
"""Filter bank-level mcp_enabled_tools from both tools/list and tool invocation.
Wraps _tool_manager.get_tools() so that:
- tools/list only returns permitted tools (they are hidden, not just blocked)
- tools/call for a disabled tool raises NotFoundError (via the manager) before run()
tool.run wrappers are kept as defense-in-depth for any caller that bypasses the manager.
Compatible with FastMCP 2.x (_tool_manager pattern) and 3.x (provider pattern).
"""
async def _get_enabled_tools() -> set[str] | None:
"""Return the enabled tool set for the current bank, or None if unrestricted."""
bank_id = config.bank_id_resolver()
if not bank_id:
return None
request_context = _get_request_context(config)
bank_cfg = await memory._config_resolver.get_bank_config(bank_id, request_context)
enabled: list[str] | None = bank_cfg.get("mcp_enabled_tools")
if enabled is None:
return None
return set(enabled)
if hasattr(mcp, "list_tools"):
# FastMCP 3.x: wrap list_tools() and get_tool() on the instance
original_list_tools = mcp.list_tools
original_get_tool = mcp.get_tool
async def _filtered_list_tools(**kwargs):
tools = await original_list_tools(**kwargs)
enabled_set = await _get_enabled_tools()
if enabled_set is None:
return tools
return [t for t in tools if t.name in enabled_set]
async def _filtered_get_tool(name, **kwargs):
enabled_set = await _get_enabled_tools()
if enabled_set is not None and name not in enabled_set:
return None # FastMCP treats None as "not found" → raises NotFoundError
return await original_get_tool(name, **kwargs)
object.__setattr__(mcp, "list_tools", _filtered_list_tools)
object.__setattr__(mcp, "get_tool", _filtered_get_tool)
elif hasattr(mcp, "_tool_manager"):
# FastMCP 2.x: wrap _tool_manager.get_tools() and tool.run()
try:
tool_manager = mcp._tool_manager
original_get_tools = tool_manager.get_tools
async def _filtered_get_tools():
all_tools = await original_get_tools()
bank_id = config.bank_id_resolver()
if not bank_id:
enabled_set = await _get_enabled_tools()
if enabled_set is None:
return all_tools
request_context = _get_request_context(config)
bank_cfg = await memory._config_resolver.get_bank_config(bank_id, request_context)
enabled: list[str] | None = bank_cfg.get("mcp_enabled_tools")
if enabled is None:
return all_tools
enabled_set = set(enabled)
return {k: v for k, v in all_tools.items() if k in enabled_set}
setattr(tool_manager, "get_tools", _filtered_get_tools)
# Defense-in-depth: also wrap tool.run for any direct caller that bypasses the manager
for name, tool in tool_manager._tools.items():
original_run = tool.run
async def _filtered_run(arguments, _name=name, _orig=original_run):
bank_id = config.bank_id_resolver()
if bank_id:
request_context = _get_request_context(config)
bank_cfg = await memory._config_resolver.get_bank_config(bank_id, request_context)
enabled: list[str] | None = bank_cfg.get("mcp_enabled_tools")
if enabled is not None and _name not in enabled:
raise ValueError(f"Tool '{_name}' is not enabled for bank '{bank_id}'")
enabled_set = await _get_enabled_tools()
if enabled_set is not None and _name not in enabled_set:
raise ValueError(f"Tool '{_name}' is not enabled for bank '{config.bank_id_resolver()}'")
return await _orig(arguments)
object.__setattr__(tool, "run", _filtered_run)
except (AttributeError, KeyError) as e:
logger.warning(f"Could not apply bank tool filtering: {e}")
logger.warning(f"Could not apply bank tool filtering (v2): {e}")
else:
logger.warning("Could not apply bank tool filtering: unknown FastMCP version")
def _register_retain(mcp: FastMCP, memory: MemoryEngine, config: MCPToolsConfig) -> None: