diff --git a/strix/agents/factory.py b/strix/agents/factory.py index b6b15f9c..e6a9ef8c 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -29,6 +29,7 @@ from strix.tools.agents_graph.tools import ( from strix.tools.coverage.tools import list_coverage, record_coverage, update_coverage from strix.tools.finish.tool import finish_scan from strix.tools.load_skill.tool import load_skill +from strix.tools.mcp import call_mcp, describe_mcp from strix.tools.notes.tools import ( create_note, delete_note, @@ -587,6 +588,8 @@ _BASE_TOOLS: tuple[Tool, ...] = ( list_sitemap, view_sitemap_entry, scope_rules, + describe_mcp, + call_mcp, view_agent_graph, send_message_to_agent, wait_for_agents, diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index 95590394..56dd2560 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -75,6 +75,16 @@ AUTHORIZED TARGETS: {% endfor %} {% endif %} +{% if system_prompt_context and system_prompt_context.mcp_connections %} +MCP CONNECTIONS (available this run): +- The user connected these MCP (Model Context Protocol) servers — external tool providers you can reach on demand. Their individual tools do NOT appear in your tool list; the two dispatch tools are the only way in. To use one: + 1. Call describe_mcp(connection="") to list that connection's tools, each with its name, description, and JSON input schema. + 2. Call call_mcp(connection="", tool="", arguments={...}) to run one, passing an arguments object that matches the schema (omit arguments for a tool that takes none). +{% for connection in system_prompt_context.mcp_connections %} +- {{ connection.name }} ({{ connection.tool_count }} tool{{ '' if connection.tool_count == 1 else 's' }}){% if connection.purpose %}: {{ connection.purpose }}{% endif %} +{% endfor %} +{% endif %} + AUTHORIZATION STATUS: - You have FULL AUTHORIZATION for authorized security validation on in-scope targets to help secure the target systems/app - All permission checks have been COMPLETED and APPROVED - never question your authority diff --git a/strix/core/runner.py b/strix/core/runner.py index ee996cd3..f2c089bf 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -91,23 +91,6 @@ def _record_mcp_connections(connections: list[ConnectedMcpServer]) -> None: report_state.record_mcp_connections([connection.name for connection in connections]) -def _mcp_connection_notes(connections: list[ConnectedMcpServer]) -> str | None: - """A block describing the connections the user left notes on, for the agent. - - Only connections with notes are listed, so the note describes the connection - once rather than being repeated onto every tool. Returns ``None`` when no - connection has notes. - """ - noted = [(c.name, c.notes) for c in connections if c.notes] - if not noted: - return None - lines = "\n".join(f"- `{name}.*` tools: {notes}" for name, notes in noted) - return ( - "The user connected these MCP servers for this run and left notes on how " - f"to use each:\n{lines}" - ) - - def _merge_root_prompt_context( scope_context: dict[str, Any], extra_system_prompt_context: dict[str, Any] | None, @@ -344,6 +327,46 @@ async def run_strix_scan( coordinator.set_budget_extender(hooks.extend_budget) scope_context = build_scope_context(scan_config) + + # Connect any MCP servers the user listed in ~/.strix/mcp-servers.json and + # hold their live sessions in a per-run registry. Nothing is registered as + # an agent tool: every agent reaches these connections on demand through + # the describe_mcp / call_mcp tools, guided by a short inventory rendered + # into its prompt. Fail-open: a missing config, or a server that will not + # connect, must never break a run. + from strix.tools.mcp import ( + McpRegistry, + connect_mcp_servers, + load_user_mcp_configs, + mcp_inventory_context, + ) + + mcp_registry = McpRegistry() + try: + user_mcp_configs = load_user_mcp_configs() + if user_mcp_configs: + connections = await connect_mcp_servers(user_mcp_configs) + mcp_servers = [c.server for c in connections] + for connection in connections: + mcp_registry.add( + name=connection.name, + server=connection.server, + purpose=connection.notes, + tool_count=connection.tool_count, + ) + # Recorded even when nothing connected, so a resumed run does not + # keep attributing tool calls to servers it no longer has. + _record_mcp_connections(connections) + if connections: + report(_mcp_startup_summary(connections)) + # The inventory reaches both the root context and the child + # factory's context (both derive from scope_context), so + # every agent renders it. Set only when a connection exists, + # so a run with no MCP leaves the prompt context unchanged. + scope_context["mcp_connections"] = mcp_inventory_context(mcp_registry) + except Exception: + logger.exception("Failed to connect user MCP servers; continuing without them") + root_context = _merge_root_prompt_context(scope_context, extra_system_prompt_context) root_instructions = _compose_root_instructions_override( root_instructions_override, @@ -355,27 +378,6 @@ async def run_strix_scan( system_prompt_context=root_context, ) - # Connect any MCP servers the user listed in ~/.strix/mcp-servers.json and - # register their tools before the agent is built. Fail-open: a missing - # config, or a server that will not connect, must never break a run. - from strix.tools.mcp import connect_mcp_servers, load_user_mcp_configs - - try: - user_mcp_configs = load_user_mcp_configs() - if user_mcp_configs: - connections = await connect_mcp_servers(user_mcp_configs) - mcp_servers = [c.server for c in connections] - # Recorded even when nothing connected, so a resumed run does not - # keep attributing tool calls to servers it no longer has. - _record_mcp_connections(connections) - if connections: - report(_mcp_startup_summary(connections)) - notes_block = _mcp_connection_notes(connections) - if notes_block: - root_task = f"{root_task}\n\n{notes_block}" - except Exception: - logger.exception("Failed to connect user MCP servers; continuing without them") - root_agent = build_strix_agent( name="Root Agent", skills=skills, @@ -427,6 +429,7 @@ async def run_strix_scan( "coordinator": coordinator, "sandbox_session": bundle["session"], "caido_client": bundle["caido_client"], + "mcp_registry": mcp_registry, "agent_id": root_id, "parent_id": None, "interactive": interactive, diff --git a/strix/tools/mcp/__init__.py b/strix/tools/mcp/__init__.py index 4b0160b1..1cd25aae 100644 --- a/strix/tools/mcp/__init__.py +++ b/strix/tools/mcp/__init__.py @@ -1,7 +1,8 @@ -"""Generic MCP client: connect MCP servers and expose their tools.""" +"""Generic MCP client: connect MCP servers and reach their tools on demand.""" from __future__ import annotations +from strix.tools.mcp.agent_tools import call_mcp, describe_mcp from strix.tools.mcp.client import ConnectedMcpServer, connect_mcp_servers from strix.tools.mcp.config import ( BearerAuth, @@ -10,16 +11,30 @@ from strix.tools.mcp.config import ( ) from strix.tools.mcp.loader import load_user_mcp_configs from strix.tools.mcp.naming import McpToolOrigin, namespaced_tool_name, resolve_mcp_tool +from strix.tools.mcp.registry import ( + MCP_REGISTRY_CONTEXT_KEY, + McpConnectionEntry, + McpConnectionSummary, + McpRegistry, + mcp_inventory_context, +) __all__ = [ + "MCP_REGISTRY_CONTEXT_KEY", "BearerAuth", "ConnectedMcpServer", "McpAuth", "McpConnectionConfig", + "McpConnectionEntry", + "McpConnectionSummary", + "McpRegistry", "McpToolOrigin", + "call_mcp", "connect_mcp_servers", + "describe_mcp", "load_user_mcp_configs", + "mcp_inventory_context", "namespaced_tool_name", "resolve_mcp_tool", ] diff --git a/strix/tools/mcp/agent_tools.py b/strix/tools/mcp/agent_tools.py new file mode 100644 index 00000000..780092dd --- /dev/null +++ b/strix/tools/mcp/agent_tools.py @@ -0,0 +1,131 @@ +"""The two generic MCP dispatch tools every agent carries. + +Under the generic-dispatch model an agent does not get one tool per MCP tool. +It gets exactly these two, plus a short inventory in its system prompt naming +which connections exist and what each is for (no schemas): + +- ``describe_mcp(connection)`` returns, as text, one connection's tools with + their names, descriptions, and JSON input schemas — the schemas the model + needs, fetched on demand instead of loaded onto every request up front. +- ``call_mcp(connection, tool, arguments)`` dispatches one call to a + connection's tool and returns its result. + +Both read the per-run :class:`~strix.tools.mcp.registry.McpRegistry` from the run +context under :data:`~strix.tools.mcp.registry.MCP_REGISTRY_CONTEXT_KEY`. They are +ordinary ``FunctionTool`` objects placed in the agent factory's base tool set, so +the factory's output-bounding and disk-spill wrapping apply to their results +automatically. +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING, Any + +from agents import RunContextWrapper, function_tool + +from strix.tools.mcp.client import dispatch_mcp_call +from strix.tools.mcp.naming import namespaced_tool_name +from strix.tools.mcp.registry import MCP_REGISTRY_CONTEXT_KEY, McpRegistry + + +if TYPE_CHECKING: + from mcp.types import Tool as MCPTool + + +def _registry_from_ctx(ctx: RunContextWrapper) -> McpRegistry | None: + context = ctx.context if isinstance(ctx.context, dict) else {} + registry = context.get(MCP_REGISTRY_CONTEXT_KEY) + return registry if isinstance(registry, McpRegistry) else None + + +_NO_CONNECTIONS = "No MCP connections are configured for this run." + + +def _unknown_connection(connection: str, registry: McpRegistry) -> str: + available = ", ".join(registry.names()) or "(none)" + return f"Unknown MCP connection {connection!r}. Available connections: {available}." + + +def _format_tool(tool: MCPTool) -> str: + schema = json.dumps(tool.inputSchema or {"type": "object"}, indent=2, ensure_ascii=False) + description = (tool.description or "").strip() or "(no description)" + return f"- {tool.name}: {description}\n input schema:\n{schema}" + + +@function_tool(timeout=60) +async def describe_mcp(ctx: RunContextWrapper, connection: str) -> str: + """List the tools one MCP connection offers, with their input schemas. + + Read-only. Look up a connection by the name shown in the MCP inventory in + your system prompt; this returns each of its tools with the tool's name, + description, and JSON input schema — the argument shape you pass to + ``call_mcp``. Call this before ``call_mcp`` on any connection you have not + used yet. Nothing is fetched from or run against the connection's data. + + Args: + connection: The connection name exactly as shown in the MCP inventory. + """ + registry = _registry_from_ctx(ctx) + if registry is None or not registry: + return _NO_CONNECTIONS + entry = registry.get(connection) + if entry is None: + return _unknown_connection(connection, registry) + tools = await entry.server.list_tools() + if not tools: + return f"MCP connection {connection!r} offers no tools." + header = f"MCP connection {connection!r} offers {len(tools)} tool(s):" + body = "\n".join(_format_tool(tool) for tool in tools) + return f"{header}\n{body}" + + +@function_tool(timeout=120, strict_mode=False) +async def call_mcp( + ctx: RunContextWrapper, + connection: str, + tool: str, + arguments: Any = None, +) -> Any: + """Call one tool on one MCP connection and return its result. + + Address the tool by the connection name from the MCP inventory and the tool + name from ``describe_mcp`` on that connection. Pass the tool's arguments as an + object matching the input schema ``describe_mcp`` showed for it (omit it, or + pass an empty object, for a tool that takes no arguments). + + Args: + connection: The connection name exactly as shown in the MCP inventory. + tool: The tool name, exactly as reported by ``describe_mcp``. + arguments: The tool's arguments as an object of names to values, or + omitted/empty for a tool that takes none. ``arguments`` is passed + through as-is, so its shape is whatever ``describe_mcp`` showed for + the tool rather than a shape this tool fixes in advance. + """ + registry = _registry_from_ctx(ctx) + if registry is None or not registry: + return _NO_CONNECTIONS + entry = registry.get(connection) + if entry is None: + return _unknown_connection(connection, registry) + if arguments is not None and not isinstance(arguments, dict): + return ( + f"Invalid arguments for {connection!r}.{tool}: expected a JSON object of " + "argument names to values, or none. Call describe_mcp for the input schema." + ) + available = await entry.server.list_tools() + valid_names = {mcp_tool.name for mcp_tool in available} + if tool not in valid_names: + offered = ", ".join(sorted(valid_names)) or "(none)" + return ( + f"Unknown tool {tool!r} on MCP connection {connection!r}. " + f"Tools this connection offers: {offered}. " + "Call describe_mcp for their input schemas." + ) + return await dispatch_mcp_call( + entry.server, + tool, + arguments or {}, + label=namespaced_tool_name(connection, tool), + result_transform=entry.result_transform, + ) diff --git a/strix/tools/mcp/client.py b/strix/tools/mcp/client.py index f1ee9bb0..df0b94df 100644 --- a/strix/tools/mcp/client.py +++ b/strix/tools/mcp/client.py @@ -1,14 +1,15 @@ -"""Connect to MCP servers and expose their tools to the agent. +"""Connect to MCP servers so a run can reach their tools on demand. Given one :class:`McpConnectionConfig` per server, :func:`connect_mcp_servers` -lists each server's tools, keeps the ones on the connection's allowlist (or all -of them when none is set), prefixes each with the connection name so servers do -not collide, and registers them through the agent factory. The factory applies -output bounding, per-call timeouts, and structured errors to every registered -tool, so this layer does not reimplement them. +connects each server, counts the tools it offers (honoring the connection's +allowlist), and returns the live sessions. It does NOT register anything as an +agent tool: under the generic-dispatch model the run holds these sessions in a +per-run :class:`~strix.tools.mcp.registry.McpRegistry`, and the agent reaches +them through the two dispatch tools (``describe_mcp`` / ``call_mcp``), which call +:func:`dispatch_mcp_call` here to run one tool and serialize its result. -A server that cannot connect, or a tool set that cannot be registered, is logged -and skipped, so one bad connection never fails the run. +A server that cannot connect is logged and skipped, so one bad connection never +fails the run. """ from __future__ import annotations @@ -18,34 +19,28 @@ import json import logging from typing import TYPE_CHECKING, Any, NamedTuple, cast -from agents.exceptions import ModelBehaviorError from agents.mcp import ( MCPServer, MCPServerStdio, MCPServerStdioParams, MCPServerStreamableHttp, MCPServerStreamableHttpParams, - MCPUtil, create_static_tool_filter, ) -from strix.agents.factory import register_agent_tools -from strix.tools.mcp.naming import namespaced_tool_name - if TYPE_CHECKING: from collections.abc import Callable - from agents.tool import FunctionTool, Tool - from mcp.types import Tool as MCPTool - from strix.tools.mcp.config import McpConnectionConfig - # Runs on each tool's structured result before it reaches the agent. Called - # ``result_transform(namespaced_tool_name, structured_result)`` and its return - # value becomes the tool's output. ``structured_result`` is the parsed - # ``CallToolResult`` as a dict (not a serialized string), so the transform can - # project or drop individual fields. + # Runs on one tool call's structured result before it reaches the agent. + # Called ``result_transform(label, structured_result)`` and its return value + # becomes the tool's output. ``label`` is the model-facing + # ``_`` name so a transform keyed on names still resolves + # the same way it did under per-tool registration; ``structured_result`` is + # the parsed ``CallToolResult`` as a dict (not a serialized string), so the + # transform can project or drop individual fields. ResultTransform = Callable[[str, Any], Any] @@ -53,12 +48,14 @@ logger = logging.getLogger(__name__) class ConnectedMcpServer(NamedTuple): - """One successfully connected MCP server and how many tools it registered. + """One successfully connected MCP server and how many tools it offers. - ``server`` is kept so the caller can clean it up when the run ends; - ``name`` and ``tool_count`` let the caller show the user a startup summary; + ``server`` is kept so the caller can clean it up when the run ends, and so + the caller can hand the live session to the run's + :class:`~strix.tools.mcp.registry.McpRegistry`; ``name`` and ``tool_count`` + let the caller show the user a startup summary and fill the prompt inventory; ``notes`` carries the connection's optional free-text description so the - caller can surface it to the agent as context about the connection. + caller can surface it as the connection's purpose in the inventory. """ server: MCPServer @@ -79,9 +76,9 @@ def _build_server(config: McpConnectionConfig) -> MCPServer: """Construct (but do not connect) the SDK server for one connection. When ``allowed_tools`` is a list the static filter means the server will not - even list tools outside it; :func:`_register_server_tools` re-applies the - same allowlist as the authoritative gate on what gets registered. When it is - ``None`` no filter is applied and every listed tool is registered. + even list tools outside it, so it is the authoritative gate on what + ``describe_mcp`` and ``call_mcp`` can see. When it is ``None`` no filter is + applied and every listed tool is reachable. """ tool_filter = ( create_static_tool_filter(allowed_tool_names=config.allowed_tools) @@ -114,105 +111,14 @@ def _build_server(config: McpConnectionConfig) -> MCPServer: ) -def _build_tool( - config: McpConnectionConfig, - server: MCPServer, - mcp_tool: MCPTool, - result_transform: ResultTransform | None, -) -> FunctionTool: - """Build one namespaced FunctionTool from a listed MCP tool. - - The SDK builds the tool (so name override, input schema, approval policy, - error-as-result handling, and tool-origin metadata are unchanged). With a - ``result_transform`` we route the underlying MCP call through - :func:`_install_result_transform` so the transform sees the structured result - and decides the tool's output. Without one (the stock path), we still route - the call, through :func:`_install_error_status_capture`, so an errored result - reads as failed in the TUI while the agent's content is unchanged. - """ - namespaced_name = namespaced_tool_name(config.name, mcp_tool.name) - tool = MCPUtil.to_function_tool( - mcp_tool, - server, - convert_schemas_to_strict=False, - tool_name_override=namespaced_name, - ) - if result_transform is not None: - _install_result_transform(tool, server, mcp_tool.name, namespaced_name, result_transform) - else: - _install_error_status_capture(tool, server, mcp_tool.name, namespaced_name) - return tool - - -def _install_result_transform( - tool: FunctionTool, - server: MCPServer, - base_tool_name: str, - namespaced_name: str, - result_transform: ResultTransform, -) -> None: - """Route a tool's MCP call through ``result_transform``, innermost. - - ``MCPUtil.to_function_tool`` serializes the result inside its own invoke, so - the structured result cannot be intercepted through it. Instead we call - ``server.call_tool`` ourselves, hand the parsed :class:`CallToolResult` to the - transform, and return the transform's output as the tool result. - - This runs INSIDE the tool's invoke. The agent factory wraps a registered - tool's ``on_invoke_tool`` with output bounding, disk spill, and tracing at - agent-build time, which is OUTSIDE this invoke, so the transform is genuinely - the innermost step: nothing sees the raw result before the transform does. - - ``to_function_tool`` wraps the real invoke in the SDK's failure-handling - invoker, which stores the inner coroutine on ``_invoke_tool_impl`` and calls - it inside its try/except. Swapping that inner impl keeps the SDK's - error-as-result handling and all tool metadata while inserting the transform. - If the SDK ever renames that attribute we fail loudly rather than silently - skip the transform. - """ - - async def _invoke(_ctx: Any, input_json: str) -> Any: - parsed: Any = json.loads(input_json) if input_json else {} - if not isinstance(parsed, dict): - raise ModelBehaviorError( - f"Invalid JSON input for tool {namespaced_name}: expected a JSON object" - ) - args = cast("dict[str, Any]", parsed) - result = await server.call_tool(base_tool_name, args) - structured_result = result.model_dump(mode="json") - return result_transform(namespaced_name, structured_result) - - _replace_tool_invoke(tool, _invoke) - - -def _replace_tool_invoke(tool: FunctionTool, invoke: Callable[[Any, str], Any]) -> None: - """Swap a FunctionTool's inner invoke, failing loudly if the SDK shape changed. - - ``to_function_tool`` wraps the real invoke in the SDK's failure-handling - invoker, which stores the inner coroutine on ``_invoke_tool_impl`` and calls - it inside its own try/except. Swapping that inner impl keeps the SDK's - error-as-result handling and every piece of tool metadata intact. It is a - plain object with the coroutine as an attribute, not a function, so we treat - it as untyped to swap it. If the SDK ever renames that attribute we raise - rather than silently leave the swap un-applied. - """ - invoker = cast("Any", tool.on_invoke_tool) - if not hasattr(invoker, "_invoke_tool_impl"): - raise RuntimeError( - "agents SDK FunctionTool invoker shape changed: cannot swap the tool " - "invoke without risking it being silently skipped." - ) - invoker._invoke_tool_impl = invoke - - def _mcp_result_to_tool_output(server: MCPServer, result: Any) -> Any: """Serialize a ``CallToolResult`` to a tool output, mirroring the agents SDK. This reproduces the serialization in ``agents.mcp.util.MCPUtil.invoke_mcp_tool`` (structured-content JSON when the server asks for it, otherwise text/image - content blocks, unwrapping a single block). Because the stock path now routes + content blocks, unwrapping a single block). Because the dispatch tool routes its own call, this is what makes the agent see byte-identical content to what - the SDK would have produced on its own. + the SDK would have produced building the tool itself. """ if getattr(server, "use_structured_content", False) and result.structuredContent: return json.dumps(result.structuredContent) @@ -232,85 +138,65 @@ def _mcp_result_to_tool_output(server: MCPServer, result: Any) -> Any: return outputs -def _install_error_status_capture( - tool: FunctionTool, - server: MCPServer, - base_tool_name: str, - namespaced_name: str, -) -> None: - """Make an errored MCP result read as failed in the TUI, agent content unchanged. - - The stock SDK invoke returns only the text/image tool output and drops the - ``CallToolResult.isError`` flag, so the TUI cannot tell an errored MCP call - (which it renders as a green "done") from a successful one. We route the call - the same way :func:`_install_result_transform` does, read ``isError`` off the - full result, and on an error tag the returned output dict with - ``success: False``. - - That tag reaches the human-facing status but not the agent. The SDK stores the - raw return value on the run item's ``output`` (which the TUI reads to derive a - tool's status), but hands the agent the value re-projected through its - ToolOutput schema, which keeps only the known ``type``/``text`` fields and - drops the extra ``success`` key. So the status flips to failed while the agent - still receives exactly the same error content it does today. Non-error calls - return the stock output unchanged and keep rendering as done. - """ - - async def _invoke(_ctx: Any, input_json: str) -> Any: - parsed: Any = json.loads(input_json) if input_json else {} - if not isinstance(parsed, dict): - raise ModelBehaviorError( - f"Invalid JSON input for tool {namespaced_name}: expected a JSON object" - ) - args = cast("dict[str, Any]", parsed) - result = await server.call_tool(base_tool_name, args) - tool_output = _mcp_result_to_tool_output(server, result) - if getattr(result, "isError", False) and isinstance(tool_output, dict): - return {**tool_output, "success": False} - return tool_output - - _replace_tool_invoke(tool, _invoke) - - -async def _register_server_tools( - config: McpConnectionConfig, +async def dispatch_mcp_call( server: MCPServer, + tool_name: str, + arguments: dict[str, Any], + *, + label: str, result_transform: ResultTransform | None = None, -) -> list[Tool]: - """List a connected server's tools, prefix + filter them, and register them. +) -> Any: + """Run one MCP tool call and convert its result to a tool output. - ``allowed_tools`` of ``None`` registers every listed tool; a list restricts - to exactly those names. + Shared single dispatch point for the generic ``call_mcp`` tool. Calls + ``server.call_tool`` with the tool's unprefixed name, then: + + - with a ``result_transform`` (strix-pro's sanitizer), hands the parsed + :class:`CallToolResult` to it as ``result_transform(label, structured)`` and + returns whatever the transform returns; or + - without one, serializes the result the way the agents SDK does (see + :func:`_mcp_result_to_tool_output`) and, when the result is an MCP error, + tags the returned output dict with ``success: False`` so the TUI can tell it + from a success. That tag rides on the human-facing status only: the SDK + re-projects the value through its ToolOutput schema before the agent sees + it, which keeps just the known ``type``/``text`` fields and drops + ``success``, so the agent receives exactly the error content it would have. + """ + result = await server.call_tool(tool_name, arguments) + if result_transform is not None: + return result_transform(label, result.model_dump(mode="json")) + tool_output = _mcp_result_to_tool_output(server, result) + if getattr(result, "isError", False) and isinstance(tool_output, dict): + return {**tool_output, "success": False} + return tool_output + + +async def _count_server_tools(config: McpConnectionConfig, server: MCPServer) -> int: + """Count a connected server's reachable tools for the startup summary. + + ``allowed_tools`` of ``None`` counts every listed tool; a list counts only + those names. The count matches what ``describe_mcp`` will show, because the + static tool filter built in :func:`_build_server` restricts the server's own + ``list_tools`` to the same allowlist. """ allowed = config.allowed_tools mcp_tools = await server.list_tools() - - tools: list[Tool] = [ - _build_tool(config, server, mcp_tool, result_transform) - for mcp_tool in mcp_tools - if allowed is None or mcp_tool.name in allowed - ] - - register_agent_tools(*tools) - return tools + return sum(1 for mcp_tool in mcp_tools if allowed is None or mcp_tool.name in allowed) async def connect_mcp_servers( configs: list[McpConnectionConfig], - result_transform: ResultTransform | None = None, ) -> list[ConnectedMcpServer]: - """Connect to each MCP server and register its tools. - - When ``result_transform`` is given, every registered tool routes its result - through it before the result reaches the agent (see - :func:`_install_result_transform`). When it is ``None`` the tools behave - exactly as the SDK builds them. + """Connect to each MCP server and return its live session. Returns one :class:`ConnectedMcpServer` per server that connected, carrying - the SDK server (so the caller can clean it up when the run ends) plus the - server name and how many tools it registered (so the caller can show the - user a startup summary). Connections that fail are skipped rather than - raised. + the SDK server (so the caller can clean it up when the run ends and hand it to + the run's registry) plus the server name, how many tools it offers, and the + connection's notes. Connections that fail are skipped rather than raised. + + Nothing is registered as an agent tool: the caller builds a per-run + :class:`~strix.tools.mcp.registry.McpRegistry` from these sessions, and the + agent reaches each tool on demand through ``describe_mcp`` / ``call_mcp``. """ connected: list[ConnectedMcpServer] = [] for config in configs: @@ -318,7 +204,7 @@ async def connect_mcp_servers( try: server = _build_server(config) await server.connect() # type: ignore[no-untyped-call] - tools = await _register_server_tools(config, server, result_transform) + tool_count = await _count_server_tools(config, server) except Exception: logger.exception("Skipping MCP connection %r", config.name) if server is not None: @@ -339,10 +225,10 @@ async def connect_mcp_servers( await established.server.cleanup() # type: ignore[no-untyped-call] raise - logger.info("Connected MCP server %r (%d tools)", config.name, len(tools)) + logger.info("Connected MCP server %r (%d tools)", config.name, tool_count) connected.append( ConnectedMcpServer( - server=server, name=config.name, tool_count=len(tools), notes=config.notes + server=server, name=config.name, tool_count=tool_count, notes=config.notes ) ) diff --git a/strix/tools/mcp/config.py b/strix/tools/mcp/config.py index a57121e7..b0cfffd1 100644 --- a/strix/tools/mcp/config.py +++ b/strix/tools/mcp/config.py @@ -61,9 +61,9 @@ class McpConnectionConfig(BaseModel): notes: str | None = None """Free-text notes for the agent describing what this connection is and how - to use it. When set, the runner collects the notes of every connection into - a single block on the root task, so a note describes its connection once - rather than being repeated onto each of its tools.""" + to use it. When set, the note becomes the connection's purpose line in the + MCP inventory every agent renders in its prompt, so it describes the + connection once rather than being repeated onto each of its tools.""" @model_validator(mode="after") def _check_transport_fields(self) -> McpConnectionConfig: diff --git a/strix/tools/mcp/loader.py b/strix/tools/mcp/loader.py index c179981d..168ec8fb 100644 --- a/strix/tools/mcp/loader.py +++ b/strix/tools/mcp/loader.py @@ -1,9 +1,10 @@ """Read the open-source user's MCP servers from ``~/.strix/mcp-servers.json``. An open-source user lists the MCP servers they want the agent to reach in a -small JSON file. Strix reads it at the start of a run, connects to each server, -and registers its tools. The file is optional; without it the run simply gets -no MCP tools. +small JSON file. Strix reads it at the start of a run and connects to each +server, holding the live sessions in the run's registry for the agent to reach +on demand. The file is optional; without it the run simply gets no MCP +connections. Parsing is fail-open. A single malformed entry is logged and skipped rather than raising, so one bad row never blocks the servers that are valid, and a missing @@ -46,9 +47,9 @@ def _resolve_path(path: Path | None) -> Path: def _dedupe_by_name(configs: list[McpConnectionConfig]) -> list[McpConnectionConfig]: """Keep the first connection of each name, dropping later duplicates. - Names namespace a server's tools (``.``), so two connections - sharing a name would collide and the second's tools would be silently - rejected at registration. Drop the duplicate here, with a warning, instead. + A connection's name is its key in the run's registry, so two connections + sharing a name would collide and the second would overwrite the first. Drop + the duplicate here, with a warning, instead. """ seen: set[str] = set() unique: list[McpConnectionConfig] = [] diff --git a/strix/tools/mcp/registry.py b/strix/tools/mcp/registry.py new file mode 100644 index 00000000..318516ef --- /dev/null +++ b/strix/tools/mcp/registry.py @@ -0,0 +1,143 @@ +"""Per-run registry of the MCP connections a scan may reach. + +Replaces per-tool registration. The old model turned every tool of every +connected MCP server into its own agent tool, so a run with a handful of +connections put dozens of provider tool schemas on the root agent's first LLM +request. Instead, a run holds its live connections here, keyed by the name the +user gave each connection, and every agent reaches them through two generic +dispatch tools: ``describe_mcp`` to learn one connection's tool schemas on +demand, and ``call_mcp`` to run one of its tools. + +One :class:`McpRegistry` is built per run in :mod:`strix.core.runner`, stored in +the run context under :data:`MCP_REGISTRY_CONTEXT_KEY`, and shared by the root +agent and every child (the child context is a copy of the parent's, so it +carries the same registry object). + +strix-pro imports :class:`McpRegistry` to add its cloud connections into the +same registry and to attach a per-connection ``result_transform`` (its +sanitizer), which :func:`strix.tools.mcp.client.dispatch_mcp_call` applies at the +single dispatch point. +""" + +from __future__ import annotations + +import dataclasses +from typing import TYPE_CHECKING, Any + + +if TYPE_CHECKING: + from agents.mcp import MCPServer + + from strix.tools.mcp.client import ResultTransform + + +# The run-context key under which the runner stores the per-run registry, and +# the two dispatch tools read it back. Kept here so the tools, the runner, and +# strix-pro all agree on one name. +MCP_REGISTRY_CONTEXT_KEY = "mcp_registry" + + +@dataclasses.dataclass(frozen=True) +class McpConnectionEntry: + """One live MCP connection a scan may reach, keyed by ``name``. + + ``server`` is the connected SDK session the dispatch tools list tools on and + call tools through. ``purpose`` is the human label shown in the prompt + inventory (the user's connection notes, or whatever the caller supplies). + ``tool_count`` is how many tools the connection offers, for the inventory + line. ``result_transform``, when set, runs on each call's structured result + at the single dispatch point (strix-pro's sanitizer uses it). + """ + + server: MCPServer + name: str + purpose: str | None = None + tool_count: int = 0 + result_transform: ResultTransform | None = None + + +@dataclasses.dataclass(frozen=True) +class McpConnectionSummary: + """One inventory line: what an agent needs to decide whether to + ``describe_mcp`` a connection, with no tool schemas.""" + + name: str + purpose: str | None + tool_count: int + + +class McpRegistry: + """Connection name -> live MCP connection, built per run and shared by every + agent in the run. + + Public API (strix-pro builds against it): the constructor, :meth:`add`, + :meth:`get`, and :meth:`summaries`. + """ + + def __init__(self) -> None: + self._entries: dict[str, McpConnectionEntry] = {} + + def add( + self, + *, + name: str, + server: MCPServer, + purpose: str | None = None, + tool_count: int = 0, + result_transform: ResultTransform | None = None, + ) -> McpConnectionEntry: + """Register one connection under ``name`` (last write wins).""" + entry = McpConnectionEntry( + server=server, + name=name, + purpose=purpose, + tool_count=tool_count, + result_transform=result_transform, + ) + self._entries[name] = entry + return entry + + def get(self, name: str) -> McpConnectionEntry | None: + """The connection registered under ``name``, or ``None``.""" + return self._entries.get(name) + + def names(self) -> list[str]: + """The registered connection names, in insertion order.""" + return list(self._entries) + + def summaries(self) -> list[McpConnectionSummary]: + """One inventory summary per connection, in insertion order.""" + return [ + McpConnectionSummary( + name=entry.name, purpose=entry.purpose, tool_count=entry.tool_count + ) + for entry in self._entries.values() + ] + + def clear(self) -> None: + """Drop every connection (the sessions themselves are closed by the + runner).""" + self._entries.clear() + + def __len__(self) -> int: + return len(self._entries) + + def __bool__(self) -> bool: + return bool(self._entries) + + +def mcp_inventory_context(registry: McpRegistry | None) -> list[dict[str, Any]]: + """Build the prompt inventory data for a run's connections. + + Returns one dict per connection with its ``name``, ``purpose``, and + ``tool_count`` (no tool schemas), ready to thread into the system-prompt + context under ``mcp_connections`` so both the root agent and every child + render the same inventory. Returns an empty list when there is no registry + or no connection, and the template's inventory section is then not rendered. + """ + if not registry: + return [] + return [ + {"name": summary.name, "purpose": summary.purpose, "tool_count": summary.tool_count} + for summary in registry.summaries() + ] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 00000000..a840a37b --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,26 @@ +"""Shared test fixtures.""" + +from __future__ import annotations + +import pytest + + +@pytest.fixture(autouse=True) +def _isolate_mcp_config( + monkeypatch: pytest.MonkeyPatch, tmp_path_factory: pytest.TempPathFactory +) -> None: + """Keep the whole suite from reading the developer's real MCP config. + + ``run_strix_scan`` connects the MCP servers listed in + ``~/.strix/mcp-servers.json`` and threads an inventory of them into the + prompt context. Without isolation, any test that drives the runner on a + machine that has a real config would do real network I/O and see MCP + connections it never asked for. Point the loader at a path that does not + exist so it resolves to "no connections", and clear the per-run selection + env vars. Tests that exercise the loader itself set their own + ``STRIX_MCP_CONFIG`` after this runs and so override it. + """ + missing = tmp_path_factory.mktemp("mcp-isolation") / "no-servers.json" + monkeypatch.setenv("STRIX_MCP_CONFIG", str(missing)) + monkeypatch.delenv("STRIX_MCP_ONLY", raising=False) + monkeypatch.delenv("STRIX_MCP_EXCLUDE", raising=False) diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index ec70882b..18eb91e2 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -1,4 +1,9 @@ -"""Tests for the generic MCP client: config contract, namespacing, and filtering.""" +"""Tests for the generic MCP dispatch model. + +Connections are connected without being registered as agent tools; their live +sessions go into a per-run registry; and every agent reaches them through the +two dispatch tools ``describe_mcp`` and ``call_mcp``. +""" from __future__ import annotations @@ -9,30 +14,31 @@ from typing import TYPE_CHECKING, Any import pytest from agents.mcp import MCPServer, MCPServerStdio, MCPServerStreamableHttp +from agents.tool_context import ToolContext from mcp.types import CallToolResult, TextContent from mcp.types import Tool as MCPTool from pydantic import ValidationError from strix.agents import factory -from strix.core.runner import _mcp_connection_notes -from strix.interface.tui.live_view import TuiLiveView, _tool_status_from_result +from strix.interface.tui.live_view import TuiLiveView from strix.tools.mcp import ( + MCP_REGISTRY_CONTEXT_KEY, BearerAuth, - ConnectedMcpServer, McpConnectionConfig, + McpRegistry, + call_mcp, + describe_mcp, load_user_mcp_configs, + mcp_inventory_context, namespaced_tool_name, resolve_mcp_tool, ) from strix.tools.mcp import client as mcp_client -from strix.tools.mcp.client import _auth_headers, _build_server, _register_server_tools if TYPE_CHECKING: from pathlib import Path - from agents.tool import Tool - class FakeMCPServer(MCPServer): """A connected MCP server stand-in, so tests never touch the network.""" @@ -76,15 +82,31 @@ class FakeMCPServer(MCPServer): raise NotImplementedError -def _mcp_tool(name: str) -> MCPTool: +class ErroringMCPServer(FakeMCPServer): + """A connected server whose calls come back as MCP errors (isError=True).""" + + async def call_tool( + self, + tool_name: str, + arguments: dict[str, Any] | None, + meta: dict[str, Any] | None = None, + ) -> CallToolResult: + self.calls.append((tool_name, arguments)) + return CallToolResult( + content=[TextContent(type="text", text=f"boom:{tool_name}")], + isError=True, + ) + + +def _mcp_tool(name: str, *, description: str | None = None) -> MCPTool: return MCPTool( name=name, - description=f"remote tool {name}", - inputSchema={"type": "object", "properties": {}}, + description=description if description is not None else f"remote tool {name}", + inputSchema={"type": "object", "properties": {"path": {"type": "string"}}}, ) -def _config(name: str, allowed_tools: list[str]) -> McpConnectionConfig: +def _config(name: str, allowed_tools: list[str] | None) -> McpConnectionConfig: return McpConnectionConfig( name=name, url="https://mcp.example.com", @@ -93,28 +115,23 @@ def _config(name: str, allowed_tools: list[str]) -> McpConnectionConfig: ) +def _ctx(registry: McpRegistry | None) -> ToolContext[dict[str, Any]]: + context: dict[str, Any] = {} if registry is None else {MCP_REGISTRY_CONTEXT_KEY: registry} + return ToolContext( + context=context, + tool_name="mcp", + tool_call_id="call-1", + tool_arguments="{}", + ) + + @pytest.fixture(autouse=True) def _clear_mcp_env(monkeypatch: pytest.MonkeyPatch) -> None: - """Hide any MCP settings the developer has exported in their own shell. - - The loader reads these to resolve the config path and the per-run - include/exclude selection, so a shell that has them set (from using - --mcp-config or --mcp-server) would otherwise filter what these tests see. - """ + """Hide any MCP settings the developer has exported in their own shell.""" for name in ("STRIX_MCP_CONFIG", "STRIX_MCP_ONLY", "STRIX_MCP_EXCLUDE"): monkeypatch.delenv(name, raising=False) -@pytest.fixture(autouse=True) -def _reset_registry() -> Any: - saved = list(factory._EXTRA_TOOLS) - factory._EXTRA_TOOLS.clear() - try: - yield - finally: - factory._EXTRA_TOOLS[:] = saved - - # --- config contract --------------------------------------------------------- @@ -160,7 +177,6 @@ def test_stdio_config_parses_from_dict() -> None: assert config.command == "npx" assert config.args == ["-y", "@modelcontextprotocol/server-filesystem", "/srv/data"] assert config.env == {"FOO": "bar"} - # A local stdio server needs no auth, and omitting allowed_tools means "all". assert config.auth is None assert config.allowed_tools is None @@ -178,12 +194,7 @@ def test_http_config_without_url_is_rejected() -> None: def test_stdio_config_without_command_is_rejected() -> None: with pytest.raises(ValidationError): - McpConnectionConfig.model_validate( - { - "name": "x", - "transport": "stdio", - } - ) + McpConnectionConfig.model_validate({"name": "x", "transport": "stdio"}) def test_empty_name_is_rejected() -> None: @@ -213,221 +224,61 @@ def test_unknown_field_is_rejected() -> None: def test_bearer_auth_builds_authorization_header() -> None: - headers = _auth_headers(_config("files_main", [])) + headers = mcp_client._auth_headers(_config("files_main", [])) assert headers == {"Authorization": "Bearer run-token"} -# --- namespacing and filtering ----------------------------------------------- - - -def _registered_names() -> list[str]: - return [tool.name for tool in factory.registered_agent_tools()] +# --- connect without global registration ------------------------------------- @pytest.mark.asyncio -async def test_tools_are_namespaced_per_connection() -> None: - server_a = FakeMCPServer("conn_a", [_mcp_tool("describe")]) - server_b = FakeMCPServer("conn_b", [_mcp_tool("describe")]) +async def test_connect_returns_sessions_without_registering_agent_tools( + monkeypatch: pytest.MonkeyPatch, +) -> None: + before = list(factory.registered_agent_tools()) + servers = { + "fs": FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]), + "db": FakeMCPServer("db", [_mcp_tool("query")]), + } + monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) - await _register_server_tools(_config("conn_a", ["describe"]), server_a) - await _register_server_tools(_config("conn_b", ["describe"]), server_b) - - # Same remote tool name on two connections does not collide. - assert _registered_names() == ["conn_a_describe", "conn_b_describe"] - - -@pytest.mark.asyncio -async def test_registered_names_are_valid_tool_names() -> None: - # Model APIs reject a tool name containing anything but letters, digits, - # underscores and hyphens, and reject the whole request rather than the one - # tool. A server naming its own tools with dots, or a connection named with - # a space in the user's config, must not be able to break a run. - server = FakeMCPServer("my server", [_mcp_tool("db.query"), _mcp_tool("ok_tool")]) - - await _register_server_tools(_config("my server", None), server) - - names = _registered_names() - assert names == ["my_server_db_query", "my_server_ok_tool"] - assert all(re.fullmatch(r"[a-zA-Z0-9_-]{1,128}", name) for name in names) - - -@pytest.mark.asyncio -async def test_a_rename_does_not_change_which_tool_is_called() -> None: - # Only the model-facing name is sanitized; the server is always asked for the - # tool name it reported. - server = FakeMCPServer("my server", [_mcp_tool("db.query")]) - - tools = await _register_server_tools(_config("my server", None), server) - - assert tools[0].name == "my_server_db_query" - await tools[0].on_invoke_tool(None, "{}") - assert server.calls == [("db.query", {})] - - -@pytest.mark.asyncio -async def test_disallowed_tool_is_not_registered() -> None: - server = FakeMCPServer( - "files_main", - [_mcp_tool("list_files"), _mcp_tool("search")], + connections = await mcp_client.connect_mcp_servers( + [_config("fs", None), _config("db", ["query"])] ) - await _register_server_tools(_config("files_main", ["list_files"]), server) - - names = _registered_names() - assert "files_main_list_files" in names - assert "files_main_search" not in names + # The live sessions come back with their tool counts, and nothing was added + # to the global agent-tool registry that pro shares. + assert [(c.name, c.tool_count) for c in connections] == [("fs", 2), ("db", 1)] + assert list(factory.registered_agent_tools()) == before @pytest.mark.asyncio -async def test_allowed_tools_none_registers_every_listed_tool() -> None: - server = FakeMCPServer( - "local_fs", - [_mcp_tool("read_file"), _mcp_tool("write_file")], - ) - config = McpConnectionConfig(name="local_fs", url="https://mcp.example.com", allowed_tools=None) +async def test_tool_count_honors_the_allowlist(monkeypatch: pytest.MonkeyPatch) -> None: + server = FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) - await _register_server_tools(config, server) + connections = await mcp_client.connect_mcp_servers([_config("fs", ["read_file"])]) - names = _registered_names() - assert "local_fs_read_file" in names - assert "local_fs_write_file" in names + assert connections[0].tool_count == 1 @pytest.mark.asyncio -async def test_allowed_tools_list_restricts_registration() -> None: - server = FakeMCPServer( - "local_fs", - [_mcp_tool("read_file"), _mcp_tool("write_file")], +async def test_connection_notes_ride_on_the_connection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = FakeMCPServer("db", [_mcp_tool("query")]) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + config = McpConnectionConfig( + name="db", + url="https://mcp.example.com", + notes="Staging analytics DB; read-only.", + allowed_tools=["query"], ) - await _register_server_tools(_config("local_fs", ["read_file"]), server) + connections = await mcp_client.connect_mcp_servers([config]) - names = _registered_names() - assert names == ["local_fs_read_file"] - - -@pytest.mark.asyncio -async def test_registered_tool_routes_to_its_server_with_the_original_name() -> None: - server = FakeMCPServer("files_main", [_mcp_tool("list_files")]) - - tools: list[Tool] = await _register_server_tools(_config("files_main", ["list_files"]), server) - tool = tools[0] - - output = await tool.on_invoke_tool(None, "{}") # type: ignore[union-attr] - - # The call reaches the right server, addressed by the unprefixed remote name. - assert server.calls == [("list_files", {})] - assert output == {"type": "text", "text": "routed:list_files"} - - -# --- result transform -------------------------------------------------------- - - -@pytest.mark.asyncio -async def test_result_transform_receives_namespaced_name_and_structured_result() -> None: - server = FakeMCPServer("files_main", [_mcp_tool("list_files")]) - seen: list[tuple[str, Any]] = [] - - def transform(name: str, structured: Any) -> Any: - seen.append((name, structured)) - return {"kept": structured["content"][0]["text"]} - - tools: list[Tool] = await _register_server_tools( - _config("files_main", ["list_files"]), server, result_transform=transform - ) - - output = await tools[0].on_invoke_tool(None, "{}") # type: ignore[union-attr] - - # The underlying MCP call still routes by the unprefixed remote name. - assert server.calls == [("list_files", {})] - - # The transform is called with the namespaced name and the parsed result. - assert len(seen) == 1 - name, structured = seen[0] - assert name == "files_main_list_files" - # A parsed CallToolResult (dict/list), not a pre-serialized string. - assert structured["content"][0]["text"] == "routed:list_files" - assert structured["isError"] is False - - # The transform's return value is exactly what the tool yields. - assert output == {"kept": "routed:list_files"} - - -@pytest.mark.asyncio -async def test_result_transform_can_rewrite_the_tool_output() -> None: - server = FakeMCPServer("files_main", [_mcp_tool("list_files")]) - - def transform(_name: str, structured: Any) -> Any: - # Keep only a truncated view of the text field. - return structured["content"][0]["text"][:6] - - tools: list[Tool] = await _register_server_tools( - _config("files_main", ["list_files"]), server, result_transform=transform - ) - - output = await tools[0].on_invoke_tool(None, "{}") # type: ignore[union-attr] - - assert output == "routed" - - -@pytest.mark.asyncio -async def test_without_result_transform_output_is_unchanged() -> None: - server = FakeMCPServer("files_main", [_mcp_tool("list_files")]) - - tools: list[Tool] = await _register_server_tools( - _config("files_main", ["list_files"]), server, result_transform=None - ) - - output = await tools[0].on_invoke_tool(None, "{}") # type: ignore[union-attr] - - # Same shape the SDK produces today: no transform in the path. - assert server.calls == [("list_files", {})] - assert output == {"type": "text", "text": "routed:list_files"} - - -# --- error status capture ---------------------------------------------------- - - -class ErroringMCPServer(FakeMCPServer): - """A connected server whose calls come back as MCP errors (isError=True).""" - - async def call_tool( - self, - tool_name: str, - arguments: dict[str, Any] | None, - meta: dict[str, Any] | None = None, - ) -> CallToolResult: - self.calls.append((tool_name, arguments)) - return CallToolResult( - content=[TextContent(type="text", text=f"boom:{tool_name}")], - isError=True, - ) - - -@pytest.mark.asyncio -async def test_errored_mcp_result_is_flagged_failed_for_the_tui() -> None: - server = ErroringMCPServer("files_main", [_mcp_tool("list_files")]) - - tools: list[Tool] = await _register_server_tools(_config("files_main", ["list_files"]), server) - output = await tools[0].on_invoke_tool(None, "{}") # type: ignore[union-attr] - - # The error text stays exactly what the agent gets today; a success:False tag - # rides alongside it purely so the TUI can tell the call apart from a success. - assert output == {"type": "text", "text": "boom:list_files", "success": False} - assert _tool_status_from_result(output) == "failed" - - -@pytest.mark.asyncio -async def test_successful_mcp_result_stays_completed_for_the_tui() -> None: - server = FakeMCPServer("files_main", [_mcp_tool("list_files")]) - - tools: list[Tool] = await _register_server_tools(_config("files_main", ["list_files"]), server) - output = await tools[0].on_invoke_tool(None, "{}") # type: ignore[union-attr] - - # A non-error result is untouched and keeps rendering as done. - assert output == {"type": "text", "text": "routed:list_files"} - assert _tool_status_from_result(output) == "completed" + assert connections[0].notes == "Staging analytics DB; read-only." # --- server build branch ----------------------------------------------------- @@ -442,9 +293,8 @@ def test_build_server_stdio_branch() -> None: env={"TOKEN": "x"}, ) - server = _build_server(config) + server = mcp_client._build_server(config) - # Built, not connected: no subprocess is launched here. assert isinstance(server, MCPServerStdio) assert server.name == "local_fs" assert server.params.command == "my-server" @@ -453,12 +303,233 @@ def test_build_server_stdio_branch() -> None: def test_build_server_http_branch() -> None: - server = _build_server(_config("files_main", ["list_files"])) + server = mcp_client._build_server(_config("files_main", ["list_files"])) assert isinstance(server, MCPServerStreamableHttp) assert server.name == "files_main" +# --- registry ---------------------------------------------------------------- + + +def test_registry_add_get_and_names() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file")]) + + registry.add(name="fs", server=server, purpose="local files", tool_count=1) + + entry = registry.get("fs") + assert entry is not None + assert entry.server is server + assert entry.purpose == "local files" + assert entry.tool_count == 1 + assert registry.get("missing") is None + assert registry.names() == ["fs"] + assert bool(registry) is True + assert len(registry) == 1 + + +def test_registry_summaries_and_inventory() -> None: + registry = McpRegistry() + registry.add(name="fs", server=FakeMCPServer("fs", []), purpose="local files", tool_count=2) + registry.add(name="db", server=FakeMCPServer("db", []), purpose=None, tool_count=1) + + summaries = registry.summaries() + assert [(s.name, s.purpose, s.tool_count) for s in summaries] == [ + ("fs", "local files", 2), + ("db", None, 1), + ] + + # The prompt inventory carries name/purpose/tool_count and no schemas. + assert mcp_inventory_context(registry) == [ + {"name": "fs", "purpose": "local files", "tool_count": 2}, + {"name": "db", "purpose": None, "tool_count": 1}, + ] + + +def test_inventory_is_empty_without_a_registry() -> None: + assert mcp_inventory_context(None) == [] + assert mcp_inventory_context(McpRegistry()) == [] + + +# --- describe_mcp ------------------------------------------------------------ + + +@pytest.mark.asyncio +async def test_describe_mcp_returns_tool_names_and_schemas() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file", description="Read a file")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + out = await describe_mcp.on_invoke_tool(_ctx(registry), json.dumps({"connection": "fs"})) + + assert "read_file" in out + assert "Read a file" in out + # The tool's JSON input schema is shown so the model can build call arguments. + assert '"path"' in out + + +@pytest.mark.asyncio +async def test_describe_mcp_errors_clearly_on_unknown_connection() -> None: + registry = McpRegistry() + registry.add(name="fs", server=FakeMCPServer("fs", []), purpose=None, tool_count=0) + + out = await describe_mcp.on_invoke_tool(_ctx(registry), json.dumps({"connection": "nope"})) + + assert "Unknown MCP connection 'nope'" in out + assert "fs" in out + + +@pytest.mark.asyncio +async def test_describe_mcp_without_any_connections() -> None: + out = await describe_mcp.on_invoke_tool(_ctx(None), json.dumps({"connection": "fs"})) + + assert out == "No MCP connections are configured for this run." + + +# --- call_mcp ---------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_call_mcp_dispatches_and_returns_converted_output() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), + json.dumps({"connection": "fs", "tool": "read_file", "arguments": {"path": "/etc/hosts"}}), + ) + + # The call reaches the server by the unprefixed tool name with its arguments. + assert server.calls == [("read_file", {"path": "/etc/hosts"})] + assert out == {"type": "text", "text": "routed:read_file"} + + +@pytest.mark.asyncio +async def test_call_mcp_defaults_missing_arguments_to_empty_object() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("ping")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + await call_mcp.on_invoke_tool(_ctx(registry), json.dumps({"connection": "fs", "tool": "ping"})) + + assert server.calls == [("ping", {})] + + +@pytest.mark.asyncio +async def test_call_mcp_errors_on_unknown_connection() -> None: + registry = McpRegistry() + registry.add(name="fs", server=FakeMCPServer("fs", []), purpose=None, tool_count=0) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "nope", "tool": "x"}) + ) + + assert "Unknown MCP connection 'nope'" in out + assert "fs" in out + + +@pytest.mark.asyncio +async def test_call_mcp_errors_on_unknown_tool() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "delete_everything"}) + ) + + assert "Unknown tool 'delete_everything'" in out + assert "read_file" in out + # A rejected tool name never reaches the server. + assert server.calls == [] + + +@pytest.mark.asyncio +async def test_call_mcp_errors_on_non_dict_arguments() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), + json.dumps({"connection": "fs", "tool": "read_file", "arguments": ["not", "a", "dict"]}), + ) + + assert "expected a JSON object" in out + assert server.calls == [] + + +@pytest.mark.asyncio +async def test_call_mcp_applies_a_connection_result_transform() -> None: + registry = McpRegistry() + server = FakeMCPServer("fs", [_mcp_tool("read_file")]) + seen: list[tuple[str, Any]] = [] + + def transform(label: str, structured: Any) -> Any: + seen.append((label, structured)) + return {"kept": structured["content"][0]["text"]} + + registry.add(name="fs", server=server, purpose=None, tool_count=1, result_transform=transform) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"}) + ) + + # The transform sees the model-facing _ label and the + # parsed CallToolResult, and its return becomes the tool output. + assert seen[0][0] == "fs_read_file" + assert seen[0][1]["content"][0]["text"] == "routed:read_file" + assert out == {"kept": "routed:read_file"} + + +@pytest.mark.asyncio +async def test_call_mcp_flags_an_errored_result_failed_for_the_tui() -> None: + registry = McpRegistry() + server = ErroringMCPServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, purpose=None, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"}) + ) + + # The agent content is unchanged; success:False rides alongside so the TUI + # can tell an errored call from a done one. + assert out == {"type": "text", "text": "boom:read_file", "success": False} + + +# --- the two tools are the only MCP surface every agent gets ----------------- + + +def test_agent_carries_exactly_the_two_dispatch_tools_regardless_of_connections() -> None: + """No matter how many MCP connections a run makes, an agent's tool list gains + exactly describe_mcp and call_mcp and never a per-connection provider tool.""" + root = factory.build_strix_agent(is_root=True) + child = factory.build_strix_agent(is_root=False) + + root_names = [t.name for t in root.tools] + child_names = [t.name for t in child.tools] + + assert {"describe_mcp", "call_mcp"} <= set(root_names) + assert {"describe_mcp", "call_mcp"} <= set(child_names) + + # Five hypothetical connections would once have added ~all their tools as + # namespaced provider tools; none of those names may appear now. + provider_names = { + namespaced_tool_name(f"conn{i}", tool) + for i in range(5) + for tool in ("read_file", "write_file", "query") + } + assert provider_names.isdisjoint(root_names) + assert provider_names.isdisjoint(child_names) + + # The tool list does not grow with connection count: it is the same set of + # names whether or not any connection exists, because connections never + # contribute tools. + assert root_names == [t.name for t in factory.build_strix_agent(is_root=True).tools] + + # --- loader ------------------------------------------------------------------ @@ -497,12 +568,8 @@ def test_loader_skips_bad_entry_but_keeps_good_ones(tmp_path: Path) -> None: config_file.write_text( json.dumps( [ - {"name": "broken", "transport": "http"}, # missing url - { - "name": "local_fs", - "transport": "stdio", - "command": "npx", - }, + {"name": "broken", "transport": "http"}, + {"name": "local_fs", "transport": "stdio", "command": "npx"}, ] ), encoding="utf-8", @@ -530,51 +597,54 @@ def test_loader_reads_env_var_override(tmp_path: Path, monkeypatch: pytest.Monke assert [c.name for c in configs] == ["local_fs"] -# --- connection notes -------------------------------------------------------- +def _names_file(tmp_path: Path, *names: str) -> Path: + config_file = tmp_path / "mcp-servers.json" + config_file.write_text( + json.dumps([{"name": n, "transport": "stdio", "command": "npx"} for n in names]), + encoding="utf-8", + ) + return config_file -@pytest.mark.asyncio -async def test_connection_notes_are_carried_on_the_connection( - monkeypatch: pytest.MonkeyPatch, -) -> None: - server = FakeMCPServer("db", [_mcp_tool("query")]) - monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) - config = McpConnectionConfig( - name="db", - url="https://mcp.example.com", - notes="Staging analytics DB; read-only.", - allowed_tools=["query"], +def test_loader_drops_duplicate_named_connections(tmp_path: Path) -> None: + config_file = tmp_path / "mcp-servers.json" + config_file.write_text( + json.dumps( + [ + {"name": "dup", "transport": "stdio", "command": "first"}, + {"name": "dup", "transport": "stdio", "command": "second"}, + {"name": "other", "transport": "stdio", "command": "npx"}, + ] + ), + encoding="utf-8", ) - connections = await mcp_client.connect_mcp_servers([config]) + configs = load_user_mcp_configs(config_file) - # Notes ride on the connection (surfaced once), not stapled onto each tool. - assert connections[0].notes == "Staging analytics DB; read-only." + assert [c.name for c in configs] == ["dup", "other"] + assert configs[0].command == "first" -def test_connection_notes_block_lists_only_noted_connections() -> None: - connections = [ - ConnectedMcpServer( - server=FakeMCPServer("db", []), name="db", tool_count=2, notes="staging, read-only" - ), - ConnectedMcpServer(server=FakeMCPServer("fs", []), name="fs", tool_count=1, notes=None), - ] +def test_loader_include_selection_keeps_only_named( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + config_file = _names_file(tmp_path, "a", "b", "c") + monkeypatch.setenv("STRIX_MCP_ONLY", "a,c") - block = _mcp_connection_notes(connections) + configs = load_user_mcp_configs(config_file) - assert block is not None - assert "db" in block - assert "staging, read-only" in block - # A connection without notes is not listed. - assert "fs" not in block + assert [c.name for c in configs] == ["a", "c"] -def test_connection_notes_block_is_none_without_notes() -> None: - connections = [ - ConnectedMcpServer(server=FakeMCPServer("db", []), name="db", tool_count=1, notes=None) - ] +def test_loader_exclude_selection_drops_named( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + config_file = _names_file(tmp_path, "a", "b", "c") + monkeypatch.setenv("STRIX_MCP_EXCLUDE", "b") - assert _mcp_connection_notes(connections) is None + configs = load_user_mcp_configs(config_file) + + assert [c.name for c in configs] == ["a", "c"] # --- cancellation cleanup ---------------------------------------------------- @@ -614,61 +684,10 @@ async def test_connect_cleans_up_when_cancelled_mid_connect( assert cleaned == ["bad", "good"] -# --- duplicate names and run selection --------------------------------------- - - -def _names_file(tmp_path: Path, *names: str) -> Path: - config_file = tmp_path / "mcp-servers.json" - config_file.write_text( - json.dumps([{"name": n, "transport": "stdio", "command": "npx"} for n in names]), - encoding="utf-8", - ) - return config_file - - -def test_loader_drops_duplicate_named_connections(tmp_path: Path) -> None: - config_file = tmp_path / "mcp-servers.json" - config_file.write_text( - json.dumps( - [ - {"name": "dup", "transport": "stdio", "command": "first"}, - {"name": "dup", "transport": "stdio", "command": "second"}, - {"name": "other", "transport": "stdio", "command": "npx"}, - ] - ), - encoding="utf-8", - ) - - configs = load_user_mcp_configs(config_file) - - # Duplicate name is dropped; the first entry wins. - assert [c.name for c in configs] == ["dup", "other"] - assert configs[0].command == "first" - - -def test_loader_include_selection_keeps_only_named( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - config_file = _names_file(tmp_path, "a", "b", "c") - monkeypatch.setenv("STRIX_MCP_ONLY", "a,c") - - configs = load_user_mcp_configs(config_file) - - assert [c.name for c in configs] == ["a", "c"] - - -def test_loader_exclude_selection_drops_named( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: - config_file = _names_file(tmp_path, "a", "b", "c") - monkeypatch.setenv("STRIX_MCP_EXCLUDE", "b") - - configs = load_user_mcp_configs(config_file) - - assert [c.name for c in configs] == ["a", "c"] - - # --- reading a tool call back to the server it went out to ------------------- +# resolve_mcp_tool / namespaced_tool_name stay in strix.tools.mcp.naming: the +# TUI reads them to attribute a call to its connection, and call_mcp builds the +# result_transform label with namespaced_tool_name. def test_resolve_mcp_tool_splits_against_the_run_connections() -> None: @@ -679,12 +698,10 @@ def test_resolve_mcp_tool_splits_against_the_run_connections() -> None: def test_resolve_mcp_tool_prefers_the_longest_matching_connection() -> None: - # One connection's name being a prefix of another's must not misattribute. assert resolve_mcp_tool("files_main_list", ["files", "files_main"]) == ("files_main", "list") def test_resolve_mcp_tool_matches_a_connection_name_it_had_to_sanitize() -> None: - # "my server" reaches the model as "my_server_db_query". tool_name = namespaced_tool_name("my server", "db.query") assert resolve_mcp_tool(tool_name, ["my server"]) == ("my server", "db_query") @@ -692,10 +709,18 @@ def test_resolve_mcp_tool_matches_a_connection_name_it_had_to_sanitize() -> None def test_resolve_mcp_tool_ignores_tools_that_are_not_a_connection_s() -> None: assert resolve_mcp_tool("exec_command", ["local_fs"]) is None - # A name that merely starts like a connection is not one of its tools. assert resolve_mcp_tool("local_fsx", ["local_fs"]) is None +def test_namespaced_name_is_a_valid_tool_name() -> None: + # A connection named with a space and a server tool named with a dot still + # sanitize to a valid model-facing label for the result_transform. + name = namespaced_tool_name("my server", "db.query") + + assert name == "my_server_db_query" + assert re.fullmatch(r"[a-zA-Z0-9_-]{1,128}", name) + + def test_projected_tool_call_names_the_server_it_went_out_to() -> None: view = TuiLiveView() view.set_mcp_connections(["local_fs"]) @@ -711,5 +736,4 @@ def test_projected_tool_call_names_the_server_it_went_out_to() -> None: mcp_call, built_in = (event["data"] for event in view.events) assert (mcp_call["mcp_connection"], mcp_call["mcp_tool"]) == ("local_fs", "read_file") - # A built-in call carries no connection, which is what keeps it rendering as one. assert "mcp_connection" not in built_in