a3s_code_core/mcp/transport/
stdio.rs1use super::McpTransport;
6use crate::mcp::protocol::{JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, McpNotification};
7use crate::tools::process::{configure_process_group, ProcessGroupGuard};
8use anyhow::{anyhow, Context, Result};
9use async_trait::async_trait;
10use futures::StreamExt;
11use std::collections::HashMap;
12use std::process::Stdio;
13use std::sync::atomic::{AtomicBool, Ordering};
14use std::sync::{Arc, Mutex as StdMutex};
15use std::time::Duration;
16use tokio::io::{AsyncReadExt, AsyncWriteExt};
17use tokio::process::{Child, ChildStderr, Command};
18use tokio::sync::{mpsc, oneshot, RwLock};
19use tokio::task::JoinHandle;
20use tokio_util::codec::{FramedRead, LinesCodec};
21use tokio_util::sync::CancellationToken;
22
23const DEFAULT_REQUEST_TIMEOUT_SECS: u64 = 60;
25const PROCESS_SETTLEMENT_TIMEOUT: Duration = Duration::from_secs(1);
26const MAX_MCP_STDIO_LINE_BYTES: usize = 8 * 1024 * 1024;
27
28pub struct StdioTransport {
30 process_group: Arc<StdMutex<ProcessGroupGuard>>,
33 process_task: StdMutex<Option<JoinHandle<std::io::Result<()>>>>,
35 io_tasks: StdMutex<Vec<JoinHandle<()>>>,
38 stdin_tx: mpsc::Sender<String>,
40 pending: Arc<RwLock<HashMap<u64, oneshot::Sender<JsonRpcResponse>>>>,
42 notification_rx: RwLock<Option<mpsc::Receiver<McpNotification>>>,
44 connected: Arc<AtomicBool>,
46 shutdown: CancellationToken,
48 request_timeout_secs: u64,
50}
51
52impl StdioTransport {
53 pub async fn spawn(
55 command: &str,
56 args: &[String],
57 env: &HashMap<String, String>,
58 ) -> Result<Self> {
59 Self::spawn_with_timeout(command, args, env, DEFAULT_REQUEST_TIMEOUT_SECS).await
60 }
61
62 pub async fn spawn_with_timeout(
64 command: &str,
65 args: &[String],
66 env: &HashMap<String, String>,
67 request_timeout_secs: u64,
68 ) -> Result<Self> {
69 let mut cmd = Command::new(command);
71 cmd.args(args)
72 .stdin(Stdio::piped())
73 .stdout(Stdio::piped())
74 .stderr(Stdio::piped())
75 .kill_on_drop(true);
76 configure_process_group(&mut cmd);
77
78 for (key, value) in env {
80 cmd.env(key, value);
81 }
82
83 let mut child = cmd
84 .spawn()
85 .with_context(|| format!("Failed to spawn MCP server: {} {:?}", command, args))?;
86 let process_group = ProcessGroupGuard::for_child(&child);
87
88 let stdin = child.stdin.take().ok_or_else(|| anyhow!("No stdin"))?;
89 let stdout = child.stdout.take().ok_or_else(|| anyhow!("No stdout"))?;
90 let stderr = child.stderr.take().ok_or_else(|| anyhow!("No stderr"))?;
91
92 let (stdin_tx, mut stdin_rx) = mpsc::channel::<String>(100);
94 let (notification_tx, notification_rx) = mpsc::channel::<McpNotification>(100);
95 let pending: Arc<RwLock<HashMap<u64, oneshot::Sender<JsonRpcResponse>>>> =
96 Arc::new(RwLock::new(HashMap::new()));
97 let connected = Arc::new(AtomicBool::new(true));
98 let shutdown = CancellationToken::new();
99 let process_group = Arc::new(StdMutex::new(process_group));
100 let process_task = tokio::spawn(monitor_child(
101 child,
102 Arc::clone(&process_group),
103 shutdown.clone(),
104 ));
105
106 let mut stdin_writer = stdin;
108 let writer_connected = Arc::clone(&connected);
109 let writer_pending = Arc::clone(&pending);
110 let writer_shutdown = shutdown.clone();
111 let writer_task = tokio::spawn(async move {
112 loop {
113 let message = tokio::select! {
114 _ = writer_shutdown.cancelled() => break,
115 message = stdin_rx.recv() => message,
116 };
117 let Some(message) = message else {
118 break;
119 };
120 let write = async {
121 stdin_writer.write_all(message.as_bytes()).await?;
122 stdin_writer.flush().await
123 };
124 let result = tokio::select! {
125 _ = writer_shutdown.cancelled() => break,
126 result = write => result,
127 };
128 if let Err(error) = result {
129 tracing::error!("Failed to write to MCP stdin: {}", error);
130 break;
131 }
132 }
133 writer_connected.store(false, Ordering::SeqCst);
134 writer_pending.write().await.clear();
135 writer_shutdown.cancel();
136 });
137
138 let pending_clone = pending.clone();
140 let reader_connected = Arc::clone(&connected);
141 let reader_shutdown = shutdown.clone();
142 let reader_task = tokio::spawn(async move {
143 let mut reader = FramedRead::new(
144 stdout,
145 LinesCodec::new_with_max_length(MAX_MCP_STDIO_LINE_BYTES),
146 );
147 loop {
148 let read = tokio::select! {
149 _ = reader_shutdown.cancelled() => break,
150 read = reader.next() => read,
151 };
152 match read {
153 None => {
154 tracing::debug!("MCP stdout closed");
155 break;
156 }
157 Some(Ok(line)) => {
158 let trimmed = line.trim();
159 if trimmed.is_empty() {
160 continue;
161 }
162
163 if let Ok(response) = serde_json::from_str::<JsonRpcResponse>(trimmed) {
165 if let Some(id) = response.id {
166 let mut pending = pending_clone.write().await;
167 if let Some(tx) = pending.remove(&id) {
168 let _ = tx.send(response);
169 }
170 }
171 continue;
172 }
173
174 if let Ok(notification) =
176 serde_json::from_str::<JsonRpcNotification>(trimmed)
177 {
178 let mcp_notif = McpNotification::from_json_rpc(¬ification);
179 tokio::select! {
180 _ = reader_shutdown.cancelled() => break,
181 _ = notification_tx.send(mcp_notif) => {}
182 }
183 continue;
184 }
185
186 tracing::warn!("Unknown MCP message: {}", trimmed);
187 }
188 Some(Err(e)) => {
189 tracing::error!("Failed to read MCP stdout: {}", e);
190 break;
191 }
192 }
193 }
194 reader_connected.store(false, Ordering::SeqCst);
195 pending_clone.write().await.clear();
196 reader_shutdown.cancel();
197 });
198 let stderr_task = tokio::spawn(drain_stderr(stderr, shutdown.clone()));
199
200 Ok(Self {
201 process_group,
202 process_task: StdMutex::new(Some(process_task)),
203 io_tasks: StdMutex::new(vec![writer_task, reader_task, stderr_task]),
204 stdin_tx,
205 pending,
206 notification_rx: RwLock::new(Some(notification_rx)),
207 connected,
208 shutdown,
209 request_timeout_secs,
210 })
211 }
212
213 fn kill_process_group(&self) {
214 self.process_group
215 .lock()
216 .unwrap_or_else(std::sync::PoisonError::into_inner)
217 .kill();
218 }
219}
220
221impl Drop for StdioTransport {
222 fn drop(&mut self) {
223 self.connected.store(false, Ordering::SeqCst);
224 self.shutdown.cancel();
225 self.process_group
226 .lock()
227 .unwrap_or_else(std::sync::PoisonError::into_inner)
228 .kill();
229 }
230}
231
232#[async_trait]
233impl McpTransport for StdioTransport {
234 async fn request(&self, request: JsonRpcRequest) -> Result<JsonRpcResponse> {
235 if !self.connected.load(Ordering::SeqCst) {
236 return Err(anyhow!("Transport not connected"));
237 }
238
239 let (tx, rx) = oneshot::channel();
241 let request_id = request.id;
242
243 {
245 let mut pending = self.pending.write().await;
246 pending.insert(request_id, tx);
247 }
248 if !self.connected.load(Ordering::SeqCst) {
249 self.pending.write().await.remove(&request_id);
250 return Err(anyhow!("Transport not connected"));
251 }
252
253 let msg = serde_json::to_string(&request)? + "\n";
255 self.stdin_tx
256 .send(msg)
257 .await
258 .map_err(|_| anyhow!("Failed to send request"))?;
259
260 let response = match tokio::time::timeout(
262 std::time::Duration::from_secs(self.request_timeout_secs),
263 rx,
264 )
265 .await
266 {
267 Ok(Ok(resp)) => resp,
268 Ok(Err(_)) => {
269 self.pending.write().await.remove(&request_id);
271 return Err(anyhow!("Response channel closed"));
272 }
273 Err(_) => {
274 self.pending.write().await.remove(&request_id);
276 return Err(anyhow!(
277 "MCP request timed out after {}s",
278 self.request_timeout_secs
279 ));
280 }
281 };
282
283 Ok(response)
284 }
285
286 async fn notify(&self, notification: JsonRpcNotification) -> Result<()> {
287 if !self.connected.load(Ordering::SeqCst) {
288 return Err(anyhow!("Transport not connected"));
289 }
290
291 let msg = serde_json::to_string(¬ification)? + "\n";
292 self.stdin_tx
293 .send(msg)
294 .await
295 .map_err(|_| anyhow!("Failed to send notification"))?;
296
297 Ok(())
298 }
299
300 fn notifications(&self) -> mpsc::Receiver<McpNotification> {
301 let mut rx_guard = self.notification_rx.blocking_write();
304 rx_guard.take().unwrap_or_else(|| {
305 let (_, rx) = mpsc::channel(1);
306 rx
307 })
308 }
309
310 async fn close(&self) -> Result<()> {
311 self.connected.store(false, Ordering::SeqCst);
312 self.shutdown.cancel();
313 self.pending.write().await.clear();
314 self.kill_process_group();
315
316 let process_task = self
317 .process_task
318 .lock()
319 .unwrap_or_else(std::sync::PoisonError::into_inner)
320 .take();
321 let process_result = if let Some(process_task) = process_task {
322 match tokio::time::timeout(PROCESS_SETTLEMENT_TIMEOUT * 2, process_task).await {
323 Ok(Ok(Ok(()))) => Ok(()),
324 Ok(Ok(Err(error))) => {
325 Err(error).context("Failed to reap MCP server after termination")
326 }
327 Ok(Err(error)) => Err(anyhow!("MCP server monitor task failed: {error}")),
328 Err(_) => Err(anyhow!(
329 "MCP server monitor did not settle after termination"
330 )),
331 }
332 } else {
333 Ok(())
334 };
335 let io_tasks = self
336 .io_tasks
337 .lock()
338 .unwrap_or_else(std::sync::PoisonError::into_inner)
339 .drain(..)
340 .collect();
341 let io_result = settle_io_tasks(io_tasks).await;
342
343 process_result?;
344 io_result
345 }
346
347 fn is_connected(&self) -> bool {
348 self.connected.load(Ordering::SeqCst)
349 }
350}
351
352async fn settle_io_tasks(tasks: Vec<JoinHandle<()>>) -> Result<()> {
353 let mut first_error = None;
354 for mut task in tasks {
355 match tokio::time::timeout(PROCESS_SETTLEMENT_TIMEOUT, &mut task).await {
356 Ok(Ok(())) => {}
357 Ok(Err(error)) => {
358 first_error.get_or_insert_with(|| anyhow!("MCP stdio task failed: {error}"));
359 }
360 Err(_) => {
361 task.abort();
362 let _ = task.await;
363 first_error
364 .get_or_insert_with(|| anyhow!("MCP stdio task did not settle during close"));
365 }
366 }
367 }
368 first_error.map_or(Ok(()), Err)
369}
370
371async fn monitor_child(
372 mut child: Child,
373 process_group: Arc<StdMutex<ProcessGroupGuard>>,
374 shutdown: CancellationToken,
375) -> std::io::Result<()> {
376 let result = tokio::select! {
377 result = child.wait() => result,
378 _ = shutdown.cancelled() => {
379 process_group
380 .lock()
381 .unwrap_or_else(std::sync::PoisonError::into_inner)
382 .kill();
383 let _ = child.start_kill();
384 match tokio::time::timeout(PROCESS_SETTLEMENT_TIMEOUT, child.wait()).await {
385 Ok(result) => result,
386 Err(_) => {
387 return Err(std::io::Error::new(
388 std::io::ErrorKind::TimedOut,
389 "MCP server did not exit after process-group termination",
390 ));
391 }
392 }
393 }
394 };
395 process_group
397 .lock()
398 .unwrap_or_else(std::sync::PoisonError::into_inner)
399 .kill();
400 result.map(|_| ())
401}
402
403async fn drain_stderr(mut stderr: ChildStderr, shutdown: CancellationToken) {
404 let mut chunk = [0_u8; 4096];
405 loop {
406 let read = tokio::select! {
407 _ = shutdown.cancelled() => break,
408 read = stderr.read(&mut chunk) => read,
409 };
410 match read {
411 Ok(0) => break,
412 Ok(count) => {
413 tracing::debug!(
414 "MCP server stderr: {}",
415 String::from_utf8_lossy(&chunk[..count]).trim_end()
416 );
417 }
418 Err(error) => {
419 tracing::debug!("Failed to read MCP stderr: {}", error);
420 break;
421 }
422 }
423 }
424}
425
426#[cfg(test)]
427mod tests {
428 use super::*;
429
430 #[cfg(unix)]
431 async fn wait_for_path(path: &std::path::Path) {
432 tokio::time::timeout(Duration::from_secs(1), async {
433 while !path.exists() {
434 tokio::time::sleep(Duration::from_millis(10)).await;
435 }
436 })
437 .await
438 .expect("MCP test process did not start");
439 }
440
441 #[cfg(unix)]
442 async fn spawn_descendant_writer(
443 started: &std::path::Path,
444 leaked: &std::path::Path,
445 ) -> StdioTransport {
446 let args = vec![
447 "-c".to_string(),
448 "touch \"$1\"; (sleep 0.30; touch \"$2\") & wait".to_string(),
449 "mcp-process-tree-test".to_string(),
450 started.to_string_lossy().into_owned(),
451 leaked.to_string_lossy().into_owned(),
452 ];
453 StdioTransport::spawn("/bin/sh", &args, &HashMap::new())
454 .await
455 .unwrap()
456 }
457
458 #[tokio::test]
459 async fn test_stdio_transport_spawn_invalid_command() {
460 let result = StdioTransport::spawn("nonexistent_command_12345", &[], &HashMap::new()).await;
461 assert!(result.is_err());
462 }
463
464 #[tokio::test]
465 async fn test_stdio_transport_spawn_echo() {
466 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
468
469 if let Ok(transport) = result {
470 assert!(transport.is_connected());
471 transport.close().await.unwrap();
472 assert!(!transport.is_connected());
473 }
474 }
476
477 #[tokio::test]
478 async fn test_stdio_transport_is_connected_initial() {
479 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
480 if let Ok(transport) = result {
481 assert!(transport.is_connected());
482 let _ = transport.close().await;
483 }
484 }
485
486 #[tokio::test]
487 async fn test_stdio_transport_close_disconnects() {
488 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
489 if let Ok(transport) = result {
490 assert!(transport.is_connected());
491 transport.close().await.unwrap();
492 assert!(!transport.is_connected());
493 }
494 }
495
496 #[tokio::test]
497 async fn test_stdio_transport_spawn_with_args() {
498 let args = vec!["--version".to_string()];
499 let result = StdioTransport::spawn("cat", &args, &HashMap::new()).await;
500 let _ = result;
502 }
503
504 #[tokio::test]
505 async fn test_stdio_transport_spawn_with_env() {
506 let mut env = HashMap::new();
507 env.insert("TEST_VAR".to_string(), "test_value".to_string());
508 let result = StdioTransport::spawn("cat", &[], &env).await;
509 if let Ok(transport) = result {
510 let _ = transport.close().await;
511 }
512 }
513
514 #[tokio::test]
515 async fn test_stdio_transport_double_close() {
516 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
517 if let Ok(transport) = result {
518 transport.close().await.unwrap();
519 let result = transport.close().await;
521 assert!(result.is_ok());
522 }
523 }
524
525 #[tokio::test]
526 async fn test_stdio_transport_request_after_close() {
527 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
528 if let Ok(transport) = result {
529 transport.close().await.unwrap();
530
531 let request = JsonRpcRequest::new(1, "test", None);
532 let result = transport.request(request).await;
533 assert!(result.is_err());
534 assert!(result.unwrap_err().to_string().contains("not connected"));
535 }
536 }
537
538 #[tokio::test]
539 async fn test_stdio_transport_notify_after_close() {
540 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
541 if let Ok(transport) = result {
542 transport.close().await.unwrap();
543
544 let notification = JsonRpcNotification::new("test", None);
545 let result = transport.notify(notification).await;
546 assert!(result.is_err());
547 assert!(result.unwrap_err().to_string().contains("not connected"));
548 }
549 }
550
551 #[test]
552 fn test_json_rpc_request_creation() {
553 let request =
554 JsonRpcRequest::new(1, "test_method", Some(serde_json::json!({"key": "value"})));
555 assert_eq!(request.id, 1);
556 assert_eq!(request.method, "test_method");
557 assert!(request.params.is_some());
558 }
559
560 #[test]
561 fn test_json_rpc_notification_creation() {
562 let notification = JsonRpcNotification::new("test_notification", None);
563 assert_eq!(notification.method, "test_notification");
564 assert!(notification.params.is_none());
565 }
566
567 #[tokio::test]
568 async fn test_stdio_transport_custom_timeout() {
569 let result = StdioTransport::spawn_with_timeout("cat", &[], &HashMap::new(), 1).await;
571 if let Ok(transport) = result {
572 assert_eq!(transport.request_timeout_secs, 1);
573 let _ = transport.close().await;
574 }
575 }
576
577 #[tokio::test]
578 async fn test_stdio_transport_default_timeout() {
579 let result = StdioTransport::spawn("cat", &[], &HashMap::new()).await;
580 if let Ok(transport) = result {
581 assert_eq!(transport.request_timeout_secs, DEFAULT_REQUEST_TIMEOUT_SECS);
582 let _ = transport.close().await;
583 }
584 }
585
586 #[cfg(unix)]
587 #[tokio::test]
588 async fn close_kills_the_entire_mcp_process_group() {
589 let directory = tempfile::tempdir().unwrap();
590 let started = directory.path().join("started");
591 let leaked = directory.path().join("close-leak");
592 let transport = spawn_descendant_writer(&started, &leaked).await;
593 wait_for_path(&started).await;
594
595 transport.close().await.unwrap();
596 tokio::time::sleep(Duration::from_millis(400)).await;
597
598 assert!(
599 !leaked.exists(),
600 "closing an MCP transport must kill server descendants"
601 );
602 }
603
604 #[cfg(unix)]
605 #[tokio::test]
606 async fn drop_kills_the_entire_mcp_process_group() {
607 let directory = tempfile::tempdir().unwrap();
608 let started = directory.path().join("started");
609 let leaked = directory.path().join("drop-leak");
610 let transport = spawn_descendant_writer(&started, &leaked).await;
611 wait_for_path(&started).await;
612
613 drop(transport);
614 tokio::time::sleep(Duration::from_millis(400)).await;
615
616 assert!(
617 !leaked.exists(),
618 "dropping an MCP transport must kill server descendants"
619 );
620 }
621
622 #[cfg(unix)]
623 #[tokio::test]
624 async fn protocol_eof_reaps_a_still_running_server_tree() {
625 let directory = tempfile::tempdir().unwrap();
626 let descendant_started = directory.path().join("descendant-started");
627 let leaked = directory.path().join("protocol-eof-leak");
628 let args = vec![
629 "-c".to_string(),
630 "(: > \"$1\"; sleep 0.30; : > \"$2\") >/dev/null 2>&1 & \
631 while [ ! -e \"$1\" ]; do :; done; exec 1>&- 2>&-; wait"
632 .to_string(),
633 "mcp-protocol-eof-test".to_string(),
634 descendant_started.to_string_lossy().into_owned(),
635 leaked.to_string_lossy().into_owned(),
636 ];
637 let transport = StdioTransport::spawn("/bin/sh", &args, &HashMap::new())
638 .await
639 .unwrap();
640 wait_for_path(&descendant_started).await;
641 tokio::time::timeout(Duration::from_secs(1), async {
642 while transport.is_connected() {
643 tokio::task::yield_now().await;
644 }
645 })
646 .await
647 .expect("protocol EOF did not disconnect the MCP transport");
648
649 transport.close().await.unwrap();
650 tokio::time::sleep(Duration::from_millis(400)).await;
651
652 assert!(
653 !leaked.exists(),
654 "protocol EOF must reap the MCP server and every descendant"
655 );
656 }
657}