From e158eab3f869a0911f0697e3311f8c47a7625bea Mon Sep 17 00:00:00 2001 From: yoni-at-strix Date: Tue, 22 Sep 2026 14:02:02 -0400 Subject: [PATCH] feat(mcp): initialize connections lazily (#1347) * feat(mcp): initialize connections lazily * fix(mcp): replace terminally dead sessions * fix(mcp): improve targeted tool discovery * fix(mcp): limit active tool fallback --- strix/agents/factory.py | 10 +- strix/agents/prompts/system_prompt.jinja | 12 +- strix/core/runner.py | 107 +++------ strix/tools/mcp/__init__.py | 14 +- strix/tools/mcp/agent_tools.py | 181 +++++++++++--- strix/tools/mcp/client.py | 55 +++-- strix/tools/mcp/config.py | 3 + strix/tools/mcp/registry.py | 274 ++++++++++++++++++--- tests/test_mcp_client.py | 290 ++++++++++++++++++++++- tests/test_runner_mcp.py | 62 +++-- tests/test_runner_root_prompt.py | 13 +- 11 files changed, 804 insertions(+), 217 deletions(-) diff --git a/strix/agents/factory.py b/strix/agents/factory.py index b2fcbf08c..ac31444dc 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -29,7 +29,13 @@ 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, list_mcps +from strix.tools.mcp import ( + call_mcp, + describe_mcp, + get_mcp_tool_schema, + list_mcps, + search_mcp_tools, +) from strix.tools.notes.tools import ( create_note, delete_note, @@ -592,6 +598,8 @@ _BASE_TOOLS: tuple[Tool, ...] = ( view_sitemap_entry, scope_rules, list_mcps, + search_mcp_tools, + get_mcp_tool_schema, describe_mcp, call_mcp, view_agent_graph, diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index 18fffa99d..7934dfb99 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -77,18 +77,20 @@ AUTHORIZED TARGETS: {% if system_prompt_context and system_prompt_context.mcp_available %} MCP CONNECTIONS (available this run): -- The user connected one or more MCP (Model Context Protocol) servers — external tool providers you can reach on demand. Their individual tools do NOT appear in your tool list; three dispatch tools are the only way in. +- The user connected one or more MCP (Model Context Protocol) servers. Their individual tools do NOT appear in your tool list. Use the four discovery and dispatch tools to reach them. {% if system_prompt_context.mcp_connections %} -- Connected this run (call describe_mcp on one to see its tools): +- Connected this run (search one to find relevant tools): {% for connection in system_prompt_context.mcp_connections %} - {{ connection.name }} ({{ connection.tool_count }} tools){% if connection.purpose %}: {{ connection.purpose }}{% endif %} {% endfor %} {% endif %} - Reach for a connection whenever the target itself cannot give you information a connection could: its database schema and access policies, real deployment or infrastructure configuration, known issues or prior findings, or server logs. In those cases call list_mcps early to see what is available, and prefer a connection's authoritative data over inferring from the target's responses. Do not wait to be told a connection exists. 1. Call list_mcps() to discover the available connections. - 2. Call describe_mcp(connection="") to inspect one connection's tools, each with its name, description, and JSON input schema. - 3. Call call_mcp(connection="", tool="", arguments={...}) to run one, passing an arguments object that matches the schema (omit arguments for a tool that takes none). -- Do not assume a connection or tool exists; discover it with list_mcps and describe it with describe_mcp before calling. + 2. Call search_mcp_tools(connection="", query="") for a short candidate list. + 3. Call get_mcp_tool_schema(connection="", tool="") for the one schema you need. + 4. Call call_mcp(connection="", tool="", arguments={...}) to run it, passing an arguments object that matches the schema (omit arguments for a tool that takes none). +- Do not assume a connection or tool exists; discover it with list_mcps and search_mcp_tools before calling. +- Use describe_mcp only as a compatibility fallback when targeted search cannot identify an expected tool. Its full catalog can be large. {% endif %} AUTHORIZATION STATUS: diff --git a/strix/core/runner.py b/strix/core/runner.py index b40a36b0e..42df1f00a 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -59,10 +59,8 @@ if TYPE_CHECKING: from strix.runtime.status import StatusSink from strix.tools.mcp import ( - ConnectedMcpServer, McpConnectionRequest, McpRegistry, - SupervisedMcpSession, ) @@ -71,8 +69,8 @@ logger = logging.getLogger(__name__) StreamEventSink = Callable[[str, Any], None] # Receives the run's MCP connection roster as a list of non-secret status dicts -# ({"name", "provider", "tool_count", "dead"}), once when the connections are -# established and again each time a connection transitions to dead. An interface +# ({"name", "provider", "tool_count", "dead", "state"}), once when the connections +# are registered and again on lifecycle transitions. An interface # can persist it, render it, or forward it on as connection status. Kept as a # snapshot of the whole roster (not a per- # connection delta) so every call carries a consistent, current picture. @@ -80,30 +78,21 @@ McpStatusSink = Callable[[list[dict[str, Any]]], None] def _mcp_roster_payload(registry: McpRegistry) -> list[dict[str, Any]]: - """The run's MCP roster as non-secret status dicts (name/provider/tool_count/dead).""" + """The run's MCP roster as non-secret lifecycle status dicts.""" return [ { "name": status.name, "provider": status.provider, "tool_count": status.tool_count, "dead": status.dead, + "state": status.state, } for status in registry.statuses() ] -def _mcp_startup_summary(connections: list[ConnectedMcpServer]) -> str: - """One user-facing line summarizing the MCP servers that connected.""" - server_count = len(connections) - tool_count = sum(c.tool_count for c in connections) - servers_word = "server" if server_count == 1 else "servers" - tools_word = "tool" if tool_count == 1 else "tools" - names = ", ".join(c.name for c in connections) - return f"MCP: connected {server_count} {servers_word} ({tool_count} {tools_word}): {names}" - - -def _record_mcp_connections(connections: list[ConnectedMcpServer]) -> None: - """Record which MCP servers this run connected, for the interfaces. +def _record_mcp_connections(connection_names: list[str]) -> None: + """Record which MCP servers this run configured, for the interfaces. A server's tools are offered to the model under a name built from the connection name and the tool's own name, which cannot be split back apart, so @@ -114,7 +103,7 @@ def _record_mcp_connections(connections: list[ConnectedMcpServer]) -> None: report_state = get_global_report_state() if report_state is None: return - report_state.record_mcp_connections([connection.name for connection in connections]) + report_state.record_mcp_connections(connection_names) def _note_exit_reason(reason: str) -> None: @@ -348,7 +337,7 @@ async def run_strix_scan( configure_spill_writer(_spill_to_workspace) sessions_to_close: list[SQLiteSession] = [] - mcp_sessions: list[SupervisedMcpSession] = [] + mcp_registry: McpRegistry | None = None try: targets = scan_config.get("targets") or [] @@ -400,7 +389,6 @@ async def run_strix_scan( from strix.tools.mcp import ( McpConnectionRequest, McpRegistry, - attach_mcp_requests, load_user_mcp_configs, ) @@ -416,56 +404,37 @@ async def run_strix_scan( else: mcp_requests = mcp_connection_requests if mcp_requests: - connections = await attach_mcp_requests(mcp_requests, mcp_registry) - mcp_sessions = [c.session 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)) - # Name the connected servers in the prompt so every agent - # (root and children, both deriving from scope_context) sees - # what is available at the start; they can still re-list or - # inspect them at run time via list_mcps / describe_mcp. Set - # only when a connection exists, so a run with no MCP leaves - # the prompt context unchanged. - scope_context["mcp_available"] = bool(mcp_registry) - scope_context["mcp_connections"] = [ - { - "name": summary.name, - "purpose": summary.purpose, - "tool_count": summary.tool_count, - } - for summary in mcp_registry.summaries() - ] + for request in mcp_requests: + mcp_registry.register(request) + _record_mcp_connections(mcp_registry.names()) + report( + f"MCP: configured {len(mcp_registry)} connection(s); " + "warming them in the background" + ) + scope_context["mcp_available"] = True + scope_context["mcp_connections"] = [ + { + "name": summary.name, + "purpose": summary.purpose, + "tool_count": summary.tool_count, + } + for summary in mcp_registry.summaries() + ] - # Feed a non-secret connection roster (name / provider / - # tool_count / dead) to two consumers: once now (all - # currently healthy) and again whenever a connection later - # dies. It is always persisted to run.json so the viewer, - # which re-reads the run's files from disk, can render the - # MCP connections panel and health without an in-memory - # sink. When an interface sink is attached (the TUI backend, - # or pro forwarding into the app's event stream) it also - # receives the same snapshot. In-use is derived separately by - # each interface from the connection-tagged tool-call events, - # so it is not carried here. - def _emit_mcp_status() -> None: - roster = _mcp_roster_payload(mcp_registry) - _persist_mcp_status(roster) - if mcp_status_sink is not None: - try: - mcp_status_sink(roster) - except Exception: - logger.exception("MCP status sink failed") + def _emit_mcp_status() -> None: + roster = _mcp_roster_payload(mcp_registry) + _persist_mcp_status(roster) + if mcp_status_sink is not None: + try: + mcp_status_sink(roster) + except Exception: + logger.exception("MCP status sink failed") - for connection_name in mcp_registry.names(): - entry = mcp_registry.get(connection_name) - if entry is not None: - entry.session.set_on_dead(_emit_mcp_status) - _emit_mcp_status() + mcp_registry.set_status_sink(_emit_mcp_status) + _emit_mcp_status() + mcp_registry.start_warmup(max_concurrency=6) except Exception: - logger.exception("Failed to connect user MCP servers; continuing without them") + logger.exception("Failed to configure user MCP servers; continuing without them") root_context = _merge_root_prompt_context(scope_context, extra_system_prompt_context) root_instructions = _compose_root_instructions_override( @@ -661,9 +630,9 @@ async def run_strix_scan( for s in sessions_to_close: with contextlib.suppress(Exception): s.close() - for mcp_session in mcp_sessions: + if mcp_registry is not None: with contextlib.suppress(Exception): - await mcp_session.aclose() + await mcp_registry.close() with contextlib.suppress(Exception): await coordinator._maybe_snapshot() if cleanup_on_exit: diff --git a/strix/tools/mcp/__init__.py b/strix/tools/mcp/__init__.py index eb1de2f89..a0de571a8 100644 --- a/strix/tools/mcp/__init__.py +++ b/strix/tools/mcp/__init__.py @@ -2,7 +2,13 @@ from __future__ import annotations -from strix.tools.mcp.agent_tools import call_mcp, describe_mcp, list_mcps +from strix.tools.mcp.agent_tools import ( + call_mcp, + describe_mcp, + get_mcp_tool_schema, + list_mcps, + search_mcp_tools, +) from strix.tools.mcp.client import ( ConnectedMcpServer, attach_mcp_requests, @@ -19,8 +25,10 @@ from strix.tools.mcp.naming import namespaced_tool_name from strix.tools.mcp.registry import ( CALL_MCP_TOOL, DESCRIBE_MCP_TOOL, + GET_MCP_TOOL_SCHEMA_TOOL, MCP_DISPATCH_TOOLS, MCP_REGISTRY_CONTEXT_KEY, + SEARCH_MCP_TOOLS_TOOL, McpCallInfo, McpConnectionEntry, McpConnectionRequest, @@ -35,8 +43,10 @@ from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcp __all__ = [ "CALL_MCP_TOOL", "DESCRIBE_MCP_TOOL", + "GET_MCP_TOOL_SCHEMA_TOOL", "MCP_DISPATCH_TOOLS", "MCP_REGISTRY_CONTEXT_KEY", + "SEARCH_MCP_TOOLS_TOOL", "BearerAuth", "ConnectedMcpServer", "FailureInfo", @@ -56,8 +66,10 @@ __all__ = [ "classify", "connect_mcp_servers", "describe_mcp", + "get_mcp_tool_schema", "list_mcps", "load_user_mcp_configs", "namespaced_tool_name", "resolve_mcp_call", + "search_mcp_tools", ] diff --git a/strix/tools/mcp/agent_tools.py b/strix/tools/mcp/agent_tools.py index 7f8679a67..ab068c644 100644 --- a/strix/tools/mcp/agent_tools.py +++ b/strix/tools/mcp/agent_tools.py @@ -1,18 +1,21 @@ -"""The three generic MCP dispatch tools every agent carries. +"""The generic MCP discovery and dispatch tools every agent carries. Under the generic-dispatch model an agent does not get one tool per MCP tool. -It gets exactly these three and discovers connections on demand: +It gets four primary tools and discovers connections on demand: - ``list_mcps()`` returns the connections available this run — each connection's id, name, description, and tool count, with no tool schemas — so the model can discover what it can reach without any inventory in the system prompt. -- ``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. +- ``search_mcp_tools(connection, query)`` returns a small ranked candidate set. +- ``get_mcp_tool_schema(connection, tool)`` returns one exact input schema. - ``call_mcp(connection, tool, arguments)`` dispatches one call to a connection's tool and returns its result. -All three read the per-run :class:`~strix.tools.mcp.registry.McpRegistry` from the +``describe_mcp`` remains as a compatibility path for older prompts, but current +agents use targeted search and one-schema lookup instead of receiving a whole +provider catalog in one model turn. + +All four 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 @@ -61,12 +64,11 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: """List the MCP connections available this run, so you can discover them. Read-only. Returns one entry per connection with its ``id`` (the exact name - you pass to ``describe_mcp`` and ``call_mcp``), ``name``, ``description``, and - ``tool_count`` — no tool schemas. The three MCP tools work in order: call - ``list_mcps`` to discover the available connections, then ``describe_mcp`` on - one connection to inspect its tools and their input schemas, then ``call_mcp`` - to run one of its tools. Returns an empty ``connections`` list when the run has - no MCP connections. + you pass to the other MCP tools), ``name``, ``description``, and + ``tool_count``. The response does not include tool schemas. Call + ``search_mcp_tools`` next. Then call ``get_mcp_tool_schema`` for one selected + tool before you call ``call_mcp``. Returns an empty ``connections`` list when + the run has no MCP connections. """ registry = _registry_from_ctx(ctx) if registry is None or not registry: @@ -80,6 +82,7 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: "description": summary.purpose, "tool_count": summary.tool_count, "dead": dead_by_name.get(summary.name, False), + "state": summary.state, } for summary in registry.summaries() ] @@ -88,13 +91,12 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: @function_tool(timeout=60) async def describe_mcp(ctx: RunContextWrapper, connection: str) -> str: - """List the tools one MCP connection offers, with their input schemas. + """Return one connection's full tool catalog as a compatibility fallback. - Read-only. Look up a connection by the id ``list_mcps`` reported for it; 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. + Read-only. This result can be large because it includes every tool schema. + First use ``search_mcp_tools`` and ``get_mcp_tool_schema``. Use this fallback + only when targeted search cannot identify an expected tool. Nothing is read + from the connected account. Args: connection: The connection name exactly as reported by ``list_mcps``. @@ -106,7 +108,7 @@ async def describe_mcp(ctx: RunContextWrapper, connection: str) -> str: if entry is None: return _unknown_connection(connection, registry) try: - tools = await entry.session.list_tools() + tools = await entry.ensure_catalog() except McpConnectionUnavailableError as exc: return str(exc) if not tools: @@ -116,6 +118,121 @@ async def describe_mcp(ctx: RunContextWrapper, connection: str) -> str: return f"{header}\n{body}" +def _search_score(tool: MCPTool, query_terms: list[str], active: bool) -> tuple[int, str]: + name = tool.name.lower() + description = (tool.description or "").lower() + if not query_terms: + return (100 if active else 0, name) + score = 100 if active else 0 + matched_terms = 0 + for term in query_terms: + if name == term: + score += 60 + matched_terms += 1 + elif name.startswith(term): + score += 35 + matched_terms += 1 + elif term in name: + score += 25 + matched_terms += 1 + elif term in description: + score += 10 + matched_terms += 1 + if matched_terms == 0: + return (-1, name) + return (score + (matched_terms * 5), name) + + +@function_tool(timeout=60) +async def search_mcp_tools( + ctx: RunContextWrapper, + connection: str, + query: str, + limit: int = 8, +) -> dict[str, Any] | str: + """Search one connection's tools without returning their full schemas. + + The first search lazily connects the provider and caches its catalog. Results + contain only names, short descriptions, and whether each tool is in the + scan's active set. Call ``get_mcp_tool_schema`` for the selected tool. + + Args: + connection: The connection name exactly as reported by ``list_mcps``. + query: Capability words such as ``page content`` or ``workspace identity``. + limit: Maximum candidates to return, from 1 through 20. + """ + 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) + try: + tools = await entry.ensure_catalog() + except McpConnectionUnavailableError as exc: + return str(exc) + terms = [term for term in query.lower().split() if term] + bounded_limit = min(max(limit, 1), 20) + ranked = [ + (score, tool) + for tool in tools + if (score := _search_score(tool, terms, tool.name in entry.active_tools))[0] >= 0 + ] + if not ranked: + ranked = [ + ((100, tool.name.lower()), tool) for tool in tools if tool.name in entry.active_tools + ] + ranked.sort(key=lambda item: (-item[0][0], item[0][1])) + return { + "connection": connection, + "query": query, + "matches": [ + { + "name": tool.name, + "description": (tool.description or "").strip() or None, + "active": tool.name in entry.active_tools, + } + for _, tool in ranked[:bounded_limit] + ], + } + + +@function_tool(timeout=60) +async def get_mcp_tool_schema( + ctx: RunContextWrapper, + connection: str, + tool: str, +) -> dict[str, Any] | str: + """Return one MCP tool's exact input schema. + + Args: + connection: The connection name exactly as reported by ``list_mcps``. + tool: One exact tool name returned by ``search_mcp_tools``. + """ + 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) + try: + tools = await entry.ensure_catalog() + except McpConnectionUnavailableError as exc: + return str(exc) + match = next((candidate for candidate in tools if candidate.name == tool), None) + if match is None: + return ( + f"Unknown tool {tool!r} on MCP connection {connection!r}. " + "Call search_mcp_tools to find an available tool." + ) + return { + "connection": connection, + "name": match.name, + "description": (match.description or "").strip() or None, + "input_schema": match.inputSchema or {"type": "object"}, + } + + @function_tool(timeout=120, strict_mode=False) async def call_mcp( ctx: RunContextWrapper, @@ -125,19 +242,18 @@ async def call_mcp( ) -> Any: """Call one tool on one MCP connection and return its result. - Address the tool by the connection id from ``list_mcps`` 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). + Use the connection id from ``list_mcps``. Use the tool name from + ``search_mcp_tools``. Pass an object that matches the schema from + ``get_mcp_tool_schema``. Omit the arguments for a tool that takes no + arguments. Args: connection: The connection name exactly as reported by ``list_mcps``. - tool: The tool name, exactly as reported by ``describe_mcp``. + tool: The tool name exactly as reported by ``search_mcp_tools``. arguments: The tool's arguments as a JSON object of names to values (for example ``{"path": "app.py"}``), or omitted/empty for a tool that - takes none. Pass an object, not a stringified one. Its shape is - whatever ``describe_mcp`` showed for the tool rather than a shape this - tool fixes in advance. + takes none. Pass an object, not a stringified object. Match the shape + that ``get_mcp_tool_schema`` returned. """ registry = _registry_from_ctx(ctx) if registry is None or not registry: @@ -147,7 +263,7 @@ async def call_mcp( return _unknown_connection(connection, registry) invalid_arguments = ( 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." + "argument names to values, or none. Call get_mcp_tool_schema for the schema." ) if isinstance(arguments, str): # The ``arguments`` parameter is schema-less (an open object is not @@ -162,18 +278,17 @@ async def call_mcp( if arguments is not None and not isinstance(arguments, dict): return invalid_arguments try: - available = await entry.session.list_tools() + available = await entry.ensure_catalog() except McpConnectionUnavailableError as exc: return _errored_tool_output(str(exc)) 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." + "Call search_mcp_tools, then get_mcp_tool_schema." ) - return await entry.session.dispatch( + session = await entry.ensure_connected() + return await session.dispatch( tool, arguments or {}, label=namespaced_tool_name(connection, tool), diff --git a/strix/tools/mcp/client.py b/strix/tools/mcp/client.py index 7335ffd43..c0de7fab5 100644 --- a/strix/tools/mcp/client.py +++ b/strix/tools/mcp/client.py @@ -14,6 +14,7 @@ fails the run. from __future__ import annotations +import asyncio import contextlib import json import logging @@ -281,8 +282,10 @@ async def _count_session_tools(config: McpConnectionConfig, session: SupervisedM async def connect_mcp_servers( configs: list[McpConnectionConfig], + *, + max_concurrency: int = 6, ) -> list[ConnectedMcpServer]: - """Connect each MCP config on its own supervising task and return the sessions. + """Connect MCP configs concurrently under a fixed bound. Each connection becomes a :class:`~strix.tools.mcp.session.SupervisedMcpSession` that owns ``connect()``, the held-open session, and ``cleanup()`` on one @@ -301,41 +304,43 @@ async def connect_mcp_servers( :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] = [] + semaphore = asyncio.Semaphore(max(1, max_concurrency)) sessions: list[SupervisedMcpSession] = [] - try: - for config in configs: + + async def connect_one(config: McpConnectionConfig) -> ConnectedMcpServer | None: + async with semaphore: session = SupervisedMcpSession(config) sessions.append(session) - if not await session.start(): - # Initial connect failed; already logged inside the session. Drop it. - await session.aclose() - sessions.remove(session) - continue try: + if not await session.start(): + await session.aclose() + return None tool_count = await _count_session_tools(config, session) except McpConnectionUnavailableError: - # The session died between connecting and its first listing; skip it. logger.warning("MCP connection %r died before its first listing", config.name) await session.aclose() - sessions.remove(session) - continue + return None + except BaseException: + with contextlib.suppress(BaseException): + await session.aclose() + raise logger.info("Connected MCP server %r (%d tools)", config.name, tool_count) - connected.append( - ConnectedMcpServer( - session=session, name=config.name, tool_count=tool_count, notes=config.notes - ) + return ConnectedMcpServer( + session=session, + name=config.name, + tool_count=tool_count, + notes=config.notes, ) - except BaseException: - # Cancelled or errored mid-attach: close every session started so far, - # each on its own task, then re-raise. The runner only receives the list - # on a clean return, so on an abnormal exit this function owns the cleanup. - for session in sessions: - with contextlib.suppress(BaseException): - await session.aclose() - raise - return connected + try: + results = await asyncio.gather(*(connect_one(config) for config in configs)) + except BaseException: + await asyncio.gather( + *(session.aclose() for session in sessions), + return_exceptions=True, + ) + raise + return [result for result in results if result is not None] async def attach_mcp_requests( diff --git a/strix/tools/mcp/config.py b/strix/tools/mcp/config.py index 795a58a32..a5caa170b 100644 --- a/strix/tools/mcp/config.py +++ b/strix/tools/mcp/config.py @@ -62,6 +62,9 @@ class McpConnectionConfig(BaseModel): """Tool allowlist, applied after the server lists its tools. ``None`` (the default) exposes every tool the server lists; a list restricts to it.""" + active_tools: list[str] = Field(default_factory=list) + """Small scan-relevant subset ranked ahead of the broader allowed catalog.""" + notes: str | None = None """Free-text notes for the agent describing what this connection is and how to use it. When set, the note becomes the connection's purpose line in the diff --git a/strix/tools/mcp/registry.py b/strix/tools/mcp/registry.py index 8533807b8..65652ab10 100644 --- a/strix/tools/mcp/registry.py +++ b/strix/tools/mcp/registry.py @@ -22,13 +22,18 @@ single dispatch point. from __future__ import annotations +import asyncio +import contextlib import dataclasses -from typing import TYPE_CHECKING, Any, NamedTuple +import time +from typing import TYPE_CHECKING, Any, Literal, NamedTuple -from strix.tools.mcp.session import SupervisedMcpSession +from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcpSession if TYPE_CHECKING: + from collections.abc import Callable + from agents.mcp import MCPServer from strix.tools.mcp.client import ResultTransform @@ -49,37 +54,60 @@ MCP_REGISTRY_CONTEXT_KEY = "mcp_registry" # tracer all recognise a connection-scoped 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}) +SEARCH_MCP_TOOLS_TOOL = "search_mcp_tools" +GET_MCP_TOOL_SCHEMA_TOOL = "get_mcp_tool_schema" +MCP_DISPATCH_TOOLS = frozenset( + { + CALL_MCP_TOOL, + DESCRIBE_MCP_TOOL, + SEARCH_MCP_TOOLS_TOOL, + GET_MCP_TOOL_SCHEMA_TOOL, + } +) + +McpConnectionState = Literal[ + "configured", + "connecting", + "connected", + "catalog_loading", + "catalog_ready", + "unavailable", +] +_RETRY_DELAY_SECONDS = 5.0 -@dataclasses.dataclass(frozen=True) +@dataclasses.dataclass class McpConnectionEntry: - """One live MCP connection a scan may reach, keyed by ``name``. + """One configured MCP connection a scan may reach, keyed by ``name``. - ``session`` is the :class:`~strix.tools.mcp.session.SupervisedMcpSession` that - owns the connection on its own task; the dispatch tools list tools and call - tools through it (``session.list_tools`` / ``session.dispatch``) so a session - failure is contained and can reconnect. ``purpose`` is the human label - ``list_mcps`` reports as the connection's description (the user's connection - notes, or whatever the caller supplies). ``tool_count`` is how many tools the - connection offers, also reported by ``list_mcps``. ``result_transform``, when - set, runs on each call's structured result 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. - - The connection config the session reconnects with (and its bearer token) lives - on ``session`` in memory only. It is reached via :attr:`config` for the - reconnect path and is never logged, serialized into the event stream, or - written to disk. + Registration is inert. The first warm-up, search, schema lookup, or call + creates one shared connection task; the first catalog operation creates one + shared listing task. Root and child agents therefore reuse the same session + and catalog even when they request a cold connection concurrently. """ - session: SupervisedMcpSession name: str + connection_config: McpConnectionConfig | None = dataclasses.field( + default=None, + repr=False, + ) + session: SupervisedMcpSession | None = dataclasses.field(default=None, repr=False) purpose: str | None = None tool_count: int = 0 result_transform: ResultTransform | None = None provider: str | None = None + state: McpConnectionState = "configured" + _catalog: list[Any] | None = dataclasses.field(default=None, repr=False) + _connect_task: asyncio.Task[SupervisedMcpSession] | None = dataclasses.field( + default=None, + repr=False, + ) + _catalog_task: asyncio.Task[list[Any]] | None = dataclasses.field( + default=None, + repr=False, + ) + _retry_after: float = dataclasses.field(default=0.0, repr=False) + _status_sink: Callable[[], None] | None = dataclasses.field(default=None, repr=False) @property def server(self) -> MCPServer | None: @@ -88,12 +116,143 @@ class McpConnectionEntry: Kept so existing callers that read ``entry.server`` keep working; new code should call through ``entry.session`` so reconnect and containment apply. """ - return self.session.server + return self.session.server if self.session is not None else None @property def config(self) -> McpConnectionConfig | None: - """The session's reconnect config. Carries the bearer token; never log it.""" - return self.session.config + """The reconnect config. Carries the bearer token; never log it.""" + if self.connection_config is not None: + return self.connection_config + return self.session.config if self.session is not None else None + + @property + def active_tools(self) -> frozenset[str]: + config = self.config + return frozenset(config.active_tools if config is not None else ()) + + def set_status_sink(self, sink: Callable[[], None] | None) -> None: + self._status_sink = sink + if self.session is not None: + self.session.set_on_dead(self._on_dead) + + def _set_state(self, state: McpConnectionState) -> None: + if self.state == state: + return + self.state = state + if self._status_sink is not None: + self._status_sink() + + def _on_dead(self) -> None: + self.session = None + self._catalog = None + self._catalog_task = None + self._retry_after = time.monotonic() + _RETRY_DELAY_SECONDS + self._set_state("unavailable") + + async def ensure_connected(self) -> SupervisedMcpSession: + """Return this entry's live session, connecting it once when needed.""" + if self.session is not None: + return self.session + if self.connection_config is None: + raise McpConnectionUnavailableError( + f"MCP connection {self.name!r} is unavailable and cannot reconnect." + ) + if self._connect_task is None: + if self.state == "unavailable" and time.monotonic() < self._retry_after: + raise McpConnectionUnavailableError( + f"MCP connection {self.name!r} is temporarily unavailable." + ) + self._connect_task = asyncio.create_task( + self._connect(), + name=f"mcp-connect-{self.name}", + ) + task = self._connect_task + try: + return await asyncio.shield(task) + finally: + if self._connect_task is task and task.done(): + self._connect_task = None + + async def _connect(self) -> SupervisedMcpSession: + config = self.connection_config + if config is None: + raise McpConnectionUnavailableError( + f"MCP connection {self.name!r} has no connection configuration." + ) + self._set_state("connecting") + session = SupervisedMcpSession(config) + try: + started = await session.start() + except BaseException: + with contextlib.suppress(BaseException): + await session.aclose() + self._retry_after = time.monotonic() + _RETRY_DELAY_SECONDS + self._set_state("unavailable") + raise + if not started: + await session.aclose() + self._retry_after = time.monotonic() + _RETRY_DELAY_SECONDS + self._set_state("unavailable") + raise McpConnectionUnavailableError(f"MCP connection {self.name!r} could not connect.") + self.session = session + session.set_on_dead(self._on_dead) + self._retry_after = 0.0 + self._set_state("connected") + return session + + async def ensure_catalog(self) -> list[Any]: + """Return the filtered catalog, listing it once on first use.""" + if ( + self._catalog is not None + and self.session is not None + and not self.session.is_dead + and not self.session.is_unavailable + ): + return self._catalog + if self.session is not None and (self.session.is_dead or self.session.is_unavailable): + self._catalog = None + if self._catalog_task is None: + self._catalog_task = asyncio.create_task( + self._load_catalog(), + name=f"mcp-catalog-{self.name}", + ) + task = self._catalog_task + try: + return await asyncio.shield(task) + finally: + if self._catalog_task is task and task.done(): + self._catalog_task = None + + async def _load_catalog(self) -> list[Any]: + session = await self.ensure_connected() + self._set_state("catalog_loading") + try: + catalog = await session.list_tools() + except BaseException: + if session.is_dead: + self._set_state("unavailable") + else: + self._set_state("connected") + raise + self._catalog = list(catalog) + self.tool_count = len(self._catalog) + self._set_state("catalog_ready") + return self._catalog + + async def close(self) -> None: + """Cancel pending initialization and close an opened session.""" + tasks = [task for task in (self._catalog_task, self._connect_task) if task is not None] + for task in tasks: + task.cancel() + for task in tasks: + with contextlib.suppress(BaseException): + await task + self._catalog_task = None + self._connect_task = None + if self.session is not None: + with contextlib.suppress(BaseException): + await self.session.aclose() + self.session = None @dataclasses.dataclass(frozen=True) @@ -105,6 +264,7 @@ class McpConnectionSummary: purpose: str | None tool_count: int provider: str | None = None + state: McpConnectionState = "configured" @dataclasses.dataclass(frozen=True) @@ -123,6 +283,7 @@ class McpConnectionStatus: provider: str | None tool_count: int dead: bool + state: McpConnectionState @dataclasses.dataclass(frozen=True) @@ -165,6 +326,21 @@ class McpRegistry: def __init__(self) -> None: self._entries: dict[str, McpConnectionEntry] = {} + self._status_sink: Callable[[], None] | None = None + self._warmup_task: asyncio.Task[None] | None = None + + def register(self, request: McpConnectionRequest) -> McpConnectionEntry: + """Register an inert request without opening a network connection.""" + entry = McpConnectionEntry( + name=request.config.name, + connection_config=request.config, + purpose=request.purpose or request.config.notes, + result_transform=request.result_transform, + provider=request.provider, + ) + entry.set_status_sink(self._status_sink) + self._entries[entry.name] = entry + return entry def add( self, @@ -191,13 +367,16 @@ class McpRegistry: raise ValueError("McpRegistry.add requires either 'session' or 'server'") session = SupervisedMcpSession.adopt(server, name=name, config=config) entry = McpConnectionEntry( - session=session, name=name, + connection_config=config or session.config, + session=session, purpose=purpose, tool_count=tool_count, result_transform=result_transform, provider=provider, + state="connected", ) + entry.set_status_sink(self._status_sink) self._entries[name] = entry return entry @@ -217,6 +396,7 @@ class McpRegistry: purpose=entry.purpose, tool_count=entry.tool_count, provider=entry.provider, + state=entry.state, ) for entry in self._entries.values() ] @@ -233,7 +413,9 @@ class McpRegistry: name=entry.name, provider=entry.provider, tool_count=entry.tool_count, - dead=entry.session.is_dead, + dead=entry.state == "unavailable" + or (entry.session is not None and entry.session.is_dead), + state=entry.state, ) for entry in self._entries.values() ] @@ -243,6 +425,42 @@ class McpRegistry: runner).""" self._entries.clear() + def set_status_sink(self, sink: Callable[[], None] | None) -> None: + """Receive a callback after any connection lifecycle transition.""" + self._status_sink = sink + for entry in self._entries.values(): + entry.set_status_sink(sink) + + def start_warmup(self, *, max_concurrency: int = 6) -> asyncio.Task[None]: + """Connect every configured entry in the background with a fixed bound.""" + if self._warmup_task is not None: + return self._warmup_task + + async def warm() -> None: + semaphore = asyncio.Semaphore(max(1, max_concurrency)) + + async def connect(entry: McpConnectionEntry) -> None: + async with semaphore: + with contextlib.suppress(McpConnectionUnavailableError): + await entry.ensure_connected() + + await asyncio.gather(*(connect(entry) for entry in self._entries.values())) + + self._warmup_task = asyncio.create_task(warm(), name="mcp-warmup") + return self._warmup_task + + async def close(self) -> None: + """Stop warm-up and close only sessions this registry opened.""" + if self._warmup_task is not None: + self._warmup_task.cancel() + with contextlib.suppress(BaseException): + await self._warmup_task + self._warmup_task = None + await asyncio.gather( + *(entry.close() for entry in self._entries.values()), + return_exceptions=True, + ) + def __len__(self) -> int: return len(self._entries) @@ -282,6 +500,6 @@ def resolve_mcp_call( if entry is None: return None provider = entry.provider - raw_tool = args.get("tool") if tool_name == CALL_MCP_TOOL else "" + raw_tool = args.get("tool") if tool_name in {CALL_MCP_TOOL, GET_MCP_TOOL_SCHEMA_TOOL} else "" tool = raw_tool if isinstance(raw_tool, str) else "" return McpCallInfo(connection=connection, tool=tool, provider=provider) diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 82161ad5c..54f49f54d 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -36,12 +36,15 @@ from strix.tools.mcp import ( attach_mcp_requests, call_mcp, describe_mcp, + get_mcp_tool_schema, list_mcps, load_user_mcp_configs, namespaced_tool_name, resolve_mcp_call, + search_mcp_tools, ) from strix.tools.mcp import client as mcp_client +from strix.tools.mcp import registry as mcp_registry_mod from strix.tools.mcp import session as mcp_session_mod @@ -434,8 +437,16 @@ async def test_list_mcps_returns_connections_with_ids_and_descriptions() -> None "description": "local files", "tool_count": 2, "dead": False, + "state": "connected", + }, + { + "id": "db", + "name": "db", + "description": None, + "tool_count": 1, + "dead": False, + "state": "connected", }, - {"id": "db", "name": "db", "description": None, "tool_count": 1, "dead": False}, ] } @@ -485,6 +496,247 @@ async def test_describe_mcp_without_any_connections() -> None: assert out == "No MCP connections are configured for this run." +@pytest.mark.asyncio +async def test_search_then_get_one_schema_without_returning_other_schemas() -> None: + registry = McpRegistry() + registry.add( + name="docs", + server=FakeMCPServer( + "docs", + [ + _mcp_tool("find_pages"), + MCPTool( + name="fetch_page", + description="Fetch page content", + inputSchema={ + "type": "object", + "properties": {"page_id": {"type": "string"}}, + }, + ), + ], + ), + config=McpConnectionConfig( + name="docs", + url="https://example.invalid/mcp", + active_tools=["fetch_page"], + ), + ) + + searched = await search_mcp_tools.on_invoke_tool( + _ctx(registry), + json.dumps({"connection": "docs", "query": "page content"}), + ) + assert [match["name"] for match in searched["matches"]] == [ + "fetch_page", + "find_pages", + ] + assert all("input_schema" not in match for match in searched["matches"]) + + schema = await get_mcp_tool_schema.on_invoke_tool( + _ctx(registry), + json.dumps({"connection": "docs", "tool": "fetch_page"}), + ) + assert schema["input_schema"]["properties"] == {"page_id": {"type": "string"}} + + +@pytest.mark.asyncio +async def test_search_matches_any_query_term_and_prioritizes_active_tools() -> None: + registry = McpRegistry() + registry.add( + name="docs", + server=FakeMCPServer( + "docs", + [ + MCPTool( + name="fetch_page", + description="Read page content", + inputSchema={"type": "object"}, + ), + MCPTool( + name="search_documents", + description="Search documents", + inputSchema={"type": "object"}, + ), + MCPTool( + name="list_records", + description="List records", + inputSchema={"type": "object"}, + ), + ], + ), + config=McpConnectionConfig( + name="docs", + url="https://example.invalid/mcp", + active_tools=["fetch_page"], + ), + ) + + searched = await search_mcp_tools.on_invoke_tool( + _ctx(registry), + json.dumps( + { + "connection": "docs", + "query": "read list search document", + } + ), + ) + + assert [match["name"] for match in searched["matches"]] == [ + "fetch_page", + "search_documents", + "list_records", + ] + + +@pytest.mark.asyncio +async def test_search_returns_active_tools_when_text_does_not_match() -> None: + registry = McpRegistry() + registry.add( + name="docs", + server=FakeMCPServer( + "docs", + [ + _mcp_tool("fetch_page"), + _mcp_tool("search_documents"), + ], + ), + config=McpConnectionConfig( + name="docs", + url="https://example.invalid/mcp", + active_tools=["fetch_page"], + ), + ) + + searched = await search_mcp_tools.on_invoke_tool( + _ctx(registry), + json.dumps( + { + "connection": "docs", + "query": "unrelated capability", + } + ), + ) + + assert searched["matches"] == [ + { + "name": "fetch_page", + "description": "remote tool fetch_page", + "active": True, + } + ] + + +@pytest.mark.asyncio +async def test_registered_entry_connects_and_lists_once_for_concurrent_catalog_requests( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls = {"build": 0, "list": 0} + + class _CountingServer(FakeMCPServer): + async def list_tools( + self, + run_context: Any = None, + agent: Any = None, + ) -> list[MCPTool]: + calls["list"] += 1 + await asyncio.sleep(0) + return await super().list_tools(run_context, agent) + + server = _CountingServer("docs", [_mcp_tool("fetch_page")]) + + def _build(_config: McpConnectionConfig) -> mcp_client.BuiltMcpServer: + calls["build"] += 1 + return _built_server(server) + + monkeypatch.setattr(mcp_client, "_build_server", _build) + registry = McpRegistry() + entry = registry.register(McpConnectionRequest(config=_config("docs", ["fetch_page"]))) + assert entry.state == "configured" + assert calls == {"build": 0, "list": 0} + + first, second = await asyncio.gather(entry.ensure_catalog(), entry.ensure_catalog()) + + assert [tool.name for tool in first] == ["fetch_page"] + assert second is first + assert calls == {"build": 1, "list": 1} + assert registry.statuses()[0].state == "catalog_ready" + await registry.close() + + +@pytest.mark.asyncio +async def test_registered_entry_replaces_a_terminally_dead_session( + monkeypatch: pytest.MonkeyPatch, +) -> None: + config = _config("docs", ["fetch_page"]) + original = SupervisedMcpSession.adopt( + FakeMCPServer("docs", [_mcp_tool("fetch_page")]), + name="docs", + config=config, + ) + replacement_server = FakeMCPServer("docs", [_mcp_tool("fetch_page")]) + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda _config: _built_server(replacement_server), + ) + monkeypatch.setattr(mcp_registry_mod, "_RETRY_DELAY_SECONDS", 0) + registry = McpRegistry() + entry = registry.add(name="docs", session=original, config=config) + + original._mark_dead() + + assert entry.session is None + replacement = await entry.ensure_connected() + assert replacement is not original + assert replacement.server is replacement_server + assert entry.state == "connected" + await registry.close() + + +@pytest.mark.asyncio +async def test_registry_warmup_bounds_parallel_connections( + monkeypatch: pytest.MonkeyPatch, +) -> None: + active = 0 + maximum = 0 + two_started = asyncio.Event() + release = asyncio.Event() + + class _BlockedConnectServer(FakeMCPServer): + async def connect(self) -> None: + nonlocal active, maximum + active += 1 + maximum = max(maximum, active) + if active == 2: + two_started.set() + try: + await release.wait() + finally: + active -= 1 + + servers = { + f"docs-{index}": _BlockedConnectServer(f"docs-{index}", [_mcp_tool("fetch_page")]) + for index in range(4) + } + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda config: _built_server(servers[config.name]), + ) + registry = McpRegistry() + for name in servers: + registry.register(McpConnectionRequest(config=_config(name, ["fetch_page"]))) + + warmup = registry.start_warmup(max_concurrency=2) + await asyncio.wait_for(two_started.wait(), timeout=1) + assert maximum == 2 + release.set() + await warmup + + assert [summary.state for summary in registry.summaries()] == ["connected"] * 4 + await registry.close() + + # --- call_mcp ---------------------------------------------------------------- @@ -573,7 +825,8 @@ async def test_call_mcp_errors_on_unknown_tool() -> None: ) assert "Unknown tool 'delete_everything'" in out - assert "read_file" in out + assert "search_mcp_tools" in out + assert "read_file" not in out # A rejected tool name never reaches the server. assert server.calls == [] @@ -631,21 +884,27 @@ async def test_call_mcp_flags_an_errored_result_failed_for_the_tui() -> None: assert out == {"type": "text", "text": "boom:read_file", "success": False} -# --- the two tools are the only MCP surface every agent gets ----------------- +# --- generic MCP tools are the only MCP surface every agent gets ------------- -def test_agent_carries_exactly_the_dispatch_tools_regardless_of_connections() -> None: +def test_agent_carries_exactly_the_generic_mcp_tools_regardless_of_connections() -> None: """No matter how many MCP connections a run makes, an agent's tool list gains - exactly list_mcps, describe_mcp, and call_mcp and never a per-connection - provider tool.""" + exactly the generic MCP tools 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 {"list_mcps", "describe_mcp", "call_mcp"} <= set(root_names) - assert {"list_mcps", "describe_mcp", "call_mcp"} <= set(child_names) + expected = { + "list_mcps", + "search_mcp_tools", + "get_mcp_tool_schema", + "describe_mcp", + "call_mcp", + } + assert expected <= set(root_names) + assert expected <= set(child_names) # Five hypothetical connections would once have added ~all their tools as # namespaced provider tools; none of those names may appear now. @@ -666,16 +925,25 @@ def test_agent_carries_exactly_the_dispatch_tools_regardless_of_connections() -> # --- prompt guidance replaces the old per-connection inventory --------------- -def test_prompt_renders_static_three_tool_guidance_when_mcp_available() -> None: +def test_prompt_renders_targeted_tool_guidance_when_mcp_available() -> None: prompt = render_system_prompt(system_prompt_context={"mcp_available": True}) assert "MCP CONNECTIONS" in prompt - # The three discovery/dispatch tools are named as the way in. + # Targeted discovery, one-schema lookup, dispatch, and compatibility are named. assert "list_mcps" in prompt + assert "search_mcp_tools" in prompt + assert "get_mcp_tool_schema" in prompt assert "describe_mcp" in prompt assert "call_mcp" in prompt +def test_model_facing_mcp_guidance_uses_targeted_discovery() -> None: + assert "describe_mcp" not in list_mcps.description + assert "describe_mcp" not in call_mcp.description + assert "full tool catalog" in describe_mcp.description + assert "compatibility fallback" in describe_mcp.description + + def test_prompt_has_no_mcp_section_without_availability() -> None: assert "MCP CONNECTIONS" not in render_system_prompt(system_prompt_context={}) @@ -683,7 +951,7 @@ def test_prompt_has_no_mcp_section_without_availability() -> None: def test_prompt_renders_named_connection_inventory() -> None: """With mcp_available set, the prompt names each connected server (name, tool count, purpose) so every agent sees what is available at the start, alongside - the three dispatch tools for re-listing and inspecting them at run time.""" + the discovery and dispatch tools for re-listing and inspecting them at run time.""" prompt = render_system_prompt( system_prompt_context={ "mcp_available": True, diff --git a/tests/test_runner_mcp.py b/tests/test_runner_mcp.py index 3b5a79e90..25a5c888d 100644 --- a/tests/test_runner_mcp.py +++ b/tests/test_runner_mcp.py @@ -80,13 +80,14 @@ async def test_none_default_attaches_from_the_user_config_file( ) monkeypatch.setattr(mcp_pkg, "load_user_mcp_configs", lambda: [file_config]) - captured: list[list[McpConnectionRequest]] = [] + captured: list[McpConnectionRequest] = [] + original_register = mcp_pkg.McpRegistry.register - async def _capture(requests: list[McpConnectionRequest], _registry: Any) -> list[Any]: - captured.append(requests) - return [] + def _capture(registry: Any, request: McpConnectionRequest) -> Any: + captured.append(request) + return original_register(registry, request) - monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _capture) + monkeypatch.setattr(mcp_pkg.McpRegistry, "register", _capture) await runner.run_strix_scan( scan_config={"targets": [], "scan_mode": "deep"}, @@ -97,12 +98,11 @@ async def test_none_default_attaches_from_the_user_config_file( # 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 + (request,) = captured + assert request.config is file_config + assert request.provider is None + assert request.result_transform is None + assert request.purpose is None @pytest.mark.asyncio @@ -116,13 +116,14 @@ async def test_supplied_requests_are_attached_and_the_user_file_is_not_read( monkeypatch.setattr(mcp_pkg, "load_user_mcp_configs", _fail_if_read) - captured: list[list[McpConnectionRequest]] = [] + captured: list[McpConnectionRequest] = [] + original_register = mcp_pkg.McpRegistry.register - async def _capture(requests: list[McpConnectionRequest], _registry: Any) -> list[Any]: - captured.append(requests) - return [] + def _capture(registry: Any, request: McpConnectionRequest) -> Any: + captured.append(request) + return original_register(registry, request) - monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _capture) + monkeypatch.setattr(mcp_pkg.McpRegistry, "register", _capture) supplied = [ McpConnectionRequest( @@ -139,7 +140,7 @@ async def test_supplied_requests_are_attached_and_the_user_file_is_not_read( mcp_connection_requests=supplied, ) - assert captured == [supplied] + assert captured == supplied @pytest.mark.asyncio @@ -147,8 +148,8 @@ async def test_roster_is_persisted_even_without_a_status_sink( monkeypatch: pytest.MonkeyPatch, tmp_path: Any ) -> None: """The viewer reads the roster off disk, so persistence must not depend on the - interface status sink: with ``mcp_status_sink=None`` the connect-time roster is - still written, carrying only the non-secret name/provider/tool_count/dead.""" + interface status sink: with ``mcp_status_sink=None`` the configured roster is + still written, carrying only non-secret lifecycle fields.""" _wire_runner(monkeypatch, tmp_path) monkeypatch.setattr( mcp_pkg, @@ -156,19 +157,6 @@ async def test_roster_is_persisted_even_without_a_status_sink( lambda: [McpConnectionConfig(name="local_fs", transport="stdio", command="npx")], ) - class _FakeSession: - is_dead = False - - def set_on_dead(self, _callback: Any) -> None: - return None - - async def _attach(_requests: list[McpConnectionRequest], registry: Any) -> list[Any]: - registry.add(name="local_fs", session=_FakeSession(), tool_count=3, provider=None) - entry = registry.get("local_fs") - return [types.SimpleNamespace(name="local_fs", tool_count=3, session=entry.session)] - - monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _attach) - persisted: list[list[dict[str, Any]]] = [] def _capture_persist(roster: list[dict[str, Any]]) -> None: @@ -185,4 +173,12 @@ async def test_roster_is_persisted_even_without_a_status_sink( ) assert persisted, "roster must persist even when no status sink is attached" - assert persisted[-1] == [{"name": "local_fs", "provider": None, "tool_count": 3, "dead": False}] + assert persisted[0] == [ + { + "name": "local_fs", + "provider": None, + "tool_count": 0, + "dead": False, + "state": "configured", + } + ] diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index b85cf56aa..1b8573055 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -194,21 +194,12 @@ async def test_mcp_available_flag_set_when_a_connection_attaches( scope_context: dict[str, Any] = {"scope": "built-in"} captured = _patch_engine_scaffold(monkeypatch, tmp_path, scope_context) - async def _aclose() -> None: - return None - - async def _attach(_requests: Any, registry: Any) -> list[Any]: - registry.add(name="fs", server=object(), purpose="local files", tool_count=2) - session = types.SimpleNamespace(aclose=_aclose) - return [types.SimpleNamespace(name="fs", tool_count=2, session=session)] - - monkeypatch.setattr(mcp_pkg, "attach_mcp_requests", _attach) - request = McpConnectionRequest( config=McpConnectionConfig( name="fs", url="https://mcp.example.com", auth=BearerAuth(token="run-token"), + notes="local files", ) ) @@ -224,7 +215,7 @@ async def test_mcp_available_flag_set_when_a_connection_attaches( assert kwargs["system_prompt_context"]["mcp_available"] is True # The named inventory names each connected server for the prompt. assert kwargs["system_prompt_context"]["mcp_connections"] == [ - {"name": "fs", "purpose": "local files", "tool_count": 2} + {"name": "fs", "purpose": "local files", "tool_count": 0} ]