use crate::{
cancellation::AgentCancellation,
config::{McpServerConfig, McpStdioServerConfig},
mcp::{
McpError, McpResult,
http::HttpConnection,
jsonrpc::RequestId,
protocol::{
CallToolParams, CallToolResult, InitializeRequestParams, InitializeResult,
ListToolsParams, ListToolsResult, METHOD_INITIALIZE, METHOD_INITIALIZED,
METHOD_TOOLS_CALL, METHOD_TOOLS_LIST, PROTOCOL_VERSION, Tool,
},
stdio::StdioConnection,
},
};
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::{
collections::HashSet,
path::Path,
sync::{Arc, Mutex, atomic::AtomicU64},
};
const MAX_TOOLS_LIST_PAGES: usize = 100;
pub(crate) struct McpClient {
connection: McpConnection,
next_id: AtomicU64,
server_info: Arc<Mutex<Option<crate::mcp::protocol::Implementation>>>,
}
impl std::fmt::Debug for McpClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpClient")
.field("connection", &self.connection)
.finish_non_exhaustive()
}
}
impl std::fmt::Debug for McpConnection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Stdio(_) => f.debug_tuple("Stdio").field(&"<stdio connection>").finish(),
Self::Http(connection) => f.debug_tuple("Http").field(connection).finish(),
}
}
}
enum McpConnection {
Stdio(StdioConnection),
Http(HttpConnection),
}
impl McpConnection {
fn send_request(
&self,
id: RequestId,
method: &str,
params: Option<Value>,
cancellation: Option<&AgentCancellation>,
) -> McpResult<Value> {
match self {
Self::Stdio(connection) => connection.send_request(id, method, params, cancellation),
Self::Http(connection) => connection.send_request(id, method, params, cancellation),
}
}
fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
match self {
Self::Stdio(connection) => connection.send_notification(method, params),
Self::Http(connection) => connection.send_notification(method, params),
}
}
fn request_shutdown(&self) {
match self {
Self::Stdio(connection) => connection.request_shutdown(),
Self::Http(connection) => connection.request_shutdown(),
}
}
fn shutdown(&self) {
match self {
Self::Stdio(connection) => connection.shutdown(),
Self::Http(connection) => connection.shutdown(),
}
}
}
impl McpClient {
pub(crate) fn connect_stdio(config: &McpStdioServerConfig) -> McpResult<Self> {
Ok(Self {
connection: McpConnection::Stdio(StdioConnection::connect(config)?),
next_id: AtomicU64::new(1),
server_info: Arc::new(Mutex::new(None)),
})
}
pub(crate) fn connect_named(
server_name: Option<&str>,
config: &McpServerConfig,
mc_home: Option<&Path>,
) -> McpResult<Self> {
if let Some(server_name) = server_name {
crate::config::validate_mcp_server_name(server_name)
.map_err(|error| McpError::Config(error.to_string()))?;
}
match config {
McpServerConfig::Stdio(config) => Self::connect_stdio(config),
McpServerConfig::Http(config) => Ok(Self {
connection: McpConnection::Http(HttpConnection::connect_named(
server_name,
config,
mc_home,
)?),
next_id: AtomicU64::new(1),
server_info: Arc::new(Mutex::new(None)),
}),
}
}
fn next_id(&self) -> RequestId {
let id = self
.next_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
RequestId::Number(i64::try_from(id).unwrap_or(i64::MAX))
}
fn send_request_inner(
&self,
method: &str,
params: Option<Value>,
cancellation: Option<&AgentCancellation>,
) -> McpResult<Value> {
let id = self.next_id();
self.connection
.send_request(id, method, params, cancellation)
}
pub(crate) fn send_request(&self, method: &str, params: Option<Value>) -> McpResult<Value> {
self.send_request_inner(method, params, None)
}
pub(crate) fn send_request_cancellable(
&self,
method: &str,
params: Option<Value>,
cancellation: &AgentCancellation,
) -> McpResult<Value> {
self.send_request_inner(method, params, Some(cancellation))
}
pub(crate) fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
self.connection.send_notification(method, params)
}
pub(crate) fn initialize(&self) -> McpResult<InitializeResult> {
self.initialize_inner(None)
}
pub(crate) fn initialize_cancellable(
&self,
cancellation: &AgentCancellation,
) -> McpResult<InitializeResult> {
self.initialize_inner(Some(cancellation))
}
fn initialize_inner(
&self,
cancellation: Option<&AgentCancellation>,
) -> McpResult<InitializeResult> {
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
let params = serde_json::to_value(InitializeRequestParams::default())
.map_err(McpError::transport)?;
let response = match cancellation {
Some(cancellation) => {
self.send_request_cancellable(METHOD_INITIALIZE, Some(params), cancellation)
}
None => self.send_request(METHOD_INITIALIZE, Some(params)),
}?;
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
let result: InitializeResult = decode(response)?;
if result.protocol_version != PROTOCOL_VERSION {
return Err(McpError::Protocol {
code: -32602,
message: format!(
"unsupported MCP protocol version {}; expected {PROTOCOL_VERSION}",
result.protocol_version
),
});
}
if let Ok(mut server_info) = self.server_info.lock() {
*server_info = Some(result.server_info.clone());
}
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
self.send_notification(METHOD_INITIALIZED, None)?;
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
Ok(result)
}
pub(crate) fn list_tools(&self) -> McpResult<Vec<Tool>> {
self.list_tools_inner(None)
}
pub(crate) fn list_tools_cancellable(
&self,
cancellation: &AgentCancellation,
) -> McpResult<Vec<Tool>> {
self.list_tools_inner(Some(cancellation))
}
fn list_tools_inner(&self, cancellation: Option<&AgentCancellation>) -> McpResult<Vec<Tool>> {
let mut tools = Vec::new();
let mut cursor = None;
let mut seen_cursors = HashSet::new();
for page in 0..MAX_TOOLS_LIST_PAGES {
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
let params = serde_json::to_value(ListToolsParams {
cursor: cursor.clone(),
})
.map_err(McpError::transport)?;
let response = match cancellation {
Some(cancellation) => {
self.send_request_cancellable(METHOD_TOOLS_LIST, Some(params), cancellation)
}
None => self.send_request(METHOD_TOOLS_LIST, Some(params)),
}?;
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
let result: ListToolsResult = decode(response)?;
tools.extend(result.tools);
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
cursor = result.next_cursor;
if let Some(next_cursor) = cursor.as_ref() {
if !seen_cursors.insert(next_cursor.clone()) {
return Err(McpError::Protocol {
code: -32603,
message: "MCP tools/list returned repeated pagination cursor".to_string(),
});
}
} else {
return Ok(tools);
}
if page + 1 == MAX_TOOLS_LIST_PAGES {
return Err(McpError::Protocol {
code: -32603,
message: format!("MCP tools/list exceeded {MAX_TOOLS_LIST_PAGES} pages"),
});
}
}
unreachable!("tools/list pagination loop returns from inside bounded range")
}
pub(crate) fn call_tool_cancellable(
&self,
name: &str,
arguments: Option<Value>,
cancellation: &AgentCancellation,
) -> McpResult<CallToolResult> {
let params = serde_json::to_value(CallToolParams {
name: name.to_string(),
arguments,
})
.map_err(McpError::transport)?;
decode(self.send_request_cancellable(METHOD_TOOLS_CALL, Some(params), cancellation)?)
}
pub(crate) fn request_shutdown(&self) {
self.connection.request_shutdown();
}
pub(crate) fn shutdown(&mut self) {
self.shutdown_shared();
}
pub(crate) fn shutdown_shared(&self) {
self.connection.shutdown();
}
}
impl Drop for McpClient {
fn drop(&mut self) {
self.shutdown();
}
}
fn decode<T: DeserializeOwned>(value: Value) -> McpResult<T> {
serde_json::from_value(value).map_err(McpError::transport)
}