Skip to main content

a3s_code_core/mcp/transport/
stdio.rs

1//! Stdio Transport for MCP
2//!
3//! Implements MCP transport over standard input/output for local process communication.
4
5use 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
23/// Default request timeout for MCP tool calls
24const 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
28/// Stdio transport for MCP servers
29pub struct StdioTransport {
30    /// Process group owned by the child. The synchronous guard is also the
31    /// drop-time backstop when an async close cannot run.
32    process_group: Arc<StdMutex<ProcessGroupGuard>>,
33    /// Monitor that owns and reaps the direct child.
34    process_task: StdMutex<Option<JoinHandle<std::io::Result<()>>>>,
35    /// Stdin/stdout/stderr workers. Async close awaits them so notification
36    /// EOF and pending-request settlement happen before close returns.
37    io_tasks: StdMutex<Vec<JoinHandle<()>>>,
38    /// Stdin writer
39    stdin_tx: mpsc::Sender<String>,
40    /// Pending requests (id -> response sender)
41    pending: Arc<RwLock<HashMap<u64, oneshot::Sender<JsonRpcResponse>>>>,
42    /// Notification receiver
43    notification_rx: RwLock<Option<mpsc::Receiver<McpNotification>>>,
44    /// Connected flag
45    connected: Arc<AtomicBool>,
46    /// Stops the stdin/stdout/stderr tasks during close and drop.
47    shutdown: CancellationToken,
48    /// Per-request timeout in seconds
49    request_timeout_secs: u64,
50}
51
52impl StdioTransport {
53    /// Create a new stdio transport by spawning a process
54    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    /// Create a new stdio transport with a custom request timeout
63    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        // Spawn the process
70        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        // Add environment variables
79        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        // Create channels
93        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        // Spawn stdin writer task
107        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        // Spawn stdout reader task
139        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                        // Try to parse as response
164                        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                        // Try to parse as notification
175                        if let Ok(notification) =
176                            serde_json::from_str::<JsonRpcNotification>(trimmed)
177                        {
178                            let mcp_notif = McpNotification::from_json_rpc(&notification);
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        // Create response channel
240        let (tx, rx) = oneshot::channel();
241        let request_id = request.id;
242
243        // Register pending request
244        {
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        // Serialize and send request
254        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        // Wait for response with timeout
261        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                // Channel closed — clean up pending entry
270                self.pending.write().await.remove(&request_id);
271                return Err(anyhow!("Response channel closed"));
272            }
273            Err(_) => {
274                // Timeout — clean up pending entry to prevent memory leak
275                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(&notification)? + "\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        // This is a bit awkward - we need to take ownership of the receiver
302        // In practice, this should only be called once
303        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    // A server may leave helpers alive even after its direct process exits.
396    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        // Use a simple command that exists on most systems
467        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        // If cat doesn't exist, that's fine - skip the test
475    }
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        // May fail depending on system, but should not panic
501        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            // Second close should not panic
520            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        // Spawn with a very short timeout (1 second)
570        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}