返回 CodeWhale
sse.rs
根目录 / crates / tui / src / mcp / sse.rs
1 use std::time::Duration;
2
3 use anyhow::{Context, Result};
4
5 use super::headers::{apply_safe_custom_headers, with_default_mcp_http_headers};
6 use super::{
7 ERROR_BODY_PREVIEW_BYTES, McpHttpAuth, McpTransport, bounded_body_excerpt,
8 find_sse_event_separator_bytes, is_mcp_stale_session_body, mask_url_secrets, sse_field_value,
9 };
10
11 const SSE_INBOUND_CHANNEL_CAPACITY: usize = 4;
12
13 pub(super) struct SseTransport {
14 pub(super) client: reqwest::Client,
15 pub(super) base_url: String,
16 pub(super) auth: McpHttpAuth,
17 pub(super) endpoint_url: Option<String>,
18 pub(super) receiver: tokio::sync::mpsc::Receiver<SseInbound>,
19 #[allow(dead_code)]
20 pub(super) sse_task: tokio::task::JoinHandle<()>,
21 }
22
23 pub(super) enum SseInbound {
24 Endpoint(String),
25 Message(Vec<u8>),
26 }
27
28 impl SseTransport {
29 pub(super) async fn connect(
30 client: reqwest::Client,
31 url: String,
32 auth: McpHttpAuth,
33 cancel_token: tokio_util::sync::CancellationToken,
34 endpoint_timeout: Duration,
35 ) -> Result<Self> {
36 let (tx, rx) = tokio::sync::mpsc::channel(SSE_INBOUND_CHANNEL_CAPACITY);
37 let client_clone = client.clone();
38 let url_clone = url.clone();
39 let auth_clone = auth.clone();
40 let wait_cancel_token = cancel_token.clone();
41
42 let sse_task = tokio::spawn(async move {
43 if cancel_token.is_cancelled() {
44 return;
45 }
46 use futures_util::FutureExt;
47 let result = std::panic::AssertUnwindSafe(Self::run_sse_loop(
48 client_clone,
49 url_clone,
50 auth_clone,
51 tx,
52 cancel_token,
53 ))
54 .catch_unwind()
55 .await;
56 match result {
57 Ok(res) => {
58 if let Err(e) = res {
59 tracing::error!("SSE loop error: {}", e);
60 }
61 }
62 Err(panic_err) => {
63 if let Some(msg) = panic_err.downcast_ref::<&str>() {
64 tracing::error!("SSE loop panicked: {}", msg);
65 } else if let Some(msg) = panic_err.downcast_ref::<String>() {
66 tracing::error!("SSE loop panicked: {}", msg);
67 } else {
68 tracing::error!("SSE loop panicked with unknown error");
69 }
70 }
71 }
72 });
73
74 let mut transport = Self {
75 client,
76 base_url: url,
77 auth,
78 endpoint_url: None,
79 receiver: rx,
80 sse_task,
81 };
82 transport
83 .wait_for_endpoint(&wait_cancel_token, endpoint_timeout)
84 .await?;
85 Ok(transport)
86 }
87
88 async fn run_sse_loop(
89 client: reqwest::Client,
90 url: String,
91 auth: McpHttpAuth,
92 tx: tokio::sync::mpsc::Sender<SseInbound>,
93 cancel_token: tokio_util::sync::CancellationToken,
94 ) -> Result<()> {
95 let headers = tokio::select! {
96 biased;
97 _ = cancel_token.cancelled() => {
98 anyhow::bail!("MCP SSE connect cancelled before authentication completed")
99 }
100 headers = auth.resolved_headers() => headers?,
101 };
102 let request = apply_safe_custom_headers(
103 with_default_mcp_http_headers(client.get(&url), false),
104 &headers,
105 );
106 let response = tokio::select! {
107 biased;
108 _ = cancel_token.cancelled() => {
109 anyhow::bail!("MCP SSE connect cancelled before the request completed")
110 }
111 response = request.send() => response.with_context(|| {
112 format!(
113 "MCP SSE connect failed (transport=http url={})",
114 mask_url_secrets(&url),
115 )
116 })?,
117 };
118 let status = response.status();
119 if !status.is_success() {
120 let body_excerpt = bounded_body_excerpt(response, ERROR_BODY_PREVIEW_BYTES).await;
121 let body_excerpt = auth.server_error_preview(&body_excerpt);
122 anyhow::bail!(
123 "MCP SSE rejected (transport=http url={} status={}): {}",
124 mask_url_secrets(&url),
125 status,
126 body_excerpt,
127 );
128 }
129
130 let mut stream = response.bytes_stream();
131 use futures_util::StreamExt;
132 // Raw byte buffer so a multi-byte UTF-8 char split across reads is not
133 // corrupted, and bounded so a separator-less server cannot OOM us.
134 let mut buffer: Vec<u8> = Vec::new();
135
136 loop {
137 if cancel_token.is_cancelled() {
138 tracing::debug!("SSE loop cancelled");
139 break;
140 }
141 let item = tokio::select! {
142 _ = cancel_token.cancelled() => {
143 tracing::debug!("SSE loop shutting down");
144 break;
145 }
146 item = stream.next() => {
147 match item {
148 Some(i) => i,
149 None => break,
150 }
151 }
152 };
153 let chunk = item?;
154 buffer.extend_from_slice(&chunk);
155 if buffer.len() > super::MAX_SSE_FRAME_BYTES {
156 anyhow::bail!(
157 "MCP SSE frame exceeded {} bytes without a separator — aborting",
158 super::MAX_SSE_FRAME_BYTES
159 );
160 }
161
162 while let Some((pos, separator_len)) = find_sse_event_separator_bytes(&buffer) {
163 // Complete block: decoding cannot split a multi-byte char.
164 let event_block = String::from_utf8_lossy(&buffer[..pos]).into_owned();
165 buffer.drain(..pos + separator_len);
166
167 let mut event_type = "message";
168 let mut data = String::new();
169
170 for line in event_block.lines() {
171 if let Some(value) = sse_field_value(line, "event:") {
172 event_type = value;
173 } else if let Some(value) = sse_field_value(line, "data:") {
174 if !data.is_empty() {
175 data.push('\n');
176 }
177 data.push_str(value);
178 }
179 }
180
181 let inbound = match event_type {
182 "endpoint" => Some(SseInbound::Endpoint(data)),
183 "message" if !data.trim().is_empty() => {
184 Some(SseInbound::Message(data.into_bytes()))
185 }
186 _ => None,
187 };
188 if let Some(inbound) = inbound {
189 let sent = tokio::select! {
190 biased;
191 _ = cancel_token.cancelled() => return Ok(()),
192 sent = tx.send(inbound) => sent,
193 };
194 if sent.is_err() {
195 return Ok(());
196 }
197 }
198 }
199 }
200 Ok(())
201 }
202
203 async fn wait_for_endpoint(
204 &mut self,
205 cancel_token: &tokio_util::sync::CancellationToken,
206 endpoint_timeout: Duration,
207 ) -> Result<()> {
208 let timeout = tokio::time::sleep(endpoint_timeout);
209 tokio::pin!(timeout);
210
211 let msg = tokio::select! {
212 _ = cancel_token.cancelled() => {
213 anyhow::bail!("SSE transport cancelled before endpoint was discovered");
214 }
215 _ = &mut timeout => {
216 anyhow::bail!(
217 "SSE endpoint not received within {}ms",
218 endpoint_timeout.as_millis()
219 );
220 }
221 msg = self.receiver.recv() => {
222 msg.context("SSE transport closed before endpoint was discovered")?
223 }
224 };
225
226 match msg {
227 SseInbound::Endpoint(endpoint) => self.store_endpoint(&endpoint),
228 SseInbound::Message(_) => {
229 anyhow::bail!("MCP SSE server sent a message before declaring its endpoint");
230 }
231 }
232 }
233
234 fn store_endpoint(&mut self, endpoint: &str) -> Result<()> {
235 self.endpoint_url = Some(Self::resolve_endpoint_url(&self.base_url, endpoint)?);
236 Ok(())
237 }
238
239 fn resolve_endpoint_url(base_url: &str, endpoint_url: &str) -> Result<String> {
240 let base = reqwest::Url::parse(base_url)?;
241 let resolved =
242 if endpoint_url.starts_with("http://") || endpoint_url.starts_with("https://") {
243 reqwest::Url::parse(endpoint_url)?
244 } else {
245 base.join(endpoint_url)?
246 };
247 // Security: the server-supplied `endpoint` event must stay same-origin
248 // as the connect URL. The connect host is vetted by network policy
249 // once, but the endpoint host is never re-checked — so an absolute
250 // cross-origin endpoint would let a malicious MCP server redirect the
251 // client's *authenticated* POSTs (Bearer/OAuth headers attached) to an
252 // internal host (169.254.169.254, localhost admin ports, …): an SSRF /
253 // policy bypass. Relative endpoints are same-origin by construction.
254 if resolved.scheme() != base.scheme()
255 || resolved.host_str() != base.host_str()
256 || resolved.port_or_known_default() != base.port_or_known_default()
257 {
258 anyhow::bail!(
259 "MCP SSE endpoint {} is not same-origin as {} — refusing to send \
260 authenticated requests cross-origin",
261 mask_url_secrets(resolved.as_str()),
262 mask_url_secrets(base.as_str()),
263 );
264 }
265 Ok(resolved.to_string())
266 }
267 }
268
269 #[async_trait::async_trait]
270 impl McpTransport for SseTransport {
271 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
272 let endpoint = self
273 .endpoint_url
274 .as_ref()
275 .context("SSE endpoint not yet discovered")?
276 .clone();
277 let headers = self.auth.resolved_headers().await?;
278 let response = apply_safe_custom_headers(
279 with_default_mcp_http_headers(self.client.post(&endpoint), true),
280 &headers,
281 )
282 .body(msg)
283 .send()
284 .await
285 .with_context(|| {
286 format!(
287 "MCP SSE POST send failed (transport=sse endpoint={})",
288 mask_url_secrets(&endpoint)
289 )
290 })?;
291 let status = response.status();
292 if !status.is_success() {
293 let body_excerpt = bounded_body_excerpt(response, ERROR_BODY_PREVIEW_BYTES).await;
294 let stale_session = is_mcp_stale_session_body(&body_excerpt);
295 let body_excerpt = self.auth.server_error_preview(&body_excerpt);
296 if stale_session {
297 anyhow::bail!(
298 "MCP session expired (transport=sse endpoint={} status={}): {}",
299 mask_url_secrets(&endpoint),
300 status,
301 body_excerpt
302 );
303 }
304 anyhow::bail!(
305 "MCP SSE POST rejected (transport=sse endpoint={} status={}): {}",
306 mask_url_secrets(&endpoint),
307 status,
308 body_excerpt
309 );
310 }
311 Ok(())
312 }
313
314 async fn recv(&mut self) -> Result<Vec<u8>> {
315 loop {
316 match self.receiver.recv().await.context("SSE transport closed")? {
317 SseInbound::Endpoint(endpoint) => {
318 self.store_endpoint(&endpoint)?;
319 }
320 SseInbound::Message(msg) => return Ok(msg),
321 }
322 }
323 }
324
325 async fn shutdown(&mut self) {
326 self.sse_task.abort();
327 }
328 }
329
330 impl Drop for SseTransport {
331 fn drop(&mut self) {
332 // Dropping a JoinHandle detaches its task. Abort explicitly so a
333 // cancelled connection cannot leave an auth refresh, connect, or SSE
334 // body stream running without an authority owner.
335 self.sse_task.abort();
336 }
337 }
338
339 #[cfg(test)]
340 mod endpoint_tests {
341 use std::time::Duration;
342
343 use super::{McpHttpAuth, SseInbound, SseTransport};
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 SseTransport::resolve_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 SseTransport::resolve_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!(SseTransport::resolve_endpoint_url(base, "http://169.254.169.254/latest").is_err());
365 // Different scheme -> rejected.
366 assert!(
367 SseTransport::resolve_endpoint_url(base, "http://mcp.example.com/messages").is_err()
368 );
369 // Different port -> rejected.
370 assert!(
371 SseTransport::resolve_endpoint_url(base, "https://mcp.example.com:8443/x").is_err()
372 );
373 }
374
375 #[tokio::test]
376 async fn message_before_endpoint_is_rejected_instead_of_buffered() {
377 let (tx, rx) = tokio::sync::mpsc::channel(1);
378 tx.send(SseInbound::Message(br#"{"jsonrpc":"2.0"}"#.to_vec()))
379 .await
380 .unwrap();
381 let mut transport = SseTransport {
382 client: reqwest::Client::new(),
383 base_url: "https://example.invalid/sse".to_string(),
384 auth: McpHttpAuth::default(),
385 endpoint_url: None,
386 receiver: rx,
387 sse_task: tokio::spawn(async {}),
388 };
389
390 let error = transport
391 .wait_for_endpoint(
392 &tokio_util::sync::CancellationToken::new(),
393 Duration::from_secs(1),
394 )
395 .await
396 .expect_err("pre-endpoint message must fail closed");
397 assert!(error.to_string().contains("before declaring its endpoint"));
398 }
399 }
400
400 lines RUST