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>, 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
38impl<R> SseTransport<R>
40where
41 R: Clone + Send + Sync + DeserializeOwned + 'static,
42{
43 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] impl Transport<ClientMessages, MessageFromServer, ClientMessage, ServerMessages, ServerMessage>
125 for SseTransport<ClientMessage>
126{
127 async fn start(&self) -> TransportResult<tokio_stream::wrappers::ReceiverStream<ClientMessages>>
138 where
139 MessageDispatcher<ClientMessage>:
140 McpDispatch<ClientMessages, ServerMessages, ClientMessage, ServerMessage>,
141 {
142 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 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 async fn shut_down(&self) -> TransportResult<()> {
207 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; 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; 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}