use crate::schema::schema_utils::{
ClientMessage, ClientMessages, MessageFromClient, MessageFromServer, SdkError, ServerMessage,
ServerMessages,
};
use crate::schema::RequestId;
use async_trait::async_trait;
use serde::de::DeserializeOwned;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::process::Command;
use tokio::sync::oneshot::Sender;
use tokio::sync::{oneshot, Mutex};
use tokio::task::JoinHandle;
use crate::error::{TransportError, TransportResult};
use crate::mcp_stream::MCPStream;
use crate::message_dispatcher::MessageDispatcher;
use crate::transport::Transport;
use crate::utils::CancellationTokenSource;
use crate::{IoStream, McpDispatch, TransportDispatcher, TransportOptions};
pub struct StdioTransport<R>
where
R: Clone + Send + Sync + DeserializeOwned + 'static,
{
command: Option<String>,
args: Option<Vec<String>>,
env: Option<HashMap<String, String>>,
options: TransportOptions,
shutdown_source: tokio::sync::RwLock<Option<CancellationTokenSource>>,
is_shut_down: Mutex<bool>,
message_sender: Arc<tokio::sync::RwLock<Option<MessageDispatcher<R>>>>,
error_stream: tokio::sync::RwLock<Option<IoStream>>,
pending_requests: Arc<Mutex<HashMap<RequestId, tokio::sync::oneshot::Sender<R>>>>,
}
impl<R> StdioTransport<R>
where
R: Clone + Send + Sync + DeserializeOwned + 'static,
{
pub fn new(options: TransportOptions) -> TransportResult<Self> {
Ok(Self {
args: None,
command: None,
env: None,
options,
shutdown_source: tokio::sync::RwLock::new(None),
is_shut_down: Mutex::new(false),
message_sender: Arc::new(tokio::sync::RwLock::new(None)),
error_stream: tokio::sync::RwLock::new(None),
pending_requests: Arc::new(Mutex::new(HashMap::new())),
})
}
pub fn create_with_server_launch<C: Into<String>>(
command: C,
args: Vec<String>,
env: Option<HashMap<String, String>>,
options: TransportOptions,
) -> TransportResult<Self> {
Ok(Self {
args: Some(args),
command: Some(command.into()),
env,
options,
shutdown_source: tokio::sync::RwLock::new(None),
is_shut_down: Mutex::new(false),
message_sender: Arc::new(tokio::sync::RwLock::new(None)),
error_stream: tokio::sync::RwLock::new(None),
pending_requests: Arc::new(Mutex::new(HashMap::new())),
})
}
fn launch_commands(&self) -> (String, Vec<std::string::String>) {
#[cfg(windows)]
{
let command = "cmd.exe".to_string();
let mut command_args = vec!["/c".to_string(), self.command.clone().unwrap_or_default()];
command_args.extend(self.args.clone().unwrap_or_default());
(command, command_args)
}
#[cfg(unix)]
{
let command = self.command.clone().unwrap_or_default();
let command_args = self.args.clone().unwrap_or_default();
(command, command_args)
}
}
pub(crate) async fn set_message_sender(&self, sender: MessageDispatcher<R>) {
let mut lock = self.message_sender.write().await;
*lock = Some(sender);
}
pub(crate) async fn set_error_stream(&self, error_stream: IoStream) {
let mut lock = self.error_stream.write().await;
*lock = Some(error_stream);
}
}
#[async_trait]
impl<R, S, M, OR, OM> Transport<R, S, M, OR, OM> for StdioTransport<M>
where
R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
S: Clone + Send + Sync + serde::Serialize + 'static,
M: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
OR: Clone + Send + Sync + serde::Serialize + 'static,
OM: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
{
async fn start(&self) -> TransportResult<tokio_stream::wrappers::ReceiverStream<R>>
where
MessageDispatcher<M>: McpDispatch<R, OR, M, OM>,
{
let (cancellation_source, cancellation_token) = CancellationTokenSource::new();
let mut lock = self.shutdown_source.write().await;
*lock = Some(cancellation_source);
if self.command.is_some() {
let (command_name, command_args) = self.launch_commands();
let mut command = Command::new(command_name);
command
.envs(self.env.as_ref().unwrap_or(&HashMap::new()))
.args(&command_args)
.stdout(std::process::Stdio::piped())
.stdin(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.kill_on_drop(true);
#[cfg(windows)]
command.creation_flags(0x08000000);
#[cfg(unix)]
command.process_group(0);
let mut process = command.spawn().map_err(TransportError::Io)?;
let stdin = process
.stdin
.take()
.ok_or_else(|| TransportError::Internal("Unable to retrieve stdin.".into()))?;
let stdout = process
.stdout
.take()
.ok_or_else(|| TransportError::Internal("Unable to retrieve stdout.".into()))?;
let stderr = process
.stderr
.take()
.ok_or_else(|| TransportError::Internal("Unable to retrieve stderr.".into()))?;
let pending_requests_clone = self.pending_requests.clone();
tokio::spawn(async move {
let _ = process.wait().await;
let mut pending_requests = pending_requests_clone.lock().await;
pending_requests.clear();
});
let (stream, sender, error_stream) = MCPStream::create(
Box::pin(stdout),
Mutex::new(Box::pin(stdin)),
IoStream::Readable(Box::pin(stderr)),
self.pending_requests.clone(),
self.options.timeout,
self.options.max_line_length,
cancellation_token,
self.options.channel_capacity,
);
self.set_message_sender(sender).await;
self.set_error_stream(error_stream).await;
Ok(stream)
} else {
let (stream, sender, error_stream) = MCPStream::create(
Box::pin(tokio::io::stdin()),
Mutex::new(Box::pin(tokio::io::stdout())),
IoStream::Writable(Box::pin(tokio::io::stderr())),
self.pending_requests.clone(),
self.options.timeout,
self.options.max_line_length,
cancellation_token,
self.options.channel_capacity,
);
self.set_message_sender(sender).await;
self.set_error_stream(error_stream).await;
Ok(stream)
}
}
async fn pending_request_tx(&self, request_id: &RequestId) -> Option<Sender<M>> {
let mut pending_requests = self.pending_requests.lock().await;
pending_requests.remove(request_id)
}
async fn is_shut_down(&self) -> bool {
let result = self.is_shut_down.lock().await;
*result
}
fn message_sender(&self) -> Arc<tokio::sync::RwLock<Option<MessageDispatcher<M>>>> {
self.message_sender.clone() as _
}
fn error_stream(&self) -> &tokio::sync::RwLock<Option<IoStream>> {
&self.error_stream as _
}
async fn consume_string_payload(&self, _payload: &str) -> TransportResult<()> {
Err(TransportError::Internal(
"Invalid invocation of consume_string_payload() function in StdioTransport".to_string(),
))
}
async fn keep_alive(
&self,
_interval: Duration,
_disconnect_tx: oneshot::Sender<()>,
) -> TransportResult<JoinHandle<()>> {
Err(TransportError::Internal(
"Invalid invocation of keep_alive() function for StdioTransport".to_string(),
))
}
async fn shut_down(&self) -> TransportResult<()> {
let mut cancellation_lock = self.shutdown_source.write().await;
if let Some(source) = cancellation_lock.as_ref() {
source.cancel()?;
}
*cancellation_lock = None;
let mut is_shut_down_lock = self.is_shut_down.lock().await;
*is_shut_down_lock = true;
Ok(())
}
}
#[async_trait]
impl McpDispatch<ClientMessages, ServerMessages, ClientMessage, ServerMessage>
for StdioTransport<ClientMessage>
{
async fn send_message(
&self,
message: ServerMessages,
request_timeout: Option<Duration>,
) -> TransportResult<Option<ClientMessages>> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.send_message(message, request_timeout).await
}
async fn send(
&self,
message: ServerMessage,
request_timeout: Option<Duration>,
) -> TransportResult<Option<ClientMessage>> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.send(message, request_timeout).await
}
async fn write_str(&self, payload: &str, skip_store: bool) -> TransportResult<()> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.write_str(payload, skip_store).await
}
}
impl
TransportDispatcher<
ClientMessages,
MessageFromServer,
ClientMessage,
ServerMessages,
ServerMessage,
> for StdioTransport<ClientMessage>
{
}
#[async_trait]
impl McpDispatch<ServerMessages, ClientMessages, ServerMessage, ClientMessage>
for StdioTransport<ServerMessage>
{
async fn send_message(
&self,
message: ClientMessages,
request_timeout: Option<Duration>,
) -> TransportResult<Option<ServerMessages>> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.send_message(message, request_timeout).await
}
async fn send(
&self,
message: ClientMessage,
request_timeout: Option<Duration>,
) -> TransportResult<Option<ServerMessage>> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.send(message, request_timeout).await
}
async fn write_str(&self, payload: &str, skip_store: bool) -> TransportResult<()> {
let sender = self.message_sender.read().await;
let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
sender.write_str(payload, skip_store).await
}
}
impl
TransportDispatcher<
ServerMessages,
MessageFromClient,
ServerMessage,
ClientMessages,
ClientMessage,
> for StdioTransport<ServerMessage>
{
}