use anyhow::{bail, Context, Result};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
use std::process::Stdio;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, Command};
use tokio::sync::{oneshot, Mutex};
use tracing::{debug, info};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Framing {
ContentLength,
#[default]
NewlineDelimited,
}
pub(crate) async fn detect_framing<R: tokio::io::AsyncRead + Unpin>(
reader: &mut BufReader<R>,
) -> Result<Option<Framing>> {
let buf = reader.fill_buf().await?;
if buf.is_empty() {
return Ok(None);
}
Ok(Some(if buf[0] == b'C' {
Framing::ContentLength
} else {
Framing::NewlineDelimited
}))
}
pub(crate) async fn read_message<R: tokio::io::AsyncRead + Unpin>(
reader: &mut BufReader<R>,
) -> Result<Option<String>> {
match detect_framing(reader).await? {
Some(Framing::ContentLength) => read_content_length_message(reader).await,
Some(Framing::NewlineDelimited) => read_newline_message(reader).await,
None => Ok(None),
}
}
pub(crate) async fn read_content_length_message<R: tokio::io::AsyncRead + Unpin>(
reader: &mut BufReader<R>,
) -> Result<Option<String>> {
let mut content_length: Option<usize> = None;
loop {
let mut header_line = String::new();
let bytes_read = reader.read_line(&mut header_line).await?;
if bytes_read == 0 {
return Ok(None);
}
let trimmed = header_line.trim();
if trimmed.is_empty() {
break;
}
if let Some(value) = trimmed.strip_prefix("Content-Length:") {
content_length = Some(
value
.trim()
.parse::<usize>()
.context("Invalid Content-Length value")?,
);
}
}
let length = content_length.context("Missing Content-Length header")?;
let mut buf = vec![0u8; length];
reader.read_exact(&mut buf).await?;
String::from_utf8(buf)
.context("Message body is not valid UTF-8")
.map(Some)
}
pub(crate) async fn read_newline_message<R: tokio::io::AsyncRead + Unpin>(
reader: &mut BufReader<R>,
) -> Result<Option<String>> {
let mut line = String::new();
let bytes_read = reader.read_line(&mut line).await?;
if bytes_read == 0 {
return Ok(None);
}
Ok(Some(line.trim().to_string()))
}
pub(crate) async fn write_message<W: tokio::io::AsyncWrite + Unpin>(
writer: &mut W,
body: &str,
) -> Result<()> {
let header = format!("Content-Length: {}\r\n\r\n", body.len());
writer.write_all(header.as_bytes()).await?;
writer.write_all(body.as_bytes()).await?;
writer.flush().await?;
Ok(())
}
pub(crate) async fn write_framed_message<W: tokio::io::AsyncWrite + Unpin>(
writer: &mut W,
body: &str,
framing: Framing,
) -> Result<()> {
match framing {
Framing::ContentLength => write_message(writer, body).await,
Framing::NewlineDelimited => {
writer.write_all(body.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}
}
}
#[derive(Debug, Serialize)]
pub struct JsonRpcRequest {
pub jsonrpc: &'static str,
pub id: u64,
pub method: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub params: Option<Value>,
}
#[derive(Debug, Deserialize)]
pub struct JsonRpcResponse {
pub jsonrpc: String,
pub id: Option<u64>,
pub result: Option<Value>,
pub error: Option<JsonRpcError>,
}
#[derive(Debug, Deserialize)]
pub struct JsonRpcError {
pub code: i64,
pub message: String,
pub data: Option<Value>,
}
impl std::fmt::Display for JsonRpcError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "JSON-RPC error {}: {}", self.code, self.message)
}
}
#[async_trait]
pub trait Transport: Send + Sync {
async fn request(&self, method: &str, params: Option<Value>) -> Result<Value>;
async fn notify(&self, method: &str, params: Option<Value>) -> Result<()>;
async fn shutdown(&self) -> Result<()>;
}
pub struct StdioTransport {
stdin: Arc<Mutex<tokio::process::ChildStdin>>,
pending: Arc<Mutex<HashMap<u64, oneshot::Sender<JsonRpcResponse>>>>,
next_id: AtomicU64,
child: Arc<Mutex<Child>>,
reader_handle: Mutex<Option<tokio::task::JoinHandle<()>>>,
framing: Framing,
}
impl StdioTransport {
pub async fn spawn(
command: &str,
args: &[String],
env: &HashMap<String, String>,
) -> Result<Self> {
info!("Spawning MCP server: {} {:?}", command, args);
let mut cmd = Command::new(command);
cmd.args(args)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped());
crate::safety::process_env::sanitize_command_env(&mut cmd);
for (key, value) in env {
cmd.env(key, value);
}
let mut child = cmd
.spawn()
.with_context(|| format!("Failed to spawn MCP server: {} {:?}", command, args))?;
let stdin = child
.stdin
.take()
.context("Failed to capture MCP server stdin")?;
let stdout = child
.stdout
.take()
.context("Failed to capture MCP server stdout")?;
let stderr = child
.stderr
.take()
.context("Failed to capture MCP server stderr")?;
let pending: Arc<Mutex<HashMap<u64, oneshot::Sender<JsonRpcResponse>>>> =
Arc::new(Mutex::new(HashMap::new()));
let pending_clone = Arc::clone(&pending);
let reader_handle = tokio::spawn(async move {
let mut reader = BufReader::new(stdout);
loop {
match read_message(&mut reader).await {
Ok(Some(body)) => {
match serde_json::from_str::<JsonRpcResponse>(&body) {
Ok(response) => {
if let Some(id) = response.id {
let mut pending = pending_clone.lock().await;
if let Some(tx) = pending.remove(&id) {
let _ = tx.send(response);
} else {
debug!(
"Received response for unknown request ID {}: {:?}",
id, response
);
}
} else {
debug!("MCP server notification: {:?}", response);
}
}
Err(e) => {
debug!("Failed to parse JSON-RPC message from MCP server: {} (body: {})", e, body);
}
}
}
Ok(None) => {
break;
}
Err(e) => {
debug!("MCP stdout framing error: {}", e);
break;
}
}
}
debug!("MCP stdout reader exited");
});
tokio::spawn(async move {
let mut reader = BufReader::new(stderr);
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => break, Ok(_) => {
let trimmed = line.trim();
if !trimmed.is_empty() {
debug!("MCP server stderr: {}", trimmed);
}
}
Err(e) => {
debug!("MCP stderr read error: {}", e);
break;
}
}
}
debug!("MCP stderr drain exited");
});
Ok(Self {
stdin: Arc::new(Mutex::new(stdin)),
pending,
next_id: AtomicU64::new(1),
child: Arc::new(Mutex::new(child)),
reader_handle: Mutex::new(Some(reader_handle)),
framing: Framing::default(),
})
}
pub fn with_framing(mut self, framing: Framing) -> Self {
self.framing = framing;
self
}
}
#[async_trait]
impl Transport for StdioTransport {
async fn request(&self, method: &str, params: Option<Value>) -> Result<Value> {
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let request = JsonRpcRequest {
jsonrpc: "2.0",
id,
method: method.to_string(),
params,
};
let body = serde_json::to_string(&request)?;
let (tx, rx) = oneshot::channel();
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
{
let mut stdin = self.stdin.lock().await;
if let Err(e) = write_framed_message(&mut *stdin, &body, self.framing).await {
self.pending.lock().await.remove(&id);
return Err(e);
}
}
debug!("Sent JSON-RPC request: {} (id={})", method, id);
let response = match tokio::time::timeout(std::time::Duration::from_secs(60), rx).await {
Ok(Ok(resp)) => resp,
Ok(Err(_)) => {
self.pending.lock().await.remove(&id);
bail!("MCP response channel closed for '{}'", method);
}
Err(_) => {
self.pending.lock().await.remove(&id);
bail!("MCP request '{}' timed out after 60s", method);
}
};
if let Some(error) = response.error {
bail!("MCP error for '{}': {}", method, error);
}
response
.result
.ok_or_else(|| anyhow::anyhow!("MCP response for '{}' has no result", method))
}
async fn notify(&self, method: &str, params: Option<Value>) -> Result<()> {
let notification = serde_json::json!({
"jsonrpc": "2.0",
"method": method,
"params": params,
});
let body = serde_json::to_string(¬ification)?;
let mut stdin = self.stdin.lock().await;
write_framed_message(&mut *stdin, &body, self.framing).await?;
debug!("Sent JSON-RPC notification: {}", method);
Ok(())
}
async fn shutdown(&self) -> Result<()> {
info!("Shutting down MCP transport");
let _ = self.request("shutdown", None).await;
let _ = self.notify("notifications/exit", None).await;
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
let mut child = self.child.lock().await;
let _ = child.kill().await;
let mut handle = self.reader_handle.lock().await;
if let Some(h) = handle.take() {
h.abort();
}
Ok(())
}
}
impl Drop for StdioTransport {
fn drop(&mut self) {
if let Ok(mut child) = self.child.try_lock() {
let _ = child.start_kill();
}
if let Ok(mut handle) = self.reader_handle.try_lock() {
if let Some(h) = handle.take() {
h.abort();
}
}
}
}
#[cfg(test)]
#[path = "../../tests/unit/mcp/transport/transport_test.rs"]
mod tests;