use rmcp::ServiceExt;
use rmcp::service::{RoleClient, RunningService};
use rmcp::transport::TokioChildProcess;
use rmcp::transport::streamable_http_client::{
StreamableHttpClientTransport, StreamableHttpClientTransportConfig,
};
use serde_json::Value;
use tokio::process::Command;
use super::{EffectOverrides, McpTool, effect_for};
pub struct McpServer {
service: RunningService<RoleClient, ()>,
tools: Vec<McpTool>,
}
impl McpServer {
pub async fn connect(command: Command, overrides: &EffectOverrides) -> Result<Self, McpError> {
let transport = TokioChildProcess::new(command).map_err(McpError::Spawn)?;
let service = ().serve(transport).await.map_err(|e| McpError::Initialize(Box::new(e)))?;
Self::from_service(service, overrides).await
}
pub async fn connect_http(
url: &str,
bearer_token: Option<&str>,
overrides: &EffectOverrides,
) -> Result<Self, McpError> {
let config = StreamableHttpClientTransportConfig::with_uri(url.to_owned());
let config = match bearer_token {
Some(token) => config.auth_header(token.to_owned()),
None => config,
};
let transport = StreamableHttpClientTransport::from_config(config);
let service = ().serve(transport).await.map_err(|e| McpError::Initialize(Box::new(e)))?;
Self::from_service(service, overrides).await
}
async fn from_service(
service: RunningService<RoleClient, ()>,
overrides: &EffectOverrides,
) -> Result<Self, McpError> {
let listed = service
.list_all_tools()
.await
.map_err(McpError::ListTools)?;
let peer = service.peer().clone();
let tools = listed
.into_iter()
.map(|tool| {
let name = tool.name.to_string();
let description = tool.description.map(|d| d.to_string()).unwrap_or_default();
let input_schema = Value::Object((*tool.input_schema).clone());
let effect = effect_for(&name, tool.annotations.as_ref(), overrides);
McpTool::new(peer.clone(), name, description, input_schema, effect)
})
.collect();
Ok(Self { service, tools })
}
pub fn tools(&self) -> &[McpTool] {
&self.tools
}
pub fn take_tools(&mut self) -> Vec<McpTool> {
std::mem::take(&mut self.tools)
}
pub async fn close(mut self) -> Result<(), McpError> {
self.service.close().await.map_err(McpError::Shutdown)?;
Ok(())
}
}
#[derive(Debug, thiserror::Error)]
pub enum McpError {
#[error("failed to spawn the MCP server process")]
Spawn(#[source] std::io::Error),
#[error("failed to initialize the MCP session")]
Initialize(#[source] Box<rmcp::service::ClientInitializeError>),
#[error("failed to list the MCP server's tools")]
ListTools(#[source] rmcp::service::ServiceError),
#[error("failed to shut the MCP session down cleanly")]
Shutdown(#[source] tokio::task::JoinError),
}