Skip to content
Open
Show file tree
Hide file tree
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
10 changes: 10 additions & 0 deletions crates/rmcp/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -414,6 +414,16 @@ required-features = [
]
path = "tests/test_streamable_http_session_store.rs"

[[test]]
name = "test_streamable_http_event_store"
required-features = [
"client",
"server",
"transport-streamable-http-client-reqwest",
"transport-streamable-http-server",
]
path = "tests/test_streamable_http_event_store.rs"

[[test]]
name = "test_streamable_http_connection_reuse"
required-features = [
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ where
async fn get_stream(
&self,
uri: std::sync::Arc<str>,
session_id: std::sync::Arc<str>,
session_id: Option<std::sync::Arc<str>>,
last_event_id: Option<String>,
mut auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
Expand All @@ -50,7 +50,7 @@ where
async fn get_stream_with_max_sse_event_size(
&self,
uri: std::sync::Arc<str>,
session_id: std::sync::Arc<str>,
session_id: Option<std::sync::Arc<str>>,
last_event_id: Option<String>,
mut auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
Expand Down
49 changes: 49 additions & 0 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,7 @@ pin_project_lite::pin_project! {
where R: SseStreamReconnect
{
retry_policy: Arc<dyn SseRetryPolicy>,
reconnect_only_after_event_id: bool,
last_event_id: Option<String>,
server_retry_interval: Option<Duration>,
connector: R,
Expand All @@ -304,6 +305,22 @@ impl<R: SseStreamReconnect> SseAutoReconnectStream<R> {
) -> Self {
Self {
retry_policy,
reconnect_only_after_event_id: false,
last_event_id: None,
server_retry_interval: None,
connector,
state: SseAutoReconnectStreamState::Connected { stream },
}
}

pub fn new_after_event_id(
stream: BoxedSseResponse,
connector: R,
retry_policy: Arc<dyn SseRetryPolicy>,
) -> Self {
Self {
retry_policy,
reconnect_only_after_event_id: true,
last_event_id: None,
server_retry_interval: None,
connector,
Expand All @@ -317,6 +334,7 @@ impl<E: std::error::Error + Send> SseAutoReconnectStream<NeverReconnect<E>> {
pub(crate) fn never_reconnect(stream: BoxedSseResponse, error_when_reconnect: E) -> Self {
Self {
retry_policy: Arc::new(NeverRetry),
reconnect_only_after_event_id: false,
last_event_id: None,
server_retry_interval: None,
connector: NeverReconnect {
Expand Down Expand Up @@ -409,6 +427,10 @@ where
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
Expand All @@ -420,6 +442,13 @@ where
}
}
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.
Expand Down Expand Up @@ -686,4 +715,24 @@ mod tests {
&& attempts.load(Ordering::Relaxed) == 0
);
}

#[tokio::test]
async fn response_without_event_id_does_not_reconnect() {
let attempts = Arc::new(AtomicUsize::new(0));
let connector = CountingReconnect {
attempts: attempts.clone(),
};
let stream = SseAutoReconnectStream::new_after_event_id(
futures::stream::empty().boxed(),
connector,
Arc::new(FixedInterval {
max_times: Some(1),
duration: Duration::ZERO,
}),
);
let mut stream = std::pin::pin!(stream);

assert!(stream.next().await.is_none());
assert_eq!(attempts.load(Ordering::Relaxed), 0);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ impl StreamableHttpClient for reqwest::Client {
async fn get_stream(
&self,
uri: Arc<str>,
session_id: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
Expand All @@ -71,16 +71,18 @@ impl StreamableHttpClient for reqwest::Client {
async fn get_stream_with_max_sse_event_size(
&self,
uri: Arc<str>,
session_id: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
max_sse_event_size: usize,
) -> Result<BoxStream<'static, Result<Sse, SseError>>, StreamableHttpError<Self::Error>> {
let mut request_builder = self
.get(uri.as_ref())
.header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "))
.header(HEADER_SESSION_ID, session_id.as_ref());
.header(ACCEPT, [EVENT_STREAM_MIME_TYPE, JSON_MIME_TYPE].join(", "));
if let Some(session_id) = session_id {
request_builder = request_builder.header(HEADER_SESSION_ID, session_id.as_ref());
}
if let Some(last_event_id) = last_event_id {
request_builder = request_builder.header(HEADER_LAST_EVENT_ID, last_event_id);
}
Expand Down
9 changes: 9 additions & 0 deletions crates/rmcp/src/transport/common/server_side_http.rs
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,15 @@ impl ServerSseMessage {
retry: Some(retry),
}
}

/// Create a retry hint without changing the client's last event ID.
pub fn retry(retry: Duration) -> Self {
Self {
event_id: None,
message: None,
retry: Some(retry),
}
}
}

pub(crate) fn sse_stream_response(
Expand Down
11 changes: 7 additions & 4 deletions crates/rmcp/src/transport/common/unix_socket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -376,7 +376,7 @@ impl StreamableHttpClient for UnixSocketHttpClient {
async fn get_stream(
&self,
uri: Arc<str>,
session_id: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
Expand All @@ -396,7 +396,7 @@ impl StreamableHttpClient for UnixSocketHttpClient {
async fn get_stream_with_max_sse_event_size(
&self,
uri: Arc<str>,
session_id: Arc<str>,
session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_token: Option<String>,
custom_headers: HashMap<HeaderName, HeaderValue>,
Expand All @@ -410,8 +410,11 @@ impl StreamableHttpClient for UnixSocketHttpClient {
.header(
http::header::ACCEPT,
format!("{EVENT_STREAM_MIME_TYPE}, {JSON_MIME_TYPE}"),
)
.header(HEADER_SESSION_ID, session_id.as_ref());
);

if let Some(session_id) = session_id {
builder = builder.header(HEADER_SESSION_ID, session_id.as_ref());
}

if let Some(last_id) = last_event_id {
builder = builder.header(HEADER_LAST_EVENT_ID, last_id);
Expand Down
Loading