| 1 | //! MCP wire-format helpers shared by the HTTP, SSE, streamable-HTTP, and |
| 2 | //! stdio transports: frame/response size ceilings, SSE event framing and |
| 3 | //! field parsing, and the error-text classifiers that decide whether a |
| 4 | //! failure is a stale session or a closed connection. |
| 5 | /// Hard ceiling on the SSE frame-assembly buffer. A server that never emits a |
| 6 | /// frame separator would otherwise grow it without bound (OOM DoS). |
| 7 | pub(crate) const MAX_SSE_FRAME_BYTES: usize = 8 * 1024 * 1024; |
| 8 | |
| 9 | /// Hard ceiling on a single MCP HTTP response body / stdio line. A misbehaving |
| 10 | /// or malicious server could otherwise stream an unbounded body (or a |
| 11 | /// newline-free multi-GB "line") and OOM the process at transport-read time, |
| 12 | /// before any transcript-level spillover applies. |
| 13 | pub(crate) const MAX_MCP_RESPONSE_BYTES: usize = 16 * 1024 * 1024; |
| 14 | |
| 15 | pub(crate) fn is_mcp_stale_session_body(body: &str) -> bool { |
| 16 | let body = body.to_ascii_lowercase(); |
| 17 | body.contains("session") && (body.contains("expired") || body.contains("invalid")) |
| 18 | } |
| 19 | |
| 20 | /// A transport-level refusal of the session id, raised only where the |
| 21 | /// HTTP layer turned the request away before handing it to the server's |
| 22 | /// method dispatch: a Streamable HTTP stale-session status, or a legacy SSE |
| 23 | /// POST rejected with a stale-session body. A JSON-RPC error response is |
| 24 | /// never this type: it answers the request id, so the server processed it. |
| 25 | #[derive(Debug)] |
| 26 | pub(crate) struct McpSessionRejected(pub(crate) String); |
| 27 | |
| 28 | impl std::fmt::Display for McpSessionRejected { |
| 29 | fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { |
| 30 | f.write_str(&self.0) |
| 31 | } |
| 32 | } |
| 33 | |
| 34 | impl std::error::Error for McpSessionRejected {} |
| 35 | |
| 36 | /// The transport refused the session id, so the server provably did not run |
| 37 | /// the request. This is the only failure after which a non-idempotent |
| 38 | /// `tools/call` may be replayed on a fresh connection. Typed, not matched on |
| 39 | /// text, so a tool error that merely mentions an expired session cannot |
| 40 | /// qualify. |
| 41 | pub(super) fn is_mcp_session_rejected_error(err: &anyhow::Error) -> bool { |
| 42 | err.downcast_ref::<McpSessionRejected>().is_some() |
| 43 | } |
| 44 | |
| 45 | /// The connection is unusable: either the server rejected the session id, |
| 46 | /// or the transport itself is gone (dead pipe/socket) rather than merely |
| 47 | /// idle. The connection must be rebuilt, but a request already written to |
| 48 | /// a transport that then died may have run, so this alone does not make a |
| 49 | /// `tools/call` safe to replay. |
| 50 | pub(super) fn is_mcp_connection_lost_error(err: &anyhow::Error) -> bool { |
| 51 | if is_mcp_stale_session_error(err) { |
| 52 | return true; |
| 53 | } |
| 54 | let lower = format!("{err:#}").to_ascii_lowercase(); |
| 55 | is_connection_closed_error_text(&lower) |
| 56 | } |
| 57 | |
| 58 | pub(super) fn is_mcp_stale_session_error(err: &anyhow::Error) -> bool { |
| 59 | let err = format!("{err:#}"); |
| 60 | let lower_err = err.to_ascii_lowercase(); |
| 61 | err.contains("MCP Streamable HTTP session expired") |
| 62 | || err.contains("MCP session expired") |
| 63 | || err.contains("SSE transport closed") |
| 64 | // The exact bail text of a stdio transport whose child died (the |
| 65 | // EOF arm of `StdioTransport::recv`); without this arm a dead-child |
| 66 | // error missed the drop→reconnect→retry path that SSE closes get. |
| 67 | || err.contains("Stdio transport closed") |
| 68 | || (err.contains("MCP SSE POST send failed") && is_connection_closed_error_text(&lower_err)) |
| 69 | || is_mcp_stale_session_body(&err) |
| 70 | } |
| 71 | |
| 72 | pub(super) fn is_connection_closed_error_text(err: &str) -> bool { |
| 73 | err.contains("connection closed") |
| 74 | || err.contains("connection reset") |
| 75 | || err.contains("broken pipe") |
| 76 | || err.contains("unexpected eof") |
| 77 | || err.contains("forcibly closed") |
| 78 | } |
| 79 | |
| 80 | pub(super) fn parse_sse_message_data(body: &str) -> Vec<Vec<u8>> { |
| 81 | let normalized = body.replace("\r\n", "\n"); |
| 82 | let mut messages = Vec::new(); |
| 83 | |
| 84 | for block in normalized.split("\n\n") { |
| 85 | let mut event_type = "message"; |
| 86 | let mut data = String::new(); |
| 87 | |
| 88 | for line in block.lines() { |
| 89 | if let Some(value) = sse_field_value(line, "event:") { |
| 90 | event_type = value; |
| 91 | } else if let Some(value) = sse_field_value(line, "data:") { |
| 92 | if !data.is_empty() { |
| 93 | data.push('\n'); |
| 94 | } |
| 95 | data.push_str(value); |
| 96 | } |
| 97 | } |
| 98 | |
| 99 | if event_type != "message" || data.trim().is_empty() { |
| 100 | continue; |
| 101 | } |
| 102 | |
| 103 | messages.push(data.trim().as_bytes().to_vec()); |
| 104 | } |
| 105 | |
| 106 | messages |
| 107 | } |
| 108 | |
| 109 | // Retained for tests; the SSE transport now uses the byte-oriented twin. |
| 110 | #[cfg(test)] |
| 111 | pub(super) fn find_sse_event_separator(buffer: &str) -> Option<(usize, usize)> { |
| 112 | match (buffer.find("\n\n"), buffer.find("\r\n\r\n")) { |
| 113 | (Some(lf), Some(crlf)) if crlf < lf => Some((crlf, 4)), |
| 114 | (Some(lf), _) => Some((lf, 2)), |
| 115 | (_, Some(crlf)) => Some((crlf, 4)), |
| 116 | _ => None, |
| 117 | } |
| 118 | } |
| 119 | |
| 120 | /// Byte-oriented twin of `find_sse_event_separator`. Used by the SSE |
| 121 | /// transport so it can accumulate RAW bytes and decode only complete event |
| 122 | /// blocks — a multi-byte UTF-8 char split across two network reads is never |
| 123 | /// corrupted to U+FFFD (the `\n`/`\r` separators are ASCII and can never fall |
| 124 | /// inside a multi-byte sequence). |
| 125 | pub(crate) fn find_sse_event_separator_bytes(buffer: &[u8]) -> Option<(usize, usize)> { |
| 126 | let lf = buffer.windows(2).position(|w| w == b"\n\n"); |
| 127 | let crlf = buffer.windows(4).position(|w| w == b"\r\n\r\n"); |
| 128 | match (lf, crlf) { |
| 129 | (Some(lf), Some(crlf)) if crlf < lf => Some((crlf, 4)), |
| 130 | (Some(lf), _) => Some((lf, 2)), |
| 131 | (_, Some(crlf)) => Some((crlf, 4)), |
| 132 | _ => None, |
| 133 | } |
| 134 | } |
| 135 | |
| 136 | pub(crate) fn sse_field_value<'a>(line: &'a str, field: &str) -> Option<&'a str> { |
| 137 | let value = line.strip_prefix(field)?; |
| 138 | Some(value.strip_prefix(' ').unwrap_or(value)) |
| 139 | } |
| 140 | |
| 141 | pub(crate) fn is_streamable_http_incompatible_status(status: reqwest::StatusCode) -> bool { |
| 142 | matches!( |
| 143 | status, |
| 144 | reqwest::StatusCode::NOT_FOUND |
| 145 | | reqwest::StatusCode::METHOD_NOT_ALLOWED |
| 146 | | reqwest::StatusCode::NOT_ACCEPTABLE |
| 147 | | reqwest::StatusCode::UNSUPPORTED_MEDIA_TYPE |
| 148 | | reqwest::StatusCode::NOT_IMPLEMENTED |
| 149 | ) |
| 150 | } |
| 151 | |
| 152 | pub(crate) fn is_streamable_http_stale_session_status( |
| 153 | status: reqwest::StatusCode, |
| 154 | body_excerpt: &str, |
| 155 | ) -> bool { |
| 156 | if status == reqwest::StatusCode::NOT_FOUND { |
| 157 | return true; |
| 158 | } |
| 159 | if status != reqwest::StatusCode::BAD_REQUEST && status != reqwest::StatusCode::UNAUTHORIZED { |
| 160 | return false; |
| 161 | } |
| 162 | let body = body_excerpt.to_ascii_lowercase(); |
| 163 | body.contains("session") && (body.contains("expired") || body.contains("invalid")) |
| 164 | } |
| 165 | |
| 166 | /// Continue one newline-terminated line in caller-owned `out`, aborting if it |
| 167 | /// exceeds `max` bytes. Cancellation retains consumed bytes; the caller clears |
| 168 | /// the buffer only after receiving a complete frame. Returns the total bytes |
| 169 | /// accumulated; 0 means EOF. |
| 170 | pub(crate) async fn read_line_capped<R>( |
| 171 | reader: &mut R, |
| 172 | out: &mut Vec<u8>, |
| 173 | max: usize, |
| 174 | ) -> std::io::Result<usize> |
| 175 | where |
| 176 | R: tokio::io::AsyncBufRead + Unpin, |
| 177 | { |
| 178 | use tokio::io::AsyncBufReadExt; |
| 179 | loop { |
| 180 | let (chunk, consumed, done) = { |
| 181 | let available = reader.fill_buf().await?; |
| 182 | if available.is_empty() { |
| 183 | (Vec::new(), 0usize, true) |
| 184 | } else if let Some(pos) = available.iter().position(|&b| b == b'\n') { |
| 185 | (available[..=pos].to_vec(), pos + 1, true) |
| 186 | } else { |
| 187 | (available.to_vec(), available.len(), false) |
| 188 | } |
| 189 | }; |
| 190 | if consumed > 0 { |
| 191 | reader.consume(consumed); |
| 192 | } |
| 193 | out.extend_from_slice(&chunk); |
| 194 | if out.len() > max { |
| 195 | return Err(std::io::Error::new( |
| 196 | std::io::ErrorKind::InvalidData, |
| 197 | format!("MCP stdio line exceeded {max} bytes"), |
| 198 | )); |
| 199 | } |
| 200 | if done { |
| 201 | break; |
| 202 | } |
| 203 | } |
| 204 | Ok(out.len()) |
| 205 | } |
| 206 | |
| 207 | #[cfg(test)] |
| 208 | mod read_cap_tests { |
| 209 | use super::read_line_capped; |
| 210 | |
| 211 | #[tokio::test] |
| 212 | async fn cancelled_partial_read_preserves_next_frame() { |
| 213 | use futures_util::FutureExt; |
| 214 | use tokio::io::AsyncWriteExt; |
| 215 | let (mut writer, reader) = tokio::io::duplex(4096); |
| 216 | let mut reader = tokio::io::BufReader::new(reader); |
| 217 | let prefix = br#"{"jsonrpc":"2.0","id":"1","result":"#; |
| 218 | writer.write_all(prefix).await.unwrap(); |
| 219 | let mut pending = Vec::new(); |
| 220 | // Poll through the consumed prefix to Pending, then drop the future. |
| 221 | assert!( |
| 222 | read_line_capped(&mut reader, &mut pending, 1024) |
| 223 | .now_or_never() |
| 224 | .is_none() |
| 225 | ); |
| 226 | assert_eq!(pending, prefix); |
| 227 | writer.write_all(b"null}\n").await.unwrap(); |
| 228 | read_line_capped(&mut reader, &mut pending, 1024) |
| 229 | .await |
| 230 | .unwrap(); |
| 231 | let first: serde_json::Value = |
| 232 | serde_json::from_slice(&std::mem::take(&mut pending)).unwrap(); |
| 233 | assert_eq!(first["id"], "1"); |
| 234 | writer |
| 235 | .write_all(b"{\"id\":\"2\",\"result\":true}\n") |
| 236 | .await |
| 237 | .unwrap(); |
| 238 | read_line_capped(&mut reader, &mut pending, 1024) |
| 239 | .await |
| 240 | .unwrap(); |
| 241 | let second: serde_json::Value = serde_json::from_slice(&pending).unwrap(); |
| 242 | assert_eq!(second["id"], "2"); |
| 243 | assert_eq!(second["result"], true); |
| 244 | } |
| 245 | |
| 246 | #[tokio::test] |
| 247 | async fn resumed_frame_still_enforces_cap_at_newline() { |
| 248 | use futures_util::FutureExt; |
| 249 | use tokio::io::AsyncWriteExt; |
| 250 | let (mut writer, reader) = tokio::io::duplex(4096); |
| 251 | let mut reader = tokio::io::BufReader::new(reader); |
| 252 | let mut pending = Vec::new(); |
| 253 | writer.write_all(b"1234").await.unwrap(); |
| 254 | assert!( |
| 255 | read_line_capped(&mut reader, &mut pending, 6) |
| 256 | .now_or_never() |
| 257 | .is_none() |
| 258 | ); |
| 259 | writer.write_all(b"567\n").await.unwrap(); |
| 260 | assert_eq!( |
| 261 | read_line_capped(&mut reader, &mut pending, 6) |
| 262 | .await |
| 263 | .unwrap_err() |
| 264 | .kind(), |
| 265 | std::io::ErrorKind::InvalidData |
| 266 | ); |
| 267 | } |
| 268 | |
| 269 | #[tokio::test] |
| 270 | async fn reads_a_line_and_reports_eof() { |
| 271 | let data = b"hello\nworld\n".to_vec(); |
| 272 | let mut reader = tokio::io::BufReader::new(std::io::Cursor::new(data)); |
| 273 | let mut out = Vec::new(); |
| 274 | assert_eq!( |
| 275 | read_line_capped(&mut reader, &mut out, 1024).await.unwrap(), |
| 276 | 6 |
| 277 | ); |
| 278 | assert_eq!(out, b"hello\n"); |
| 279 | out.clear(); |
| 280 | assert_eq!( |
| 281 | read_line_capped(&mut reader, &mut out, 1024).await.unwrap(), |
| 282 | 6 |
| 283 | ); |
| 284 | assert_eq!(out, b"world\n"); |
| 285 | out.clear(); |
| 286 | // EOF. |
| 287 | assert_eq!( |
| 288 | read_line_capped(&mut reader, &mut out, 1024).await.unwrap(), |
| 289 | 0 |
| 290 | ); |
| 291 | } |
| 292 | |
| 293 | #[tokio::test] |
| 294 | async fn aborts_on_newline_free_line_over_cap() { |
| 295 | let data = vec![b'x'; 4096]; // no newline |
| 296 | let mut reader = tokio::io::BufReader::new(std::io::Cursor::new(data)); |
| 297 | let mut out = Vec::new(); |
| 298 | let err = read_line_capped(&mut reader, &mut out, 1024) |
| 299 | .await |
| 300 | .unwrap_err(); |
| 301 | assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); |
| 302 | } |
| 303 | } |
| 304 | |
| 305 | pub(crate) fn resolve_sse_endpoint_url( |
| 306 | base_url: &str, |
| 307 | endpoint_url: &str, |
| 308 | ) -> anyhow::Result<String> { |
| 309 | let base = reqwest::Url::parse(base_url)?; |
| 310 | let resolved = if endpoint_url.starts_with("http://") || endpoint_url.starts_with("https://") { |
| 311 | reqwest::Url::parse(endpoint_url)? |
| 312 | } else { |
| 313 | base.join(endpoint_url)? |
| 314 | }; |
| 315 | // reqwest converts userinfo into Basic Authorization while building a |
| 316 | // request, before the request-time guard can inspect the original URL. |
| 317 | if !resolved.username().is_empty() || resolved.password().is_some() { |
| 318 | anyhow::bail!("MCP SSE endpoint must not contain URL credentials"); |
| 319 | } |
| 320 | // Security: the server-supplied `endpoint` event must stay same-origin |
| 321 | // as the connect URL. The connect host is vetted by network policy |
| 322 | // once, but the endpoint host is never re-checked — so an absolute |
| 323 | // cross-origin endpoint would let a malicious MCP server redirect the |
| 324 | // client's *authenticated* POSTs (Bearer/OAuth headers attached) to an |
| 325 | // internal host (169.254.169.254, localhost admin ports, …): an SSRF / |
| 326 | // policy bypass. Relative endpoints are same-origin by construction. |
| 327 | if resolved.scheme() != base.scheme() |
| 328 | || resolved.host_str() != base.host_str() |
| 329 | || resolved.port_or_known_default() != base.port_or_known_default() |
| 330 | { |
| 331 | anyhow::bail!( |
| 332 | "MCP SSE endpoint {} is not same-origin as {} — refusing to send \ |
| 333 | authenticated requests cross-origin", |
| 334 | super::mask_url_secrets(resolved.as_str()), |
| 335 | super::mask_url_secrets(base.as_str()), |
| 336 | ); |
| 337 | } |
| 338 | Ok(resolved.to_string()) |
| 339 | } |
| 340 | |
| 341 | #[cfg(test)] |
| 342 | mod endpoint_tests { |
| 343 | use super::resolve_sse_endpoint_url; |
| 344 | |
| 345 | #[test] |
| 346 | fn resolve_endpoint_accepts_relative_and_same_origin() { |
| 347 | let base = "https://mcp.example.com/v1/sse"; |
| 348 | // Relative path -> same origin. |
| 349 | assert_eq!( |
| 350 | resolve_sse_endpoint_url(base, "/messages?sid=1").unwrap(), |
| 351 | "https://mcp.example.com/messages?sid=1" |
| 352 | ); |
| 353 | // Absolute but same origin -> allowed. |
| 354 | assert_eq!( |
| 355 | resolve_sse_endpoint_url(base, "https://mcp.example.com/messages").unwrap(), |
| 356 | "https://mcp.example.com/messages" |
| 357 | ); |
| 358 | } |
| 359 | |
| 360 | #[test] |
| 361 | fn resolve_endpoint_rejects_cross_origin_ssrf() { |
| 362 | let base = "https://mcp.example.com/v1/sse"; |
| 363 | // Different host (metadata endpoint) -> rejected. |
| 364 | assert!(resolve_sse_endpoint_url(base, "http://169.254.169.254/latest").is_err()); |
| 365 | // Different scheme -> rejected. |
| 366 | assert!(resolve_sse_endpoint_url(base, "http://mcp.example.com/messages").is_err()); |
| 367 | // Different port -> rejected. |
| 368 | assert!(resolve_sse_endpoint_url(base, "https://mcp.example.com:8443/x").is_err()); |
| 369 | // Same-origin userinfo must not become an implicit Basic credential. |
| 370 | for endpoint in [ |
| 371 | "https://fixture-user:fixture-password@mcp.example.com/messages", |
| 372 | "//fixture-user@mcp.example.com/messages", |
| 373 | ] { |
| 374 | let error = resolve_sse_endpoint_url(base, endpoint).unwrap_err(); |
| 375 | assert!( |
| 376 | error |
| 377 | .to_string() |
| 378 | .contains("must not contain URL credentials") |
| 379 | ); |
| 380 | assert!(!error.to_string().contains("fixture-user")); |
| 381 | assert!(!error.to_string().contains("fixture-password")); |
| 382 | } |
| 383 | } |
| 384 | } |
| 385 |