Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
288 changes: 164 additions & 124 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -376,145 +376,157 @@ where
mut self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> Poll<Option<Self::Item>> {
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::<ServerJsonRpcMessage>(&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::<ServerJsonRpcMessage>(&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);
}
}
}

Expand Down Expand Up @@ -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));
Expand Down
Loading