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:
parent
66423b85f5
commit
f17406fdf0
1 changed files with 59 additions and 36 deletions
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in a new issue