From 5ad0a3b47b9b08a951b2f365bb2a6d86f3f233f2 Mon Sep 17 00:00:00 2001 From: Jeremy Schoemaker Date: Thu, 6 Aug 2026 23:54:28 -0500 Subject: [PATCH] fix(sse): loop instead of recursing when skipping SSE events MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `SseAutoReconnectStream::poll_next` recursed into itself for every event it skips — control frames, data-less frames, and frames whose data fails to deserialize — plus once more on every state transition. The inner stream returns `Poll::Ready` for each event it can parse out of already-buffered bytes, so a burst of skipped events has no yield point between them. Each one adds a stack frame, and `poll_next` is a large frame. A client connected to a server that emits a run of non-JSON `message` frames overflows the stack and aborts the process. Wrap the body in a `loop` and replace the four `self.poll_next(cx)` tail calls with `continue`. `this` is re-derived from `self.as_mut().project()` at the top of each iteration, so the borrows end cleanly per iteration; no other logic changes. The diff is mostly the resulting re-indent — review with `?w=1`. Adds `skipped_events_do_not_grow_the_stack`, which feeds 50,000 undeserializable frames through the stream. Before this change it aborts with `fatal runtime error: stack overflow` (SIGABRT); after, it passes. --- .../src/transport/common/client_side_sse.rs | 288 ++++++++++-------- 1 file changed, 164 insertions(+), 124 deletions(-) diff --git a/crates/rmcp/src/transport/common/client_side_sse.rs b/crates/rmcp/src/transport/common/client_side_sse.rs index 4ed05b44f..e668d63df 100644 --- a/crates/rmcp/src/transport/common/client_side_sse.rs +++ b/crates/rmcp/src/transport/common/client_side_sse.rs @@ -376,145 +376,157 @@ where mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> Poll> { - let mut this = self.as_mut().project(); - // let this_state = this.state.as_mut().project() - let state = this.state.as_mut().project(); - let next_state = match state { - SseAutoReconnectStreamStateProj::Connected { stream } => { - match ready!(stream.poll_next(cx)) { - Some(Ok(sse)) => { - if let Some(new_server_retry) = sse.retry { - *this.server_retry_interval = - Some(Duration::from_millis(new_server_retry)); - } - if let Some(ref event_id) = sse.id { - *this.last_event_id = Some(event_id.clone()); - } - // Only treat blank/`message` events as JSON-RPC payloads. - // Other control frames (endpoint, ping, etc.) are passed to - // the reconnection handler. - let is_message_event = - matches!(sse.event.as_deref(), None | Some("") | Some("message")); - if !is_message_event { - match this.connector.handle_control_event(&sse) { - Ok(()) => return self.poll_next(cx), - Err(e) => { - this.state.set(SseAutoReconnectStreamState::Terminated); - return Poll::Ready(Some(Err(e))); - } + loop { + let mut this = self.as_mut().project(); + // let this_state = this.state.as_mut().project() + let state = this.state.as_mut().project(); + let next_state = match state { + SseAutoReconnectStreamStateProj::Connected { stream } => { + match ready!(stream.poll_next(cx)) { + Some(Ok(sse)) => { + if let Some(new_server_retry) = sse.retry { + *this.server_retry_interval = + Some(Duration::from_millis(new_server_retry)); } - } - if let Some(data) = sse.data { - match serde_json::from_str::(&data) { - Err(e) => { - // Downgrade to debug to avoid noisy logs when servers emit - // non-JSON payloads as message frames. Include last_event_id - // to aid troubleshooting while keeping default behaviour. - let last_id = this.last_event_id.as_deref().unwrap_or(""); - tracing::debug!(last_event_id=%last_id, "failed to deserialize server message: {e}"); - return self.poll_next(cx); - } - Ok(message) => { - return Poll::Ready(Some(Ok(message))); + if let Some(ref event_id) = sse.id { + *this.last_event_id = Some(event_id.clone()); + } + // Only treat blank/`message` events as JSON-RPC payloads. + // Other control frames (endpoint, ping, etc.) are passed to + // the reconnection handler. + let is_message_event = + matches!(sse.event.as_deref(), None | Some("") | Some("message")); + if !is_message_event { + match this.connector.handle_control_event(&sse) { + Ok(()) => continue, + Err(e) => { + this.state.set(SseAutoReconnectStreamState::Terminated); + return Poll::Ready(Some(Err(e))); + } } - }; - } else { - return self.poll_next(cx); - } - } - Some(Err(e)) => { - if is_event_too_large_error(&e) { - this.state.set(SseAutoReconnectStreamState::Terminated); - return Poll::Ready(this.connector.map_fatal_stream_error(e).map(Err)); - } - if *this.reconnect_only_after_event_id && this.last_event_id.is_none() { - this.state.set(SseAutoReconnectStreamState::Terminated); - return Poll::Ready(this.connector.map_fatal_stream_error(e).map(Err)); - } - this.connector - .handle_stream_error(&e, this.last_event_id.as_deref()); - let retrying = this - .connector - .retry_connection(this.last_event_id.as_deref()); - SseAutoReconnectStreamState::Retrying { - retry_times: 0, - retrying, - } - } - None => { - if *this.reconnect_only_after_event_id && this.last_event_id.is_none() { - tracing::debug!( - "sse response ended before an event ID was received; cannot resume" - ); - this.state.set(SseAutoReconnectStreamState::Terminated); - return Poll::Ready(None); + } + if let Some(data) = sse.data { + match serde_json::from_str::(&data) { + Err(e) => { + // Downgrade to debug to avoid noisy logs when servers emit + // non-JSON payloads as message frames. Include last_event_id + // to aid troubleshooting while keeping default behaviour. + let last_id = this.last_event_id.as_deref().unwrap_or(""); + tracing::debug!(last_event_id=%last_id, "failed to deserialize server message: {e}"); + continue; + } + Ok(message) => { + return Poll::Ready(Some(Ok(message))); + } + }; + } else { + continue; + } } - // Per SEP-1699, a graceful stream close is - // reconnectable. If the server sent a `retry` field - // we MUST wait that long before reconnecting. - let interval = this - .server_retry_interval - .take() - .or_else(|| this.retry_policy.retry(0)); - if let Some(interval) = interval { - tracing::debug!(?interval, "sse stream ended gracefully, reconnecting"); - SseAutoReconnectStreamState::WaitingNextRetry { - sleep: tokio::time::sleep(interval), + Some(Err(e)) => { + if is_event_too_large_error(&e) { + this.state.set(SseAutoReconnectStreamState::Terminated); + return Poll::Ready( + this.connector.map_fatal_stream_error(e).map(Err), + ); + } + if *this.reconnect_only_after_event_id && this.last_event_id.is_none() { + this.state.set(SseAutoReconnectStreamState::Terminated); + return Poll::Ready( + this.connector.map_fatal_stream_error(e).map(Err), + ); + } + this.connector + .handle_stream_error(&e, this.last_event_id.as_deref()); + let retrying = this + .connector + .retry_connection(this.last_event_id.as_deref()); + SseAutoReconnectStreamState::Retrying { retry_times: 0, + retrying, } - } else { - tracing::debug!("sse stream terminated, no reconnect policy"); - return Poll::Ready(None); } - } - } - } - SseAutoReconnectStreamStateProj::Retrying { - retry_times, - retrying, - } => { - let retry_result = ready!(retrying.poll(cx)); - match retry_result { - Ok(new_stream) => SseAutoReconnectStreamState::Connected { stream: new_stream }, - Err(e) => { - tracing::debug!("retry sse stream error: {e}"); - *retry_times += 1; - if let Some(interval) = this.retry_policy.retry(*retry_times) { + None => { + if *this.reconnect_only_after_event_id && this.last_event_id.is_none() { + tracing::debug!( + "sse response ended before an event ID was received; cannot resume" + ); + this.state.set(SseAutoReconnectStreamState::Terminated); + return Poll::Ready(None); + } + // Per SEP-1699, a graceful stream close is + // reconnectable. If the server sent a `retry` field + // we MUST wait that long before reconnecting. let interval = this .server_retry_interval - .map(|server_retry_interval| server_retry_interval.max(interval)) - .unwrap_or(interval); - let sleep = tokio::time::sleep(interval); - SseAutoReconnectStreamState::WaitingNextRetry { - sleep, - retry_times: *retry_times, + .take() + .or_else(|| this.retry_policy.retry(0)); + if let Some(interval) = interval { + tracing::debug!( + ?interval, + "sse stream ended gracefully, reconnecting" + ); + SseAutoReconnectStreamState::WaitingNextRetry { + sleep: tokio::time::sleep(interval), + retry_times: 0, + } + } else { + tracing::debug!("sse stream terminated, no reconnect policy"); + return Poll::Ready(None); } - } else { - tracing::error!("sse stream error: {e}, max retry times reached"); - this.state.set(SseAutoReconnectStreamState::Terminated); - return Poll::Ready(Some(Err(e))); } } } - } - SseAutoReconnectStreamStateProj::WaitingNextRetry { sleep, retry_times } => { - ready!(sleep.poll(cx)); - let retrying = this - .connector - .retry_connection(this.last_event_id.as_deref()); - let retry_times = *retry_times; - SseAutoReconnectStreamState::Retrying { + SseAutoReconnectStreamStateProj::Retrying { retry_times, retrying, + } => { + let retry_result = ready!(retrying.poll(cx)); + match retry_result { + Ok(new_stream) => { + SseAutoReconnectStreamState::Connected { stream: new_stream } + } + Err(e) => { + tracing::debug!("retry sse stream error: {e}"); + *retry_times += 1; + if let Some(interval) = this.retry_policy.retry(*retry_times) { + let interval = this + .server_retry_interval + .map(|server_retry_interval| { + server_retry_interval.max(interval) + }) + .unwrap_or(interval); + let sleep = tokio::time::sleep(interval); + SseAutoReconnectStreamState::WaitingNextRetry { + sleep, + retry_times: *retry_times, + } + } else { + tracing::error!("sse stream error: {e}, max retry times reached"); + this.state.set(SseAutoReconnectStreamState::Terminated); + return Poll::Ready(Some(Err(e))); + } + } + } } - } - SseAutoReconnectStreamStateProj::Terminated => { - return Poll::Ready(None); - } - }; - // update the state - this.state.set(next_state); - self.poll_next(cx) + SseAutoReconnectStreamStateProj::WaitingNextRetry { sleep, retry_times } => { + ready!(sleep.poll(cx)); + let retrying = this + .connector + .retry_connection(this.last_event_id.as_deref()); + let retry_times = *retry_times; + SseAutoReconnectStreamState::Retrying { + retry_times, + retrying, + } + } + SseAutoReconnectStreamStateProj::Terminated => { + return Poll::Ready(None); + } + }; + // update the state + this.state.set(next_state); + } } } @@ -716,6 +728,34 @@ mod tests { ); } + #[tokio::test] + async fn skipped_events_do_not_grow_the_stack() { + // Every frame here fails to deserialize, so each one takes the "skip this + // event" path. The inner stream is always ready, so nothing yields in + // between: when that path recursed into poll_next instead of looping, this + // overflowed the stack and aborted the process. + const SKIPPED_EVENTS: usize = 50_000; + let mut payload = Vec::with_capacity(SKIPPED_EVENTS * 9); + for _ in 0..SKIPPED_EVENTS { + payload.extend_from_slice(b"data: x\n\n"); + } + + let source = futures::stream::iter([Ok::<_, std::io::Error>(Bytes::from(payload))]); + let attempts = Arc::new(AtomicUsize::new(0)); + let connector = CountingReconnect { + attempts: attempts.clone(), + }; + let stream = SseAutoReconnectStream::new( + bounded_sse_stream(source, 1024), + connector, + Arc::new(NeverRetry), + ); + let mut stream = std::pin::pin!(stream); + + assert!(stream.next().await.is_none()); + assert_eq!(attempts.load(Ordering::Relaxed), 0); + } + #[tokio::test] async fn response_without_event_id_does_not_reconnect() { let attempts = Arc::new(AtomicUsize::new(0));