Skip to main content

mcp_utils/client/
call_tool.rs

1use crate::client::McpClient;
2use crate::client::elicitation::{ElicitInputsError, elicit_inputs};
3use crate::client::mrtr::{AbortReason, MrtrAction, MrtrState};
4use crate::client::task::{TaskDriver, TaskErrorReason};
5use async_stream::stream;
6use futures::Stream;
7use futures::StreamExt;
8use futures::future::{Either, select};
9use rmcp::RoleClient;
10use rmcp::model::{
11    CallToolRequestParams, CallToolResponse, CallToolResult, ClientRequest, CreateTaskResult, InputRequests,
12    InputResponses, ProgressNotificationParam, Request, RequestMetaObject, ServerResult, Task,
13};
14use rmcp::service::{PeerRequestOptions, RequestHandle, RunningService, ServiceError};
15use std::pin::pin;
16use std::sync::Arc;
17use std::time::Duration;
18use thiserror::Error;
19use tokio::time::sleep;
20use tokio_util::sync::CancellationToken;
21
22#[derive(Debug, Default)]
23pub struct CallToolOptions {
24    pub timeout: Duration,
25    pub meta: Option<RequestMetaObject>,
26    pub cancel: CancellationToken,
27}
28
29#[derive(Debug)]
30pub enum ToolCallEvent {
31    Progress(ProgressNotificationParam),
32    TaskCreated(CreateTaskResult),
33    TaskStatus(Task),
34    Complete(Result<CallToolResult, CallToolError>),
35    TaskComplete { task: Task, result: Result<CallToolResult, CallToolError> },
36    Cancelled { task_id: Option<String> },
37}
38
39#[derive(Debug, Error)]
40pub enum CallToolError {
41    #[error("Failed to send tool request: {0}")]
42    Send(#[source] ServiceError),
43    #[error("Tool execution failed: {0}")]
44    Call(#[source] ServiceError),
45    #[error("Server '{server}' requested an input kind this client does not support (sampling or roots)")]
46    UnsupportedInput { server: String },
47    #[error("Server '{server}' failed to serialize an elicitation response: {source}")]
48    SerializeResponse {
49        server: String,
50        #[source]
51        source: serde_json::Error,
52    },
53    #[error("{}", reason.message(server, *timeout))]
54    Aborted { server: String, reason: AbortReason, timeout: Duration },
55    #[error("Task '{task_id}' from server '{server}' {reason}")]
56    Task {
57        server: String,
58        task_id: String,
59        #[source]
60        reason: Box<TaskErrorReason>,
61    },
62    #[error("Server '{server}' returned a tool call response kind this client does not support")]
63    UnsupportedResponse { server: String },
64    #[error("{message}")]
65    Unavailable { message: String },
66}
67
68pub fn call_tool(
69    client: Arc<RunningService<RoleClient, McpClient>>,
70    mut params: CallToolRequestParams,
71    options: CallToolOptions,
72) -> impl Stream<Item = ToolCallEvent> {
73    stream! {
74        let server_name = client.service().server_name().to_string();
75        let mut mrtr_state = MrtrState::new(options.timeout);
76
77        loop {
78            let request = ClientRequest::CallToolRequest(Request::new(params.clone()));
79            let send = pin!(client.send_cancellable_request(request, peer_request_options(&options)));
80            let handle = match select(send, pin!(options.cancel.cancelled())).await {
81                Either::Left((Ok(handle), _)) => handle,
82                Either::Left((Err(e), _)) => {
83                    yield ToolCallEvent::Complete(Err(CallToolError::Send(e)));
84                    return;
85                }
86                Either::Right(((), _)) => {
87                    yield ToolCallEvent::Cancelled { task_id: None };
88                    return;
89                }
90            };
91
92            let mut progress = client.service().progress_dispatcher.subscribe(handle.progress_token.clone()).await;
93            let mut response_or_cancel = pin!(await_response_or_cancel(handle, &options.cancel));
94            let response = loop {
95                match select(progress.next(), response_or_cancel.as_mut()).await {
96                    Either::Left((Some(progress), _)) => yield ToolCallEvent::Progress(progress),
97                    Either::Left((None, response_or_cancel)) => break response_or_cancel.await,
98                    Either::Right((response, _)) => break response,
99                }
100            };
101            let Some(response) = response else {
102                yield ToolCallEvent::Cancelled { task_id: None };
103                return;
104            };
105
106            match response {
107                Ok(CallToolResponse::Complete(result)) => {
108                    yield ToolCallEvent::Complete(Ok(result));
109                    return;
110                }
111                Ok(CallToolResponse::InputRequired(input_required)) => {
112                    match mrtr_state.tick(input_required) {
113                        MrtrAction::Poll { backoff, request_state } => {
114                            let backoff = pin!(sleep(backoff));
115                            if let Either::Right(((), _)) = select(backoff, pin!(options.cancel.cancelled())).await {
116                                yield ToolCallEvent::Cancelled { task_id: None };
117                                return;
118                            }
119                            params.input_responses = None;
120                            params.request_state = Some(request_state);
121                        }
122                        MrtrAction::Elicit { input_requests, request_state } => {
123                            let elicit = pin!(elicit_input(client.service(), &mut mrtr_state, &server_name, input_requests));
124                            match select(elicit, pin!(options.cancel.cancelled())).await {
125                                Either::Left((Ok(responses), _)) => {
126                                    params.input_responses = Some(responses);
127                                    params.request_state = request_state;
128                                }
129                                Either::Left((Err(e), _)) => {
130                                    yield ToolCallEvent::Complete(Err(e));
131                                    return;
132                                }
133                                Either::Right(((), _)) => {
134                                    yield ToolCallEvent::Cancelled { task_id: None };
135                                    return;
136                                }
137                            }
138                        }
139                        MrtrAction::Abort(reason) => {
140                            yield ToolCallEvent::Complete(Err(CallToolError::Aborted {
141                                server: server_name,
142                                reason,
143                                timeout: options.timeout,
144                            }));
145                            return;
146                        }
147                    }
148                }
149                Ok(CallToolResponse::Task(task)) => {
150                    let driver = TaskDriver::new(&server_name, client.as_ref(), options.timeout, options.cancel.clone());
151                    let mut events = Box::pin(driver.stream(task, progress));
152                    while let Some(event) = events.next().await {
153                        yield event;
154                    }
155                    return;
156                }
157                Ok(_) => {
158                    yield ToolCallEvent::Complete(Err(CallToolError::UnsupportedResponse { server: server_name }));
159                    return;
160                }
161                Err(e) => {
162                    yield ToolCallEvent::Complete(Err(CallToolError::Call(e)));
163                    return;
164                }
165            }
166        }
167    }
168}
169
170fn peer_request_options(options: &CallToolOptions) -> PeerRequestOptions {
171    let request_options = PeerRequestOptions::with_timeout(options.timeout);
172    match &options.meta {
173        Some(meta) => request_options.with_meta(meta.clone()),
174        None => request_options,
175    }
176}
177
178async fn await_response_or_cancel(
179    handle: RequestHandle<RoleClient>,
180    cancel: &CancellationToken,
181) -> Option<Result<CallToolResponse, ServiceError>> {
182    let response = pin!(await_tool_response(handle));
183    match select(response, pin!(cancel.cancelled())).await {
184        Either::Left((response, _)) => Some(response),
185        Either::Right(((), _)) => None,
186    }
187}
188
189async fn await_tool_response(handle: RequestHandle<RoleClient>) -> Result<CallToolResponse, ServiceError> {
190    match handle.await_response().await? {
191        ServerResult::CallToolResult(result) => Ok(CallToolResponse::Complete(result)),
192        ServerResult::InputRequiredResult(result) => Ok(CallToolResponse::InputRequired(result)),
193        ServerResult::CreateTaskResult(result) => Ok(CallToolResponse::Task(result)),
194        _ => Err(ServiceError::UnexpectedResponse),
195    }
196}
197
198async fn elicit_input(
199    client: &McpClient,
200    mrtr_state: &mut MrtrState,
201    server_name: &str,
202    requests: InputRequests,
203) -> Result<InputResponses, CallToolError> {
204    let (responses, results) = elicit_inputs(client, requests).await.map_err(|error| match error {
205        ElicitInputsError::UnsupportedInput => CallToolError::UnsupportedInput { server: server_name.to_string() },
206        ElicitInputsError::Serialize(source) => {
207            CallToolError::SerializeResponse { server: server_name.to_string(), source }
208        }
209    })?;
210    for result in &results {
211        mrtr_state.record_response(result);
212    }
213    Ok(responses)
214}