Skip to main content

rust_mcp_transport/
sse.rs

1use crate::schema::schema_utils::{
2    ClientMessage, ClientMessages, MessageFromServer, SdkError, ServerMessage, ServerMessages,
3};
4use crate::schema::RequestId;
5use async_trait::async_trait;
6use serde::de::DeserializeOwned;
7use std::collections::HashMap;
8use std::pin::Pin;
9use std::sync::Arc;
10use std::time::Duration;
11use tokio::io::{AsyncWriteExt, DuplexStream};
12use tokio::sync::oneshot::Sender;
13use tokio::sync::{oneshot, Mutex};
14use tokio::task::JoinHandle;
15use tokio::time::{self, Interval};
16
17use crate::error::{TransportError, TransportResult};
18use crate::mcp_stream::MCPStream;
19use crate::message_dispatcher::MessageDispatcher;
20use crate::transport::Transport;
21use crate::utils::{endpoint_with_session_id, CancellationTokenSource};
22use crate::{IoStream, McpDispatch, SessionId, TransportDispatcher, TransportOptions};
23
24pub struct SseTransport<R>
25where
26    R: Clone + Send + Sync + DeserializeOwned + 'static,
27{
28    shutdown_source: tokio::sync::RwLock<Option<CancellationTokenSource>>,
29    is_shut_down: Mutex<bool>,
30    read_write_streams: Mutex<Option<(DuplexStream, DuplexStream)>>,
31    receiver_tx: Mutex<DuplexStream>, // receiving string payload
32    options: Arc<TransportOptions>,
33    message_sender: Arc<tokio::sync::RwLock<Option<MessageDispatcher<R>>>>,
34    error_stream: tokio::sync::RwLock<Option<IoStream>>,
35    pending_requests: Arc<Mutex<HashMap<RequestId, tokio::sync::oneshot::Sender<R>>>>,
36}
37
38/// Server-Sent Events (SSE) transport implementation
39impl<R> SseTransport<R>
40where
41    R: Clone + Send + Sync + DeserializeOwned + 'static,
42{
43    /// Creates a new SseTransport instance
44    ///
45    /// Initializes the transport with provided read and write duplex streams and options.
46    ///
47    /// # Arguments
48    /// * `read_rx` - Duplex stream for receiving messages
49    /// * `write_tx` - Duplex stream for sending messages
50    /// * `receiver_tx` - Duplex stream for receiving string payload
51    /// * `options` - Shared transport configuration options
52    ///
53    /// # Returns
54    /// * `TransportResult<Self>` - The initialized transport or an error
55    pub fn new(
56        read_rx: DuplexStream,
57        write_tx: DuplexStream,
58        receiver_tx: DuplexStream,
59        options: Arc<TransportOptions>,
60    ) -> TransportResult<Self> {
61        Ok(Self {
62            read_write_streams: Mutex::new(Some((read_rx, write_tx))),
63            options,
64            shutdown_source: tokio::sync::RwLock::new(None),
65            is_shut_down: Mutex::new(false),
66            receiver_tx: Mutex::new(receiver_tx),
67            message_sender: Arc::new(tokio::sync::RwLock::new(None)),
68            error_stream: tokio::sync::RwLock::new(None),
69            pending_requests: Arc::new(Mutex::new(HashMap::new())),
70        })
71    }
72
73    pub fn message_endpoint(endpoint: &str, session_id: &SessionId) -> String {
74        endpoint_with_session_id(endpoint, session_id)
75    }
76
77    pub(crate) async fn set_message_sender(&self, sender: MessageDispatcher<R>) {
78        let mut lock = self.message_sender.write().await;
79        *lock = Some(sender);
80    }
81
82    pub(crate) async fn set_error_stream(
83        &self,
84        error_stream: Pin<Box<dyn tokio::io::AsyncWrite + Send + Sync>>,
85    ) {
86        let mut lock = self.error_stream.write().await;
87        *lock = Some(IoStream::Writable(error_stream));
88    }
89}
90
91#[async_trait]
92impl McpDispatch<ClientMessages, ServerMessages, ClientMessage, ServerMessage>
93    for SseTransport<ClientMessage>
94{
95    async fn send_message(
96        &self,
97        message: ServerMessages,
98        request_timeout: Option<Duration>,
99    ) -> TransportResult<Option<ClientMessages>> {
100        let sender = self.message_sender.read().await;
101        let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
102
103        sender.send_message(message, request_timeout).await
104    }
105
106    async fn send(
107        &self,
108        message: ServerMessage,
109        request_timeout: Option<Duration>,
110    ) -> TransportResult<Option<ClientMessage>> {
111        let sender = self.message_sender.read().await;
112        let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
113        sender.send(message, request_timeout).await
114    }
115
116    async fn write_str(&self, payload: &str, skip_store: bool) -> TransportResult<()> {
117        let sender = self.message_sender.read().await;
118        let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
119        sender.write_str(payload, skip_store).await
120    }
121}
122
123#[async_trait] //RSMX
124impl Transport<ClientMessages, MessageFromServer, ClientMessage, ServerMessages, ServerMessage>
125    for SseTransport<ClientMessage>
126{
127    /// Starts the transport, initializing streams and message dispatcher
128    ///
129    /// Sets up the MCP stream and dispatcher using the provided duplex streams.
130    ///
131    /// # Returns
132    /// * `TransportResult<(Pin<Box<dyn Stream<Item = R> + Send>>, MessageDispatcher<R>, IoStream)>`
133    ///   - The message stream, dispatcher, and error stream
134    ///
135    /// # Errors
136    /// * Returns `TransportError` if streams are already taken or not initialized
137    async fn start(&self) -> TransportResult<tokio_stream::wrappers::ReceiverStream<ClientMessages>>
138    where
139        MessageDispatcher<ClientMessage>:
140            McpDispatch<ClientMessages, ServerMessages, ClientMessage, ServerMessage>,
141    {
142        // Create CancellationTokenSource and token
143        let (cancellation_source, cancellation_token) = CancellationTokenSource::new();
144        let mut lock = self.shutdown_source.write().await;
145        *lock = Some(cancellation_source);
146
147        let mut lock = self.read_write_streams.lock().await;
148        let (read_rx, write_tx) = lock.take().ok_or_else(|| {
149            TransportError::Internal(
150                "SSE streams already taken or transport not initialized".to_string(),
151            )
152        })?;
153
154        let (stream, sender, error_stream) = MCPStream::create::<ClientMessages, ClientMessage>(
155            Box::pin(read_rx),
156            Mutex::new(Box::pin(write_tx)),
157            IoStream::Writable(Box::pin(tokio::io::stderr())),
158            self.pending_requests.clone(),
159            self.options.timeout,
160            self.options.max_line_length,
161            cancellation_token,
162            self.options.channel_capacity,
163        );
164
165        self.set_message_sender(sender).await;
166
167        if let IoStream::Writable(error_stream) = error_stream {
168            self.set_error_stream(error_stream).await;
169        }
170
171        Ok(stream)
172    }
173
174    /// Checks if the transport has been shut down
175    ///
176    /// # Returns
177    /// * `bool` - True if the transport is shut down, false otherwise
178    async fn is_shut_down(&self) -> bool {
179        let result = self.is_shut_down.lock().await;
180        *result
181    }
182
183    fn message_sender(&self) -> Arc<tokio::sync::RwLock<Option<MessageDispatcher<ClientMessage>>>> {
184        self.message_sender.clone() as _
185    }
186
187    fn error_stream(&self) -> &tokio::sync::RwLock<Option<IoStream>> {
188        &self.error_stream as _
189    }
190
191    async fn consume_string_payload(&self, payload: &str) -> TransportResult<()> {
192        let mut transmit = self.receiver_tx.lock().await;
193        transmit
194            .write_all(format!("{payload}\n").as_bytes())
195            .await?;
196        transmit.flush().await?;
197        Ok(())
198    }
199
200    /// Shuts down the transport, terminating tasks and signaling closure
201    ///
202    /// Cancels any running tasks and clears the cancellation source.
203    ///
204    /// # Returns
205    /// * `TransportResult<()>` - Ok if shutdown is successful, Err if cancellation fails
206    async fn shut_down(&self) -> TransportResult<()> {
207        // Trigger cancellation
208        let mut cancellation_lock = self.shutdown_source.write().await;
209        if let Some(source) = cancellation_lock.as_ref() {
210            source.cancel()?;
211        }
212        *cancellation_lock = None; // Clear cancellation_source
213
214        // Mark as shut down
215        let mut is_shut_down_lock = self.is_shut_down.lock().await;
216        *is_shut_down_lock = true;
217        Ok(())
218    }
219
220    async fn keep_alive(
221        &self,
222        interval: Duration,
223        disconnect_tx: oneshot::Sender<()>,
224    ) -> TransportResult<JoinHandle<()>> {
225        let sender = self.message_sender();
226
227        let handle = tokio::spawn(async move {
228            let mut interval: Interval = time::interval(interval);
229            interval.tick().await; // Skip the first immediate tick
230            loop {
231                interval.tick().await;
232                let sender = sender.read().await;
233                if let Some(sender) = sender.as_ref() {
234                    match sender.write_str("\n", true).await {
235                        Ok(_) => {}
236                        Err(TransportError::Io(error))
237                            if error.kind() == std::io::ErrorKind::BrokenPipe =>
238                        {
239                            let _ = disconnect_tx.send(());
240                            break;
241                        }
242                        _ => {}
243                    }
244                }
245            }
246        });
247        Ok(handle)
248    }
249    async fn pending_request_tx(&self, request_id: &RequestId) -> Option<Sender<ClientMessage>> {
250        let mut pending_requests = self.pending_requests.lock().await;
251        pending_requests.remove(request_id)
252    }
253}
254
255impl
256    TransportDispatcher<
257        ClientMessages,
258        MessageFromServer,
259        ClientMessage,
260        ServerMessages,
261        ServerMessage,
262    > for SseTransport<ClientMessage>
263{
264}