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
This commit is contained in:
yoni-at-strix 2026-09-22 14:02:02 -04:00 • committed by GitHub
parent 56e9ae982c
commit e158eab3f8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 804 additions and 217 deletions

View file

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

View file

@ -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="<name>") to inspect one connection's tools, each with its name, description, and JSON input schema.
3. Call call_mcp(connection="<name>", tool="<tool>", arguments={...}) to run one, passing an arguments object that matches the schema (omit arguments for a tool that takes none).
- 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="<name>", query="<capability>") for a short candidate list.
3. Call get_mcp_tool_schema(connection="<name>", tool="<tool>") for the one schema you need.
4. Call call_mcp(connection="<name>", tool="<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:

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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