use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use serde_json::Value;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, Command};
use tokio::sync::Mutex;
use xz_mcp_core::{McpClient, McpError, McpTool, McpToolResult};
static NEXT_ID: AtomicU64 = AtomicU64::new(1);
type Writer = tokio::io::BufWriter<tokio::process::ChildStdin>;
type Reader = BufReader<tokio::process::ChildStdout>;
struct Connection {
_child: Child,
writer: Mutex<Writer>,
reader: Mutex<Reader>,
}
pub struct StdioMcpClient {
command: String,
args: Vec<String>,
env: HashMap<String, String>,
use_content_length: bool,
conn: Mutex<Option<Arc<Connection>>>,
}
impl StdioMcpClient {
pub fn new(
command: impl Into<String>,
args: Vec<String>,
env: HashMap<String, String>,
) -> Self {
Self {
command: command.into(),
args,
env,
use_content_length: true,
conn: Mutex::new(None),
}
}
pub fn with_newline_framing(mut self) -> Self {
self.use_content_length = false;
self
}
fn next_id() -> u64 {
NEXT_ID.fetch_add(1, Ordering::SeqCst)
}
}
#[async_trait]
impl McpClient for StdioMcpClient {
async fn connect(&mut self) -> Result<(), McpError> {
let mut cmd = Command::new(&self.command);
cmd.args(&self.args);
for (k, v) in &self.env {
cmd.env(k, v);
}
cmd.stdin(std::process::Stdio::piped());
cmd.stdout(std::process::Stdio::piped());
cmd.stderr(std::process::Stdio::piped());
cmd.kill_on_drop(true);
let mut child = cmd
.spawn()
.map_err(|e| McpError::Connection(format!("spawn: {e}")))?;
let stdin = child
.stdin
.take()
.ok_or_else(|| McpError::Connection("no stdin".into()))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| McpError::Connection("no stdout".into()))?;
let conn = Arc::new(Connection {
_child: child,
writer: Mutex::new(tokio::io::BufWriter::new(stdin)),
reader: Mutex::new(BufReader::new(stdout)),
});
let init_req = rpc(
"initialize",
serde_json::json!({
"protocolVersion": "2024-11-05",
"capabilities": {},
"clientInfo": { "name": "xz-mcp", "version": "0.1" }
}),
);
send(&conn, &init_req, self.use_content_length).await?;
let notif = serde_json::json!({"jsonrpc":"2.0","method":"notifications/initialized"});
let _ = send_raw(&conn, ¬if, self.use_content_length).await;
*self.conn.lock().await = Some(conn);
Ok(())
}
async fn list_tools(&self) -> Result<Vec<McpTool>, McpError> {
let conn = self
.conn
.lock()
.await
.clone()
.ok_or_else(|| McpError::Connection("not connected".into()))?;
let req = rpc("tools/list", serde_json::json!({}));
let resp = send(&conn, &req, self.use_content_length).await?;
let tools = resp["result"]["tools"]
.as_array()
.ok_or_else(|| McpError::Protocol("missing tools array".into()))?;
tools
.iter()
.map(|t| {
Ok(McpTool {
name: t["name"].as_str().unwrap_or("").into(),
description: t["description"].as_str().unwrap_or("").into(),
input_schema: t
.get("inputSchema")
.cloned()
.unwrap_or_else(|| serde_json::json!({})),
})
})
.collect()
}
async fn call_tool(&self, name: &str, args: Value) -> Result<McpToolResult, McpError> {
let conn = self
.conn
.lock()
.await
.clone()
.ok_or_else(|| McpError::Connection("not connected".into()))?;
let req = rpc(
"tools/call",
serde_json::json!({"name":name,"arguments":args}),
);
let resp = send(&conn, &req, self.use_content_length).await?;
let result = &resp["result"];
let content: Vec<Value> = result["content"].as_array().cloned().unwrap_or_default();
Ok(McpToolResult {
content: serde_json::from_value(Value::Array(content)).unwrap_or_default(),
is_error: result["isError"].as_bool().unwrap_or(false),
})
}
async fn is_alive(&self) -> bool {
self.conn.lock().await.is_some()
}
}
fn rpc(method: &str, params: Value) -> Value {
serde_json::json!({
"jsonrpc": "2.0",
"id": StdioMcpClient::next_id(),
"method": method,
"params": params,
})
}
fn encode_content_length(body: &str) -> Vec<u8> {
let header = format!("Content-Length: {}\r\n\r\n", body.len());
let mut out = Vec::with_capacity(header.len() + body.len());
out.extend_from_slice(header.as_bytes());
out.extend_from_slice(body.as_bytes());
out
}
async fn send_raw(conn: &Connection, msg: &Value, content_length: bool) -> Result<(), McpError> {
let body = serde_json::to_string(msg)?;
let bytes = if content_length {
encode_content_length(&body)
} else {
let mut b = body.into_bytes();
b.push(b'\n');
b
};
let mut w = conn.writer.lock().await;
w.write_all(&bytes).await.map_err(McpError::Io)?;
w.flush().await.map_err(McpError::Io)
}
async fn send(conn: &Connection, req: &Value, content_length: bool) -> Result<Value, McpError> {
let req_id = req["id"].clone();
send_raw(conn, req, content_length).await?;
let mut r = conn.reader.lock().await;
loop {
let msg = read_json_message(&mut r).await?;
if msg.get("id").is_none() {
continue; }
if msg["id"] != req_id {
continue;
}
if let Some(err) = msg.get("error") {
return Err(McpError::Server(err.to_string()));
}
return Ok(msg);
}
}
async fn read_json_message(reader: &mut Reader) -> Result<Value, McpError> {
let mut first = String::new();
let n = reader
.read_line(&mut first)
.await
.map_err(McpError::Io)?;
if n == 0 {
return Err(McpError::Connection(
"server closed stdout while waiting for response".into(),
));
}
let mut header_line = first;
loop {
let trimmed = header_line.trim_end_matches(['\r', '\n']);
if !trimmed.is_empty() {
break;
}
header_line.clear();
let n = reader.read_line(&mut header_line).await.map_err(McpError::Io)?;
if n == 0 {
return Err(McpError::Connection(
"server closed stdout while waiting for response".into(),
));
}
}
let trimmed = header_line.trim_end_matches(['\r', '\n']);
if let Some(rest) = trimmed
.strip_prefix("Content-Length:")
.or_else(|| trimmed.strip_prefix("content-length:"))
{
let content_length: usize = rest.trim().parse().map_err(|e| {
McpError::Protocol(format!("Invalid Content-Length '{rest}': {e}"))
})?;
loop {
let mut line = String::new();
let n = reader.read_line(&mut line).await.map_err(McpError::Io)?;
if n == 0 {
return Err(McpError::Connection(
"server closed stdout mid-headers".into(),
));
}
if line == "\n" || line == "\r\n" || line.trim().is_empty() {
break;
}
}
let mut buf = vec![0u8; content_length];
reader.read_exact(&mut buf).await.map_err(McpError::Io)?;
let response: Value =
serde_json::from_slice(&buf).map_err(|e| McpError::Protocol(e.to_string()))?;
return Ok(response);
}
serde_json::from_str(trimmed).map_err(|e| McpError::Protocol(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn content_length_encode_roundtrip_shape() {
let body = r#"{"jsonrpc":"2.0","id":1}"#;
let enc = encode_content_length(body);
let s = String::from_utf8_lossy(&enc);
assert!(s.starts_with("Content-Length: "));
assert!(s.contains("\r\n\r\n"));
assert!(s.ends_with(body));
}
}