fix(mcp): load catalogs before agents discover connections (#1492)

* fix(mcp): load catalogs before agents discover connections

* fix(mcp): bound discovery waits without cancelling catalog loads

* fix(mcp): keep discovery waiters alive across deadlines
This commit is contained in:
alex s 2026-10-08 12:42:39 -04:00 • committed by GitHub
parent 4695290682
commit f5a900416b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 421 additions and 25 deletions

View file

@ -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.

View file

@ -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")

View file

@ -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,
}

View file

@ -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

View file

@ -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."""

View file

@ -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

View file

@ -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",
}
]