diff --git a/strix/agents/prompts/scope.jinja b/strix/agents/prompts/scope.jinja index 4211fd53..c9dc5c72 100644 --- a/strix/agents/prompts/scope.jinja +++ b/strix/agents/prompts/scope.jinja @@ -21,7 +21,7 @@ MCP CONNECTIONS (available this run): {% if system_prompt_context.mcp_connections %} - 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 %} + - {{ connection.name }}{% if connection.state == "unavailable" %} (unavailable right now){% elif connection.tool_count is not none %} ({{ connection.tool_count }} tools){% endif %}{% 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. diff --git a/strix/core/runner.py b/strix/core/runner.py index 048a59de..58124be2 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -68,6 +68,8 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_MCP_PROMPT_WARMUP_TIMEOUT_SECONDS = 30.0 + StreamEventSink = Callable[[str, Any], None] # Receives the run's MCP connection roster as a list of non-secret status dicts @@ -93,6 +95,19 @@ def _mcp_roster_payload(registry: McpRegistry) -> list[dict[str, Any]]: ] +def _mcp_prompt_roster(registry: McpRegistry) -> list[dict[str, Any]]: + """The MCP roster rendered in agent prompts, with only verified counts.""" + return [ + { + "name": summary.name, + "purpose": summary.purpose, + "tool_count": summary.tool_count if summary.state == "catalog_ready" else None, + "state": summary.state, + } + for summary in registry.summaries() + ] + + def _record_mcp_connections(connection_names: list[str]) -> None: """Record which MCP servers this run configured, for the interfaces. @@ -438,19 +453,13 @@ async def run_strix_scan( _record_mcp_connections(mcp_registry.names()) report( f"MCP: configured {len(mcp_registry)} connection(s); " - "warming them in the background" + "connecting and listing tools" ) 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() - ] + scope_context["mcp_connections"] = _mcp_prompt_roster(mcp_registry) def _emit_mcp_status() -> None: + scope_context["mcp_connections"] = _mcp_prompt_roster(mcp_registry) roster = _mcp_roster_payload(mcp_registry) _persist_mcp_status(roster) if mcp_status_sink is not None: @@ -461,7 +470,11 @@ async def run_strix_scan( mcp_registry.set_status_sink(_emit_mcp_status) _emit_mcp_status() - mcp_registry.start_warmup(max_concurrency=6) + warmup_task = mcp_registry.start_warmup(max_concurrency=6) + await asyncio.wait( + {warmup_task}, + timeout=_MCP_PROMPT_WARMUP_TIMEOUT_SECONDS, + ) except Exception: logger.exception("Failed to configure user MCP servers; continuing without them") diff --git a/strix/tools/mcp/agent_tools.py b/strix/tools/mcp/agent_tools.py index 7fa2c392..1033d334 100644 --- a/strix/tools/mcp/agent_tools.py +++ b/strix/tools/mcp/agent_tools.py @@ -24,6 +24,7 @@ automatically. from __future__ import annotations +import asyncio import json from typing import TYPE_CHECKING, Any @@ -46,6 +47,7 @@ def _registry_from_ctx(ctx: RunContextWrapper) -> McpRegistry | None: _NO_CONNECTIONS = "No MCP connections are configured for this run." +_MCP_DISCOVERY_WAIT_SECONDS = 30.0 def _unknown_connection(connection: str, registry: McpRegistry) -> str: @@ -61,11 +63,13 @@ def _format_tool(tool: MCPTool) -> str: @function_tool(timeout=60) async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: - """List the MCP connections available this run, so you can discover them. + """List available MCP connections, loading each connection's catalog on first use. Read-only. Returns one entry per connection with its ``id`` (the exact name you pass to the other MCP tools), ``name``, ``description``, and - ``tool_count``. The response does not include tool schemas. Call + ``tool_count`` (null while the catalog is unknown). Waits at most 30 seconds + for catalogs; unfinished catalogs keep loading in the background. + 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. @@ -73,6 +77,23 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: registry = _registry_from_ctx(ctx) if registry is None or not registry: return {"connections": []} + pending = [ + entry + for name in registry.names() + if (entry := registry.get(name)) is not None + and entry.state != "catalog_ready" + and (entry.session is None or not entry.session.is_dead) + ] + if pending: + await asyncio.wait( + { + asyncio.gather( + *(entry.ensure_catalog() for entry in pending), + return_exceptions=True, + ) + }, + timeout=_MCP_DISCOVERY_WAIT_SECONDS, + ) dead_by_name = {status.name: status.dead for status in registry.statuses()} return { "connections": [ @@ -80,7 +101,7 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]: "id": summary.name, "name": summary.name, "description": summary.purpose, - "tool_count": summary.tool_count, + "tool_count": summary.tool_count if summary.state == "catalog_ready" else None, "dead": dead_by_name.get(summary.name, False), "state": summary.state, } diff --git a/strix/tools/mcp/registry.py b/strix/tools/mcp/registry.py index 65652ab1..6e87c958 100644 --- a/strix/tools/mcp/registry.py +++ b/strix/tools/mcp/registry.py @@ -432,19 +432,19 @@ class McpRegistry: 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.""" + """Connect and list every configured entry with a fixed concurrency 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 def load_catalog(entry: McpConnectionEntry) -> None: async with semaphore: - with contextlib.suppress(McpConnectionUnavailableError): - await entry.ensure_connected() + with contextlib.suppress(Exception): + await entry.ensure_catalog() - await asyncio.gather(*(connect(entry) for entry in self._entries.values())) + await asyncio.gather(*(load_catalog(entry) for entry in self._entries.values())) self._warmup_task = asyncio.create_task(warm(), name="mcp-warmup") return self._warmup_task diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 56b96887..fcbac374 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -43,6 +43,7 @@ from strix.tools.mcp import ( resolve_mcp_call, search_mcp_tools, ) +from strix.tools.mcp import agent_tools as mcp_agent_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 @@ -421,8 +422,18 @@ def test_registry_summaries() -> None: @pytest.mark.asyncio async def test_list_mcps_returns_connections_with_ids_and_descriptions() -> 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) + registry.add( + name="fs", + server=FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]), + purpose="local files", + tool_count=2, + ) + registry.add( + name="db", + server=FakeMCPServer("db", [_mcp_tool("query")]), + purpose=None, + tool_count=1, + ) out = await list_mcps.on_invoke_tool(_ctx(registry), "{}") @@ -437,7 +448,7 @@ async def test_list_mcps_returns_connections_with_ids_and_descriptions() -> None "description": "local files", "tool_count": 2, "dead": False, - "state": "connected", + "state": "catalog_ready", }, { "id": "db", @@ -445,10 +456,143 @@ async def test_list_mcps_returns_connections_with_ids_and_descriptions() -> None "description": None, "tool_count": 1, "dead": False, - "state": "connected", + "state": "catalog_ready", }, ] } + await registry.close() + + +@pytest.mark.asyncio +async def test_list_mcps_loads_catalog_for_unwarmed_connection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + server = FakeMCPServer("fs", [_mcp_tool("read_file"), _mcp_tool("write_file")]) + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda _config: _built_server(server), + ) + registry = McpRegistry() + entry = registry.register(McpConnectionRequest(config=_config("fs", None))) + + out = await list_mcps.on_invoke_tool(_ctx(registry), "{}") + + assert entry.state == "catalog_ready" + assert entry.tool_count == 2 + assert out["connections"][0]["tool_count"] == 2 + assert out["connections"][0]["state"] == "catalog_ready" + await registry.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("slow_stage", ["connect", "catalog"]) +@pytest.mark.parametrize("warmup", [False, True]) +async def test_list_mcps_returns_healthy_connections_while_slow_catalog_keeps_loading( + monkeypatch: pytest.MonkeyPatch, + slow_stage: str, + warmup: bool, +) -> None: + release = asyncio.Event() + started = asyncio.Event() + catalog_calls = 0 + + class _SlowServer(FakeMCPServer): + async def connect(self) -> None: + if slow_stage == "connect": + started.set() + await release.wait() + + async def list_tools(self, run_context: Any = None, agent: Any = None) -> list[MCPTool]: + nonlocal catalog_calls + catalog_calls += 1 + if slow_stage == "catalog": + started.set() + await release.wait() + return await super().list_tools(run_context, agent) + + servers = { + "healthy": FakeMCPServer("healthy", [_mcp_tool("read_file")]), + "slow": _SlowServer("slow", [_mcp_tool("query")]), + } + monkeypatch.setattr( + mcp_client, "_build_server", lambda config: _built_server(servers[config.name]) + ) + monkeypatch.setattr(mcp_agent_tools, "_MCP_DISCOVERY_WAIT_SECONDS", 0.01) + registry = McpRegistry() + for name in servers: + registry.register(McpConnectionRequest(config=_config(name, None))) + if warmup: + registry.start_warmup() + entry = registry.get("slow") + assert entry is not None + try: + out = await asyncio.wait_for(list_mcps.on_invoke_tool(_ctx(registry), "{}"), timeout=1) + assert started.is_set() + assert out["connections"][0]["tool_count"] == 1 + assert out["connections"][0]["state"] == "catalog_ready" + assert out["connections"][1]["tool_count"] is None + assert out["connections"][1]["state"] in {"connecting", "catalog_loading"} + assert entry._catalog_task is not None + assert not entry._catalog_task.done() + + repeated = await asyncio.wait_for(list_mcps.on_invoke_tool(_ctx(registry), "{}"), timeout=1) + assert repeated == out + release.set() + await asyncio.wait_for(entry.ensure_catalog(), timeout=1) + finished = await list_mcps.on_invoke_tool(_ctx(registry), "{}") + assert finished["connections"][1]["tool_count"] == 1 + assert finished["connections"][1]["state"] == "catalog_ready" + assert catalog_calls == 1 + finally: + release.set() + await registry.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_after_deadline", [False, True]) +async def test_list_mcps_retains_and_cleans_up_background_catalog_tasks( + monkeypatch: pytest.MonkeyPatch, + fail_after_deadline: bool, +) -> None: + release = asyncio.Event() + finished = asyncio.Event() + loop_errors: list[dict[str, Any]] = [] + + async def load_catalog() -> list[MCPTool]: + await release.wait() + raise RuntimeError("catalog unavailable") + + registry = McpRegistry() + entry = registry.add( + name="blocked", server=FakeMCPServer("blocked", []), purpose=None, tool_count=0 + ) + assert entry.session is not None + monkeypatch.setattr(entry.session, "list_tools", load_catalog) + monkeypatch.setattr(mcp_agent_tools, "_MCP_DISCOVERY_WAIT_SECONDS", 0.01) + loop = asyncio.get_running_loop() + previous_handler = loop.get_exception_handler() + loop.set_exception_handler(lambda _loop, context: loop_errors.append(context)) + try: + out = await asyncio.wait_for(list_mcps.on_invoke_tool(_ctx(registry), "{}"), timeout=1) + assert out["connections"][0]["tool_count"] is None + assert entry._catalog_task is not None + assert not entry._catalog_task.done() + if fail_after_deadline: + entry._catalog_task.add_done_callback(lambda _task: finished.set()) + release.set() + await asyncio.wait_for(finished.wait(), timeout=1) + assert entry.state == "connected" + # Drop the finished task without awaiting it; its error must already be consumed. + entry._catalog_task = None + else: + task = entry._catalog_task + await registry.close() + assert task.cancelled() + assert loop_errors == [] + finally: + await registry.close() + loop.set_exception_handler(previous_handler) @pytest.mark.asyncio @@ -733,7 +877,37 @@ async def test_registry_warmup_bounds_parallel_connections( release.set() await warmup - assert [summary.state for summary in registry.summaries()] == ["connected"] * 4 + assert [summary.state for summary in registry.summaries()] == ["catalog_ready"] * 4 + assert [summary.tool_count for summary in registry.summaries()] == [1] * 4 + await registry.close() + + +@pytest.mark.asyncio +async def test_registry_warmup_swallows_catalog_listing_errors( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _FailingListServer(FakeMCPServer): + async def list_tools( + self, + run_context: Any = None, + agent: Any = None, + ) -> list[MCPTool]: + raise RuntimeError("catalog unavailable") + + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda _config: _built_server(_FailingListServer("db", [])), + ) + registry = McpRegistry() + entry = registry.register(McpConnectionRequest(config=_config("db", None))) + + warmup = registry.start_warmup() + await warmup + + assert warmup.done() + assert entry.state == "connected" + assert entry.tool_count == 0 await registry.close() @@ -1012,6 +1186,53 @@ def test_prompt_renders_named_connection_inventory() -> None: assert "read the app's schema" in prompt +def test_prompt_omits_unloaded_tool_count() -> None: + prompt = render_system_prompt( + system_prompt_context={ + "mcp_available": True, + "mcp_connections": [ + {"name": "pending_conn", "purpose": None, "tool_count": None, "state": "connected"} + ], + } + ) + + connection_line = next(line for line in prompt.splitlines() if "pending_conn" in line) + assert "(0 tools)" not in connection_line + assert "tools)" not in connection_line + + +def test_prompt_renders_loaded_tool_count() -> None: + prompt = render_system_prompt( + system_prompt_context={ + "mcp_available": True, + "mcp_connections": [ + {"name": "ready_conn", "purpose": None, "tool_count": 8, "state": "catalog_ready"} + ], + } + ) + + assert "ready_conn (8 tools)" in prompt + + +def test_prompt_renders_unavailable_connection_without_tool_count() -> None: + prompt = render_system_prompt( + system_prompt_context={ + "mcp_available": True, + "mcp_connections": [ + { + "name": "unavailable_conn", + "purpose": None, + "tool_count": None, + "state": "unavailable", + } + ], + } + ) + + assert "unavailable_conn (unavailable right now)" in prompt + assert "unavailable_conn (0 tools)" not in prompt + + def test_prompt_inventory_is_gated_on_availability() -> None: """The block is gated on ``mcp_available``; an ``mcp_connections`` payload without it renders nothing, so a stale or spoofed list cannot leak names.""" diff --git a/tests/test_runner_mcp.py b/tests/test_runner_mcp.py index 25a5c888..62d6fab1 100644 --- a/tests/test_runner_mcp.py +++ b/tests/test_runner_mcp.py @@ -8,6 +8,7 @@ routine does the connecting. from __future__ import annotations +import importlib import types from typing import Any @@ -17,10 +18,21 @@ from agents import ModelSettings import strix.tools.mcp as mcp_pkg import strix.tools.notes.tools as notes_tools import strix.tools.todo.tools as todo_tools +from strix.agents import factory from strix.core import runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager from strix.tools.mcp import McpConnectionConfig, McpConnectionRequest +from strix.tools.mcp import client as mcp_client + + +_test_mcp_client = importlib.import_module("tests.test_mcp_client") +FakeMCPServer: Any = _test_mcp_client.FakeMCPServer +_mcp_tool: Any = _test_mcp_client._mcp_tool + + +def _built_server(server: Any) -> Any: + return mcp_client.BuiltMcpServer(server, None) def _settings() -> Any: @@ -62,6 +74,11 @@ def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: monkeypatch.setattr(runner, "build_strix_agent", lambda **_k: object()) monkeypatch.setattr(runner, "make_child_factory", lambda **_k: lambda **_kk: object()) monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object()) + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda config: _built_server(FakeMCPServer(config.name, [])), + ) async def _run_agent_loop(**_kwargs: Any) -> None: return None @@ -156,6 +173,11 @@ async def test_roster_is_persisted_even_without_a_status_sink( "load_user_mcp_configs", lambda: [McpConnectionConfig(name="local_fs", transport="stdio", command="npx")], ) + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda _config: _built_server(FakeMCPServer("local_fs", [])), + ) persisted: list[list[dict[str, Any]]] = [] @@ -182,3 +204,106 @@ async def test_roster_is_persisted_even_without_a_status_sink( "state": "configured", } ] + + +@pytest.mark.asyncio +async def test_warmup_listing_failure_does_not_interrupt_scan( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Any, +) -> None: + _wire_runner(monkeypatch, tmp_path) + + server = FakeMCPServer("db", []) + + async def _failing_list_tools( + run_context: Any = None, + agent: Any = None, + ) -> list[Any]: + del run_context, agent + raise RuntimeError("catalog unavailable") + + server.list_tools = _failing_list_tools + + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda _config: _built_server(server), + ) + root_builds: list[dict[str, Any]] = [] + monkeypatch.setattr(runner, "build_strix_agent", lambda **kwargs: root_builds.append(kwargs)) + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-mcp-list-failure", + image="img", + coordinator=AgentCoordinator(), + mcp_connection_requests=[ + McpConnectionRequest( + config=McpConnectionConfig(name="db", url="https://mcp.example.com") + ) + ], + ) + + assert len(root_builds) == 1 + + +@pytest.mark.asyncio +async def test_warmed_tool_counts_reach_root_and_child_agent_builds( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Any, +) -> None: + _wire_runner(monkeypatch, tmp_path) + server = FakeMCPServer("db", [_mcp_tool(f"tool_{index}") for index in range(3)]) + monkeypatch.setattr(mcp_client, "_build_server", lambda _config: _built_server(server)) + + root_builds: list[dict[str, Any]] = [] + + def _build_root(**kwargs: Any) -> object: + root_builds.append(kwargs) + return object() + + monkeypatch.setattr( + runner, + "build_strix_agent", + _build_root, + ) + child_builds: list[dict[str, Any]] = [] + + def _build_child(**kwargs: Any) -> object: + child_builds.append(kwargs) + return object() + + monkeypatch.setattr( + factory, + "build_strix_agent", + _build_child, + ) + real_make_child_factory = factory.make_child_factory + child_factory_capture: dict[str, Any] = {} + + def _make_child_factory(**kwargs: Any) -> Any: + child_factory_capture["context"] = kwargs["system_prompt_context"] + child_factory_capture["factory"] = real_make_child_factory(**kwargs) + return child_factory_capture["factory"] + + monkeypatch.setattr(runner, "make_child_factory", _make_child_factory) + request = McpConnectionRequest( + config=McpConnectionConfig(name="db", url="https://mcp.example.com", notes="database") + ) + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-mcp-prompt-counts", + image="img", + coordinator=AgentCoordinator(), + mcp_connection_requests=[request], + ) + + root_context = root_builds[0]["system_prompt_context"] + assert root_context["mcp_connections"][0]["tool_count"] == 3 + assert child_factory_capture["context"] is root_context + assert child_factory_capture["context"]["mcp_connections"][0]["tool_count"] == 3 + + child_factory_capture["factory"](name="child", skills=[]) + assert child_builds[0]["system_prompt_context"] is child_factory_capture["context"] + assert child_builds[0]["system_prompt_context"]["mcp_connections"][0]["tool_count"] == 3 diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index e25d09ab..a564a3ea 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -6,6 +6,7 @@ flow through to the root agent's ``build_strix_agent`` call. from __future__ import annotations +import importlib import os import types from typing import Any @@ -27,6 +28,11 @@ from strix.core.inputs import make_model_settings from strix.runtime import session_manager from strix.tools.load_skill.tool import load_skill from strix.tools.mcp import BearerAuth, McpConnectionConfig, McpConnectionRequest +from strix.tools.mcp import client as mcp_client + + +_test_mcp_client = importlib.import_module("tests.test_mcp_client") +FakeMCPServer: Any = _test_mcp_client.FakeMCPServer def _make_rate_limit_error() -> RateLimitError: @@ -84,6 +90,11 @@ def _patch_engine_scaffold( monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: scope_context) monkeypatch.setattr(runner, "make_model_settings", lambda *_args, **_kwargs: ModelSettings()) + monkeypatch.setattr( + mcp_client, + "_build_server", + lambda config: mcp_client.BuiltMcpServer(FakeMCPServer(config.name, []), None), + ) captured: dict[str, Any] = {} @@ -222,7 +233,12 @@ 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": 0} + { + "name": "fs", + "purpose": "local files", + "tool_count": 0, + "state": "catalog_ready", + } ]