use anda_core::{
BoxError, BoxFut, CancellationToken, FunctionDefinition, Json, StateFeatures, ToolGroup,
ToolInput, ToolOutput, ToolProvider, validate_function_name,
};
use parking_lot::{Mutex as SyncMutex, RwLock};
use rmcp::{
Peer, RoleClient,
model::{CallToolRequestParams, ServerPeerInfo, Tool as McpTool},
serve_client_with_lifecycle,
service::ClientLifecycleMode,
transport::{AuthClient, AuthorizationManager, StreamableHttpClientTransport},
};
use serde::{Deserialize, Serialize};
use serde_json::Map;
pub use rmcp::transport::StoredCredentials;
use std::{
collections::{BTreeMap, BTreeSet, HashMap},
sync::{
Arc,
atomic::{AtomicBool, AtomicU64, Ordering},
},
time::Duration,
};
use tokio::sync::Mutex;
use crate::context::BaseCtx;
mod auth;
mod bounded;
mod catalog;
mod http_client;
mod interaction;
mod policy;
mod presentation;
mod resources;
pub use interaction::McpElicitationHandler;
use catalog::{DirtyGuard, Registration, Snapshot, collect_pages};
pub use policy::{McpConcurrency, McpLimits, McpServerStatus, McpStartup, McpTimeouts};
mod router;
mod session;
pub use auth::{
InMemoryMcpCredentialStore, McpAuthorizationRequired, McpCredentialStore, McpOAuthConfig,
McpOAuthMetadata, OAuthAuthorizationCodeConfig, OAuthClientCredentialsConfig,
};
pub use router::McpToolRoute;
pub use session::{McpStdioTransport, McpStreamableHttpTransport, McpTransportConfig};
use auth::{ScopedCredentialStore, authorization_required_hint, is_authorization_error};
use router::{
DEFAULT_TASK_MAX_WAIT_SECS, MAX_LOCAL_NAME_ATTEMPTS, MAX_TASK_MAX_WAIT_SECS, call_tool_rounds,
mcp_result_to_tool_output, sanitize_name_part, shorten_with_hash,
};
use session::{
AndaMcpClient, DISCOVERY_PROBE_TIMEOUT, McpSession, SUBSCRIPTION_ACK_TIMEOUT,
legacy_protocol_version, needs_tool_subscription, preferred_protocol_versions,
pump_tool_subscription, serve_bounded, tool_subscription_filter,
};
pub const DEFAULT_MCP_TOOL_PREFIX: &str = "mcp";
#[derive(Clone)]
pub struct McpToolProvider {
inner: Arc<McpToolProviderInner>,
}
impl std::fmt::Debug for McpToolProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpToolProvider")
.field("name", &self.inner.name)
.field("tool_prefix", &self.inner.tool_prefix)
.field("servers", &self.server_ids())
.finish()
}
}
impl McpToolProvider {
pub fn new(servers: Vec<McpServerConfig>) -> Result<Self, BoxError> {
Self::builder().servers(servers).build()
}
pub fn builder() -> McpToolProviderBuilder {
McpToolProviderBuilder::default()
}
pub async fn add_server(&self, server: McpServerConfig) -> Result<(), BoxError> {
let registration = self.insert_server(server)?;
let refresh = async {
let snapshot = self
.fetch_registration_snapshot(registration.clone(), true)
.await?;
self.publish_snapshot(snapshot).await
};
let result = tokio::select! {
biased;
_ = registration.cancelled.cancelled() => Err("MCP server removed during registration".into()),
result = refresh => result,
};
if let Err(err) = result {
self.remove_registration(®istration);
return Err(err);
}
Ok(())
}
pub fn register_server(&self, server: McpServerConfig) -> Result<(), BoxError> {
self.insert_server(server).map(|_| ())
}
pub fn remove_server(&self, server_id: &str) -> bool {
let existed = self.contains_server(server_id);
self.remove_server_state(server_id);
existed
}
pub fn contains_server(&self, server_id: &str) -> bool {
self.inner.servers.read().contains_key(server_id)
}
pub fn server_ids(&self) -> Vec<String> {
self.inner.servers.read().keys().cloned().collect()
}
pub async fn refresh_server(&self, server_id: &str) -> Result<(), BoxError> {
self.refresh_server_inner(server_id, true).await
}
async fn refresh_server_inner(
&self,
server_id: &str,
clear_dirty: bool,
) -> Result<(), BoxError> {
let snapshot = self.fetch_server_snapshot(server_id, clear_dirty).await?;
self.publish_snapshot(snapshot).await
}
async fn publish_snapshot(&self, snapshot: Option<Snapshot>) -> Result<(), BoxError> {
let Some(mut snapshot) = snapshot else {
return Ok(());
};
let registration = &snapshot.registration;
let _catalog = registration.catalog.write().await;
let servers = self.inner.servers.read();
if !servers
.get(®istration.id)
.is_some_and(|current| Arc::ptr_eq(current, registration))
{
return Err("MCP server registration changed during discovery".into());
}
let mut index = self.inner.index.write();
if !index
.sessions
.get(®istration.id)
.is_some_and(|session| Arc::ptr_eq(session, &snapshot.session))
{
return Err("MCP session changed during discovery".into());
}
let known = index
.names
.keys()
.filter(|(id, _)| id == ®istration.id)
.count();
let added = snapshot
.routes
.iter()
.filter(|route| {
!index
.names
.contains_key(&(route.server_id.clone(), route.remote_name.clone()))
})
.count();
if known.saturating_add(added) > registration.limits.catalog_items {
return Err(
"MCP catalog identity budget exhausted; re-register the server to reset it".into(),
);
}
let revision = registration.revision.fetch_add(1, Ordering::SeqCst) + 1;
for route in &mut snapshot.routes {
route.catalog_revision = revision;
}
index.replace_server_routes(®istration.id, snapshot.routes);
index.metas.insert(registration.id.clone(), snapshot.meta);
*registration.status.lock() = McpServerStatus::Ready;
snapshot.dirty.commit();
Ok(())
}
async fn fetch_server_snapshot(
&self,
server_id: &str,
clear_dirty: bool,
) -> Result<Option<Snapshot>, BoxError> {
self.fetch_registered_snapshot(self.server_config(server_id)?, clear_dirty)
.await
}
async fn fetch_registered_snapshot(
&self,
config: Arc<Registration>,
clear_dirty: bool,
) -> Result<Option<Snapshot>, BoxError> {
let result = tokio::select! {
biased;
_ = config.cancelled.cancelled() => Err("MCP server removed".into()),
result = self.fetch_registration_snapshot(config.clone(), clear_dirty) => result,
};
if let Err(err) = &result {
*config.status.lock() = if is_authorization_error(err.as_ref()) {
McpServerStatus::AuthorizationRequired
} else {
McpServerStatus::Failed
};
}
result
}
async fn fetch_registration_snapshot(
&self,
config: Arc<Registration>,
clear_dirty: bool,
) -> Result<Option<Snapshot>, BoxError> {
let guard = config.refresh.clone().lock_owned().await;
let session = self.ensure_session(&config).await?;
let peer = session.service.lock().await.peer().clone();
if !clear_dirty && !session.dirty.load(Ordering::SeqCst) {
return Ok(None);
}
let dirty = DirtyGuard::new(session.dirty.clone());
session.dirty.store(false, Ordering::SeqCst);
if clear_dirty {
peer.clear_response_cache().await;
}
let meta = McpServerMeta::from_peer_info(&config.id, peer.peer_info().as_deref());
let metadata_bytes = [&meta.title, &meta.description, &meta.instructions]
.into_iter()
.flatten()
.map(String::len)
.sum::<usize>();
if metadata_bytes > config.limits.server_metadata_bytes {
return Err("MCP server metadata limit exceeded".into());
}
let tools = tokio::time::timeout(
Duration::from_secs(config.timeouts.list_secs),
collect_pages(&config.limits, |params| {
let peer = peer.clone();
async move {
let page = http_client::retry_read(|| peer.list_tools(params.clone())).await?;
Ok((page.tools, page.next_cursor))
}
}),
)
.await
.map_err(|_| format!("MCP server {} tools/list timed out", config.id))??;
let routes = self.routes_for_tools(&config.id, tools)?;
Ok(Some(Snapshot {
registration: config,
session,
routes,
meta,
dirty,
_refresh: guard,
}))
}
pub async fn server_statuses(&self) -> BTreeMap<String, McpServerStatus> {
let registrations: Vec<_> = self.inner.servers.read().values().cloned().collect();
let mut states = BTreeMap::new();
for registration in registrations {
let mut status = *registration.status.lock();
if matches!(status, McpServerStatus::Ready | McpServerStatus::Connected)
&& self.live_session(®istration.id).await.is_none()
{
status = McpServerStatus::Disconnected;
}
states.insert(registration.id.clone(), status);
}
states
}
pub fn tool_groups(&self) -> Vec<ToolGroup> {
let index = self.inner.index.read();
let mut members: BTreeMap<String, Vec<String>> = BTreeMap::new();
for route in index.routes.values() {
members
.entry(route.server_id.clone())
.or_default()
.push(route.name.clone());
}
members
.into_iter()
.map(|(server_id, mut names)| {
names.sort();
let meta = index.metas.get(&server_id).cloned().unwrap_or_default();
ToolGroup {
id: format!("{}:{}", self.inner.name, server_id),
title: meta.resolved_title(&server_id),
description: meta.resolved_description(&server_id),
instructions: meta.instructions,
members: names,
}
})
.collect()
}
pub fn routes(&self) -> Vec<McpToolRoute> {
self.inner.index.read().routes.values().cloned().collect()
}
async fn refresh_servers(&self, tolerant: bool) -> Result<(), BoxError> {
let registrations = self.inner.servers.read().values().cloned().collect();
self.refresh_registrations(registrations, tolerant).await
}
async fn refresh_registrations(
&self,
registrations: Vec<Arc<Registration>>,
tolerant: bool,
) -> Result<(), BoxError> {
let results = futures::future::join_all(
registrations
.iter()
.map(|registration| self.fetch_registered_snapshot(registration.clone(), true)),
)
.await;
let mut errors = Vec::new();
for (registration, result) in registrations.iter().zip(results) {
let result = match result {
Ok(snapshot) => self.publish_snapshot(snapshot).await,
Err(err) => Err(err),
};
if let Err(err) = result {
if tolerant && !registration.required {
log::warn!("MCP server {} discovery failed: {err}", registration.id);
} else {
errors.push(format!("{}: {err}", registration.id));
}
}
}
if errors.is_empty() {
Ok(())
} else {
Err(format!("failed to refresh MCP servers: {}", errors.join("; ")).into())
}
}
fn server_config(&self, server_id: &str) -> Result<Arc<Registration>, BoxError> {
self.inner
.servers
.read()
.get(server_id)
.cloned()
.ok_or_else(|| format!("MCP server {} not configured", server_id).into())
}
async fn live_session(&self, server_id: &str) -> Option<Arc<McpSession>> {
let session = self.inner.index.read().sessions.get(server_id).cloned()?;
if session.is_closed().await {
None
} else {
Some(session)
}
}
async fn ensure_session(
&self,
config: &Arc<Registration>,
) -> Result<Arc<McpSession>, BoxError> {
if let Some(session) = self.live_session(&config.id).await {
return Ok(session);
}
let _guard = config.connect.lock().await;
if config.cancelled.is_cancelled() {
return Err("MCP server removed".into());
}
if let Some(session) = self.live_session(&config.id).await {
return Ok(session);
}
*config.status.lock() = McpServerStatus::Connecting;
let _status_guard = catalog::ConnectingGuard(config);
let attempt = tokio::time::timeout(Duration::from_secs(config.timeouts.setup_secs), async { match self.connect(config, config.lifecycle.into_mode()).await {
Ok(session) => Ok(session),
Err(err)
if config.lifecycle == McpLifecycle::Auto
&& !is_authorization_error(err.as_ref()) =>
{
log::info!(
"MCP server {}: discovery lifecycle failed ({err}); retrying with the legacy initialize handshake",
config.id
);
self.connect(config, ClientLifecycleMode::Initialize).await
}
Err(err) => Err(err),
} }).await.map_err(|_| format!("MCP server {} setup timed out", config.id))?;
let session = attempt.map_err(|err| authorization_required_hint(config, err))?;
let servers = self.inner.servers.read();
if !servers
.get(&config.id)
.is_some_and(|current| Arc::ptr_eq(current, config))
{
return Err("MCP server registration changed during connection".into());
}
self.inner
.index
.write()
.sessions
.insert(config.id.clone(), session.clone());
*config.status.lock() = McpServerStatus::Connected;
Ok(session)
}
async fn connect(
&self,
config: &McpServerConfig,
lifecycle: ClientLifecycleMode,
) -> Result<Arc<McpSession>, BoxError> {
tokio::time::timeout(session::SESSION_SETUP_TIMEOUT, async {
for delay in [250, 1_000] {
match self.connect_inner(config, lifecycle.clone()).await {
Ok(session) => return Ok(session),
Err(err) if http_client::is_transient(err.as_ref()) => {
tokio::time::sleep(Duration::from_millis(delay)).await
}
Err(err) => return Err(err),
}
}
self.connect_inner(config, lifecycle).await
})
.await
.map_err(|_| format!("MCP server {} session setup timed out", config.id))?
}
async fn connect_inner(
&self,
config: &McpServerConfig,
lifecycle: ClientLifecycleMode,
) -> Result<Arc<McpSession>, BoxError> {
let dirty = Arc::new(AtomicBool::new(true));
let session_cancelled = CancellationToken::new();
let elicitation = if config.elicitation {
Some(interaction::ElicitationDispatcher {
handler: self
.inner
.elicitation_handler
.clone()
.ok_or("MCP elicitation requires an application handler")?,
server_id: config.id.clone(),
timeout: Duration::from_secs(config.timeouts.elicitation_secs),
session_cancelled: session_cancelled.clone(),
})
} else {
None
};
let handler = AndaMcpClient::new(dirty.clone(), config.tasks.is_some())
.with_elicitation(elicitation.clone());
let probe_timeout = (!matches!(lifecycle, ClientLifecycleMode::Initialize))
.then_some(DISCOVERY_PROBE_TIMEOUT);
let mut expires_at = None;
let mut process = None;
let service = match &config.transport {
McpTransportConfig::Stdio(stdio) => {
let (child, transport) =
bounded::spawn(stdio.command(), config.limits.message_bytes)?;
process = Some(child);
serve_bounded(
serve_client_with_lifecycle(handler, transport, lifecycle),
probe_timeout,
)
.await?
}
McpTransportConfig::StreamableHttp(http) => match &http.auth {
None => {
let transport = StreamableHttpClientTransport::with_client(
http_client::McpHttpClient::new(config.limits.message_bytes)?,
http.transport_config()?
.max_sse_event_size(config.limits.message_bytes),
);
serve_bounded(
serve_client_with_lifecycle(handler, transport, lifecycle),
probe_timeout,
)
.await?
}
Some(McpOAuthConfig::ClientCredentials(cc)) => {
let (manager, deadline) =
auth::authorize_client_credentials(http.url.as_str(), cc).await?;
expires_at = deadline;
let transport = StreamableHttpClientTransport::with_client(
AuthClient::new(
http_client::McpHttpClient::new(config.limits.message_bytes)?,
manager,
),
http.base_transport_config()?
.max_sse_event_size(config.limits.message_bytes),
);
serve_bounded(
serve_client_with_lifecycle(handler, transport, lifecycle),
probe_timeout,
)
.await?
}
Some(McpOAuthConfig::AuthorizationCode(_)) => {
let manager = auth::authorize_from_store(
&config.id,
http.url.as_str(),
self.scoped_store(&config.id),
)
.await?;
let transport = StreamableHttpClientTransport::with_client(
AuthClient::new(
http_client::McpHttpClient::new(config.limits.message_bytes)?,
manager,
),
http.base_transport_config()?
.max_sse_event_size(config.limits.message_bytes),
);
serve_bounded(
serve_client_with_lifecycle(handler, transport, lifecycle),
probe_timeout,
)
.await?
}
},
};
let subscription = self
.subscribe_tool_changes(&config.id, service.peer(), dirty.clone())
.await;
Ok(Arc::new(McpSession {
retired: AtomicBool::new(false),
_process: process,
session_cancelled,
elicitation,
service: Mutex::new(service),
dirty,
expires_at,
subscription,
}))
}
async fn subscribe_tool_changes(
&self,
server_id: &str,
peer: &Peer<RoleClient>,
dirty: Arc<AtomicBool>,
) -> Option<tokio::task::JoinHandle<()>> {
if !needs_tool_subscription(peer.peer_info().as_deref()) {
return None;
}
match tokio::time::timeout(
SUBSCRIPTION_ACK_TIMEOUT,
peer.listen(tool_subscription_filter()),
)
.await
{
Ok(Ok(subscription)) => {
let peer = peer.clone();
let server_id = server_id.to_string();
Some(tokio::spawn(pump_tool_subscription(
server_id,
peer,
subscription,
dirty,
)))
}
Ok(Err(err)) => {
log::warn!(
"MCP server {server_id}: could not subscribe to tools/list_changed: {err}"
);
None
}
Err(_) => {
log::warn!(
"MCP server {server_id}: tools/list_changed subscription was not acknowledged \
within {}s",
SUBSCRIPTION_ACK_TIMEOUT.as_secs()
);
None
}
}
}
pub async fn discover_http_oauth(url: &str) -> Result<Option<McpOAuthMetadata>, BoxError> {
auth::discover_http_oauth(url).await
}
pub async fn begin_authorization(&self, server_id: &str) -> Result<String, BoxError> {
let config = self.server_config(server_id)?;
let McpTransportConfig::StreamableHttp(http) = &config.transport else {
return Err(format!("MCP server {server_id} does not use the HTTP transport").into());
};
let Some(McpOAuthConfig::AuthorizationCode(ac)) = &http.auth else {
return Err(format!(
"MCP server {server_id} is not configured for the OAuth authorization_code flow"
)
.into());
};
let (manager, auth_url) =
auth::begin_authorization_manager(http.url.as_str(), ac, self.scoped_store(server_id))
.await?;
let servers = self.inner.servers.read();
if !servers
.get(server_id)
.is_some_and(|current| Arc::ptr_eq(current, &config))
{
return Err("MCP server changed during authorization".into());
}
self.inner
.pending_auth
.lock()
.insert(server_id.to_string(), (config.clone(), manager));
Ok(auth_url)
}
pub async fn complete_authorization(
&self,
server_id: &str,
redirect_url: &str,
) -> Result<(), BoxError> {
let (registration, manager) = self
.inner
.pending_auth
.lock()
.remove(server_id)
.ok_or_else(|| format!("no pending OAuth authorization for MCP server {server_id}"))?;
let credential_guard = self
.inner
.credential_store
.acquire_refresh_guard(server_id)
.await?;
tokio::select! {
biased;
_ = registration.cancelled.cancelled() => return Err("MCP server removed during authorization".into()),
result = auth::complete_authorization_exchange(manager, redirect_url) => result?,
}
drop(credential_guard);
if !self
.server_config(server_id)
.is_ok_and(|current| Arc::ptr_eq(¤t, ®istration))
{
return Err("MCP server changed during authorization".into());
}
self.disconnect_server(server_id).await;
Ok(())
}
pub fn cancel_authorization(&self, server_id: &str) -> bool {
self.inner.pending_auth.lock().remove(server_id).is_some()
}
pub async fn disconnect_server(&self, server_id: &str) -> bool {
let Ok(config) = self.server_config(server_id) else {
return false;
};
let _refresh = config.refresh.lock().await;
let _guard = config.connect.lock().await;
let servers = self.inner.servers.read();
if !servers
.get(server_id)
.is_some_and(|current| Arc::ptr_eq(current, &config))
{
return false;
}
*config.status.lock() = McpServerStatus::Disconnected;
let session = self.inner.index.write().sessions.remove(server_id);
if let Some(session) = &session {
session.retired.store(true, Ordering::SeqCst);
}
session.is_some()
}
pub async fn clear_credentials(&self, server_id: &str) -> Result<(), BoxError> {
let _guard = self
.inner
.credential_store
.acquire_refresh_guard(server_id)
.await?;
self.inner.credential_store.clear(server_id).await?;
drop(_guard);
self.disconnect_server(server_id).await;
Ok(())
}
fn scoped_store(&self, server_id: &str) -> ScopedCredentialStore {
ScopedCredentialStore {
server_id: server_id.to_string(),
inner: self.inner.credential_store.clone(),
}
}
async fn refresh_if_dirty(&self, server_id: &str) -> Result<(), BoxError> {
let snapshot = self.fetch_server_snapshot(server_id, false).await?;
self.publish_snapshot(snapshot).await
}
fn routes_for_tools(
&self,
server_id: &str,
tools: Vec<McpTool>,
) -> Result<Vec<McpToolRoute>, BoxError> {
let config = self.server_config(server_id)?;
let mut routes = Vec::new();
let mut used = BTreeSet::new();
let mut tools = tools;
tools.sort_by(|a, b| a.name.cmp(&b.name));
let mut counts = BTreeMap::<String, usize>::new();
for tool in &tools {
*counts.entry(sanitize_name_part(&tool.name)).or_default() += 1;
}
let mut remote_names = BTreeSet::new();
for tool in tools {
if !remote_names.insert(tool.name.to_string()) {
return Err("MCP duplicate remote tool name".into());
}
let remote_name = tool.name.to_string();
if !router::tool_is_model_visible(&tool) || !self.includes_tool(server_id, &remote_name)
{
continue;
}
if serde_json::to_vec(tool.input_schema.as_ref())?.len() > config.limits.schema_bytes
|| tool.output_schema.as_ref().is_some_and(|schema| {
serde_json::to_vec(schema.as_ref())
.map_or(true, |v| v.len() > config.limits.schema_bytes)
})
|| tool
.description
.as_ref()
.is_some_and(|v| v.len() > config.limits.description_bytes)
{
return Err(format!(
"MCP tool {} exceeds schema or description limits",
tool.name
)
.into());
}
if tool
.title
.as_ref()
.is_some_and(|title| title.len() > config.limits.description_bytes)
{
return Err("MCP tool title limit exceeded".into());
}
let collision_key =
(counts[&sanitize_name_part(&remote_name)] > 1).then_some(remote_name.as_str());
let mut local_name = self.local_tool_name(server_id, &remote_name, collision_key)?;
let mut attempt = 0usize;
while used.contains(&local_name) {
if attempt >= MAX_LOCAL_NAME_ATTEMPTS {
log::warn!(
"skipping MCP tool {remote_name:?} on server {server_id:?}: could not derive a unique local name"
);
break;
}
let key = format!("{remote_name}#{attempt}");
local_name = self.local_tool_name(server_id, &remote_name, Some(&key))?;
attempt += 1;
}
if used.contains(&local_name) {
continue;
}
used.insert(local_name.clone());
let definition = self.function_definition(server_id, &local_name, &tool);
routes.push(McpToolRoute {
name: local_name,
server_id: server_id.to_string(),
remote_name,
definition,
tool,
server_generation: config.generation,
catalog_revision: config.revision.load(Ordering::SeqCst),
});
}
Ok(routes)
}
fn includes_tool(&self, server_id: &str, remote_name: &str) -> bool {
let Some(config) = self.inner.servers.read().get(server_id).cloned() else {
return false;
};
if config.exclude.contains(remote_name) {
return false;
}
config.include.is_empty() || config.include.contains(remote_name)
}
fn function_definition(
&self,
server_id: &str,
local_name: &str,
tool: &McpTool,
) -> FunctionDefinition {
let mut description = format!("MCP server `{server_id}` tool `{}`.", tool.name);
if let Some(title) = tool.title.as_ref().filter(|title| !title.trim().is_empty()) {
description.push_str(" Title: ");
description.push_str(title.trim());
description.push('.');
}
if let Some(remote_description) = tool
.description
.as_ref()
.map(|description| description.trim())
.filter(|description| !description.is_empty())
{
description.push(' ');
description.push_str(remote_description);
}
FunctionDefinition {
name: local_name.to_string(),
description,
parameters: router::model_schema(&tool.input_schema),
strict: Some(false),
}
}
fn local_tool_name(
&self,
server_id: &str,
remote_name: &str,
collision_key: Option<&str>,
) -> Result<String, BoxError> {
let server = sanitize_name_part(server_id);
let tool = sanitize_name_part(remote_name);
let base = format!("{}_{}_{}", self.inner.tool_prefix, server, tool);
let name = match collision_key {
Some(key) => shorten_with_hash(&base, &format!("{server_id}:{key}")),
None if base.len() > 64 => {
shorten_with_hash(&base, &format!("{server_id}:{remote_name}"))
}
None => base,
};
validate_function_name(&name)?;
Ok(name)
}
#[cfg(test)]
async fn call_route(
&self,
route: McpToolRoute,
input: ToolInput<Json>,
) -> Result<ToolOutput<Json>, BoxError> {
self.call_route_with_cancellation(route, input, CancellationToken::new())
.await
}
async fn call_route_with_cancellation(
&self,
route: McpToolRoute,
input: ToolInput<Json>,
cancellation: CancellationToken,
) -> Result<ToolOutput<Json>, BoxError> {
let config = self.server_config(&route.server_id)?;
if route.server_generation != config.generation {
return Err("MCP server registration changed; select the tool again".into());
}
tokio::select! {
biased;
_ = cancellation.cancelled() => Err("MCP tool call cancelled".into()),
_ = config.cancelled.cancelled() => Err("MCP server removed".into()),
result = tokio::time::timeout(config.timeouts.call(), self.execute_route(&config, route, input, &cancellation)) =>
result.map_err(|_| "MCP logical tool call timed out")?,
}
}
async fn execute_route(
&self,
config: &Arc<Registration>,
route: McpToolRoute,
input: ToolInput<Json>,
cancellation: &CancellationToken,
) -> Result<ToolOutput<Json>, BoxError> {
self.refresh_if_dirty(&route.server_id).await?;
let _catalog = config.catalog.read().await;
let current = self
.inner
.index
.read()
.routes
.get(&route.name)
.cloned()
.ok_or("MCP tool was removed; select a tool again")?;
if current.server_generation != route.server_generation
|| current.remote_name != route.remote_name
|| current.tool != route.tool
{
return Err("MCP tool catalog changed; select the tool again".into());
}
let session = self
.live_session(&route.server_id)
.await
.ok_or("MCP session closed; refresh the tool catalog")?;
let parallel = config.concurrency == McpConcurrency::Parallel
|| (config.concurrency == McpConcurrency::ReadOnlyParallel
&& current
.tool
.annotations
.as_ref()
.and_then(|a| a.read_only_hint)
== Some(true));
let (_read, _write) = if parallel {
(Some(config.calls.read().await), None)
} else {
(None, Some(config.calls.write().await))
};
if session.retired.load(Ordering::SeqCst) {
return Err("MCP session was disconnected before execution".into());
}
let arguments = match input.args {
Json::Object(map) => map,
Json::Null => Map::new(),
_ => {
return Err(
format!("MCP tool {} expects JSON object arguments", route.name).into(),
);
}
};
let params =
CallToolRequestParams::new(current.remote_name.clone()).with_arguments(arguments);
let peer = session.service.lock().await.peer().clone();
let result = call_tool_rounds(
¤t,
&peer,
params,
config.tasks.as_ref(),
cancellation,
Duration::from_secs(config.timeouts.request_secs),
session.elicitation.as_ref(),
)
.await;
let result = match result {
Ok(result) => result,
Err(err) if is_authorization_error(err.as_ref()) => {
*config.status.lock() = McpServerStatus::AuthorizationRequired;
let mut output = ToolOutput::new(
serde_json::json!({"error": {"code": "authorization_required", "server_id": config.id}}),
);
output.is_error = Some(true);
return Ok(output);
}
Err(err) => return Err(err),
};
Ok(mcp_result_to_tool_output(¤t, result, &config.limits))
}
fn insert_server(&self, server: McpServerConfig) -> Result<Arc<Registration>, BoxError> {
server.validate()?;
let server_id = server.id.clone();
let sanitized = sanitize_name_part(&server_id);
let mut servers = self.inner.servers.write();
if servers.contains_key(&server_id) {
return Err(format!("MCP server {} already exists", server_id).into());
}
if servers
.keys()
.map(|existing| sanitize_name_part(existing))
.any(|existing| existing == sanitized)
{
return Err(format!(
"MCP server id {} collides with another server after normalization to {}",
server_id, sanitized
)
.into());
}
let generation = self.inner.next_generation.fetch_add(1, Ordering::SeqCst);
let registration = Arc::new(Registration::new(server, generation));
servers.insert(server_id, registration.clone());
Ok(registration)
}
fn remove_registration(&self, expected: &Arc<Registration>) {
let mut servers = self.inner.servers.write();
if !servers
.get(&expected.id)
.is_some_and(|current| Arc::ptr_eq(current, expected))
{
return;
}
servers.remove(&expected.id);
expected.cancelled.cancel();
self.inner.index.write().remove_server(&expected.id);
self.inner.pending_auth.lock().remove(&expected.id);
}
fn remove_server_state(&self, server_id: &str) {
if let Ok(registration) = self.server_config(server_id) {
self.remove_registration(®istration);
}
}
}
impl ToolProvider<BaseCtx> for McpToolProvider {
fn name(&self) -> String {
self.inner.name.clone()
}
fn definitions(&self, names: Option<&[String]>) -> Vec<FunctionDefinition> {
let index = self.inner.index.read();
match names {
Some([]) => Vec::new(),
Some(names) => names
.iter()
.filter_map(|name| {
index
.routes
.get(&name.to_ascii_lowercase())
.map(|route| route.definition.clone())
})
.collect(),
None => index
.routes
.values()
.map(|route| route.definition.clone())
.collect(),
}
}
fn contains_lowercase(&self, lowercase_name: &str) -> bool {
self.inner.index.read().routes.contains_key(lowercase_name)
}
fn groups(&self) -> Vec<ToolGroup> {
self.tool_groups()
}
fn init(&self, ctx: BaseCtx) -> BoxFut<'_, Result<(), BoxError>> {
Box::pin(async move {
let cancellation = ctx.cancellation_token();
let registrations: Vec<_> = self.inner.servers.read().values().cloned().collect();
let (background, eager): (Vec<_>, Vec<_>) = registrations
.into_iter()
.partition(|server| server.startup == McpStartup::Background);
tokio::select! {
biased;
_ = cancellation.cancelled() => return Err("MCP initialization cancelled".into()),
result = self.refresh_registrations(eager, true) => result?,
}
if !background.is_empty() {
let provider = self.clone();
tokio::spawn(async move {
tokio::select! {
biased;
_ = cancellation.cancelled() => {},
result = provider.refresh_registrations(background, true) => {
if let Err(err) = result { log::warn!("MCP background discovery failed: {err}"); }
}
}
});
}
Ok(())
})
}
fn refresh(&self) -> BoxFut<'_, Result<(), BoxError>> {
Box::pin(async move { self.refresh_servers(false).await })
}
fn call(
&self,
ctx: BaseCtx,
mut input: ToolInput<Json>,
) -> BoxFut<'_, Result<ToolOutput<Json>, BoxError>> {
Box::pin(async move {
input.name.make_ascii_lowercase();
let route = self
.inner
.index
.read()
.routes
.get(&input.name)
.cloned()
.ok_or_else(|| format!("MCP tool {} not found", input.name))?;
self.call_route_with_cancellation(route, input, ctx.cancellation_token())
.await
})
}
}
#[derive(Default)]
pub struct McpToolProviderBuilder {
name: Option<String>,
tool_prefix: Option<String>,
servers: Vec<McpServerConfig>,
credential_store: Option<Arc<dyn McpCredentialStore>>,
elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
}
impl McpToolProviderBuilder {
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
pub fn tool_prefix(mut self, prefix: impl Into<String>) -> Self {
self.tool_prefix = Some(prefix.into());
self
}
pub fn server(mut self, server: McpServerConfig) -> Self {
self.servers.push(server);
self
}
pub fn servers(mut self, servers: Vec<McpServerConfig>) -> Self {
self.servers = servers;
self
}
pub fn credential_store(mut self, store: Arc<dyn McpCredentialStore>) -> Self {
self.credential_store = Some(store);
self
}
pub fn elicitation_handler(mut self, handler: Arc<dyn McpElicitationHandler>) -> Self {
self.elicitation_handler = Some(handler);
self
}
pub fn build(self) -> Result<McpToolProvider, BoxError> {
let name = self
.name
.unwrap_or_else(|| DEFAULT_MCP_TOOL_PREFIX.to_string());
let name = name.to_ascii_lowercase();
validate_function_name(&name)?;
let tool_prefix = self
.tool_prefix
.unwrap_or_else(|| DEFAULT_MCP_TOOL_PREFIX.to_string());
let tool_prefix = sanitize_name_part(&tool_prefix);
validate_function_name(&tool_prefix)?;
let mut servers = BTreeMap::new();
let mut sanitized_ids = BTreeSet::new();
for server in self.servers {
server.validate()?;
if servers.contains_key(&server.id) {
return Err(format!("duplicate MCP server id {}", server.id).into());
}
let sanitized = sanitize_name_part(&server.id);
if !sanitized_ids.insert(sanitized.clone()) {
return Err(format!(
"MCP server id {} collides with another server after normalization to {}",
server.id, sanitized
)
.into());
}
servers.insert(
server.id.clone(),
Arc::new(Registration::new(server, servers.len() as u64 + 1)),
);
}
let next_generation = AtomicU64::new(servers.len() as u64 + 1);
let credential_store = self
.credential_store
.unwrap_or_else(|| Arc::new(InMemoryMcpCredentialStore::new()));
Ok(McpToolProvider {
inner: Arc::new(McpToolProviderInner {
name,
tool_prefix,
servers: RwLock::new(servers),
next_generation,
index: RwLock::new(McpToolIndex::default()),
credential_store,
elicitation_handler: self.elicitation_handler,
pending_auth: SyncMutex::new(HashMap::new()),
}),
})
}
}
impl std::fmt::Debug for McpToolProviderBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpToolProviderBuilder")
.field("name", &self.name)
.field("tool_prefix", &self.tool_prefix)
.field("servers", &self.servers)
.field("credential_store", &self.credential_store.is_some())
.finish()
}
}
struct McpToolProviderInner {
name: String,
tool_prefix: String,
servers: RwLock<BTreeMap<String, Arc<Registration>>>,
next_generation: AtomicU64,
index: RwLock<McpToolIndex>,
credential_store: Arc<dyn McpCredentialStore>,
elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
pending_auth: SyncMutex<HashMap<String, (Arc<Registration>, AuthorizationManager)>>,
}
#[derive(Default)]
struct McpToolIndex {
routes: BTreeMap<String, McpToolRoute>,
names: BTreeMap<(String, String), String>,
sessions: BTreeMap<String, Arc<McpSession>>,
metas: BTreeMap<String, McpServerMeta>,
}
impl McpToolIndex {
fn replace_server_routes(&mut self, server_id: &str, routes: Vec<McpToolRoute>) {
self.routes.retain(|_, route| route.server_id != server_id);
for mut route in routes {
let identity = (route.server_id.clone(), route.remote_name.clone());
if let Some(name) = self.names.get(&identity) {
route.name.clone_from(name);
} else {
let base = route.name.clone();
for attempt in 0..MAX_LOCAL_NAME_ATTEMPTS {
if !self.names.values().any(|name| name == &route.name) {
break;
}
route.name = shorten_with_hash(
&base,
&format!("{}:{}#{attempt}", route.server_id, route.remote_name),
);
}
if self.names.values().any(|name| name == &route.name) {
continue;
}
self.names.insert(identity, route.name.clone());
}
route.definition.name.clone_from(&route.name);
self.routes.insert(route.name.clone(), route);
}
}
fn remove_server(&mut self, server_id: &str) {
self.routes.retain(|_, route| route.server_id != server_id);
self.sessions.remove(server_id);
self.metas.remove(server_id);
self.names.retain(|(id, _), _| id != server_id);
}
}
#[derive(Debug, Clone, Default)]
struct McpServerMeta {
title: Option<String>,
description: Option<String>,
instructions: Option<String>,
}
impl McpServerMeta {
fn from_peer_info(server_id: &str, info: Option<&ServerPeerInfo>) -> Self {
let Some(info) = info else {
return Self::default();
};
let implementation = info.server_info.as_ref();
let title = implementation
.and_then(|implementation| {
non_empty(implementation.title.as_deref())
.or_else(|| non_empty(Some(implementation.name.as_str())))
})
.filter(|title| title != server_id);
Self {
title,
description: implementation
.and_then(|implementation| non_empty(implementation.description.as_deref())),
instructions: non_empty(info.instructions.as_deref()),
}
}
fn resolved_title(&self, server_id: &str) -> String {
self.title
.clone()
.unwrap_or_else(|| format!("MCP server `{server_id}`"))
}
fn resolved_description(&self, server_id: &str) -> String {
self.description
.clone()
.unwrap_or_else(|| format!("Tools provided by MCP server `{server_id}`."))
}
}
fn non_empty(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct McpServerConfig {
pub id: String,
pub transport: McpTransportConfig,
#[serde(default)]
pub include: BTreeSet<String>,
#[serde(default)]
pub exclude: BTreeSet<String>,
#[serde(default)]
pub lifecycle: McpLifecycle,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tasks: Option<McpTasksConfig>,
#[serde(default)]
pub limits: McpLimits,
#[serde(default)]
pub timeouts: McpTimeouts,
#[serde(default)]
pub concurrency: McpConcurrency,
#[serde(default)]
pub required: bool,
#[serde(default)]
pub startup: McpStartup,
#[serde(default)]
pub elicitation: bool,
#[serde(default)]
pub resources: bool,
}
impl McpServerConfig {
pub fn stdio(id: impl Into<String>, command: impl Into<String>) -> Self {
Self {
id: id.into(),
transport: McpTransportConfig::Stdio(McpStdioTransport {
command: command.into(),
..Default::default()
}),
include: BTreeSet::new(),
exclude: BTreeSet::new(),
lifecycle: McpLifecycle::default(),
tasks: None,
limits: McpLimits::default(),
timeouts: McpTimeouts::default(),
concurrency: McpConcurrency::default(),
required: false,
startup: McpStartup::default(),
elicitation: false,
resources: false,
}
}
pub fn streamable_http(id: impl Into<String>, url: impl Into<String>) -> Self {
Self {
id: id.into(),
transport: McpTransportConfig::StreamableHttp(McpStreamableHttpTransport {
url: url.into(),
..Default::default()
}),
include: BTreeSet::new(),
exclude: BTreeSet::new(),
lifecycle: McpLifecycle::default(),
tasks: None,
limits: McpLimits::default(),
timeouts: McpTimeouts::default(),
concurrency: McpConcurrency::default(),
required: false,
startup: McpStartup::default(),
elicitation: false,
resources: false,
}
}
fn validate(&self) -> Result<(), BoxError> {
validate_function_name(&sanitize_name_part(&self.id))?;
if self.id.trim().is_empty() {
return Err("MCP server id must not be empty".into());
}
self.limits.validate()?;
self.timeouts.validate()?;
if self.required && self.startup == McpStartup::Background {
return Err("required MCP servers cannot start in the background".into());
}
self.transport.validate()
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum McpLifecycle {
#[default]
Auto,
Discover,
Initialize,
}
impl McpLifecycle {
fn into_mode(self) -> ClientLifecycleMode {
match self {
Self::Auto => ClientLifecycleMode::Auto {
preferred_versions: preferred_protocol_versions(),
legacy_version: Some(legacy_protocol_version()),
},
Self::Discover => ClientLifecycleMode::Discover {
preferred_versions: preferred_protocol_versions(),
},
Self::Initialize => ClientLifecycleMode::Initialize,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct McpTasksConfig {
#[serde(default = "default_task_max_wait_secs")]
pub max_wait_secs: u64,
}
fn default_task_max_wait_secs() -> u64 {
DEFAULT_TASK_MAX_WAIT_SECS
}
impl Default for McpTasksConfig {
fn default() -> Self {
Self {
max_wait_secs: DEFAULT_TASK_MAX_WAIT_SECS,
}
}
}
impl McpTasksConfig {
fn max_wait(&self) -> Duration {
Duration::from_secs(self.max_wait_secs.clamp(1, MAX_TASK_MAX_WAIT_SECS))
}
}
#[cfg(test)]
mod tests {
use super::*;
use rmcp::model::{CallToolResult, InputRequiredResult, ProtocolVersion};
use rmcp::transport::AuthError;
use serde_json::json;
use std::borrow::Cow;
use std::path::PathBuf;
use super::auth::client_credentials_deadline;
use super::router::{
TASK_POLL_INTERVAL, TASK_POLL_INTERVAL_MAX, TASK_POLL_INTERVAL_MIN, input_required_error,
task_poll_interval,
};
use std::time::Instant;
fn tool(name: &'static str, description: &'static str) -> McpTool {
McpTool::new(
Cow::Borrowed(name),
Cow::Borrowed(description),
Arc::new(Map::from_iter([
("type".to_string(), json!("object")),
("properties".to_string(), json!({})),
])),
)
}
#[test]
fn sanitizes_and_bounds_tool_names() {
let provider = McpToolProvider::builder()
.server(McpServerConfig::stdio("GitHub-Prod", "server"))
.build()
.unwrap();
let short = provider
.local_tool_name("GitHub-Prod", "issues.get-by-id", None)
.unwrap();
assert_eq!(short, "mcp_github_prod_issues_get_by_id");
let long = provider
.local_tool_name("server", &"x".repeat(120), None)
.unwrap();
assert!(long.len() <= 64);
validate_function_name(&long).unwrap();
}
#[test]
fn include_exclude_filter_remote_tool_names() {
let mut server = McpServerConfig::stdio("repo", "server");
server.include.insert("allowed".to_string());
server.exclude.insert("blocked".to_string());
let provider = McpToolProvider::new(vec![server]).unwrap();
assert!(provider.includes_tool("repo", "allowed"));
assert!(!provider.includes_tool("repo", "other"));
assert!(!provider.includes_tool("repo", "blocked"));
}
#[tokio::test]
async fn refresh_tolerates_unreachable_servers_only_in_tolerant_mode() {
let provider = McpToolProvider::new(vec![McpServerConfig::stdio(
"down",
"anda_nonexistent_mcp_command_xyz",
)])
.unwrap();
provider.refresh_servers(true).await.unwrap();
assert!(provider.routes().is_empty());
let err = provider.refresh_servers(false).await.unwrap_err();
assert!(err.to_string().contains("down"));
}
#[tokio::test]
async fn add_server_rolls_back_when_initial_refresh_fails() {
let provider = McpToolProvider::new(Vec::new()).unwrap();
let err = provider
.add_server(McpServerConfig::stdio(
"down",
"anda_nonexistent_mcp_command_xyz",
))
.await
.unwrap_err();
assert!(!err.to_string().is_empty());
assert!(!provider.contains_server("down"));
assert!(provider.server_ids().is_empty());
assert!(provider.routes().is_empty());
}
#[cfg(unix)]
#[tokio::test]
async fn add_server_discovers_and_calls_tools_at_runtime() {
use std::os::unix::fs::PermissionsExt;
let script_path = std::env::temp_dir().join(format!(
"anda_fake_mcp_server_{}_{}",
std::process::id(),
"runtime_add"
));
let script = r#"#!/bin/sh
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32601,"message":"Method not found"}}\n' "$id"
;;
*"initialize"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":"2025-11-25","capabilities":{"tools":{"listChanged":true}},"serverInfo":{"name":"fake","version":"1.0.0"}}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"tools":[{"name":"echo","description":"Echoes input.","inputSchema":{"type":"object","properties":{"text":{"type":"string"}}}}]}}\n' "$id"
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"content":[{"type":"text","text":"ok"}],"isError":false}}\n' "$id"
;;
esac
done
"#;
std::fs::write(&script_path, script).unwrap();
let mut permissions = std::fs::metadata(&script_path).unwrap().permissions();
permissions.set_mode(0o700);
std::fs::set_permissions(&script_path, permissions).unwrap();
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"runtime",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
assert_eq!(provider.server_ids(), vec!["runtime".to_string()]);
let routes = provider.routes();
assert_eq!(routes.len(), 1);
assert_eq!(routes[0].name, "mcp_runtime_echo");
assert_eq!(routes[0].remote_name, "echo");
let output = provider
.call_route(
routes[0].clone(),
ToolInput::new("mcp_runtime_echo".to_string(), json!({"text": "hi"})),
)
.await
.unwrap();
assert_eq!(output.output["server_id"], "runtime");
assert_eq!(output.output["tool"], "echo");
assert_eq!(output.is_error, Some(false));
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn captures_server_metadata_into_a_tool_group() {
use std::os::unix::fs::PermissionsExt;
let script_path = std::env::temp_dir().join(format!(
"anda_fake_mcp_server_{}_{}",
std::process::id(),
"group_meta"
));
let script = r#"#!/bin/sh
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32601,"message":"Method not found"}}\n' "$id"
;;
*"initialize"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":"2025-11-25","capabilities":{"tools":{"listChanged":true}},"serverInfo":{"name":"fs","title":"Filesystem","version":"1.0.0"},"instructions":"Call list_dir before read_file."}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"tools":[{"name":"read_file","description":"Read a file.","inputSchema":{"type":"object","properties":{}}},{"name":"list_dir","description":"List a directory.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
;;
esac
done
"#;
std::fs::write(&script_path, script).unwrap();
let mut permissions = std::fs::metadata(&script_path).unwrap().permissions();
permissions.set_mode(0o700);
std::fs::set_permissions(&script_path, permissions).unwrap();
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"files",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
let groups = provider.tool_groups();
assert_eq!(groups.len(), 1);
let group = &groups[0];
assert_eq!(group.id, "mcp:files");
assert_eq!(group.title, "Filesystem");
assert_eq!(
group.instructions.as_deref(),
Some("Call list_dir before read_file.")
);
assert_eq!(
group.members,
vec![
"mcp_files_list_dir".to_string(),
"mcp_files_read_file".to_string()
]
);
let _ = std::fs::remove_file(script_path);
}
#[test]
fn server_meta_falls_back_to_server_id_without_peer_info() {
let meta = McpServerMeta::from_peer_info("files", None);
assert!(meta.title.is_none());
assert!(meta.instructions.is_none());
assert_eq!(meta.resolved_title("files"), "MCP server `files`");
assert_eq!(
meta.resolved_description("files"),
"Tools provided by MCP server `files`."
);
}
#[test]
fn tool_groups_is_empty_without_discovered_routes() {
let provider =
McpToolProvider::new(vec![McpServerConfig::stdio("files", "server")]).unwrap();
assert!(provider.tool_groups().is_empty());
}
#[test]
fn add_server_rejects_duplicate_and_normalized_colliding_ids() {
let provider =
McpToolProvider::new(vec![McpServerConfig::stdio("GitHub-Prod", "server")]).unwrap();
let err = provider
.insert_server(McpServerConfig::stdio("GitHub-Prod", "server"))
.unwrap_err();
assert!(err.to_string().contains("already exists"));
let err = provider
.insert_server(McpServerConfig::stdio("github_prod", "server"))
.unwrap_err();
assert!(err.to_string().contains("collides"));
}
#[test]
fn rejects_server_ids_that_collide_after_normalization() {
let err = McpToolProvider::new(vec![
McpServerConfig::stdio("GitHub-Prod", "server"),
McpServerConfig::stdio("github_prod", "server"),
])
.unwrap_err();
assert!(err.to_string().contains("collides"));
}
#[test]
fn converts_mcp_tools_to_function_definitions() {
let provider =
McpToolProvider::new(vec![McpServerConfig::stdio("repo", "server")]).unwrap();
let routes = provider
.routes_for_tools("repo", vec![tool("issues/get", "Fetch an issue")])
.unwrap();
assert_eq!(routes.len(), 1);
assert_eq!(routes[0].name, "mcp_repo_issues_get");
assert_eq!(routes[0].remote_name, "issues/get");
assert!(routes[0].definition.description.contains("Fetch an issue"));
assert_eq!(routes[0].definition.strict, Some(false));
}
#[test]
fn cross_server_local_name_collision_is_disambiguated() {
let provider = McpToolProvider::new(vec![
McpServerConfig::stdio("github", "server"),
McpServerConfig::stdio("github_prod", "server"),
])
.unwrap();
let a = provider
.routes_for_tools("github", vec![tool("prod_list_issues", "A")])
.unwrap();
let b = provider
.routes_for_tools("github_prod", vec![tool("list_issues", "B")])
.unwrap();
assert_eq!(a[0].name, "mcp_github_prod_list_issues");
assert_eq!(a[0].name, b[0].name, "the two servers' local names collide");
{
let mut index = provider.inner.index.write();
index.replace_server_routes("github", a);
index.replace_server_routes("github_prod", b);
}
let routes = provider.routes();
let github = routes
.iter()
.find(|route| route.server_id == "github")
.expect("github route retained");
let github_prod = routes
.iter()
.find(|route| route.server_id == "github_prod")
.expect("github_prod route retained");
assert_eq!(github.name, "mcp_github_prod_list_issues");
assert_ne!(github_prod.name, github.name);
assert_eq!(github_prod.remote_name, "list_issues");
assert_eq!(github_prod.definition.name, github_prod.name);
}
#[cfg(unix)]
#[tokio::test]
async fn concurrent_refresh_keeps_colliding_routes_stable() {
fn server(id: &str, tool: &str, delay: &str) -> McpServerConfig {
let script = r#"
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32601,"message":"Method not found"}}\n' "$id"
;;
*"initialize"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":"2025-11-25","capabilities":{"tools":{}},"serverInfo":{"name":"%s","version":"1.0.0"}}}\n' "$id" "$0"
;;
*"tools/list"*)
sleep "$2"
printf '{"jsonrpc":"2.0","id":%s,"result":{"tools":[{"name":"%s","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id" "$1"
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"content":[{"type":"text","text":"%s"}],"isError":false}}\n' "$id" "$1"
;;
esac
done
"#;
let mut config = McpServerConfig::stdio(id, "/bin/sh");
if let McpTransportConfig::Stdio(stdio) = &mut config.transport {
stdio.args = vec![
"-c".into(),
script.into(),
id.into(),
tool.into(),
delay.into(),
];
}
config
}
let mut expected = None;
for (github_delay, prod_delay) in [("0.2", "0"), ("0", "0.2")] {
let provider = McpToolProvider::new(vec![
server("github", "prod_list_issues", github_delay),
server("github_prod", "list_issues", prod_delay),
McpServerConfig::stdio("down", "anda_nonexistent_mcp_command_xyz"),
])
.unwrap();
provider.refresh_servers(true).await.unwrap();
let routes = provider.routes();
assert_eq!(routes.len(), 2);
let github = routes.iter().find(|r| r.server_id == "github").unwrap();
assert_eq!(github.name, "mcp_github_prod_list_issues");
assert_eq!(github.remote_name, "prod_list_issues");
let names: Vec<_> = routes
.iter()
.map(|r| (r.name.clone(), r.server_id.clone(), r.remote_name.clone()))
.collect();
assert_eq!(expected.get_or_insert_with(|| names.clone()), &names);
assert_eq!(provider.tool_groups().len(), 2);
for route in routes {
let output = provider
.call_route(route.clone(), ToolInput::new(route.name, json!({})))
.await
.unwrap();
assert_eq!(output.output["server_id"], route.server_id);
assert_eq!(output.output["tool"], route.remote_name);
assert_eq!(output.is_error, Some(false));
}
let err = provider.refresh_servers(false).await.unwrap_err();
assert!(err.to_string().contains("down"));
let refreshed: Vec<_> = provider
.routes()
.into_iter()
.map(|r| (r.name, r.server_id, r.remote_name))
.collect();
assert_eq!(refreshed, names);
}
}
fn http_auth_code_server(id: &str, client_id: Option<&str>) -> McpServerConfig {
let mut server = McpServerConfig::streamable_http(id, "https://example.com/mcp");
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport {
http.auth = Some(McpOAuthConfig::AuthorizationCode(
OAuthAuthorizationCodeConfig {
redirect_uri: "http://127.0.0.1:8080/callback".to_string(),
scopes: vec!["mcp:tools".to_string()],
client_name: Some("test".to_string()),
client_id: client_id.map(str::to_string),
},
));
}
server
}
#[test]
fn oauth_config_round_trips_through_serde() {
let server = http_auth_code_server("gh", None);
let json = serde_json::to_value(&server).unwrap();
assert_eq!(json["transport"]["auth"]["flow"], "authorization_code");
let parsed: McpServerConfig = serde_json::from_value(json).unwrap();
let McpTransportConfig::StreamableHttp(http) = &parsed.transport else {
panic!("expected streamable http transport");
};
match &http.auth {
Some(McpOAuthConfig::AuthorizationCode(ac)) => {
assert_eq!(ac.redirect_uri, "http://127.0.0.1:8080/callback");
assert_eq!(ac.scopes, vec!["mcp:tools".to_string()]);
assert!(ac.client_id.is_none());
}
other => panic!("unexpected auth config: {other:?}"),
}
}
#[test]
fn client_credentials_config_round_trips_through_serde() {
let mut server = McpServerConfig::streamable_http("svc", "https://example.com/mcp");
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport {
http.auth = Some(McpOAuthConfig::ClientCredentials(
OAuthClientCredentialsConfig {
client_id: "cid".to_string(),
client_secret: "secret".to_string(),
scopes: vec!["a".to_string()],
resource: Some("https://api.example.com".to_string()),
},
));
}
let json = serde_json::to_value(&server).unwrap();
assert_eq!(json["transport"]["auth"]["flow"], "client_credentials");
let parsed: McpServerConfig = serde_json::from_value(json).unwrap();
assert!(parsed.validate().is_ok());
}
#[test]
fn validate_rejects_bearer_token_combined_with_oauth() {
let mut server = http_auth_code_server("gh", None);
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport {
http.bearer_token = Some("tok".to_string());
}
let err = server.validate().unwrap_err().to_string();
assert!(err.contains("cannot set both"), "{err}");
}
#[test]
fn debug_output_redacts_secrets() {
let cfg = OAuthClientCredentialsConfig {
client_id: "cid".to_string(),
client_secret: "super-secret-value".to_string(),
scopes: vec![],
resource: None,
};
let rendered = format!("{cfg:?}");
assert!(!rendered.contains("super-secret-value"), "{rendered}");
assert!(rendered.contains("cid"), "{rendered}");
assert!(rendered.contains("[REDACTED]"), "{rendered}");
let mut server = McpServerConfig::streamable_http("svc", "https://example.com/mcp");
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport {
http.bearer_token = Some("super-secret-token".to_string());
}
let rendered = format!("{server:?}");
assert!(!rendered.contains("super-secret-token"), "{rendered}");
assert!(rendered.contains("[REDACTED]"), "{rendered}");
let mut server = McpServerConfig::streamable_http("svc", "https://example.com/mcp");
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport {
http.auth = Some(McpOAuthConfig::ClientCredentials(cfg));
}
let rendered = format!("{server:?}");
assert!(!rendered.contains("super-secret-value"), "{rendered}");
let mut server = McpServerConfig::stdio("files", "mcp-files");
if let McpTransportConfig::Stdio(stdio) = &mut server.transport {
stdio
.env
.insert("GITHUB_TOKEN".to_string(), "ghp_live_value".to_string());
}
let rendered = format!("{server:?}");
assert!(!rendered.contains("ghp_live_value"), "{rendered}");
assert!(rendered.contains("GITHUB_TOKEN"), "{rendered}");
assert!(rendered.contains("[REDACTED]"), "{rendered}");
}
#[test]
fn client_credentials_deadlines_preserve_short_token_lifetime_and_do_not_overflow() {
let now = Instant::now();
let after = |secs| {
client_credentials_deadline(now, Duration::from_secs(secs))
.unwrap()
.duration_since(now)
};
assert_eq!(after(3_600), Duration::from_secs(3_480));
assert_eq!(after(240), Duration::from_secs(120));
assert_eq!(after(120), Duration::from_secs(60));
assert_eq!(after(60), Duration::from_secs(29));
assert_eq!(after(30), Duration::ZERO);
assert!(
client_credentials_deadline(now, Duration::MAX).is_none(),
"an untrusted, unrepresentable expires_in must not panic"
);
}
#[test]
fn validate_rejects_incomplete_oauth_configs() {
let mut server = http_auth_code_server("gh", None);
if let McpTransportConfig::StreamableHttp(http) = &mut server.transport
&& let Some(McpOAuthConfig::AuthorizationCode(ac)) = &mut http.auth
{
ac.redirect_uri = " ".to_string();
}
assert!(server.validate().is_err());
let mut creds = McpServerConfig::streamable_http("svc", "https://example.com/mcp");
if let McpTransportConfig::StreamableHttp(http) = &mut creds.transport {
http.auth = Some(McpOAuthConfig::ClientCredentials(
OAuthClientCredentialsConfig {
client_id: "cid".to_string(),
client_secret: String::new(),
scopes: vec![],
resource: None,
},
));
}
assert!(creds.validate().is_err());
}
#[tokio::test]
async fn begin_authorization_rejects_non_oauth_servers() {
let provider = McpToolProvider::new(vec![
McpServerConfig::stdio("cli", "server"),
McpServerConfig::streamable_http("plain", "https://example.com/mcp"),
])
.unwrap();
let err = provider.begin_authorization("cli").await.unwrap_err();
assert!(err.to_string().contains("does not use the HTTP transport"));
let err = provider.begin_authorization("plain").await.unwrap_err();
assert!(
err.to_string()
.contains("not configured for the OAuth authorization_code flow")
);
}
#[tokio::test]
async fn authorization_code_without_credentials_reports_auth_required() {
let provider = McpToolProvider::new(vec![http_auth_code_server("gh", None)]).unwrap();
let err = provider.refresh_server("gh").await.unwrap_err();
let required = err
.downcast_ref::<McpAuthorizationRequired>()
.expect("expected McpAuthorizationRequired");
assert_eq!(required.server_id, "gh");
}
#[tokio::test]
async fn register_server_registers_without_connecting() {
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.register_server(http_auth_code_server("gh", None))
.unwrap();
assert!(provider.contains_server("gh"));
assert!(provider.routes().is_empty());
assert!(
provider
.register_server(http_auth_code_server("gh", None))
.is_err()
);
assert!(provider.remove_server("gh"));
assert!(!provider.contains_server("gh"));
assert!(!provider.remove_server("gh"));
}
#[tokio::test]
async fn complete_and_cancel_authorization_without_pending_state() {
let provider = McpToolProvider::new(vec![http_auth_code_server("gh", None)]).unwrap();
assert!(!provider.cancel_authorization("gh"));
let err = provider
.complete_authorization("gh", "http://127.0.0.1:8080/callback?code=x&state=y")
.await
.unwrap_err();
assert!(err.to_string().contains("no pending OAuth authorization"));
}
#[cfg(unix)]
#[tokio::test]
async fn disconnect_drops_the_session_but_keeps_the_server_and_routes() {
let script_path = write_fake_server("disconnect", DISCOVER_SERVER);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"stateless",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
assert!(
provider
.inner
.index
.read()
.sessions
.contains_key("stateless")
);
assert!(!provider.disconnect_server("missing").await);
assert!(provider.disconnect_server("stateless").await);
assert!(!provider.disconnect_server("stateless").await);
assert!(provider.inner.index.read().sessions.is_empty());
assert!(provider.contains_server("stateless"));
let route = provider.routes().remove(0);
let output = provider
.call_route(
route,
ToolInput::new("mcp_stateless_echo".to_string(), json!({})),
)
.await
.unwrap();
assert_eq!(output.is_error, Some(false));
assert!(
provider
.inner
.index
.read()
.sessions
.contains_key("stateless")
);
let _ = std::fs::remove_file(script_path);
}
#[tokio::test]
async fn clear_credentials_forces_reauthorization() {
let store = Arc::new(InMemoryMcpCredentialStore::new());
let provider = McpToolProvider::builder()
.server(http_auth_code_server("gh", None))
.credential_store(store.clone())
.build()
.unwrap();
let creds = StoredCredentials::new("client".to_string(), None, vec!["a".to_string()], None);
store.save("gh", creds).await.unwrap();
provider.clear_credentials("gh").await.unwrap();
assert!(store.load("gh").await.unwrap().is_none());
let err = provider.refresh_server("gh").await.unwrap_err();
assert!(err.downcast_ref::<McpAuthorizationRequired>().is_some());
}
#[tokio::test]
async fn in_memory_credential_store_round_trip() {
let store = InMemoryMcpCredentialStore::new();
assert!(store.load("gh").await.unwrap().is_none());
let creds = StoredCredentials::new("client".to_string(), None, vec!["a".to_string()], None);
store.save("gh", creds).await.unwrap();
let loaded = store.load("gh").await.unwrap().expect("stored");
assert_eq!(loaded.client_id, "client");
store.clear("gh").await.unwrap();
assert!(store.load("gh").await.unwrap().is_none());
}
#[test]
fn converts_mcp_call_result_to_audited_output() {
let route = McpToolRoute {
tool: tool("echo", "Echo"),
server_generation: 0,
catalog_revision: 0,
name: "mcp_repo_echo".to_string(),
server_id: "repo".to_string(),
remote_name: "echo".to_string(),
definition: FunctionDefinition::default(),
};
let result = CallToolResult::structured(json!({"ok": true}));
let output = mcp_result_to_tool_output(&route, result, &McpLimits::default());
assert_eq!(output.is_error, Some(false));
assert_eq!(output.usage.requests, 1);
assert_eq!(output.output["server_id"], "repo");
assert_eq!(output.output["tool"], "echo");
assert_eq!(output.output["structured_content"], json!({"ok": true}));
}
#[test]
fn lifecycles_map_to_the_matching_rmcp_modes() {
match McpLifecycle::Auto.into_mode() {
ClientLifecycleMode::Auto {
preferred_versions,
legacy_version,
} => {
assert_eq!(preferred_versions[0], ProtocolVersion::V_2026_07_28);
assert_eq!(legacy_version, Some(ProtocolVersion::V_2025_11_25));
}
other => panic!("unexpected lifecycle mode: {other:?}"),
}
match McpLifecycle::Discover.into_mode() {
ClientLifecycleMode::Discover { preferred_versions } => {
assert!(preferred_versions.contains(&ProtocolVersion::V_2026_07_28));
}
other => panic!("unexpected lifecycle mode: {other:?}"),
}
assert_eq!(
McpLifecycle::Initialize.into_mode(),
ClientLifecycleMode::Initialize
);
assert_eq!(
AndaMcpClient::new(Arc::new(AtomicBool::new(false)), false)
.info
.protocol_version,
ProtocolVersion::V_2025_11_25
);
}
#[test]
fn declares_the_tasks_extension_only_when_configured() {
let dirty = Arc::new(AtomicBool::new(false));
assert!(
AndaMcpClient::new(dirty.clone(), false)
.info
.capabilities
.extensions
.is_none()
);
let capabilities = AndaMcpClient::new(dirty, true).info.capabilities;
assert!(capabilities.supports_tasks());
}
#[test]
fn tool_subscriptions_are_required_only_from_2026_07_28() {
let peer_info = |version: &str, list_changed: bool| -> ServerPeerInfo {
serde_json::from_value(json!({
"protocolVersion": version,
"capabilities": {"tools": {"listChanged": list_changed}},
}))
.unwrap()
};
assert!(needs_tool_subscription(Some(&peer_info(
"2026-07-28",
true
))));
assert!(!needs_tool_subscription(Some(&peer_info(
"2025-11-25",
true
))));
assert!(!needs_tool_subscription(Some(&peer_info(
"2026-07-28",
false
))));
assert!(!needs_tool_subscription(None));
}
#[test]
fn task_poll_intervals_clamp_untrusted_server_hints() {
assert_eq!(task_poll_interval(None), TASK_POLL_INTERVAL);
assert_eq!(task_poll_interval(Some(2_000)), Duration::from_secs(2));
assert_eq!(task_poll_interval(Some(0)), TASK_POLL_INTERVAL_MIN);
assert_eq!(task_poll_interval(Some(u64::MAX)), TASK_POLL_INTERVAL_MAX);
}
#[test]
fn task_max_wait_is_clamped_so_the_deadline_cannot_overflow() {
assert_eq!(
McpTasksConfig::default().max_wait(),
Duration::from_secs(DEFAULT_TASK_MAX_WAIT_SECS)
);
assert_eq!(
McpTasksConfig { max_wait_secs: 0 }.max_wait(),
Duration::from_secs(1)
);
let huge = McpTasksConfig {
max_wait_secs: u64::MAX,
};
assert_eq!(huge.max_wait(), Duration::from_secs(MAX_TASK_MAX_WAIT_SECS));
let _ = Instant::now() + huge.max_wait();
}
#[tokio::test]
async fn revoked_authorization_is_reported_as_authorization_required() {
let config = http_auth_code_server("gh", None);
let err: BoxError = AuthError::AuthorizationRequired.into();
assert!(is_authorization_error(err.as_ref()));
let mapped = authorization_required_hint(&config, err);
let required = mapped
.downcast_ref::<McpAuthorizationRequired>()
.expect("expected McpAuthorizationRequired");
assert_eq!(required.server_id, "gh");
let other: BoxError = "transport closed".into();
assert!(
authorization_required_hint(&config, other)
.downcast_ref::<McpAuthorizationRequired>()
.is_none()
);
let stdio = McpServerConfig::stdio("cli", "server");
let err: BoxError = McpAuthorizationRequired {
server_id: "other".to_string(),
}
.into();
assert_eq!(
authorization_required_hint(&stdio, err)
.downcast_ref::<McpAuthorizationRequired>()
.map(|required| required.server_id.clone()),
Some("other".to_string())
);
}
#[test]
fn authorization_failures_are_not_treated_as_lifecycle_failures() {
let err: BoxError = McpAuthorizationRequired {
server_id: "gh".to_string(),
}
.into();
assert!(is_authorization_error(err.as_ref()));
let err: BoxError = "transport closed".into();
assert!(!is_authorization_error(err.as_ref()));
}
#[test]
fn unsupported_input_rounds_become_tool_level_errors() {
let route = McpToolRoute {
tool: tool("echo", "Echo"),
server_generation: 0,
catalog_revision: 0,
name: "mcp_repo_echo".to_string(),
server_id: "repo".to_string(),
remote_name: "echo".to_string(),
definition: FunctionDefinition::default(),
};
let result: InputRequiredResult = serde_json::from_value(json!({
"resultType": "input_required",
"inputRequests": {"pick_root": {"method": "roots/list"}},
}))
.unwrap();
let output = mcp_result_to_tool_output(
&route,
input_required_error(&route, &result),
&McpLimits::default(),
);
assert_eq!(output.is_error, Some(true));
let rendered = output.output.to_string();
assert!(rendered.contains("pick_root"), "{rendered}");
assert!(rendered.contains("does not provide"), "{rendered}");
}
#[test]
fn server_config_defaults_to_the_auto_lifecycle_without_tasks() {
let parsed: McpServerConfig = serde_json::from_value(json!({
"id": "files",
"transport": {"type": "stdio", "command": "server"},
}))
.unwrap();
assert_eq!(parsed.lifecycle, McpLifecycle::Auto);
assert!(parsed.tasks.is_none());
let mut server = McpServerConfig::stdio("files", "server");
server.lifecycle = McpLifecycle::Discover;
server.tasks = Some(McpTasksConfig::default());
let json = serde_json::to_value(&server).unwrap();
assert_eq!(json["lifecycle"], "discover");
assert_eq!(json["tasks"]["max_wait_secs"], 300);
let parsed: McpServerConfig = serde_json::from_value(json).unwrap();
assert_eq!(parsed.lifecycle, McpLifecycle::Discover);
assert_eq!(
parsed.tasks.map(|tasks| tasks.max_wait()),
Some(Duration::from_secs(300))
);
}
#[cfg(unix)]
fn write_fake_server(name: &str, script: &str) -> PathBuf {
use std::os::unix::fs::PermissionsExt;
let path = std::env::temp_dir().join(format!(
"anda_fake_mcp_server_{}_{name}",
std::process::id()
));
std::fs::write(&path, script).unwrap();
let mut permissions = std::fs::metadata(&path).unwrap().permissions();
permissions.set_mode(0o700);
std::fs::set_permissions(&path, permissions).unwrap();
path
}
#[cfg(unix)]
const DISCOVER_SERVER: &str = r#"#!/bin/sh
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","supportedVersions":["2026-07-28"],"capabilities":{"tools":{"listChanged":false}},"instructions":"Call list_dir first.","ttlMs":0,"cacheScope":"private","_meta":{"io.modelcontextprotocol/serverInfo":{"name":"fs","title":"Stateless Files","version":"1.0.0"}}}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","ttlMs":0,"cacheScope":"private","tools":[{"name":"echo","description":"Echoes input.","inputSchema":{"type":"object","properties":{"text":{"type":"string"}}}}]}}\n' "$id"
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","content":[{"type":"text","text":"ok"}],"isError":false}}\n' "$id"
;;
esac
done
"#;
#[cfg(unix)]
#[tokio::test]
async fn negotiates_the_2026_lifecycle_and_calls_tools() {
let script_path = write_fake_server("discover", DISCOVER_SERVER);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"stateless",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
let routes = provider.routes();
assert_eq!(routes.len(), 1);
assert_eq!(routes[0].name, "mcp_stateless_echo");
let groups = provider.tool_groups();
assert_eq!(groups[0].title, "Stateless Files");
assert_eq!(
groups[0].instructions.as_deref(),
Some("Call list_dir first.")
);
let output = provider
.call_route(
routes[0].clone(),
ToolInput::new("mcp_stateless_echo".to_string(), json!({"text": "hi"})),
)
.await
.unwrap();
assert_eq!(output.is_error, Some(false));
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn falls_back_to_the_legacy_handshake_when_discovery_is_refused() {
let script_path = write_fake_server(
"legacy_only",
r#"#!/bin/sh
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
exit 1
;;
*"initialize"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":"2025-11-25","capabilities":{"tools":{"listChanged":true}},"serverInfo":{"name":"legacy","version":"1.0.0"}}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"tools":[{"name":"echo","description":"Echoes input.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
;;
esac
done
"#,
);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"legacy",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
assert_eq!(provider.routes().len(), 1);
let mut pinned =
McpServerConfig::stdio("pinned", script_path.to_string_lossy().to_string());
pinned.lifecycle = McpLifecycle::Initialize;
provider.add_server(pinned).await.unwrap();
assert_eq!(provider.routes().len(), 2);
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn input_required_results_surface_as_failed_tool_calls() {
let script_path = write_fake_server(
"mrtr",
r#"#!/bin/sh
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","supportedVersions":["2026-07-28"],"capabilities":{"tools":{"listChanged":false}},"ttlMs":0,"cacheScope":"private"}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","ttlMs":0,"cacheScope":"private","tools":[{"name":"ask","description":"Asks first.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"input_required","inputRequests":{"pick_root":{"method":"roots/list"}},"requestState":"opaque"}}\n' "$id"
;;
esac
done
"#,
);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"asker",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
let route = provider.routes().remove(0);
let output = provider
.call_route(
route,
ToolInput::new("mcp_asker_ask".to_string(), json!({})),
)
.await
.unwrap();
assert_eq!(output.is_error, Some(true));
assert!(output.output.to_string().contains("pick_root"));
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
const TASK_SERVER: &str = r#"#!/bin/sh
polls=0
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","supportedVersions":["2026-07-28"],"capabilities":{"tools":{"listChanged":false},"extensions":{"io.modelcontextprotocol/tasks":{}}},"ttlMs":0,"cacheScope":"private"}}\n' "$id"
;;
*"tools/list"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","ttlMs":0,"cacheScope":"private","tools":[{"name":"slow","description":"Takes a while.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"task","taskId":"t1","status":"working","createdAt":"2026-07-28T00:00:00Z","lastUpdatedAt":"2026-07-28T00:00:00Z","ttlMs":null,"pollIntervalMs":10}}\n' "$id"
;;
*"tasks/get"*)
polls=$((polls+1))
if [ "$polls" -le 1 ]; then
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","taskId":"t1","status":"working","createdAt":"2026-07-28T00:00:00Z","lastUpdatedAt":"2026-07-28T00:00:01Z","ttlMs":null,"pollIntervalMs":10}}\n' "$id"
else
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","taskId":"t1","status":"completed","createdAt":"2026-07-28T00:00:00Z","lastUpdatedAt":"2026-07-28T00:00:02Z","ttlMs":null,"result":{"resultType":"complete","content":[{"type":"text","text":"done"}],"structuredContent":{"ok":true},"isError":false}}}\n' "$id"
fi
;;
*"tasks/cancel"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete"}}\n' "$id"
;;
esac
done
"#;
#[cfg(unix)]
#[tokio::test]
async fn polls_tasks_to_completion_when_the_extension_is_enabled() {
let script_path = write_fake_server("tasks", TASK_SERVER);
let mut server =
McpServerConfig::stdio("worker", script_path.to_string_lossy().to_string());
server.tasks = Some(McpTasksConfig::default());
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider.add_server(server).await.unwrap();
let route = provider.routes().remove(0);
let output = provider
.call_route(
route,
ToolInput::new("mcp_worker_slow".to_string(), json!({})),
)
.await
.unwrap();
assert_eq!(output.is_error, Some(false));
assert_eq!(output.output["structured_content"], json!({"ok": true}));
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn rejects_task_handles_when_the_extension_is_not_enabled() {
let script_path = write_fake_server("tasks_undeclared", TASK_SERVER);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"worker",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
let route = provider.routes().remove(0);
let err = provider
.call_route(
route,
ToolInput::new("mcp_worker_slow".to_string(), json!({})),
)
.await
.unwrap_err()
.to_string();
assert!(err.contains("tasks extension is not enabled"), "{err}");
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn task_waiting_is_bounded_by_max_wait() {
let script_path = write_fake_server(
"tasks_slow",
&TASK_SERVER.replace("\"pollIntervalMs\":10", "\"pollIntervalMs\":9000"),
);
let mut server =
McpServerConfig::stdio("worker", script_path.to_string_lossy().to_string());
server.tasks = Some(McpTasksConfig { max_wait_secs: 1 });
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider.add_server(server).await.unwrap();
let route = provider.routes().remove(0);
let err = provider
.call_route(
route,
ToolInput::new("mcp_worker_slow".to_string(), json!({})),
)
.await
.unwrap_err()
.to_string();
assert!(err.contains("did not finish within 1s"), "{err}");
let _ = std::fs::remove_file(script_path);
}
#[cfg(unix)]
#[tokio::test]
async fn tool_list_changes_arrive_on_the_subscription_stream() {
let script_path = write_fake_server(
"subscriptions",
r#"#!/bin/sh
count=0
sub=0
while IFS= read -r line; do
id=$(printf '%s\n' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$line" in
*"server/discover"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","supportedVersions":["2026-07-28"],"capabilities":{"tools":{"listChanged":true}},"ttlMs":0,"cacheScope":"private"}}\n' "$id"
;;
*"subscriptions/listen"*)
sub=$id
printf '{"jsonrpc":"2.0","method":"notifications/subscriptions/acknowledged","params":{"_meta":{"io.modelcontextprotocol/subscriptionId":%s},"notifications":{"toolsListChanged":true}}}\n' "$sub"
;;
*"tools/list"*)
count=$((count+1))
if [ "$count" -le 1 ]; then
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","ttlMs":0,"cacheScope":"private","tools":[{"name":"echo","description":"Echoes.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
else
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","ttlMs":0,"cacheScope":"private","tools":[{"name":"echo","description":"Echoes.","inputSchema":{"type":"object","properties":{}}},{"name":"ping","description":"Pings.","inputSchema":{"type":"object","properties":{}}}]}}\n' "$id"
fi
;;
*"tools/call"*)
printf '{"jsonrpc":"2.0","id":%s,"result":{"resultType":"complete","content":[{"type":"text","text":"ok"}],"isError":false}}\n' "$id"
printf '{"jsonrpc":"2.0","method":"notifications/tools/list_changed","params":{"_meta":{"io.modelcontextprotocol/subscriptionId":%s}}}\n' "$sub"
;;
esac
done
"#,
);
let provider = McpToolProvider::new(Vec::new()).unwrap();
provider
.add_server(McpServerConfig::stdio(
"live",
script_path.to_string_lossy().to_string(),
))
.await
.unwrap();
assert_eq!(provider.routes().len(), 1);
let route = provider.routes().remove(0);
for _ in 0..20 {
provider
.call_route(route.clone(), ToolInput::new(route.name.clone(), json!({})))
.await
.unwrap();
if provider.routes().len() == 2 {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
let names: Vec<String> = provider
.routes()
.into_iter()
.map(|route| route.name)
.collect();
assert_eq!(names, vec!["mcp_live_echo", "mcp_live_ping"]);
let _ = std::fs::remove_file(script_path);
}
}
#[cfg(test)]
#[path = "mcp/regression_tests.rs"]
mod regression_tests;