use std::{
collections::{BTreeMap, HashSet},
future::Future,
pin::Pin,
sync::Arc,
};
use rho_sdk::tool::Tool;
use super::sdk_registry::ToolBundle;
use config::{McpConfig, McpServerConfig};
pub(crate) mod catalog;
pub(crate) mod client;
pub(crate) mod config;
pub(crate) mod definition;
pub(crate) mod elicitation;
pub(crate) mod elicitation_form;
pub(crate) mod inflight;
pub(crate) mod oauth;
pub(crate) mod progress;
pub(crate) mod report;
pub(crate) mod result;
pub(crate) mod roots;
pub(crate) mod sampling;
pub(crate) mod session;
pub(crate) mod tool;
pub(crate) mod validate;
pub(crate) use catalog::{
McpCatalog, McpCatalogError, McpCompletionSupport, McpResource, McpResourceContent,
};
pub(crate) use elicitation::McpElicitationSupport;
pub(crate) use oauth::McpAuthorizationMode;
pub(crate) use report::{
McpLoadMode, McpServerReport, McpServerStatus, McpSessionReport, McpToolReport,
McpTransportSummary,
};
pub(crate) use roots::McpRoots;
pub(crate) use sampling::{McpSamplingBridge, McpSamplingModel};
pub(crate) use validate::{
parse_remote_url, validate_environment_header_names, validate_identity,
validate_literal_headers, validate_oauth_client, validate_stdio_environment,
};
use definition::McpToolDefinition;
use session::{ConnectResult, ConnectedServer, McpSession, SessionMaintenance};
use tool::{namespaced_tool_name, McpTool, McpToolSlot};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum McpSessionPlan {
Connect,
Inventory(McpLoadMode),
}
#[derive(Clone, Debug)]
pub(crate) struct McpSessionOptions {
pub(crate) max_output_bytes: usize,
pub(crate) roots: McpRoots,
pub(crate) authorization: McpAuthorizationMode,
services: session::McpSessionServices,
}
impl McpSessionOptions {
pub(crate) fn new(
max_output_bytes: usize,
roots: McpRoots,
authorization: McpAuthorizationMode,
) -> Self {
Self {
max_output_bytes: max_output_bytes.max(1),
roots,
authorization,
services: session::McpSessionServices {
elicitation: McpElicitationSupport::Unavailable,
sampling: None,
},
}
}
pub(crate) fn with_elicitation(mut self, support: McpElicitationSupport) -> Self {
self.services.elicitation = support;
self
}
pub(crate) fn with_sampling(mut self, bridge: McpSamplingBridge) -> Self {
self.services.sampling = Some(bridge);
self
}
}
pub(crate) struct McpConnectOutcome {
pub(crate) report: McpSessionReport,
pub(crate) bundle: Option<McpBundle>,
pub(crate) catalog: McpCatalog,
}
impl McpConnectOutcome {
pub(crate) async fn run(
plan: McpSessionPlan,
config: &McpConfig,
options: McpSessionOptions,
) -> Self {
match plan {
McpSessionPlan::Connect => McpBundle::connect(config, options).await,
McpSessionPlan::Inventory(mode) => Self {
report: McpSessionReport::from_config_unloaded(config, mode),
bundle: None,
catalog: McpCatalog::default(),
},
}
}
}
pub(crate) struct McpBundle {
tools: Vec<Arc<dyn Tool>>,
sessions: tokio::sync::Mutex<Vec<McpSession>>,
maintenance: tokio::sync::Mutex<Vec<tokio::task::JoinHandle<()>>>,
}
impl McpBundle {
pub(crate) async fn connect(
config: &McpConfig,
options: McpSessionOptions,
) -> McpConnectOutcome {
let mut servers = Vec::with_capacity(config.servers.len() + config.invalid_servers.len());
for invalid in &config.invalid_servers {
tracing::warn!(
server = %invalid.identity,
error = %invalid.error,
"ignoring invalid MCP server configuration"
);
servers.push(McpServerReport::invalid(
invalid.identity.clone(),
invalid.error.clone(),
));
}
for (identity, server) in &config.servers {
if !server.enabled {
servers.push(McpServerReport::disabled(identity.clone(), server));
}
}
if !config.has_enabled_servers() {
servers.sort_by(|left, right| left.identity.cmp(&right.identity));
return McpConnectOutcome {
report: McpSessionReport {
mode: McpLoadMode::Native,
servers,
},
bundle: None,
catalog: McpCatalog::default(),
};
}
#[cfg(test)]
MCP_RUNTIME_CONSTRUCTIONS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let connect_jobs = config
.servers
.iter()
.filter(|(_, server)| server.enabled)
.map(|(identity, server)| {
let identity = identity.clone();
let server = server.clone();
let roots = options.roots.clone();
let services = options.services.clone();
let authorization = options.authorization;
async move {
let transport = McpTransportSummary::from_server(&server);
let result = session::connect_server_bounded(
&identity,
&server,
&roots,
&services,
authorization,
)
.await;
(identity, server, transport, result)
}
});
let connect_results = futures_util::future::join_all(connect_jobs).await;
let mut bundle = McpBundleBuilder::new(options.max_output_bytes);
for (identity, server, transport, result) in connect_results {
let connected = match result {
ConnectResult::Ready(connected) => connected,
ConnectResult::Failed { error } => {
tracing::warn!(server = %identity, error = %error, "MCP server failed to initialize");
servers.push(McpServerReport::failed(
identity,
transport,
error.to_string(),
));
continue;
}
ConnectResult::TimedOut => {
tracing::warn!(
server = %identity,
limit_seconds = session::MCP_SERVER_STARTUP_BUDGET.as_secs(),
"MCP server exceeded its startup budget"
);
servers.push(McpServerReport::timed_out(
identity,
transport,
session::MCP_SERVER_STARTUP_BUDGET.as_secs(),
));
continue;
}
};
servers.push(
bundle
.register(identity, server, transport, *connected)
.await,
);
}
servers.sort_by(|left, right| left.identity.cmp(&right.identity));
let catalog = bundle.catalog.clone();
McpConnectOutcome {
report: McpSessionReport {
mode: McpLoadMode::Native,
servers,
},
bundle: bundle.build(),
catalog,
}
}
pub(crate) async fn close(&self) {
for task in std::mem::take(&mut *self.maintenance.lock().await) {
task.abort();
}
let sessions = {
let mut guard = self.sessions.lock().await;
std::mem::take(&mut *guard)
};
let close_jobs = sessions.into_iter().map(session::close_session);
futures_util::future::join_all(close_jobs).await;
}
}
impl ToolBundle for McpBundle {
fn tools(&self) -> &[Arc<dyn Tool>] {
&self.tools
}
fn shutdown(&self) -> Pin<Box<dyn Future<Output = ()> + Send + '_>> {
Box::pin(self.close())
}
}
struct McpBundleBuilder {
max_output_bytes: usize,
tools: Vec<Arc<dyn Tool>>,
sessions: Vec<McpSession>,
maintenance: Vec<tokio::task::JoinHandle<()>>,
registered_names: HashSet<String>,
catalog: McpCatalog,
}
impl McpBundleBuilder {
fn new(max_output_bytes: usize) -> Self {
Self {
max_output_bytes,
tools: Vec::new(),
sessions: Vec::new(),
maintenance: Vec::new(),
registered_names: HashSet::new(),
catalog: McpCatalog::default(),
}
}
async fn register(
&mut self,
identity: String,
server: McpServerConfig,
transport: McpTransportSummary,
connected: ConnectedServer,
) -> McpServerReport {
let ConnectedServer {
session,
discovered,
instructions,
progress,
calls,
events,
offers,
} = connected;
let mut exported = Vec::new();
let mut slots = BTreeMap::new();
let mut filtered_out_count = 0usize;
let mut collision_skipped_count = 0usize;
for remote in discovered {
let remote_name = remote.name.to_string();
if !server.tools.includes(&remote_name) {
filtered_out_count += 1;
continue;
}
let name = namespaced_tool_name(&identity, &remote_name);
if !self.registered_names.insert(name.clone()) {
tracing::warn!(server = %identity, tool = %remote_name, exported = %name, "MCP tool name collision; ignoring tool");
collision_skipped_count += 1;
continue;
}
let slot = Arc::new(McpToolSlot::new(McpToolDefinition::from_remote(
&identity,
&remote_name,
&remote,
)));
slots.insert(remote_name.clone(), Arc::clone(&slot));
self.tools.push(Arc::new(McpTool {
slot,
identity: identity.clone(),
remote_name: remote_name.clone(),
peer: session.peer().clone(),
progress: progress.clone(),
calls: calls.clone(),
transport: server.transport.clone(),
max_output_bytes: self.max_output_bytes,
}));
exported.push(McpToolReport {
remote_name,
exported_name: name,
});
}
let catalog = self
.catalog
.register(identity.clone(), session.peer().clone(), offers);
session::list_offers(&catalog, offers).await;
let live = report::McpLiveServerState::default();
self.maintenance
.push(tokio::spawn(session::maintain_session(
SessionMaintenance {
identity: identity.clone(),
peer: session.peer().clone(),
server,
slots,
live: live.clone(),
events,
catalog,
offers,
},
)));
self.sessions.push(session);
McpServerReport::connected(report::ConnectedServerReport {
identity,
transport,
tools: exported,
instructions,
live,
filtered_out_count,
collision_skipped_count,
})
}
fn build(self) -> Option<McpBundle> {
if self.sessions.is_empty() {
return None;
}
Some(McpBundle {
tools: self.tools,
sessions: tokio::sync::Mutex::new(self.sessions),
maintenance: tokio::sync::Mutex::new(self.maintenance),
})
}
}
#[cfg(test)]
static MCP_RUNTIME_CONSTRUCTIONS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
#[path = "mcp_tests.rs"]
mod tests;