use crate::shell::child_pipe_writer::ChildPipeWriter;
use crate::{
config::McpStdioServerConfig,
mcp::{
McpError, McpResult,
jsonrpc::{ErrorData, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, RequestId},
},
process_environment::{self, SubprocessEnvProfile},
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,
io::{BufRead, BufReader},
process::{Child, 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: Mutex<Option<StdioReader>>,
timeout: Duration,
closed: Arc<AtomicBool>,
}
struct StdioReader {
thread: JoinHandle<()>,
done: Receiver<()>,
}
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: Mutex::new(Some(StdioReader {
thread: reader_thread,
done: 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::cancellation::AgentCancellation>,
) -> McpResult<Value> {
let started = std::time::Instant::now();
if let Some(cancellation) = cancellation {
cancellation.check().map_err(McpError::transport)?;
}
let message = serde_json::to_value(JsonRpcRequest::new(id.clone(), method, params))
.map_err(McpError::transport)?;
let encoded = encode_message(&message)?;
let (sender, receiver) = mpsc::sync_channel(1);
{
let mut pending = self
.pending
.lock()
.map_err(|_| McpError::Transport("pending request map poisoned".to_string()))?;
if self.closed.load(Ordering::SeqCst) {
return Err(McpError::Transport("MCP server is closed".to_string()));
}
pending.insert(id.clone(), sender);
}
if let Err(error) =
self.transport
.writer()
.write(&encoded, started + self.timeout, cancellation)
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
self.request_shutdown();
return Err(if started.elapsed() >= self.timeout {
McpError::Timeout {
seconds: self.timeout.as_secs(),
}
} else {
McpError::transport(error)
});
}
loop {
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
let _ = self.pending.lock().map(|mut pending| pending.remove(&id));
self.request_shutdown();
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.request_shutdown();
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.request_shutdown();
return Err(McpError::Transport(
"MCP reader thread disconnected before response".to_string(),
));
}
}
}
}
pub(crate) fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
if self.closed.load(Ordering::SeqCst) {
return Err(McpError::Transport("MCP server is closed".to_string()));
}
let message = serde_json::to_value(JsonRpcNotification::new(method, params))
.map_err(McpError::transport)?;
let result = self
.transport
.writer()
.write(
&encode_message(&message)?,
std::time::Instant::now() + self.timeout.min(Duration::from_secs(1)),
None,
)
.map_err(McpError::transport);
if result.is_err() {
self.request_shutdown();
}
result
}
pub(crate) fn request_shutdown(&self) {
self.closed.store(true, Ordering::SeqCst);
fail_pending(
&self.pending,
McpError::Transport("MCP server closed".to_string()),
);
self.transport.terminate();
}
pub(crate) fn shutdown(&self) {
let mut reader = self
.reader
.lock()
.unwrap_or_else(|error| error.into_inner());
self.request_shutdown();
self.transport.shutdown();
if let Some(reader) = reader.take() {
match reader.done.recv_timeout(PIPE_READER_JOIN_TIMEOUT) {
Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => {
let _ = reader.thread.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<ChildPipeWriter>,
stdout: Option<BufReader<ChildStdout>>,
stderr_receiver: Mutex<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());
process_environment::apply_configured_profile(
&mut command,
SubprocessEnvProfile::McpStdio,
&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 stdin = match ChildPipeWriter::new(stdin) {
Ok(writer) => writer,
Err(error) => {
let _ = terminate_child_tree_and_wait(&mut child);
return Err(McpError::transport(error));
}
};
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(stdin),
stdout: Some(BufReader::new(stdout)),
stderr_receiver: Mutex::new(Some(stderr_receiver)),
})
}
fn cleanup_handle(&self) -> McpStdioCleanupHandle {
McpStdioCleanupHandle {
child: Arc::clone(&self.child),
}
}
pub(crate) fn writer(&self) -> Arc<ChildPipeWriter> {
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 terminate(&self) {
self.cleanup_handle().terminate();
}
pub(crate) fn shutdown(&self) {
let mut receiver = self
.stderr_receiver
.lock()
.unwrap_or_else(|error| error.into_inner());
self.terminate();
if let Some(receiver) = receiver.take() {
let _ = recv_pipe_reader_with_timeout(receiver);
}
}
}
impl Drop for McpStdioTransport {
fn drop(&mut self) {
self.shutdown();
}
}
fn encode_message(message: &Value) -> McpResult<Vec<u8>> {
let mut bytes = serde_json::to_vec(message).map_err(McpError::transport)?;
bytes.push(b'\n');
Ok(bytes)
}
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());
}
}
}