diff --git a/crates/goose-cli/src/commands/session.rs b/crates/goose-cli/src/commands/session.rs index 23e3a2acf8..b6d39638e7 100644 --- a/crates/goose-cli/src/commands/session.rs +++ b/crates/goose-cli/src/commands/session.rs @@ -9,6 +9,8 @@ use crate::config::Config; use crate::prompt::rustyline::RustylinePrompt; use crate::prompt::Prompt; use crate::session::{ensure_session_dir, get_most_recent_session, Session}; +use goose::agents::extension::ExtensionError; +use mcp_client::transport::Error as McpClientError; /// Get the provider and model to use, following priority: /// 1. CLI arguments @@ -91,7 +93,17 @@ pub async fn build_session<'a>( agent .add_extension(extension_config.clone()) .await - .unwrap_or_else(|_| panic!("Failed to start extension: {}", name)); + .unwrap_or_else(|e| { + let err = match e { + ExtensionError::Transport(McpClientError::StdioProcessError(inner)) => { + inner + } + _ => e.to_string(), + }; + println!("Failed to start extension: {}, {:?}", name, err); + println!("Please check extension configuration for {}.", name); + process::exit(1); + }); } } diff --git a/crates/goose/src/agents/capabilities.rs b/crates/goose/src/agents/capabilities.rs index b986737b5d..f6b888b57b 100644 --- a/crates/goose/src/agents/capabilities.rs +++ b/crates/goose/src/agents/capabilities.rs @@ -143,7 +143,7 @@ impl Capabilities { let init_result = client .initialize(info, capabilities) .await - .map_err(|_| ExtensionError::Initialization(config.clone()))?; + .map_err(|e| ExtensionError::Initialization(config.clone(), e))?; // Store instructions if provided if let Some(instructions) = init_result.instructions { diff --git a/crates/goose/src/agents/extension.rs b/crates/goose/src/agents/extension.rs index 9ba5bfc596..cbcc489109 100644 --- a/crates/goose/src/agents/extension.rs +++ b/crates/goose/src/agents/extension.rs @@ -7,8 +7,8 @@ use thiserror::Error; /// Errors from Extension operation #[derive(Error, Debug)] pub enum ExtensionError { - #[error("Failed to start the MCP server from configuration `{0}` within 60 seconds")] - Initialization(ExtensionConfig), + #[error("Failed to start the MCP server from configuration `{0}` `{1}`")] + Initialization(ExtensionConfig, ClientError), #[error("Failed a client call to an MCP server: {0}")] Client(#[from] ClientError), #[error("User Message exceeded context-limit. History could not be truncated to accomodate.")] diff --git a/crates/mcp-client/src/transport/stdio.rs b/crates/mcp-client/src/transport/stdio.rs index 879421783d..5789640f60 100644 --- a/crates/mcp-client/src/transport/stdio.rs +++ b/crates/mcp-client/src/transport/stdio.rs @@ -1,11 +1,11 @@ use std::collections::HashMap; use std::sync::Arc; -use tokio::process::{Child, ChildStdin, ChildStdout, Command}; +use tokio::process::{Child, ChildStderr, ChildStdin, ChildStdout, Command}; use async_trait::async_trait; use mcp_core::protocol::JsonRpcMessage; -use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; -use tokio::sync::mpsc; +use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; +use tokio::sync::{mpsc, Mutex}; use super::{send_message, Error, PendingRequests, Transport, TransportHandle, TransportMessage}; @@ -16,20 +16,59 @@ pub struct StdioActor { receiver: mpsc::Receiver, pending_requests: Arc, _process: Child, // we store the process to keep it alive + error_sender: mpsc::Sender, stdin: ChildStdin, stdout: ChildStdout, + stderr: ChildStderr, } impl StdioActor { - pub async fn run(self) { - tokio::join!( - Self::handle_incoming_messages(self.stdout, self.pending_requests.clone()), - Self::handle_outgoing_messages( - self.receiver, - self.stdin, - self.pending_requests.clone() - ) + pub async fn run(mut self) { + use tokio::pin; + + let incoming = Self::handle_incoming_messages(self.stdout, self.pending_requests.clone()); + let outgoing = Self::handle_outgoing_messages( + self.receiver, + self.stdin, + self.pending_requests.clone(), ); + + // take ownership of futures for tokio::select + pin!(incoming); + pin!(outgoing); + + // Use select! to wait for either I/O completion or process exit + tokio::select! { + result = &mut incoming => { + tracing::debug!("Stdin handler completed: {:?}", result); + } + result = &mut outgoing => { + tracing::debug!("Stdout handler completed: {:?}", result); + } + // capture the status so we don't need to wait for a timeout + status = self._process.wait() => { + tracing::debug!("Process exited with status: {:?}", status); + } + } + + // Then always try to read stderr before cleaning up + let mut stderr_buffer = Vec::new(); + if let Ok(bytes) = self.stderr.read_to_end(&mut stderr_buffer).await { + let err_msg = if bytes > 0 { + String::from_utf8_lossy(&stderr_buffer).to_string() + } else { + "Process ended unexpectedly".to_string() + }; + + tracing::error!("Process stderr: {}", err_msg); + let _ = self + .error_sender + .send(Error::StdioProcessError(err_msg)) + .await; + } + + // Clean up regardless of which path we took + self.pending_requests.clear().await; } async fn handle_incoming_messages(stdout: ChildStdout, pending_requests: Arc) { @@ -38,7 +77,7 @@ impl StdioActor { loop { match reader.read_line(&mut line).await { Ok(0) => { - eprintln!("Child process ended (EOF on stdout)"); + tracing::error!("Child process ended (EOF on stdout)"); break; } // EOF Ok(_) => { @@ -69,11 +108,11 @@ impl StdioActor { mut stdin: ChildStdin, pending_requests: Arc, ) { - while let Some(transport_msg) = receiver.recv().await { + while let Some(mut transport_msg) = receiver.recv().await { let message_str = match serde_json::to_string(&transport_msg.message) { Ok(s) => s, Err(e) => { - if let Some(tx) = transport_msg.response_tx { + if let Some(tx) = transport_msg.response_tx.take() { let _ = tx.send(Err(Error::Serialization(e))); } continue; @@ -82,7 +121,7 @@ impl StdioActor { tracing::debug!(message = ?transport_msg.message, "Sending outgoing message"); - if let Some(response_tx) = transport_msg.response_tx { + if let Some(response_tx) = transport_msg.response_tx.take() { if let JsonRpcMessage::Request(request) = &transport_msg.message { if let Some(id) = &request.id { pending_requests.insert(id.to_string(), response_tx).await; @@ -111,12 +150,29 @@ impl StdioActor { #[derive(Clone)] pub struct StdioTransportHandle { sender: mpsc::Sender, + error_receiver: Arc>>, } #[async_trait::async_trait] impl TransportHandle for StdioTransportHandle { async fn send(&self, message: JsonRpcMessage) -> Result { - send_message(&self.sender, message).await + let result = send_message(&self.sender, message).await; + // Check for any pending errors even if send is successful + self.check_for_errors().await?; + result + } +} + +impl StdioTransportHandle { + /// Check if there are any process errors + pub async fn check_for_errors(&self) -> Result<(), Error> { + match self.error_receiver.lock().await.try_recv() { + Ok(error) => { + tracing::debug!("Found error: {:?}", error); + Err(error) + } + Err(_) => Ok(()), + } } } @@ -139,13 +195,13 @@ impl StdioTransport { } } - async fn spawn_process(&self) -> Result<(Child, ChildStdin, ChildStdout), Error> { + async fn spawn_process(&self) -> Result<(Child, ChildStdin, ChildStdout, ChildStderr), Error> { let mut process = Command::new(&self.command) .envs(&self.env) .args(&self.args) .stdin(std::process::Stdio::piped()) .stdout(std::process::Stdio::piped()) - .stderr(std::process::Stdio::inherit()) + .stderr(std::process::Stdio::piped()) .kill_on_drop(true) // 0 sets the process group ID equal to the process ID .process_group(0) // don't inherit signal handling from parent process @@ -162,7 +218,12 @@ impl StdioTransport { .take() .ok_or_else(|| Error::StdioProcessError("Failed to get stdout".into()))?; - Ok((process, stdin, stdout)) + let stderr = process + .stderr + .take() + .ok_or_else(|| Error::StdioProcessError("Failed to get stderr".into()))?; + + Ok((process, stdin, stdout, stderr)) } } @@ -171,20 +232,26 @@ impl Transport for StdioTransport { type Handle = StdioTransportHandle; async fn start(&self) -> Result { - let (process, stdin, stdout) = self.spawn_process().await?; + let (process, stdin, stdout, stderr) = self.spawn_process().await?; let (message_tx, message_rx) = mpsc::channel(32); + let (error_tx, error_rx) = mpsc::channel(1); let actor = StdioActor { receiver: message_rx, pending_requests: Arc::new(PendingRequests::new()), _process: process, + error_sender: error_tx, stdin, stdout, + stderr, }; tokio::spawn(actor.run()); - let handle = StdioTransportHandle { sender: message_tx }; + let handle = StdioTransportHandle { + sender: message_tx, + error_receiver: Arc::new(Mutex::new(error_rx)), + }; Ok(handle) }