Skip to content

Commit aad4d4e

Browse files
authored
feat!: add distributed SSE event store (#1024)
1 parent 50dd8e2 commit aad4d4e

14 files changed

Lines changed: 953 additions & 182 deletions

crates/rmcp/Cargo.toml

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -414,6 +414,16 @@ required-features = [
414414
]
415415
path = "tests/test_streamable_http_session_store.rs"
416416

417+
[[test]]
418+
name = "test_streamable_http_event_store"
419+
required-features = [
420+
"client",
421+
"server",
422+
"transport-streamable-http-client-reqwest",
423+
"transport-streamable-http-server",
424+
]
425+
path = "tests/test_streamable_http_event_store.rs"
426+
417427
[[test]]
418428
name = "test_streamable_http_connection_reuse"
419429
required-features = [

crates/rmcp/src/transport/common/auth/streamable_http_client.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@ where
3131
async fn get_stream(
3232
&self,
3333
uri: std::sync::Arc<str>,
34-
session_id: std::sync::Arc<str>,
34+
session_id: Option<std::sync::Arc<str>>,
3535
last_event_id: Option<String>,
3636
mut auth_token: Option<String>,
3737
custom_headers: HashMap<HeaderName, HeaderValue>,
@@ -50,7 +50,7 @@ where
5050
async fn get_stream_with_max_sse_event_size(
5151
&self,
5252
uri: std::sync::Arc<str>,
53-
session_id: std::sync::Arc<str>,
53+
session_id: Option<std::sync::Arc<str>>,
5454
last_event_id: Option<String>,
5555
mut auth_token: Option<String>,
5656
custom_headers: HashMap<HeaderName, HeaderValue>,

crates/rmcp/src/transport/common/client_side_sse.rs

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,7 @@ pin_project_lite::pin_project! {
288288
where R: SseStreamReconnect
289289
{
290290
retry_policy: Arc<dyn SseRetryPolicy>,
291+
reconnect_only_after_event_id: bool,
291292
last_event_id: Option<String>,
292293
server_retry_interval: Option<Duration>,
293294
connector: R,
@@ -304,6 +305,22 @@ impl<R: SseStreamReconnect> SseAutoReconnectStream<R> {
304305
) -> Self {
305306
Self {
306307
retry_policy,
308+
reconnect_only_after_event_id: false,
309+
last_event_id: None,
310+
server_retry_interval: None,
311+
connector,
312+
state: SseAutoReconnectStreamState::Connected { stream },
313+
}
314+
}
315+
316+
pub fn new_after_event_id(
317+
stream: BoxedSseResponse,
318+
connector: R,
319+
retry_policy: Arc<dyn SseRetryPolicy>,
320+
) -> Self {
321+
Self {
322+
retry_policy,
323+
reconnect_only_after_event_id: true,
307324
last_event_id: None,
308325
server_retry_interval: None,
309326
connector,
@@ -317,6 +334,7 @@ impl<E: std::error::Error + Send> SseAutoReconnectStream<NeverReconnect<E>> {
317334
pub(crate) fn never_reconnect(stream: BoxedSseResponse, error_when_reconnect: E) -> Self {
318335
Self {
319336
retry_policy: Arc::new(NeverRetry),
337+
reconnect_only_after_event_id: false,
320338
last_event_id: None,
321339
server_retry_interval: None,
322340
connector: NeverReconnect {
@@ -409,6 +427,10 @@ where
409427
this.state.set(SseAutoReconnectStreamState::Terminated);
410428
return Poll::Ready(this.connector.map_fatal_stream_error(e).map(Err));
411429
}
430+
if *this.reconnect_only_after_event_id && this.last_event_id.is_none() {
431+
this.state.set(SseAutoReconnectStreamState::Terminated);
432+
return Poll::Ready(this.connector.map_fatal_stream_error(e).map(Err));
433+
}
412434
this.connector
413435
.handle_stream_error(&e, this.last_event_id.as_deref());
414436
let retrying = this
@@ -420,6 +442,13 @@ where
420442
}
421443
}
422444
None => {
445+
if *this.reconnect_only_after_event_id && this.last_event_id.is_none() {
446+
tracing::debug!(
447+
"sse response ended before an event ID was received; cannot resume"
448+
);
449+
this.state.set(SseAutoReconnectStreamState::Terminated);
450+
return Poll::Ready(None);
451+
}
423452
// Per SEP-1699, a graceful stream close is
424453
// reconnectable. If the server sent a `retry` field
425454
// we MUST wait that long before reconnecting.
@@ -686,4 +715,24 @@ mod tests {
686715
&& attempts.load(Ordering::Relaxed) == 0
687716
);
688717
}
718+
719+
#[tokio::test]
720+
async fn response_without_event_id_does_not_reconnect() {
721+
let attempts = Arc::new(AtomicUsize::new(0));
722+
let connector = CountingReconnect {
723+
attempts: attempts.clone(),
724+
};
725+
let stream = SseAutoReconnectStream::new_after_event_id(
726+
futures::stream::empty().boxed(),
727+
connector,
728+
Arc::new(FixedInterval {
729+
max_times: Some(1),
730+
duration: Duration::ZERO,
731+
}),
732+
);
733+
let mut stream = std::pin::pin!(stream);
734+
735+
assert!(stream.next().await.is_none());
736+
assert_eq!(attempts.load(Ordering::Relaxed), 0);
737+
}
689738
}

crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@ impl StreamableHttpClient for reqwest::Client {
5252
async fn get_stream(
5353
&self,
5454
uri: Arc<str>,
55-
session_id: Arc<str>,
55+
session_id: Option<Arc<str>>,
5656
last_event_id: Option<String>,
5757
auth_token: Option<String>,
5858
custom_headers: HashMap<HeaderName, HeaderValue>,
@@ -71,16 +71,18 @@ impl StreamableHttpClient for reqwest::Client {
7171
async fn get_stream_with_max_sse_event_size(
7272
&self,
7373
uri: Arc<str>,
74-
session_id: Arc<str>,
74+
session_id: Option<Arc<str>>,
7575
last_event_id: Option<String>,
7676
auth_token: Option<String>,
7777
custom_headers: HashMap<HeaderName, HeaderValue>,
7878
max_sse_event_size: usize,
7979
) -> Result<BoxStream<'static, Result<Sse, SseError>>, StreamableHttpError<Self::Error>> {
8080
let mut request_builder = self
8181
.get(uri.as_ref())
82-
.header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "))
83-
.header(HEADER_SESSION_ID, session_id.as_ref());
82+
.header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "));
83+
if let Some(session_id) = session_id {
84+
request_builder = request_builder.header(HEADER_SESSION_ID, session_id.as_ref());
85+
}
8486
if let Some(last_event_id) = last_event_id {
8587
request_builder = request_builder.header(HEADER_LAST_EVENT_ID, last_event_id);
8688
}

crates/rmcp/src/transport/common/server_side_http.rs

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -119,6 +119,15 @@ impl ServerSseMessage {
119119
retry: Some(retry),
120120
}
121121
}
122+
123+
/// Create a retry hint without changing the client's last event ID.
124+
pub fn retry(retry: Duration) -> Self {
125+
Self {
126+
event_id: None,
127+
message: None,
128+
retry: Some(retry),
129+
}
130+
}
122131
}
123132

124133
pub(crate) fn sse_stream_response(

crates/rmcp/src/transport/common/unix_socket.rs

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -376,7 +376,7 @@ impl StreamableHttpClient for UnixSocketHttpClient {
376376
async fn get_stream(
377377
&self,
378378
uri: Arc<str>,
379-
session_id: Arc<str>,
379+
session_id: Option<Arc<str>>,
380380
last_event_id: Option<String>,
381381
auth_token: Option<String>,
382382
custom_headers: HashMap<HeaderName, HeaderValue>,
@@ -396,7 +396,7 @@ impl StreamableHttpClient for UnixSocketHttpClient {
396396
async fn get_stream_with_max_sse_event_size(
397397
&self,
398398
uri: Arc<str>,
399-
session_id: Arc<str>,
399+
session_id: Option<Arc<str>>,
400400
last_event_id: Option<String>,
401401
auth_token: Option<String>,
402402
custom_headers: HashMap<HeaderName, HeaderValue>,
@@ -410,8 +410,11 @@ impl StreamableHttpClient for UnixSocketHttpClient {
410410
.header(
411411
http::header::ACCEPT,
412412
format!("{EVENT_STREAM_MIME_TYPE}, {JSON_MIME_TYPE}"),
413-
)
414-
.header(HEADER_SESSION_ID, session_id.as_ref());
413+
);
414+
415+
if let Some(session_id) = session_id {
416+
builder = builder.header(HEADER_SESSION_ID, session_id.as_ref());
417+
}
415418

416419
if let Some(last_id) = last_event_id {
417420
builder = builder.header(HEADER_LAST_EVENT_ID, last_id);

0 commit comments

Comments
 (0)