Skip to main content

soma_mcp_client/upstream/
http_body_cap.rs

1use std::{borrow::Cow, collections::HashMap, sync::Arc};
2
3use futures::{stream::BoxStream, StreamExt};
4use http::{HeaderName, HeaderValue};
5use reqwest::header::{ACCEPT, WWW_AUTHENTICATE};
6use rmcp::{
7    model::{ClientJsonRpcMessage, JsonRpcMessage, ServerJsonRpcMessage},
8    transport::{
9        common::http_header::{
10            EVENT_STREAM_MIME_TYPE, HEADER_LAST_EVENT_ID, HEADER_MCP_PROTOCOL_VERSION,
11            HEADER_SESSION_ID, JSON_MIME_TYPE,
12        },
13        streamable_http_client::{
14            AuthRequiredError, InsufficientScopeError, SseError, StreamableHttpClient,
15            StreamableHttpError, StreamableHttpPostResponse,
16        },
17    },
18};
19use sse_stream::{Sse, SseStream};
20
21#[derive(Clone)]
22pub struct BodyCappedHttpClient {
23    inner: reqwest::Client,
24    json_max_bytes: usize,
25    sse_event_max_bytes: usize,
26}
27
28impl BodyCappedHttpClient {
29    #[must_use]
30    pub fn new(inner: reqwest::Client, json_max_bytes: usize, sse_event_max_bytes: usize) -> Self {
31        Self {
32            inner,
33            json_max_bytes,
34            sse_event_max_bytes,
35        }
36    }
37
38    #[must_use]
39    pub fn default_with_caps(json_max_bytes: usize, sse_event_max_bytes: usize) -> Self {
40        let inner = reqwest::Client::builder()
41            .pool_max_idle_per_host(0)
42            .redirect(reqwest::redirect::Policy::none())
43            .build()
44            .expect("failed to build gateway HTTP client");
45        Self::new(inner, json_max_bytes, sse_event_max_bytes)
46    }
47}
48
49impl StreamableHttpClient for BodyCappedHttpClient {
50    type Error = reqwest::Error;
51
52    async fn post_message(
53        &self,
54        uri: Arc<str>,
55        message: ClientJsonRpcMessage,
56        session_id: Option<Arc<str>>,
57        auth_token: Option<String>,
58        custom_headers: HashMap<HeaderName, HeaderValue>,
59    ) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
60        let session_was_attached = session_id.is_some();
61        let mut request = self
62            .inner
63            .post(uri.as_ref())
64            .header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "));
65        if let Some(token) = auth_token {
66            request = request.bearer_auth(token);
67        }
68        if let Some(session_id) = session_id {
69            request = request.header(HEADER_SESSION_ID, session_id.as_ref());
70        }
71        let response = apply_custom_headers(request, custom_headers)?
72            .json(&message)
73            .send()
74            .await
75            .map_err(StreamableHttpError::Client)?;
76        response_to_post_result(response, message, session_was_attached, self.json_max_bytes).await
77    }
78
79    async fn delete_session(
80        &self,
81        uri: Arc<str>,
82        session_id: Arc<str>,
83        auth_token: Option<String>,
84        custom_headers: HashMap<HeaderName, HeaderValue>,
85    ) -> Result<(), StreamableHttpError<Self::Error>> {
86        let mut request = self
87            .inner
88            .delete(uri.as_ref())
89            .header(HEADER_SESSION_ID, session_id.as_ref());
90        if let Some(token) = auth_token {
91            request = request.bearer_auth(token);
92        }
93        let response = apply_custom_headers(request, custom_headers)?
94            .send()
95            .await
96            .map_err(StreamableHttpError::Client)?;
97        if response.status() == reqwest::StatusCode::METHOD_NOT_ALLOWED {
98            return Err(StreamableHttpError::ServerDoesNotSupportDeleteSession);
99        }
100        response
101            .error_for_status()
102            .map(|_| ())
103            .map_err(StreamableHttpError::Client)
104    }
105
106    async fn get_stream(
107        &self,
108        uri: Arc<str>,
109        session_id: Arc<str>,
110        last_event_id: Option<String>,
111        auth_token: Option<String>,
112        custom_headers: HashMap<HeaderName, HeaderValue>,
113    ) -> Result<BoxStream<'static, Result<Sse, SseError>>, StreamableHttpError<Self::Error>> {
114        let mut request = self
115            .inner
116            .get(uri.as_ref())
117            .header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "))
118            .header(HEADER_SESSION_ID, session_id.as_ref());
119        if let Some(last_event_id) = last_event_id {
120            request = request.header(HEADER_LAST_EVENT_ID, last_event_id);
121        }
122        if let Some(token) = auth_token {
123            request = request.bearer_auth(token);
124        }
125        let response = apply_custom_headers(request, custom_headers)?
126            .send()
127            .await
128            .map_err(StreamableHttpError::Client)?;
129        if response.status() == reqwest::StatusCode::METHOD_NOT_ALLOWED {
130            return Err(StreamableHttpError::ServerDoesNotSupportSse);
131        }
132        let response = response
133            .error_for_status()
134            .map_err(StreamableHttpError::Client)?;
135        ensure_stream_content_type(&response)?;
136        let capped = per_event_capped_stream(response.bytes_stream(), self.sse_event_max_bytes);
137        Ok(SseStream::from_bytes_stream(capped).boxed())
138    }
139}
140
141fn apply_custom_headers(
142    mut request: reqwest::RequestBuilder,
143    custom_headers: HashMap<HeaderName, HeaderValue>,
144) -> Result<reqwest::RequestBuilder, StreamableHttpError<reqwest::Error>> {
145    for (name, value) in custom_headers {
146        validate_custom_header(&name).map_err(StreamableHttpError::ReservedHeaderConflict)?;
147        request = request.header(name, value);
148    }
149    Ok(request)
150}
151
152fn validate_custom_header(name: &HeaderName) -> Result<(), String> {
153    let reserved = [
154        "accept",
155        HEADER_SESSION_ID,
156        HEADER_LAST_EVENT_ID,
157        HEADER_MCP_PROTOCOL_VERSION,
158    ];
159    if reserved
160        .iter()
161        .any(|reserved| name.as_str().eq_ignore_ascii_case(reserved))
162        && !name
163            .as_str()
164            .eq_ignore_ascii_case(HEADER_MCP_PROTOCOL_VERSION)
165    {
166        return Err(name.to_string());
167    }
168    Ok(())
169}
170
171async fn response_to_post_result(
172    response: reqwest::Response,
173    message: ClientJsonRpcMessage,
174    session_was_attached: bool,
175    max_bytes: usize,
176) -> Result<StreamableHttpPostResponse, StreamableHttpError<reqwest::Error>> {
177    if let Some(error) = auth_error(&response) {
178        return Err(error);
179    }
180    let status = response.status();
181    if matches!(
182        status,
183        reqwest::StatusCode::ACCEPTED | reqwest::StatusCode::NO_CONTENT
184    ) {
185        return Ok(StreamableHttpPostResponse::Accepted);
186    }
187    if status == reqwest::StatusCode::NOT_FOUND && session_was_attached {
188        return Err(StreamableHttpError::SessionExpired);
189    }
190    let content_type = content_type(&response);
191    let session_id = response
192        .headers()
193        .get(HEADER_SESSION_ID)
194        .and_then(|value| value.to_str().ok())
195        .map(ToOwned::to_owned);
196    if status.is_success() && response.content_length() == Some(0) && is_empty_ok(&message) {
197        return Ok(StreamableHttpPostResponse::Accepted);
198    }
199    if !status.is_success() {
200        return non_success_response(status, content_type, response, max_bytes).await;
201    }
202    match content_type.as_deref() {
203        Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => {
204            let capped = per_event_capped_stream(response.bytes_stream(), max_bytes);
205            Ok(StreamableHttpPostResponse::Sse(
206                SseStream::from_bytes_stream(capped).boxed(),
207                session_id,
208            ))
209        }
210        Some(ct) if ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()) => {
211            let bytes = read_body_capped(response, max_bytes).await?;
212            match serde_json::from_slice::<ServerJsonRpcMessage>(&bytes) {
213                Ok(message) => Ok(StreamableHttpPostResponse::Json(message, session_id)),
214                Err(_) => Ok(StreamableHttpPostResponse::Accepted),
215            }
216        }
217        _ => Err(StreamableHttpError::UnexpectedContentType(content_type)),
218    }
219}
220
221fn auth_error(response: &reqwest::Response) -> Option<StreamableHttpError<reqwest::Error>> {
222    let header = response.headers().get(WWW_AUTHENTICATE)?.to_str().ok()?;
223    match response.status() {
224        reqwest::StatusCode::UNAUTHORIZED => Some(StreamableHttpError::AuthRequired(
225            AuthRequiredError::new(header.to_owned()),
226        )),
227        reqwest::StatusCode::FORBIDDEN => Some(StreamableHttpError::InsufficientScope(
228            InsufficientScopeError::new(header.to_owned(), extract_scope(header)),
229        )),
230        _ => None,
231    }
232}
233
234async fn non_success_response(
235    status: reqwest::StatusCode,
236    content_type: Option<String>,
237    response: reqwest::Response,
238    max_bytes: usize,
239) -> Result<StreamableHttpPostResponse, StreamableHttpError<reqwest::Error>> {
240    let bytes = read_body_capped(response, max_bytes).await?;
241    let body = String::from_utf8_lossy(&bytes);
242    if content_type
243        .as_deref()
244        .is_some_and(|ct| ct.as_bytes().starts_with(JSON_MIME_TYPE.as_bytes()))
245    {
246        if let Some(message) = parse_json_rpc_error(&body) {
247            return Ok(StreamableHttpPostResponse::Json(message, None));
248        }
249    }
250    Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned(
251        format!("HTTP {status}: {body}"),
252    )))
253}
254
255async fn read_body_capped(
256    response: reqwest::Response,
257    max_bytes: usize,
258) -> Result<Vec<u8>, StreamableHttpError<reqwest::Error>> {
259    if let Some(length) = response.content_length() {
260        if length > max_bytes as u64 {
261            return Err(too_large(format!(
262                "response_too_large: declared {length} bytes, max {max_bytes}"
263            )));
264        }
265    }
266    let mut stream = response.bytes_stream();
267    let mut bytes = Vec::new();
268    while let Some(chunk) = stream.next().await {
269        let chunk = chunk.map_err(StreamableHttpError::Client)?;
270        if bytes.len().saturating_add(chunk.len()) > max_bytes {
271            return Err(too_large(format!(
272                "response_too_large: streamed {} bytes, max {max_bytes}",
273                bytes.len() + chunk.len()
274            )));
275        }
276        bytes.extend_from_slice(&chunk);
277    }
278    Ok(bytes)
279}
280
281fn per_event_capped_stream(
282    inner: impl futures::Stream<Item = reqwest::Result<bytes::Bytes>> + Send + 'static,
283    max_bytes: usize,
284) -> BoxStream<'static, Result<bytes::Bytes, CappedStreamError>> {
285    inner
286        .scan((0usize, false), move |state, item| {
287            let result = match item {
288                Ok(chunk) => account_event_bytes(&chunk, state, max_bytes).map(|_| chunk),
289                Err(error) => Err(CappedStreamError::Reqwest(error)),
290            };
291            futures::future::ready(Some(result))
292        })
293        .boxed()
294}
295
296fn account_event_bytes(
297    chunk: &[u8],
298    state: &mut (usize, bool),
299    max_bytes: usize,
300) -> Result<(), CappedStreamError> {
301    for byte in chunk {
302        if state.1 && *byte == b'\n' {
303            state.0 = 0;
304            state.1 = false;
305            continue;
306        }
307        state.0 = state.0.saturating_add(1);
308        if state.0 > max_bytes {
309            return Err(CappedStreamError::TooLarge {
310                event_bytes: state.0,
311                max_bytes,
312            });
313        }
314        state.1 = *byte == b'\n';
315    }
316    Ok(())
317}
318
319#[derive(Debug)]
320enum CappedStreamError {
321    Reqwest(reqwest::Error),
322    TooLarge {
323        event_bytes: usize,
324        max_bytes: usize,
325    },
326}
327
328impl std::fmt::Display for CappedStreamError {
329    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
330        match self {
331            Self::Reqwest(error) => write!(f, "upstream stream error: {error}"),
332            Self::TooLarge {
333                event_bytes,
334                max_bytes,
335            } => write!(
336                f,
337                "response_too_large: single SSE event reached {event_bytes} bytes, max {max_bytes}"
338            ),
339        }
340    }
341}
342
343impl std::error::Error for CappedStreamError {}
344
345fn ensure_stream_content_type(
346    response: &reqwest::Response,
347) -> Result<(), StreamableHttpError<reqwest::Error>> {
348    match response.headers().get(reqwest::header::CONTENT_TYPE) {
349        Some(value) => {
350            let raw = value.as_bytes();
351            if raw.starts_with(EVENT_STREAM_MIME_TYPE.as_bytes())
352                || raw.starts_with(JSON_MIME_TYPE.as_bytes())
353            {
354                Ok(())
355            } else {
356                Err(StreamableHttpError::UnexpectedContentType(Some(
357                    String::from_utf8_lossy(raw).to_string(),
358                )))
359            }
360        }
361        None => Err(StreamableHttpError::UnexpectedContentType(None)),
362    }
363}
364
365fn content_type(response: &reqwest::Response) -> Option<String> {
366    response
367        .headers()
368        .get(reqwest::header::CONTENT_TYPE)
369        .map(|value| String::from_utf8_lossy(value.as_bytes()).to_string())
370}
371
372fn is_empty_ok(message: &ClientJsonRpcMessage) -> bool {
373    matches!(
374        message,
375        ClientJsonRpcMessage::Notification(_)
376            | ClientJsonRpcMessage::Response(_)
377            | ClientJsonRpcMessage::Error(_)
378    )
379}
380
381fn parse_json_rpc_error(body: &str) -> Option<ServerJsonRpcMessage> {
382    match serde_json::from_str::<ServerJsonRpcMessage>(body) {
383        Ok(message @ JsonRpcMessage::Error(_)) => Some(message),
384        _ => None,
385    }
386}
387
388fn extract_scope(header: &str) -> Option<String> {
389    let lower = header.to_ascii_lowercase();
390    let start = lower.find("scope=")? + "scope=".len();
391    let value = &header[start..];
392    if let Some(quoted) = value.strip_prefix('"') {
393        return quoted.split('"').next().map(ToOwned::to_owned);
394    }
395    let end = value
396        .find(|ch: char| ch == ',' || ch == ';' || ch.is_whitespace())
397        .unwrap_or(value.len());
398    (end > 0).then(|| value[..end].to_owned())
399}
400
401fn too_large(message: String) -> StreamableHttpError<reqwest::Error> {
402    StreamableHttpError::UnexpectedServerResponse(Cow::Owned(message))
403}
404
405#[cfg(test)]
406#[path = "http_body_cap_tests.rs"]
407mod tests;