tower_mcp/client/
stdio.rs1use std::process::Stdio;
19
20use async_trait::async_trait;
21use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
22use tokio::process::{Child, Command};
23
24use super::transport::ClientTransport;
25use crate::error::{Error, Result};
26
27pub struct StdioClientTransport {
33 child: Option<Child>,
34 stdin: Option<tokio::process::ChildStdin>,
35 stdout: BufReader<tokio::process::ChildStdout>,
36}
37
38impl StdioClientTransport {
39 pub async fn spawn(program: &str, args: &[&str]) -> Result<Self> {
46 let mut cmd = Command::new(program);
47 cmd.args(args);
48 Self::spawn_command(&mut cmd).await
49 }
50
51 pub async fn spawn_command(cmd: &mut Command) -> Result<Self> {
74 cmd.stdin(Stdio::piped())
75 .stdout(Stdio::piped())
76 .stderr(Stdio::inherit());
77
78 let mut child = cmd
79 .spawn()
80 .map_err(|e| Error::Transport(format!("Failed to spawn process: {}", e)))?;
81
82 let stdin = child
83 .stdin
84 .take()
85 .ok_or_else(|| Error::Transport("Failed to get child stdin".to_string()))?;
86 let stdout = child
87 .stdout
88 .take()
89 .ok_or_else(|| Error::Transport("Failed to get child stdout".to_string()))?;
90
91 tracing::info!("Spawned MCP server process");
92
93 Ok(Self {
94 child: Some(child),
95 stdin: Some(stdin),
96 stdout: BufReader::new(stdout),
97 })
98 }
99
100 pub fn from_child(mut child: Child) -> Result<Self> {
104 let stdin = child
105 .stdin
106 .take()
107 .ok_or_else(|| Error::Transport("Failed to get child stdin".to_string()))?;
108 let stdout = child
109 .stdout
110 .take()
111 .ok_or_else(|| Error::Transport("Failed to get child stdout".to_string()))?;
112
113 Ok(Self {
114 child: Some(child),
115 stdin: Some(stdin),
116 stdout: BufReader::new(stdout),
117 })
118 }
119}
120
121#[async_trait]
122impl ClientTransport for StdioClientTransport {
123 async fn send(&mut self, message: &str) -> Result<()> {
124 let stdin = self
125 .stdin
126 .as_mut()
127 .ok_or_else(|| Error::Transport("Transport closed".to_string()))?;
128
129 stdin
130 .write_all(message.as_bytes())
131 .await
132 .map_err(|e| Error::Transport(format!("Failed to write: {}", e)))?;
133 stdin
134 .write_all(b"\n")
135 .await
136 .map_err(|e| Error::Transport(format!("Failed to write newline: {}", e)))?;
137 stdin
138 .flush()
139 .await
140 .map_err(|e| Error::Transport(format!("Failed to flush: {}", e)))?;
141 Ok(())
142 }
143
144 async fn recv(&mut self) -> Result<Option<String>> {
145 let mut line = String::new();
146 let bytes = self
147 .stdout
148 .read_line(&mut line)
149 .await
150 .map_err(|e| Error::Transport(format!("Failed to read: {}", e)))?;
151
152 if bytes == 0 {
153 return Ok(None); }
155
156 Ok(Some(line.trim().to_string()))
157 }
158
159 fn is_connected(&self) -> bool {
160 self.child.is_some() && self.stdin.is_some()
161 }
162
163 async fn close(&mut self) -> Result<()> {
164 self.stdin.take();
166
167 if let Some(mut child) = self.child.take() {
168 let result =
169 tokio::time::timeout(std::time::Duration::from_secs(5), child.wait()).await;
170
171 match result {
172 Ok(Ok(status)) => {
173 tracing::info!(status = ?status, "Child process exited");
174 }
175 Ok(Err(e)) => {
176 tracing::error!(error = %e, "Error waiting for child");
177 }
178 Err(_) => {
179 tracing::warn!("Timeout waiting for child, killing");
180 let _ = child.kill().await;
181 }
182 }
183 }
184
185 Ok(())
186 }
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 #[tokio::test]
194 async fn test_spawn_nonexistent_program() {
195 let result = StdioClientTransport::spawn("nonexistent-program-xyz", &[]).await;
196 assert!(result.is_err());
197 }
198
199 #[tokio::test]
200 async fn test_send_and_recv_via_cat() {
201 let mut transport = StdioClientTransport::spawn("cat", &[]).await.unwrap();
203
204 assert!(transport.is_connected());
205
206 let msg = r#"{"jsonrpc":"2.0","id":1,"method":"test"}"#;
208 transport.send(msg).await.unwrap();
209
210 let received = transport.recv().await.unwrap();
212 assert_eq!(received.as_deref(), Some(msg));
213 }
214
215 #[tokio::test]
216 async fn test_close_signals_eof() {
217 let mut transport = StdioClientTransport::spawn("cat", &[]).await.unwrap();
218 assert!(transport.is_connected());
219
220 transport.close().await.unwrap();
221 assert!(!transport.is_connected());
222 }
223
224 #[tokio::test]
225 async fn test_recv_returns_none_on_eof() {
226 let mut transport = StdioClientTransport::spawn("true", &[]).await.unwrap();
228
229 let result = transport.recv().await.unwrap();
231 assert_eq!(result, None);
232 }
233
234 #[tokio::test]
235 async fn test_send_after_close_fails() {
236 let mut transport = StdioClientTransport::spawn("cat", &[]).await.unwrap();
237 transport.close().await.unwrap();
238
239 let result = transport.send("hello").await;
240 assert!(result.is_err());
241 }
242
243 #[tokio::test]
244 async fn test_spawn_command_with_env() {
245 let mut cmd = Command::new("sh");
246 cmd.args(["-c", "echo $TEST_VAR"]);
247 cmd.env("TEST_VAR", "hello_from_test");
248
249 let mut transport = StdioClientTransport::spawn_command(&mut cmd).await.unwrap();
250
251 let received = transport.recv().await.unwrap();
252 assert_eq!(received.as_deref(), Some("hello_from_test"));
253 }
254
255 #[tokio::test]
256 async fn test_multiple_send_recv_roundtrips() {
257 let mut transport = StdioClientTransport::spawn("cat", &[]).await.unwrap();
258
259 for i in 0..5 {
260 let msg = format!(r#"{{"id":{i},"msg":"test"}}"#);
261 transport.send(&msg).await.unwrap();
262 let received = transport.recv().await.unwrap();
263 assert_eq!(received.as_deref(), Some(msg.as_str()));
264 }
265
266 transport.close().await.unwrap();
267 }
268}