mirror of
https://github.com/usestrix/strix.git
synced 2026-10-10 03:28:11 +00:00
reach MCP tools on demand instead of registering every one
This commit is contained in:
parent
bfaaa904f2
commit
cb5691d6e2
11 changed files with 808 additions and 566 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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="<name>") to list that connection's tools, each with its name, description, and JSON input schema.
|
||||
2. Call call_mcp(connection="<name>", tool="<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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
131
strix/tools/mcp/agent_tools.py
Normal file
131
strix/tools/mcp/agent_tools.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
# ``<connection>_<tool>`` 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
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (``<name>.<tool>``), 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] = []
|
||||
|
|
|
|||
143
strix/tools/mcp/registry.py
Normal file
143
strix/tools/mcp/registry.py
Normal file
|
|
@ -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()
|
||||
]
|
||||
26
tests/conftest.py
Normal file
26
tests/conftest.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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 <connection>_<tool> 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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue