Files
Kigi-CLI/crates/codegen/kigi-mcp/src/servers.rs
T
ZacharyZhang-NY 6f31415ed6 §9 acceptance: grep-zero sweep — every internal x.ai/grok identifier renamed
The PRD's first acceptance gate now holds: grep -RinE '\bx\.ai\b|grok'
crates/ --include='*.rs' → 0 matches (exempt: NOTICE and third-party
license archives, README provenance, and the required 'Based on Grok
Build Open Source' attribution, now sourced from version_attribution.txt).

Wire-visible renames (both sides in this repo, changed in lockstep):
- Auth method id 'grok.com' → 'kimi-code' (AuthMethodKind::KimiCode).
- Every x.ai/* and _x.ai/* ACP ext method and meta key → kigi/* /
  _kigi/* (~200 names; grokShell → kigiShell). Session-file replay keeps
  a read-side alias for the legacy '_x.ai/session/update' method so
  existing updates.jsonl histories load; writes emit only the new name
  (both directions test-pinned).
- Agent types grok-build* → kigi* with a documented legacy-prefix alias
  at resolution time so persisted sessions keep resolving.
- ToolNamespace/BuiltinAgentName GrokBuild* → Kigi* (wire snake_case
  kigi/kigi_concise/kigi_hashline; schema regenerated); grok_build
  implementation dirs renamed to kigi*.
- x-grok-* headers → x-kigi-*, __GROK_* sentinels → __KIGI_*, themes
  grokday/groknight → kigiday/kiginight (old persisted values fall back
  to the default theme), web_fetch allowlist xAI hosts → kimi.com +
  moonshot platforms, changelog CDN → this repo, grok-build changelog
  archives deleted.
- BYOK default endpoint removed: [endpoints] api_base_url is now truly
  optional with NO default — consumers fail fast with the flag name when
  unset (no silent x.ai egress). Mock harnesses inject it explicitly.
- System-prompt identity fixed: 'released by xAI' → 'an unofficial
  community CLI for Kimi' (template + regenerated encrypted form).

Also repaired pre-existing grok-era test debt found by the sweep: the
stale trace_classify default-model pin, the grok-pager UA label test,
pty-harness stale-binary reuse and non-hermetic moonshot routing (a PTY
test could previously reach the real api.moonshot.cn), and the outdated
oauth fixture scope key.

Gates: §9 grep 0; fmt clean; workspace check/clippy 0/0 (-D warnings);
FULL cargo test --workspace: 234 suites, 21,961 passed, 0 failed;
deny advisories ok.
2026-07-18 02:48:46 -04:00

7523 lines
293 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! MCP server integration using the official rmcp SDK.
use std::collections::HashMap;
use std::ffi::OsString;
use std::future::Future;
use std::sync::{Arc, LazyLock};
use agent_client_protocol as acp;
use futures::StreamExt;
use regex::Regex;
use tokio::{
io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader},
process::{ChildStderr, Command},
sync::{Mutex, Notify},
};
use rmcp::{
ClientHandler, ServiceExt,
model::{
CallToolRequestParams, ClientCapabilities, ClientInfo, Implementation,
PaginatedRequestParams,
},
service::{
ClientInitializeError, NotificationContext, RoleClient, RunningService, ServiceError,
},
service::{RxJsonRpcMessage, TxJsonRpcMessage},
transport::{
StreamableHttpClientTransport, Transport,
streamable_http_client::StreamableHttpClientTransportConfig,
},
};
use crate::oauth_config::McpOAuthConfig;
use kigi_tools::types::{
output::{MCPOutput, MCPOutputDetails, ToolOutput},
tool::{ToolKind, ToolNamespace},
tool_metadata::ToolMetadata,
};
use kigi_tools::util::ProcessGroup;
/// MCP tool name delimiter: server names are qualified as `"server__tool"`.
/// Canonical definition lives in `kigi_workspace_types`; re-exported here
/// for callers that historically imported it from this module.
pub use kigi_workspace_types::MCP_TOOL_NAME_DELIMITER;
/// Normalize an MCP server URL for comparison: strip trailing slashes.
/// Must match the normalization the host's managed-config layer uses
/// (e.g. shell's `session::managed_mcp::normalize_url`) so URL
/// lookup keys agree.
fn normalize_url(url: &str) -> String {
url.trim_end_matches('/').to_string()
}
/// Regex for strictest cross-provider tool name validation.
///
/// Requirements across providers:
/// - Anthropic/OpenAI: `^[a-zA-Z0-9_-]{1,64}$` (allows starting with digit/hyphen)
/// - Google Gemini: `^[a-zA-Z_][a-zA-Z0-9_.-]{0,63}$` (must start with letter/underscore, allows dots)
///
/// Strictest common denominator: must start with letter/underscore, only alphanumeric/_/- allowed, max 64 chars.
static TOOL_NAME_REGEX: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"^[a-zA-Z_][a-zA-Z0-9_-]{0,63}$").unwrap());
/// Validate that a tool name matches the strictest cross-provider LLM API requirements.
///
/// Pattern: `^[a-zA-Z_][a-zA-Z0-9_-]{0,63}$`
/// - Must start with a letter or underscore (Gemini requirement)
/// - Only letters, digits, underscores, hyphens allowed (no dots — Anthropic/OpenAI requirement)
/// - Maximum 64 characters
///
/// Returns `Ok(())` if valid, or `Err(reason)` if invalid.
pub fn validate_tool_name(name: &str) -> Result<(), String> {
if name.is_empty() {
return Err("tool name cannot be empty".to_string());
}
if !TOOL_NAME_REGEX.is_match(name) {
return Err(format!(
"tool name '{}' is invalid — must match ^[a-zA-Z_][a-zA-Z0-9_-]{{0,63}}$ (start with letter/underscore, max 64 chars)",
name
));
}
Ok(())
}
/// Sanitize an MCP server or tool name into a single safe path segment
/// (e.g. `"user-Hugging Face"` becomes `user-Hugging_Face`). Shared so the
/// per-server folder advertised in the prompt matches the tool files on disk.
pub fn sanitize_descriptor_segment(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if c.is_ascii_alphanumeric() || c == '-' || c == '_' || c == '.' {
out.push(c);
} else {
out.push('_');
}
}
if out.is_empty() {
out.push('_');
}
out
}
/// Result of a diff-based MCP config update.
#[derive(Debug)]
pub struct McpConfigDiff {
/// Server names that are new or had their config changed.
pub added: Vec<McpServerName>,
/// Server names that were removed or had their config changed (old instance torn down).
pub removed: Vec<McpServerName>,
/// Server names whose config is identical — clients kept alive.
pub retained: Vec<McpServerName>,
}
/// MCP server name used as the key in client/tool maps (e.g. `"github"`).
pub type McpServerName = String;
/// Unqualified MCP tool name (e.g. `"create_issue"`, without the `server__` prefix).
type ToolName = String;
/// Typed state machine for MCP-pool initialization.
///
/// Replaces the previous trio of correlated fields — `initialized: bool`,
/// `initializing: bool`, `initializing_servers: HashSet<McpServerName>` —
/// whose product space could represent nonsensical combinations such as
/// "initialized AND initializing" or "no init started AND per-server
/// handshakes outstanding". With one enum field, every legal state has
/// exactly one representation and the compiler enforces exhaustiveness
/// at every match site.
///
/// Lifecycle:
///
/// ```text
/// ┌─────────────┐ try_start ┌──────────────────────┐
/// │ NotStarted │ ──────────▶ │ Starting{handshakes}│
/// └─────────────┘ ◀── cancel ─┴──────────┬───────────┘
/// ▲ │ finish
/// │ cancel ▼
/// │ ┌──────────────────────────┐
/// └──────────────────┤ Finished{handshakes} │
/// └──────────────────────────┘
/// ```
///
/// `Starting` is the pre-`finish_init` window; `Finished` is the post-
/// `finish_init` window where per-server background handshakes may still
/// be draining. `is_complete()` requires `Finished` with an empty
/// handshaking set.
#[derive(Debug, Default)]
pub enum InitProgress {
/// Init has never been started, or was cancelled / reset by a
/// config change.
#[default]
NotStarted,
/// `try_start_init` was called; per-server tasks may be spawning;
/// `finish_init` has NOT yet fired. `handshaking` tracks the set of
/// servers whose background handshake is in flight.
Starting {
handshaking: std::collections::HashSet<McpServerName>,
},
/// `finish_init` fired (deliberately early, so the session is not
/// blocked on MCP for non-MCP work). Background per-server
/// handshakes may still be running; `handshaking` shrinks as each
/// completes. `is_complete()` returns `true` only when it is empty.
Finished {
handshaking: std::collections::HashSet<McpServerName>,
},
}
impl InitProgress {
/// True iff every per-server handshake has settled and `finish_init`
/// has fired. Pairs with [`Self::is_in_progress`].
pub fn is_complete(&self) -> bool {
matches!(self, Self::Finished { handshaking } if handshaking.is_empty())
}
/// True iff any init work is outstanding — either we are pre-
/// `finish_init`, or per-server handshakes are still in flight in
/// the background.
pub fn is_in_progress(&self) -> bool {
match self {
Self::Starting { .. } => true,
Self::Finished { handshaking } => !handshaking.is_empty(),
Self::NotStarted => false,
}
}
/// True iff `finish_init` has fired, regardless of whether
/// background handshakes are still draining. Used for diagnostic
/// logging where the caller wants to distinguish pre-finish from
/// post-finish-with-bg-work.
pub fn has_finished_init(&self) -> bool {
matches!(self, Self::Finished { .. })
}
/// True iff the named server is currently handshaking.
pub fn is_server_handshaking(&self, name: &str) -> bool {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => {
handshaking.contains(name)
}
Self::NotStarted => false,
}
}
/// Iterate over server names whose background handshake is still
/// in flight. Empty when [`Self::NotStarted`] or fully complete.
pub fn handshaking_servers(&self) -> impl Iterator<Item = &McpServerName> {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => {
Some(handshaking.iter())
}
Self::NotStarted => None,
}
.into_iter()
.flatten()
}
/// Number of in-flight per-server handshakes.
pub fn handshaking_count(&self) -> usize {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => handshaking.len(),
Self::NotStarted => 0,
}
}
/// Transition `NotStarted` → `Starting { ∅ }`. Returns `true` on
/// successful transition, `false` if init was already started or
/// finished (mirrors the pre-refactor `try_start_init` contract).
pub fn try_start(&mut self) -> bool {
if matches!(self, Self::NotStarted) {
*self = Self::Starting {
handshaking: std::collections::HashSet::new(),
};
true
} else {
false
}
}
/// Transition `Starting { hs }` → `Finished { hs }`, preserving the
/// handshaking set. No-op if already `Finished`; no-op-with-log if
/// called from `NotStarted` (defensive — that would be a caller bug).
pub fn finish(&mut self) {
match self {
Self::Starting { handshaking } => {
// Move only the inner set, leaving the outer `&mut self`
// ready to be reassigned without going through
// `mem::take(self)` (which would force a `NotStarted`
// placeholder and a redundant put-back in the
// already-Finished arm below).
let handshaking = std::mem::take(handshaking);
*self = Self::Finished { handshaking };
}
// Idempotent: already past the finish boundary. Per-server
// handshakes continue draining via `mark_handshake_complete`.
Self::Finished { .. } => {}
Self::NotStarted => {
tracing::warn!(
"InitProgress::finish called from NotStarted; staying in NotStarted"
);
}
}
}
/// Transition any state → `NotStarted`. Clears all per-server
/// progress. Used on generation mismatch (config change racing
/// active init) and on full reset.
pub fn cancel(&mut self) {
*self = Self::NotStarted;
}
/// Add names to the handshaking set. Only meaningful in `Starting`
/// or `Finished`; warns if called from `NotStarted` (that would
/// mean a per-server handshake started without a `try_start_init`,
/// which is a caller bug).
pub fn mark_handshaking(&mut self, names: impl IntoIterator<Item = McpServerName>) {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => {
handshaking.extend(names);
}
Self::NotStarted => {
tracing::warn!("InitProgress::mark_handshaking called from NotStarted; ignoring");
}
}
}
/// Remove a server from the handshaking set (on success or failure
/// of its handshake). No-op if not present or if `NotStarted`.
pub fn mark_handshake_complete(&mut self, name: &str) {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => {
handshaking.remove(name);
}
Self::NotStarted => {}
}
}
/// Clear the handshaking set entirely. Used by the proxy-mode and
/// bg-handshake completion paths as a defensive sweep after the
/// per-server `mark_handshake_complete` calls — ensures the set is
/// empty before/after `finish_init` fires.
pub fn clear_handshaking(&mut self) {
match self {
Self::Starting { handshaking } | Self::Finished { handshaking } => {
handshaking.clear();
}
Self::NotStarted => {}
}
}
}
/// One in-process SDK MCP server registration: its tool-namespace name and the
/// SDK-side id echoed back in `kigi/mcp/sdk_call`. A named struct (rather than a
/// `(String, String)` tuple) so callers can't transpose the two strings.
///
/// `Deserialize`d straight from a `_meta["kigi/mcp/servers"]` entry, so the
/// `serverId` wire field name is declared (and serde-checked) exactly once here.
#[derive(Debug, Clone, serde::Deserialize)]
pub struct AcpServerEntry {
pub name: McpServerName,
#[serde(rename = "serverId")]
pub server_id: String,
}
/// The session's in-process SDK MCP servers (declared via `_meta["kigi/mcp/servers"]`,
/// reached over the ACP reverse channel), bundled with the shared reverse-RPC invoker.
/// Held as `McpState::acp_mcp: Option<_>` so the set is one atom — present together or
/// absent, never "servers without an invoker" — and survives `update_configs` clears
/// (config reloads only touch `configs`/`owned_clients`). Per-server config.toml overrides
/// are NOT cached here — they are re-resolved per init (see [`McpState::build_pending_acp_clients`]).
struct AcpMcpRegistry {
/// Registered servers (`name -> serverId`).
servers: Vec<AcpServerEntry>,
/// Shared reverse-RPC invoker all these servers' tools are called through (emits
/// `kigi/mcp/sdk_call` over the ACP connection).
invoker: Arc<dyn crate::acp_transport::AcpReverseInvoker>,
}
/// Consolidated MCP state behind a single lock. Generation counter detects stale inits.
pub struct McpState {
pub configs: Vec<acp::McpServer>,
pub meta_config_map: McpMetaConfigMap,
/// Clients owned by this session; cleared on config changes.
pub owned_clients: HashMap<McpServerName, Arc<McpClient>>,
/// Clients inherited from parent via `SharedMcpPool`; never cleared by config changes.
pub shared_clients: HashMap<McpServerName, Arc<McpClient>>,
/// The session's in-process SDK MCP servers + their shared invoker/overrides; `None`
/// when the session has none. See [`AcpMcpRegistry`]. Kept out of `configs` (the closed
/// `acp::McpServer` enum) so it survives `update_configs` clears.
acp_mcp: Option<AcpMcpRegistry>,
/// Encapsulated init lifecycle. Access via [`Self::is_initialized`],
/// [`Self::is_initializing`], [`Self::try_start_init`],
/// [`Self::finish_init`], etc. — those route through a single
/// [`InitProgress`] state machine that rules out nonsensical
/// combinations like "initialized AND initializing".
///
/// Private on purpose: external callers must go through the typed
/// transition methods, not poke the variant directly.
init_progress: InitProgress,
pub generation: u64,
/// Qualified tool name → `_meta` from MCP tools/list. Populated during init.
pub mcp_tool_meta: HashMap<String, serde_json::Value>,
/// HTTP servers that support OAuth but haven't been authenticated yet.
pub auth_required: std::collections::HashSet<McpServerName>,
/// Servers whose background init failed (handshake error, `tools/list`
/// error, or overall init timeout) even though a client object exists,
/// mapped to a short failure cause surfaced to the model in the MCP
/// reminder. Surfaced as `Unavailable` in status snapshots so a server
/// that connected but never finished initializing — e.g. wedged on
/// `tools/list` and registered zero tools — does not misleadingly show
/// as `Ready`. Cleared when the server begins a fresh init attempt.
pub init_failed: std::collections::HashMap<McpServerName, String>,
/// Per-server set of unqualified tool names that the user has disabled.
/// Persisted to `~/.kigi/config.toml` under `[mcp_servers.<name>].disabled_tools`.
pub disabled_tools: HashMap<McpServerName, std::collections::HashSet<ToolName>>,
/// Stashed registrations for disabled tools so they can be re-enabled
/// without a full MCP re-init (no need to call `list_tools` again).
pub disabled_tool_registrations: HashMap<String, McpToolRegistration>,
event_writer: kigi_file_utils::events::EventWriter,
/// Sender wired by the session actor to its `StatusDispatcher`
/// task. When `Some`, the state — and every [`McpClient`] reached
/// through [`Self::all_clients`] / [`Self::get_client`] — forwards
/// [`McpClientEvent`]s here for coalescing and fan-out as ACP
/// `kigi/mcp/server_status` notifications.
///
/// Intentionally `None` in subagent-pool / shared-pool snapshots
/// ([`SharedMcpPool`]) where the **parent** session is the
/// single owner of liveness/notification flow. Clients in those
/// snapshots inherit the parent's `Arc<McpClient>` (with the
/// parent's `notify_tx` slot still pointing at the parent), so
/// duplicating event flow into a subagent would just double-push
/// every event.
///
/// Populated by [`Self::set_client_event_tx`], which fans the
/// sender into every existing client's `notify_tx` slot.
///
/// **Private on purpose.** Callers MUST go through
/// [`Self::set_client_event_tx`] so the sender is fanned out into
/// every existing `owned_clients` entry; a direct field write
/// (`state.client_event_tx = Some(tx)`) would leave all
/// already-owned clients with `notify_tx = None`, silently
/// dropping `tools/list_changed`, `Ready`, and `HandshakeFailed`
/// emits for them. Read access is via [`Self::client_event_tx`].
client_event_tx: Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>>,
}
impl McpState {
pub fn new(configs: Vec<acp::McpServer>) -> Self {
Self::new_with_meta(configs, McpMetaConfigMap::new())
}
pub fn new_with_meta(configs: Vec<acp::McpServer>, meta_config_map: McpMetaConfigMap) -> Self {
Self {
configs,
meta_config_map,
owned_clients: HashMap::new(),
shared_clients: HashMap::new(),
acp_mcp: None,
init_progress: InitProgress::default(),
generation: 0,
mcp_tool_meta: HashMap::new(),
auth_required: std::collections::HashSet::new(),
init_failed: HashMap::new(),
disabled_tools: HashMap::new(),
disabled_tool_registrations: HashMap::new(),
event_writer: kigi_file_utils::events::EventWriter::noop(),
client_event_tx: None,
}
}
/// Install (or remove) the [`McpClientEvent`] sender owned by the
/// session-actor `StatusDispatcher`.
///
/// Synchronous: the per-client slot is a `parking_lot::Mutex`, so
/// the iteration no longer holds `&mut McpState` across `.await`.
///
/// Side effect: clones the sender into every existing client's
/// shared `notify_tx` slot. New clients added later (e.g. on a
/// config diff that re-spawns a server) MUST be wired by the
/// caller post-construction — typically by calling
/// [`McpClient::set_event_tx`] **before**
/// `get_tool_registrations` (so `ensure_initialized`'s
/// `Ready`/`HandshakeFailed` emit fires with `Some(tx)` and the
/// `KigiClientHandler` cloned during `try_handshake` reads
/// through the same Arc).
pub fn set_client_event_tx(
&mut self,
tx: Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>>,
) {
self.client_event_tx = tx.clone();
for client in self.owned_clients.values() {
client.set_event_tx(tx.clone());
}
// Shared clients are intentionally NOT wired here: see the
// `client_event_tx` doc-comment for why a subagent must not
// duplicate the parent's event flow.
}
/// Read-only access to the installed [`McpClientEvent`] sender.
///
/// Returns a clone of the sender wired by
/// [`Self::set_client_event_tx`], or `None` if no dispatcher is
/// attached (subagent / shared-pool snapshot). Exposed as a getter
/// rather than a `pub` field so the fan-out contract documented on
/// `client_event_tx` cannot be bypassed by a direct assignment.
pub fn client_event_tx(&self) -> Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>> {
self.client_event_tx.clone()
}
pub fn set_event_writer(&mut self, writer: kigi_file_utils::events::EventWriter) {
self.event_writer = writer;
}
pub fn event_writer(&self) -> &kigi_file_utils::events::EventWriter {
&self.event_writer
}
/// Register the session's in-process SDK MCP servers (`name -> serverId`) plus the
/// reverse-RPC invoker. Held across `update_configs` clears so each init re-adds them.
pub fn set_acp_servers(
&mut self,
servers: Vec<AcpServerEntry>,
invoker: Arc<dyn crate::acp_transport::AcpReverseInvoker>,
) {
self.acp_mcp = Some(AcpMcpRegistry { servers, invoker });
}
/// Whether any in-process SDK MCP servers are registered (so the session knows to
/// run MCP init even with no `configs`).
pub fn has_acp_servers(&self) -> bool {
self.acp_mcp
.as_ref()
.is_some_and(|acp| !acp.servers.is_empty())
}
/// Registered SDK servers not yet connected (no owned/shared client) — the ones an
/// init pass should build. Shared by [`build_pending_acp_clients`] and
/// [`pending_acp_server_names`] so the "what to build" filter lives in one place.
fn pending_acp_entries(&self) -> impl Iterator<Item = &AcpServerEntry> {
self.acp_mcp.iter().flat_map(|acp| {
acp.servers.iter().filter(|entry| {
!self.owned_clients.contains_key(&entry.name)
&& !self.shared_clients.contains_key(&entry.name)
})
})
}
/// Names of the SDK servers [`build_pending_acp_clients`] will build — used to mark
/// them initializing before the (async) build.
pub fn pending_acp_server_names(&self) -> Vec<String> {
self.pending_acp_entries()
.map(|entry| entry.name.to_string())
.collect()
}
/// Build [`McpClient`]s for registered ACP servers not already connected. Appended to the
/// init handshake batch so they register tools + land in `owned_clients` on the SAME path
/// as HTTP/stdio servers.
///
/// `overrides` is the per-server config.toml tuning (keyed by server name), resolved by
/// the caller per init — kept caller-side so this method stays pure (no file I/O under
/// the `McpState` lock).
pub fn build_pending_acp_clients(
&self,
overrides: &HashMap<String, McpClientTimeoutOverrides>,
) -> Vec<McpClient> {
let Some(acp) = &self.acp_mcp else {
return Vec::new();
};
self.pending_acp_entries()
.map(|entry| {
McpClient::new_acp(
entry.name.clone(),
entry.server_id.clone(),
acp.invoker.clone(),
overrides.get(&entry.name),
self.meta_config_map.get(&entry.name),
)
})
.collect()
}
/// Check if a specific tool is disabled for a server (unqualified tool name).
pub fn is_tool_disabled(&self, server_name: &str, tool_name: &str) -> bool {
self.disabled_tools
.get(server_name)
.is_some_and(|set| set.contains(tool_name))
}
/// Update configs and reset initialization state.
/// Returns true if the configs actually changed, false if they were identical.
pub fn update_configs(&mut self, new_configs: Vec<acp::McpServer>) -> bool {
if mcp_servers_equal(&self.configs, &new_configs) {
tracing::debug!("MCP configs unchanged, skipping update");
return false;
}
// Clear owned clients only — shared (inherited) clients are untouched.
self.owned_clients.clear();
self.mcp_tool_meta.clear();
self.disabled_tool_registrations.clear();
self.configs = new_configs;
self.init_progress.cancel();
self.auth_required.clear();
self.generation = self.generation.wrapping_add(1);
true
}
/// Diff-based config update: only tears down servers whose config changed
/// or were removed, keeps healthy unchanged servers alive.
///
/// Returns `None` if configs are identical (no work needed), or `Some(diff)`
/// describing which servers to add/remove.
pub fn update_configs_diff(
&mut self,
new_configs: Vec<acp::McpServer>,
) -> Option<McpConfigDiff> {
if mcp_servers_equal(&self.configs, &new_configs) {
tracing::debug!("MCP configs unchanged, skipping update");
return None;
}
let old_by_name: HashMap<&str, String> = self
.configs
.iter()
.filter_map(|c| match serde_json::to_string(c) {
Ok(json) => Some((mcp_server_name(c), json)),
Err(e) => {
tracing::warn!(server = mcp_server_name(c), error = %e, "Failed to serialize MCP server config for diff");
None
}
})
.collect();
let new_by_name: HashMap<&str, String> = new_configs
.iter()
.filter_map(|c| match serde_json::to_string(c) {
Ok(json) => Some((mcp_server_name(c), json)),
Err(e) => {
tracing::warn!(server = mcp_server_name(c), error = %e, "Failed to serialize MCP server config for diff");
None
}
})
.collect();
let mut removed = Vec::new();
let mut added = Vec::new();
let mut retained = Vec::new();
for (name, old_json) in &old_by_name {
match new_by_name.get(name) {
None => removed.push(name.to_string()),
Some(new_json) if new_json != old_json => {
removed.push(name.to_string());
}
Some(_) => retained.push(name.to_string()),
}
}
for (name, new_json) in &new_by_name {
match old_by_name.get(name) {
None => added.push(name.to_string()),
Some(old_json) if old_json != new_json => {
added.push(name.to_string());
}
Some(_) => {}
}
}
for name in &removed {
self.owned_clients.remove(name);
self.auth_required.remove(name);
self.init_progress.mark_handshake_complete(name);
let prefix = format!("{}{}", name, MCP_TOOL_NAME_DELIMITER);
self.mcp_tool_meta.retain(|k, _| !k.starts_with(&prefix));
self.disabled_tool_registrations
.retain(|k, _| !k.starts_with(&prefix));
}
tracing::info!(
retained = retained.len(),
added = added.len(),
removed = removed.len(),
"MCP config diff: {} retained, {} added, {} removed",
retained.len(),
added.len(),
removed.len(),
);
self.configs = new_configs;
self.init_progress.cancel();
self.generation = self.generation.wrapping_add(1);
Some(McpConfigDiff {
added,
removed,
retained,
})
}
/// Returns `true` only when MCP setup is fully complete: the
/// init lifecycle reached [`InitProgress::Finished`] AND every
/// per-server background handshake has settled (success or failure).
///
/// The strict per-server check matters because session actors call
/// [`Self::finish_init`] **early** (right after spawning processes,
/// before any handshake completes) so the session isn't blocked on
/// MCP for non-MCP work. Callers that gate MCP-tool dispatch on
/// "is MCP actually ready" — e.g. the Blocking-strategy waits in
/// `prepare_tool_definitions_timed`, `wait_for_mcp_initialized`,
/// and the tool-dispatch fast path — therefore need the *combined*
/// check or they'd race the in-flight per-server handshakes and the
/// first tool call would land inside the
/// [`ClientState::Initializing`] window.
///
/// Delegates to [`InitProgress::is_complete`]; see that doc for the
/// full state machine.
pub fn is_initialized(&self) -> bool {
self.init_progress.is_complete()
}
/// Returns `true` whenever any initialization work is still
/// outstanding: pre-`finish_init` OR at least one per-server
/// handshake is still running in the background.
///
/// Pairs with [`Self::is_initialized`]: during the window between
/// the early [`Self::finish_init`] and the background task draining
/// the per-server handshaking set, `is_initialized()` is still
/// `false` (per-server work remains) AND `is_initializing()` is
/// `true` (so wait-loops keep waiting instead of kicking off a
/// second init).
pub fn is_initializing(&self) -> bool {
self.init_progress.is_in_progress()
}
/// Returns `true` once `finish_init` has fired, regardless of
/// whether per-server background handshakes are still draining.
/// Used for diagnostic logging where the caller wants to
/// distinguish "pre-finish window" from "post-finish, bg work
/// outstanding".
pub fn has_finished_init(&self) -> bool {
self.init_progress.has_finished_init()
}
/// Borrow the underlying [`InitProgress`] state machine, primarily
/// for tests that want to assert against the discriminant directly.
pub fn init_progress(&self) -> &InitProgress {
&self.init_progress
}
/// Try to start initialization. Returns `true` if we transitioned
/// from [`InitProgress::NotStarted`] to [`InitProgress::Starting`];
/// returns `false` if init is already in progress or finished.
pub fn try_start_init(&mut self) -> bool {
self.init_progress.try_start()
}
/// Transition [`InitProgress::Starting`] → [`InitProgress::Finished`],
/// preserving the per-server handshaking set. Called early (before
/// per-server handshakes complete) so the session is unblocked for
/// non-MCP work — `is_initialized()` still returns `false` until
/// every handshake has reported via [`Self::mark_server_ready`].
pub fn finish_init(&mut self) {
self.init_progress.finish();
}
/// Cancel initialization back to [`InitProgress::NotStarted`].
/// Used when generation changed during init (config change races
/// with active init) and on full reset.
pub fn cancel_init(&mut self) {
self.init_progress.cancel();
}
/// Add server names to the handshaking set. Call after filtering
/// `configs_to_start`, before spawning per-server tasks. Only
/// meaningful in [`InitProgress::Starting`] / [`InitProgress::Finished`];
/// logs a warning otherwise.
pub fn mark_servers_initializing(&mut self, names: impl IntoIterator<Item = McpServerName>) {
let names: Vec<McpServerName> = names.into_iter().collect();
// A fresh init attempt clears any prior failure for these servers so
// a server that recovers on retry stops showing as `Unavailable`.
for name in &names {
self.init_failed.remove(name);
}
self.init_progress.mark_handshaking(names);
}
/// Remove a server from the handshaking set (on success or failure
/// of its handshake). Safe if not present.
pub fn mark_server_ready(&mut self, name: &str) {
self.init_progress.mark_handshake_complete(name);
}
/// Record a per-server background-init failure for status reporting,
/// routing it to the correct set so the two stay disjoint.
///
/// `needs_auth` failures are owned by the auth state machine: its recovery
/// paths (`handle_mcp_auth_trigger`, `retry_auth_required_servers`)
/// re-handshake and clear `auth_required`. Such servers must therefore NOT
/// also land in `init_failed`, or a server that successfully authenticates
/// would stay reported as `Unavailable` with zero tools. Every other
/// failure (handshake / `tools/list` error or init timeout) goes to
/// `init_failed` so the server surfaces as `Unavailable`.
///
/// `detail` is a short cause stored for non-auth failures (the value in
/// [`Self::init_failed`]); ignored for `needs_auth`.
pub fn record_init_failure(&mut self, name: &str, needs_auth: bool, detail: Option<String>) {
if needs_auth {
self.auth_required.insert(name.to_string());
} else {
self.init_failed
.insert(name.to_string(), detail.unwrap_or_default());
}
}
/// Clear a prior init failure for `name` (symmetric with
/// [`Self::record_init_failure`]). Used by the reactive managed re-auth
/// path so a server that recovers is no longer reported as `Unavailable`
/// with a stale non-auth `detail`.
pub fn clear_init_failed(&mut self, name: &str) {
self.init_failed.remove(name);
}
/// Clear the entire handshaking set in one shot. Used by the
/// proxy-mode "init complete" path and the bg-handshake completion
/// path as a defensive sweep after the per-server
/// [`Self::mark_server_ready`] calls; cheap no-op if already empty.
pub fn mark_all_servers_ready(&mut self) {
self.init_progress.clear_handshaking();
}
/// True iff the named server's handshake is still in flight.
/// Used by status snapshots and tool-dispatch gating that need to
/// know per-server progress without cloning the whole set.
pub fn is_server_handshaking(&self, name: &str) -> bool {
self.init_progress.is_server_handshaking(name)
}
/// Iterate over server names whose background handshake is still
/// in flight. Empty when init has not started or has fully
/// completed.
pub fn handshaking_servers_iter(&self) -> impl Iterator<Item = &McpServerName> {
self.init_progress.handshaking_servers()
}
/// Snapshot of the handshaking set as a cloned `HashSet`. Used by
/// the status snapshot API where the caller wants an owned copy
/// that survives lock release.
pub fn handshaking_servers_cloned(&self) -> std::collections::HashSet<McpServerName> {
self.init_progress.handshaking_servers().cloned().collect()
}
/// Number of in-flight per-server handshakes.
pub fn handshaking_servers_count(&self) -> usize {
self.init_progress.handshaking_count()
}
/// Get current generation (for stale check after async init)
pub fn generation(&self) -> u64 {
self.generation
}
/// Replace managed MCP clients whose URL matches a fresh config entry.
///
/// Caller passes `(endpoint, headers)` pairs from whatever source it uses
/// (e.g. shell's cli-chat-proxy `ManagedMcpConfig` cache). The MCP crate
/// stays free of the host's managed-config schema.
///
/// Old `Arc<McpClient>` holders (in-flight tool calls) finish naturally;
/// new calls look up the fresh client from the map.
pub fn refresh_managed_clients<'a, I>(&mut self, fresh_configs: I)
where
I: IntoIterator<Item = (&'a str, &'a HashMap<String, String>)>,
{
let fresh_by_url: HashMap<String, (&'a str, &'a HashMap<String, String>)> = fresh_configs
.into_iter()
.map(|(endpoint, headers)| (normalize_url(endpoint), (endpoint, headers)))
.collect();
for (client_name, client) in &mut self.owned_clients {
let Some(client_url) = self.configs.iter().find_map(|cfg| match cfg {
acp::McpServer::Http(acp::McpServerHttp { name, url, .. })
| acp::McpServer::Sse(acp::McpServerSse { name, url, .. })
if name == client_name =>
{
Some(normalize_url(url))
}
_ => None,
}) else {
continue;
};
let Some(&(fresh_endpoint, fresh_headers)) = fresh_by_url.get(&client_url) else {
continue;
};
if fresh_headers.is_empty() {
continue;
}
// Rebuilding drops the warm connection and forces a full
// re-handshake on next use; skip it when the token is unchanged.
if client.http_headers_match(fresh_headers) {
continue;
}
let headers = fresh_headers
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
*client = Arc::new(McpClient::new_http(
client_name.clone(),
HttpConfig {
url: fresh_endpoint.to_string(),
headers,
},
None,
self.meta_config_map.get(client_name.as_str()),
));
tracing::info!(server = %client_name, "Refreshed managed MCP client with fresh token");
}
}
/// Look up a client by server name.
/// Owned clients take priority (they can override inherited ones).
pub fn get_client(&self, name: &str) -> Option<&Arc<McpClient>> {
self.owned_clients
.get(name)
.or_else(|| self.shared_clients.get(name))
}
/// Iterate over all clients (owned first, then shared — skipping shared
/// entries whose name is overridden by an owned client).
pub fn all_clients(&self) -> impl Iterator<Item = (&McpServerName, &Arc<McpClient>)> {
self.owned_clients.iter().chain(
self.shared_clients
.iter()
.filter(|(name, _)| !self.owned_clients.contains_key(name.as_str())),
)
}
/// Import shared clients from a parent pool snapshot.
/// Clients whose name collides with an agent-definition-owned server
/// are skipped (the owned server takes priority).
pub fn import_shared_clients(&mut self, pool: &SharedMcpPool) {
let config_names: std::collections::HashSet<&str> =
self.configs.iter().map(mcp_server_name).collect();
for (name, client) in &pool.clients {
if !config_names.contains(name.as_str()) {
self.shared_clients.insert(name.clone(), Arc::clone(client));
}
}
}
}
/// Snapshot of an MCP connection pool, taken at subagent spawn time.
///
/// The HashMap is cloned (cheap — values are `Arc<McpClient>`), so the
/// subagent's map is independent of the parent's. The `Arc<McpClient>`
/// entries are shared — both parent and child use the same transport.
/// This is intentionally snapshot-based, not live-updating.
#[derive(Clone)]
pub struct SharedMcpPool {
clients: HashMap<McpServerName, Arc<McpClient>>,
configs: Vec<acp::McpServer>,
meta_config_map: McpMetaConfigMap,
}
impl SharedMcpPool {
/// Create a snapshot from an existing `McpState`.
/// Captures both owned and shared clients (deduped — owned wins).
pub fn from_state(state: &McpState) -> Self {
Self {
clients: state
.all_clients()
.map(|(k, v)| (k.clone(), Arc::clone(v)))
.collect(),
configs: state.configs.clone(),
meta_config_map: state.meta_config_map.clone(),
}
}
pub fn get_client(&self, name: &str) -> Option<&Arc<McpClient>> {
self.clients.get(name)
}
pub fn len(&self) -> usize {
self.clients.len()
}
pub fn is_empty(&self) -> bool {
self.clients.is_empty()
}
pub fn server_names(&self) -> impl Iterator<Item = &str> {
self.clients.keys().map(String::as_str)
}
pub fn configs(&self) -> &[acp::McpServer] {
&self.configs
}
pub fn meta_config_map(&self) -> &McpMetaConfigMap {
&self.meta_config_map
}
/// Retain only clients whose name satisfies `predicate`.
///
/// Only filters the `clients` map. `configs` and `meta_config_map` are
/// left unchanged — callers that need config-level consistency should
/// filter those separately. In the subagent inheritance path this is
/// fine because `import_shared_clients` only iterates `clients`.
pub fn retain_clients(&mut self, predicate: impl Fn(&str) -> bool) {
self.clients.retain(|name, _| predicate(name));
}
}
/// Compare two MCP server config lists for equality.
///
/// Since `acp::McpServer` may not implement PartialEq, we serialize to JSON and compare.
/// This is only called during config updates, so the overhead is acceptable.
pub(crate) fn mcp_servers_equal(a: &[acp::McpServer], b: &[acp::McpServer]) -> bool {
if a.len() != b.len() {
return false;
}
// Compare JSON serializations
match (serde_json::to_string(a), serde_json::to_string(b)) {
(Ok(a_json), Ok(b_json)) => a_json == b_json,
_ => false, // If serialization fails, assume not equal
}
}
/// Default timeout for an MCP server's `initialize` handshake & initial tool
/// listing, used when no override is supplied. 30s is generous enough that
/// cold-start `uvx` / `uv run --with` stdio servers that download deps on
/// first launch aren't killed mid-handshake. The shell resolves env / config /
/// requirements / remote overrides and injects them via `McpClientTimeoutOverrides`.
const DEFAULT_STARTUP_TIMEOUT_SECS: u64 = 30;
/// Default timeout for individual tool calls.
const DEFAULT_TOOL_TIMEOUT_SECS: u64 = 6000;
/// How long a stdio server gets to exit after its transport closes before
/// its process group is killed.
const STDIO_SHUTDOWN_GRACE: std::time::Duration = std::time::Duration::from_secs(3);
/// Timeout for OAuth metadata discovery when building an HTTP transport.
/// Bounds transport setup for servers without OAuth support.
const OAUTH_DISCOVERY_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
/// Per-MCP-server config overrides from `_meta.mcpConfig` in session/new or session/load.
#[derive(Debug, Clone, Default, serde::Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct McpServerMetaConfig {
/// Init handshake timeout in ms. Overrides config.toml `startup_timeout_sec`.
#[serde(default)]
pub startup_timeout_ms: Option<u64>,
/// Per-tool-call timeout in ms. Overrides config.toml `tool_timeout_sec`.
#[serde(default)]
pub tool_timeout_ms: Option<u64>,
/// Per-tool timeout overrides in ms: `{ "create_issue": 120000, "search": 30000 }`.
/// Overrides config.toml `tool_timeouts` (and `tool_timeout_sec`) for matching tools.
#[serde(default)]
pub tool_timeouts_ms: Option<HashMap<ToolName, u64>>,
/// Also keep the raw base64 in tool-result text (in addition to the
/// vision-token rendering) so the agent can decode + forward it via
/// path-based tools like `send_file`. Costs ~2× tokens per image.
/// Default `false`. See [`format_mcp_image`].
#[serde(default)]
pub expose_image_base64: Option<bool>,
}
/// MCP server name → per-server config overrides from `_meta.mcpConfig`.
pub type McpMetaConfigMap = HashMap<McpServerName, McpServerMetaConfig>;
/// Parse `mcpConfig` from a session request's `_meta`. Empty map if absent/invalid.
pub fn parse_mcp_meta_config(
meta: Option<&serde_json::Map<String, serde_json::Value>>,
) -> McpMetaConfigMap {
meta.and_then(|m| m.get("mcpConfig"))
.and_then(|v| serde_json::from_value::<McpMetaConfigMap>(v.clone()).ok())
.unwrap_or_default()
}
/// MCP initialization strategy
#[derive(Debug, Clone, Copy, PartialEq, serde::Serialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum McpInitStrategy {
/// Wait for MCP initialization before first LLM call
#[default]
Blocking,
/// Start immediately, advertise tools as they become available
Progressive,
}
impl<S: AsRef<str>> From<S> for McpInitStrategy {
fn from(s: S) -> Self {
match s.as_ref() {
"progressive" => McpInitStrategy::Progressive,
_ => McpInitStrategy::Blocking,
}
}
}
/// Parse MCP tool name in format "server__tool"
/// Returns (server_name, tool_name) if valid MCP tool, None otherwise
pub fn parse_mcp_tool_name(name: &str) -> Option<(String, String)> {
let parts: Vec<&str> = name.splitn(2, MCP_TOOL_NAME_DELIMITER).collect();
if parts.len() == 2 {
Some((parts[0].to_string(), parts[1].to_string()))
} else {
None
}
}
#[derive(Debug, thiserror::Error)]
pub enum McpError {
#[error("MCP client error: {0}")]
ClientError(String),
#[error("MCP server '{server}' timed out after {timeout_secs}s")]
Timeout { server: String, timeout_secs: u64 },
#[error("Failed to spawn MCP server '{server}': {source}")]
SpawnFailed {
server: String,
source: std::io::Error,
},
#[error("MCP server '{server}' handshake failed: {source}")]
HandshakeFailed {
server: String,
source: Box<ClientInitializeError>,
},
/// Pre-spawn gate: server needs OAuth but this session cannot complete interactive auth.
#[error(
"MCP server '{server}': Auth required (non-interactive session; authenticate in TUI or set Authorization header)"
)]
AuthRequired { server: String },
#[error("MCP service error: {0}")]
ServiceError(#[from] ServiceError),
}
impl McpError {
fn timeout(server: &str, duration: std::time::Duration) -> Self {
Self::Timeout {
server: server.to_string(),
timeout_secs: duration.as_secs(),
}
}
pub fn is_timeout(&self) -> bool {
matches!(self, Self::Timeout { .. })
}
pub fn error_category(&self) -> kigi_file_utils::events::McpErrorCategory {
use kigi_file_utils::events::McpErrorCategory;
match self {
Self::SpawnFailed { .. } => McpErrorCategory::SpawnFailed,
Self::Timeout { .. } => McpErrorCategory::Timeout,
Self::HandshakeFailed { .. } => McpErrorCategory::HandshakeFailed,
Self::AuthRequired { .. } => McpErrorCategory::AuthRequired,
Self::ClientError(_) | Self::ServiceError(_) => McpErrorCategory::ClientError,
}
}
pub fn server_name(&self) -> Option<&str> {
match self {
Self::SpawnFailed { server, .. }
| Self::Timeout { server, .. }
| Self::HandshakeFailed { server, .. }
| Self::AuthRequired { server } => Some(server),
Self::ClientError(_) | Self::ServiceError(_) => None,
}
}
/// True if this error indicates the server rejected us for auth reasons (a
/// credential re-fetch could help). Timeout/spawn failures can't be cured by
/// re-fetching credentials, so they're never auth.
pub fn is_auth_rejection(&self) -> bool {
match self {
Self::AuthRequired { .. } => true,
Self::HandshakeFailed { source, .. } => is_auth_rejection_message(&source.to_string()),
Self::ServiceError(e) => is_auth_rejection_message(&e.to_string()),
Self::ClientError(s) => is_auth_rejection_message(s),
Self::Timeout { .. } | Self::SpawnFailed { .. } => false,
}
}
}
/// True if an MCP error *message* indicates an auth rejection (vs. a transport
/// drop, timeout, or protocol error), so host recovery can decide whether a
/// credential re-fetch would help.
///
/// Matches auth wording and context-anchored 401 patterns only, so a bare digit
/// ("took 401ms", ports) can't trip it. Excludes 403/forbidden — a non-auth
/// policy denial here, not a credential problem.
pub fn is_auth_rejection_message(s: &str) -> bool {
let l = s.to_ascii_lowercase();
// Auth wording has no numeric component, so plain substrings are safe.
if l.contains("auth required")
|| l.contains("authorizationrequired")
|| l.contains("authrequired")
|| l.contains("authentication")
|| l.contains("unauthorized")
{
return true;
}
// Require a non-alphanumeric (or end) after "401" so "http 401" matches but
// "http 4012" (other status) and "http 401ms" (a duration) do not.
[
"status: 401",
"status code 401",
"http status 401",
"http 401",
"error 401",
]
.iter()
.any(|token| token_at_word_boundary(&l, token))
}
/// Whether `haystack` contains `token` at a right word boundary, so a
/// digit-terminated token (`...401`) doesn't match a longer run (`4012`) or an
/// adjacent unit (`401ms`).
///
/// Invariant: `token` must be ASCII (all callers pass ASCII literals). A
/// non-ASCII token could advance `from` mid-UTF-8 and panic on the next slice.
fn token_at_word_boundary(haystack: &str, token: &str) -> bool {
debug_assert!(
token.is_ascii(),
"token_at_word_boundary requires an ASCII token"
);
let mut from = 0;
while let Some(idx) = haystack[from..].find(token) {
let end = from + idx + token.len();
if !haystack[end..].starts_with(|c: char| c.is_ascii_alphanumeric()) {
return true;
}
from += idx + 1;
}
false
}
#[derive(Clone)]
pub struct McpTool {
name: String,
description: String,
server_name: String,
mcp_state: Arc<Mutex<McpState>>,
schema: serde_json::Value,
meta: Option<serde_json::Value>,
}
/// Data needed to register an MCP tool via `register_erased()`.
///
/// MCP tools have two visibility audiences controlled by `_meta.ui.visibility`:
///
/// - **Model-visible** (default, or `["model", "app"]`): registered in `ToolBridge`
/// so the LLM can invoke them during a conversation.
/// - **App-visible only** (`["app"]`): not registered in `ToolBridge`, so the LLM
/// never sees them. These are UI-only actions (e.g. refresh buttons) surfaced to
/// the frontend via `kigi/mcp/tools_changed` notifications and callable via
/// `kigi/mcp/call`.
pub struct McpToolRegistration {
pub name: String,
pub description: String,
pub input_schema: serde_json::Value,
pub tool: McpErasedTool,
pub meta: Option<serde_json::Value>,
pub model_visible: bool,
}
impl McpTool {
/// Reconstruct an `McpTool` from its constituent parts. Used when stashing
/// a disabled tool at runtime so it can be re-enabled without a full re-init.
pub fn new(
name: String,
description: String,
server_name: String,
mcp_state: Arc<Mutex<McpState>>,
schema: serde_json::Value,
meta: Option<serde_json::Value>,
) -> Self {
Self {
name,
description,
server_name,
mcp_state,
schema,
meta,
}
}
/// Convert into the data needed for `ToolBridge::register_erased()`.
///
/// Returns `None` if the tool name is invalid (doesn't match LLM API requirements).
/// Invalid tools are logged and skipped — fix the upstream connector.
///
/// Also rejects qualified names that contain the delimiter
/// (`MCP_TOOL_NAME_DELIMITER`) more than once. The underlying tool-name
/// regex permits underscores in each segment, so a server like
/// `"foo__bar"`, a tool like `"my__thing"`, or even a `"foo_"`/`"_bar"`
/// pair (which concatenates to `"foo___bar"` — two valid `__`
/// positions) would produce a qualified name that downstream
/// `split_once("__")` consumers would split at the wrong boundary.
/// The "exactly one delimiter" check covers all three cases with a
/// single rule.
pub fn into_registration(self) -> Option<McpToolRegistration> {
// Qualify MCP tool name with server name: "server__tool"
let qualified_name = format!(
"{}{}{}",
self.server_name, MCP_TOOL_NAME_DELIMITER, self.name
);
// Reject ambiguous qualified names — see doc-comment above.
if qualified_name.matches(MCP_TOOL_NAME_DELIMITER).count() != 1 {
tracing::error!(
server = %self.server_name,
tool = %self.name,
qualified = %qualified_name,
"Skipping MCP tool: qualified name contains '{MCP_TOOL_NAME_DELIMITER}' more than once (server, tool, or their boundary collides with the reserved delimiter)"
);
return None;
}
if let Err(reason) = validate_tool_name(&qualified_name) {
tracing::error!(
tool_name = %qualified_name,
server = %self.server_name,
reason = %reason,
"Skipping MCP tool with invalid name"
);
return None;
}
let description = self.description.clone();
let input_schema = self.schema.clone();
let meta = self.meta.clone();
let model_visible = meta
.as_ref()
.and_then(|m| m.get("ui"))
.and_then(|ui| ui.get("visibility"))
.and_then(|v| v.as_array())
.map(|arr| arr.iter().any(|s| s.as_str() == Some("model")))
.unwrap_or(true); // default: visible to model
Some(McpToolRegistration {
name: qualified_name,
description,
input_schema,
tool: McpErasedTool { tool: self },
meta,
model_visible,
})
}
}
/// MCP tool wrapper for runtime dispatch.
///
/// MCP tools are already untyped (JSON → JSON), so they implement
/// `kigi_tool_runtime::Tool` directly instead of going through typed wrappers.
pub struct McpErasedTool {
tool: McpTool,
}
impl std::fmt::Debug for McpErasedTool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpErasedTool")
.field("name", &self.tool.name)
.field("server", &self.tool.server_name)
.finish()
}
}
impl ToolMetadata for McpErasedTool {
fn kind(&self) -> ToolKind {
ToolKind::Other
}
fn tool_namespace(&self) -> ToolNamespace {
ToolNamespace::MCP
}
fn description_template(&self) -> &str {
&self.tool.description
}
}
impl kigi_tool_runtime::Tool for McpErasedTool {
type Args = serde_json::Value;
type Output = ToolOutput;
fn id(&self) -> kigi_tool_protocol::ToolId {
// Use the qualified name (server__tool) so that two MCP servers
// exposing the same raw tool name get distinct LocalRegistry entries.
let qualified = format!(
"{}{}{}",
self.tool.server_name, MCP_TOOL_NAME_DELIMITER, self.tool.name
);
kigi_tool_protocol::ToolId::new(&qualified)
.unwrap_or_else(|_| kigi_tool_protocol::ToolId::new("mcp_tool").expect("valid"))
}
fn description(
&self,
_ctx: &::kigi_tool_runtime::ListToolsContext,
) -> kigi_tool_types::ToolDescription {
kigi_tool_types::ToolDescription::new(&self.tool.name, &self.tool.description)
}
async fn run(
&self,
_ctx: kigi_tool_runtime::ToolCallContext,
raw: serde_json::Value,
) -> Result<ToolOutput, kigi_tool_runtime::ToolError> {
let mcp_call_start = std::time::Instant::now();
let (client, event_writer) = {
let state = self.tool.mcp_state.lock().await;
let c = Arc::clone(state.get_client(&self.tool.server_name).ok_or_else(|| {
kigi_tool_runtime::ToolError::custom(
"process_manager",
format!("MCP server '{}' not found", self.tool.server_name),
)
})?);
(c, state.event_writer().clone())
};
let server = &self.tool.server_name;
let tool = &self.tool.name;
let tool_timeout = client.tool_timeout_for(tool);
let qualified_name = format!("{}{}{}", server, MCP_TOOL_NAME_DELIMITER, tool);
event_writer.emit(kigi_file_utils::events::Event::McpToolCallStarted {
server_name: server.clone(),
tool_name: tool.clone(),
call_id: qualified_name.clone(),
timeout_sec: tool_timeout,
});
let mut auth_retry_attempted = false;
let mut reconnect_attempted = false;
let mut is_timeout = false;
let ew = &event_writer;
let dispatch_result = match self
.try_call_tool(&client, &raw, &mut reconnect_attempted, &mut is_timeout, ew)
.await
{
Ok(result) => Ok(result),
Err(first_err) if client.has_auth() => {
auth_retry_attempted = true;
let reauth_ok = client.force_reauth(false).await;
ew.emit(kigi_file_utils::events::Event::McpAuthRetry {
server_name: server.clone(),
trigger: "tool_call_failed".to_string(),
success: reauth_ok,
});
if reauth_ok {
self.try_call_tool(&client, &raw, &mut reconnect_attempted, &mut is_timeout, ew)
.await
.map_err(|e| {
kigi_tool_runtime::ToolError::custom("process_manager", e.to_string())
})
} else {
Err(first_err)
}
}
Err(e) => Err(e),
};
let call_result = match dispatch_result {
Ok(result) => result,
Err(e) => {
ew.emit(kigi_file_utils::events::Event::McpToolCallCompleted {
server_name: server.clone(),
tool_name: tool.clone(),
call_id: qualified_name,
duration_ms: mcp_call_start.elapsed().as_millis() as u64,
success: false,
is_timeout,
error: Some(e.to_string()),
reconnect_attempted,
auth_retry_attempted,
});
return Err(e);
}
};
let is_error = call_result.is_error.unwrap_or(false);
let mut output = if is_error {
let error_msg = call_result
.content
.iter()
.filter_map(|c| match c {
rmcp::model::ContentBlock::Text(t) => Some(t.text.clone()),
_ => None,
})
.collect::<Vec<_>>()
.join("\n");
ToolOutput::MCP(MCPOutput::errored(tool.clone(), server.clone(), error_msg))
} else {
let expose_base64 = client.expose_image_base64();
let parts: Vec<String> = call_result
.content
.into_iter()
.filter_map(|c| match c {
rmcp::model::ContentBlock::Text(t) => Some(t.text),
rmcp::model::ContentBlock::Image(img) => {
Some(format_mcp_image(&img.mime_type, &img.data, expose_base64))
}
rmcp::model::ContentBlock::Resource(r) => match &r.resource {
rmcp::model::ResourceContents::BlobResourceContents {
mime_type,
blob,
..
} if mime_type
.as_deref()
.is_some_and(|m| m.starts_with("image/")) =>
{
let mime = mime_type.as_deref().unwrap();
Some(format_mcp_image(mime, blob, expose_base64))
}
_ => serde_json::to_string(&r).ok(),
},
_ => None,
})
.collect();
let text = parts.join("\n");
ToolOutput::MCP(MCPOutput::okay_output(tool.clone(), server.clone(), text))
};
if let ToolOutput::MCP(ref mut mcp_out) = output {
mcp_out.auth_retry_attempted = auth_retry_attempted;
mcp_out.reconnect_attempted = reconnect_attempted;
mcp_out.is_timeout = is_timeout;
}
let success = !is_error;
let duration_ms = mcp_call_start.elapsed().as_millis() as u64;
let error_text = if is_error {
match &output {
ToolOutput::MCP(mcp) => match mcp.output() {
MCPOutputDetails::Error(e) => Some(e.clone()),
_ => None,
},
_ => None,
}
} else {
None
};
event_writer.emit(kigi_file_utils::events::Event::McpToolCallCompleted {
server_name: server.clone(),
tool_name: tool.clone(),
call_id: qualified_name,
duration_ms,
success,
is_timeout,
error: error_text,
reconnect_attempted,
auth_retry_attempted,
});
Ok(output)
}
}
/// Render an MCP image content block. The data URI is consumed by the
/// session-layer `extract_base64_images` and rendered as vision tokens.
/// When `expose_base64`, also emit a `<mcp_image_base64>` wrapper that
/// survives extraction (wrapper has no `data:image/` prefix → regex skips
/// it), exposing the raw bytes to the agent for path-based forwarding.
fn format_mcp_image(mime: &str, base64_data: &str, expose_base64: bool) -> String {
if expose_base64 {
format!(
"data:{mime};base64,{base64_data}\n\
<mcp_image_base64 mime=\"{mime}\">\n\
{base64_data}\n\
</mcp_image_base64>"
)
} else {
format!("data:{mime};base64,{base64_data}")
}
}
/// Check whether a `ServiceError` indicates the underlying transport has died
/// and a fresh connection could recover it.
fn is_retriable_transport_error(err: &ServiceError) -> bool {
matches!(
err,
ServiceError::TransportClosed | ServiceError::TransportSend(_)
)
}
/// Recover for every JSON-RPC code except the deterministic client set
/// {-32700, -32600, -32601, -32602} (those mean the request was wrong, not the session).
fn should_recover_mcp_error(code: i32) -> bool {
use rmcp::model::ErrorCode;
let deterministic_client_error = code == ErrorCode::PARSE_ERROR.0
|| code == ErrorCode::INVALID_REQUEST.0
|| code == ErrorCode::METHOD_NOT_FOUND.0
|| code == ErrorCode::INVALID_PARAMS.0;
!deterministic_client_error
}
/// Recovers transport errors, and an HTTP `McpError` once per dispatch except
/// deterministic client codes and auth-class errors — a rebuild reuses stale
/// creds, so auth is routed to the re-auth paths instead.
fn should_recover_service_error(
err: &ServiceError,
is_http: bool,
reconnect_attempted: bool,
) -> bool {
is_retriable_transport_error(err)
|| matches!(
err,
ServiceError::McpError(e)
if is_http
&& !reconnect_attempted
&& should_recover_mcp_error(e.code.0)
&& !is_auth_rejection_message(e.message.as_ref())
)
}
impl McpErasedTool {
async fn try_call_tool(
&self,
client: &Arc<McpClient>,
raw: &serde_json::Value,
reconnect_attempted: &mut bool,
is_timeout: &mut bool,
ew: &kigi_file_utils::events::EventWriter,
) -> Result<rmcp::model::CallToolResult, kigi_tool_runtime::ToolError> {
let mcp_service = client
.ensure_initialized()
.await
.map_err(|e| kigi_tool_runtime::ToolError::custom("process_manager", e.to_string()))?;
let tool_timeout = client.tool_timeout_for(&self.tool.name);
let timeout_duration = std::time::Duration::from_secs(tool_timeout);
let mut params = CallToolRequestParams::new(self.tool.name.clone());
params.arguments = raw.as_object().cloned();
let result =
tokio::time::timeout(timeout_duration, mcp_service.call_tool(params.clone())).await;
match result {
Ok(Ok(call_result)) => Ok(call_result),
Ok(Err(service_err))
if should_recover_service_error(
&service_err,
client.is_http(),
*reconnect_attempted,
) =>
{
self.recover_and_retry(
client,
params,
timeout_duration,
tool_timeout,
service_err,
reconnect_attempted,
is_timeout,
ew,
)
.await
}
Ok(Err(e)) => Err(kigi_tool_runtime::ToolError::custom(
"process_manager",
e.to_string(),
)),
Err(_) => {
*is_timeout = true;
// Reset for the next call but don't retry — a slow side-effecting tool must not run twice.
if client.is_http() && !*reconnect_attempted {
client.reset_transport().await;
*reconnect_attempted = true;
}
Err(kigi_tool_runtime::ToolError::custom(
"process_manager",
format!(
"MCP tool '{}' timed out after {} seconds",
self.tool.name, tool_timeout
),
))
}
}
}
/// On `recover()` failure surface the original error, else the retry error
/// (preserves the auth signal managed re-auth reads from the string).
#[allow(clippy::too_many_arguments)]
async fn recover_and_retry(
&self,
client: &Arc<McpClient>,
params: CallToolRequestParams,
timeout_duration: std::time::Duration,
tool_timeout: u64,
original_err: ServiceError,
reconnect_attempted: &mut bool,
is_timeout: &mut bool,
ew: &kigi_file_utils::events::EventWriter,
) -> Result<rmcp::model::CallToolResult, kigi_tool_runtime::ToolError> {
*reconnect_attempted = true;
tracing::warn!(
server = self.tool.server_name.as_str(),
tool = self.tool.name.as_str(),
error = %original_err,
"MCP transport error, attempting reconnect"
);
ew.emit(kigi_file_utils::events::Event::McpTransportError {
server_name: self.tool.server_name.clone(),
tool_name: self.tool.name.clone(),
error: original_err.to_string(),
});
let mcp_service = match client.recover().await {
Ok(service) => {
ew.emit(kigi_file_utils::events::Event::McpTransportReconnect {
server_name: self.tool.server_name.clone(),
success: true,
error: None,
});
service
}
Err(e) => {
ew.emit(kigi_file_utils::events::Event::McpTransportReconnect {
server_name: self.tool.server_name.clone(),
success: false,
error: Some(e.to_string()),
});
return Err(kigi_tool_runtime::ToolError::custom(
"process_manager",
original_err.to_string(),
));
}
};
match tokio::time::timeout(timeout_duration, mcp_service.call_tool(params)).await {
Ok(Ok(call_result)) => Ok(call_result),
Ok(Err(retry_err)) => Err(kigi_tool_runtime::ToolError::custom(
"process_manager",
retry_err.to_string(),
)),
Err(_) => {
*is_timeout = true;
Err(kigi_tool_runtime::ToolError::custom(
"process_manager",
format!(
"MCP tool '{}' timed out after {} seconds",
self.tool.name, tool_timeout
),
))
}
}
}
}
/// Whether this session can complete an interactive (browser) OAuth flow.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OauthInteractivity {
Interactive,
NonInteractive,
}
impl OauthInteractivity {
/// Headless/SDK sessions set `non_interactive = true` and cannot complete browser OAuth.
pub fn from_non_interactive(non_interactive: bool) -> Self {
if non_interactive {
Self::NonInteractive
} else {
Self::Interactive
}
}
}
/// Outcome of probing whether an HTTP/SSE MCP server needs OAuth and whether
/// we have credentials usable without an interactive browser flow.
enum HttpOauthPrep {
/// Server does not advertise OAuth (or discovery failed conservatively).
NoOauthSupport,
/// Ready to connect with an auth manager (stored token works, or interactive deferred auth).
ManagerReady(Arc<tokio::sync::Mutex<rmcp::transport::auth::AuthorizationManager>>),
/// OAuth is required but cannot complete in non-interactive mode — do not start unauthenticated.
NeedsInteractiveLogin,
}
impl HttpOauthPrep {
/// Inconclusive OAuth probe (manager-create error, discovery error, or timeout):
/// interactive proceeds as plain HTTP; non-interactive fails closed to avoid rmcp
/// auth-worker stderr noise.
fn on_probe_failure(mode: OauthInteractivity) -> Self {
match mode {
OauthInteractivity::Interactive => Self::NoOauthSupport,
OauthInteractivity::NonInteractive => Self::NeedsInteractiveLogin,
}
}
}
/// Proactive OAuth discovery per RFC 8414 + 9728.
///
/// Creates an `AuthorizationManager` with our `CredentialStoreAdapter`,
/// discovers server metadata, and loads stored tokens if available.
///
/// With no stored tokens but server OAuth support, behavior splits on `mode`:
/// `Interactive` spawns the browser flow in the background (non-blocking; the
/// first tool call picks up tokens via `force_reauth` → `initialize_from_store`
/// once the user consents), while `NonInteractive` fails closed
/// (`NeedsInteractiveLogin`) rather than start an unauthenticated worker that
/// fatals with `Auth(AuthorizationRequired)` on stderr while the prompt still
/// succeeds.
///
/// Known gap: rmcp `get_access_token` returns stored tokens as-is when expiry
/// metadata is absent, so an expiry-less revoked token can still reach
/// `ManagerReady`.
async fn discover_and_prepare_auth(
server_name: &str,
server_url: &str,
mode: OauthInteractivity,
) -> HttpOauthPrep {
let Ok(parsed_url) = url::Url::parse(server_url) else {
return HttpOauthPrep::NoOauthSupport;
};
let adapter =
crate::credentials::McpCredentialStoreAdapter::new(server_name.to_string(), parsed_url);
let mut manager = match rmcp::transport::auth::AuthorizationManager::new(server_url).await {
Ok(m) => m,
Err(e) => {
tracing::warn!(server = server_name, %e, "Failed to create OAuth manager");
// Non-interactive: fail closed — unauthenticated HTTP may still fatal in rmcp.
return HttpOauthPrep::on_probe_failure(mode);
}
};
manager.set_credential_store(adapter);
if let Ok(true) = manager.initialize_from_store().await {
// Stored creds may be expired/unrefreshable; in non-interactive mode that
// still yields rmcp worker fatal AuthorizationRequired on stderr. This probe
// shares the caller's 5s discovery timeout budget.
if mode == OauthInteractivity::NonInteractive
&& let Err(e) = manager.get_access_token().await
{
tracing::warn!(
server = server_name,
error = %e,
"Skipping OAuth MCP in non-interactive mode (stored credentials unusable); re-authenticate in TUI"
);
return HttpOauthPrep::NeedsInteractiveLogin;
}
tracing::info!(server = server_name, "Loaded stored OAuth credentials");
return HttpOauthPrep::ManagerReady(Arc::new(tokio::sync::Mutex::new(manager)));
}
match manager.discover_metadata().await {
Ok(metadata) => {
manager.set_metadata(metadata);
if mode == OauthInteractivity::NonInteractive {
tracing::warn!(
server = server_name,
"Skipping OAuth MCP in non-interactive mode (no stored tokens); authenticate in TUI or set an Authorization header"
);
return HttpOauthPrep::NeedsInteractiveLogin;
}
tracing::info!(
server = server_name,
"Server supports OAuth but has no stored tokens"
);
HttpOauthPrep::ManagerReady(Arc::new(tokio::sync::Mutex::new(manager)))
}
Err(rmcp::transport::auth::AuthError::NoAuthorizationSupport) => {
tracing::debug!(server = server_name, "Server does not support OAuth");
HttpOauthPrep::NoOauthSupport
}
Err(e) => {
tracing::warn!(server = server_name, %e, "OAuth discovery failed");
HttpOauthPrep::on_probe_failure(mode)
}
}
}
/// Configuration for HTTP MCP server connection.
#[derive(Clone)]
pub struct HttpConfig {
pub url: String,
pub headers: Vec<(String, String)>,
}
/// Newline-delimited JSON-RPC stdio transport whose read side survives a
/// single undecodable line.
///
/// Used instead of rmcp's `AsyncRwTransport` for two reasons:
/// - **Wire silence:** a bad line is skipped without replying, whereas rmcp
/// answers shape-mismatched JSON with a -32600 error — a reply an off-spec
/// server could echo back as more invalid input.
/// - **Telemetry:** each skip emits an `McpTransportDecodeError` event (with a
/// truncated sample of the offending line) so the failure is visible in the
/// session trace — rmcp's own tracing is not captured there.
///
/// We read lines ourselves (rather than via `FramedRead` + rmcp's codec) so
/// reading continues after a bad line; only a genuine end-of-stream returns
/// `None`. A stray non-JSON stdout line, a JSON-RPC batch array, or an
/// off-spec response therefore never collapses the transport ("Transport
/// closed" failing every in-flight request — the "connector shows but doesn't
/// work" report).
///
/// Generic over `R`/`W` so it can be unit-tested with in-memory pipes; the
/// production transport binds `ChildStdout`/`ChildStdin`.
struct ResilientRwTransport<R, W>
where
R: AsyncRead,
W: AsyncWrite,
{
read: BufReader<R>,
/// `Arc<Mutex<Option<…>>>` so `send` can return a `Send + 'static` future
/// (the `Transport` contract) without borrowing `self`, and so `close` can
/// drop the writer — mirrors rmcp's own `AsyncRwTransport`.
write: Arc<Mutex<Option<W>>>,
server_name: String,
event_writer: kigi_file_utils::events::EventWriter,
}
/// Max bytes of an offending line copied into the decode-error event.
const DECODE_ERROR_SAMPLE_LEN: usize = 200;
/// A line that failed to deserialize but is a JSON *notification* (an object
/// with a `method` and no `id`) is benign — many servers emit non-MCP / unknown
/// notifications (e.g. LSP-style). Skip those quietly instead of flagging a
/// decode error, mirroring rmcp's compatibility handling.
fn is_ignorable_notification(line: &[u8]) -> bool {
match serde_json::from_slice::<serde_json::Value>(line) {
Ok(v) => v.get("id").is_none() && v.get("method").and_then(|m| m.as_str()).is_some(),
Err(_) => false,
}
}
impl<R, W> ResilientRwTransport<R, W>
where
R: AsyncRead + Send + Unpin,
W: AsyncWrite + Send + Unpin + 'static,
{
fn new(
read: R,
write: W,
server_name: String,
event_writer: kigi_file_utils::events::EventWriter,
) -> Self {
Self {
read: BufReader::new(read),
write: Arc::new(Mutex::new(Some(write))),
server_name,
event_writer,
}
}
/// Record a skipped, undecodable stdout line: a `warn!` log plus an
/// `McpTransportDecodeError` event carrying the serde error and a truncated
/// sample of the raw line (the diagnostic the untagged-enum serde error
/// alone lacks).
fn record_decode_error(&self, line: &[u8], err: &serde_json::Error) {
let sample: String = String::from_utf8_lossy(line)
.chars()
.take(DECODE_ERROR_SAMPLE_LEN)
.collect();
tracing::warn!(
server = %self.server_name,
error = %err,
sample = %sample,
"Skipping undecodable MCP stdout line; keeping transport alive",
);
self.event_writer
.emit(kigi_file_utils::events::Event::McpTransportDecodeError {
server_name: self.server_name.clone(),
error: err.to_string(),
sample,
});
}
}
impl<R, W> Transport<RoleClient> for ResilientRwTransport<R, W>
where
R: AsyncRead + Send + Unpin,
W: AsyncWrite + Send + Unpin + 'static,
{
type Error = std::io::Error;
fn send(
&mut self,
item: TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let lock = self.write.clone();
async move {
let mut bytes = serde_json::to_vec(&item).map_err(std::io::Error::other)?;
bytes.push(b'\n');
let mut guard = lock.lock().await;
match guard.as_mut() {
Some(write) => {
write.write_all(&bytes).await?;
write.flush().await
}
None => Err(std::io::Error::new(
std::io::ErrorKind::NotConnected,
"transport is closed",
)),
}
}
}
async fn receive(&mut self) -> Option<RxJsonRpcMessage<RoleClient>> {
loop {
let mut line = Vec::new();
match self.read.read_until(b'\n', &mut line).await {
Ok(0) => return None, // genuine end-of-stream
Ok(_) => {}
Err(e) => {
tracing::debug!(
server = %self.server_name,
error = %e,
"MCP stdio read error; closing transport",
);
return None;
}
}
if line.last() == Some(&b'\n') {
line.pop();
}
if line.last() == Some(&b'\r') {
line.pop();
}
if line.is_empty() {
continue;
}
match serde_json::from_slice::<RxJsonRpcMessage<RoleClient>>(&line) {
Ok(msg) => return Some(msg),
// The whole point: a single undecodable line must not
// collapse the transport — skip it and keep reading.
Err(err) => {
if is_ignorable_notification(&line) {
tracing::trace!(
server = %self.server_name,
"Ignoring unrecognized MCP notification",
);
} else {
self.record_decode_error(&line, &err);
}
continue;
}
}
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
let mut guard = self.write.lock().await;
drop(guard.take());
Ok(())
}
}
/// Stdio MCP transport with a non-panicking cleanup path.
///
/// Unlike `rmcp`'s `TokioChildProcess` (which `tokio::spawn`s from `Drop` and so
/// panics when dropped without an entered runtime), this wrapper's `Drop` is
/// best-effort: it reaps via the current runtime if present, else a short-lived
/// cleanup thread, so the child never leaks as a zombie. Since the caller's
/// `detach_command` `setsid`s the child into its own group, teardown also
/// `killpg`s the whole group via [`ProcessGroup`] to avoid orphaning
/// grandchildren (e.g. `npx` -> `node`) before reaping the leader.
pub struct SafeTokioChildProcess {
child: Option<tokio::process::Child>,
process_group: Option<ProcessGroup>,
transport: ResilientRwTransport<tokio::process::ChildStdout, tokio::process::ChildStdin>,
}
impl SafeTokioChildProcess {
/// `server_name` + `event_writer` are threaded into the transport so a
/// skipped (undecodable) stdout line emits an `McpTransportDecodeError`
/// event for that server.
fn spawn(
mut cmd: Command,
server_name: String,
event_writer: kigi_file_utils::events::EventWriter,
) -> std::io::Result<(Self, Option<ChildStderr>)> {
cmd.stdin(std::process::Stdio::piped())
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn()?;
let stdin = child
.stdin
.take()
.ok_or_else(|| std::io::Error::other("stdin was already taken"))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| std::io::Error::other("stdout was already taken"))?;
let stderr = child.stderr.take();
// Best-effort: a missing group just degrades to direct-child-only cleanup.
let process_group = match ProcessGroup::new() {
Ok(mut group) => match group.attach(&child) {
Ok(()) => Some(group),
Err(e) => {
tracing::warn!("Failed to attach MCP child to process group: {e}");
None
}
},
Err(e) => {
tracing::warn!("Failed to create MCP child process group: {e}");
None
}
};
Ok((
Self {
child: Some(child),
process_group,
transport: ResilientRwTransport::new(stdout, stdin, server_name, event_writer),
},
stderr,
))
}
fn id(&self) -> Option<u32> {
self.child.as_ref()?.id()
}
/// SIGKILLs the whole process group (child + grandchildren). Synchronous, so
/// it's safe from `Drop`; the leader still needs reaping afterwards.
fn kill_process_group(&self) {
if let Some(group) = &self.process_group
&& let Err(e) = group.kill()
{
tracing::warn!("Error killing MCP child process group: {e}");
}
}
async fn graceful_shutdown(&mut self) -> std::io::Result<()> {
let Some(mut child) = self.child.take() else {
return Ok(());
};
self.transport.close().await?;
let result = tokio::select! {
_ = tokio::time::sleep(STDIO_SHUTDOWN_GRACE) => {
self.kill_process_group();
match child.kill().await {
Ok(()) => Ok(()),
Err(e) => {
tracing::warn!("Error killing MCP child: {e}");
Err(e)
}
}
}
res = child.wait() => {
// Reap any grandchildren now, while the pgid is still kept alive
// by them and before the reaped leader's pid can be reused.
self.kill_process_group();
match res {
Ok(status) => {
tracing::info!("MCP child exited gracefully {status}");
Ok(())
}
Err(e) => {
tracing::warn!("Error waiting for MCP child: {e}");
Err(e)
}
}
}
};
// Leader is reaped; drop the group so `Drop` can't `killpg` a now-reusable pid.
self.process_group = None;
result
}
}
impl Drop for SafeTokioChildProcess {
fn drop(&mut self) {
// Group teardown is synchronous, so do it first regardless of runtime.
self.kill_process_group();
let Some(mut child) = self.child.take() else {
return;
};
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
if let Err(e) = child.kill().await {
tracing::warn!("Error killing MCP child process: {e}");
}
});
} else if let Err(e) = std::thread::Builder::new()
.name("mcp-stdio-child-cleanup".to_string())
.spawn(move || {
let rt = match tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
{
Ok(rt) => rt,
Err(e) => {
tracing::warn!("Error creating runtime to clean up MCP child process: {e}");
if let Err(e) = child.start_kill() {
tracing::warn!("Error signaling MCP child process during drop: {e}");
}
return;
}
};
rt.block_on(async move {
if let Err(e) = child.kill().await {
tracing::warn!("Error killing MCP child process: {e}");
}
});
})
{
tracing::warn!("Error spawning MCP child cleanup thread: {e}");
}
}
}
impl Transport<RoleClient> for SafeTokioChildProcess {
type Error = std::io::Error;
fn send(
&mut self,
item: TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
self.transport.send(item)
}
fn receive(&mut self) -> impl Future<Output = Option<RxJsonRpcMessage<RoleClient>>> + Send {
self.transport.receive()
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.graceful_shutdown().await
}
}
/// Transport configuration before connection is established.
enum PendingTransport {
Stdio(Box<SafeTokioChildProcess>),
Http(HttpConfig),
HttpAuth {
config: HttpConfig,
auth_manager: Arc<tokio::sync::Mutex<rmcp::transport::auth::AuthorizationManager>>,
},
/// In-process SDK MCP server reached over the ACP reverse channel
/// (`kigi/mcp/sdk_call`). Rebuildable from its `server_id` + invoker, so handshake
/// failures restore like Http (unlike the consumed Stdio child).
Acp {
server_id: String,
invoker: Arc<dyn crate::acp_transport::AcpReverseInvoker>,
},
}
/// A connected MCP service (rmcp's RunningService wrapped in Arc).
/// Uses [`KigiClientHandler`] rather than rmcp's default `ClientInfo`
/// handler: rmcp 2.1 parameterizes `RunningService` over the handler
/// type, and `ClientInfo` is only a `ClientHandler` impl with no
/// notification routing. The custom handler keeps the same protocol
/// behavior (same `get_info`) while plumbing
/// `tools/list_changed` / `resources/list_changed` notifications
/// through to the session-actor dispatcher.
pub type McpService = Arc<RunningService<RoleClient, KigiClientHandler>>;
/// MCP client connection state machine.
///
/// Single-flight handshake invariant: at most one task at a time may run
/// the handshake. While the handshake is in flight the state is
/// [`ClientState::Initializing`]; the holder owns the transport for the
/// duration of [`McpClient::try_handshake`]. Concurrent callers of
/// [`McpClient::ensure_initialized`] observe [`ClientState::Initializing`]
/// and park on [`McpClient::init_done`] until the holder publishes a
/// result, instead of failing fast with
/// `"MCP client already initializing"` as in earlier versions.
enum ClientState {
/// No transport configured. Reachable from:
/// - [`McpClient::stub`] (test placeholder; `ensure_initialized`
/// returns a configuration error).
/// - Stdio handshake failure (the spawned child process is consumed
/// by `client.serve` and cannot be reused — Http/HttpAuth keep
/// their `HttpConfig` clone and transition back to `Pending`).
Empty,
/// Transport is configured and ready for the next handshake.
Pending(PendingTransport),
/// A caller is currently inside [`McpClient::try_handshake`] and owns
/// the transport. New callers MUST park on
/// [`McpClient::init_done`] (with a bounded timeout) rather than
/// attempt a parallel handshake. If the holder is cancelled or
/// panics before publishing a result, an [`InitGuard`] restores the
/// transport on a best-effort basis so other callers can retry.
Initializing,
/// Handshake completed; the service is reference-counted via `Arc`.
Ready(McpService),
}
/// `Copy` projection of [`ClientState`] used for cheap state-machine
/// inspection (see [`McpClient::state_kind`]). Mirrors the variants
/// 1:1, dropping the payloads so callers can pattern-match without
/// borrowing the state mutex's inner data.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ClientStateKind {
Empty,
Pending,
Initializing,
Ready,
}
/// Classification used by [`crate::liveness::spawn_transport_liveness`].
///
/// Returned by [`McpClient::liveness_check`] under a single state-mutex
/// acquisition, so the watcher's per-tick predicate is atomic.
///
/// `Transient` covers states the watcher should silently exit on
/// (re-handshake races, externally-reset transports, post-failure
/// `Empty` slots). Only `Ready + transport closed` produces a
/// `TransportClosed` ACP push.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LivenessCheck {
/// `Ready` + `is_transport_closed() == false`. Keep polling.
Healthy,
/// `Ready` + `is_transport_closed() == true`. Emit + exit.
TransportClosed,
/// Anything else (`Initializing`, `Pending`, `Empty`). The
/// watcher exits silently — the new state is being managed
/// externally; if it returns to `Ready` the owner can re-arm.
Transient,
}
/// Events emitted by a live MCP client to its session-side dispatcher.
///
/// Produced by three sources:
///
/// 1. [`crate::liveness::spawn_transport_liveness`] when an `is_healthy`
/// poll observes that the rmcp service loop has shut down its receiver
/// (`TransportClosed`).
/// 2. [`KigiClientHandler`] when the server pushes a notification we
/// care about — currently `notifications/tools/list_changed` and
/// `notifications/resources/list_changed`.
/// 3. The session/managed-config layer when a server is added, removed,
/// or successfully (re-)initialized.
///
/// Consumers fan these out to ACP `kigi/mcp/server_status` after 50 ms
/// of tumbling-window coalescing keyed by `(server, kind)`; see the
/// session-actor `StatusDispatcher`.
#[derive(Debug, Clone)]
pub enum McpClientEvent {
/// The rmcp service loop has terminated; the client is no longer
/// usable for tool calls and must be torn down (or restarted).
TransportClosed {
server: McpServerName,
/// Identity of the client whose transport closed (see
/// [`McpClient::client_id`]). A mismatch with the client
/// currently registered under `server` marks the event stale —
/// it must not tear down the replacement. Every emitter holds the
/// closing `McpClient`, so the id is always known.
client_id: u64,
},
/// `ensure_initialized` returned `Err(_)`; `reason` is the full
/// stringified error, surfaced verbatim to the client (no
/// sanitization) so failures are easy to debug.
HandshakeFailed {
server: McpServerName,
reason: String,
},
/// Server pushed `notifications/tools/list_changed`.
ToolsChanged { server: McpServerName },
/// Server pushed `notifications/resources/list_changed`.
ResourcesChanged { server: McpServerName },
/// Client transitioned to [`ClientState::Ready`]; dispatcher uses
/// this to surface "ready" status without polling. Emitted from
/// `ensure_initialized`; the dispatcher maps it to
/// `reason=initialized` (NOT `reason=restart_succeeded`, which is
/// reserved for the restart path).
Ready { server: McpServerName },
/// Managed/local config diff resolved. The dispatcher fans this
/// out into one [`Self::ConfigAdded`] / [`Self::ConfigRemoved`]
/// event per affected server before buffering.
ConfigDiff {
added: Vec<McpServerName>,
removed: Vec<McpServerName>,
},
/// Per-server `(server, ConfigAdded)` fan-out variant produced by
/// the dispatcher from a [`Self::ConfigDiff`]. Keeps the
/// `kind ↔ event payload` invariant: storing a fake `Ready`
/// payload at a `ConfigAdded` key would be a footgun whenever a
/// real `Ready` and a `ConfigDiff` collided in the same coalesce
/// window.
ConfigAdded { server: McpServerName },
/// Per-server `(server, ConfigRemoved)` fan-out variant — the
/// dispatched analogue of [`Self::ConfigAdded`] for the removed
/// set of a [`Self::ConfigDiff`].
ConfigRemoved { server: McpServerName },
}
/// Discriminant for [`McpClientEvent`], used as the second half of the
/// coalescing key `(server, kind)`. Two events with the same
/// `(server, kind)` collapse into the latest one inside the
/// dispatcher's 50 ms window.
///
/// Distinct from [`McpClientEvent`] because the latter carries
/// payload (e.g. `reason` on `HandshakeFailed`) that we don't want
/// participating in equality / hashing.
#[derive(Debug, Clone, Copy, Hash, Eq, PartialEq)]
pub enum McpClientEventKind {
TransportClosed,
HandshakeFailed,
ToolsChanged,
ResourcesChanged,
Ready,
ConfigAdded,
ConfigRemoved,
}
impl McpClientEvent {
/// Server name carried by the event, if any.
///
/// Returns `None` only for [`McpClientEvent::ConfigDiff`] — that
/// variant is fanned out per-server by the dispatcher into
/// [`Self::ConfigAdded`] / [`Self::ConfigRemoved`], where each
/// fan-out child has a single server name.
pub fn server_name(&self) -> Option<&str> {
match self {
Self::TransportClosed { server, .. }
| Self::HandshakeFailed { server, .. }
| Self::ToolsChanged { server }
| Self::ResourcesChanged { server }
| Self::Ready { server }
| Self::ConfigAdded { server }
| Self::ConfigRemoved { server } => Some(server.as_str()),
Self::ConfigDiff { .. } => None,
}
}
}
/// RAII guard that restores [`ClientState::Pending`] if the
/// [`McpClient::ensure_initialized`] holder is dropped before publishing
/// its handshake result (task cancellation, panic). Without this guard a
/// cancellation mid-handshake would leave `state` stuck in
/// [`ClientState::Initializing`] and every subsequent caller would block
/// until the wait-timeout fallback fires, then return an error — the
/// caller would have to call [`McpClient::reset_transport`]
/// manually to recover.
///
/// On the success path the holder calls [`Self::disarm`] before storing
/// `Ready`/`Empty` under the state lock, which converts the `Drop` into
/// a no-op. The guard never restores on the success path.
///
/// Drop uses [`tokio::sync::Mutex::try_lock`] because `Drop` runs
/// synchronously and we cannot block the runtime here. If the lock is
/// contended (extremely rare — the only competing locker is another
/// `ensure_initialized` caller which holds the lock for the duration of
/// a match arm, microseconds), the restore is skipped and the
/// inflight-wait timeout in `ensure_initialized` becomes the
/// last-resort recovery path.
struct InitGuard<'a> {
state: &'a Mutex<ClientState>,
init_done: &'a Notify,
/// `Some` until [`Self::disarm`] is called. Holds the restorable
/// transport (HTTP / HttpAuth) or `None` for Stdio (whose child
/// process is consumed by `client.serve` and cannot be reused).
restore: Option<PendingTransport>,
}
impl InitGuard<'_> {
/// Mark the guard as having published a result. Subsequent `Drop`
/// becomes a no-op so it doesn't fight with the holder's own
/// state-store-under-the-lock or wake waiters twice.
fn disarm(&mut self) {
self.restore = None;
}
}
impl Drop for InitGuard<'_> {
fn drop(&mut self) {
let Some(restore) = self.restore.take() else {
return;
};
// Best-effort restore. `try_lock` cannot block the runtime from
// inside Drop; on contention the slot stays Initializing and the
// inflight-wait timeout becomes the recovery path.
if let Ok(mut guard) = self.state.try_lock()
&& matches!(&*guard, ClientState::Initializing)
{
*guard = ClientState::Pending(restore);
}
// Notify whether or not we managed to restore — parked waiters
// need to wake up and either retry against the restored
// transport or hit the wait-timeout error path.
self.init_done.notify_waiters();
}
}
/// Build a restorable handle for a pending transport, or `None` if the
/// transport cannot be reused after a handshake failure.
///
/// `PendingTransport` deliberately does not implement `Clone`: the
/// `Stdio` variant owns a [`tokio::process::Child`] that is consumed by
/// `client.serve`, and a "restored" Stdio entry would be a dead handle.
/// HTTP and HttpAuth, by contrast, only carry config + an `Arc` to a
/// shared auth manager, so a clone is the canonical way to retry.
fn restorable_transport(pending: &PendingTransport) -> Option<PendingTransport> {
match pending {
PendingTransport::Http(cfg) => Some(PendingTransport::Http(cfg.clone())),
PendingTransport::HttpAuth {
config,
auth_manager,
} => Some(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_manager.clone(),
}),
PendingTransport::Acp { server_id, invoker } => Some(PendingTransport::Acp {
server_id: server_id.clone(),
invoker: invoker.clone(),
}),
PendingTransport::Stdio(_) => None,
}
}
/// Monotonic source for [`McpClient::client_id`]. Process-global so every
/// client instance — including test stubs — gets a unique identity.
static NEXT_CLIENT_ID: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
fn next_client_id() -> u64 {
NEXT_CLIENT_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
}
pub struct McpClient {
/// Unique identity of this client *instance*. See [`Self::client_id`].
client_id: u64,
server_name: McpServerName,
state: Mutex<ClientState>,
/// Wakes [`Self::ensure_initialized`] callers that observed
/// [`ClientState::Initializing`] and parked. Notified after each
/// handshake attempt finishes (success **or** failure) and `state`
/// has been updated. See [`ClientState`] for the single-flight
/// invariant this preserves.
///
/// Replaces the previous fail-fast
/// `McpError::ClientError("MCP client already initializing")` branch
/// which leaked into model-visible tool results whenever the model's
/// first tool dispatch raced the session actor's background
/// `get_tool_registrations` handshake.
init_done: Notify,
startup_timeout_sec: u64,
tool_timeout_sec: u64,
/// Per-tool timeout overrides in seconds. Looked up by tool name;
/// falls back to `tool_timeout_sec` when a tool isn't listed.
tool_timeouts: HashMap<ToolName, u64>,
/// See [`McpServerMetaConfig::expose_image_base64`].
expose_image_base64: bool,
/// Shared `AuthorizationManager` for OAuth-enabled servers. `AuthClient`
/// inside the transport holds a clone of this Arc so token updates are
/// visible to both the transport and the re-auth path.
auth_manager: Option<Arc<tokio::sync::Mutex<rmcp::transport::auth::AuthorizationManager>>>,
/// Stored for OAuth clients so we can rebuild the transport after re-auth.
http_config: Option<HttpConfig>,
/// BYO OAuth config for the full browser flow fallback (when refresh fails).
byo_oauth_config: Option<McpOAuthConfig>,
/// Rate limit on this server's reconnect warnings; passed to each HTTP
/// transport so rebuilds keep the limit.
warn_budget: crate::mcp_http_client::WarnBudget,
/// The transport to rebuild on a dead connection — see
/// [`McpClient::reset_transport`]. `None` for transports that can't
/// reconnect, e.g. Stdio (whose child process is consumed by the
/// handshake and can't be restarted from here).
reconnect: Option<PendingTransport>,
/// Event sink for transport-closed pollers, server-pushed
/// `tools/list_changed` / `resources/list_changed` notifications,
/// and handshake failures.
///
/// The slot is `Some` after [`Self::set_event_tx`] is called and
/// `None` otherwise. The `Arc<Mutex<...>>` is **shared with
/// [`KigiClientHandler`]** constructed by
/// [`Self::make_client_handler`]: the handler holds a clone of
/// the same Arc and reads through it on every notification.
/// Snapshotting the slot at handshake time instead would mean any
/// session that wired `notify_tx` post-handshake silently lost
/// every `tools/list_changed` and `resources/list_changed` for the
/// life of the connection.
///
/// `None` in three cases:
/// 1. Test stubs and standalone-pool fixtures that don't need
/// cross-component event flow.
/// 2. Subagent / shared-pool snapshots — only the **parent** session
/// is the owner of these events. A subagent that inherits a
/// shared `Arc<McpClient>` reads tools through it but does not
/// install its own dispatcher; the parent's
/// [`crate::liveness::TransportLivenessHandle`] already covers it.
/// 3. Brand-new clients before the session's per-server task has
/// called [`Self::set_event_tx`].
///
/// `parking_lot::Mutex` is sufficient (and lighter than the previous
/// `tokio::sync::Mutex`): the lock is never held across an
/// `.await`, and the handler's `emit` path is short and
/// allocation-free.
notify_tx: SharedEventTx,
/// RAII handle for the per-client transport-liveness poller.
///
/// `Some` after [`Self::arm_liveness_watcher`] succeeds; `None`
/// initially. The slot is also cleared by the poller itself when
/// it exits (whether on `TransportClosed` or because the state
/// machine drifted out of `Ready` during a re-handshake — see
/// [`crate::liveness::spawn_transport_liveness`]) so subsequent
/// `arm_liveness_watcher` calls aren't silently blocked by a
/// dead-but-still-present handle.
///
/// `parking_lot::Mutex` is sufficient because the lock is only ever
/// held for the duration of a slot swap. The poller task uses an
/// internal `Arc` clone of this mutex (the same memory) so it
/// can clear the slot before `break`.
liveness_handle: Arc<parking_lot::Mutex<Option<crate::liveness::TransportLivenessHandle>>>,
}
/// Shared sender slot type — the same Arc lives on the [`McpClient`]
/// and the [`KigiClientHandler`] it constructs during
/// [`McpClient::try_handshake`]. Mutating the slot via
/// [`McpClient::set_event_tx`] is observed by the live rmcp service
/// loop on the next notification, so there's no "snapshot at
/// handshake" hazard.
pub type SharedEventTx =
Arc<parking_lot::Mutex<Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>>>>;
/// External-config overrides for an MCP server, surfaced to kigi-mcp
/// from whatever loader the host crate uses (e.g. the host's `config.toml` parser).
///
/// All fields are pre-precedence: [`McpClient::load_timeouts`] (and
/// [`McpClient::load_expose_image_base64`]) still apply the
/// `_meta > overrides > default` precedence on top. Owning these types here
/// keeps MCP transport state free of the host's TOML schema.
///
/// Name retained for call-site stability; struct now carries non-timeout
/// config too (e.g. [`Self::expose_image_base64`]).
#[derive(Default, Debug, Clone)]
pub struct McpClientTimeoutOverrides {
/// Server startup timeout in seconds.
pub startup_timeout_sec: Option<u64>,
/// Default per-tool timeout in seconds (used when a tool has no entry in `tool_timeouts`).
pub tool_timeout_sec: Option<u64>,
/// Per-tool timeout overrides in seconds, keyed by tool name.
pub tool_timeouts: Option<HashMap<String, u64>>,
/// See [`McpServerMetaConfig::expose_image_base64`].
pub expose_image_base64: Option<bool>,
}
impl McpClient {
fn load_timeouts(
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> (u64, u64, HashMap<ToolName, u64>) {
// _meta > overrides > default; env / config / requirements / remote are
// resolved by the shell and injected via `overrides.startup_timeout_sec`.
let startup = meta_config
.and_then(|mc| mc.startup_timeout_ms)
.map(|ms| ms.div_ceil(1000))
.or_else(|| overrides.and_then(|o| o.startup_timeout_sec))
.unwrap_or(DEFAULT_STARTUP_TIMEOUT_SECS);
let tool = meta_config
.and_then(|mc| mc.tool_timeout_ms)
.map(|ms| ms.div_ceil(1000))
.or_else(|| overrides.and_then(|o| o.tool_timeout_sec))
.unwrap_or(DEFAULT_TOOL_TIMEOUT_SECS);
// Per-tool overrides: external base, _meta overrides on top.
// Precedence: _meta per-tool > overrides per-tool > (falls back to server-level `tool`)
let mut tool_timeouts = HashMap::new();
// Layer 1: overrides (already in seconds)
if let Some(o) = overrides
&& let Some(ref tt) = o.tool_timeouts
{
tool_timeouts.extend(tt.iter().map(|(k, v)| (k.clone(), *v)));
}
// Layer 2: _meta tool_timeouts_ms (milliseconds → seconds), overrides external config
if let Some(mc) = meta_config
&& let Some(ref tt) = mc.tool_timeouts_ms
{
for (k, v) in tt {
tool_timeouts.insert(k.clone(), v.div_ceil(1000));
}
}
(startup, tool, tool_timeouts)
}
/// `_meta > overrides > default(false)`, mirroring [`Self::load_timeouts`].
fn load_expose_image_base64(
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> bool {
meta_config
.and_then(|mc| mc.expose_image_base64)
.or_else(|| overrides.and_then(|o| o.expose_image_base64))
.unwrap_or(false)
}
/// The ONLY place that writes the `McpClient { .. }` struct literal.
/// Every constructor funnels through here so adding a field touches one
/// site. `reconnect` is snapshotted from the transport before it is
/// moved into [`ClientState::Pending`] (`None` for non-reconnectable
/// transports like Stdio — see [`restorable_transport`]).
#[allow(clippy::too_many_arguments)]
fn new_with_transport(
server_name: String,
transport: PendingTransport,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
auth_manager: Option<Arc<tokio::sync::Mutex<rmcp::transport::auth::AuthorizationManager>>>,
http_config: Option<HttpConfig>,
byo_oauth_config: Option<McpOAuthConfig>,
) -> Self {
let reconnect = restorable_transport(&transport);
let (startup_timeout_sec, tool_timeout_sec, tool_timeouts) =
Self::load_timeouts(overrides, meta_config);
let expose_image_base64 = Self::load_expose_image_base64(overrides, meta_config);
Self {
client_id: next_client_id(),
server_name,
state: Mutex::new(ClientState::Pending(transport)),
init_done: Notify::new(),
startup_timeout_sec,
tool_timeout_sec,
tool_timeouts,
expose_image_base64,
auth_manager,
http_config,
byo_oauth_config,
warn_budget: crate::mcp_http_client::WarnBudget::default(),
reconnect,
notify_tx: Arc::new(parking_lot::Mutex::new(None)),
liveness_handle: Arc::new(parking_lot::Mutex::new(None)),
}
}
pub fn new_http_auth(
server_name: String,
config: HttpConfig,
auth_manager: Arc<tokio::sync::Mutex<rmcp::transport::auth::AuthorizationManager>>,
byo_oauth_config: Option<McpOAuthConfig>,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> Self {
Self::new_with_transport(
server_name,
PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_manager.clone(),
},
overrides,
meta_config,
Some(auth_manager),
Some(config),
byo_oauth_config,
)
}
pub fn has_auth(&self) -> bool {
self.auth_manager.is_some()
}
/// Try to recover tokens from disk or via refresh — no browser flow.
///
/// Returns true if valid tokens were found (from another session/process
/// writing to the credential store, or a successful token refresh).
/// Used by `retry_auth_required_servers` on overlay refresh.
pub async fn try_reauth_from_disk(&self) -> bool {
let (Some(auth_mgr), Some(config)) = (&self.auth_manager, &self.http_config) else {
return false;
};
// Token-changed gate: rmcp's `initialize_from_store` returns Ok(true)
// for any disk-resident creds regardless of expiry, so without
// comparing against the in-memory token we'd claim "fresh tokens from
// disk" on the same stale token that triggered this retry. The
// downstream handshake would catch it, but the log line would lie
// during incident debugging — and the divergence from `force_reauth`'s
// gate is the exact invariant drift we just fixed there.
{
use oauth2::TokenResponse as _;
let mut mgr = auth_mgr.lock().await;
let token_before = mgr
.get_credentials()
.await
.ok()
.and_then(|(_, tok)| tok)
.map(|t| t.access_token().secret().to_string());
if let Ok(true) = mgr.initialize_from_store().await {
let token_after = mgr
.get_credentials()
.await
.ok()
.and_then(|(_, tok)| tok)
.map(|t| t.access_token().secret().to_string());
if token_after.is_some() && token_after != token_before {
tracing::info!(
server = self.server_name.as_str(),
"Loaded fresh tokens from disk"
);
drop(mgr);
self.replace_state(ClientState::Pending(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
}))
.await;
return true;
}
}
}
let refresh_ok = {
let mgr = auth_mgr.lock().await;
mgr.refresh_token().await.is_ok()
};
if refresh_ok {
self.replace_state(ClientState::Pending(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
}))
.await;
return true;
}
false
}
/// Force token acquisition and reset the transport so the next
/// `ensure_initialized` rebuilds it with the fresh token.
///
/// Tries in order:
/// 1. Reload from disk (picks up tokens from background auth task)
/// 2. Refresh via refresh_token grant
/// 3. Full browser-based OAuth flow
pub async fn force_reauth(&self, force: bool) -> bool {
let (Some(auth_mgr), Some(config)) = (&self.auth_manager, &self.http_config) else {
return false;
};
// Check if another process/session wrote *fresh* tokens to disk. We
// must compare against the token we already had in memory — rmcp's
// `initialize_from_store` returns Ok(true) for any disk-resident
// credentials regardless of expiry, so without the token-changed
// check we'd short-circuit on the same stale token that triggered
// this re-auth in the first place (real bug: pressing the auth
// shortcut on a server with an expired bearer + no refresh_token
// would no-op and then 401 on the next handshake).
//
// Hold a single lock guard across `token_before` → `initialize_from_store`
// → `token_after` so the comparison's invariant ("snapshot, reload,
// re-read") can't be torn by an interleaved mutation.
{
use oauth2::TokenResponse as _;
let mut mgr = auth_mgr.lock().await;
let token_before = mgr
.get_credentials()
.await
.ok()
.and_then(|(_, tok)| tok)
.map(|t| t.access_token().secret().to_string());
if let Ok(true) = mgr.initialize_from_store().await {
let token_after = mgr
.get_credentials()
.await
.ok()
.and_then(|(_, tok)| tok)
.map(|t| t.access_token().secret().to_string());
if token_after.is_some() && token_after != token_before {
tracing::info!(
server = self.server_name.as_str(),
"Loaded fresh tokens from disk (background auth or other process)"
);
drop(mgr);
self.replace_state(ClientState::Pending(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
}))
.await;
return true;
}
}
}
// Try token refresh.
let refresh_ok = {
let mgr = auth_mgr.lock().await;
mgr.refresh_token().await.is_ok()
};
if refresh_ok {
tracing::info!(
server = self.server_name.as_str(),
"Token refreshed successfully (no browser)"
);
self.replace_state(ClientState::Pending(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
}))
.await;
return true;
}
// Full browser-based OAuth flow.
{
tracing::info!(
server = self.server_name.as_str(),
"Falling back to browser auth"
);
if let Err(e) = crate::oauth::authenticate_mcp_server_dedup(
&self.server_name,
&config.url,
auth_mgr,
self.byo_oauth_config.as_ref(),
force,
)
.await
{
tracing::warn!(
server = self.server_name.as_str(),
%e,
"Full re-authentication failed"
);
return false;
}
}
self.replace_state(ClientState::Pending(PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
}))
.await;
true
}
/// Reset the transport so the next `ensure_initialized` rebuilds it with a
/// fresh connection.
///
/// Called when a tool call fails with a transport error (`TransportClosed`,
/// `TransportSend`) — the underlying connection is dead but the server's
/// addressing (URL/headers for HTTP, `server_id`/invoker for ACP) is still
/// valid.
///
/// Returns `true` if the transport was reset: HTTP/HttpAuth/ACP rebuild
/// from the `reconnect` snapshot taken at construction. Returns `false`
/// for clients whose `reconnect` is `None` (e.g. Stdio — dead child
/// processes can't be restarted from here).
async fn reset_transport(&self) -> bool {
let Some(t) = self.reconnect.as_ref().and_then(restorable_transport) else {
return false;
};
self.replace_state(ClientState::Pending(t)).await;
tracing::info!(
server = %self.server_name,
"Reset transport for reconnect after transport failure"
);
true
}
/// `true` if this client has an HTTP/SSE transport. The explicit predicate
/// for recovery gates (the proactive path only recovers HTTP clients).
pub fn is_http(&self) -> bool {
self.http_config.is_some()
}
/// `true` for an in-process SDK client reached over the ACP reverse channel
/// (rather than HTTP/stdio). Gates liveness watching — see
/// [`Self::arm_liveness_watcher`].
pub fn is_acp(&self) -> bool {
matches!(self.reconnect, Some(PendingTransport::Acp { .. }))
}
/// Read-only: do `headers` equal this client's current HTTP transport
/// headers? Compares the full set order-insensitively (the caller's
/// headers originate from a `HashMap`). Returns `false` for a client
/// with no HTTP config.
pub fn http_headers_match(&self, headers: &HashMap<String, String>) -> bool {
let Some(config) = &self.http_config else {
return false;
};
// Materialize into a map so a duplicate stored key collapses to one
// entry, keeping the length comparison honest. HTTP header names are
// case-insensitive, so normalize names to lowercase on both sides (the
// crate already does this for `authorization`) and avoid a needless
// rebuild on a pure casing difference. Values stay case-sensitive.
let stored: HashMap<String, &str> = config
.headers
.iter()
.map(|(k, v)| (k.to_ascii_lowercase(), v.as_str()))
.collect();
stored.len() == headers.len()
&& headers
.iter()
.all(|(k, v)| stored.get(&k.to_ascii_lowercase()) == Some(&v.as_str()))
}
/// Recover a dead transport in place: reset → re-handshake → re-arm the
/// liveness watcher. Returns the live [`McpService`].
///
/// The single recovery path for both the proactive HTTP recovery
/// (`SessionActor::reset_http_client`, gated on [`Self::is_http`]) and the
/// lazy `try_call_tool` retry. Rebuilds from the `reconnect` snapshot, so it
/// covers HTTP/HttpAuth/ACP; `arm_liveness_watcher` self-gates for ACP.
///
/// `Err` for a client with no restorable transport (e.g. Stdio — its child
/// was consumed by the handshake).
pub async fn recover(self: &Arc<Self>) -> Result<McpService, McpError> {
// Coalesce concurrent recoveries: reset only when Ready; if already
// non-Ready a recovery is in flight, so join its single-flight
// ensure_initialized instead of racing a reset.
if matches!(self.state_kind().await, ClientStateKind::Ready)
&& !self.reset_transport().await
{
return Err(McpError::ClientError(format!(
"MCP client {} has no transport to recover",
self.server_name,
)));
}
let service = self.ensure_initialized().await?;
// Re-arm liveness so the next close is detected again. A `false` return
// with a wired sender is unexpected only for watched (non-ACP) clients.
if !self
.arm_liveness_watcher(crate::liveness::DEFAULT_POLL_INTERVAL)
.await
&& !self.is_acp()
&& self.event_tx_clone().is_some()
{
tracing::warn!(
server = %self.server_name,
"recovery: liveness watcher not re-armed despite a wired event sender",
);
}
Ok(service)
}
/// Replace [`Self::state`] under the lock and wake any
/// [`Self::ensure_initialized`] callers parked on [`Self::init_done`]
/// so they re-check the new state on their next loop iteration.
///
/// Use this for *external* state transitions that need to invalidate
/// in-flight waits (re-auth completions, transport resets) so a parked
/// waiter doesn't sit on a stale [`ClientState::Initializing`] view of
/// the world. `ensure_initialized` itself doesn't go through this
/// helper because it already mints the new state under the lock it's
/// holding and notifies waiters once at the end of the handshake
/// attempt.
async fn replace_state(&self, new_state: ClientState) {
{
let mut guard = self.state.lock().await;
*guard = new_state;
}
self.init_done.notify_waiters();
}
pub fn new_stdio(
server_name: String,
transport: SafeTokioChildProcess,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> Self {
Self::new_with_transport(
server_name,
PendingTransport::Stdio(Box::new(transport)),
overrides,
meta_config,
None,
None,
None,
)
}
/// Build a client for an in-process SDK MCP server reached over the ACP reverse
/// channel. `server_id` is the id the agent echoes back in `kigi/mcp/sdk_call`; the
/// `invoker` performs the reverse request. Same downstream path as HTTP/stdio.
pub fn new_acp(
server_name: String,
server_id: String,
invoker: Arc<dyn crate::acp_transport::AcpReverseInvoker>,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> Self {
Self::new_with_transport(
server_name,
PendingTransport::Acp { server_id, invoker },
overrides,
meta_config,
None,
None,
None,
)
}
pub fn new_http(
server_name: String,
config: HttpConfig,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
) -> Self {
Self::new_with_transport(
server_name,
PendingTransport::Http(config.clone()),
overrides,
meta_config,
None,
Some(config),
None,
)
}
pub fn server_name(&self) -> &str {
&self.server_name
}
/// Unique identity of this client *instance*. Two clients for the
/// same server name (e.g. a dead client and its replacement after
/// a config remove+re-add) have different ids. Carried on
/// [`McpClientEvent::TransportClosed`] so consumers can tell a
/// death event for the current client from a stale predecessor's.
pub fn client_id(&self) -> u64 {
self.client_id
}
pub fn startup_timeout_sec(&self) -> u64 {
self.startup_timeout_sec
}
pub fn tool_timeout_sec(&self) -> u64 {
self.tool_timeout_sec
}
/// See [`McpServerMetaConfig::expose_image_base64`].
pub fn expose_image_base64(&self) -> bool {
self.expose_image_base64
}
/// Resolve the timeout for a specific tool.
///
/// Precedence (highest → lowest):
/// 1. `_meta.mcpConfig.<server>.toolTimeoutsMs.<tool>`
/// 2. `config.toml [mcp_servers.<server>].tool_timeouts.<tool>`
/// 3. `_meta.mcpConfig.<server>.toolTimeoutMs`
/// 4. `config.toml [mcp_servers.<server>].tool_timeout_sec`
/// 5. Default (60s)
///
/// Steps 12 are already merged into `self.tool_timeouts` at construction;
/// steps 35 are already resolved into `self.tool_timeout_sec`.
pub fn tool_timeout_for(&self, tool_name: &str) -> u64 {
self.tool_timeouts
.get(tool_name)
.copied()
.unwrap_or(self.tool_timeout_sec)
}
/// Drive the MCP handshake to completion (or return the cached
/// service if one is already established), with single-flight
/// semantics that are safe under arbitrary concurrent callers.
///
/// ## Concurrency contract
///
/// At most one task at a time runs [`Self::try_handshake`]; that task
/// holds the transport and observes [`ClientState::Initializing`].
/// Other concurrent callers park on [`Self::init_done`] (with a
/// bounded timeout) instead of issuing parallel handshakes or
/// failing immediately. When the holder publishes a result, all
/// parked waiters re-check `state` and either:
///
/// - return the freshly-stored [`McpService`] (handshake succeeded),
/// - take ownership of the freshly-restored transport and run their
/// own handshake (handshake failed but transport is restorable),
/// - or surface the error (Stdio handshake failed → no restorable
/// transport → [`ClientState::Empty`]).
///
/// This replaces the pre-fix behavior where concurrent callers got
/// an immediate `McpError::ClientError("MCP client already
/// initializing")` — surfaced inside model-visible tool results
/// whenever the model's first tool call landed inside the session
/// actor's background `get_tool_registrations` handshake, causing
/// repeated retries that exhausted prompt budgets without ever
/// reaching the actual MCP server.
///
/// ## Cancellation safety
///
/// If the holder is dropped (parent task cancelled, panic) before
/// publishing a result, [`InitGuard`]'s `Drop` impl best-effort
/// restores the transport (so future callers can retry without an
/// explicit `reset_transport`) and wakes parked waiters. The
/// restore uses [`tokio::sync::Mutex::try_lock`] because `Drop` is
/// synchronous; on the rare contention case the slot stays
/// `Initializing` and the wait-timeout fallback below surfaces a
/// clear error rather than blocking forever.
pub async fn ensure_initialized(&self) -> Result<McpService, McpError> {
// Bound how long a parked caller waits on `init_done` before
// surfacing an error. `try_handshake` is itself bounded by
// `startup_timeout_sec`, so anything beyond that plus a 1 s margin
// means the holder was dropped without restoring the transport
// (cancellation under heavy contention) — wedging silently would
// turn this into the exact "stuck client" failure mode the rest of
// this rewrite is designed to eliminate.
let inflight_wait =
std::time::Duration::from_secs(self.startup_timeout_sec.saturating_add(1));
// Drive the loop body until we either return directly or break
// out with an owned `PendingTransport`. We deliberately use a
// labelled `loop` with a `break <expr>` so the compiler proves
// every arm of the inner match either diverges (return /
// continue) or yields the transport — no `unreachable!()`
// escape hatch needed.
let pending: PendingTransport = loop {
// Subscribe to `init_done` BEFORE inspecting `state` so a
// wake-up fired between our state check and the wait can't be
// lost. `tokio::sync::Notify` only delivers a permit to
// notify-futures that exist at the time of `notify_waiters`.
let notified = self.init_done.notified();
tokio::pin!(notified);
let mut guard = self.state.lock().await;
// Swap the current state for `Initializing` up front and
// match on the OWNED previous value. This avoids the
// `match-by-ref → mem::replace → re-match → unreachable!()`
// dance — the compiler can bind `ClientState::Pending(t)`
// directly from an owned value with no irrefutable-let
// hole. Non-Pending arms restore their original variant
// before falling through; the lock is held the entire
// window so the brief `Initializing` placeholder is
// invisible to other callers. Cost is one trivial unit-
// variant write per non-Pending call (plus an `Arc::clone`
// on the Ready path) — negligible.
match std::mem::replace(&mut *guard, ClientState::Initializing) {
ClientState::Ready(service) => {
*guard = ClientState::Ready(service.clone());
return Ok(service);
}
ClientState::Empty => {
*guard = ClientState::Empty;
return Err(McpError::ClientError(format!(
"MCP client {} has no transport configured",
self.server_name,
)));
}
ClientState::Initializing => {
// Another caller already owns the slot. Restore the
// placeholder we just swapped in (semantically a
// no-op since `Initializing` is a unit variant),
// drop the lock, park on `init_done`.
*guard = ClientState::Initializing;
drop(guard);
match tokio::time::timeout(inflight_wait, notified.as_mut()).await {
Ok(()) => continue,
Err(_) => {
return Err(McpError::ClientError(format!(
"MCP client {} init still in progress after {}s",
self.server_name,
inflight_wait.as_secs(),
)));
}
}
}
// The single arm that KEEPS the `Initializing`
// placeholder we swapped in — this caller becomes the
// single-flight handshake holder for the duration of
// `try_handshake` below.
ClientState::Pending(transport) => break transport,
}
};
// Lock released. Run the handshake outside the lock so other
// callers can park on `init_done` instead of stalling on
// `state.lock()`.
// Clone the transport's restorable handle twice — once for
// the failure-path retry below, once for the drop guard.
// `PendingTransport` is intentionally not `Clone` (Stdio's
// `TokioChildProcess` is unique), so the helper returns
// `None` for Stdio (whose handshake failures cannot be
// recovered without a fresh spawn).
let restore = restorable_transport(&pending);
let restore_for_guard = restorable_transport(&pending);
// Drop guard: if `try_handshake` panics or is cancelled before
// we publish a result, restore `Pending(restore)` so other
// callers don't stall on `Initializing` forever. Disarm via
// `disarm()` immediately before storing the real result.
let mut init_guard = InitGuard {
state: &self.state,
init_done: &self.init_done,
restore: restore_for_guard,
};
let handshake_start = std::time::Instant::now();
let mut result = self.try_handshake(pending).await;
let handshake_elapsed = handshake_start.elapsed().as_micros() as u64;
tracing::info!(target: kigi_log::instrumentation::TARGET, event = "timing", name = "mcp_try_handshake", elapsed_us = handshake_elapsed);
// On handshake failure, if we have an auth_manager, try
// refreshing the token and retrying once. Handles expired
// access tokens loaded from disk — the handshake fails at the
// transport layer before rmcp's transparent 401 refresh can
// kick in. We attempt refresh on any failure (not just auth
// errors) because the cost is low and error strings from
// different MCP servers are not reliable to match.
if result.is_err()
&& let (Some(auth_mgr), Some(config)) = (&self.auth_manager, &self.http_config)
{
tracing::info!(
server = %self.server_name,
"Handshake failed, attempting token refresh and retry"
);
let refresh_ok = {
let mgr = auth_mgr.lock().await;
mgr.refresh_token().await.is_ok()
};
if refresh_ok {
let retry_transport = PendingTransport::HttpAuth {
config: config.clone(),
auth_manager: auth_mgr.clone(),
};
result = self.try_handshake(retry_transport).await;
}
}
// Disarm before publishing the result so the drop guard
// doesn't double-restore on the success path or fight with
// the failure-path assignment below.
init_guard.disarm();
// Snapshot the event sender before we commit Ready/Pending/Empty
// under the lock. We want to emit `HandshakeFailed` (on `Err`) or
// signal the dispatcher to set status=ready (on `Ok`) AFTER
// releasing the state lock, so a `state.lock().await` inside the
// dispatcher (should one ever exist — none today) can't deadlock.
//
// The snapshot reads through the SHARED `Arc<Mutex<...>>`
// slot. If the per-server task wired [`Self::set_event_tx`]
// BEFORE invoking `get_tool_registrations` (the pattern in
// `acp_session.rs`), this snapshot picks up the sender even
// for the very first handshake.
let event_tx = self.event_tx_clone();
let outcome = {
let mut guard = self.state.lock().await;
match result {
Ok(service) => {
let service = Arc::new(service);
*guard = ClientState::Ready(service.clone());
tracing::info!(
server = %self.server_name,
"MCP server initialized successfully"
);
Ok(service)
}
Err(e) => {
*guard = match restore {
Some(transport) => ClientState::Pending(transport),
None => ClientState::Empty,
};
tracing::warn!(
server = %self.server_name,
error = %e,
"MCP server init failed"
);
Err(e)
}
}
};
// Wake parked callers AFTER releasing the state lock so they
// observe the freshly-stored Ready/Pending/Empty value.
self.init_done.notify_waiters();
// Notify the session-actor StatusDispatcher of the handshake
// outcome, AFTER releasing the state lock. Best-effort: if the
// receiver is gone (dispatcher torn down, subagent without
// wiring) the send fails silently. The dispatcher is the only
// path that turns these into ACP pushes — see the
// `client_event_tx` field on `McpState`.
if let Some(tx) = &event_tx {
match &outcome {
Ok(_) => {
let _ = tx.send(McpClientEvent::Ready {
server: self.server_name.clone(),
});
}
Err(e) => {
let _ = tx.send(McpClientEvent::HandshakeFailed {
server: self.server_name.clone(),
reason: e.to_string(),
});
}
}
}
outcome
}
/// Run the MCP handshake (no lock held).
async fn try_handshake(
&self,
pending: PendingTransport,
) -> Result<rmcp::service::RunningService<RoleClient, KigiClientHandler>, McpError> {
let timeout = std::time::Duration::from_secs(self.startup_timeout_sec);
let name = &self.server_name;
match pending {
PendingTransport::Stdio(process) => {
let handler = self.make_client_handler();
tokio::time::timeout(timeout, handler.serve(*process))
.await
.map_err(|_| McpError::timeout(name, timeout))?
.map_err(|e| McpError::HandshakeFailed {
server: name.to_string(),
source: Box::new(e),
})
}
PendingTransport::Http(config) => {
let transport =
Self::build_http_transport(&config, name, self.warn_budget.clone())?;
let handler = self.make_client_handler();
tokio::time::timeout(timeout, handler.serve(transport))
.await
.map_err(|_| McpError::timeout(name, timeout))?
.map_err(|e| McpError::HandshakeFailed {
server: name.to_string(),
source: Box::new(e),
})
}
PendingTransport::HttpAuth {
config,
auth_manager,
} => {
let mut headers = reqwest::header::HeaderMap::new();
for (key, value) in &config.headers {
if key.eq_ignore_ascii_case("Authorization") {
continue;
}
if let (Ok(n), Ok(v)) = (
reqwest::header::HeaderName::from_bytes(key.as_bytes()),
value.parse::<reqwest::header::HeaderValue>(),
) {
headers.insert(n, v);
}
}
ensure_figma_user_agent(&mut headers, name, &config.url);
let http_client = reqwest::Client::builder()
.default_headers(headers)
.build()
.map_err(|e| {
McpError::ClientError(format!("Failed to build HTTP client: {e}"))
})?;
// `AuthClient::new` wants an owned manager, but ours is shared
// (`Arc`) with the OAuth flow; the struct is non_exhaustive, so
// build with a throwaway manager and swap in the shared one.
let placeholder_manager =
rmcp::transport::auth::AuthorizationManager::new(config.url.as_str())
.await
.map_err(|e| {
McpError::ClientError(format!("Failed to build OAuth client: {e}"))
})?;
let mut auth_client =
rmcp::transport::auth::AuthClient::new(http_client, placeholder_manager);
auth_client.auth_manager = auth_manager.clone();
let mcp_http_client = crate::mcp_http_client::McpHttpClient::new(
auth_client,
name.as_str(),
self.warn_budget.clone(),
);
let transport_config =
StreamableHttpClientTransportConfig::with_uri(config.url.as_str());
let transport =
StreamableHttpClientTransport::with_client(mcp_http_client, transport_config);
let handler = self.make_client_handler();
tokio::time::timeout(timeout, handler.serve(transport))
.await
.map_err(|_| McpError::timeout(name, timeout))?
.map_err(|e| McpError::HandshakeFailed {
server: name.to_string(),
source: Box::new(e),
})
}
PendingTransport::Acp { server_id, invoker } => {
// Per-reverse-call backstop on `kigi/mcp/sdk_call`: the larger of the
// startup and tool timeouts, so it never undercuts the real outer bound
// (the handshake `initialize` is bounded by the serve `timeout` below;
// tool calls by `tool_timeout_for` in `try_call_tool`). The bridge
// forwards raw JSON-RPC without the tool name, so per-TOOL overrides
// aren't applied here in v1; the HTTP path still honors them.
let invoke_timeout = std::time::Duration::from_secs(
self.startup_timeout_sec.max(self.tool_timeout_sec),
);
let transport =
crate::acp_transport::acp_bridge_transport(server_id, invoker, invoke_timeout);
let handler = self.make_client_handler();
tokio::time::timeout(timeout, handler.serve(transport))
.await
.map_err(|_| McpError::timeout(name, timeout))?
.map_err(|e| McpError::HandshakeFailed {
server: name.to_string(),
source: Box::new(e),
})
}
}
}
fn make_client_info(server_name: &str) -> ClientInfo {
let mut extensions = rmcp::model::ExtensionCapabilities::new();
extensions.insert(
"io.modelcontextprotocol/ui".to_string(),
serde_json::from_value(serde_json::json!({
"mimeTypes": ["text/html;profile=mcp-app"]
}))
.unwrap_or_default(),
);
let mut capabilities = ClientCapabilities::default();
capabilities.extensions = Some(extensions);
ClientInfo::new(
capabilities,
Implementation::new(
format!("kigi-shell-{server_name}"),
kigi_version::VERSION.to_string(),
),
)
// rmcp's default `ProtocolVersion` tracks its LATEST; pin explicitly
// so the advertised protocol only changes deliberately, never as a
// side effect of an rmcp bump.
.with_protocol_version(rmcp::model::ProtocolVersion::V_2025_06_18)
}
/// Build the [`KigiClientHandler`] that drives `client.serve(...)`.
///
/// The handler holds a **clone of `Arc<Mutex<Option<Sender>>>`**,
/// not a snapshot — so any subsequent call to
/// [`Self::set_event_tx`] is observed by the live rmcp service
/// loop on its next notification.
fn make_client_handler(&self) -> KigiClientHandler {
KigiClientHandler {
info: Self::make_client_info(&self.server_name),
server_name: self.server_name.clone(),
notify_tx: Arc::clone(&self.notify_tx),
}
}
/// Wire a sender for [`McpClientEvent`]s emitted by this client.
///
/// Mutates the shared slot synchronously. All previously-cloned
/// references (the [`KigiClientHandler`] handed to
/// `client.serve`, the [`crate::liveness::spawn_transport_liveness`]
/// task) read through the same Arc, so this is observed
/// session-wide on the next event.
pub fn set_event_tx(&self, tx: Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>>) {
*self.notify_tx.lock() = tx;
}
/// Snapshot the current event sender, if any.
///
/// Used by [`crate::liveness::spawn_transport_liveness`] (which
/// captures a `Sender` clone at spawn time) and by
/// [`Self::ensure_initialized`]'s post-handshake emit. Synchronous
/// because the shared slot is a `parking_lot::Mutex`.
pub fn event_tx_clone(&self) -> Option<tokio::sync::mpsc::UnboundedSender<McpClientEvent>> {
self.notify_tx.lock().clone()
}
/// Install or replace this client's transport-liveness handle.
/// Dropping the previous handle (if any) cancels its task; the
/// new handle starts polling on its own schedule. Pass `None` to
/// stop watching without installing a new one.
///
/// Synchronous: the slot is a `parking_lot::Mutex`. The poller
/// task uses an `Arc` clone of this same mutex so it can clear the
/// slot from inside the task before exiting.
pub fn set_liveness_handle(&self, handle: Option<crate::liveness::TransportLivenessHandle>) {
*self.liveness_handle.lock() = handle;
}
/// Arm the per-client transport-liveness watcher.
///
/// Idempotent and gated:
/// - Returns `false` for in-process SDK ([`Self::is_acp`]) clients: the
/// watcher's only output is `TransportClosed`, which the dispatcher can't
/// recover for ACP (not in `configs`), so it would evict the client. ACP
/// recovers lazily via [`Self::reset_transport`] instead. Gated here so no
/// caller can forget it.
/// - Returns `false` if there's no `notify_tx` wired (subagent
/// snapshot or pre-dispatcher state) — nothing to do.
/// - Returns `false` if the client isn't `Ready` — armed pollers
/// would just exit silently on their first poll, but skipping
/// the spawn entirely is cheaper.
/// - Returns `false` if a live handle is already installed.
/// - Otherwise spawns the poller and stores the handle.
///
/// **TOCTOU note**: the state check is performed before the
/// liveness lock is acquired. A concurrent re-handshake could move
/// the state to `Initializing` between the check and the spawn.
/// This is benign — the poller's first tick observes the
/// non-`Ready` state and exits silently without emitting. So the
/// worst case under TOCTOU is "the poller starts and immediately
/// stops"; it never produces a spurious `TransportClosed`.
///
/// Lifecycle: when the watcher emits `TransportClosed` it clears
/// the slot itself; the next `arm_liveness_watcher` call can
/// install a fresh handle without a manual
/// [`Self::set_liveness_handle`] reset.
pub async fn arm_liveness_watcher(
self: &Arc<Self>,
poll_interval: std::time::Duration,
) -> bool {
if self.is_acp() {
return false;
}
let Some(event_tx) = self.event_tx_clone() else {
return false;
};
if !matches!(self.state_kind().await, ClientStateKind::Ready) {
return false;
}
let mut slot = self.liveness_handle.lock();
if slot.is_some() {
return false;
}
let handle = crate::liveness::spawn_transport_liveness(
self.server_name.clone(),
Arc::clone(self),
poll_interval,
event_tx,
Arc::clone(&self.liveness_handle),
);
*slot = Some(handle);
true
}
fn build_http_transport(
config: &HttpConfig,
server_name: &str,
warn_budget: crate::mcp_http_client::WarnBudget,
) -> Result<
StreamableHttpClientTransport<crate::mcp_http_client::McpHttpClient<reqwest::Client>>,
McpError,
> {
let mut headers = reqwest::header::HeaderMap::new();
for (key, value) in &config.headers {
match (
reqwest::header::HeaderName::from_bytes(key.as_bytes()),
value.parse::<reqwest::header::HeaderValue>(),
) {
(Ok(name), Ok(val)) => {
headers.insert(name, val);
}
_ => {
tracing::warn!("Skipping invalid MCP HTTP header: {key}");
}
}
}
ensure_figma_user_agent(&mut headers, server_name, &config.url);
let client = reqwest::Client::builder()
.default_headers(headers)
.build()
.map_err(|e| McpError::ClientError(format!("Failed to build HTTP client: {e}")))?;
let mcp_http_client =
crate::mcp_http_client::McpHttpClient::new(client, server_name, warn_budget);
let transport_config = StreamableHttpClientTransportConfig::with_uri(config.url.as_str());
Ok(StreamableHttpClientTransport::with_client(
mcp_http_client,
transport_config,
))
}
/// Cheap, non-blocking liveness predicate.
///
/// Inspects the current [`ClientState`] under the state mutex only —
/// it MUST NOT call [`Self::ensure_initialized`] or any other path
/// that can trigger a network round-trip. The previous implementation
/// went through `ensure_initialized`, which could block UI callers
/// (e.g. an MCP status modal) for up to `startup_timeout_sec` seconds
/// on a dead stdio server.
///
/// Semantics:
/// - `Ready(service)` with an open transport → `true`.
/// - `Ready(service)` whose receiver-side has been dropped (typically
/// because the rmcp service loop terminated) → `false`. rmcp 2.1
/// `Peer::is_transport_closed` reports `self.tx.is_closed()` at
/// `service.rs:703-705`; `RunningService` derefs to `Peer` at
/// `service.rs:716-722`.
/// - Any other variant (`Empty`, `Pending`, `Initializing`) →
/// `false`.
///
/// HTTP idle caveat: for [`StreamableHttpClientTransport`] the rmcp
/// service loop only terminates on an outgoing send failure or an
/// explicit shutdown. A long-idle HTTP server therefore keeps
/// `is_transport_closed()` returning `false`, and this method
/// continues to report `true`. That is the desired semantics — a
/// liveness probe would belong in a separate watcher, not here.
pub async fn is_healthy(&self) -> bool {
let guard = self.state.lock().await;
match &*guard {
ClientState::Ready(service) => !service.is_transport_closed(),
_ => false,
}
}
/// Atomic classification for the liveness watcher.
///
/// Reads `state` once and projects onto
/// [`LivenessCheck`]: distinguishes "transport actually closed"
/// (emit + exit) from "state moved out of `Ready`" (exit
/// silently). The watcher depends on this distinction: a plain
/// `is_healthy`-based predicate cannot tell the cases apart and
/// would false-fire `TransportClosed` on re-handshake transitions.
pub async fn liveness_check(&self) -> LivenessCheck {
let guard = self.state.lock().await;
match &*guard {
ClientState::Ready(service) => {
if service.is_transport_closed() {
LivenessCheck::TransportClosed
} else {
LivenessCheck::Healthy
}
}
_ => LivenessCheck::Transient,
}
}
/// State-machine snapshot for diagnostics and downstream UI.
///
/// Like [`Self::is_healthy`], this is a cheap state inspection
/// (no handshake, no network I/O). Maps [`ClientState`] onto a
/// `Copy` enum so callers can match without holding a reference to
/// the inner [`McpService`] / [`PendingTransport`].
pub async fn state_kind(&self) -> ClientStateKind {
let guard = self.state.lock().await;
match &*guard {
ClientState::Empty => ClientStateKind::Empty,
ClientState::Pending(_) => ClientStateKind::Pending,
ClientState::Initializing => ClientStateKind::Initializing,
ClientState::Ready(_) => ClientStateKind::Ready,
}
}
/// Materialize this server's tool descriptors as JSON files under
/// `<server_dir>/tools/`.
///
/// The model reads these before issuing an MCP tool call.
/// Each tool becomes `<server_dir>/tools/<sanitized_tool_name>.json`
/// with `{name, description, inputSchema}`. Resources are intentionally
/// not materialized: this harness exposes only MCP tool calls, so resource
/// descriptors would advertise MCP-resource tools the model can't use.
///
/// Best-effort: errors writing individual descriptors are logged but
/// don't abort the materialization. Returns the number of files written.
pub async fn materialize_descriptors(
&self,
server_dir: &std::path::Path,
) -> Result<usize, McpError> {
let mcp_service = self.ensure_initialized().await?;
// Collect descriptors via the async MCP API, then defer all filesystem
// work to a single `spawn_blocking` so the executor is never blocked on
// `std::fs` (this runs on every MCP tool-set change, not just startup).
let mut files: Vec<(String, Vec<u8>)> = Vec::new();
let mut cursor: Option<String> = None;
loop {
let result = mcp_service
.list_tools(Some(
PaginatedRequestParams::default().with_cursor(cursor.clone()),
))
.await?;
for tool in result.tools {
let descriptor = serde_json::json!({
"name": tool.name.as_ref(),
"description": tool.description.as_deref(),
"inputSchema": tool.input_schema.as_ref(),
});
match serde_json::to_vec_pretty(&descriptor) {
Ok(bytes) => files.push((
format!("{}.json", sanitize_descriptor_segment(tool.name.as_ref())),
bytes,
)),
Err(e) => tracing::warn!(
tool = %tool.name.as_ref(),
error = %e,
"failed to serialize MCP tool descriptor",
),
}
}
match result.next_cursor {
Some(next) => cursor = Some(next),
None => break,
}
}
// Write each descriptor atomically (temp file + rename) so a concurrent
// reader never sees a half-written JSON and overlapping writers converge
// without a lock.
let tools_dir = server_dir.join("tools");
tokio::task::spawn_blocking(move || -> Result<usize, McpError> {
std::fs::create_dir_all(&tools_dir).map_err(|e| {
McpError::ClientError(format!(
"failed to create MCP tools descriptor dir {}: {e}",
tools_dir.display()
))
})?;
let mut written = 0usize;
for (file_name, bytes) in files {
let path = tools_dir.join(&file_name);
let write_result =
tempfile::NamedTempFile::new_in(&tools_dir).and_then(|mut tmp| {
std::io::Write::write_all(&mut tmp, &bytes)?;
tmp.persist(&path).map_err(|e| e.error)
});
match write_result {
Ok(_) => written += 1,
Err(e) => tracing::warn!(
path = %path.display(),
error = %e,
"failed to write MCP tool descriptor",
),
}
}
Ok(written)
})
.await
.map_err(|e| McpError::ClientError(format!("descriptor write task panicked: {e}")))?
}
/// Read the server's `instructions` from the MCP initialize handshake.
/// Returns `None` if the client isn't ready yet.
pub async fn server_instructions(&self) -> Option<String> {
let guard = self.state.lock().await;
if let ClientState::Ready(service) = &*guard {
service
.peer_info()?
.instructions
.as_deref()
.filter(|s| !s.trim().is_empty())
.map(String::from)
} else {
None
}
}
pub async fn get_tool_registrations(
&self,
mcp_state: Arc<Mutex<McpState>>,
) -> Result<Vec<McpToolRegistration>, McpError> {
let _ensure_init_timer = kigi_log::instrumentation::timer("mcp_ensure_initialized");
let mcp_service = self.ensure_initialized().await?;
let mut all_tools = Vec::new();
let mut cursor: Option<String> = None;
let _list_tools_timer = kigi_log::instrumentation::timer("mcp_list_tools");
loop {
let list_tools_result = mcp_service
.list_tools(Some(
PaginatedRequestParams::default().with_cursor(cursor.clone()),
))
.await?;
all_tools.extend(list_tools_result.tools);
match list_tools_result.next_cursor {
Some(next) => cursor = Some(next),
None => break,
}
}
let registrations: Vec<_> = all_tools
.into_iter()
.filter_map(|tool| {
let meta: Option<serde_json::Value> = tool
.meta
.as_ref()
.and_then(|m| serde_json::to_value(m).ok());
let name = tool.name.to_string();
let description = tool.description.map(|d| d.to_string()).unwrap_or_default();
let mut schema = serde_json::to_value(tool.input_schema.as_ref())
.unwrap_or_else(|_| serde_json::json!({"type": "object"}));
// Ensure the schema has "type": "object" — some MCP servers
// (e.g., VSCode) send `inputSchema: {}` for tools with no
// parameters. Azure's OpenAI API rejects schemas without a
// `type` field with: 'schema must be a JSON Schema of type:
// "object", got type: "None"'.
if let Some(obj) = schema.as_object_mut() {
obj.entry("type")
.or_insert_with(|| serde_json::json!("object"));
obj.entry("properties")
.or_insert_with(|| serde_json::json!({}));
}
let mcp_tool = McpTool {
name,
description,
server_name: self.server_name.clone(),
mcp_state: Arc::clone(&mcp_state),
schema,
meta,
};
// Invalid tools (bad names) return None and are skipped
mcp_tool.into_registration()
})
.collect();
// Warn about tool_timeouts keys that don't match any discovered tool.
// This catches typos like `creat_issue` instead of `create_issue`.
if !self.tool_timeouts.is_empty() {
// Registration names are qualified ("server__tool"); tool_timeouts
// keys are raw tool names. Strip the server prefix for comparison.
let prefix = format!("{}{}", self.server_name, MCP_TOOL_NAME_DELIMITER);
let raw_names: Vec<&str> = registrations
.iter()
.map(|r| r.name.strip_prefix(prefix.as_str()).unwrap_or(&r.name))
.collect();
let discovered: std::collections::HashSet<&str> = raw_names.iter().copied().collect();
for key in self.tool_timeouts.keys() {
if !discovered.contains(key.as_str()) {
tracing::info!(
server = %self.server_name,
tool_timeout_key = %key,
"tool_timeouts entry '{}' does not match any tool exposed by MCP server '{}' \
(available: {}). The per-tool timeout will have no effect — check for typos.",
key,
self.server_name,
raw_names.join(", "),
);
}
}
}
Ok(registrations)
}
/// Call a tool directly on the MCP server (for testing/debugging).
pub async fn call_tool(
&self,
tool_name: &str,
arguments: serde_json::Value,
) -> Result<rmcp::model::CallToolResult, McpError> {
let mcp_service = self.ensure_initialized().await?;
let result = mcp_service
.call_tool({
let mut params = CallToolRequestParams::new(tool_name.to_string());
params.arguments = arguments.as_object().cloned();
params
})
.await?;
Ok(result)
}
}
fn contains_session_placeholder(value: &str) -> bool {
value.contains("{{session_id}}") || value.contains("${session_id}")
}
/// Sanitize an MCP server name into a safe filename component.
fn sanitize_mcp_log_filename(name: &str) -> String {
let sanitized: String = name
.chars()
.take(96)
.map(|c| match c {
c if c.is_ascii_alphanumeric() => c,
'.' | '_' | '-' => c,
_ => '_',
})
.collect();
if sanitized.is_empty() {
"server".into()
} else {
sanitized
}
}
/// Copy an MCP server's stderr to `~/.kigi/logs/mcp/<server>.stderr.log`
/// in a background task. Truncated per spawn.
fn drain_mcp_stderr_to_log(server_name: &str, mut stderr: tokio::process::ChildStderr) {
let log_dir = kigi_config::kigi_home().join("logs").join("mcp");
if let Err(e) = std::fs::create_dir_all(&log_dir) {
tracing::warn!("MCP stderr drain: failed to create log dir: {e}");
return;
}
let log_path = log_dir.join(format!(
"{}.stderr.log",
sanitize_mcp_log_filename(server_name)
));
let file = match std::fs::OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(&log_path)
{
Ok(f) => f,
Err(e) => {
tracing::warn!(
"MCP stderr drain: failed to open {}: {e}",
log_path.display()
);
return;
}
};
let mut file = tokio::fs::File::from_std(file);
let server_name = server_name.to_string();
tokio::spawn(async move {
if let Err(e) = tokio::io::copy(&mut stderr, &mut file).await {
tracing::warn!("MCP stderr drain '{server_name}': {e}");
}
});
}
fn expand_session_id_headers(
headers: Vec<acp::HttpHeader>,
session_id: Option<&str>,
) -> Vec<(String, String)> {
headers
.into_iter()
.filter_map(|header| {
let value = header.value;
if let Some(session_id) = session_id {
let expanded = value
.replace("{{session_id}}", session_id)
.replace("${session_id}", session_id);
Some((header.name, expanded))
} else if contains_session_placeholder(&value) {
None
} else {
Some((header.name, value))
}
})
.collect()
}
/// Decide the actual (program, args) to spawn for a stdio MCP server.
///
/// On Windows, npm ships launchers like `npx`/`npm`/`pnpm`/`yarn` as `.cmd`
/// batch shims (there is no `npx.exe`). `CreateProcessW` only appends `.exe`
/// and ignores `PATHEXT`, so `Command::new("npx")` fails with "file not
/// found". We resolve the bare name on `PATH` (honoring `PATHEXT`, via the
/// `resolve` closure) so std spawns the real launcher path (e.g. `npx.cmd`) —
/// std then runs `.cmd`/`.bat` through `cmd.exe` with hardened arg escaping. On
/// non-Windows we never touch the command (verified working). A command
/// containing a path separator is used as-is. The resolved path is returned as
/// an `OsString` so it reaches `Command::new` without a lossy UTF-8 round-trip.
fn plan_stdio_spawn(
command: &str,
args: &[String],
is_windows: bool,
resolve: impl Fn(&str) -> Option<std::path::PathBuf>,
) -> (OsString, Vec<String>) {
if is_windows
&& !command.contains('/')
&& !command.contains('\\')
&& let Some(resolved) = resolve(command)
{
return (resolved.into_os_string(), args.to_vec());
}
(OsString::from(command), args.to_vec())
}
fn is_figma_mcp(server_name: &str, url: &str) -> bool {
if server_name.eq_ignore_ascii_case("figma") {
return true;
}
reqwest::Url::parse(url)
.ok()
.and_then(|u| u.host_str().map(|h| h.to_ascii_lowercase()))
.is_some_and(|h| h == "figma.com" || h.ends_with(".figma.com"))
}
fn ensure_figma_user_agent(headers: &mut reqwest::header::HeaderMap, server_name: &str, url: &str) {
if !is_figma_mcp(server_name, url) {
return;
}
if headers.contains_key(reqwest::header::USER_AGENT) {
return;
}
headers.insert(
reqwest::header::USER_AGENT,
reqwest::header::HeaderValue::from_static("kigi-cli"),
);
}
fn stdio_path_override(env: &[acp::EnvVariable]) -> Option<&str> {
env.iter()
.find(|e| e.name.eq_ignore_ascii_case("PATH"))
.map(|e| e.value.as_str())
}
pub async fn start_mcp_server(
mcp_server: acp::McpServer,
session_id: Option<&str>,
overrides: Option<&McpClientTimeoutOverrides>,
meta_config: Option<&McpServerMetaConfig>,
byo_config: Option<&McpOAuthConfig>,
event_writer: &kigi_file_utils::events::EventWriter,
mode: OauthInteractivity,
) -> Result<McpClient, McpError> {
let _per_server_timer = kigi_log::instrumentation::timer("mcp_start_one_server");
match mcp_server {
acp::McpServer::Stdio(acp::McpServerStdio {
name,
command,
args,
env,
..
}) => {
if let Some(mc) = meta_config {
tracing::info!(server = %name, ?mc, "MCP stdio: meta config override");
}
let command_str = command.to_string_lossy().into_owned();
let _stdio_spawn_timer = kigi_log::instrumentation::timer("mcp_stdio_spawn");
let path_override = stdio_path_override(&env);
let (program, spawn_args) = plan_stdio_spawn(&command_str, &args, cfg!(windows), |c| {
if let Some(path) = path_override
&& let Ok(cwd) = std::env::current_dir()
{
which::which_in(c, Some(path), cwd).ok()
} else {
which::which(c).ok()
}
});
let mut cmd = Command::new(&program);
cmd.kill_on_drop(true).args(&spawn_args);
for env_variable in &env {
cmd.env(&env_variable.name, &env_variable.value);
}
kigi_tools::util::detach_command(&mut cmd);
let (transport, stderr_handle) =
SafeTokioChildProcess::spawn(cmd, name.clone(), event_writer.clone()).map_err(
|e| {
tracing::error!("Failed to spawn MCP server '{}': {}", name, e);
McpError::SpawnFailed {
server: name.clone(),
source: e,
}
},
)?;
tracing::debug!("MCP server '{}' spawned: PID={:?}", name, transport.id());
if let Some(stderr) = stderr_handle {
drain_mcp_stderr_to_log(&name, stderr);
}
Ok(McpClient::new_stdio(
name.clone(),
transport,
overrides,
meta_config,
))
}
acp::McpServer::Http(acp::McpServerHttp {
name, url, headers, ..
})
| acp::McpServer::Sse(acp::McpServerSse {
name, url, headers, ..
}) => {
if let Some(mc) = meta_config {
tracing::info!(server = %name, %url, ?mc, "MCP http: meta config override");
}
let headers = expand_session_id_headers(headers, session_id);
let http_config = HttpConfig {
url: url.clone(),
headers,
};
let has_existing_auth = http_config
.headers
.iter()
.any(|(k, _)| k.eq_ignore_ascii_case("authorization"));
let auth_prep = if has_existing_auth {
tracing::debug!(
server = %name,
"Skipping OAuth discovery: server already has Authorization header"
);
HttpOauthPrep::NoOauthSupport
} else {
let _auth_discovery_timer =
kigi_log::instrumentation::timer("mcp_http_auth_discovery");
match tokio::time::timeout(
OAUTH_DISCOVERY_TIMEOUT,
discover_and_prepare_auth(&name, &url, mode),
)
.await
{
Ok(result) => result,
Err(_) => {
tracing::warn!(
server = %name,
url = %url,
?mode,
timeout_secs = OAUTH_DISCOVERY_TIMEOUT.as_secs(),
"OAuth discovery timed out"
);
event_writer.emit(
kigi_file_utils::events::Event::McpOAuthDiscoveryTimeout {
server_name: name.clone(),
url: url.clone(),
},
);
HttpOauthPrep::on_probe_failure(mode)
}
}
};
match auth_prep {
HttpOauthPrep::ManagerReady(auth_mgr) => Ok(McpClient::new_http_auth(
name.clone(),
http_config,
auth_mgr,
byo_config.cloned(),
overrides,
meta_config,
)),
HttpOauthPrep::NoOauthSupport => Ok(McpClient::new_http(
name.clone(),
http_config,
overrides,
meta_config,
)),
// Avoid starting an unauthenticated HTTP worker that fatals on server OAuth challenge.
HttpOauthPrep::NeedsInteractiveLogin => Err(McpError::AuthRequired {
server: name.clone(),
}),
}
}
// TODO(acp-0.10): `McpServer` is #[non_exhaustive]; reject unknown transports.
other => Err(McpError::ClientError(format!(
"unsupported MCP server transport: {other:?}"
))),
}
}
pub async fn start_mcp_servers(
mcp_servers: Vec<acp::McpServer>,
session_id: Option<&str>,
overrides_map: &HashMap<String, McpClientTimeoutOverrides>,
meta_config_map: &McpMetaConfigMap,
oauth_config_map: &crate::oauth_config::McpOAuthConfigMap,
event_writer: &kigi_file_utils::events::EventWriter,
mode: OauthInteractivity,
) -> Vec<Result<McpClient, McpError>> {
let _mcp_start_timer = kigi_log::instrumentation::timer("mcp_start_servers");
if !meta_config_map.is_empty() {
tracing::info!(
count = mcp_servers.len(),
overrides = ?meta_config_map.keys().collect::<Vec<_>>(),
"Starting MCP servers with meta config"
);
}
futures::stream::iter(mcp_servers)
.map(|server| {
let server_name = mcp_server_name(&server);
let overrides = overrides_map.get(server_name);
let mc = meta_config_map.get(server_name);
let byo = oauth_config_map.get(server_name);
start_mcp_server(server, session_id, overrides, mc, byo, event_writer, mode)
})
.buffer_unordered(8)
.collect::<Vec<_>>()
.await
}
/// Extract the name from an MCP server enum variant.
pub fn mcp_server_name(server: &acp::McpServer) -> &str {
match server {
acp::McpServer::Stdio(stdio) => &stdio.name,
acp::McpServer::Http(http) => &http.name,
acp::McpServer::Sse(sse) => &sse.name,
// TODO(acp-0.10): `McpServer` is #[non_exhaustive].
_ => "unknown",
}
}
pub fn mcp_transport_str(server: &acp::McpServer) -> &'static str {
match server {
acp::McpServer::Stdio(_) => "stdio",
acp::McpServer::Http(_) => "http",
acp::McpServer::Sse(_) => "sse",
// TODO(acp-0.10): `McpServer` is #[non_exhaustive].
_ => "unknown",
}
}
pub fn mcp_target_str(server: &acp::McpServer) -> String {
match server {
acp::McpServer::Stdio(acp::McpServerStdio { command, args, .. }) => {
let cmd = command.to_string_lossy();
if args.is_empty() {
cmd.to_string()
} else {
format!("{} {}", cmd, args.join(" "))
}
}
acp::McpServer::Http(acp::McpServerHttp { url, .. })
| acp::McpServer::Sse(acp::McpServerSse { url, .. }) => url.clone(),
// TODO(acp-0.10): `McpServer` is #[non_exhaustive].
_ => String::new(),
}
}
impl McpClient {
/// Minimal stub for unit tests in dependent crates. Hidden from rustdoc;
/// not gated behind `#[cfg(test)]` so cross-crate test code can construct
/// it without the host crate enabling a feature.
#[doc(hidden)]
pub fn stub(name: &str) -> Self {
// Route through the single constructor (so new fields never need
// touching here), then downgrade to the no-transport placeholder:
// `Empty` state makes `ensure_initialized` error, and `reconnect =
// None` makes `reset_transport` return false — i.e. a client that
// can't reconnect, like a dead Stdio child. Overrides preserve the
// historical stub timeouts (10s startup / 60s tool).
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(10),
tool_timeout_sec: Some(60),
..Default::default()
};
let mut client = Self::new_with_transport(
name.to_string(),
PendingTransport::Http(HttpConfig {
url: String::new(),
headers: Vec::new(),
}),
Some(&overrides),
None,
None,
None,
None,
);
*client.state.get_mut() = ClientState::Empty;
client.reconnect = None;
client
}
}
/// rmcp [`ClientHandler`] used by all MCP transports.
///
/// Replaces the previous bare [`ClientInfo`] handler at the three
/// `client.serve(...)` call sites in [`McpClient::try_handshake`].
/// Plumbs server-pushed notifications through an
/// [`tokio::sync::mpsc::UnboundedSender<McpClientEvent>`] so the
/// session-actor dispatcher can fan them out as ACP
/// `kigi/mcp/server_status` events.
///
/// ## RPIT, not `#[async_trait]`
///
/// rmcp 2.1's [`ClientHandler`] declares its async methods as
/// return-position `impl Future` (see
/// `~/.cargo/registry/src/.../rmcp-2.1.0/src/handler/client.rs`,
/// lines 202217). Applying `#[async_trait]` here would produce
/// methods whose signature mismatches the trait, and the impl would
/// not satisfy the bound. The macro path is also unnecessary — the
/// trait already supports `async fn` syntax indirectly via
/// `impl Future<Output = ()> + Send + '_`, which is what we mirror.
///
/// Future contributor reading this: do **not** add `#[async_trait]`.
/// The methods below intentionally return `impl Future` directly.
///
/// ## Notification routing
///
/// `on_tool_list_changed` / `on_resource_list_changed` push an
/// [`McpClientEvent`] into [`Self::notify_tx`]. If the receiver has
/// been dropped (subagent teardown, session shutdown, or the field
/// was `None` to begin with — see [`McpClient::notify_tx`] doc), the
/// send fails silently; rmcp must not see an error from a
/// notification handler or the service loop tears down.
#[derive(Debug)]
pub struct KigiClientHandler {
/// Static `ClientInfo` returned by [`Self::get_info`]; built once
/// at handshake time and stored to avoid re-allocating per call.
info: ClientInfo,
/// MCP server name this handler is bound to. Cloned into emitted
/// events so the dispatcher can route per-server.
server_name: McpServerName,
/// **Shared** event sink — the same Arc lives on the owning
/// [`McpClient`]. Mutating the slot via [`McpClient::set_event_tx`]
/// is observed here on the next read, so wiring the sender
/// post-handshake is supported without restarting the rmcp
/// service loop.
notify_tx: SharedEventTx,
}
impl KigiClientHandler {
/// Best-effort event emit. Reads the shared `notify_tx` slot on
/// every call (so the handler picks up any post-handshake wiring
/// done by [`McpClient::set_event_tx`]). Drops the send error: if
/// the receiver is gone, the consumer has shut down and there's
/// nothing useful to do here. Splitting this out keeps the trait
/// methods short.
fn emit(&self, ev: McpClientEvent) {
let sender = self.notify_tx.lock().clone();
if let Some(tx) = sender {
let _ = tx.send(ev);
}
}
}
impl ClientHandler for KigiClientHandler {
// NOTE: `async fn` here is sugar for the trait's
// `-> impl Future<Output = ()> + Send + '_`. We INTENTIONALLY do
// not use `#[async_trait]` — rmcp 2.1's `ClientHandler` declares
// its notification methods as return-position `impl Future`, and
// async_trait would produce a different (incompatible) signature.
// See the [`KigiClientHandler`] doc-comment for the full RPIT
// contract.
async fn on_tool_list_changed(&self, _context: NotificationContext<RoleClient>) {
self.emit(McpClientEvent::ToolsChanged {
server: self.server_name.clone(),
});
}
async fn on_resource_list_changed(&self, _context: NotificationContext<RoleClient>) {
self.emit(McpClientEvent::ResourcesChanged {
server: self.server_name.clone(),
});
}
fn get_info(&self) -> ClientInfo {
self.info.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
/// A single undecodable line on an MCP stdio server's stdout must NOT
/// collapse the transport: if the decode error surfaced as `None`, the
/// service would read it as EOF → "Transport closed" → `tools/list` fails
/// and the connector "shows but doesn't work". The resilient transport
/// skips the bad line and keeps reading, so a stray stdout log line never
/// takes the whole server down.
#[tokio::test]
async fn resilient_transport_skips_undecodable_line_and_keeps_stream_alive() {
// `server_out` is the writer half (the fake server's stdout); the
// transport reads framed JSON-RPC from `client_in`.
let (mut server_out, client_in) = tokio::io::duplex(64 * 1024);
let mut transport = ResilientRwTransport::new(
client_in,
tokio::io::sink(),
"fwbuild".to_string(),
kigi_file_utils::events::EventWriter::noop(),
);
let valid = r#"{"jsonrpc":"2.0","method":"notifications/tools/list_changed"}"#;
// A stray non-JSON log line — the shape that, under rmcp's stock
// transport, decodes to an error and closes the connection.
let garbage = "info: fwbuild started, listening on stdio";
server_out
.write_all(format!("{valid}\n{garbage}\n{valid}\n").as_bytes())
.await
.unwrap();
// Dropping the writer half signals a clean end-of-stream.
drop(server_out);
assert!(
transport.receive().await.is_some(),
"first valid message must be received"
);
assert!(
transport.receive().await.is_some(),
"the undecodable line must be skipped and the next valid message delivered"
);
assert!(
transport.receive().await.is_none(),
"only a genuine end-of-stream yields None"
);
}
fn make_stdio_server(name: &str, command: &str) -> acp::McpServer {
acp::McpServer::Stdio(acp::McpServerStdio::new(name, PathBuf::from(command)))
}
fn make_http_server(name: &str, url: &str) -> acp::McpServer {
acp::McpServer::Http(acp::McpServerHttp::new(name, url))
}
#[test]
fn plan_stdio_spawn_windows_resolves_bare_launcher_to_cmd_shim() {
let args = vec!["-y".to_string(), "@scope/pkg".to_string()];
let (program, spawn_args) = plan_stdio_spawn("npx", &args, true, |c| {
assert_eq!(c, "npx");
Some(PathBuf::from(r"C:\path\npx.cmd"))
});
assert_eq!(program, OsString::from(r"C:\path\npx.cmd"));
assert_eq!(spawn_args, args);
}
#[test]
fn plan_stdio_spawn_windows_unresolved_falls_back_to_raw_command() {
let args = vec!["-y".to_string(), "@scope/pkg".to_string()];
let (program, spawn_args) = plan_stdio_spawn("npx", &args, true, |_| None);
assert_eq!(program, OsString::from("npx"));
assert_eq!(spawn_args, args);
}
#[test]
fn plan_stdio_spawn_windows_backslash_path_command_used_as_is_without_resolving() {
let args = vec!["--config".to_string(), "x.json".to_string()];
let (program, spawn_args) = plan_stdio_spawn(r"C:\tools\server.exe", &args, true, |_| {
panic!("resolver must not be consulted for a command with a backslash separator")
});
assert_eq!(program, OsString::from(r"C:\tools\server.exe"));
assert_eq!(spawn_args, args);
}
#[test]
fn plan_stdio_spawn_windows_forward_slash_path_command_used_as_is_without_resolving() {
let args = vec!["--port".to_string(), "8080".to_string()];
let (program, spawn_args) = plan_stdio_spawn("C:/tools/server.exe", &args, true, |_| {
panic!("resolver must not be consulted for a command with a forward-slash separator")
});
assert_eq!(program, OsString::from("C:/tools/server.exe"));
assert_eq!(spawn_args, args);
}
#[test]
fn plan_stdio_spawn_non_windows_never_resolves() {
let args = vec!["-y".to_string(), "pkg".to_string()];
let (program, spawn_args) = plan_stdio_spawn("npx", &args, false, |_| {
panic!("resolver must not be consulted on non-Windows")
});
assert_eq!(program, OsString::from("npx"));
assert_eq!(spawn_args, args);
}
#[test]
fn stdio_path_override_matches_path_case_insensitively() {
let mk = |name: &str, value: &str| acp::EnvVariable::new(name, value);
let env = vec![mk("FOO", "bar"), mk("Path", r"C:\node")];
assert_eq!(stdio_path_override(&env), Some(r"C:\node"));
let env_upper = vec![mk("PATH", "/custom/bin")];
assert_eq!(stdio_path_override(&env_upper), Some("/custom/bin"));
let env_none = vec![mk("FOO", "bar")];
assert_eq!(stdio_path_override(&env_none), None);
}
#[test]
fn is_figma_mcp_matches_name_and_host() {
assert!(is_figma_mcp("figma", "https://example.com/mcp"));
assert!(is_figma_mcp("Figma", "https://example.com/mcp"));
assert!(is_figma_mcp("other", "https://mcp.figma.com/mcp"));
assert!(is_figma_mcp("other", "https://figma.com/mcp"));
assert!(!is_figma_mcp("linear", "https://mcp.linear.app/mcp"));
assert!(!is_figma_mcp("figma_extra", "https://example.com/mcp"));
assert!(!is_figma_mcp("linear", "not-a-url"));
assert!(!is_figma_mcp("linear", "https://notfigma.com/mcp"));
assert!(!is_figma_mcp("linear", "https://figma.com.evil/mcp"));
}
#[test]
fn ensure_figma_user_agent_sets_kigi_cli_when_missing() {
let mut headers = reqwest::header::HeaderMap::new();
ensure_figma_user_agent(&mut headers, "figma", "https://mcp.figma.com/mcp");
assert_eq!(
headers.get(reqwest::header::USER_AGENT).unwrap(),
"kigi-cli"
);
let mut host_only = reqwest::header::HeaderMap::new();
ensure_figma_user_agent(&mut host_only, "other", "https://mcp.figma.com/mcp");
assert_eq!(
host_only.get(reqwest::header::USER_AGENT).unwrap(),
"kigi-cli"
);
}
#[test]
fn ensure_figma_user_agent_does_not_overwrite_existing() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::USER_AGENT,
reqwest::header::HeaderValue::from_static("custom-ua"),
);
ensure_figma_user_agent(&mut headers, "figma", "https://mcp.figma.com/mcp");
assert_eq!(
headers.get(reqwest::header::USER_AGENT).unwrap(),
"custom-ua"
);
}
#[test]
fn ensure_figma_user_agent_skips_non_figma() {
let mut headers = reqwest::header::HeaderMap::new();
ensure_figma_user_agent(&mut headers, "linear", "https://mcp.linear.app/mcp");
assert!(!headers.contains_key(reqwest::header::USER_AGENT));
let mut invalid_url = reqwest::header::HeaderMap::new();
ensure_figma_user_agent(&mut invalid_url, "linear", "not-a-url");
assert!(!invalid_url.contains_key(reqwest::header::USER_AGENT));
}
#[cfg(unix)]
#[test]
fn safe_stdio_child_drop_without_entered_runtime_reaps_child() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("test runtime");
let (transport, pid) = rt.block_on(async {
let mut cmd = Command::new("sleep");
cmd.arg("30").kill_on_drop(true);
kigi_tools::util::detach_command(&mut cmd);
let (transport, _stderr) = SafeTokioChildProcess::spawn(
cmd,
"test".to_string(),
kigi_file_utils::events::EventWriter::noop(),
)
.expect("spawn test child");
let pid = transport.id().expect("spawned child pid");
(transport, pid)
});
drop(rt);
drop(transport);
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
while std::time::Instant::now() < deadline {
if !unix_process_exists(pid) {
return;
}
std::thread::sleep(std::time::Duration::from_millis(25));
}
panic!("MCP child process {pid} was not reaped after no-runtime drop");
}
#[cfg(unix)]
fn unix_process_exists(pid: u32) -> bool {
let result = unsafe { libc::kill(pid as libc::pid_t, 0) };
if result == 0 {
return true;
}
std::io::Error::last_os_error().raw_os_error() != Some(libc::ESRCH)
}
#[test]
fn test_mcp_state_new() {
let configs = vec![make_stdio_server("test", "/bin/test")];
let state = McpState::new(configs.clone());
assert_eq!(state.configs.len(), 1);
assert!(state.owned_clients.is_empty());
assert!(!state.is_initialized());
assert!(!state.is_initializing());
assert!(!state.has_finished_init());
assert!(matches!(state.init_progress(), InitProgress::NotStarted));
assert_eq!(state.generation, 0);
}
#[test]
fn test_mcp_state_update_configs_returns_false_when_unchanged() {
let configs = vec![make_stdio_server("test", "/bin/test")];
let mut state = McpState::new(configs.clone());
// Same configs should return false
let changed = state.update_configs(configs.clone());
assert!(!changed);
assert_eq!(state.generation, 0); // Generation should not change
}
#[test]
fn test_mcp_state_update_configs_returns_true_when_changed() {
let configs = vec![make_stdio_server("test", "/bin/test")];
let mut state = McpState::new(configs);
// Different configs should return true
let new_configs = vec![make_stdio_server("test2", "/bin/test2")];
let changed = state.update_configs(new_configs);
assert!(changed);
assert_eq!(state.generation, 1); // Generation should increment
}
#[test]
fn test_mcp_state_update_configs_resets_initialized() {
let configs = vec![make_stdio_server("test", "/bin/test")];
let mut state = McpState::new(configs);
// Drive the state machine into Finished{handshaking:{"a"}} so
// the reset path has both the lifecycle flag AND a per-server
// entry to clear.
assert!(state.try_start_init());
state.mark_servers_initializing(["a".to_string()]);
state.finish_init();
assert!(state.has_finished_init());
assert!(state.is_server_handshaking("a"));
let new_configs = vec![make_stdio_server("test2", "/bin/test2")];
let changed = state.update_configs(new_configs);
assert!(changed);
// update_configs must drop us back to NotStarted — neither
// lifecycle flag set nor any per-server progress carried over.
assert!(!state.is_initialized());
assert!(!state.is_initializing());
assert!(!state.has_finished_init());
assert!(matches!(state.init_progress(), InitProgress::NotStarted));
}
#[tokio::test]
async fn acp_servers_survive_update_configs_clear() {
use crate::acp_transport::AcpReverseInvoker;
use std::time::Duration;
struct NoopInvoker;
#[async_trait::async_trait]
impl AcpReverseInvoker for NoopInvoker {
async fn invoke(
&self,
_server_id: &str,
_message: serde_json::Value,
_timeout: Duration,
) -> Result<serde_json::Value, String> {
Ok(serde_json::Value::Null)
}
}
let mut state = McpState::new(vec![make_http_server("http-srv", "http://localhost")]);
state.set_acp_servers(
vec![AcpServerEntry {
name: "sdk-tools".to_string(),
server_id: "srv_0".to_string(),
}],
Arc::new(NoopInvoker),
);
assert!(state.has_acp_servers());
assert_eq!(state.build_pending_acp_clients(&HashMap::new()).len(), 1);
// A config change clears owned clients/configs (proven by the generation bump)
// but must NOT drop the separately-held acp servers — otherwise the in-process
// SDK tools would silently vanish on every `update_configs`.
let changed = state.update_configs(vec![make_http_server("other", "http://other")]);
assert!(changed);
assert_eq!(state.generation, 1);
assert!(
state.has_acp_servers(),
"acp servers must survive update_configs"
);
let pending = state.build_pending_acp_clients(&HashMap::new());
assert_eq!(pending.len(), 1, "acp clients rebuild after the clear");
assert_eq!(pending[0].server_name(), "sdk-tools");
}
#[tokio::test]
async fn acp_overrides_apply_to_built_clients() {
use crate::acp_transport::AcpReverseInvoker;
use std::time::Duration;
struct NoopInvoker;
#[async_trait::async_trait]
impl AcpReverseInvoker for NoopInvoker {
async fn invoke(
&self,
_server_id: &str,
_message: serde_json::Value,
_timeout: Duration,
) -> Result<serde_json::Value, String> {
Ok(serde_json::Value::Null)
}
}
let mut overrides = HashMap::new();
overrides.insert(
"sdk-tools".to_string(),
McpClientTimeoutOverrides {
tool_timeout_sec: Some(123),
..Default::default()
},
);
let mut state = McpState::new(vec![]);
state.set_acp_servers(
vec![AcpServerEntry {
name: "sdk-tools".to_string(),
server_id: "srv_0".to_string(),
}],
Arc::new(NoopInvoker),
);
let pending = state.build_pending_acp_clients(&overrides);
assert_eq!(pending.len(), 1);
assert_eq!(
pending[0].tool_timeout_sec(),
123,
"config.toml tool_timeout_sec override must reach the SDK client"
);
}
/// In-process SDK (ACP) clients must never get a liveness watcher: the
/// dispatcher can't recover them (no `configs` entry), so a proactive
/// `TransportClosed` would evict the client with no recovery. Guards both
/// the `is_acp` predicate (across transports) and the `arm_liveness_watcher`
/// self-gate that depends on it. HTTP/stdio must report `false` so they
/// keep their watchers.
#[tokio::test]
async fn acp_clients_are_not_liveness_watched() {
use crate::acp_transport::AcpReverseInvoker;
use std::time::Duration;
struct NoopInvoker;
#[async_trait::async_trait]
impl AcpReverseInvoker for NoopInvoker {
async fn invoke(
&self,
_server_id: &str,
_message: serde_json::Value,
_timeout: Duration,
) -> Result<serde_json::Value, String> {
Ok(serde_json::Value::Null)
}
}
let acp = McpClient::new_acp(
"sdk".to_string(),
"srv_0".to_string(),
Arc::new(NoopInvoker),
None,
None,
);
assert!(acp.is_acp());
assert!(!acp.is_http());
let http = McpClient::new_http(
"http".to_string(),
HttpConfig {
url: "http://localhost/api/mcp".to_string(),
headers: vec![],
},
None,
None,
);
assert!(!http.is_acp());
// Stub stands in for a no-transport / Stdio client (reconnect = None).
assert!(!McpClient::stub("stdio").is_acp());
// The gate that prevents the evict-on-close bug: arming is a no-op for ACP.
assert!(
!Arc::new(acp)
.arm_liveness_watcher(Duration::from_millis(500))
.await
);
}
#[test]
fn test_mark_servers_initializing_clears_prior_init_failure() {
// A server that failed a previous init is recorded in `init_failed`
// (so the status snapshot reports it Unavailable). Starting a fresh
// init attempt for that server must clear the failure flag so a
// successful retry can surface as Ready again.
let mut state = McpState::new(vec![make_stdio_server("a", "/bin/a")]);
state.init_failed.insert("a".to_string(), String::new());
state.init_failed.insert("b".to_string(), String::new());
state.mark_servers_initializing(["a".to_string()]);
assert!(
!state.init_failed.contains_key("a"),
"fresh init attempt must clear the prior failure for that server",
);
assert!(
state.init_failed.contains_key("b"),
"servers not in this init attempt must keep their failure flag",
);
}
#[test]
fn test_record_init_failure_keeps_auth_and_init_failed_disjoint() {
let mut state = McpState::new(vec![make_stdio_server("a", "/bin/a")]);
// Auth failures are owned by `auth_required` only — never `init_failed` —
// so a later successful authentication (which clears `auth_required` and
// registers tools) is not left stuck as Unavailable with zero tools.
state.record_init_failure("auth-srv", true, None);
assert!(state.auth_required.contains("auth-srv"));
assert!(
!state.init_failed.contains_key("auth-srv"),
"auth-required failures must not also be flagged init_failed",
);
// Non-auth failures (handshake/`tools/list` error or timeout) → init_failed,
// and their cause is retained for the model-facing reminder.
state.record_init_failure(
"dead-srv",
false,
Some("tools/list failed: boom".to_string()),
);
assert!(!state.auth_required.contains("dead-srv"));
assert_eq!(
state.init_failed.get("dead-srv").map(String::as_str),
Some("tools/list failed: boom"),
);
// A fresh init attempt clears the failure entry and its cause.
state.mark_servers_initializing(["dead-srv".to_string()]);
assert!(!state.init_failed.contains_key("dead-srv"));
}
#[test]
fn test_clear_init_failed_removes_entry() {
let mut state = McpState::new(vec![make_stdio_server("a", "/bin/a")]);
state.record_init_failure("dead-srv", false, Some("boom".to_string()));
assert!(state.init_failed.contains_key("dead-srv"));
// Symmetric with record_init_failure: the reactive re-auth path clears
// a prior failure so a recovered server is not stuck Unavailable.
state.clear_init_failed("dead-srv");
assert!(!state.init_failed.contains_key("dead-srv"));
// Idempotent: clearing an absent entry is a no-op.
state.clear_init_failed("never-seen");
}
#[test]
fn test_mcp_state_update_configs_increments_generation() {
let mut state = McpState::new(vec![]);
// Each change should increment generation
state.update_configs(vec![make_stdio_server("a", "/bin/a")]);
assert_eq!(state.generation, 1);
state.update_configs(vec![make_stdio_server("b", "/bin/b")]);
assert_eq!(state.generation, 2);
state.update_configs(vec![make_stdio_server("c", "/bin/c")]);
assert_eq!(state.generation, 3);
}
#[test]
fn test_mcp_servers_equal_empty_lists() {
let a: Vec<acp::McpServer> = vec![];
let b: Vec<acp::McpServer> = vec![];
assert!(mcp_servers_equal(&a, &b));
}
#[test]
fn test_mcp_servers_equal_identical_configs() {
let a = vec![make_stdio_server("test", "/bin/test")];
let b = vec![make_stdio_server("test", "/bin/test")];
assert!(mcp_servers_equal(&a, &b));
}
#[test]
fn test_mcp_servers_equal_different_names() {
let a = vec![make_stdio_server("test1", "/bin/test")];
let b = vec![make_stdio_server("test2", "/bin/test")];
assert!(!mcp_servers_equal(&a, &b));
}
#[test]
fn test_mcp_servers_equal_different_lengths() {
let a = vec![make_stdio_server("test", "/bin/test")];
let b = vec![
make_stdio_server("test", "/bin/test"),
make_stdio_server("test2", "/bin/test2"),
];
assert!(!mcp_servers_equal(&a, &b));
}
#[test]
fn test_mcp_servers_equal_different_types() {
let a = vec![make_stdio_server("test", "/bin/test")];
let b = vec![make_http_server("test", "http://localhost")];
assert!(!mcp_servers_equal(&a, &b));
}
#[test]
fn test_mcp_servers_equal_order_matters() {
let a = vec![
make_stdio_server("a", "/bin/a"),
make_stdio_server("b", "/bin/b"),
];
let b = vec![
make_stdio_server("b", "/bin/b"),
make_stdio_server("a", "/bin/a"),
];
// Order matters since we're comparing JSON serialization
assert!(!mcp_servers_equal(&a, &b));
}
#[test]
fn test_try_start_init_prevents_concurrent_init() {
let mut state = McpState::new(vec![make_stdio_server("test", "/bin/test")]);
// First call should succeed
assert!(state.try_start_init());
assert!(state.is_initializing());
assert!(!state.is_initialized());
// Second call should fail (already initializing)
assert!(!state.try_start_init());
}
#[test]
fn test_try_start_init_fails_when_initialized() {
let mut state = McpState::new(vec![make_stdio_server("test", "/bin/test")]);
// Drive to Finished{empty} via the typed API.
assert!(state.try_start_init());
state.finish_init();
assert!(state.is_initialized());
// Second `try_start_init` must be rejected: we're already done.
assert!(!state.try_start_init());
assert!(!state.is_initializing());
assert!(state.is_initialized(), "is_initialized stays true");
}
#[test]
fn test_finish_init_clears_initializing() {
let mut state = McpState::new(vec![make_stdio_server("test", "/bin/test")]);
state.try_start_init();
assert!(state.is_initializing());
assert!(!state.is_initialized());
state.finish_init();
assert!(!state.is_initializing());
assert!(state.is_initialized());
}
#[test]
fn test_cancel_init_clears_initializing() {
let mut state = McpState::new(vec![make_stdio_server("test", "/bin/test")]);
state.try_start_init();
assert!(state.is_initializing());
state.cancel_init();
assert!(!state.is_initializing());
assert!(!state.is_initialized()); // Should NOT be marked as initialized
}
#[test]
fn test_update_configs_resets_initializing() {
let mut state = McpState::new(vec![make_stdio_server("test", "/bin/test")]);
state.try_start_init();
assert!(state.is_initializing());
// Updating configs should reset initializing flag
state.update_configs(vec![make_stdio_server("test2", "/bin/test2")]);
assert!(!state.is_initializing());
assert!(!state.is_initialized());
}
#[test]
fn test_parse_mcp_meta_config_with_tool_timeouts_ms() {
let meta = serde_json::json!({
"mcpConfig": {
"github": {
"toolTimeoutMs": 60000,
"toolTimeoutsMs": {
"create_issue": 120000,
"search": 30000
}
}
}
})
.as_object()
.cloned()
.unwrap();
let map = parse_mcp_meta_config(Some(&meta));
let github = map.get("github").unwrap();
assert_eq!(github.tool_timeout_ms, Some(60000));
let tt = github.tool_timeouts_ms.as_ref().unwrap();
assert_eq!(tt.get("create_issue"), Some(&120000));
assert_eq!(tt.get("search"), Some(&30000));
}
#[test]
fn test_parse_mcp_meta_config_without_tool_timeouts_ms() {
let meta = serde_json::json!({
"mcpConfig": {
"github": {
"toolTimeoutMs": 60000
}
}
})
.as_object()
.cloned()
.unwrap();
let map = parse_mcp_meta_config(Some(&meta));
let github = map.get("github").unwrap();
assert_eq!(github.tool_timeout_ms, Some(60000));
assert!(github.tool_timeouts_ms.is_none());
assert!(github.expose_image_base64.is_none());
}
/// Locks in the `exposeImageBase64` camelCase wire-format contract.
#[test]
fn test_parse_mcp_meta_config_with_expose_image_base64() {
let meta = serde_json::json!({
"mcpConfig": {
"grafana": { "exposeImageBase64": true },
"linear": { "exposeImageBase64": false },
}
})
.as_object()
.cloned()
.unwrap();
let map = parse_mcp_meta_config(Some(&meta));
assert_eq!(map.get("grafana").unwrap().expose_image_base64, Some(true));
assert_eq!(map.get("linear").unwrap().expose_image_base64, Some(false));
}
#[test]
fn test_tool_timeout_for_returns_per_tool_override() {
let mut tool_timeouts = HashMap::new();
tool_timeouts.insert("create_issue".to_string(), 120u64);
tool_timeouts.insert("search".to_string(), 30u64);
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(10),
tool_timeout_sec: Some(60),
tool_timeouts: Some(tool_timeouts),
..Default::default()
};
let client = McpClient::new_http(
"github".to_string(),
HttpConfig {
url: String::new(),
headers: vec![],
},
Some(&overrides),
None,
);
// Per-tool overrides
assert_eq!(client.tool_timeout_for("create_issue"), 120);
assert_eq!(client.tool_timeout_for("search"), 30);
// Falls back to server-level default
assert_eq!(client.tool_timeout_for("list_repos"), 60);
assert_eq!(client.tool_timeout_for(""), 60);
}
#[test]
fn test_tool_timeout_for_empty_map_returns_default() {
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(10),
tool_timeout_sec: Some(45),
..Default::default()
};
let client = McpClient::new_http(
"test".to_string(),
HttpConfig {
url: String::new(),
headers: vec![],
},
Some(&overrides),
None,
);
// All tools should get the server-level default
assert_eq!(client.tool_timeout_for("any_tool"), 45);
assert_eq!(client.tool_timeout_sec(), 45);
}
#[test]
fn test_load_timeouts_startup_precedence() {
// No override -> the standalone default (env/config resolved by the shell).
assert_eq!(
McpClient::load_timeouts(None, None).0,
DEFAULT_STARTUP_TIMEOUT_SECS
);
// A per-server `startup_timeout_sec` (injected by the shell) wins over the default...
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(7),
..Default::default()
};
assert_eq!(McpClient::load_timeouts(Some(&overrides), None).0, 7);
// ...and `_meta.startup_timeout_ms` wins over that.
let meta = McpServerMetaConfig {
startup_timeout_ms: Some(12_000),
..Default::default()
};
assert_eq!(
McpClient::load_timeouts(Some(&overrides), Some(&meta)).0,
12
);
}
#[test]
fn test_update_configs_diff_no_change() {
let configs = vec![make_stdio_server("test", "/bin/test")];
let mut state = McpState::new(configs.clone());
assert!(state.update_configs_diff(configs).is_none());
assert_eq!(state.generation, 0);
}
#[test]
fn test_update_configs_diff_added() {
let configs = vec![make_stdio_server("a", "/bin/a")];
let mut state = McpState::new(configs);
let new_configs = vec![
make_stdio_server("a", "/bin/a"),
make_stdio_server("b", "/bin/b"),
];
let diff = state
.update_configs_diff(new_configs)
.expect("should detect change");
assert_eq!(diff.retained, vec!["a"]);
assert_eq!(diff.added, vec!["b"]);
assert!(diff.removed.is_empty());
assert_eq!(state.generation, 1);
}
#[test]
fn test_update_configs_diff_removed() {
let configs = vec![
make_stdio_server("a", "/bin/a"),
make_stdio_server("b", "/bin/b"),
];
let mut state = McpState::new(configs);
let new_configs = vec![make_stdio_server("a", "/bin/a")];
let diff = state
.update_configs_diff(new_configs)
.expect("should detect change");
assert_eq!(diff.retained, vec!["a"]);
assert!(diff.added.is_empty());
assert_eq!(diff.removed, vec!["b"]);
}
#[test]
fn test_update_configs_diff_changed() {
let configs = vec![make_stdio_server("a", "/bin/a")];
let mut state = McpState::new(configs);
let new_configs = vec![make_stdio_server("a", "/bin/a_v2")];
let diff = state
.update_configs_diff(new_configs)
.expect("should detect change");
assert!(diff.retained.is_empty());
assert_eq!(diff.added, vec!["a"]);
assert_eq!(diff.removed, vec!["a"]);
}
#[test]
fn test_update_configs_diff_auth_required_cleanup() {
let configs = vec![
make_stdio_server("keep", "/bin/keep"),
make_stdio_server("remove", "/bin/remove"),
];
let mut state = McpState::new(configs);
state.auth_required.insert("remove".to_string());
state.auth_required.insert("keep".to_string());
let new_configs = vec![make_stdio_server("keep", "/bin/keep")];
let diff = state
.update_configs_diff(new_configs)
.expect("should detect change");
assert_eq!(diff.retained, vec!["keep"]);
assert_eq!(diff.removed, vec!["remove"]);
assert!(state.auth_required.contains("keep"));
assert!(!state.auth_required.contains("remove"));
}
#[test]
fn test_update_configs_diff_empty_to_nonempty() {
let mut state = McpState::new(vec![]);
let new_configs = vec![make_stdio_server("a", "/bin/a")];
let diff = state
.update_configs_diff(new_configs)
.expect("should detect change");
assert!(diff.retained.is_empty());
assert_eq!(diff.added, vec!["a"]);
assert!(diff.removed.is_empty());
}
#[test]
fn test_update_configs_diff_nonempty_to_empty() {
let configs = vec![make_stdio_server("a", "/bin/a")];
let mut state = McpState::new(configs);
let diff = state
.update_configs_diff(vec![])
.expect("should detect change");
assert!(diff.retained.is_empty());
assert!(diff.added.is_empty());
assert_eq!(diff.removed, vec!["a"]);
}
/// Two MCP servers exposing a tool with the same raw name must produce
/// `McpErasedTool` instances with **distinct** `ToolId`s (qualified with
/// the server name). Regression test for a bug where `McpErasedTool::id()`
/// returned the unqualified name, causing the second registration to
/// silently overwrite the first in the `LocalRegistry`.
#[test]
fn test_mcp_erased_tool_id_is_qualified() {
use kigi_tool_runtime::Tool;
let mcp_state = Arc::new(Mutex::new(McpState::new(vec![])));
let tool_a = McpErasedTool {
tool: McpTool::new(
"SearchUsers".to_string(),
"Search users".to_string(),
"calendar".to_string(),
Arc::clone(&mcp_state),
serde_json::json!({"type": "object"}),
None,
),
};
let tool_b = McpErasedTool {
tool: McpTool::new(
"SearchUsers".to_string(),
"Search users".to_string(),
"teams".to_string(),
Arc::clone(&mcp_state),
serde_json::json!({"type": "object"}),
None,
),
};
let id_a = tool_a.id();
let id_b = tool_b.id();
// IDs must be qualified with the server name.
assert_eq!(id_a.as_str(), "calendar__SearchUsers");
assert_eq!(id_b.as_str(), "teams__SearchUsers");
// And therefore distinct.
assert_ne!(id_a, id_b);
}
/// Registering two MCP tools with the same raw name from different servers
/// into a `LocalRegistry` must preserve both entries (no silent overwrite).
#[test]
fn test_same_raw_name_different_servers_no_local_registry_collision() {
use kigi_tool_runtime::LocalRegistry;
use kigi_tool_runtime::Tool;
let mcp_state = Arc::new(Mutex::new(McpState::new(vec![])));
let registry = LocalRegistry::new();
let tool_a = McpErasedTool {
tool: McpTool::new(
"SearchUsers".to_string(),
"Search users on calendar".to_string(),
"calendar".to_string(),
Arc::clone(&mcp_state),
serde_json::json!({"type": "object"}),
None,
),
};
let tool_b = McpErasedTool {
tool: McpTool::new(
"SearchUsers".to_string(),
"Search users on teams".to_string(),
"teams".to_string(),
Arc::clone(&mcp_state),
serde_json::json!({"type": "object"}),
None,
),
};
let id_a = tool_a.id();
let id_b = tool_b.id();
// First registration should not displace anything.
let displaced_a = registry.register(tool_a);
assert!(
displaced_a.is_none(),
"first registration should not displace"
);
// Second registration should also not displace anything (distinct IDs).
let displaced_b = registry.register(tool_b);
assert!(
displaced_b.is_none(),
"second registration must not overwrite first"
);
// Both tools must be independently resolvable.
assert!(
registry.find(&id_a).is_some(),
"calendar tool must be found"
);
assert!(registry.find(&id_b).is_some(), "teams tool must be found");
assert_eq!(registry.len(), 2);
}
fn make_test_client(name: &str) -> Arc<McpClient> {
// Same shape as the no-transport placeholder.
Arc::new(McpClient::stub(name))
}
#[test]
fn test_shared_mcp_pool_from_empty_state() {
let state = McpState::new(vec![]);
let pool = SharedMcpPool::from_state(&state);
assert_eq!(pool.len(), 0);
assert_eq!(pool.server_names().count(), 0);
assert!(pool.configs().is_empty());
assert!(pool.meta_config_map().is_empty());
assert!(pool.get_client("anything").is_none());
}
#[test]
fn test_shared_mcp_pool_len_matches_client_count() {
let mut state = McpState::new(vec![]);
for name in ["alpha", "beta", "gamma"] {
state
.owned_clients
.insert(name.to_string(), make_test_client(name));
}
let pool = SharedMcpPool::from_state(&state);
assert_eq!(pool.len(), 3);
assert_eq!(pool.len(), pool.server_names().count());
}
#[test]
fn test_shared_mcp_pool_snapshot_shares_arc_clients() {
let mut state = McpState::new(vec![make_stdio_server("github", "/bin/gh")]);
let client = make_test_client("github");
state
.owned_clients
.insert("github".to_string(), Arc::clone(&client));
let pool = SharedMcpPool::from_state(&state);
let pool_client = pool.get_client("github").expect("should find client");
// Must point to the same allocation (shared transport)
assert!(Arc::ptr_eq(&client, pool_client));
}
#[test]
fn test_shared_mcp_pool_get_client_missing() {
let mut state = McpState::new(vec![]);
state
.owned_clients
.insert("a".to_string(), make_test_client("a"));
let pool = SharedMcpPool::from_state(&state);
assert!(pool.get_client("a").is_some());
assert!(pool.get_client("nonexistent").is_none());
assert!(pool.get_client("").is_none());
}
#[test]
fn test_shared_mcp_pool_server_names() {
let mut state = McpState::new(vec![]);
for name in ["alpha", "beta", "gamma"] {
state
.owned_clients
.insert(name.to_string(), make_test_client(name));
}
let pool = SharedMcpPool::from_state(&state);
let mut names: Vec<&str> = pool.server_names().collect();
names.sort();
assert_eq!(names, vec!["alpha", "beta", "gamma"]);
}
#[test]
fn test_shared_mcp_pool_snapshot_independent_of_state_mutations() {
let mut state = McpState::new(vec![make_stdio_server("srv", "/bin/srv")]);
state
.owned_clients
.insert("srv".to_string(), make_test_client("srv"));
let pool = SharedMcpPool::from_state(&state);
// Mutate state after snapshot
state.owned_clients.clear();
state.configs.clear();
// Pool retains original data
assert_eq!(pool.server_names().count(), 1);
assert!(pool.get_client("srv").is_some());
assert_eq!(pool.configs().len(), 1);
}
#[test]
fn test_shared_mcp_pool_meta_config_preserved() {
let mut meta = McpMetaConfigMap::new();
meta.insert(
"github".to_string(),
McpServerMetaConfig {
startup_timeout_ms: Some(5000),
tool_timeout_ms: Some(120000),
tool_timeouts_ms: None,
expose_image_base64: None,
},
);
let state =
McpState::new_with_meta(vec![make_http_server("github", "http://gh.local")], meta);
let pool = SharedMcpPool::from_state(&state);
let mc = pool
.meta_config_map()
.get("github")
.expect("should have meta config");
assert_eq!(mc.startup_timeout_ms, Some(5000));
assert_eq!(mc.tool_timeout_ms, Some(120000));
}
#[test]
fn test_shared_mcp_pool_clone_shares_arcs() {
let mut state = McpState::new(vec![]);
let client = make_test_client("svc");
state
.owned_clients
.insert("svc".to_string(), Arc::clone(&client));
let pool = SharedMcpPool::from_state(&state);
let pool2 = pool.clone();
// Both clones share the same Arc<McpClient>
let c1 = pool.get_client("svc").unwrap();
let c2 = pool2.get_client("svc").unwrap();
assert!(Arc::ptr_eq(c1, c2));
}
// ── owned/shared split behavioral tests ─────────────────────────
#[test]
fn test_get_client_owned_overrides_shared() {
let mut state = McpState::new(vec![]);
let shared = make_test_client("srv");
let owned = make_test_client("srv");
state
.shared_clients
.insert("srv".to_string(), Arc::clone(&shared));
state
.owned_clients
.insert("srv".to_string(), Arc::clone(&owned));
let got = state.get_client("srv").unwrap();
assert!(Arc::ptr_eq(got, &owned));
assert!(!Arc::ptr_eq(got, &shared));
}
#[test]
fn test_get_client_falls_through_to_shared() {
let mut state = McpState::new(vec![]);
let shared = make_test_client("srv");
state
.shared_clients
.insert("srv".to_string(), Arc::clone(&shared));
let got = state.get_client("srv").unwrap();
assert!(Arc::ptr_eq(got, &shared));
assert!(state.get_client("missing").is_none());
}
#[test]
fn test_all_clients_deduplicates_shared_by_owned() {
let mut state = McpState::new(vec![]);
state
.owned_clients
.insert("a".to_string(), make_test_client("a"));
state
.shared_clients
.insert("a".to_string(), make_test_client("a-shared"));
state
.shared_clients
.insert("b".to_string(), make_test_client("b-shared"));
let all: Vec<_> = state.all_clients().map(|(n, _)| n.as_str()).collect();
// "a" appears once (from owned), "b" from shared
assert_eq!(all.iter().filter(|&&n| n == "a").count(), 1);
assert!(all.contains(&"b"));
assert_eq!(all.len(), 2);
// The "a" entry must be the owned client, not the shared one
let (_, a_client) = state.all_clients().find(|(n, _)| *n == "a").unwrap();
assert!(Arc::ptr_eq(a_client, state.owned_clients.get("a").unwrap()));
}
#[test]
fn test_import_shared_clients_skips_config_collisions() {
// Child has a config entry named "github" — importing a shared
// client with the same name must be skipped.
let mut state = McpState::new(vec![make_stdio_server("github", "/bin/gh")]);
let mut pool_clients = HashMap::new();
pool_clients.insert("github".to_string(), make_test_client("github"));
pool_clients.insert("linear".to_string(), make_test_client("linear"));
let pool = SharedMcpPool {
clients: pool_clients,
configs: vec![],
meta_config_map: McpMetaConfigMap::new(),
};
state.import_shared_clients(&pool);
assert!(
!state.shared_clients.contains_key("github"),
"github should be skipped — collides with child config"
);
assert!(
state.shared_clients.contains_key("linear"),
"linear should be imported — no collision"
);
}
#[test]
fn test_update_configs_preserves_shared_clients() {
let mut state = McpState::new(vec![make_stdio_server("old", "/bin/old")]);
state
.owned_clients
.insert("old".to_string(), make_test_client("old"));
let shared = make_test_client("inherited");
state
.shared_clients
.insert("inherited".to_string(), Arc::clone(&shared));
let changed = state.update_configs(vec![make_stdio_server("new", "/bin/new")]);
assert!(changed);
assert!(state.owned_clients.is_empty(), "owned should be cleared");
assert_eq!(state.shared_clients.len(), 1, "shared should be untouched");
assert!(Arc::ptr_eq(
state.shared_clients.get("inherited").unwrap(),
&shared
));
}
#[test]
fn test_update_configs_diff_preserves_shared_clients() {
let mut state = McpState::new(vec![
make_stdio_server("keep", "/bin/keep"),
make_stdio_server("drop", "/bin/drop"),
]);
state
.owned_clients
.insert("keep".to_string(), make_test_client("keep"));
state
.owned_clients
.insert("drop".to_string(), make_test_client("drop"));
let shared = make_test_client("inherited");
state
.shared_clients
.insert("inherited".to_string(), Arc::clone(&shared));
// New config removes "drop", keeps "keep"
let diff = state
.update_configs_diff(vec![make_stdio_server("keep", "/bin/keep")])
.expect("configs changed");
assert!(diff.removed.contains(&"drop".to_string()));
assert!(diff.retained.contains(&"keep".to_string()));
assert!(!state.owned_clients.contains_key("drop"));
assert!(state.owned_clients.contains_key("keep"));
// Shared clients must be completely untouched
assert!(Arc::ptr_eq(
state.shared_clients.get("inherited").unwrap(),
&shared
));
}
#[test]
fn test_from_state_captures_both_owned_and_shared() {
let mut state = McpState::new(vec![]);
let owned = make_test_client("owned-srv");
let shared = make_test_client("shared-srv");
state
.owned_clients
.insert("owned-srv".to_string(), Arc::clone(&owned));
state
.shared_clients
.insert("shared-srv".to_string(), Arc::clone(&shared));
let pool = SharedMcpPool::from_state(&state);
assert!(Arc::ptr_eq(pool.get_client("owned-srv").unwrap(), &owned));
assert!(Arc::ptr_eq(pool.get_client("shared-srv").unwrap(), &shared));
assert_eq!(pool.server_names().count(), 2);
}
#[test]
fn test_retain_clients_keeps_matching() {
let mut state = McpState::new(vec![]);
for name in ["github", "linear", "slack"] {
state
.owned_clients
.insert(name.to_string(), make_test_client(name));
}
let mut pool = SharedMcpPool::from_state(&state);
pool.retain_clients(|name| name == "github" || name == "slack");
assert!(pool.get_client("github").is_some());
assert!(pool.get_client("slack").is_some());
assert!(pool.get_client("linear").is_none());
assert_eq!(pool.server_names().count(), 2);
}
#[test]
fn test_retain_clients_remove_all() {
let mut state = McpState::new(vec![]);
state
.owned_clients
.insert("srv".to_string(), make_test_client("srv"));
let mut pool = SharedMcpPool::from_state(&state);
pool.retain_clients(|_| false);
assert_eq!(pool.server_names().count(), 0);
assert!(pool.get_client("srv").is_none());
}
#[test]
fn test_retain_clients_keep_all() {
let mut state = McpState::new(vec![]);
for name in ["a", "b", "c"] {
state
.owned_clients
.insert(name.to_string(), make_test_client(name));
}
let mut pool = SharedMcpPool::from_state(&state);
pool.retain_clients(|_| true);
assert_eq!(pool.server_names().count(), 3);
}
#[test]
fn test_retain_clients_preserves_arc_identity() {
let mut state = McpState::new(vec![]);
let client = make_test_client("keep");
state
.owned_clients
.insert("keep".to_string(), Arc::clone(&client));
state
.owned_clients
.insert("drop".to_string(), make_test_client("drop"));
let mut pool = SharedMcpPool::from_state(&state);
pool.retain_clients(|name| name == "keep");
assert!(Arc::ptr_eq(pool.get_client("keep").unwrap(), &client));
}
fn make_mcp_tool(server_name: &str, name: &str) -> McpTool {
McpTool::new(
name.to_string(),
"test desc".to_string(),
server_name.to_string(),
Arc::new(Mutex::new(McpState::new(vec![]))),
serde_json::json!({}),
None,
)
}
#[test]
fn into_registration_accepts_well_formed_segments() {
// Positive guard: the count check rejects `__`-anywhere-but-the-delimiter
// names but must not reject legitimate ones. If the rejection rule is
// ever tightened too far, this test breaks before any of the negative
// cases below.
let tool = make_mcp_tool("linear", "list_issues");
let reg = tool.into_registration().expect("should register");
assert_eq!(reg.name, "linear__list_issues");
}
#[test]
fn into_registration_rejects_boundary_ambiguity() {
// `"foo_"` + `"__"` + `"_bar"` => `"foo___bar"` has two valid
// `__` positions (indices 3 and 4), so `split_once("__")` would
// misparse it as `("foo", "_bar")` and silently auto-allow a
// future legitimate `"foo"` server. The naïve per-segment check
// (each side individually has no `__`) misses this — the count
// check catches it. Same rule also rejects "__-in-segment" cases
// (`"weird__server"` + `"list"`, `"linear"` + `"my__weird__tool"`)
// which are covered by the same `count() != 1` line of code.
let tool = make_mcp_tool("foo_", "_bar");
assert!(tool.into_registration().is_none());
}
// ── is_retriable_transport_error tests ───────────────────────────
#[test]
fn test_is_retriable_transport_closed() {
assert!(is_retriable_transport_error(&ServiceError::TransportClosed));
}
#[test]
fn test_is_retriable_transport_send() {
let err = ServiceError::TransportSend(rmcp::transport::DynamicTransportError::from_parts(
"test",
std::any::TypeId::of::<()>(),
Box::new(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"connection reset",
)),
));
assert!(is_retriable_transport_error(&err));
}
#[test]
fn test_not_retriable_unexpected_response() {
assert!(!is_retriable_transport_error(
&ServiceError::UnexpectedResponse
));
}
#[test]
fn test_not_retriable_cancelled() {
assert!(!is_retriable_transport_error(&ServiceError::Cancelled {
reason: Some("shutdown".to_string()),
}));
}
#[test]
fn test_not_retriable_timeout() {
assert!(!is_retriable_transport_error(&ServiceError::Timeout {
timeout: std::time::Duration::from_secs(30),
}));
}
fn mcp_service_err(code: i32) -> ServiceError {
ServiceError::McpError(rmcp::ErrorData::new(
rmcp::model::ErrorCode(code),
"boom",
None,
))
}
#[test]
fn should_recover_mcp_error_recovers_everything_outside_excluded_set() {
assert!(should_recover_mcp_error(-32603));
assert!(should_recover_mcp_error(-32002));
assert!(should_recover_mcp_error(-32000));
assert!(should_recover_mcp_error(-32099));
assert!(should_recover_mcp_error(-32100));
assert!(should_recover_mcp_error(0));
assert!(should_recover_mcp_error(1));
assert!(should_recover_mcp_error(i32::MIN));
assert!(should_recover_mcp_error(i32::MAX));
}
#[test]
fn should_recover_mcp_error_skips_deterministic_client_errors() {
assert!(!should_recover_mcp_error(-32700));
assert!(!should_recover_mcp_error(-32600));
assert!(!should_recover_mcp_error(-32601));
assert!(!should_recover_mcp_error(-32602));
}
#[test]
fn should_recover_service_error_http_mcperror_recoverable() {
assert!(should_recover_service_error(
&mcp_service_err(-32603),
true,
false,
));
}
#[test]
fn should_recover_service_error_http_mcperror_invalid_params_skipped() {
assert!(!should_recover_service_error(
&mcp_service_err(-32602),
true,
false,
));
}
#[test]
fn should_recover_service_error_stdio_mcperror_not_recovered() {
assert!(!should_recover_service_error(
&mcp_service_err(-32603),
false,
false,
));
}
#[test]
fn should_recover_service_error_mcperror_at_most_once_per_dispatch() {
assert!(!should_recover_service_error(
&mcp_service_err(-32603),
true,
true,
));
}
#[test]
fn should_recover_service_error_http_mcperror_auth_rejection_not_recovered() {
let auth_err = ServiceError::McpError(rmcp::ErrorData::new(
rmcp::model::ErrorCode(-32603),
"Unauthorized: token expired",
None,
));
assert!(!should_recover_service_error(&auth_err, true, false));
let session_err = ServiceError::McpError(rmcp::ErrorData::new(
rmcp::model::ErrorCode(-32603),
"session not found",
None,
));
assert!(should_recover_service_error(&session_err, true, false));
}
#[test]
fn should_recover_service_error_transport_errors_always_recover() {
assert!(should_recover_service_error(
&ServiceError::TransportClosed,
true,
false
));
assert!(should_recover_service_error(
&ServiceError::TransportClosed,
false,
false
));
assert!(should_recover_service_error(
&ServiceError::TransportClosed,
true,
true
));
}
#[test]
fn should_recover_service_error_other_non_transport_not_recovered() {
assert!(!should_recover_service_error(
&ServiceError::UnexpectedResponse,
true,
false
));
assert!(!should_recover_service_error(
&ServiceError::Timeout {
timeout: std::time::Duration::from_secs(30),
},
true,
false
));
}
#[tokio::test]
async fn recover_and_retry_surfaces_original_error_when_recover_fails() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(1),
..Default::default()
};
let client = Arc::new(McpClient::new_http(
"wedged".to_string(),
config,
Some(&overrides),
None,
));
let tool = McpErasedTool {
tool: McpTool::new(
"do_thing".to_string(),
"desc".to_string(),
"wedged".to_string(),
Arc::new(Mutex::new(McpState::new(vec![]))),
serde_json::json!({"type": "object"}),
None,
),
};
let original = mcp_service_err(-32603);
let expected = original.to_string();
let params = CallToolRequestParams::new("do_thing");
let mut reconnect_attempted = false;
let mut is_timeout = false;
let ew = kigi_file_utils::events::EventWriter::noop();
let err = tool
.recover_and_retry(
&client,
params,
std::time::Duration::from_secs(1),
1,
original,
&mut reconnect_attempted,
&mut is_timeout,
&ew,
)
.await
.expect_err("recover must fail against an unreachable host");
assert_eq!(err.to_string(), expected, "original error must be surfaced");
assert!(reconnect_attempted, "reconnect attempt must be flagged");
assert!(!is_timeout, "a recover failure is not a tool timeout");
}
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Clone, Copy)]
enum CallToolBehavior {
ErrorThenOk { code: i32 },
AlwaysError { code: i32 },
HangThenOk { hang_ms: u64 },
ErrorThenHang { code: i32, hang_ms: u64 },
}
#[derive(Clone)]
struct FakeMcpState {
behavior: CallToolBehavior,
inits: Arc<AtomicUsize>,
calls: Arc<AtomicUsize>,
}
async fn fake_handle_post(
axum::extract::State(state): axum::extract::State<FakeMcpState>,
axum::Json(req): axum::Json<serde_json::Value>,
) -> axum::response::Response {
use axum::response::IntoResponse;
let id = req["id"].clone();
let ok = || {
serde_json::json!({
"jsonrpc": "2.0",
"id": id.clone(),
"result": {"content": [{"type": "text", "text": "ok"}], "isError": false},
})
};
let err = |code: i32, msg: String| {
serde_json::json!({
"jsonrpc": "2.0",
"id": id.clone(),
"error": {"code": code, "message": msg},
})
};
match req["method"].as_str() {
Some("initialize") => {
state.inits.fetch_add(1, Ordering::Relaxed);
let result = serde_json::json!({
"jsonrpc": "2.0",
"id": id.clone(),
"result": {
"protocolVersion": req["params"]["protocolVersion"].clone(),
"capabilities": {},
"serverInfo": {"name": "fake", "version": "0.0.0"},
},
});
([("mcp-session-id", "fake-session")], axum::Json(result)).into_response()
}
Some("tools/list") => axum::Json(serde_json::json!({
"jsonrpc": "2.0",
"id": id.clone(),
"result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]},
}))
.into_response(),
Some("tools/call") => {
let n = state.calls.fetch_add(1, Ordering::Relaxed);
match state.behavior {
CallToolBehavior::ErrorThenOk { code } => {
if n == 0 {
axum::Json(err(code, "session expired".to_string())).into_response()
} else {
axum::Json(ok()).into_response()
}
}
CallToolBehavior::AlwaysError { code } => {
axum::Json(err(code, format!("attempt {}", n + 1))).into_response()
}
CallToolBehavior::HangThenOk { hang_ms } => {
if n == 0 {
tokio::time::sleep(std::time::Duration::from_millis(hang_ms)).await;
}
axum::Json(ok()).into_response()
}
CallToolBehavior::ErrorThenHang { code, hang_ms } => {
if n == 0 {
axum::Json(err(code, "session expired".to_string())).into_response()
} else {
tokio::time::sleep(std::time::Duration::from_millis(hang_ms)).await;
axum::Json(ok()).into_response()
}
}
}
}
_ => axum::http::StatusCode::ACCEPTED.into_response(),
}
}
async fn fake_handle_get() -> axum::response::Response {
use axum::response::IntoResponse;
let body = axum::body::Body::from_stream(futures::stream::pending::<
Result<String, std::io::Error>,
>());
(
[(axum::http::header::CONTENT_TYPE, "text/event-stream")],
body,
)
.into_response()
}
async fn spawn_fake_mcp(
behavior: CallToolBehavior,
) -> (String, Arc<AtomicUsize>, Arc<AtomicUsize>) {
let inits = Arc::new(AtomicUsize::new(0));
let calls = Arc::new(AtomicUsize::new(0));
let app = axum::Router::new()
.route(
"/mcp",
axum::routing::get(fake_handle_get).post(fake_handle_post),
)
.with_state(FakeMcpState {
behavior,
inits: inits.clone(),
calls: calls.clone(),
});
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = axum::serve(listener, app).await;
});
(format!("http://{addr}/mcp"), inits, calls)
}
fn fake_http_client(url: &str, tool_timeout_sec: u64) -> Arc<McpClient> {
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(5),
tool_timeout_sec: Some(tool_timeout_sec),
..Default::default()
};
Arc::new(McpClient::new_http(
"fake".to_string(),
HttpConfig {
url: url.to_string(),
headers: vec![],
},
Some(&overrides),
None,
))
}
fn fake_echo_tool() -> McpErasedTool {
McpErasedTool {
tool: McpTool::new(
"echo".to_string(),
"echo desc".to_string(),
"fake".to_string(),
Arc::new(Mutex::new(McpState::new(vec![]))),
serde_json::json!({"type": "object"}),
None,
),
}
}
fn event_types(jsonl: &str) -> Vec<serde_json::Value> {
jsonl
.lines()
.filter_map(|l| serde_json::from_str::<serde_json::Value>(l).ok())
.collect()
}
#[tokio::test(flavor = "multi_thread")]
async fn try_call_tool_http_mcperror_recovers_then_retry_succeeds() {
let (url, inits, calls) =
spawn_fake_mcp(CallToolBehavior::ErrorThenOk { code: -32603 }).await;
let client = fake_http_client(&url, 5);
let tool = fake_echo_tool();
let tmp = tempfile::tempdir().unwrap();
let ew = kigi_file_utils::events::EventWriter::open(tmp.path());
let mut reconnect = false;
let mut is_timeout = false;
let raw = serde_json::json!({});
let out = tool
.try_call_tool(&client, &raw, &mut reconnect, &mut is_timeout, &ew)
.await
.expect("recovered call should succeed");
assert!(
!out.is_error.unwrap_or(false),
"retry should return a success result"
);
assert!(reconnect, "reconnect_attempted must be set");
assert!(!is_timeout);
assert_eq!(
calls.load(Ordering::Relaxed),
2,
"one failed + one retried tools/call"
);
assert_eq!(
inits.load(Ordering::Relaxed),
2,
"initial handshake + one recovery re-init"
);
let jsonl = std::fs::read_to_string(tmp.path().join("events.jsonl")).unwrap();
let events = event_types(&jsonl);
assert!(
events.iter().any(|e| e["type"] == "mcp_transport_error"),
"expected mcp_transport_error in {jsonl}"
);
assert!(
events
.iter()
.any(|e| e["type"] == "mcp_transport_reconnect" && e["success"] == true),
"expected a successful mcp_transport_reconnect in {jsonl}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn try_call_tool_http_retry_failure_surfaces_retry_error() {
let (url, _inits, calls) =
spawn_fake_mcp(CallToolBehavior::AlwaysError { code: -32603 }).await;
let client = fake_http_client(&url, 5);
let tool = fake_echo_tool();
let ew = kigi_file_utils::events::EventWriter::noop();
let mut reconnect = false;
let mut is_timeout = false;
let raw = serde_json::json!({});
let err = tool
.try_call_tool(&client, &raw, &mut reconnect, &mut is_timeout, &ew)
.await
.expect_err("both attempts fail");
let msg = err.to_string();
assert!(msg.contains("attempt 2"), "want retry error, got: {msg}");
assert!(
!msg.contains("attempt 1"),
"must not surface the original error: {msg}"
);
assert!(reconnect);
assert!(!is_timeout);
assert_eq!(
calls.load(Ordering::Relaxed),
2,
"one failed + one retried tools/call"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn try_call_tool_http_invalid_params_not_recovered() {
let (url, inits, calls) =
spawn_fake_mcp(CallToolBehavior::AlwaysError { code: -32602 }).await;
let client = fake_http_client(&url, 5);
let tool = fake_echo_tool();
let ew = kigi_file_utils::events::EventWriter::noop();
let mut reconnect = false;
let mut is_timeout = false;
let raw = serde_json::json!({});
let err = tool
.try_call_tool(&client, &raw, &mut reconnect, &mut is_timeout, &ew)
.await
.expect_err("invalid params surfaced as-is");
assert!(err.to_string().contains("attempt 1"), "got: {err}");
assert!(!reconnect, "invalid-params must not trigger recovery");
assert!(!is_timeout);
assert_eq!(calls.load(Ordering::Relaxed), 1, "no retry POST");
assert_eq!(inits.load(Ordering::Relaxed), 1, "no recovery re-init");
}
#[tokio::test(flavor = "multi_thread")]
async fn try_call_tool_http_outer_timeout_resets_transport_no_retry() {
let (url, inits, calls) =
spawn_fake_mcp(CallToolBehavior::HangThenOk { hang_ms: 3000 }).await;
let client = fake_http_client(&url, 1);
let tool = fake_echo_tool();
let ew = kigi_file_utils::events::EventWriter::noop();
let mut reconnect = false;
let mut is_timeout = false;
let raw = serde_json::json!({});
let err = tool
.try_call_tool(&client, &raw, &mut reconnect, &mut is_timeout, &ew)
.await
.expect_err("call must time out");
assert!(err.to_string().contains("timed out"), "got: {err}");
assert!(is_timeout, "is_timeout must be set");
assert!(reconnect, "timeout arm flags the reconnect after resetting");
assert_eq!(
calls.load(Ordering::Relaxed),
1,
"timed-out call is NOT retried"
);
assert!(matches!(
client.state_kind().await,
ClientStateKind::Pending
));
assert_eq!(
inits.load(Ordering::Relaxed),
1,
"no re-init during the timed-out dispatch"
);
let mut reconnect2 = false;
let mut is_timeout2 = false;
let out = tool
.try_call_tool(&client, &raw, &mut reconnect2, &mut is_timeout2, &ew)
.await
.expect("second dispatch should re-init and succeed");
assert!(!out.is_error.unwrap_or(false));
assert!(!is_timeout2);
assert_eq!(
inits.load(Ordering::Relaxed),
2,
"second dispatch re-initialized the session"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn try_call_tool_http_retry_timeout_surfaces_timeout() {
let (url, inits, calls) = spawn_fake_mcp(CallToolBehavior::ErrorThenHang {
code: -32603,
hang_ms: 3000,
})
.await;
let client = fake_http_client(&url, 1);
let tool = fake_echo_tool();
let ew = kigi_file_utils::events::EventWriter::noop();
let mut reconnect = false;
let mut is_timeout = false;
let raw = serde_json::json!({});
let err = tool
.try_call_tool(&client, &raw, &mut reconnect, &mut is_timeout, &ew)
.await
.expect_err("the retried call must time out");
assert!(err.to_string().contains("timed out"), "got: {err}");
assert!(is_timeout, "retry-timeout must set is_timeout");
assert!(reconnect, "recovery was attempted");
assert_eq!(
calls.load(Ordering::Relaxed),
2,
"the retry tools/call was attempted"
);
assert_eq!(
inits.load(Ordering::Relaxed),
2,
"recovery re-initialized before the retry"
);
}
// ── new_http stores http_config tests ────────────────────────────
#[test]
fn test_new_http_stores_http_config() {
let config = HttpConfig {
url: "http://localhost:5000/api/mcp".to_string(),
headers: vec![("x-token".to_string(), "abc".to_string())],
};
let client = McpClient::new_http("example-mcp".to_string(), config, None, None);
let stored = client
.http_config
.as_ref()
.expect("http_config should be Some");
assert_eq!(stored.url, "http://localhost:5000/api/mcp");
assert_eq!(stored.headers.len(), 1);
assert_eq!(stored.headers[0].0, "x-token");
}
#[test]
fn test_new_stdio_has_no_http_config() {
// Stdio clients must NOT have http_config — they can't reconnect via HTTP.
let client = McpClient::stub("stdio-srv");
assert!(client.http_config.is_none());
}
// ── http_headers_match / refresh_managed_clients guard tests ─────
#[test]
fn http_headers_match_compares_full_set_order_insensitively() {
let config = HttpConfig {
url: "http://localhost:5000/api/mcp".to_string(),
headers: vec![
("authorization".to_string(), "Bearer t".to_string()),
("x-scope".to_string(), "read".to_string()),
],
};
let client = McpClient::new_http("managed".to_string(), config, None, None);
let equal: HashMap<String, String> = [
("x-scope".to_string(), "read".to_string()),
("authorization".to_string(), "Bearer t".to_string()),
]
.into_iter()
.collect();
assert!(client.http_headers_match(&equal));
let changed_value: HashMap<String, String> = [
("authorization".to_string(), "Bearer NEW".to_string()),
("x-scope".to_string(), "read".to_string()),
]
.into_iter()
.collect();
assert!(!client.http_headers_match(&changed_value));
let missing_key: HashMap<String, String> =
[("authorization".to_string(), "Bearer t".to_string())]
.into_iter()
.collect();
assert!(!client.http_headers_match(&missing_key));
}
#[test]
fn http_headers_match_handles_duplicate_stored_keys() {
// Duplicate stored key must not mask a missing fresh key by inflating
// the stored length to match.
let config = HttpConfig {
url: "http://localhost:5000/api/mcp".to_string(),
headers: vec![
("authorization".to_string(), "Bearer t".to_string()),
("authorization".to_string(), "Bearer t".to_string()),
],
};
let client = McpClient::new_http("managed".to_string(), config, None, None);
let two_distinct: HashMap<String, String> = [
("authorization".to_string(), "Bearer t".to_string()),
("x-scope".to_string(), "read".to_string()),
]
.into_iter()
.collect();
assert!(!client.http_headers_match(&two_distinct));
let single: HashMap<String, String> =
[("authorization".to_string(), "Bearer t".to_string())]
.into_iter()
.collect();
assert!(client.http_headers_match(&single));
}
#[test]
fn http_headers_match_false_for_non_http_client() {
let client = McpClient::stub("stdio-srv");
let headers: HashMap<String, String> =
[("authorization".to_string(), "Bearer t".to_string())]
.into_iter()
.collect();
assert!(!client.http_headers_match(&headers));
}
#[test]
fn refresh_managed_clients_keeps_arc_when_headers_unchanged() {
let url = "http://localhost:5000/api/mcp";
let mut state = McpState::new(vec![make_http_server("managed", url)]);
let config = HttpConfig {
url: url.to_string(),
headers: vec![("authorization".to_string(), "Bearer t".to_string())],
};
state.owned_clients.insert(
"managed".to_string(),
Arc::new(McpClient::new_http(
"managed".to_string(),
config,
None,
None,
)),
);
let before = Arc::clone(state.owned_clients.get("managed").unwrap());
let fresh: HashMap<String, String> =
[("authorization".to_string(), "Bearer t".to_string())]
.into_iter()
.collect();
state.refresh_managed_clients(std::iter::once((url, &fresh)));
let after = state.owned_clients.get("managed").unwrap();
assert!(
Arc::ptr_eq(&before, after),
"unchanged headers must not rebuild the client"
);
}
#[test]
fn refresh_managed_clients_installs_new_arc_when_headers_differ() {
let url = "http://localhost:5000/api/mcp";
let mut state = McpState::new(vec![make_http_server("managed", url)]);
let config = HttpConfig {
url: url.to_string(),
headers: vec![("authorization".to_string(), "Bearer old".to_string())],
};
state.owned_clients.insert(
"managed".to_string(),
Arc::new(McpClient::new_http(
"managed".to_string(),
config,
None,
None,
)),
);
let before = Arc::clone(state.owned_clients.get("managed").unwrap());
let fresh: HashMap<String, String> =
[("authorization".to_string(), "Bearer new".to_string())]
.into_iter()
.collect();
state.refresh_managed_clients(std::iter::once((url, &fresh)));
let after = state.owned_clients.get("managed").unwrap();
assert!(
!Arc::ptr_eq(&before, after),
"changed headers must install a fresh client"
);
assert!(after.http_headers_match(&fresh));
}
// ── reset_transport tests ────────────────────────────────────────
#[tokio::test]
async fn test_reset_transport_succeeds_for_http_client() {
let config = HttpConfig {
url: "http://127.0.0.1:9/api/mcp".to_string(),
headers: vec![],
};
let client = McpClient::new_http("example-mcp".to_string(), config, None, None);
assert!(client.reset_transport().await);
}
#[tokio::test]
async fn test_reset_transport_fails_for_stub() {
// Stub has `reconnect = None`, simulating a Stdio client.
let client = McpClient::stub("stdio-srv");
assert!(!client.reset_transport().await);
}
#[tokio::test]
async fn test_reset_transport_is_idempotent() {
let config = HttpConfig {
url: "http://127.0.0.1:9/api/mcp".to_string(),
headers: vec![],
};
let client = McpClient::new_http("example-mcp".to_string(), config, None, None);
// Multiple resets should all succeed.
assert!(client.reset_transport().await);
assert!(client.reset_transport().await);
assert!(client.reset_transport().await);
}
#[tokio::test]
async fn test_reset_transport_makes_ensure_initialized_retry_handshake() {
// Port 1 on loopback refuses immediately (ECONNREFUSED -> HandshakeFailed),
// so each handshake fails fast instead of waiting out the connect timeout.
let config = HttpConfig {
url: "http://127.0.0.1:1/unreachable".to_string(),
headers: vec![],
};
let client = McpClient::new_http("test".to_string(), config, None, None);
// First ensure_initialized will fail (unreachable server) but proves
// the client attempts a handshake from the Pending state.
let err1 = client.ensure_initialized().await.unwrap_err();
assert!(
matches!(
err1,
McpError::Timeout { .. } | McpError::HandshakeFailed { .. }
),
"first init should fail: {err1}"
);
// Reset puts the client back into Pending with a fresh transport.
assert!(client.reset_transport().await);
// Second ensure_initialized should attempt another handshake (not
// return a cached error). It will fail again with the same kind of
// error, proving the reset restored the transport.
let err2 = client.ensure_initialized().await.unwrap_err();
assert!(
matches!(
err2,
McpError::Timeout { .. } | McpError::HandshakeFailed { .. }
),
"second init after reset should also attempt handshake: {err2}"
);
}
#[tokio::test]
async fn recover_errors_for_client_with_no_restorable_transport() {
// A stub has `reconnect = None` (like Stdio): `recover` can't rebuild it.
let err = Arc::new(McpClient::stub("stdio"))
.recover()
.await
.unwrap_err();
assert!(matches!(err, McpError::ClientError(_)), "got {err}");
}
#[tokio::test]
async fn reset_transport_rebuilds_acp_client() {
use crate::acp_transport::AcpReverseInvoker;
use std::time::Duration;
struct NoopInvoker;
#[async_trait::async_trait]
impl AcpReverseInvoker for NoopInvoker {
async fn invoke(
&self,
_server_id: &str,
_message: serde_json::Value,
_timeout: Duration,
) -> Result<serde_json::Value, String> {
Ok(serde_json::Value::Null)
}
}
let client = McpClient::new_acp(
"sdk-tools".to_string(),
"srv_0".to_string(),
Arc::new(NoopInvoker),
None,
None,
);
// ACP clients restore from `reconnect`, unlike Stdio.
assert!(client.reset_transport().await);
assert!(
matches!(
&*client.state.lock().await,
ClientState::Pending(PendingTransport::Acp { .. })
),
"reset_transport should restore the ACP transport to Pending"
);
}
/// End-to-end reconnect-THEN-SUCCEED for the `try_call_tool` retry arm: the one
/// piece otherwise covered only by its parts (`is_retriable_transport_error`,
/// `reset_transport_*`, `ensure_initialized_*`).
///
/// Drives the REAL `McpErasedTool::try_call_tool` against a real
/// `McpClient`. The first `call_tool` hits a real `RunningService`
/// whose transport is already closed, so it returns a genuine,
/// retriable `ServiceError::TransportClosed`; the arm must then flag
/// `reconnect_attempted`, run the real `reset_transport` +
/// `ensure_initialized` re-handshake (rebuilding the ACP transport
/// against a working echo server), and return the SECOND attempt's
/// `Ok` result.
///
/// Why a separately-built dead service instead of failing the initial
/// connection: the ACP bridge transport can only be torn down from the
/// rmcp side, so a fresh real service is built over a raw duplex whose
/// server answers `initialize` then drops — closing the transport so
/// the first `call_tool` observes `TransportClosed`. Everything from
/// the retriable-error gate through the successful retry is real code.
#[tokio::test]
async fn try_call_tool_reconnects_then_succeeds_after_retriable_transport_error() {
use crate::acp_transport::AcpReverseInvoker;
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
// Working in-process echo server for the post-reconnect retry.
struct EchoSdkServer;
#[async_trait::async_trait]
impl AcpReverseInvoker for EchoSdkServer {
async fn invoke(
&self,
_server_id: &str,
message: serde_json::Value,
_timeout: Duration,
) -> Result<serde_json::Value, String> {
let id = message
.get("id")
.cloned()
.unwrap_or(serde_json::Value::Null);
let method = message
.get("method")
.and_then(|m| m.as_str())
.unwrap_or_default();
let result = match method {
"initialize" => serde_json::json!({
"protocolVersion": message["params"]["protocolVersion"],
"capabilities": { "tools": {} },
"serverInfo": { "name": "echo", "version": "0.0.0" },
}),
"tools/call" => serde_json::json!({
"content": [{
"type": "text",
"text": message["params"]["arguments"]["text"]
.as_str()
.unwrap_or_default(),
}],
"isError": false,
}),
other => return Err(format!("unexpected method {other}")),
};
Ok(serde_json::json!({ "jsonrpc": "2.0", "id": id, "result": result }))
}
}
// A real `RunningService` whose transport is already closed: the
// server answers `initialize`, consumes the `initialized`
// notification (so the client's handshake send succeeds), then drops
// its duplex ends. The next `call_tool` therefore observes a real
// `ServiceError::TransportClosed`.
async fn dead_service() -> McpService {
let (client_read, server_write) = tokio::io::duplex(64 * 1024); // server -> client
let (server_read, client_write) = tokio::io::duplex(64 * 1024); // client -> server
tokio::spawn(async move {
let mut reader = BufReader::new(server_read);
let mut writer = server_write;
let mut line = String::new();
loop {
line.clear();
if reader.read_line(&mut line).await.unwrap_or(0) == 0 {
return;
}
let Ok(msg) = serde_json::from_str::<serde_json::Value>(line.trim()) else {
continue;
};
if msg.get("method").and_then(|m| m.as_str()) == Some("initialize") {
let id = msg.get("id").cloned().unwrap_or(serde_json::Value::Null);
let resp = serde_json::json!({ "jsonrpc": "2.0", "id": id, "result": {
"protocolVersion": msg["params"]["protocolVersion"],
"capabilities": { "tools": {} },
"serverInfo": { "name": "dead", "version": "0.0.0" },
}});
let mut encoded = serde_json::to_string(&resp).unwrap();
encoded.push('\n');
let _ = writer.write_all(encoded.as_bytes()).await;
let _ = writer.flush().await;
// Drain the `initialized` notification, then drop to close.
let _ = reader.read_line(&mut line).await;
return;
}
}
});
let handler = KigiClientHandler {
info: McpClient::make_client_info("dead"),
server_name: "dead".to_string(),
notify_tx: Arc::new(parking_lot::Mutex::new(None)),
};
let transport = rmcp::transport::async_rw::AsyncRwTransport::<RoleClient, _, _>::new(
client_read,
client_write,
);
Arc::new(
handler
.serve(transport)
.await
.expect("dead-service handshake"),
)
}
// ACP client whose `reconnect` snapshot rebuilds against the echo server.
let client = Arc::new(McpClient::new_acp(
"sdk".to_string(),
"srv_0".to_string(),
Arc::new(EchoSdkServer),
None,
None,
));
// Inject the closed real service so the FIRST `call_tool` fails retriably.
let dead = dead_service().await;
*client.state.lock().await = ClientState::Ready(dead);
let erased = McpErasedTool {
tool: McpTool::new(
"echo".to_string(),
"echo".to_string(),
"sdk".to_string(),
Arc::new(Mutex::new(McpState::new(vec![]))),
serde_json::json!({}),
None,
),
};
let raw = serde_json::json!({ "text": "after reconnect" });
let mut reconnect_attempted = false;
let mut is_timeout = false;
let ew = kigi_file_utils::events::EventWriter::noop();
let result = erased
.try_call_tool(
&client,
&raw,
&mut reconnect_attempted,
&mut is_timeout,
&ew,
)
.await
.expect("retry after reconnect should succeed");
// The Ok came from the SECOND attempt — the dead service cannot echo,
// so this text proves the rebuilt transport served the retry.
assert_eq!(
result.content[0].as_text().expect("text content").text,
"after reconnect"
);
assert!(
reconnect_attempted,
"retriable transport error must set reconnect_attempted"
);
assert!(
!is_timeout,
"successful retry must not be flagged as timeout"
);
// reset_transport + re-handshake replaced the dead service with a live one.
assert!(matches!(&*client.state.lock().await, ClientState::Ready(_)));
}
#[test]
fn is_auth_rejection_message_matches_auth_signals() {
// The verbatim string captured in production for a managed handshake.
assert!(is_auth_rejection_message(
"MCP server 'notion' handshake failed: Auth required, when send initialize request"
));
assert!(is_auth_rejection_message("401 Unauthorized"));
assert!(is_auth_rejection_message("unauthorized"));
assert!(is_auth_rejection_message("Authentication required"));
assert!(is_auth_rejection_message("authentication failed"));
assert!(is_auth_rejection_message("status: 401"));
assert!(is_auth_rejection_message("HTTP status 401"));
assert!(is_auth_rejection_message("server returned status code 401"));
assert!(is_auth_rejection_message("HTTP 401"));
assert!(is_auth_rejection_message("error 401"));
// rmcp worker fatal context uses Debug form without spaces.
assert!(is_auth_rejection_message(
"worker quit with fatal: Transport channel closed, when Auth(AuthorizationRequired)"
));
let auth_req = McpError::AuthRequired {
server: "clickhouse".into(),
};
assert!(auth_req.is_auth_rejection());
assert_eq!(auth_req.server_name(), Some("clickhouse"));
}
#[test]
fn auth_required_records_as_auth_not_init_failed_and_maps_category() {
// Pre-spawn gate is owned by the auth state machine: it lands in
// `auth_required` (recoverable via re-auth) and never `init_failed`.
let mut state = McpState::new(vec![]);
state.record_init_failure("oauth-srv", true, None);
assert!(state.auth_required.contains("oauth-srv"));
assert!(!state.init_failed.contains_key("oauth-srv"));
// AuthRequired carries the AuthRequired telemetry category, not ClientError.
let err = McpError::AuthRequired {
server: "oauth-srv".into(),
};
assert!(matches!(
err.error_category(),
kigi_file_utils::events::McpErrorCategory::AuthRequired
));
}
#[test]
fn is_auth_rejection_message_rejects_non_auth() {
// Transport / timeout / spawn wording is never an auth rejection.
assert!(!is_auth_rejection_message("Transport closed"));
assert!(!is_auth_rejection_message(
"MCP server 'x' timed out after 30s"
));
assert!(!is_auth_rejection_message(
"Failed to spawn MCP server 'x': No such file or directory"
));
// 403/forbidden is a non-auth policy denial in this stack, not auth.
assert!(!is_auth_rejection_message("403 Forbidden"));
assert!(!is_auth_rejection_message("forbidden"));
// Incidental digits must not trip the status-anchored 401 patterns.
assert!(!is_auth_rejection_message("request took 401ms"));
assert!(!is_auth_rejection_message("connect 10.0.4.01:443"));
assert!(!is_auth_rejection_message("read 401 bytes"));
// A status literal followed by another alphanumeric is a different
// token: a longer number (4012) or an adjacent unit (401ms).
assert!(!is_auth_rejection_message("http 4012"));
assert!(!is_auth_rejection_message("error 4012"));
assert!(!is_auth_rejection_message("status: 4012"));
assert!(!is_auth_rejection_message("http 401ms"));
assert!(!is_auth_rejection_message("error 401ms"));
// ...but a trailing punctuation/whitespace still matches.
assert!(is_auth_rejection_message("http 401."));
assert!(is_auth_rejection_message("error 401: token expired"));
}
#[test]
fn mcp_error_is_auth_rejection_delegates() {
assert!(McpError::ClientError("Auth required".to_string()).is_auth_rejection());
assert!(!McpError::ClientError("Transport closed".to_string()).is_auth_rejection());
assert!(
!McpError::Timeout {
server: "x".to_string(),
timeout_secs: 30,
}
.is_auth_rejection()
);
assert!(
!McpError::SpawnFailed {
server: "x".to_string(),
source: std::io::Error::new(std::io::ErrorKind::NotFound, "401 Unauthorized"),
}
.is_auth_rejection()
);
// HandshakeFailed is the production carrier: its `source` Display must
// surface the auth substring for the delegation to fire.
assert!(
McpError::HandshakeFailed {
server: "x".to_string(),
source: Box::new(ClientInitializeError::ConnectionClosed(
"Auth required, when send initialize request".to_string()
)),
}
.is_auth_rejection()
);
assert!(
!McpError::HandshakeFailed {
server: "x".to_string(),
source: Box::new(ClientInitializeError::ConnectionClosed(
"transport closed".to_string()
)),
}
.is_auth_rejection()
);
}
#[test]
fn format_mcp_image_default_emits_only_data_uri() {
let out = format_mcp_image("image/png", "AAAA", false);
assert_eq!(out, "data:image/png;base64,AAAA");
assert!(!out.contains("<mcp_image_base64"));
}
#[test]
fn format_mcp_image_expose_emits_data_uri_and_raw_block() {
let out = format_mcp_image("image/png", "AAAA", true);
assert!(out.contains("data:image/png;base64,AAAA"));
assert!(out.contains("<mcp_image_base64 mime=\"image/png\">\nAAAA\n</mcp_image_base64>"));
}
/// Wrapper must not re-match the extractor regex, else the raw copy gets stripped too.
#[test]
fn format_mcp_image_expose_raw_block_has_no_data_prefix() {
let out = format_mcp_image("image/jpeg", "ZZZZ", true);
assert_eq!(out.matches("data:image/").count(), 1);
}
#[test]
fn load_expose_image_base64_defaults_to_false() {
assert!(!McpClient::load_expose_image_base64(None, None));
}
#[test]
fn load_expose_image_base64_uses_overrides_when_meta_unset() {
let overrides = McpClientTimeoutOverrides {
expose_image_base64: Some(true),
..Default::default()
};
assert!(McpClient::load_expose_image_base64(Some(&overrides), None));
}
#[test]
fn load_expose_image_base64_meta_wins_over_overrides() {
let overrides = McpClientTimeoutOverrides {
expose_image_base64: Some(true),
..Default::default()
};
let meta = McpServerMetaConfig {
expose_image_base64: Some(false),
..Default::default()
};
assert!(!McpClient::load_expose_image_base64(
Some(&overrides),
Some(&meta)
));
}
#[test]
fn load_expose_image_base64_meta_falls_through_when_none() {
let overrides = McpClientTimeoutOverrides {
expose_image_base64: Some(true),
..Default::default()
};
let meta = McpServerMetaConfig::default(); // expose_image_base64 = None
assert!(McpClient::load_expose_image_base64(
Some(&overrides),
Some(&meta)
));
}
/// End-to-end: override → constructor → public getter.
/// New constructors should add a similar assertion.
#[test]
fn new_http_propagates_expose_image_base64_override_to_getter() {
let config = HttpConfig {
url: "http://localhost/api/mcp".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
expose_image_base64: Some(true),
..Default::default()
};
let client = McpClient::new_http(
"grafana".to_string(),
config.clone(),
Some(&overrides),
None,
);
assert!(client.expose_image_base64());
let client_default = McpClient::new_http("grafana".to_string(), config, None, None);
assert!(!client_default.expose_image_base64());
}
// ------------------------------------------------------------------
// ensure_initialized single-flight + Notify behavior (regression
// suite for the "MCP client already initializing" doom-loop).
// ------------------------------------------------------------------
/// `ensure_initialized` on a stub (no transport) must surface a
/// clear, actionable configuration error — never the legacy
/// "already initializing" sentinel which leaked into model-visible
/// tool results and triggered retry loops that exhausted the
/// per-tick prompt budget.
#[tokio::test]
async fn ensure_initialized_on_empty_client_returns_no_transport_error() {
let client = McpClient::stub("test-server");
let err = client.ensure_initialized().await.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("no transport configured"),
"expected clear 'no transport configured' error, got: {msg}"
);
assert!(
!msg.contains("already initializing"),
"regression: legacy fast-fail sentinel surfaced: {msg}"
);
}
/// Drive `N` `ensure_initialized` calls concurrently against an
/// unreachable HTTP server with a tight startup timeout. Every
/// caller must surface a real handshake error (`Timeout` or
/// `HandshakeFailed`); none may surface the legacy
/// "MCP client already initializing" sentinel which the
/// pre-fix branch emitted whenever a caller observed
/// `Pending(None)` while another caller was running the handshake.
///
/// The race window is intentionally widened by using an unreachable
/// host (`192.0.2.1:1` — TEST-NET-1, guaranteed unrouteable) so the
/// handshake stalls for `startup_timeout_sec` and every concurrent
/// caller spawned after the first observes `Initializing` instead
/// of `Pending`.
#[tokio::test]
async fn ensure_initialized_concurrent_callers_never_see_legacy_fast_fail() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(1),
..Default::default()
};
let client = Arc::new(McpClient::new_http(
"test-server".to_string(),
config,
Some(&overrides),
None,
));
let mut handles = Vec::new();
for _ in 0..5 {
let c = Arc::clone(&client);
handles.push(tokio::spawn(async move { c.ensure_initialized().await }));
}
for (idx, handle) in handles.into_iter().enumerate() {
let result = handle.await.expect("task did not panic");
let err = result.expect_err("unreachable host must fail");
let msg = err.to_string();
assert!(
!msg.contains("MCP client already initializing"),
"caller {idx}: legacy fast-fail sentinel surfaced: {msg}"
);
assert!(
matches!(
err,
McpError::Timeout { .. } | McpError::HandshakeFailed { .. }
),
"caller {idx}: expected handshake failure, got: {err}"
);
}
}
/// A caller that finds `ClientState::Initializing` must park on
/// `init_done` and wake up when the holder publishes a new state,
/// then take the freshly-restored transport for its own retry.
///
/// We exercise the wake path directly (without an actual concurrent
/// handshake) by manually transitioning state to `Initializing`,
/// spawning a parker, then transitioning back to `Pending` and
/// firing `notify_waiters`. The parker should retry against the
/// restored (still-unreachable) transport and surface a normal
/// handshake error rather than the wait-timeout error.
#[tokio::test]
async fn ensure_initialized_parked_caller_retries_after_notify() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(1),
..Default::default()
};
let client = Arc::new(McpClient::new_http(
"test-server".to_string(),
config.clone(),
Some(&overrides),
None,
));
// Simulate an in-flight handshake by another task: pretend
// that task took the transport and entered Initializing.
*client.state.lock().await = ClientState::Initializing;
// Spawn the parker. It must observe Initializing and park on
// `init_done` rather than fail-fast.
let parker_client = Arc::clone(&client);
let parker = tokio::spawn(async move { parker_client.ensure_initialized().await });
// Give the parker a chance to reach the await on `init_done`.
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
// Publish a fresh Pending transport and notify — simulates the
// holder's failure-path restore.
*client.state.lock().await = ClientState::Pending(PendingTransport::Http(config.clone()));
client.init_done.notify_waiters();
// The parker should wake, take the transport, run its own
// handshake (which fails against the unreachable host), and
// surface a regular handshake error — never the wait-timeout
// error and never the legacy fast-fail.
let err = parker
.await
.expect("parker did not panic")
.expect_err("unreachable host must fail");
let msg = err.to_string();
assert!(
!msg.contains("MCP client already initializing"),
"regression: legacy fast-fail sentinel: {msg}"
);
assert!(
!msg.contains("init still in progress"),
"parker should not hit wait-timeout when notified: {msg}"
);
assert!(
matches!(
err,
McpError::Timeout { .. } | McpError::HandshakeFailed { .. }
),
"expected handshake failure, got: {err}"
);
}
/// If a caller is parked on `Initializing` and the holder is
/// dropped without notifying (cancellation-storm edge case), the
/// parker must eventually surface a clear `init still in progress`
/// timeout error rather than block indefinitely.
///
/// Without the inflight-wait timeout, a wedged client (one whose
/// drop guard couldn't acquire the lock to restore) would silently
/// stall every future `ensure_initialized` caller until process
/// restart. The 1 s margin past `startup_timeout_sec` keeps the
/// happy path snappy while still bounding the worst case.
#[tokio::test]
async fn ensure_initialized_inflight_wait_times_out_when_holder_silent() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(0),
..Default::default()
};
let client = McpClient::new_http("test-server".to_string(), config, Some(&overrides), None);
// Wedge the slot in Initializing with no live holder.
*client.state.lock().await = ClientState::Initializing;
let err = client.ensure_initialized().await.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("init still in progress"),
"expected wait-timeout error, got: {msg}"
);
assert!(
!msg.contains("already initializing"),
"regression: legacy fast-fail sentinel: {msg}"
);
}
/// When the holder task is cancelled (`abort()`) mid-handshake, the
/// `InitGuard` drop impl restores `Pending(transport)` on a
/// best-effort basis so a follow-on caller can retry without
/// requiring an explicit `reset_transport`.
#[tokio::test]
async fn ensure_initialized_drop_guard_restores_state_after_holder_aborted() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let overrides = McpClientTimeoutOverrides {
// Long enough that the holder is guaranteed to still be
// inside try_handshake when we abort it.
startup_timeout_sec: Some(10),
..Default::default()
};
let client = Arc::new(McpClient::new_http(
"test-server".to_string(),
config,
Some(&overrides),
None,
));
let holder_client = Arc::clone(&client);
let holder = tokio::spawn(async move { holder_client.ensure_initialized().await });
// Wait for the holder to enter Initializing.
let started = std::time::Instant::now();
loop {
if matches!(&*client.state.lock().await, ClientState::Initializing) {
break;
}
assert!(
started.elapsed() < std::time::Duration::from_secs(2),
"holder never reached Initializing"
);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
// Cancel the holder mid-handshake. The drop guard should
// restore Pending so the next caller can retry.
holder.abort();
let _ = holder.await;
// The drop guard restores best-effort via `try_lock` and notifies.
// Wait briefly for it to settle.
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
match &*client.state.lock().await {
ClientState::Pending(_) => {} // expected
other => panic!(
"expected Pending after holder abort + drop guard, found {}",
state_label(other)
),
}
}
/// `McpState::is_initialized()` MUST require both the early
/// `finish_init` flag AND an empty `initializing_servers` set.
///
/// The session actor's `start_mcp_servers` path calls `finish_init`
/// **early** (right after spawning processes, before any handshake
/// completes) so non-MCP work can proceed in parallel. Tool dispatch
/// and the Blocking-strategy prompt guard, however, must NOT
/// observe "initialized" until every per-server handshake is done —
/// otherwise the model's first tool call races the background
/// `get_tool_registrations` handshake and the
/// `McpClient::ensure_initialized` window described above triggers.
#[test]
fn test_mcp_state_is_initialized_requires_empty_initializing_servers() {
let mut state = McpState::new(vec![make_stdio_server("a", "/bin/a")]);
// NotStarted: neither flag set, no per-server work.
assert!(!state.is_initialized());
assert!(!state.is_initializing());
assert!(!state.has_finished_init());
assert!(matches!(state.init_progress(), InitProgress::NotStarted));
// Starting: try_start_init fired, per-server names registered,
// finish_init has NOT yet fired. is_initializing() is true.
assert!(state.try_start_init());
state.mark_servers_initializing(["a".to_string()]);
assert!(!state.is_initialized());
assert!(state.is_initializing());
assert!(!state.has_finished_init());
assert!(matches!(
state.init_progress(),
InitProgress::Starting { .. }
));
// Finished + handshakes outstanding: actor called finish_init
// early but the per-server background handshake is still in
// flight. is_initialized() must be FALSE during this window.
state.finish_init();
assert!(
!state.is_initialized(),
"is_initialized() must wait for per-server handshakes"
);
assert!(
state.is_initializing(),
"is_initializing() must report in-flight per-server work"
);
assert!(state.has_finished_init());
assert!(state.is_server_handshaking("a"));
assert_eq!(state.handshaking_servers_count(), 1);
// Finished + empty: background task has reported the handshake
// complete. Now and only now is the pool fully initialized.
state.mark_server_ready("a");
assert!(state.is_initialized());
assert!(!state.is_initializing());
assert!(state.has_finished_init());
assert!(!state.is_server_handshaking("a"));
assert_eq!(state.handshaking_servers_count(), 0);
}
/// Locks in the typed-state contract: the `init_progress` field
/// makes nonsensical combinations like "initialized AND
/// initializing" structurally unrepresentable. Every legal state
/// has exactly one [`InitProgress`] variant; every transition is
/// driven through the typed methods.
#[test]
fn test_init_progress_state_machine_invariants() {
let mut state = McpState::new(vec![make_stdio_server("a", "/bin/a")]);
// Invariant: try_start_init is one-shot per cycle.
assert!(state.try_start_init());
assert!(!state.try_start_init(), "double try_start_init is rejected");
// Invariant: mark_all_servers_ready clears handshaking in
// both Starting and Finished states; never resurrects them.
state.mark_servers_initializing(["a".to_string(), "b".to_string()]);
assert_eq!(state.handshaking_servers_count(), 2);
state.mark_all_servers_ready();
assert_eq!(state.handshaking_servers_count(), 0);
assert!(
matches!(state.init_progress(), InitProgress::Starting { .. }),
"mark_all_servers_ready preserves the lifecycle variant"
);
// Invariant: finish_init from Starting → Finished preserves
// (or in this case, the now-empty) handshaking set.
state.finish_init();
assert!(state.is_initialized());
assert!(matches!(
state.init_progress(),
InitProgress::Finished { .. }
));
// Invariant: cancel_init returns us cleanly to NotStarted,
// ready for a new try_start_init.
state.cancel_init();
assert!(matches!(state.init_progress(), InitProgress::NotStarted));
assert!(state.try_start_init(), "cancel_init re-enables init");
}
fn state_label(s: &ClientState) -> &'static str {
match s {
ClientState::Empty => "Empty",
ClientState::Pending(_) => "Pending",
ClientState::Initializing => "Initializing",
ClientState::Ready(_) => "Ready",
}
}
// -- is_healthy / state_kind --------------------------------------
//
// These tests cover the cheap, non-blocking predicate. They focus
// on the state-machine inspection: any
// non-`Ready` variant returns `false` for `is_healthy`, and
// `state_kind` projects every variant onto the matching
// [`ClientStateKind`].
//
// The two `Ready` cases
// (`is_healthy_ready_open_returns_true` and
// `is_healthy_transport_closed_returns_false`) require a real
// `RunningService<RoleClient, InitializeRequestParams>`, which can
// only be constructed through rmcp's `serve_client` path. That
// path needs a peer that responds to the MCP initialize
// handshake, and this crate intentionally does NOT enable rmcp's
// `server` feature (see `Cargo.toml`). Wiring up a hand-rolled
// JSON-RPC responder over `tokio::io::duplex` would balloon the
// test scaffolding far beyond what these tests need. We therefore
// exercise the `Ready` arm indirectly: the cheap predicate is a
// single `match` on the state mutex plus
// `Peer::is_transport_closed`, which is upstream-tested in rmcp
// itself (`rmcp-2.1.0/tests/test_close_connection.rs`).
#[tokio::test]
async fn is_healthy_empty_returns_false() {
let client = McpClient::stub("empty");
// `stub` starts in `ClientState::Empty`.
assert!(matches!(*client.state.lock().await, ClientState::Empty));
assert!(!client.is_healthy().await);
assert_eq!(client.state_kind().await, ClientStateKind::Empty);
}
#[tokio::test]
async fn is_healthy_pending_returns_false() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
let client = McpClient::new_http("pending".to_string(), config, None, None);
// `new_http` constructs with `ClientState::Pending(_)`.
assert!(matches!(
*client.state.lock().await,
ClientState::Pending(_)
));
assert!(!client.is_healthy().await);
assert_eq!(client.state_kind().await, ClientStateKind::Pending);
}
#[tokio::test]
async fn is_healthy_initializing_returns_false() {
let client = McpClient::stub("initializing");
*client.state.lock().await = ClientState::Initializing;
assert!(!client.is_healthy().await);
assert_eq!(client.state_kind().await, ClientStateKind::Initializing);
}
/// `is_healthy` MUST NOT trigger a handshake. Regression guard:
/// the previous implementation called `ensure_initialized`, which
/// for a `Pending` HTTP client pointing at an unreachable host
/// would block for `startup_timeout_sec` seconds. The cheap
/// predicate must return immediately.
#[tokio::test]
async fn is_healthy_pending_does_not_block_on_handshake() {
let config = HttpConfig {
url: "http://192.0.2.1:1/unreachable".to_string(),
headers: vec![],
};
// Force a generous startup timeout — if the predicate
// regressed to going through ensure_initialized, this test
// would hang for ~10 s. We assert it completes in well under
// a second.
let overrides = McpClientTimeoutOverrides {
startup_timeout_sec: Some(10),
..Default::default()
};
let client = McpClient::new_http(
"pending-unreachable".to_string(),
config,
Some(&overrides),
None,
);
let start = std::time::Instant::now();
let healthy = client.is_healthy().await;
let elapsed = start.elapsed();
assert!(!healthy);
// 1 s bound: the cheap path is microseconds, so this is a 10×
// safety margin against cold-runtime / contended-CI jitter while
// still firing well inside the 10 s blocking window that a
// regressed predicate (back through `ensure_initialized`) would
// sit in.
assert!(
elapsed < std::time::Duration::from_secs(1),
"is_healthy must be a cheap state inspection, took {elapsed:?}"
);
}
// -- KigiClientHandler --------------------------------------
//
// The handler's notification routing is the only behavior worth
// unit-testing here; `get_info` is a literal `info.clone()` and
// doesn't merit a test. `NotificationContext` is non-trivial to
// construct outside of an rmcp `RunningService`, so we exercise
// the routing through the `emit` helper that the trait methods
// call. If the trait wiring (one-line `async move { self.emit(...) }`)
// ever regresses, the integration tests against a real MCP
// server will catch it.
#[tokio::test]
async fn client_handler_routes_tools_changed() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<McpClientEvent>();
let handler = KigiClientHandler {
info: McpClient::make_client_info("test"),
server_name: "test".to_string(),
notify_tx: Arc::new(parking_lot::Mutex::new(Some(tx))),
};
handler.emit(McpClientEvent::ToolsChanged {
server: handler.server_name.clone(),
});
let ev = rx.recv().await.expect("event arrived");
match ev {
McpClientEvent::ToolsChanged { server } => assert_eq!(server, "test"),
other => panic!("expected ToolsChanged, got {other:?}"),
}
}
/// Contract: when `notify_tx` is `None` (subagent snapshot,
/// no dispatcher), `emit` is a no-op and the trait methods
/// must not panic.
#[tokio::test]
async fn client_handler_no_dispatcher_is_silent() {
let handler = KigiClientHandler {
info: McpClient::make_client_info("test"),
server_name: "test".to_string(),
notify_tx: Arc::new(parking_lot::Mutex::new(None)),
};
handler.emit(McpClientEvent::ToolsChanged {
server: "test".to_string(),
});
// No assertion needed — reaching this line means no panic.
}
/// Contract: get_info returns a clone of the stored ClientInfo.
#[tokio::test]
async fn client_handler_get_info_round_trips() {
let info = McpClient::make_client_info("test-srv");
let handler = KigiClientHandler {
info: info.clone(),
server_name: "test-srv".to_string(),
notify_tx: Arc::new(parking_lot::Mutex::new(None)),
};
let got = handler.get_info();
// ClientInfo doesn't derive PartialEq; check the visible
// fields the constructor sets.
assert_eq!(got.client_info.name, info.client_info.name);
assert_eq!(got.client_info.version, info.client_info.version);
}
// A sender wired *after* the handler is constructed must still
// reach the live rmcp service loop. This test exercises the
// post-construction wiring path: build a handler from a client
// whose slot is `None`, then install a sender via
// `client.set_event_tx` and verify the handler picks it up (the
// handler holds a clone of the same shared Arc slot).
#[tokio::test]
async fn client_handler_observes_post_handshake_set_event_tx() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<McpClientEvent>();
// McpClient::stub initializes notify_tx as `Arc<Mutex<None>>`.
let client = Arc::new(McpClient::stub("test"));
// Build the handler BEFORE wiring the sender — emulates
// the production flow where `make_client_handler` is called
// during `try_handshake` and the dispatcher is wired
// separately.
let handler = client.make_client_handler();
// Confirm the slot is `None` at handler-construction time.
assert!(handler.notify_tx.lock().is_none());
// Now wire the sender on the client. Because the handler
// holds a CLONE OF THE SAME ARC, this mutation is observed
// by the handler's next `emit`.
client.set_event_tx(Some(tx));
handler.emit(McpClientEvent::ToolsChanged {
server: "test".to_string(),
});
let ev = rx.recv().await.expect("event arrived");
match ev {
McpClientEvent::ToolsChanged { server } => assert_eq!(server, "test"),
other => panic!("expected ToolsChanged, got {other:?}"),
}
}
// Mirrors the post-construction wiring on the `ensure_initialized`
// emit path: even though `Ready` / `HandshakeFailed` fire from
// inside `try_handshake`, the slot is read at emit time through the
// SAME shared Arc, so wiring `set_event_tx` BEFORE the handshake is
// sufficient to capture these events.
#[tokio::test]
async fn event_tx_clone_observes_set_event_tx() {
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel::<McpClientEvent>();
let client = McpClient::stub("test");
assert!(client.event_tx_clone().is_none());
client.set_event_tx(Some(tx));
assert!(client.event_tx_clone().is_some());
client.set_event_tx(None);
assert!(client.event_tx_clone().is_none());
}
// An `ensure_initialized`-emitted `Ready` event must NOT be
// conflated with a restart. This unit test exercises the event
// level; the wire-level mapping ("Ready → reason=initialized, NOT
// restart_succeeded") is covered by host integration tests.
#[test]
fn config_added_kind_carries_correct_server_name() {
let ev = McpClientEvent::ConfigAdded {
server: "srv".to_string(),
};
assert_eq!(ev.server_name(), Some("srv"));
}
}