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}