Add progress-aware request timeout reset (#858)

* feat: add progress-aware request timeouts

* Update crates/rmcp/src/service.rs

Co-authored-by: Dale Seo <5466341+DaleSeo@users.noreply.github.com>

* refactor(rmcp): move helpers and simplify response waiting

---------

Co-authored-by: Dale Seo <5466341+DaleSeo@users.noreply.github.com>
This commit is contained in:
ContextVM-org 2026-06-17 21:07:37 +02:00 committed by GitHub
parent 4b82e41522
commit 5d00e20f2a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 441 additions and 21 deletions

View file

@ -252,6 +252,11 @@ name = "test_progress_subscriber"
required-features = ["server", "client", "macros"] required-features = ["server", "client", "macros"]
path = "tests/test_progress_subscriber.rs" path = "tests/test_progress_subscriber.rs"
[[test]]
name = "test_request_timeout_progress"
required-features = ["server", "client", "macros"]
path = "tests/test_request_timeout_progress.rs"
[[test]] [[test]]
name = "test_elicitation" name = "test_elicitation"
required-features = ["elicitation", "client", "server"] required-features = ["elicitation", "client", "server"]

View file

@ -42,8 +42,12 @@ pub(crate) type MaybeBoxFuture<'a, T> = BoxFuture<'a, T>;
#[cfg(feature = "local")] #[cfg(feature = "local")]
pub(crate) type MaybeBoxFuture<'a, T> = LocalBoxFuture<'a, T>; pub(crate) type MaybeBoxFuture<'a, T> = LocalBoxFuture<'a, T>;
#[cfg(feature = "server")]
use crate::model::ClientNotification;
#[cfg(feature = "server")] #[cfg(feature = "server")]
use crate::model::ServerJsonRpcMessage; use crate::model::ServerJsonRpcMessage;
#[cfg(feature = "client")]
use crate::model::ServerNotification;
use crate::{ use crate::{
error::ErrorData as McpError, error::ErrorData as McpError,
model::{ model::{
@ -299,7 +303,37 @@ impl ProgressTokenProvider for AtomicU32Provider {
} }
} }
#[doc(hidden)]
pub trait ProgressNotificationToken {
fn progress_token(&self) -> Option<&ProgressToken>;
}
#[cfg(feature = "server")]
impl ProgressNotificationToken for ClientNotification {
fn progress_token(&self) -> Option<&ProgressToken> {
match self {
ClientNotification::ProgressNotification(notification) => {
Some(&notification.params.progress_token)
}
_ => None,
}
}
}
#[cfg(feature = "client")]
impl ProgressNotificationToken for ServerNotification {
fn progress_token(&self) -> Option<&ProgressToken> {
match self {
ServerNotification::ProgressNotification(notification) => {
Some(&notification.params.progress_token)
}
_ => None,
}
}
}
type Responder<T> = tokio::sync::oneshot::Sender<T>; type Responder<T> = tokio::sync::oneshot::Sender<T>;
type ProgressTimeoutWatchers = Arc<tokio::sync::RwLock<HashMap<ProgressToken, mpsc::Sender<()>>>>;
/// A handle to a remote request /// A handle to a remote request
/// ///
@ -314,40 +348,126 @@ pub struct RequestHandle<R: ServiceRole> {
pub peer: Peer<R>, pub peer: Peer<R>,
pub id: RequestId, pub id: RequestId,
pub progress_token: ProgressToken, pub progress_token: ProgressToken,
progress_timeout_watchers: ProgressTimeoutWatchers,
progress_reset_rx: Option<mpsc::Receiver<()>>,
} }
impl<R: ServiceRole> RequestHandle<R> { impl<R: ServiceRole> RequestHandle<R> {
pub const REQUEST_TIMEOUT_REASON: &str = "request timeout"; pub const REQUEST_TIMEOUT_REASON: &str = "request timeout";
pub async fn await_response(self) -> Result<R::PeerResp, ServiceError> { pub const REQUEST_MAX_TOTAL_TIMEOUT_REASON: &str = "maximum total timeout exceeded";
if let Some(timeout) = self.options.timeout {
let timeout_result = tokio::time::timeout(timeout, async move { pub async fn await_response(mut self) -> Result<R::PeerResp, ServiceError> {
self.rx.await.map_err(|_e| ServiceError::TransportClosed)? let timeout = self.options.timeout;
}) let max_total_timeout = self.options.max_total_timeout;
.await; let reset_timeout_on_progress = self.options.reset_timeout_on_progress;
match timeout_result {
Ok(response) => response, let has_progress_reset_rx = self.progress_reset_rx.is_some();
let progress_token = self.progress_token.clone();
let result = match (timeout, max_total_timeout, reset_timeout_on_progress) {
(Some(timeout), None, false) => match tokio::time::timeout(timeout, &mut self.rx).await
{
Ok(response) => response.map_err(|_e| ServiceError::TransportClosed)?,
Err(_) => { Err(_) => {
let error = Err(ServiceError::Timeout { timeout }); let error = Err(ServiceError::Timeout { timeout });
// cancel this request // cancel this request
let notification = CancelledNotification { self.send_timeout_cancel_notification(Self::REQUEST_TIMEOUT_REASON)
params: CancelledNotificationParam { .await;
request_id: self.id,
reason: Some(Self::REQUEST_TIMEOUT_REASON.to_owned()),
},
method: crate::model::CancelledNotificationMethod,
extensions: Default::default(),
};
let _ = self.peer.send_notification(notification.into()).await;
error error
} }
},
(None, None, _) => (&mut self.rx)
.await
.map_err(|_e| ServiceError::TransportClosed)?,
_ => {
self.await_response_with_progress_timeout(
timeout,
max_total_timeout,
reset_timeout_on_progress,
)
.await
}
};
Self::cleanup_progress_timeout_watcher(
&self.peer.progress_timeout_watchers,
&progress_token,
has_progress_reset_rx,
)
.await;
result
}
async fn send_timeout_cancel_notification(&self, reason: &str) {
let notification = CancelledNotification {
params: CancelledNotificationParam {
request_id: self.id.clone(),
reason: Some(reason.to_owned()),
},
method: crate::model::CancelledNotificationMethod,
extensions: Default::default(),
};
let _ = self.peer.send_notification(notification.into()).await;
}
async fn await_response_with_progress_timeout(
&mut self,
timeout: Option<Duration>,
max_total_timeout: Option<Duration>,
reset_timeout_on_progress: bool,
) -> Result<R::PeerResp, ServiceError> {
let mut idle_sleep = timeout.map(tokio::time::sleep).map(Box::pin);
let mut max_total_sleep = max_total_timeout.map(tokio::time::sleep).map(Box::pin);
loop {
tokio::select! {
biased;
response = &mut self.rx => {
return response.map_err(|_e| ServiceError::TransportClosed)?;
}
_ = async {
if let Some(sleep) = idle_sleep.as_mut() {
sleep.as_mut().await;
}
}, if idle_sleep.is_some() => {
let timeout = timeout.expect("idle timeout exists when idle sleep exists");
self.send_timeout_cancel_notification(Self::REQUEST_TIMEOUT_REASON).await;
return Err(ServiceError::Timeout { timeout });
}
_ = async {
if let Some(sleep) = max_total_sleep.as_mut() {
sleep.as_mut().await;
}
}, if max_total_sleep.is_some() => {
let timeout = max_total_timeout.expect("max total timeout exists when max total sleep exists");
self.send_timeout_cancel_notification(Self::REQUEST_MAX_TOTAL_TIMEOUT_REASON).await;
return Err(ServiceError::Timeout { timeout });
}
progress = async {
match self.progress_reset_rx.as_mut() {
Some(rx) => rx.recv().await,
None => None,
}
}, if reset_timeout_on_progress && timeout.is_some() && self.progress_reset_rx.is_some() => {
if progress.is_some() {
if let (Some(timeout), Some(sleep)) = (timeout, idle_sleep.as_mut()) {
sleep.as_mut().reset(tokio::time::Instant::now() + timeout);
}
}
}
} }
} else {
self.rx.await.map_err(|_e| ServiceError::TransportClosed)?
} }
} }
/// Cancel this request /// Cancel this request
pub async fn cancel(self, reason: Option<String>) -> Result<(), ServiceError> { pub async fn cancel(self, reason: Option<String>) -> Result<(), ServiceError> {
Self::cleanup_progress_timeout_watcher(
&self.progress_timeout_watchers,
&self.progress_token,
self.progress_reset_rx.is_some(),
)
.await;
let notification = CancelledNotification { let notification = CancelledNotification {
params: CancelledNotificationParam { params: CancelledNotificationParam {
request_id: self.id, request_id: self.id,
@ -359,6 +479,19 @@ impl<R: ServiceRole> RequestHandle<R> {
self.peer.send_notification(notification.into()).await?; self.peer.send_notification(notification.into()).await?;
Ok(()) Ok(())
} }
async fn cleanup_progress_timeout_watcher(
progress_timeout_watchers: &ProgressTimeoutWatchers,
progress_token: &ProgressToken,
has_progress_reset_rx: bool,
) {
if has_progress_reset_rx {
progress_timeout_watchers
.write()
.await
.remove(progress_token);
}
}
} }
#[derive(Debug)] #[derive(Debug)]
@ -384,6 +517,7 @@ pub struct Peer<R: ServiceRole> {
tx: mpsc::Sender<PeerSinkMessage<R>>, tx: mpsc::Sender<PeerSinkMessage<R>>,
request_id_provider: Arc<dyn RequestIdProvider>, request_id_provider: Arc<dyn RequestIdProvider>,
progress_token_provider: Arc<dyn ProgressTokenProvider>, progress_token_provider: Arc<dyn ProgressTokenProvider>,
progress_timeout_watchers: ProgressTimeoutWatchers,
info: Arc<std::sync::RwLock<Option<Arc<R::PeerInfo>>>>, info: Arc<std::sync::RwLock<Option<Arc<R::PeerInfo>>>>,
} }
@ -403,12 +537,33 @@ type ProxyOutbound<R> = mpsc::Receiver<PeerSinkMessage<R>>;
pub struct PeerRequestOptions { pub struct PeerRequestOptions {
pub timeout: Option<Duration>, pub timeout: Option<Duration>,
pub meta: Option<Meta>, pub meta: Option<Meta>,
/// Reset the request timeout when a matching progress notification is received.
pub reset_timeout_on_progress: bool,
/// Maximum total time to wait for the request, regardless of progress notifications.
pub max_total_timeout: Option<Duration>,
} }
impl PeerRequestOptions { impl PeerRequestOptions {
pub fn no_options() -> Self { pub fn no_options() -> Self {
Self::default() Self::default()
} }
pub fn with_timeout(timeout: Duration) -> Self {
Self {
timeout: Some(timeout),
..Self::default()
}
}
pub fn reset_timeout_on_progress(mut self) -> Self {
self.reset_timeout_on_progress = true;
self
}
pub fn with_max_total_timeout(mut self, timeout: Duration) -> Self {
self.max_total_timeout = Some(timeout);
self
}
} }
impl<R: ServiceRole> Peer<R> { impl<R: ServiceRole> Peer<R> {
@ -423,6 +578,7 @@ impl<R: ServiceRole> Peer<R> {
tx, tx,
request_id_provider, request_id_provider,
progress_token_provider: Arc::new(AtomicU32ProgressTokenProvider::default()), progress_token_provider: Arc::new(AtomicU32ProgressTokenProvider::default()),
progress_timeout_watchers: Default::default(),
info: Arc::new(std::sync::RwLock::new(peer_info.map(Arc::new))), info: Arc::new(std::sync::RwLock::new(peer_info.map(Arc::new))),
}, },
rx, rx,
@ -468,22 +624,68 @@ impl<R: ServiceRole> Peer<R> {
request.get_meta_mut().extend(meta); request.get_meta_mut().extend(meta);
} }
let (responder, receiver) = tokio::sync::oneshot::channel(); let (responder, receiver) = tokio::sync::oneshot::channel();
self.tx let progress_reset_rx = if options.reset_timeout_on_progress && options.timeout.is_some() {
let (sender, receiver) = mpsc::channel(1);
self.progress_timeout_watchers
.write()
.await
.insert(progress_token.clone(), sender);
Some(receiver)
} else {
None
};
if self
.tx
.send(PeerSinkMessage::Request { .send(PeerSinkMessage::Request {
request, request,
id: id.clone(), id: id.clone(),
responder, responder,
}) })
.await .await
.map_err(|_m| ServiceError::TransportClosed)?; .is_err()
{
if progress_reset_rx.is_some() {
self.progress_timeout_watchers
.write()
.await
.remove(&progress_token);
}
return Err(ServiceError::TransportClosed);
}
Ok(RequestHandle { Ok(RequestHandle {
id, id,
rx: receiver, rx: receiver,
progress_token, progress_token,
options, options,
peer: self.clone(), peer: self.clone(),
progress_timeout_watchers: self.progress_timeout_watchers.clone(),
progress_reset_rx,
}) })
} }
async fn notify_progress_timeout_watcher(&self, progress_token: &ProgressToken) {
let sender = self
.progress_timeout_watchers
.read()
.await
.get(progress_token)
.cloned();
if let Some(sender) = sender {
match sender.try_send(()) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
tracing::trace!(?progress_token, "progress timeout watcher channel is full");
}
Err(mpsc::error::TrySendError::Closed(_)) => {
self.progress_timeout_watchers
.write()
.await
.remove(progress_token);
}
}
}
}
/// Snapshot of the peer's handshake info. /// Snapshot of the peer's handshake info.
pub fn peer_info(&self) -> Option<Arc<R::PeerInfo>> { pub fn peer_info(&self) -> Option<Arc<R::PeerInfo>> {
self.info.read().expect("peer info lock poisoned").clone() self.info.read().expect("peer info lock poisoned").clone()
@ -700,6 +902,7 @@ pub fn serve_directly<R, S, T, E, A>(
) -> RunningService<R, S> ) -> RunningService<R, S>
where where
R: ServiceRole, R: ServiceRole,
R::PeerNot: ProgressNotificationToken,
S: Service<R>, S: Service<R>,
T: IntoTransport<R, E, A>, T: IntoTransport<R, E, A>,
E: std::error::Error + Send + Sync + 'static, E: std::error::Error + Send + Sync + 'static,
@ -716,6 +919,7 @@ pub fn serve_directly_with_ct<R, S, T, E, A>(
) -> RunningService<R, S> ) -> RunningService<R, S>
where where
R: ServiceRole, R: ServiceRole,
R::PeerNot: ProgressNotificationToken,
S: Service<R>, S: Service<R>,
T: IntoTransport<R, E, A>, T: IntoTransport<R, E, A>,
E: std::error::Error + Send + Sync + 'static, E: std::error::Error + Send + Sync + 'static,
@ -756,6 +960,7 @@ fn serve_inner<R, S, T>(
) -> RunningService<R, S> ) -> RunningService<R, S>
where where
R: ServiceRole, R: ServiceRole,
R::PeerNot: ProgressNotificationToken,
S: Service<R>, S: Service<R>,
T: Transport<R> + 'static, T: Transport<R> + 'static,
{ {
@ -1002,6 +1207,9 @@ where
} }
Err(notification) => notification, Err(notification) => notification,
}; };
if let Some(progress_token) = notification.progress_token() {
peer.notify_progress_timeout_watcher(progress_token).await;
}
{ {
let service = shared_service.clone(); let service = shared_service.clone();
let mut extensions = Extensions::new(); let mut extensions = Extensions::new();

View file

@ -362,6 +362,8 @@ macro_rules! method {
let options = crate::service::PeerRequestOptions { let options = crate::service::PeerRequestOptions {
timeout, timeout,
meta: None, meta: None,
reset_timeout_on_progress: false,
max_total_timeout: None,
}; };
let result = self let result = self
.send_request_with_option(request, options) .send_request_with_option(request, options)
@ -390,6 +392,8 @@ macro_rules! method {
let options = crate::service::PeerRequestOptions { let options = crate::service::PeerRequestOptions {
timeout, timeout,
meta: None, meta: None,
reset_timeout_on_progress: false,
max_total_timeout: None,
}; };
let result = self let result = self
.send_request_with_option(request, options) .send_request_with_option(request, options)

View file

@ -0,0 +1,203 @@
#![cfg(not(feature = "local"))]
use std::{
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use rmcp::{
ClientHandler, Peer, RoleServer, ServiceError, ServiceExt,
model::{CallToolRequestParams, ClientRequest, Meta, ProgressNotificationParam, Request},
service::PeerRequestOptions,
tool, tool_router,
};
#[derive(Clone, Default)]
struct ProgressCountingClient {
progress_count: Arc<AtomicUsize>,
}
impl ClientHandler for ProgressCountingClient {
async fn on_progress(
&self,
_params: ProgressNotificationParam,
_context: rmcp::service::NotificationContext<rmcp::RoleClient>,
) {
self.progress_count.fetch_add(1, Ordering::SeqCst);
}
}
struct ProgressTimeoutServer;
impl ProgressTimeoutServer {
fn new() -> Self {
Self
}
}
#[tool_router(server_handler)]
impl ProgressTimeoutServer {
#[tool]
async fn delayed_without_progress(&self) -> Result<(), rmcp::ErrorData> {
tokio::time::sleep(Duration::from_millis(250)).await;
Ok(())
}
#[tool]
async fn delayed_with_progress(
&self,
meta: Meta,
client: Peer<RoleServer>,
) -> Result<(), rmcp::ErrorData> {
let progress_token = meta
.get_progress_token()
.ok_or(rmcp::ErrorData::invalid_params(
"Progress token is required",
None,
))?;
for step in 0..4 {
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = client
.notify_progress(ProgressNotificationParam {
progress_token: progress_token.clone(),
progress: step as f64,
total: Some(4.0),
message: Some("working".into()),
})
.await;
}
Ok(())
}
#[tool]
async fn delayed_with_unrelated_progress(
&self,
client: Peer<RoleServer>,
) -> Result<(), rmcp::ErrorData> {
for step in 0..4 {
tokio::time::sleep(Duration::from_millis(50)).await;
let _ = client
.notify_progress(ProgressNotificationParam {
progress_token: rmcp::model::ProgressToken(
rmcp::model::NumberOrString::Number(999_999),
),
progress: step as f64,
total: Some(4.0),
message: Some("unrelated".into()),
})
.await;
}
Ok(())
}
}
async fn start_pair()
-> anyhow::Result<rmcp::service::RunningService<rmcp::RoleClient, ProgressCountingClient>> {
let server = ProgressTimeoutServer::new();
let client = ProgressCountingClient::default();
let (transport_server, transport_client) = tokio::io::duplex(4096);
tokio::spawn(async move {
let service = server.serve(transport_server).await?;
service.waiting().await?;
anyhow::Ok(())
});
Ok(client.serve(transport_client).await?)
}
async fn call_tool_with_options(
client: &rmcp::service::RunningService<rmcp::RoleClient, ProgressCountingClient>,
name: &str,
options: PeerRequestOptions,
) -> Result<rmcp::model::ServerResult, ServiceError> {
client
.send_request_with_option(
ClientRequest::CallToolRequest(Request::new(CallToolRequestParams::new(
name.to_owned(),
))),
options,
)
.await?
.await_response()
.await
}
#[tokio::test]
async fn request_timeout_still_expires_without_progress() -> anyhow::Result<()> {
let client = start_pair().await?;
let result = call_tool_with_options(
&client,
"delayed_without_progress",
PeerRequestOptions::with_timeout(Duration::from_millis(75)),
)
.await;
assert!(matches!(result, Err(ServiceError::Timeout { .. })));
Ok(())
}
#[tokio::test]
async fn progress_does_not_reset_timeout_by_default() -> anyhow::Result<()> {
let client = start_pair().await?;
let result = call_tool_with_options(
&client,
"delayed_with_progress",
PeerRequestOptions::with_timeout(Duration::from_millis(75)),
)
.await;
assert!(matches!(result, Err(ServiceError::Timeout { .. })));
Ok(())
}
#[tokio::test]
async fn matching_progress_resets_timeout_when_enabled() -> anyhow::Result<()> {
let client = start_pair().await?;
let result = call_tool_with_options(
&client,
"delayed_with_progress",
PeerRequestOptions::with_timeout(Duration::from_millis(75)).reset_timeout_on_progress(),
)
.await;
assert!(result.is_ok());
assert!(client.service().progress_count.load(Ordering::SeqCst) > 0);
Ok(())
}
#[tokio::test]
async fn max_total_timeout_wins_over_progress_reset() -> anyhow::Result<()> {
let client = start_pair().await?;
let result = call_tool_with_options(
&client,
"delayed_with_progress",
PeerRequestOptions::with_timeout(Duration::from_millis(75))
.reset_timeout_on_progress()
.with_max_total_timeout(Duration::from_millis(125)),
)
.await;
assert!(matches!(result, Err(ServiceError::Timeout { .. })));
Ok(())
}
#[tokio::test]
async fn unrelated_progress_does_not_reset_timeout() -> anyhow::Result<()> {
let client = start_pair().await?;
let result = call_tool_with_options(
&client,
"delayed_with_unrelated_progress",
PeerRequestOptions::with_timeout(Duration::from_millis(75)).reset_timeout_on_progress(),
)
.await;
assert!(matches!(result, Err(ServiceError::Timeout { .. })));
Ok(())
}