From 181e83d3d6303a76d188ba8cb0a7252c2aa74854 Mon Sep 17 00:00:00 2001 From: Jonathan Singer Date: Wed, 26 Aug 2026 15:07:09 -0400 Subject: [PATCH] let a run take MCP connections from any source and flag failed MCP calls --- strix/core/runner.py | 46 +++--- strix/interface/tui/live_view.py | 35 ++--- strix/tools/mcp/__init__.py | 19 ++- strix/tools/mcp/client.py | 74 ++++++++- strix/tools/mcp/registry.py | 94 +++++++++++- tests/test_mcp_client.py | 255 +++++++++++++++++++++++++++++++ tests/test_runner_mcp.py | 142 +++++++++++++++++ 7 files changed, 614 insertions(+), 51 deletions(-) create mode 100644 tests/test_runner_mcp.py diff --git a/strix/core/runner.py b/strix/core/runner.py index f2c089bf..357c830f 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -58,7 +58,7 @@ if TYPE_CHECKING: from agents.result import RunResultBase from strix.runtime.status import StatusSink - from strix.tools.mcp import ConnectedMcpServer + from strix.tools.mcp import ConnectedMcpServer, McpConnectionRequest logger = logging.getLogger(__name__) @@ -156,6 +156,7 @@ async def run_strix_scan( root_instructions_override: str | None = None, extra_system_prompt_context: dict[str, Any] | None = None, status_sink: StatusSink | None = None, + mcp_connection_requests: list[McpConnectionRequest] | None = None, ) -> RunResultBase | None: """Run or resume one Strix scan against a sandbox. @@ -167,6 +168,11 @@ async def run_strix_scan( ``extra_system_prompt_context`` is merged into the root agent's scan context before prompt rendering. Child agents keep the standard scan prompt and context. + ``mcp_connection_requests`` supplies the run's MCP connections from any + source: when given, the engine connects those requests; when ``None`` (the + command-line default) it reads ``~/.strix/mcp-servers.json`` itself. Either + way the engine does the connecting, so the caller passes inert configs plus + metadata and never live sessions. """ def report(phase: str) -> None: @@ -328,32 +334,38 @@ async def run_strix_scan( 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 + # Attach the run's MCP connections and hold their live sessions in a + # per-run registry. The connections are source-agnostic: a caller + # (the SaaS/pro product) can supply them as mcp_connection_requests, and + # when it does not the command-line path reads them from + # ~/.strix/mcp-servers.json here. Either way one shared engine routine + # does the connecting and populating. 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 ( + McpConnectionRequest, McpRegistry, - connect_mcp_servers, + attach_mcp_requests, 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) + if mcp_connection_requests is None: + # Command-line default: read the user's file and wrap each config + # in a bare request (no provider or transform), so this path is + # exactly the old behavior. + mcp_requests = [ + McpConnectionRequest(config=config) for config in load_user_mcp_configs() + ] + else: + mcp_requests = mcp_connection_requests + if mcp_requests: + connections = await attach_mcp_requests(mcp_requests, mcp_registry) 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) diff --git a/strix/interface/tui/live_view.py b/strix/interface/tui/live_view.py index 3d984883..7e8f2534 100644 --- a/strix/interface/tui/live_view.py +++ b/strix/interface/tui/live_view.py @@ -14,11 +14,7 @@ from agents.tool import ToolOutputImage from strix.core.paths import runtime_state_dir from strix.interface.tui.history import load_session_history - - -# Every MCP call the model makes goes through one of two dispatch tools, so a -# call to a user's server is recognised by the tool's name alone. -_MCP_DISPATCH_TOOLS = frozenset({"call_mcp", "describe_mcp"}) +from strix.tools.mcp import resolve_mcp_call class TuiLiveView: @@ -35,26 +31,19 @@ class TuiLiveView: def _mcp_tool_fields(self, tool_name: str, args: dict[str, Any]) -> dict[str, str]: """Event fields naming the MCP server a tool call went out to, if any. - Every MCP call the model makes goes through one of two dispatch tools: - ``call_mcp`` runs a named tool on a connection, and ``describe_mcp`` - inspects a connection's catalog. The connection, and for ``call_mcp`` the - server's own name for the tool, ride in the call's arguments rather than - in the tool name, so they are read from there. Empty for every other tool, - which is what tells an interface to render the call as one of its own - rather than as a call to a user's server. + Delegates to the shared engine resolver :func:`resolve_mcp_call` so a + dispatch call is attributed the same way here and in strix-pro's tracer. + The projection has no live registry, so it passes none: it reports the + connection and tool read from the call's arguments and leaves the provider + out. Empty for every other tool, which is what tells an interface to + render the call as one of its own rather than as a call to a user's + server. ``describe_mcp`` resolves with an empty tool, which tells both + renderers to present the row as inspecting the connection itself. """ - if tool_name not in _MCP_DISPATCH_TOOLS: + info = resolve_mcp_call(tool_name, args) + if info is None: return {} - connection = args.get("connection") - if not isinstance(connection, str) or not connection: - return {} - # describe_mcp has no underlying tool; an empty tool tells both renderers - # to present the row as inspecting the connection itself. - tool = args.get("tool") if tool_name == "call_mcp" else "" - return { - "mcp_connection": connection, - "mcp_tool": tool if isinstance(tool, str) else "", - } + return {"mcp_connection": info.connection, "mcp_tool": info.tool} def set_user_instruction(self, text: str | None, *, timestamp: str | None = None) -> None: """Open the transcript with what the user asked for. diff --git a/strix/tools/mcp/__init__.py b/strix/tools/mcp/__init__.py index 78f8eaaf..e6037356 100644 --- a/strix/tools/mcp/__init__.py +++ b/strix/tools/mcp/__init__.py @@ -3,7 +3,11 @@ 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.client import ( + ConnectedMcpServer, + attach_mcp_requests, + connect_mcp_servers, +) from strix.tools.mcp.config import ( BearerAuth, McpAuth, @@ -12,27 +16,40 @@ from strix.tools.mcp.config import ( from strix.tools.mcp.loader import load_user_mcp_configs from strix.tools.mcp.naming import namespaced_tool_name from strix.tools.mcp.registry import ( + CALL_MCP_TOOL, + DESCRIBE_MCP_TOOL, + MCP_DISPATCH_TOOLS, MCP_REGISTRY_CONTEXT_KEY, + McpCallInfo, McpConnectionEntry, + McpConnectionRequest, McpConnectionSummary, McpRegistry, mcp_inventory_context, + resolve_mcp_call, ) __all__ = [ + "CALL_MCP_TOOL", + "DESCRIBE_MCP_TOOL", + "MCP_DISPATCH_TOOLS", "MCP_REGISTRY_CONTEXT_KEY", "BearerAuth", "ConnectedMcpServer", "McpAuth", + "McpCallInfo", "McpConnectionConfig", "McpConnectionEntry", + "McpConnectionRequest", "McpConnectionSummary", "McpRegistry", + "attach_mcp_requests", "call_mcp", "connect_mcp_servers", "describe_mcp", "load_user_mcp_configs", "mcp_inventory_context", "namespaced_tool_name", + "resolve_mcp_call", ] diff --git a/strix/tools/mcp/client.py b/strix/tools/mcp/client.py index 2ba5cd6a..1c49ff26 100644 --- a/strix/tools/mcp/client.py +++ b/strix/tools/mcp/client.py @@ -36,6 +36,7 @@ if TYPE_CHECKING: from collections.abc import Callable from strix.tools.mcp.config import McpConnectionConfig + from strix.tools.mcp.registry import McpConnectionRequest, McpRegistry # Runs on one tool call's structured result before it reaches the agent. # Called ``result_transform(label, structured_result)`` and its return value @@ -188,21 +189,44 @@ async def dispatch_mcp_call( 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. + normalizes it through :func:`_errored_tool_output` so the failure reaches + the interfaces (see that function for the representation and why it does not + corrupt the content the agent receives). """ 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} + if getattr(result, "isError", False): + return _errored_tool_output(tool_output) return tool_output +def _errored_tool_output(tool_output: Any) -> dict[str, Any]: + """Tag a serialized MCP error so the interfaces render it as failed. + + Both the TUI and the run viewer decide a tool call failed by reading a + ``success`` key off a top-level dict in the result (``success is False`` means + failed). :func:`_mcp_result_to_tool_output` returns a dict only for a single + content block; a structured-content result comes back as a string and a + multi-block result as a list, and on those the failure flag had nowhere to + ride, so the interfaces showed a failed call as done. This normalizes every + errored result to a top-level dict carrying ``success: False``: + + - a single content block (already a dict) keeps its ``type``/``text`` and gains + ``success: False`` alongside. The SDK's ToolOutput projection keeps the known + ``type``/``text`` fields and drops ``success`` before the agent sees it, so + the agent still receives exactly the error content; + - a list (multiple blocks) or a string (structured content) is placed under a + stable ``content`` key so the flag has a top-level dict to ride on. The agent + still receives the full error content, under ``content``, rather than losing + it. + """ + if isinstance(tool_output, dict): + return {**tool_output, "success": False} + return {"success": False, "content": tool_output} + + async def _count_server_tools(config: McpConnectionConfig, server: MCPServer) -> int: """Count a connected server's reachable tools for the startup summary. @@ -265,3 +289,39 @@ async def connect_mcp_servers( ) return connected + + +async def attach_mcp_requests( + requests: list[McpConnectionRequest], + registry: McpRegistry, +) -> list[ConnectedMcpServer]: + """Connect a caller's MCP requests and populate the run's registry. + + The one shared attach-and-populate path both the command-line and the + SaaS/pro product go through, so all connecting and cleanup lives in one owner. + The caller supplies inert :class:`McpConnectionRequest` objects (a config plus + a provider label, an optional per-connection ``result_transform``, and an + optional ``purpose``) and never a live session: the engine connects each + config here, reusing :func:`connect_mcp_servers` so the fail-open behavior (a + connection that will not connect is logged and skipped) and the cancellation + cleanup are preserved unchanged. + + For each connection that came up, this registers it under its config name with + its tool count, its ``provider`` label, its ``result_transform``, and a purpose + of ``request.purpose`` when set else the connection's notes. Returns the + connected servers (the runner records them and cleans them up when the run + ends). + """ + request_by_name = {request.config.name: request for request in requests} + connections = await connect_mcp_servers([request.config for request in requests]) + for connection in connections: + request = request_by_name[connection.name] + registry.add( + name=connection.name, + server=connection.server, + tool_count=connection.tool_count, + purpose=request.purpose or connection.notes, + provider=request.provider, + result_transform=request.result_transform, + ) + return connections diff --git a/strix/tools/mcp/registry.py b/strix/tools/mcp/registry.py index 318516ef..ea9b55da 100644 --- a/strix/tools/mcp/registry.py +++ b/strix/tools/mcp/registry.py @@ -22,13 +22,14 @@ single dispatch point. from __future__ import annotations import dataclasses -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, NamedTuple if TYPE_CHECKING: from agents.mcp import MCPServer from strix.tools.mcp.client import ResultTransform + from strix.tools.mcp.config import McpConnectionConfig # The run-context key under which the runner stores the per-run registry, and @@ -37,6 +38,16 @@ if TYPE_CHECKING: MCP_REGISTRY_CONTEXT_KEY = "mcp_registry" +# The two generic dispatch tools every agent reaches its MCP connections through. +# ``call_mcp`` runs one tool on a connection; ``describe_mcp`` lists a +# connection's tool schemas. Kept here (not in the interface layer) so the engine, +# the OSS viewer, and strix-pro's tracer all recognise a dispatch call by the same +# names. +CALL_MCP_TOOL = "call_mcp" +DESCRIBE_MCP_TOOL = "describe_mcp" +MCP_DISPATCH_TOOLS = frozenset({CALL_MCP_TOOL, DESCRIBE_MCP_TOOL}) + + @dataclasses.dataclass(frozen=True) class McpConnectionEntry: """One live MCP connection a scan may reach, keyed by ``name``. @@ -46,7 +57,10 @@ class McpConnectionEntry: 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). + at the single dispatch point (strix-pro's sanitizer uses it). ``provider`` is + an optional source label (e.g. ``"supabase"``) the caller tags the connection + with; the command-line path leaves it ``None``, and event tagging surfaces it + when set. """ server: MCPServer @@ -54,6 +68,7 @@ class McpConnectionEntry: purpose: str | None = None tool_count: int = 0 result_transform: ResultTransform | None = None + provider: str | None = None @dataclasses.dataclass(frozen=True) @@ -64,6 +79,37 @@ class McpConnectionSummary: name: str purpose: str | None tool_count: int + provider: str | None = None + + +@dataclasses.dataclass(frozen=True) +class McpConnectionRequest: + """A source-agnostic request to attach one MCP connection to a run. + + The caller hands the engine an inert ``config`` (how to reach the server, its + name, and any auth token) plus metadata, and never a live session: the engine + owns connecting and cleaning up. ``provider`` is an optional source label + (e.g. ``"supabase"``; empty for the command-line path). ``result_transform`` + is an optional per-connection transform run on each call's structured result + at the single dispatch point (strix-pro's sanitizer; empty for the + command-line path). ``purpose`` is the human label shown in the prompt + inventory; when unset it falls back to ``config.notes``. + """ + + config: McpConnectionConfig + provider: str | None = None + result_transform: ResultTransform | None = None + purpose: str | None = None + + +class McpCallInfo(NamedTuple): + """What one MCP dispatch call resolved to: the connection name, the + underlying tool (empty for ``describe_mcp``), and the connection's provider + label (``None`` when unknown or untagged).""" + + connection: str + tool: str + provider: str | None class McpRegistry: @@ -85,6 +131,7 @@ class McpRegistry: purpose: str | None = None, tool_count: int = 0, result_transform: ResultTransform | None = None, + provider: str | None = None, ) -> McpConnectionEntry: """Register one connection under ``name`` (last write wins).""" entry = McpConnectionEntry( @@ -93,6 +140,7 @@ class McpRegistry: purpose=purpose, tool_count=tool_count, result_transform=result_transform, + provider=provider, ) self._entries[name] = entry return entry @@ -109,7 +157,10 @@ class McpRegistry: """One inventory summary per connection, in insertion order.""" return [ McpConnectionSummary( - name=entry.name, purpose=entry.purpose, tool_count=entry.tool_count + name=entry.name, + purpose=entry.purpose, + tool_count=entry.tool_count, + provider=entry.provider, ) for entry in self._entries.values() ] @@ -141,3 +192,40 @@ def mcp_inventory_context(registry: McpRegistry | None) -> list[dict[str, Any]]: {"name": summary.name, "purpose": summary.purpose, "tool_count": summary.tool_count} for summary in registry.summaries() ] + + +def resolve_mcp_call( + tool_name: str, + args: dict[str, Any], + registry: McpRegistry | None = None, +) -> McpCallInfo | None: + """Resolve one tool call to the MCP connection/tool/provider it went out to. + + The single resolver both the OSS viewer and strix-pro's tracer read a + dispatch call through, so a call is attributed the same way everywhere. Every + MCP call an agent makes goes through ``call_mcp`` or ``describe_mcp``, and the + connection (and, for ``call_mcp``, the server's own tool name) ride in the + call's ``args`` rather than the tool name, so they are read from there. + + Returns ``None`` when ``tool_name`` is not one of the two dispatch tools, when + the call carries no connection name, or when a ``registry`` is supplied and + has no connection under that name. ``tool`` is the underlying tool for + ``call_mcp`` and empty for ``describe_mcp`` (which inspects the connection + itself). ``provider`` comes from the registry entry; it is ``None`` when no + ``registry`` is supplied (the viewer projects calls without one) or when the + connection carries no provider label. + """ + if tool_name not in MCP_DISPATCH_TOOLS: + return None + connection = args.get("connection") + if not isinstance(connection, str) or not connection: + return None + provider: str | None = None + if registry is not None: + entry = registry.get(connection) + if entry is None: + return None + provider = entry.provider + raw_tool = args.get("tool") if tool_name == CALL_MCP_TOOL else "" + tool = raw_tool if isinstance(raw_tool, str) else "" + return McpCallInfo(connection=connection, tool=tool, provider=provider) diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 3f097497..218fb05b 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -24,13 +24,17 @@ from strix.interface.tui.live_view import TuiLiveView from strix.tools.mcp import ( MCP_REGISTRY_CONTEXT_KEY, BearerAuth, + McpCallInfo, McpConnectionConfig, + McpConnectionRequest, McpRegistry, + attach_mcp_requests, call_mcp, describe_mcp, load_user_mcp_configs, mcp_inventory_context, namespaced_tool_name, + resolve_mcp_call, ) from strix.tools.mcp import client as mcp_client @@ -97,6 +101,48 @@ class ErroringMCPServer(FakeMCPServer): ) +class MultiBlockErrorServer(FakeMCPServer): + """An errored call whose serialized output is a list (multiple content blocks).""" + + 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="first"), + TextContent(type="text", text="second"), + ], + isError=True, + ) + + +class StructuredErrorServer(FakeMCPServer): + """An errored call whose serialized output is a string (structured content).""" + + def __init__(self, name: str, tools: list[MCPTool]) -> None: + super().__init__(name, tools) + # The base server sets this in __init__, so flip it on the instance to + # take the structured-content serialization branch. + self.use_structured_content = 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="ignored")], + structuredContent={"error": "boom"}, + isError=True, + ) + + def _mcp_tool(name: str, *, description: str | None = None) -> MCPTool: return MCPTool( name=name, @@ -766,3 +812,212 @@ def test_projected_describe_mcp_names_the_connection_with_no_tool() -> None: # the connection rather than as a call to a tool on it. assert describe["mcp_connection"] == "local_fs" assert describe["mcp_tool"] == "" + + +# --- source-agnostic attach -------------------------------------------------- + + +@pytest.mark.asyncio +async def test_attach_populates_registry_with_provider_and_transform( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = FakeMCPServer("db", [_mcp_tool("query")]) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: server) + + def transform(_label: str, structured: Any) -> Any: + return {"kept": structured} + + registry = McpRegistry() + request = McpConnectionRequest( + config=_config("db", ["query"]), + provider="supabase", + result_transform=transform, + purpose="Customer DB", + ) + + connections = await attach_mcp_requests([request], registry) + + assert [(c.name, c.tool_count) for c in connections] == [("db", 1)] + entry = registry.get("db") + assert entry is not None + assert entry.server is server + assert entry.provider == "supabase" + assert entry.purpose == "Customer DB" + assert entry.result_transform is transform + + +@pytest.mark.asyncio +async def test_attach_bare_request_matches_the_command_line_shape( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # The command-line path wraps each config in a bare request (no provider or + # transform); purpose then falls back to the connection's notes. + 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"], + ) + + registry = McpRegistry() + await attach_mcp_requests([McpConnectionRequest(config=config)], registry) + + entry = registry.get("db") + assert entry is not None + assert entry.provider is None + assert entry.result_transform is None + assert entry.purpose == "Staging analytics DB; read-only." + + +@pytest.mark.asyncio +async def test_attach_is_fail_open_and_skips_a_failed_connection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + good = FakeMCPServer("good", [_mcp_tool("t")]) + + class _Failing(FakeMCPServer): + async def connect(self) -> None: + raise RuntimeError("cannot reach server") + + servers = {"good": good, "bad": _Failing("bad", [_mcp_tool("t")])} + monkeypatch.setattr(mcp_client, "_build_server", lambda config: servers[config.name]) + + registry = McpRegistry() + connections = await attach_mcp_requests( + [ + McpConnectionRequest(config=_config("bad", ["t"]), provider="p"), + McpConnectionRequest(config=_config("good", ["t"]), provider="q"), + ], + registry, + ) + + # The failed connection is skipped without raising; the good one is attached. + assert [c.name for c in connections] == ["good"] + assert registry.names() == ["good"] + assert registry.get("good") is not None + assert registry.get("bad") is None + + +# --- provider on the registry ------------------------------------------------ + + +def test_provider_round_trips_through_registry_and_summaries() -> None: + registry = McpRegistry() + registry.add( + name="db", + server=FakeMCPServer("db", []), + purpose="Customer DB", + tool_count=1, + provider="supabase", + ) + registry.add(name="fs", server=FakeMCPServer("fs", []), purpose=None, tool_count=0) + + assert registry.get("db").provider == "supabase" # type: ignore[union-attr] + # A connection with no provider defaults to None, not an error. + assert registry.get("fs").provider is None # type: ignore[union-attr] + + summaries = {s.name: s.provider for s in registry.summaries()} + assert summaries == {"db": "supabase", "fs": None} + + +# --- resolve_mcp_call -------------------------------------------------------- + + +def test_resolve_call_mcp_reads_connection_tool_and_provider() -> None: + registry = McpRegistry() + registry.add(name="db", server=FakeMCPServer("db", []), tool_count=1, provider="supabase") + + info = resolve_mcp_call( + "call_mcp", {"connection": "db", "tool": "query", "arguments": {}}, registry + ) + + assert info == McpCallInfo(connection="db", tool="query", provider="supabase") + + +def test_resolve_describe_mcp_has_an_empty_tool() -> None: + registry = McpRegistry() + registry.add(name="db", server=FakeMCPServer("db", []), tool_count=1, provider="supabase") + + info = resolve_mcp_call("describe_mcp", {"connection": "db"}, registry) + + assert info == McpCallInfo(connection="db", tool="", provider="supabase") + + +def test_resolve_without_a_registry_omits_the_provider() -> None: + # The OSS viewer projects calls with no live registry: it still reads the + # connection and tool, and simply leaves the provider out. + info = resolve_mcp_call("call_mcp", {"connection": "db", "tool": "query"}) + + assert info == McpCallInfo(connection="db", tool="query", provider=None) + + +def test_resolve_returns_none_for_a_non_dispatch_tool() -> None: + assert resolve_mcp_call("exec_command", {"cmd": "ls"}) is None + + +def test_resolve_returns_none_for_an_unknown_connection_with_a_registry() -> None: + registry = McpRegistry() + registry.add(name="db", server=FakeMCPServer("db", []), tool_count=1) + + assert resolve_mcp_call("call_mcp", {"connection": "nope", "tool": "x"}, registry) is None + + +def test_resolve_returns_none_when_the_connection_is_missing_from_args() -> None: + assert resolve_mcp_call("call_mcp", {"tool": "query"}) is None + + +# --- errored results surface as failed regardless of output shape ------------ + + +@pytest.mark.asyncio +async def test_errored_dict_output_carries_success_false() -> None: + registry = McpRegistry() + server = ErroringMCPServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"}) + ) + + # A single content block is a dict; success:False rides alongside and the + # SDK's ToolOutput projection drops it before the agent, so the agent keeps + # the exact error content. + assert out == {"type": "text", "text": "boom:read_file", "success": False} + + +@pytest.mark.asyncio +async def test_errored_list_output_is_wrapped_with_success_false() -> None: + registry = McpRegistry() + server = MultiBlockErrorServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"}) + ) + + # Multiple content blocks serialize to a list, which has no top-level dict to + # carry the flag, so it is wrapped under ``content`` with success:False. + assert out == { + "success": False, + "content": [ + {"type": "text", "text": "first"}, + {"type": "text", "text": "second"}, + ], + } + + +@pytest.mark.asyncio +async def test_errored_structured_output_is_wrapped_with_success_false() -> None: + registry = McpRegistry() + server = StructuredErrorServer("fs", [_mcp_tool("read_file")]) + registry.add(name="fs", server=server, tool_count=1) + + out = await call_mcp.on_invoke_tool( + _ctx(registry), json.dumps({"connection": "fs", "tool": "read_file"}) + ) + + # Structured content serializes to a JSON string; it too is wrapped under + # ``content`` so the failure flag has a top-level dict to ride on. + assert out == {"success": False, "content": json.dumps({"error": "boom"})} diff --git a/tests/test_runner_mcp.py b/tests/test_runner_mcp.py new file mode 100644 index 00000000..d9ccdd1f --- /dev/null +++ b/tests/test_runner_mcp.py @@ -0,0 +1,142 @@ +"""The runner attaches MCP connections source-agnostically. + +When a caller supplies ``mcp_connection_requests`` the runner attaches those; +when it does not, the runner reads ``~/.strix/mcp-servers.json`` itself and wraps +each config in a bare request. Either way the one shared ``attach_mcp_requests`` +routine does the connecting. +""" + +from __future__ import annotations + +import types +from typing import Any + +import pytest +from agents import ModelSettings + +import strix.tools.mcp as mcp_pkg +import strix.tools.notes.tools as notes_tools +import strix.tools.todo.tools as todo_tools +from strix.core import runner +from strix.core.agents import AgentCoordinator +from strix.runtime import session_manager +from strix.tools.mcp import McpConnectionConfig, McpConnectionRequest + + +def _settings() -> Any: + return types.SimpleNamespace( + llm=types.SimpleNamespace( + model="openai/gpt-4o", + reasoning_effort="high", + force_required_tool_choice=False, + timeout=300, + prompt_cache=True, + extra_headers=None, + ), + runtime=types.SimpleNamespace(max_context_images=3), + ) + + +def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: + monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path) + monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path) + monkeypatch.setattr(runner, "setup_scan_logging", lambda _run_dir: lambda: None) + monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None) + monkeypatch.setattr(runner, "load_settings", _settings) + monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _s: None) + monkeypatch.setattr(runner, "uses_chat_completions_tool_schema", lambda _m, _s: False) + monkeypatch.setattr(todo_tools, "hydrate_todos_from_disk", lambda _d: None) + monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _d: None) + + async def _create_or_reuse(*_a: Any, **_k: Any) -> dict[str, Any]: + return {"client": object(), "session": object(), "caido_client": None} + + async def _cleanup(*_a: Any, **_k: Any) -> None: + return None + + monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse) + monkeypatch.setattr(session_manager, "cleanup", _cleanup) + monkeypatch.setattr(runner, "build_root_task", lambda _c: "task") + monkeypatch.setattr(runner, "build_scope_context", lambda _c: {}) + monkeypatch.setattr(runner, "make_model_settings", lambda *_a, **_k: ModelSettings()) + monkeypatch.setattr(runner, "build_strix_agent", lambda **_k: object()) + monkeypatch.setattr(runner, "make_child_factory", lambda **_k: lambda **_kk: object()) + monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object()) + + async def _run_agent_loop(**_kwargs: Any) -> None: + return None + + monkeypatch.setattr(runner, "run_agent_loop", _run_agent_loop) + + +@pytest.mark.asyncio +async def test_none_default_attaches_from_the_user_config_file( + monkeypatch: pytest.MonkeyPatch, tmp_path: Any +) -> None: + _wire_runner(monkeypatch, tmp_path) + + file_config = McpConnectionConfig( + name="local_fs", transport="stdio", command="npx", notes="local files" + ) + monkeypatch.setattr(mcp_pkg, "load_user_mcp_configs", lambda: [file_config]) + + captured: list[list[McpConnectionRequest]] = [] + + async def _capture(requests: list[McpConnectionRequest], _registry: Any) -> list[Any]: + captured.append(requests) + return [] + + monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _capture) + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-none", + image="img", + coordinator=AgentCoordinator(), + ) + + # Each config from the file is wrapped in a bare request: no provider, no + # transform, no explicit purpose (purpose falls back to notes at attach time). + (requests,) = captured + assert len(requests) == 1 + assert requests[0].config is file_config + assert requests[0].provider is None + assert requests[0].result_transform is None + assert requests[0].purpose is None + + +@pytest.mark.asyncio +async def test_supplied_requests_are_attached_and_the_user_file_is_not_read( + monkeypatch: pytest.MonkeyPatch, tmp_path: Any +) -> None: + _wire_runner(monkeypatch, tmp_path) + + def _fail_if_read() -> list[Any]: + raise AssertionError("load_user_mcp_configs must not be read when requests are supplied") + + monkeypatch.setattr(mcp_pkg, "load_user_mcp_configs", _fail_if_read) + + captured: list[list[McpConnectionRequest]] = [] + + async def _capture(requests: list[McpConnectionRequest], _registry: Any) -> list[Any]: + captured.append(requests) + return [] + + monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _capture) + + supplied = [ + McpConnectionRequest( + config=McpConnectionConfig(name="db", url="https://mcp.example.com"), + provider="supabase", + ) + ] + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-supplied", + image="img", + coordinator=AgentCoordinator(), + mcp_connection_requests=supplied, + ) + + assert captured == [supplied]