use crate::{
config::McpStdioServerConfig,
mcp::{
McpError, McpResult,
jsonrpc::{ErrorData, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, RequestId},
},
tools::process::{
PIPE_READER_JOIN_TIMEOUT, PipeReaderHandle, recv_pipe_reader_with_timeout,
spawn_bounded_pipe_reader, terminate_child_tree_and_wait,
},
};
use serde_json::Value;
use std::{
collections::HashMap,
env,
io::{BufRead, BufReader, BufWriter, Write},
process::{Child, ChildStdin, ChildStdout, Command, Stdio},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
mpsc::{self, Receiver, SyncSender},
},
thread::{self, JoinHandle},
time::Duration,
};
pub(crate) const MAX_MCP_FRAME_BYTES: usize = 1_048_576;
const MAX_MCP_STDERR_BYTES: usize = 8192;
type PendingMap = Arc<Mutex<HashMap<RequestId, SyncSender<McpResult<Value>>>>>;
pub(crate) struct StdioConnection {
transport: McpStdioTransport,
pending: PendingMap,
reader_thread: Option<JoinHandle<()>>,
reader_done: Option<Receiver<()>>,
timeout: Duration,
closed: Arc<AtomicBool>,
}
impl StdioConnection {
pub(crate) fn connect(config: &McpStdioServerConfig) -> McpResult<Self> {
let mut transport = McpStdioTransport::spawn(config)?;
let reader = transport.take_stdout()?;
let pending: PendingMap = Arc::new(Mutex::new(HashMap::new()));
let reader_pending = Arc::clone(&pending);
let closed = Arc::new(AtomicBool::new(false));
let reader_closed = Arc::clone(&closed);
let reader_transport = transport.cleanup_handle();
let (reader_done_sender, reader_done) = mpsc::sync_channel(1);
let reader_thread = thread::spawn(move || {
let mut reader = reader;
loop {
let value = match read_message(&mut reader) {
Ok(value) => value,
Err(error) => {
reader_closed.store(true, Ordering::SeqCst);
fail_pending(
&reader_pending,
McpError::Transport(format!(
"MCP server stdout closed or unreadable: {error}"
)),
);
reader_transport.terminate();
break;
}
};
let message = match serde_json::from_value::<JsonRpcMessage>(value) {
Ok(message) => message,
Err(error) => {
reader_closed.store(true, Ordering::SeqCst);
fail_pending(
&reader_pending,
McpError::Transport(format!("MCP malformed JSON-RPC message: {error}")),
);
reader_transport.terminate();
break;
}
};
match message {
JsonRpcMessage::Response(response) => {
if let Ok(mut pending) = reader_pending.lock()
&& let Some(sender) = pending.remove(&response.id)
{
let _ = sender.send(Ok(response.result));
}
}
JsonRpcMessage::Error(error) => {
if let Some(id) = error.id
&& let Ok(mut pending) = reader_pending.lock()
&& let Some(sender) = pending.remove(&id)
{
let ErrorData { code, message, .. } = error.error;
let _ = sender.send(Err(McpError::Protocol { code, message }));
}
}
JsonRpcMessage::Notification(_) | JsonRpcMessage::Request(_) => {}
}
}
let _ = reader_done_sender.send(());
});
Ok(Self {
transport,
pending,
reader_thread: Some(reader_thread),
reader_done: Some(reader_done),
timeout: Duration::from_secs(
config
.timeout
.unwrap_or(crate::config::DEFAULT_MCP_TIMEOUT_SECONDS),
),
closed,
})
}
pub(crate) fn send_request(
&self,
id: RequestId,
method: &str,
params: Option<Value>,
cancellation: Option<&crate::agent::cancellation::AgentCancellation>,
) -> McpResult<Value> {
if self.closed.load(Ordering::SeqCst) {
return Err(McpError::Transport("MCP server is closed".to_string()));
}
let (sender, receiver) = mpsc::sync_channel(1);
self.pending
.lock()
.map_err(|_| McpError::Transport("pending request map poisoned".to_string()))?
.insert(id.clone(), sender);
let message = serde_json::to_value(JsonRpcRequest::new(id.clone(), method, params))
.map_err(McpError::transport)?;
if let Err(error) = McpStdioTransport::write_message_to(&self.transport.writer(), &message)
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
return Err(error);
}
let started = std::time::Instant::now();
loop {
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
self.close_after_inflight_failure();
return Err(McpError::Transport(error.to_string()));
}
let remaining = self.timeout.saturating_sub(started.elapsed());
if remaining.is_zero() {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
self.close_after_inflight_failure();
return Err(McpError::Timeout {
seconds: self.timeout.as_secs(),
});
}
match receiver.recv_timeout(remaining.min(Duration::from_millis(100))) {
Ok(result) => return result,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => {
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
self.close_after_inflight_failure();
return Err(McpError::Transport(
"MCP reader thread disconnected before response".to_string(),
));
}
}
}
}
pub(crate) fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
let message = serde_json::to_value(JsonRpcNotification::new(method, params))
.map_err(McpError::transport)?;
McpStdioTransport::write_message_to(&self.transport.writer(), &message)
}
fn close_after_inflight_failure(&self) {
self.closed.store(true, Ordering::SeqCst);
self.transport.terminate();
fail_pending(
&self.pending,
McpError::Transport("MCP server closed after in-flight request failure".to_string()),
);
}
#[cfg(test)]
pub(crate) fn is_process_running(&self) -> bool {
self.transport.is_running()
}
pub(crate) fn shutdown(&mut self) {
self.transport.shutdown();
if let Some(handle) = self.reader_thread.take()
&& let Some(receiver) = self.reader_done.take()
{
match receiver.recv_timeout(PIPE_READER_JOIN_TIMEOUT) {
Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => {
let _ = handle.join();
}
Err(mpsc::RecvTimeoutError::Timeout) => {
}
}
}
}
}
impl Drop for StdioConnection {
fn drop(&mut self) {
self.shutdown();
}
}
fn fail_pending(pending: &PendingMap, error: McpError) {
if let Ok(mut pending) = pending.lock() {
for (_, sender) in pending.drain() {
let _ = sender.send(Err(error.clone()));
}
}
}
pub(crate) struct McpStdioTransport {
child: Arc<Mutex<Child>>,
stdin: Arc<Mutex<BufWriter<ChildStdin>>>,
stdout: Option<BufReader<ChildStdout>>,
stderr_receiver: Option<PipeReaderHandle>,
}
#[derive(Clone)]
struct McpStdioCleanupHandle {
child: Arc<Mutex<Child>>,
}
impl McpStdioCleanupHandle {
fn terminate(&self) {
if let Ok(mut child) = self.child.lock() {
let _ = terminate_child_tree_and_wait(&mut child);
}
}
}
impl McpStdioTransport {
pub(crate) fn spawn(config: &McpStdioServerConfig) -> McpResult<Self> {
if config.command.trim().is_empty() {
return Err(McpError::Config(
"stdio command must not be empty".to_string(),
));
}
let mut command = Command::new(&config.command);
command
.args(&config.args)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.env_clear();
if !config.env.contains_key("PATH")
&& let Some(path) = env::var_os("PATH")
{
command.env("PATH", path);
}
command.envs(&config.env);
#[cfg(unix)]
{
use std::os::unix::process::CommandExt;
command.process_group(0);
}
let mut child = command.spawn().map_err(McpError::transport)?;
let stdin = child
.stdin
.take()
.ok_or_else(|| McpError::Transport("stdio stdin pipe unavailable".to_string()))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| McpError::Transport("stdio stdout pipe unavailable".to_string()))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| McpError::Transport("stdio stderr pipe unavailable".to_string()))?;
let stderr_truncated = Arc::new(AtomicBool::new(false));
let stderr_receiver =
spawn_bounded_pipe_reader(stderr, MAX_MCP_STDERR_BYTES, Arc::clone(&stderr_truncated));
Ok(Self {
child: Arc::new(Mutex::new(child)),
stdin: Arc::new(Mutex::new(BufWriter::new(stdin))),
stdout: Some(BufReader::new(stdout)),
stderr_receiver: Some(stderr_receiver),
})
}
fn cleanup_handle(&self) -> McpStdioCleanupHandle {
McpStdioCleanupHandle {
child: Arc::clone(&self.child),
}
}
pub(crate) fn writer(&self) -> Arc<Mutex<BufWriter<ChildStdin>>> {
Arc::clone(&self.stdin)
}
pub(crate) fn take_stdout(&mut self) -> McpResult<BufReader<ChildStdout>> {
self.stdout
.take()
.ok_or_else(|| McpError::Transport("stdio stdout already taken".to_string()))
}
pub(crate) fn write_message_to(
writer: &Arc<Mutex<BufWriter<ChildStdin>>>,
msg: &Value,
) -> McpResult<()> {
let mut writer = writer
.lock()
.map_err(|_| McpError::Transport("stdio writer lock poisoned".to_string()))?;
serde_json::to_writer(&mut *writer, msg).map_err(McpError::transport)?;
writer.write_all(b"\n").map_err(McpError::transport)?;
writer.flush().map_err(McpError::transport)
}
pub(crate) fn terminate(&self) {
self.cleanup_handle().terminate();
}
#[cfg(test)]
pub(crate) fn is_running(&self) -> bool {
self.child
.lock()
.ok()
.and_then(|mut child| child.try_wait().ok().flatten())
.is_none()
}
pub(crate) fn shutdown(&mut self) {
self.terminate();
if let Some(receiver) = self.stderr_receiver.take() {
let _ = recv_pipe_reader_with_timeout(receiver);
}
}
}
impl Drop for McpStdioTransport {
fn drop(&mut self) {
self.shutdown();
}
}
pub(crate) fn read_message(reader: &mut impl BufRead) -> McpResult<Value> {
let mut bytes = Vec::new();
loop {
bytes.clear();
let count = read_one_frame(reader, &mut bytes)?;
if count == 0 {
return Err(McpError::Transport("stdio stdout closed".to_string()));
}
while matches!(bytes.last(), Some(b'\n' | b'\r')) {
bytes.pop();
}
if bytes.is_empty() {
continue;
}
return serde_json::from_slice(&bytes).map_err(McpError::transport);
}
}
fn read_one_frame(reader: &mut impl BufRead, bytes: &mut Vec<u8>) -> McpResult<usize> {
loop {
let available = reader.fill_buf().map_err(McpError::transport)?;
if available.is_empty() {
return Ok(bytes.len());
}
let take = available
.iter()
.position(|byte| *byte == b'\n')
.map_or(available.len(), |index| index + 1);
if bytes.len() + take > MAX_MCP_FRAME_BYTES {
return Err(McpError::Transport(format!(
"MCP frame exceeded {MAX_MCP_FRAME_BYTES} bytes"
)));
}
bytes.extend_from_slice(&available[..take]);
reader.consume(take);
if bytes.ends_with(b"\n") {
return Ok(bytes.len());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
io::Cursor,
process::Command,
thread,
time::{Duration, Instant},
};
fn stdio_config(
command: &str,
args: Vec<String>,
timeout: Option<u64>,
) -> McpStdioServerConfig {
let mut config = McpStdioServerConfig {
command: command.to_string(),
args,
env: std::collections::BTreeMap::new(),
enabled: true,
timeout,
};
if let Some(path) = env::var_os("PATH") {
config
.env
.insert("PATH".to_string(), path.to_string_lossy().into_owned());
}
config
}
fn wait_until_stopped(connection: &StdioConnection, timeout: Duration) -> bool {
let started = Instant::now();
while started.elapsed() < timeout {
if !connection.is_process_running() {
return true;
}
thread::sleep(Duration::from_millis(20));
}
false
}
#[test]
fn read_message_skips_empty_and_accepts_crlf() {
let mut reader = Cursor::new(b"\n\r\n{\"ok\":true}\r\n");
assert_eq!(
read_message(&mut reader).unwrap(),
serde_json::json!({"ok": true})
);
}
#[test]
fn read_message_rejects_oversized_frame() {
let mut data = vec![b'a'; MAX_MCP_FRAME_BYTES + 1];
data.push(b'\n');
let mut reader = Cursor::new(data);
let error = read_message(&mut reader).unwrap_err().to_string();
assert!(error.contains("exceeded"), "{error}");
}
#[test]
fn terminate_and_shutdown_do_not_wait_for_writer_lock() {
let config = stdio_config(
"sh",
vec!["-c".to_string(), "sleep 30".to_string()],
Some(1),
);
let transport = McpStdioTransport::spawn(&config).unwrap();
let writer = transport.writer();
let _guard = writer.lock().unwrap();
let started = Instant::now();
transport.terminate();
assert!(
started.elapsed() < Duration::from_secs(1),
"elapsed: {:?}",
started.elapsed()
);
assert!(!transport.is_running());
let mut transport = McpStdioTransport::spawn(&config).unwrap();
let writer = transport.writer();
let _guard = writer.lock().unwrap();
let started = Instant::now();
transport.shutdown();
assert!(
started.elapsed() < Duration::from_secs(1),
"elapsed: {:?}",
started.elapsed()
);
assert!(!transport.is_running());
}
#[test]
fn reader_fatal_error_terminates_child() {
let config = stdio_config(
"sh",
vec![
"-c".to_string(),
"printf 'not-json\\n'; sleep 30".to_string(),
],
Some(1),
);
let connection = StdioConnection::connect(&config).unwrap();
assert!(wait_until_stopped(&connection, Duration::from_secs(2)));
}
#[cfg(unix)]
#[test]
fn shutdown_does_not_block_on_escaped_stdout_owner() {
if Command::new("python3").arg("--version").status().is_err() {
return;
}
let config = stdio_config(
"python3",
vec![
"-c".to_string(),
"import subprocess, sys, time; subprocess.Popen(['sleep', '5'], stdout=sys.stdout, stderr=subprocess.DEVNULL, start_new_session=True); time.sleep(30)".to_string(),
],
Some(1),
);
let mut connection = StdioConnection::connect(&config).unwrap();
let started = Instant::now();
connection.shutdown();
assert!(
started.elapsed() < Duration::from_secs(2),
"elapsed: {:?}",
started.elapsed()
);
}
}