Skip to main content

rust_mcp_transport/
client_streamable_http.rs

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    /// Optional hook invoked with the raw JSON-RPC payload of each outgoing
44    /// POST, returning extra HTTP headers to attach to that single request.
45    ///
46    /// Unlike [`custom_headers`](Self::custom_headers) (static, applied to
47    /// every request), this enables per-request headers computed from the
48    /// message being sent — e.g. SEP-2243 `Mcp-Param-*` tool-parameter
49    /// mirroring, rotating bearer tokens, or distributed-tracing headers.
50    pub request_header_provider: Option<RequestHeaderProvider>,
51}
52
53/// Hook computing extra HTTP headers for a given outgoing POST payload
54/// (e.g. SEP-2243 `Mcp-Param-*` mirroring, rotating bearer tokens).
55pub 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    /// Optional cancellation token source for shutting down the transport
76    shutdown_source: tokio::sync::RwLock<Option<CancellationTokenSource>>,
77    /// Flag indicating if the transport is shut down
78    is_shut_down: Mutex<bool>,
79    /// Timeout duration for MCP messages
80    request_timeout: Duration,
81    /// Maximum line length for incoming messages
82    max_line_length: usize,
83    /// Capacity of the incoming-message channel buffer
84    channel_capacity: usize,
85    /// HTTP client for making requests
86    client: Client,
87    /// URL for the SSE endpoint
88    mcp_server_url: String,
89    /// Delay between retry attempts
90    retry_delay: Duration,
91    /// Maximum number of retry attempts
92    max_retries: usize,
93    /// Optional custom HTTP headers
94    custom_headers: Option<HeaderMap>,
95    /// Optional hook computing extra headers for each outgoing POST payload
96    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
103/// Merge the static `custom_headers` with any headers computed by the
104/// `request_header_provider` for the given outgoing payload. Provider
105/// headers are applied last, so they take precedence on name conflicts.
106fn 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        // Create CancellationTokenSource and token
203        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        // Initiate a task to process POST requests from messages received via the writable stream.
238        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                        // trim the trailing \n before making a request
249                        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);// Ignore error if receiver dropped
253                    },
254                    None => break, // Exit if channel is closed
255                    }
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        // Create readable stream
264        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        // Trigger cancellation
299        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; // Clear cancellation_source
304
305        // Mark as shut down
306        let mut is_shut_down_lock = self.is_shut_down.lock().await;
307        *is_shut_down_lock = true;
308
309        // Get task handle
310        let post_task = self.post_task.write().await.take();
311
312        // // Wait for tasks to complete with a timeout
313        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 // result of task completion
324            }
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}