feat: add progress notification handling and related structures (#282)

* feat: add progress notification handling and related structures

also fix some lifetime bugs of tool macro

* fix: fmt
This commit is contained in:
4t145 2025-06-25 14:30:22 +08:00 committed by GitHub
parent 893969ac77
commit 8e1490bcd8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 298 additions and 18 deletions

View file

@ -207,12 +207,22 @@ pub fn tool(attr: TokenStream, input: TokenStream) -> syn::Result<TokenStream> {
// 2. make return type: `std::pin::Pin<Box<dyn Future<Output = #ReturnType> + Send + '_>>`
// 3. make body: { Box::pin(async move { #body }) }
let new_output = syn::parse2::<ReturnType>({
let mut lt = quote! { 'static };
if let Some(receiver) = fn_item.sig.receiver() {
if let Some((_, receiver_lt)) = receiver.reference.as_ref() {
if let Some(receiver_lt) = receiver_lt {
lt = quote! { #receiver_lt };
} else {
lt = quote! { '_ };
}
}
}
match &fn_item.sig.output {
syn::ReturnType::Default => {
quote! { -> std::pin::Pin<Box<dyn Future<Output = ()> + Send + '_>> }
quote! { -> std::pin::Pin<Box<dyn Future<Output = ()> + Send + #lt>> }
}
syn::ReturnType::Type(_, ty) => {
quote! { -> std::pin::Pin<Box<dyn Future<Output = #ty> + Send + '_>> }
quote! { -> std::pin::Pin<Box<dyn Future<Output = #ty> + Send + #lt>> }
}
}
})?;

View file

@ -71,7 +71,7 @@ chrono = { version = "0.4.38", default-features = false, features = ["serde", "c
[features]
default = ["base64", "macros", "server"]
client = []
client = ["dep:tokio-stream"]
server = ["transport-async-rw", "dep:schemars"]
macros = ["dep:rmcp-macros", "dep:paste"]
@ -191,3 +191,8 @@ path = "tests/test_message_protocol.rs"
name = "test_message_schema"
required-features = ["server", "client", "schemars"]
path = "tests/test_message_schema.rs"
[[test]]
name = "test_progress_subscriber"
required-features = ["server", "client", "macros"]
path = "tests/test_progress_subscriber.rs"

View file

@ -1,3 +1,4 @@
pub mod progress;
use crate::{
error::Error as McpError,
model::*,

View file

@ -0,0 +1,100 @@
use std::{collections::HashMap, sync::Arc};
use futures::{Stream, StreamExt};
use tokio::sync::RwLock;
use tokio_stream::wrappers::ReceiverStream;
use crate::model::{ProgressNotificationParam, ProgressToken};
type Dispatcher =
Arc<RwLock<HashMap<ProgressToken, tokio::sync::mpsc::Sender<ProgressNotificationParam>>>>;
/// A dispatcher for progress notifications.
#[derive(Debug, Clone, Default)]
pub struct ProgressDispatcher {
pub(crate) dispatcher: Dispatcher,
}
impl ProgressDispatcher {
const CHANNEL_SIZE: usize = 16;
pub fn new() -> Self {
Self::default()
}
/// Handle a progress notification by sending it to the appropriate subscriber
pub async fn handle_notification(&self, notification: ProgressNotificationParam) {
let token = &notification.progress_token;
if let Some(sender) = self.dispatcher.read().await.get(token).cloned() {
let send_result = sender.send(notification).await;
if let Err(e) = send_result {
tracing::warn!("Failed to send progress notification: {e}");
}
}
}
/// Subscribe to progress notifications for a specific token.
///
/// If you drop the returned `ProgressSubscriber`, it will automatically unsubscribe from notifications for that token.
pub async fn subscribe(&self, progress_token: ProgressToken) -> ProgressSubscriber {
let (sender, receiver) = tokio::sync::mpsc::channel(Self::CHANNEL_SIZE);
self.dispatcher
.write()
.await
.insert(progress_token.clone(), sender);
let receiver = ReceiverStream::new(receiver);
ProgressSubscriber {
progress_token,
receiver,
dispacher: self.dispatcher.clone(),
}
}
/// Unsubscribe from progress notifications for a specific token.
pub async fn unsubscribe(&self, token: &ProgressToken) {
self.dispatcher.write().await.remove(token);
}
/// Clear all dispachter.
pub async fn clear(&self) {
let mut dispacher = self.dispatcher.write().await;
dispacher.clear();
}
}
pub struct ProgressSubscriber {
pub(crate) progress_token: ProgressToken,
pub(crate) receiver: ReceiverStream<ProgressNotificationParam>,
pub(crate) dispacher: Dispatcher,
}
impl ProgressSubscriber {
pub fn progress_token(&self) -> &ProgressToken {
&self.progress_token
}
}
impl Stream for ProgressSubscriber {
type Item = ProgressNotificationParam;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.receiver.poll_next_unpin(cx)
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.receiver.size_hint()
}
}
impl Drop for ProgressSubscriber {
fn drop(&mut self) {
let token = self.progress_token.clone();
self.receiver.close();
let dispatcher = self.dispacher.clone();
tokio::spawn(async move {
let mut dispacher = dispatcher.write_owned().await;
dispacher.remove(&token);
});
}
}

View file

@ -99,11 +99,6 @@ pub trait FromToolCallContextPart<S>: Sized {
pub trait IntoCallToolResult {
fn into_call_tool_result(self) -> Result<CallToolResult, crate::Error>;
}
impl IntoCallToolResult for () {
fn into_call_tool_result(self) -> Result<CallToolResult, crate::Error> {
Ok(CallToolResult::success(vec![]))
}
}
impl<T: IntoContents> IntoCallToolResult for T {
fn into_call_tool_result(self) -> Result<CallToolResult, crate::Error> {
@ -120,6 +115,15 @@ impl<T: IntoContents, E: IntoContents> IntoCallToolResult for Result<T, E> {
}
}
impl<T: IntoCallToolResult> IntoCallToolResult for Result<T, crate::Error> {
fn into_call_tool_result(self) -> Result<CallToolResult, crate::Error> {
match self {
Ok(value) => value.into_call_tool_result(),
Err(error) => Err(error),
}
}
}
pin_project_lite::pin_project! {
#[project = IntoCallToolResultFutProj]
pub enum IntoCallToolResultFut<F, R> {

View file

@ -242,7 +242,7 @@ pub type RequestId = NumberOrString;
#[serde(transparent)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct ProgressToken(pub NumberOrString);
#[derive(Debug, Clone)]
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct Request<M = String, P = JsonObject> {
pub method: M,
@ -255,6 +255,16 @@ pub struct Request<M = String, P = JsonObject> {
pub extensions: Extensions,
}
impl<M: Default, P> Request<M, P> {
pub fn new(params: P) -> Self {
Self {
method: Default::default(),
params,
extensions: Extensions::default(),
}
}
}
impl<M, P> GetExtensions for Request<M, P> {
fn extensions(&self) -> &Extensions {
&self.extensions
@ -264,7 +274,7 @@ impl<M, P> GetExtensions for Request<M, P> {
}
}
#[derive(Debug, Clone)]
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct RequestOptionalParam<M = String, P = JsonObject> {
pub method: M,
@ -277,7 +287,17 @@ pub struct RequestOptionalParam<M = String, P = JsonObject> {
pub extensions: Extensions,
}
#[derive(Debug, Clone)]
impl<M: Default, P> RequestOptionalParam<M, P> {
pub fn with_param(params: P) -> Self {
Self {
method: Default::default(),
params: Some(params),
extensions: Extensions::default(),
}
}
}
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct RequestNoParam<M = String> {
pub method: M,
@ -296,7 +316,7 @@ impl<M> GetExtensions for RequestNoParam<M> {
&mut self.extensions
}
}
#[derive(Debug, Clone)]
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct Notification<M = String, P = JsonObject> {
pub method: M,
@ -308,7 +328,17 @@ pub struct Notification<M = String, P = JsonObject> {
pub extensions: Extensions,
}
#[derive(Debug, Clone)]
impl<M: Default, P> Notification<M, P> {
pub fn new(params: P) -> Self {
Self {
method: Default::default(),
params,
extensions: Extensions::default(),
}
}
}
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
pub struct NotificationNoParam<M = String> {
pub method: M,

View file

@ -165,3 +165,9 @@ impl IntoContents for String {
vec![Content::text(self)]
}
}
impl IntoContents for () {
fn into_contents(self) -> Vec<Content> {
vec![]
}
}

View file

@ -745,8 +745,9 @@ where
let mut extensions = Extensions::new();
let mut meta = Meta::new();
// avoid clone
std::mem::swap(&mut extensions, request.extensions_mut());
// swap meta firstly, otherwise progress token will be lost
std::mem::swap(&mut meta, request.get_meta_mut());
std::mem::swap(&mut extensions, request.extensions_mut());
let context = RequestContext {
ct: context_ct,
id: id.clone(),

View file

@ -0,0 +1,126 @@
use futures::StreamExt;
use rmcp::{
ClientHandler, Peer, RoleServer, ServerHandler, ServiceExt,
handler::{client::progress::ProgressDispatcher, server::tool::ToolRouter},
model::{CallToolRequestParam, ClientRequest, Meta, ProgressNotificationParam, Request},
service::PeerRequestOptions,
tool, tool_handler, tool_router,
};
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
pub struct MyClient {
progress_handler: ProgressDispatcher,
}
impl MyClient {
pub fn new() -> Self {
Self {
progress_handler: ProgressDispatcher::new(),
}
}
}
impl Default for MyClient {
fn default() -> Self {
Self::new()
}
}
impl ClientHandler for MyClient {
async fn on_progress(
&self,
params: rmcp::model::ProgressNotificationParam,
_context: rmcp::service::NotificationContext<rmcp::RoleClient>,
) {
tracing::info!("Received progress notification: {:?}", params);
self.progress_handler.handle_notification(params).await;
}
}
impl Default for MyServer {
fn default() -> Self {
Self::new()
}
}
pub struct MyServer {
tool_router: ToolRouter<Self>,
}
#[tool_router]
impl MyServer {
pub fn new() -> Self {
Self {
tool_router: Self::tool_router(),
}
}
#[tool]
pub async fn some_progress(meta: Meta, client: Peer<RoleServer>) -> Result<(), rmcp::Error> {
let progress_token = meta
.get_progress_token()
.ok_or(rmcp::Error::invalid_params(
"Progress token is required for this tool",
None,
))?;
for step in 0..10 {
let _ = client
.notify_progress(ProgressNotificationParam {
progress_token: progress_token.clone(),
progress: step,
total: Some(10),
message: Some("Some message".into()),
})
.await;
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
Ok(())
}
}
#[tool_handler]
impl ServerHandler for MyServer {}
#[tokio::test]
async fn test_progress_subscriber() -> anyhow::Result<()> {
let _ = tracing_subscriber::registry()
.with(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "debug".to_string().into()),
)
.with(tracing_subscriber::fmt::layer())
.try_init();
let client = MyClient::new();
let server = MyServer::new();
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(())
});
let client_service = client.serve(transport_client).await?;
let handle = client_service
.send_cancellable_request(
ClientRequest::CallToolRequest(Request::new(CallToolRequestParam {
name: "some_progress".into(),
arguments: None,
})),
PeerRequestOptions::no_options(),
)
.await?;
let mut progress_subscriber = client_service
.service()
.progress_handler
.subscribe(handle.progress_token.clone())
.await;
tokio::spawn(async move {
while let Some(notification) = progress_subscriber.next().await {
tracing::info!("Progress notification: {:?}", notification);
}
});
let _response = handle.await_response().await?;
// Simulate some delay to allow the async task to complete
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
Ok(())
}

View file

@ -43,10 +43,7 @@ impl<T> TestHandler<T> {
}
#[rmcp::tool]
async fn async_function<T>(
_callee: &TestHandler<T>,
Parameters(Request { fields }): Parameters<Request>,
) {
async fn async_function(Parameters(Request { fields }): Parameters<Request>) {
drop(fields)
}