use std::time::Duration;
use futures::future::BoxFuture;
use rmcp::model::{CallToolRequestParams, CallToolResult};
use rmcp::service::{Peer, PeerRequestOptions, ServiceError};
use rmcp::RoleClient;
use tokio::task::{JoinError, JoinHandle};
use crate::core::error::{Error, Result};
use crate::transport::client::ClientOpenStreamHandle;
use crate::transport::open_stream::OpenStreamSession;
use super::progress::PeerRequestOptionsExt;
type AbortFn = Box<dyn Fn(Option<String>) -> BoxFuture<'static, ()> + Send + Sync>;
pub struct ToolStreamCall {
pub progress_token: String,
pub stream: OpenStreamSession,
pub result: BoxFuture<'static, Result<CallToolResult>>,
abort_fn: AbortFn,
}
impl ToolStreamCall {
pub async fn abort(&self, reason: Option<String>) {
(self.abort_fn)(reason).await;
}
}
impl std::fmt::Debug for ToolStreamCall {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolStreamCall")
.field("progress_token", &self.progress_token)
.finish_non_exhaustive()
}
}
fn open_stream_request_options(handle: &ClientOpenStreamHandle) -> PeerRequestOptions {
let config = handle.config();
let idle_ms = config
.idle_timeout_ms
.saturating_add(config.probe_timeout_ms)
.saturating_add(config.close_grace_period_ms)
.max(1);
let mut options = PeerRequestOptions::with_timeout(Duration::from_millis(idle_ms))
.reset_timeout_on_progress();
if let Some(max_total_ms) = config.max_total_timeout_ms {
options = options.with_max_total_timeout(Duration::from_millis(max_total_ms));
}
options
}
pub async fn call_tool_stream(
peer: &Peer<RoleClient>,
transport: &ClientOpenStreamHandle,
params: CallToolRequestParams,
) -> Result<ToolStreamCall> {
let bind_guard = transport.bind_lock().clone().lock_owned().await;
let pending = transport.prepare_outbound();
let options = open_stream_request_options(transport);
let peer = peer.clone();
let mut result_handle: JoinHandle<std::result::Result<CallToolResult, ServiceError>> =
tokio::spawn(async move { peer.call_tool_with_options(params, options).await });
let (progress_token, stream) = tokio::select! {
biased;
bound = pending => match bound {
Ok(Ok(pair)) => pair,
Ok(Err(error)) => {
drop(bind_guard);
result_handle.abort();
return Err(error);
}
Err(_) => {
drop(bind_guard);
result_handle.abort();
return Err(Error::Transport(
"transport closed before the outbound open-stream session was bound"
.to_string(),
));
}
},
settled = &mut result_handle => {
transport.cancel_outbound();
drop(bind_guard);
return Err(match flatten_call_result(settled) {
Err(error) => error,
Ok(_) => Error::Other(
"open-stream tool call completed without establishing a stream".to_string(),
),
});
}
};
drop(bind_guard);
let registry = transport.registry();
let abort_session = stream.clone();
let abort_token = progress_token.clone();
let abort_fn: AbortFn = Box::new(move |reason: Option<String>| {
let registry = registry.clone();
let session = abort_session.clone();
let token = abort_token.clone();
Box::pin(async move {
session.abort(reason.clone()).await;
registry.lock().await.consumer_abort(&token, reason).await;
})
});
let result: BoxFuture<'static, Result<CallToolResult>> =
Box::pin(async move { flatten_call_result(result_handle.await) });
Ok(ToolStreamCall {
progress_token,
stream,
result,
abort_fn,
})
}
fn flatten_call_result(
settled: std::result::Result<std::result::Result<CallToolResult, ServiceError>, JoinError>,
) -> Result<CallToolResult> {
match settled {
Ok(Ok(result)) => Ok(result),
Ok(Err(service_error)) => Err(Error::Transport(service_error.to_string())),
Err(join_error) => Err(Error::Other(format!(
"call_tool_stream task failed: {join_error}"
))),
}
}