返回 CodeWhale
wire.rs
根目录 / crates / tui / src / mcp / wire.rs
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
385 lines RUST