Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 53 additions & 20 deletions crates/rmcp/src/service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,14 @@ pub trait ServiceRole: std::fmt::Debug + Send + Sync + 'static + Copy + Clone {
fn peer_cancelled_params(_notification: &Self::PeerNot) -> Option<&CancelledNotificationParam> {
None
}
#[doc(hidden)]
fn peer_cancels_subscriptions(_peer: &Peer<Self>) -> bool {
false
}
#[doc(hidden)]
fn is_subscription_request(_request: &Self::Req) -> bool {
false
}
/// Invalidate any response cache affected by an inbound peer notification.
///
/// The serve loop calls this for every notification *before* subscription
Expand Down Expand Up @@ -522,6 +530,10 @@ impl ProgressNotificationToken for ServerNotification {
}

type Responder<T> = tokio::sync::oneshot::Sender<T>;
struct PendingResponse<T> {
responder: Responder<T>,
is_subscription: bool,
}
type ProgressTimeoutWatchers = Arc<tokio::sync::RwLock<HashMap<ProgressToken, mpsc::Sender<()>>>>;
type SubscriptionChannel<N> = (mpsc::Sender<N>, usize);
type SubscriptionChannelMap<N> = HashMap<RequestId, SubscriptionChannel<N>>;
Expand Down Expand Up @@ -1348,7 +1360,7 @@ where
}

let mut local_responder_pool =
HashMap::<RequestId, Responder<Result<R::PeerResp, ServiceError>>>::new();
HashMap::<RequestId, PendingResponse<Result<R::PeerResp, ServiceError>>>::new();
let mut local_ct_pool = HashMap::<RequestId, CancellationToken>::new();
let shared_service = Arc::new(service);
// for return
Expand Down Expand Up @@ -1445,7 +1457,7 @@ where
Event::SendTaskResult(SendTaskResult::Request { id, result }) => {
if let Err(e) = result
&& let Some(responder) = local_responder_pool.remove(&id) {
let _ = responder.send(Err(ServiceError::TransportSend(e)));
let _ = responder.responder.send(Err(ServiceError::TransportSend(e)));
}
}
Event::SendTaskResult(SendTaskResult::Notification {
Expand All @@ -1463,7 +1475,7 @@ where
&& let Some(request_id) = &param.request_id
&& let Some(responder) = local_responder_pool.remove(request_id) {
tracing::info!(id = %request_id, reason = param.reason, "cancelled");
let _response_result = responder.send(Err(ServiceError::Cancelled {
let _response_result = responder.responder.send(Err(ServiceError::Cancelled {
reason: param.reason.clone(),
}));
}
Expand Down Expand Up @@ -1500,7 +1512,10 @@ where
id,
responder,
}) => {
local_responder_pool.insert(id.clone(), responder);
local_responder_pool.insert(id.clone(), PendingResponse {
responder,
is_subscription: R::is_subscription_request(&request),
});
let send = transport.send(JsonRpcMessage::request(request, id.clone()));
{
let id = id.clone();
Expand Down Expand Up @@ -1601,31 +1616,38 @@ where
})) => {
tracing::info!(?notification, "received notification");
R::invalidate_response_cache(&peer, &notification).await;
let cancellation_request_id =
let subscription_id =
if let Some(cancelled) = R::peer_cancelled_params(&notification) {
let request_id = cancelled.request_id.clone();
if let Some(request_id) = request_id.as_ref() {
if R::IS_CLIENT {
if let Some(responder) =
local_responder_pool.remove(request_id)
// Modern servers cancel listen requests; legacy peers cancel
// requests they originated, even when both directions share an ID.
if R::peer_cancels_subscriptions(&peer) {
if let Some(request_id) = request_id.as_ref() {
if local_responder_pool.get(request_id)
.is_some_and(|pending| pending.is_subscription)
&& let Some(pending) = local_responder_pool.remove(request_id)
Comment thread
aurokin marked this conversation as resolved.
{
let _ = responder.send(Err(ServiceError::Cancelled {
let _ = pending.responder.send(Err(ServiceError::Cancelled {
reason: cancelled.reason.clone(),
}));
} else {
tracing::debug!(%request_id, "ignoring cancellation of unknown subscription");
continue;
Comment thread
aurokin marked this conversation as resolved.
}
} else if let Some(ct) = local_ct_pool.remove(request_id) {
}
request_id
} else {
if let Some(request_id) = request_id.as_ref()
&& let Some(ct) = local_ct_pool.remove(request_id)
{
tracing::info!(id = %request_id, reason = cancelled.reason, "cancelled");
ct.cancel();
}
None
}
request_id
} else {
None
notification.get_meta().subscription_id()
};
let subscription_id = notification
.get_meta()
.subscription_id()
.or(cancellation_request_id);
if let Some(subscription_id) = subscription_id
&& let Some((sender, capacity)) =
peer.subscription_sender(&subscription_id)
Expand All @@ -1642,7 +1664,7 @@ where
&& let Some(responder) =
local_responder_pool.remove(&subscription_id)
{
let _ = responder
let _ = responder.responder
.send(Err(ServiceError::SubscriptionLagged { capacity }));
}
peer.unregister_subscription(&subscription_id);
Expand Down Expand Up @@ -1694,7 +1716,7 @@ where
if let Some(responder) =
remove_pending_request(&mut local_responder_pool, &id)
{
let response_result = responder.send(Ok(result));
let response_result = responder.responder.send(Ok(result));
if let Err(_error) = response_result {
tracing::warn!(%id, "Error sending response");
}
Expand All @@ -1715,7 +1737,7 @@ where
} else {
ServiceError::McpError(error)
};
let _response_result = responder.send(Err(service_error));
let _response_result = responder.responder.send(Err(service_error));
if let Err(_error) = _response_result {
tracing::warn!(%id, "Error sending response");
}
Expand Down Expand Up @@ -1749,6 +1771,17 @@ where
// Then drain any handler responses still in the channel
// (handlers that finished after the loop broke).
while let Some(m) = sink_proxy_rx.recv().await {
if let Some(id) = match &m {
JsonRpcMessage::Response(response) => Some(&response.id),
JsonRpcMessage::Error(error) => error.id.as_ref(),
_ => None,
} {
let Some(ct) = local_ct_pool.remove(id) else {
tracing::debug!(%id, "dropping response for cancelled request");
continue;
};
ct.cancel();
}
if let Err(error) = transport.send(m).await {
tracing::error!(%error, "failed to send pending response during drain");
break;
Expand Down
Loading