1use crate::error::TransportError;
2use crate::mcp_stream::MCPStream;
3
4use crate::schema::{
5 schema_utils::{
6 ClientMessage, ClientMessages, McpMessage, MessageFromClient, SdkError, ServerMessage,
7 ServerMessages,
8 },
9 RequestId,
10};
11use crate::utils::{CancellationTokenSource, ReadableChannel, StreamableHttpStream};
12use crate::{error::TransportResult, IoStream, McpDispatch, MessageDispatcher, Transport};
13use crate::{TransportDispatcher, TransportOptions};
14use async_trait::async_trait;
15use bytes::Bytes;
16use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
17use reqwest::Client;
18use std::collections::HashMap;
19use std::pin::Pin;
20use std::{sync::Arc, time::Duration};
21use tokio::io::BufReader;
22use tokio::sync::oneshot::Sender;
23use tokio::sync::{mpsc, oneshot, Mutex};
24use tokio::task::JoinHandle;
25
26const DEFAULT_CHANNEL_CAPACITY: usize = 64;
27const DEFAULT_MAX_RETRY: usize = 5;
28const DEFAULT_RETRY_TIME_SECONDS: u64 = 1;
29const SHUTDOWN_TIMEOUT_SECONDS: u64 = 5;
30
31pub struct StreamableTransportOptions {
32 pub mcp_url: String,
33 pub request_options: RequestOptions,
34}
35
36pub struct RequestOptions {
37 pub request_timeout: Duration,
38 pub max_line_length: usize,
39 pub channel_capacity: usize,
40 pub retry_delay: Option<Duration>,
41 pub max_retries: Option<usize>,
42 pub custom_headers: Option<HashMap<String, String>>,
43 pub request_header_provider: Option<RequestHeaderProvider>,
51}
52
53pub type RequestHeaderProvider = std::sync::Arc<dyn Fn(&str) -> Option<HeaderMap> + Send + Sync>;
56
57impl Default for RequestOptions {
58 fn default() -> Self {
59 Self {
60 request_timeout: TransportOptions::default().timeout,
61 max_line_length: TransportOptions::default().max_line_length,
62 channel_capacity: TransportOptions::default().channel_capacity,
63 retry_delay: None,
64 max_retries: None,
65 custom_headers: None,
66 request_header_provider: None,
67 }
68 }
69}
70
71pub struct ClientStreamableTransport<R>
72where
73 R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
74{
75 shutdown_source: tokio::sync::RwLock<Option<CancellationTokenSource>>,
77 is_shut_down: Mutex<bool>,
79 request_timeout: Duration,
81 max_line_length: usize,
83 channel_capacity: usize,
85 client: Client,
87 mcp_server_url: String,
89 retry_delay: Duration,
91 max_retries: usize,
93 custom_headers: Option<HeaderMap>,
95 request_header_provider: Option<RequestHeaderProvider>,
97 post_task: tokio::sync::RwLock<Option<tokio::task::JoinHandle<()>>>,
98 message_sender: Arc<tokio::sync::RwLock<Option<MessageDispatcher<R>>>>,
99 error_stream: tokio::sync::RwLock<Option<IoStream>>,
100 pending_requests: Arc<Mutex<HashMap<RequestId, tokio::sync::oneshot::Sender<R>>>>,
101}
102
103fn merge_request_headers(
107 custom_headers: &Option<HeaderMap>,
108 provider: &Option<RequestHeaderProvider>,
109 payload: &str,
110) -> Option<HeaderMap> {
111 let mut headers = custom_headers.clone();
112 if let Some(provider) = provider {
113 if let Some(extra) = provider(payload) {
114 headers.get_or_insert_with(HeaderMap::new).extend(extra);
115 }
116 }
117 headers
118}
119
120impl<R> ClientStreamableTransport<R>
121where
122 R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
123{
124 pub fn new(options: &StreamableTransportOptions) -> TransportResult<Self> {
125 let client = Client::new();
126
127 let headers = match &options.request_options.custom_headers {
128 Some(h) => Some(Self::validate_headers(h)?),
129 None => None,
130 };
131
132 let mcp_server_url = options.mcp_url.to_owned();
133 Ok(Self {
134 shutdown_source: tokio::sync::RwLock::new(None),
135 is_shut_down: Mutex::new(false),
136 request_timeout: options.request_options.request_timeout,
137 max_line_length: options.request_options.max_line_length,
138 channel_capacity: options.request_options.channel_capacity,
139 client,
140 mcp_server_url,
141 retry_delay: options
142 .request_options
143 .retry_delay
144 .unwrap_or(Duration::from_secs(DEFAULT_RETRY_TIME_SECONDS)),
145 max_retries: options
146 .request_options
147 .max_retries
148 .unwrap_or(DEFAULT_MAX_RETRY),
149 post_task: tokio::sync::RwLock::new(None),
150 custom_headers: headers,
151 request_header_provider: options.request_options.request_header_provider.clone(),
152 message_sender: Arc::new(tokio::sync::RwLock::new(None)),
153 error_stream: tokio::sync::RwLock::new(None),
154 pending_requests: Arc::new(Mutex::new(HashMap::new())),
155 })
156 }
157
158 fn validate_headers(headers: &HashMap<String, String>) -> TransportResult<HeaderMap> {
159 let mut header_map = HeaderMap::new();
160 for (key, value) in headers {
161 let header_name =
162 key.parse::<HeaderName>()
163 .map_err(|e| TransportError::Configuration {
164 message: format!("Invalid header name: {e}"),
165 })?;
166 let header_value =
167 HeaderValue::from_str(value).map_err(|e| TransportError::Configuration {
168 message: format!("Invalid header value: {e}"),
169 })?;
170 header_map.insert(header_name, header_value);
171 }
172 Ok(header_map)
173 }
174
175 pub(crate) async fn set_message_sender(&self, sender: MessageDispatcher<R>) {
176 let mut lock = self.message_sender.write().await;
177 *lock = Some(sender);
178 }
179
180 pub(crate) async fn set_error_stream(
181 &self,
182 error_stream: Pin<Box<dyn tokio::io::AsyncRead + Send + Sync>>,
183 ) {
184 let mut lock = self.error_stream.write().await;
185 *lock = Some(IoStream::Readable(error_stream));
186 }
187}
188
189#[async_trait]
190impl<R, S, M, OR, OM> Transport<R, S, M, OR, OM> for ClientStreamableTransport<M>
191where
192 R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
193 S: McpMessage + Clone + Send + Sync + serde::Serialize + 'static,
194 M: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
195 OR: Clone + Send + Sync + serde::Serialize + 'static,
196 OM: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
197{
198 async fn start(&self) -> TransportResult<tokio_stream::wrappers::ReceiverStream<R>>
199 where
200 MessageDispatcher<M>: McpDispatch<R, OR, M, OM>,
201 {
202 let (cancellation_source, cancellation_token) = CancellationTokenSource::new();
204 let mut lock = self.shutdown_source.write().await;
205 *lock = Some(cancellation_source);
206
207 let (write_tx, mut write_rx): (
208 tokio::sync::mpsc::Sender<(
209 String,
210 tokio::sync::oneshot::Sender<crate::error::TransportResult<()>>,
211 )>,
212 tokio::sync::mpsc::Receiver<(
213 String,
214 tokio::sync::oneshot::Sender<crate::error::TransportResult<()>>,
215 )>,
216 ) = tokio::sync::mpsc::channel(DEFAULT_CHANNEL_CAPACITY);
217 let (read_tx, read_rx) = mpsc::channel::<Bytes>(DEFAULT_CHANNEL_CAPACITY);
218
219 let max_retries = self.max_retries;
220 let retry_delay = self.retry_delay;
221
222 let post_url = self.mcp_server_url.clone();
223 let custom_headers = self.custom_headers.clone();
224 let request_header_provider = self.request_header_provider.clone();
225 let cancellation_token_post = cancellation_token.clone();
226 let cancellation_token_sse = cancellation_token.clone();
227
228 let mut streamable_http = StreamableHttpStream {
229 client: self.client.clone(),
230 mcp_url: post_url,
231 max_retries,
232 retry_delay,
233 read_tx,
234 session_id: Arc::new(tokio::sync::RwLock::new(None)),
235 };
236
237 let post_task_handle = tokio::spawn(async move {
239 loop {
240 tokio::select! {
241 _ = cancellation_token_post.cancelled() =>
242 {
243 break;
244 },
245 data = write_rx.recv() => {
246 match data{
247 Some((data, ack_tx)) => {
248 let payload = data.trim().to_string();
250 let headers = merge_request_headers(&custom_headers, &request_header_provider, &payload);
251 let result = streamable_http.run(payload, &cancellation_token_sse, &headers).await;
252 let _ = ack_tx.send(result);},
254 None => break, }
256 }
257 }
258 }
259 });
260 let mut post_task_lock = self.post_task.write().await;
261 *post_task_lock = Some(post_task_handle);
262
263 let readable: Pin<Box<dyn tokio::io::AsyncRead + Send + Sync>> =
265 Box::pin(BufReader::new(ReadableChannel {
266 read_rx,
267 buffer: Bytes::new(),
268 }));
269
270 let (stream, sender, error_stream) = MCPStream::create_with_ack(
271 readable,
272 write_tx,
273 IoStream::Writable(Box::pin(tokio::io::stderr())),
274 self.pending_requests.clone(),
275 self.request_timeout,
276 self.max_line_length,
277 cancellation_token,
278 self.channel_capacity,
279 );
280
281 self.set_message_sender(sender).await;
282
283 if let IoStream::Readable(error_stream) = error_stream {
284 self.set_error_stream(error_stream).await;
285 }
286
287 Ok(stream)
288 }
289
290 fn message_sender(&self) -> Arc<tokio::sync::RwLock<Option<MessageDispatcher<M>>>> {
291 self.message_sender.clone() as _
292 }
293
294 fn error_stream(&self) -> &tokio::sync::RwLock<Option<IoStream>> {
295 &self.error_stream as _
296 }
297 async fn shut_down(&self) -> TransportResult<()> {
298 let mut cancellation_lock = self.shutdown_source.write().await;
300 if let Some(source) = cancellation_lock.as_ref() {
301 source.cancel()?;
302 }
303 *cancellation_lock = None; let mut is_shut_down_lock = self.is_shut_down.lock().await;
307 *is_shut_down_lock = true;
308
309 let post_task = self.post_task.write().await.take();
311
312 let timeout = Duration::from_secs(SHUTDOWN_TIMEOUT_SECONDS);
314 let shutdown_future = async {
315 if let Some(post_handle) = post_task {
316 let _ = post_handle.await;
317 }
318 Ok::<(), TransportError>(())
319 };
320
321 tokio::select! {
322 result = shutdown_future => {
323 result }
325 _ = tokio::time::sleep(timeout) => {
326 tracing::warn!("Shutdown timed out after {:?}", timeout);
327 Err(TransportError::ShutdownTimeout)
328 }
329 }
330 }
331 async fn is_shut_down(&self) -> bool {
332 let result = self.is_shut_down.lock().await;
333 *result
334 }
335 async fn consume_string_payload(&self, _: &str) -> TransportResult<()> {
336 Err(TransportError::Internal(
337 "Invalid invocation of consume_string_payload() function for ClientStreamableTransport"
338 .to_string(),
339 ))
340 }
341
342 async fn pending_request_tx(&self, request_id: &RequestId) -> Option<Sender<M>> {
343 let mut pending_requests = self.pending_requests.lock().await;
344 pending_requests.remove(request_id)
345 }
346
347 async fn keep_alive(
348 &self,
349 _: Duration,
350 _: oneshot::Sender<()>,
351 ) -> TransportResult<JoinHandle<()>> {
352 Err(TransportError::Internal(
353 "Invalid invocation of keep_alive() function for ClientStreamableTransport".to_string(),
354 ))
355 }
356}
357
358#[async_trait]
359impl McpDispatch<ServerMessages, ClientMessages, ServerMessage, ClientMessage>
360 for ClientStreamableTransport<ServerMessage>
361{
362 async fn send_message(
363 &self,
364 message: ClientMessages,
365 request_timeout: Option<Duration>,
366 ) -> TransportResult<Option<ServerMessages>> {
367 let sender = self.message_sender.read().await;
368
369 let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
370
371 sender.send_message(message, request_timeout).await
372 }
373
374 async fn send(
375 &self,
376 message: ClientMessage,
377 request_timeout: Option<Duration>,
378 ) -> TransportResult<Option<ServerMessage>> {
379 let sender = self.message_sender.read().await;
380
381 let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
382
383 sender.send(message, request_timeout).await
384 }
385
386 async fn write_str(&self, payload: &str, skip_store: bool) -> TransportResult<()> {
387 let sender = self.message_sender.read().await;
388 let sender = sender.as_ref().ok_or(SdkError::connection_closed())?;
389 sender.write_str(payload, skip_store).await
390 }
391}
392
393impl
394 TransportDispatcher<
395 ServerMessages,
396 MessageFromClient,
397 ServerMessage,
398 ClientMessages,
399 ClientMessage,
400 > for ClientStreamableTransport<ServerMessage>
401{
402}