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;