mirror of
https://github.com/usestrix/strix.git
synced 2026-10-09 03:18:31 +00:00
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:
parent
4695290682
commit
f5a900416b
7 changed files with 421 additions and 25 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue