Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
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
47 changes: 42 additions & 5 deletions crates/rmcp/src/transport/common/client_side_sse.rs
Original file line numberDiff line numberDiff line change
Expand Up@@ -98,10 +98,29 @@ impl<E: std::error::Error + Send> SseStreamReconnect for NeverReconnect<E> {
}
}

/// Abstraction for SSE reconnection logic. Implementors can hook into
/// [`handle_control_event`](Self::handle_control_event) to consume control
/// frames (e.g. `event: endpoint`) that arrive when a server restarts an SSE
/// stream. The default implementation is a no-op, keeping existing behaviour
/// intact.
pub(crate) trait SseStreamReconnect {
type Error: std::error::Error;
type Future: Future<Output = Result<BoxedSseResponse, Self::Error>> + Send;
fn retry_connection(&mut self, last_event_id: Option<&str>) -> Self::Future;
fn handle_control_event(&mut self, _event: &Sse) -> Result<(), Self::Error> {
Ok(())
}
fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
if let Some(id) = last_event_id {
tracing::warn!(%id, "sse stream error: {error}");
} else {
tracing::warn!("sse stream error: {error}");
}
}
}

pin_project_lite::pin_project! {
Expand DownExpand Up@@ -189,14 +208,31 @@ where
*this.server_retry_interval =
Some(Duration::from_millis(new_server_retry));
}
if let Some(event_id) = sse.id {
*this.last_event_id = Some(event_id);
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"));
Comment thread
4t145 marked this conversation as resolved.
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)));
}
}
}
if let Some(data) = sse.data {
match serde_json::from_str::<ServerJsonRpcMessage>(&data) {
Err(e) => {
// not sure should this be a hard error
tracing::warn!("failed to deserialize server message: {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) => {
Expand All@@ -208,7 +244,8 @@ where
}
}
Some(Err(e)) => {
tracing::warn!("sse stream error: {e}");
this.connector
.handle_stream_error(&e, this.last_event_id.as_deref());
let retrying = this
.connector
.retry_connection(this.last_event_id.as_deref());
Expand Down
156 changes: 146 additions & 10 deletions crates/rmcp/src/transport/sse_client.rs
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,12 @@
//! reference: https://html.spec.whatwg.org/multipage/server-sent-events.html
use std::{pin::Pin, sync::Arc};
//! Reference: <https://html.spec.whatwg.org/multipage/server-sent-events.html>
use std::{
pin::Pin,
sync::{Arc, RwLock},
};

use futures::{StreamExt, future::BoxFuture};
use http::Uri;
use sse_stream::Error as SseError;
use sse_stream::{Error as SseError, Sse};
use thiserror::Error;

use super::{
Expand DownExpand Up@@ -54,9 +57,13 @@ pub trait SseClient: Clone + Send + Sync + 'static {
) -> impl Future<Output = Result<BoxedSseResponse, SseTransportError<Self::Error>>> + Send + '_;
}

/// Helper that refreshes the POST endpoint whenever the server emits
/// control frames during SSE reconnect; used together with
/// [`SseAutoReconnectStream`].
struct SseClientReconnect<C> {
pub client: C,
pub uri: Uri,
pub message_endpoint: Arc<RwLock<Uri>>,
}

impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
Expand All@@ -68,6 +75,37 @@ impl<C: SseClient> SseStreamReconnect for SseClientReconnect<C> {
let last_event_id = last_event_id.map(|s| s.to_owned());
Box::pin(async move { client.get_stream(uri, last_event_id, None).await })
}

fn handle_control_event(&mut self, event: &Sse) -> Result<(), Self::Error> {
if event.event.as_deref() != Some("endpoint") {
return Ok(());
}
let Some(data) = event.data.as_ref() else {
return Ok(());
};
// Servers typically resend the message POST endpoint (often with a new
// sessionId) when a stream reconnects. Reuse `message_endpoint` helper
// to resolve it and update the shared URI.
let new_endpoint = message_endpoint(self.uri.clone(), data.clone())
.map_err(SseTransportError::InvalidUri)?;
*self
.message_endpoint
.write()
.expect("message endpoint lock poisoned") = new_endpoint;
Ok(())
}

fn handle_stream_error(
&mut self,
error: &(dyn std::error::Error + 'static),
last_event_id: Option<&str>,
) {
tracing::warn!(
uri = %self.uri,
last_event_id = last_event_id.unwrap_or(""),
"sse stream error: {error}"
);
}
}
type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<C>>>>;

Expand All@@ -81,7 +119,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using reqwest
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::SseClientTransport;
///
/// // Enable the reqwest feature in Cargo.toml:
Expand All@@ -95,7 +133,7 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
///
/// ## Using a custom HTTP client
///
Comment thread
4t145 marked this conversation as resolved.
/// ```rust
/// ```rust,ignore
/// use rmcp::transport::sse_client::{SseClient, SseClientTransport, SseClientConfig};
/// use std::sync::Arc;
/// use futures::stream::BoxStream;
Expand DownExpand Up@@ -154,7 +192,9 @@ type ServerMessageStream<C> = Pin<Box<SseAutoReconnectStream<SseClientReconnect<
pub struct SseClientTransport<C: SseClient> {
client: C,
config: SseClientConfig,
message_endpoint: Uri,
/// Current POST endpoint; refreshed when the server sends new endpoint
/// control frames.
message_endpoint: Arc<RwLock<Uri>>,
stream: Option<ServerMessageStream<C>>,
}

Expand All@@ -168,8 +208,16 @@ impl<C: SseClient> Transport<RoleClient> for SseClientTransport<C> {
item: crate::service::TxJsonRpcMessage<RoleClient>,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let client = self.client.clone();
let uri = self.message_endpoint.clone();
async move { client.post_message(uri, item, None).await }
let message_endpoint = self.message_endpoint.clone();
async move {
let uri = {
let guard = message_endpoint
.read()
.expect("message endpoint lock poisoned");
guard.clone()
};
client.post_message(uri, item, None).await
}
}
async fn close(&mut self) -> Result<(), Self::Error> {
self.stream.take();
Expand All@@ -194,7 +242,7 @@ impl<C: SseClient> SseClientTransport<C> {
let sse_endpoint = config.sse_endpoint.as_ref().parse::<http::Uri>()?;

let mut sse_stream = client.get_stream(sse_endpoint.clone(), None, None).await?;
let message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let initial_message_endpoint = if let Some(endpoint) = config.use_message_endpoint.clone() {
let ep = endpoint.parse::<http::Uri>()?;
let mut sse_endpoint_parts = sse_endpoint.clone().into_parts();
sse_endpoint_parts.path_and_query = ep.into_parts().path_and_query;
Expand All@@ -214,12 +262,14 @@ impl<C: SseClient> SseClientTransport<C> {
break message_endpoint(sse_endpoint.clone(), ep)?;
}
};
let message_endpoint = Arc::new(RwLock::new(initial_message_endpoint));

let stream = Box::pin(SseAutoReconnectStream::new(
sse_stream,
SseClientReconnect {
client: client.clone(),
uri: sse_endpoint.clone(),
message_endpoint: message_endpoint.clone(),
},
config.retry_policy.clone(),
));
Expand DownExpand Up@@ -274,7 +324,7 @@ pub struct SseClientConfig {
/// and the server send the message endpoint event as `message?session_id=123`,
/// then the message endpoint will be `http://example.com/message`.
///
/// This follow the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/zh-CN/docs/Web/API/URL/URL)
/// This follows the rules of JavaScript's [`new URL(url, base)`](https://developer.mozilla.org/en-US/docs/Web/API/URL/URL)
pub sse_endpoint: Arc<str>,
pub retry_policy: Arc<dyn SseRetryPolicy>,
/// if this is settled, the client will use this endpoint to send message and skip get the endpoint event
Expand All@@ -293,8 +343,40 @@ impl Default for SseClientConfig {

#[cfg(test)]
mod tests {
use futures::StreamExt;
use serde_json::{Value, json};

use super::*;

#[derive(Clone)]
struct DummyClient;

#[derive(Debug, thiserror::Error)]
#[error("dummy error")]
struct DummyError;

impl SseClient for DummyClient {
type Error = DummyError;

async fn post_message(
&self,
_uri: Uri,
_message: ClientJsonRpcMessage,
_auth_token: Option<String>,
) -> Result<(), SseTransportError<Self::Error>> {
Ok(())
}

async fn get_stream(
&self,
_uri: Uri,
_last_event_id: Option<String>,
_auth_token: Option<String>,
) -> Result<BoxedSseResponse, SseTransportError<Self::Error>> {
unreachable!("get_stream should not be called in this test")
}
}

#[test]
fn test_message_endpoint() {
let base_url = "https://localhost/sse".parse::<http::Uri>().unwrap();
Expand All@@ -319,4 +401,58 @@ mod tests {
.unwrap();
assert_eq!(result.to_string(), "http://example.com/xxx?sessionId=x");
}

#[test]
fn handle_endpoint_control_event_updates_uri() {
let initial_endpoint = "https://example.com/message?sessionId=old"
.parse::<Uri>()
.unwrap();
let shared_endpoint = Arc::new(RwLock::new(initial_endpoint));
let mut reconnect = SseClientReconnect {
client: DummyClient,
uri: "https://example.com/sse".parse::<Uri>().unwrap(),
message_endpoint: shared_endpoint.clone(),
};

let control_event = Sse::default()
.event("endpoint")
.data("/message?sessionId=new");

reconnect.handle_control_event(&control_event).unwrap();

let guard = shared_endpoint.read().expect("lock poisoned");
assert_eq!(
guard.to_string(),
"https://example.com/message?sessionId=new"
);
}

#[tokio::test]
async fn control_event_frames_are_skipped() {
let payload = json!({
"jsonrpc": "2.0",
"id": 1,
"result": {"ok": true}
})
.to_string();

let events = vec![
Ok(Sse::default()
.event("endpoint")
.data("/message?sessionId=reconnect")),
Ok(Sse::default().event("message").data(payload.clone())),
];

let sse_src: BoxedSseResponse = futures::stream::iter(events).boxed();
let reconn_stream = SseAutoReconnectStream::never_reconnect(sse_src, DummyError);
futures::pin_mut!(reconn_stream);

let message = reconn_stream.next().await.expect("stream item").unwrap();
let actual: Value = serde_json::to_value(message).expect("serialize actual message");
// We only need to assert that a valid JSON-RPC response came through after
// skipping control frames. The exact `result` shape depends on the SDK's
// typed result enums and is not asserted here.
assert_eq!(actual.get("jsonrpc"), Some(&Value::String("2.0".into())));
assert_eq!(actual.get("id"), Some(&Value::Number(1u64.into())));
}
}
Loading