返回 CodeWhale
sse.rs
根目录 / crates / tui / src / mcp / sse.rs
1 use std::time::Duration;
2
3 use anyhow::{Context, Result};
4
5 use super::http_client::McpHttpClient;
6 use super::wire::{
7 MAX_SSE_FRAME_BYTES, McpSessionRejected, find_sse_event_separator_bytes,
8 is_mcp_stale_session_body, resolve_sse_endpoint_url, sse_field_value,
9 };
10 use super::{ERROR_BODY_PREVIEW_BYTES, McpTransport, bounded_body_excerpt, mask_url_secrets};
11
12 const SSE_INBOUND_CHANNEL_CAPACITY: usize = 4;
13
14 pub(crate) struct SseTransport {
15 pub(super) client: McpHttpClient,
16 pub(super) base_url: String,
17 pub(super) endpoint_url: Option<String>,
18 pub(super) receiver: tokio::sync::mpsc::Receiver<SseInbound>,
19 pub(super) sse_task: tokio::task::JoinHandle<()>,
20 }
21
22 pub(super) enum SseInbound {
23 Endpoint(String),
24 Message(Vec<u8>),
25 }
26
27 impl SseTransport {
28 pub(super) async fn connect(
29 client: McpHttpClient,
30 url: String,
31 cancel_token: tokio_util::sync::CancellationToken,
32 endpoint_timeout: Duration,
33 ) -> Result<Self> {
34 let (tx, rx) = tokio::sync::mpsc::channel(SSE_INBOUND_CHANNEL_CAPACITY);
35 let client_clone = client.clone();
36 let url_clone = url.clone();
37 let wait_cancel_token = cancel_token.clone();
38
39 let sse_task = tokio::spawn(async move {
40 if cancel_token.is_cancelled() {
41 return;
42 }
43 use futures_util::FutureExt;
44 let result = std::panic::AssertUnwindSafe(Self::run_sse_loop(
45 client_clone,
46 url_clone,
47 tx,
48 cancel_token,
49 ))
50 .catch_unwind()
51 .await;
52 match result {
53 Ok(res) => {
54 if let Err(e) = res {
55 tracing::error!("SSE loop error: {}", e);
56 }
57 }
58 Err(panic_err) => {
59 if let Some(msg) = panic_err.downcast_ref::<&str>() {
60 tracing::error!("SSE loop panicked: {}", msg);
61 } else if let Some(msg) = panic_err.downcast_ref::<String>() {
62 tracing::error!("SSE loop panicked: {}", msg);
63 } else {
64 tracing::error!("SSE loop panicked with unknown error");
65 }
66 }
67 }
68 });
69
70 let mut transport = Self {
71 client,
72 base_url: url,
73 endpoint_url: None,
74 receiver: rx,
75 sse_task,
76 };
77 transport
78 .wait_for_endpoint(&wait_cancel_token, endpoint_timeout)
79 .await?;
80 Ok(transport)
81 }
82
83 async fn run_sse_loop(
84 client: McpHttpClient,
85 url: String,
86 tx: tokio::sync::mpsc::Sender<SseInbound>,
87 cancel_token: tokio_util::sync::CancellationToken,
88 ) -> Result<()> {
89 let request = tokio::select! {
90 biased;
91 _ = cancel_token.cancelled() => {
92 anyhow::bail!("MCP SSE connect cancelled before authentication completed")
93 }
94 request = client.prepare_mcp_request(client.get(&url), false) => request?,
95 };
96 let response = tokio::select! {
97 biased;
98 _ = cancel_token.cancelled() => {
99 anyhow::bail!("MCP SSE connect cancelled before the request completed")
100 }
101 response = client.send_event_stream(request) => response.with_context(|| {
102 format!(
103 "MCP SSE connect failed (transport=http url={})",
104 mask_url_secrets(&url),
105 )
106 })?,
107 };
108 let status = response.status();
109 if !status.is_success() {
110 let body_excerpt = bounded_body_excerpt(response, ERROR_BODY_PREVIEW_BYTES).await;
111 let body_excerpt = client.server_error_preview(&body_excerpt);
112 anyhow::bail!(
113 "MCP SSE rejected (transport=http url={} status={}): {}",
114 mask_url_secrets(&url),
115 status,
116 body_excerpt,
117 );
118 }
119
120 let mut stream = response.bytes_stream();
121 use futures_util::StreamExt;
122 // Raw byte buffer so a multi-byte UTF-8 char split across reads is not
123 // corrupted, and bounded so a separator-less server cannot OOM us.
124 let mut buffer: Vec<u8> = Vec::new();
125
126 loop {
127 if cancel_token.is_cancelled() {
128 tracing::debug!("SSE loop cancelled");
129 break;
130 }
131 let item = tokio::select! {
132 _ = cancel_token.cancelled() => {
133 tracing::debug!("SSE loop shutting down");
134 break;
135 }
136 item = stream.next() => {
137 match item {
138 Some(i) => i,
139 None => break,
140 }
141 }
142 };
143 let chunk = item?;
144 buffer.extend_from_slice(&chunk);
145 if buffer.len() > MAX_SSE_FRAME_BYTES {
146 anyhow::bail!(
147 "MCP SSE frame exceeded {} bytes without a separator — aborting",
148 MAX_SSE_FRAME_BYTES
149 );
150 }
151
152 while let Some((pos, separator_len)) = find_sse_event_separator_bytes(&buffer) {
153 // Complete block: decoding cannot split a multi-byte char.
154 let event_block = String::from_utf8_lossy(&buffer[..pos]).into_owned();
155 buffer.drain(..pos + separator_len);
156
157 let mut event_type = "message";
158 let mut data = String::new();
159
160 for line in event_block.lines() {
161 if let Some(value) = sse_field_value(line, "event:") {
162 event_type = value;
163 } else if let Some(value) = sse_field_value(line, "data:") {
164 if !data.is_empty() {
165 data.push('\n');
166 }
167 data.push_str(value);
168 }
169 }
170
171 let inbound = match event_type {
172 "endpoint" => Some(SseInbound::Endpoint(data)),
173 "message" if !data.trim().is_empty() => {
174 Some(SseInbound::Message(data.into_bytes()))
175 }
176 _ => None,
177 };
178 if let Some(inbound) = inbound {
179 let sent = tokio::select! {
180 biased;
181 _ = cancel_token.cancelled() => return Ok(()),
182 sent = tx.send(inbound) => sent,
183 };
184 if sent.is_err() {
185 return Ok(());
186 }
187 }
188 }
189 }
190 Ok(())
191 }
192
193 async fn wait_for_endpoint(
194 &mut self,
195 cancel_token: &tokio_util::sync::CancellationToken,
196 endpoint_timeout: Duration,
197 ) -> Result<()> {
198 let timeout = tokio::time::sleep(endpoint_timeout);
199 tokio::pin!(timeout);
200
201 let msg = tokio::select! {
202 _ = cancel_token.cancelled() => {
203 anyhow::bail!("SSE transport cancelled before endpoint was discovered");
204 }
205 _ = &mut timeout => {
206 anyhow::bail!(
207 "SSE endpoint not received within {}ms",
208 endpoint_timeout.as_millis()
209 );
210 }
211 msg = self.receiver.recv() => {
212 msg.context("SSE transport closed before endpoint was discovered")?
213 }
214 };
215
216 match msg {
217 SseInbound::Endpoint(endpoint) => self.store_endpoint(&endpoint),
218 SseInbound::Message(_) => {
219 anyhow::bail!("MCP SSE server sent a message before declaring its endpoint");
220 }
221 }
222 }
223
224 fn store_endpoint(&mut self, endpoint: &str) -> Result<()> {
225 self.endpoint_url = Some(resolve_sse_endpoint_url(&self.base_url, endpoint)?);
226 Ok(())
227 }
228 }
229
230 #[async_trait::async_trait]
231 impl McpTransport for SseTransport {
232 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
233 let endpoint = self
234 .endpoint_url
235 .as_ref()
236 .context("SSE endpoint not yet discovered")?
237 .clone();
238 let request = self
239 .client
240 .prepare_mcp_request(self.client.post(&endpoint), true)
241 .await?
242 .body(msg);
243 let response = self.client.send(request).await.with_context(|| {
244 format!(
245 "MCP SSE POST send failed (transport=sse endpoint={})",
246 mask_url_secrets(&endpoint)
247 )
248 })?;
249 let status = response.status();
250 if !status.is_success() {
251 let body_excerpt = bounded_body_excerpt(response, ERROR_BODY_PREVIEW_BYTES).await;
252 let stale_session = is_mcp_stale_session_body(&body_excerpt);
253 let body_excerpt = self.client.server_error_preview(&body_excerpt);
254 if stale_session {
255 return Err(McpSessionRejected(format!(
256 "MCP session expired (transport=sse endpoint={} status={}): {}",
257 mask_url_secrets(&endpoint),
258 status,
259 body_excerpt
260 ))
261 .into());
262 }
263 anyhow::bail!(
264 "MCP SSE POST rejected (transport=sse endpoint={} status={}): {}",
265 mask_url_secrets(&endpoint),
266 status,
267 body_excerpt
268 );
269 }
270 Ok(())
271 }
272
273 async fn recv(&mut self) -> Result<Vec<u8>> {
274 loop {
275 match self.receiver.recv().await.context("SSE transport closed")? {
276 SseInbound::Endpoint(endpoint) => {
277 self.store_endpoint(&endpoint)?;
278 }
279 SseInbound::Message(msg) => return Ok(msg),
280 }
281 }
282 }
283
284 /// The event stream is the only inbound channel: once its task has
285 /// ended (server closed it, network error, oversize frame), POSTs may
286 /// still be accepted while every reply is lost. Report it dead so the
287 /// pool reconnects before dispatching instead of losing a tool result.
288 fn probe_dead(&self) -> bool {
289 self.sse_task.is_finished()
290 }
291
292 async fn shutdown(&mut self) {
293 self.sse_task.abort();
294 }
295 }
296
297 impl Drop for SseTransport {
298 fn drop(&mut self) {
299 // Dropping a JoinHandle detaches its task. Abort explicitly so a
300 // cancelled connection cannot leave an auth refresh, connect, or SSE
301 // body stream running without an authority owner.
302 self.sse_task.abort();
303 }
304 }
305
306 #[cfg(test)]
307 mod endpoint_tests {
308 use std::time::Duration;
309
310 use super::{McpHttpClient, SseInbound, SseTransport};
311
312 #[tokio::test]
313 async fn message_before_endpoint_is_rejected_instead_of_buffered() {
314 // Building a reqwest client needs the process-wide rustls provider;
315 // production installs it at startup, and this test must not depend
316 // on another test in the same process having done so first.
317 crate::tls::ensure_rustls_crypto_provider();
318 let (tx, rx) = tokio::sync::mpsc::channel(1);
319 tx.send(SseInbound::Message(br#"{"jsonrpc":"2.0"}"#.to_vec()))
320 .await
321 .unwrap();
322 let mut transport = SseTransport {
323 client: McpHttpClient::new(
324 "https://example.invalid/sse",
325 false,
326 false,
327 false,
328 None,
329 Duration::from_secs(10),
330 Duration::from_secs(120),
331 )
332 .unwrap(),
333 base_url: "https://example.invalid/sse".to_string(),
334 endpoint_url: None,
335 receiver: rx,
336 sse_task: tokio::spawn(async {}),
337 };
338
339 let error = transport
340 .wait_for_endpoint(
341 &tokio_util::sync::CancellationToken::new(),
342 Duration::from_secs(1),
343 )
344 .await
345 .expect_err("pre-endpoint message must fail closed");
346 assert!(error.to_string().contains("before declaring its endpoint"));
347 }
348
349 /// Serve one legacy SSE stream: the endpoint event at once, then each
350 /// `(delay, frame)` in order, then close the stream or hold it open.
351 async fn serve_sse_stream(frames: Vec<(Duration, &'static [u8])>, close: bool) -> String {
352 use tokio::io::{AsyncReadExt, AsyncWriteExt};
353 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
354 let url = format!("http://{}/sse", listener.local_addr().unwrap());
355 tokio::spawn(async move {
356 let (mut socket, _) = listener.accept().await.unwrap();
357 let mut request = Vec::new();
358 let mut buf = [0; 1024];
359 while !request.windows(4).any(|window| window == b"\r\n\r\n") {
360 let n = socket.read(&mut buf).await.unwrap();
361 assert!(n > 0, "client closed before sending its request");
362 request.extend_from_slice(&buf[..n]);
363 }
364 socket
365 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\nevent: endpoint\ndata: /messages\n\n")
366 .await
367 .unwrap();
368 for (delay, frame) in frames {
369 tokio::time::sleep(delay).await;
370 if socket.write_all(frame).await.is_err() {
371 return;
372 }
373 }
374 if !close {
375 std::future::pending::<()>().await;
376 }
377 });
378 url
379 }
380
381 async fn connect_with_read_timeout(url: String, read_timeout: Duration) -> SseTransport {
382 crate::tls::ensure_rustls_crypto_provider();
383 let client = McpHttpClient::new(
384 &url,
385 false,
386 false,
387 false,
388 None,
389 Duration::from_secs(5),
390 read_timeout,
391 )
392 .unwrap();
393 SseTransport::connect(
394 client,
395 url,
396 tokio_util::sync::CancellationToken::new(),
397 Duration::from_secs(5),
398 )
399 .await
400 .unwrap()
401 }
402
403 #[tokio::test]
404 async fn event_stream_outlives_the_request_read_timeout() {
405 use crate::mcp::McpTransport as _;
406 let _env = crate::test_support::lock_test_env();
407 let _no_proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
408 // A quiet stream for longer than read_timeout is healthy, not dead.
409 let url = serve_sse_stream(
410 vec![(
411 Duration::from_millis(900),
412 b"event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1}\n\n",
413 )],
414 false,
415 )
416 .await;
417 let mut transport = connect_with_read_timeout(url, Duration::from_millis(300)).await;
418 let message = tokio::time::timeout(Duration::from_secs(5), transport.recv())
419 .await
420 .expect("stream message within the test bound")
421 .expect("stream still open after read_timeout");
422 assert_eq!(message, br#"{"jsonrpc":"2.0","id":1}"#);
423 assert!(!transport.probe_dead());
424 }
425
426 #[tokio::test]
427 async fn closed_event_stream_reads_as_dead() {
428 use crate::mcp::McpTransport as _;
429 let _env = crate::test_support::lock_test_env();
430 let _no_proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
431 let url = serve_sse_stream(Vec::new(), true).await;
432 let transport = connect_with_read_timeout(url, Duration::from_secs(5)).await;
433 let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
434 while !transport.probe_dead() {
435 assert!(
436 tokio::time::Instant::now() < deadline,
437 "a closed SSE stream must stop reading as alive"
438 );
439 tokio::time::sleep(Duration::from_millis(20)).await;
440 }
441 }
442 }
443
443 lines RUST