mirror of
https://github.com/usestrix/strix.git
synced 2026-10-07 02:58:26 +00:00
let a run take MCP connections from any source and flag failed MCP calls
This commit is contained in:
parent
e3f95d6cfd
commit
181e83d3d6
7 changed files with 614 additions and 51 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"})}
|
||||
|
|
|
|||
142
tests/test_runner_mcp.py
Normal file
142
tests/test_runner_mcp.py
Normal file
|
|
@ -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]
|
||||
Loading…
Add table
Reference in a new issue