1use crate::client::McpClient;
2use crate::client::call_tool::{CallToolError, ToolCallEvent};
3use crate::client::elicitation::{ElicitInputsError, elicit_inputs};
4use async_stream::stream;
5use futures::future::{Either, select};
6use futures::{Stream, StreamExt, pin_mut};
7use rmcp::RoleClient;
8use rmcp::model::{
9 CancelTaskParams, CreateTaskResult, GetTaskParams, InputRequests, ProgressNotificationParam, Task, TaskPayload,
10 TaskStatus, UpdateTaskParams,
11};
12use rmcp::service::{RunningService, ServiceError};
13use std::collections::HashSet;
14use std::future::Future;
15use std::pin::pin;
16use std::time::Duration;
17use thiserror::Error;
18use tokio::time::error::Elapsed;
19use tokio::time::{Instant, sleep, timeout, timeout_at};
20use tokio_util::sync::CancellationToken;
21
22#[derive(Debug, Error)]
23pub enum TaskErrorReason {
24 #[error("failed to get task: {0}")]
25 Get(#[source] ServiceError),
26 #[error("failed to update task: {0}")]
27 Update(#[source] ServiceError),
28 #[error("expired before completion")]
29 Expired,
30 #[error("exceeded the {timeout:?} execution deadline")]
31 TimedOut { timeout: Duration },
32 #[error("repeated input requests that were already answered")]
33 RepeatedInput,
34 #[error("failed: {error}")]
35 Failed { error: serde_json::Value },
36 #[error("was cancelled")]
37 Cancelled,
38 #[error("returned a malformed result: {0}")]
39 MalformedResult(#[source] serde_json::Error),
40 #[error("requested an input kind this client does not support")]
41 UnsupportedInput,
42 #[error("produced an elicitation response that could not be serialized: {0}")]
43 Serialize(#[source] serde_json::Error),
44 #[error("returned a task payload this client does not support (status {status:?})")]
45 UnsupportedPayload { status: TaskStatus },
46}
47
48pub(crate) struct TaskDriver<'a> {
49 server_name: &'a str,
50 client: &'a RunningService<RoleClient, McpClient>,
51 timeout: Duration,
52 cancellation_token: CancellationToken,
53 default_poll_interval: Duration,
54}
55
56impl<'a> TaskDriver<'a> {
57 pub(crate) fn new(
58 server_name: &'a str,
59 client: &'a RunningService<RoleClient, McpClient>,
60 timeout: Duration,
61 cancellation_token: CancellationToken,
62 ) -> Self {
63 Self { client, server_name, timeout, cancellation_token, default_poll_interval: Duration::from_secs(1) }
64 }
65
66 pub(crate) fn stream<T: Stream<Item = ProgressNotificationParam> + Send + 'a>(
67 self,
68 created: CreateTaskResult,
69 progress_events: T,
70 ) -> impl Stream<Item = ToolCallEvent> + 'a {
71 stream! {
72 yield ToolCallEvent::TaskCreated(created.clone());
73 let task_events = self.stream_task_events(created.task);
74 pin_mut!(task_events);
75 pin_mut!(progress_events);
76
77 loop {
78 tokio::select! {
79 progress_event = progress_events.next() => {
80 let Some(progress_event) = progress_event else {
81 while let Some(event) = task_events.next().await {
82 yield event;
83 }
84 return;
85 };
86 yield ToolCallEvent::Progress(progress_event);
87 }
88 event = task_events.next() => {
89 let Some(event) = event else {
90 return;
91 };
92 yield event;
93 }
94 }
95 }
96 }
97 }
98
99 fn stream_task_events(self, mut task: Task) -> impl Stream<Item = ToolCallEvent> + 'a {
100 stream! {
101 let bounds = TaskBounds::new(self.timeout, self.cancellation_token.clone());
102 let mut answered_input_keys = HashSet::new();
103
104 loop {
105 if is_task_expired(&task) {
106 yield self.fail(task, TaskErrorReason::Expired);
107 return;
108 }
109
110 let detailed_task = match bounds
111 .run(self.client.get_task(GetTaskParams::new(task.task_id.clone())))
112 .await
113 {
114 Ok(Ok(result)) => result.task,
115 Ok(Err(source)) => {
116 yield self.cancel(task, TaskErrorReason::Get(source)).await;
117 return;
118 }
119 Err(interrupt) => {
120 yield self.interrupted(task, interrupt).await;
121 return;
122 }
123 };
124
125 task = detailed_task.task;
126 if !task.status.is_terminal() {
127 yield ToolCallEvent::TaskStatus(task.clone());
128 }
129
130 match detailed_task.payload {
131 TaskPayload::Working => {}
132 TaskPayload::InputRequired { input_requests } => {
133 match bounds.run(self.elicit_inputs(input_requests, &mut answered_input_keys, &task.task_id)).await {
134 Ok(Ok(())) => {}
135 Ok(Err(reason)) => {
136 yield self.cancel(task, reason).await;
137 return;
138 }
139 Err(interrupt) => {
140 yield self.interrupted(task, interrupt).await;
141 return;
142 }
143 }
144 }
145 TaskPayload::Completed { result } => {
146 let result = serde_json::from_value(serde_json::Value::Object(result))
147 .map_err(|source| self.error(&task.task_id, TaskErrorReason::MalformedResult(source)));
148 yield ToolCallEvent::TaskComplete { task, result };
149 return;
150 }
151 TaskPayload::Failed { error } => {
152 yield self.fail(task, TaskErrorReason::Failed { error: serde_json::Value::Object(error) });
153 return;
154 }
155 TaskPayload::Cancelled => {
156 yield self.fail(task, TaskErrorReason::Cancelled);
157 return;
158 }
159 _ => {
160 let status = task.status;
161 yield self.cancel(task, TaskErrorReason::UnsupportedPayload { status }).await;
162 return;
163 }
164 }
165
166 let duration = task.poll_interval_ms.map_or(self.default_poll_interval, Duration::from_millis);
167 if let Err(interrupt) = bounds.run(sleep(duration)).await {
168 yield self.interrupted(task, interrupt).await;
169 return;
170 }
171 }
172 }
173 }
174
175 async fn elicit_inputs(
176 &self,
177 input_requests: InputRequests,
178 answered_input_keys: &mut HashSet<String>,
179 task_id: &str,
180 ) -> Result<(), TaskErrorReason> {
181 if input_requests.keys().any(|key| answered_input_keys.contains(key)) {
182 return Err(TaskErrorReason::RepeatedInput);
183 }
184
185 let (responses, _) = elicit_inputs(self.client.service(), input_requests).await?;
186 answered_input_keys.extend(responses.keys().cloned());
187
188 self.client.update_task(UpdateTaskParams::new(task_id, responses)).await.map_err(TaskErrorReason::Update)
189 }
190
191 async fn interrupted(&self, task: Task, interrupt: InterruptedReason) -> ToolCallEvent {
192 match interrupt {
193 InterruptedReason::TimedOut => self.cancel(task, TaskErrorReason::TimedOut { timeout: self.timeout }).await,
194 InterruptedReason::Cancelled => {
195 cancel_server_task(self.client, self.server_name, &task.task_id).await;
196 ToolCallEvent::Cancelled { task_id: Some(task.task_id) }
197 }
198 }
199 }
200
201 async fn cancel(&self, task: Task, reason: TaskErrorReason) -> ToolCallEvent {
202 cancel_server_task(self.client, self.server_name, &task.task_id).await;
203 self.fail(task, reason)
204 }
205
206 fn fail(&self, task: Task, reason: TaskErrorReason) -> ToolCallEvent {
207 let error = self.error(&task.task_id, reason);
208 ToolCallEvent::TaskComplete { task, result: Err(error) }
209 }
210
211 fn error(&self, task_id: &str, reason: TaskErrorReason) -> CallToolError {
212 CallToolError::Task {
213 server: self.server_name.to_string(),
214 task_id: task_id.to_string(),
215 reason: Box::new(reason),
216 }
217 }
218}
219
220pub(crate) async fn cancel_server_task(
221 client: &RunningService<RoleClient, McpClient>,
222 server_name: &str,
223 task_id: &str,
224) {
225 match timeout(Duration::from_secs(1), client.cancel_task(CancelTaskParams::new(task_id))).await {
226 Ok(Ok(())) => {}
227 Ok(Err(error)) => {
228 tracing::warn!(server = %server_name, %task_id, "Failed to cancel abandoned MCP task: {error}");
229 }
230 Err(_) => tracing::warn!(server = %server_name, %task_id, "Timed out cancelling abandoned MCP task"),
231 }
232}
233
234impl From<ElicitInputsError> for TaskErrorReason {
235 fn from(error: ElicitInputsError) -> Self {
236 match error {
237 ElicitInputsError::UnsupportedInput => Self::UnsupportedInput,
238 ElicitInputsError::Serialize(source) => Self::Serialize(source),
239 }
240 }
241}
242
243struct TaskBounds {
244 deadline: TaskDeadline,
245 cancel: CancellationToken,
246}
247
248enum InterruptedReason {
249 TimedOut,
250 Cancelled,
251}
252
253impl TaskBounds {
254 fn new(timeout: Duration, cancel: CancellationToken) -> Self {
255 Self { deadline: TaskDeadline::after(timeout), cancel }
256 }
257
258 async fn run<T>(&self, future: impl Future<Output = T>) -> Result<T, InterruptedReason> {
259 let timedout = pin!(self.deadline.timeout(future));
260 let cancelled = pin!(self.cancel.cancelled());
261 match select(timedout, cancelled).await {
262 Either::Left((Ok(value), _)) => Ok(value),
263 Either::Left((Err(_), _)) => Err(InterruptedReason::TimedOut),
264 Either::Right(((), _)) => Err(InterruptedReason::Cancelled),
265 }
266 }
267}
268
269enum TaskDeadline {
270 At(Instant),
271 FarFuture,
272}
273
274impl TaskDeadline {
275 fn after(timeout: Duration) -> Self {
276 Instant::now().checked_add(timeout).map_or(Self::FarFuture, Self::At)
277 }
278
279 async fn timeout<T>(&self, future: impl Future<Output = T>) -> Result<T, Elapsed> {
280 match self {
281 Self::At(deadline) => timeout_at(*deadline, future).await,
282 Self::FarFuture => Ok(future.await),
283 }
284 }
285}
286
287fn is_task_expired(task: &Task) -> bool {
288 if task.status.is_terminal() {
289 return false;
290 }
291 let Some(ttl_ms) = task.ttl_ms else {
292 return false;
293 };
294 let Ok(created_at) = chrono::DateTime::parse_from_rfc3339(&task.created_at) else {
295 tracing::warn!(task_id = %task.task_id, created_at = %task.created_at, "Ignoring malformed MCP task creation timestamp");
296 return false;
297 };
298 let Ok(ttl_ms) = i64::try_from(ttl_ms) else {
299 return false;
300 };
301 created_at
302 .with_timezone(&chrono::Utc)
303 .checked_add_signed(chrono::Duration::milliseconds(ttl_ms))
304 .is_some_and(|expires_at| chrono::Utc::now() > expires_at)
305}
306
307#[cfg(test)]
308mod tests {
309 use super::*;
310 use crate::client::call_tool::{CallToolOptions, call_tool};
311 use crate::client::{McpClientEvent, client_capabilities};
312 use crate::testing::{FakeMcpServer, FakeMcpState, FakeTool, FakeToolResponse, connect};
313 use futures::StreamExt;
314 use rmcp::model::{
315 CallToolRequestParams, CallToolResult, ClientInfo, CreateTaskResult, DetailedTask, ElicitRequest,
316 ElicitRequestParams, Implementation, InputRequest, ProtocolVersion,
317 };
318 use serde_json::json;
319 use std::sync::Arc;
320 use tokio::sync::mpsc;
321
322 #[tokio::test]
323 async fn call_tool_drives_created_task_to_completion() {
324 let result = task_test([completed_task()]).run().await;
325
326 assert!(
327 matches!(result.events.first(), Some(ToolCallEvent::TaskCreated(created)) if created.task.task_id == "task-1")
328 );
329 assert!(matches!(
330 result.events.last(),
331 Some(ToolCallEvent::TaskComplete { task, result: Ok(result) })
332 if task.task_id == "task-1"
333 && result.content.first().and_then(|content| content.as_text()).is_some_and(|text| text.text == "finished")
334 ));
335 assert_eq!(result.state.task_get_ids(), ["task-1"]);
336 }
337
338 #[tokio::test]
339 async fn call_tool_forwards_progress_after_task_creation() {
340 let seed = task(TaskStatus::Working);
341 let server = FakeMcpServer::new()
342 .with_tool(
343 FakeTool::new("deferred")
344 .responds(FakeToolResponse::task(CreateTaskResult::new(seed)).task_progress(1.0, Some(2.0))),
345 )
346 .with_task(
347 "task-1",
348 [DetailedTask::new(task(TaskStatus::Working), TaskPayload::Working), completed_task()],
349 );
350 let (event_tx, _event_rx) = mpsc::channel::<McpClientEvent>(4);
351 let client = McpClient::new(
352 ClientInfo::new(client_capabilities(), Implementation::new("test-client", "0.1.0")),
353 "task-server".into(),
354 event_tx,
355 );
356 let (_server, client) = connect(server, client).await.expect("connect task server");
357
358 let events = call_tool(
359 Arc::new(client),
360 CallToolRequestParams::new("deferred"),
361 CallToolOptions { timeout: Duration::from_secs(1), ..CallToolOptions::default() },
362 )
363 .collect::<Vec<_>>()
364 .await;
365
366 assert!(matches!(events.first(), Some(ToolCallEvent::TaskCreated(_))));
367 assert!(events.iter().any(|event| matches!(
368 event,
369 ToolCallEvent::Progress(progress)
370 if (progress.progress - 1.0).abs() < f64::EPSILON
371 && progress.total.is_some_and(|total| (total - 2.0).abs() < f64::EPSILON)
372 )));
373 assert!(matches!(events.last(), Some(ToolCallEvent::TaskComplete { result: Ok(_), .. })));
374 }
375
376 #[tokio::test]
377 async fn call_tool_handles_huge_task_ttl() {
378 let result =
379 task_test([completed_task()]).with_task(task(TaskStatus::Working).with_ttl_ms(u64::MAX)).run().await;
380 assert!(matches!(result.events.last(), Some(ToolCallEvent::TaskComplete { result: Ok(_), .. })));
381 }
382
383 #[tokio::test]
384 async fn call_tool_handles_huge_execution_timeout() {
385 let result = task_test([completed_task()]).with_timeout(Duration::MAX).run().await;
386 assert!(matches!(result.events.last(), Some(ToolCallEvent::TaskComplete { result: Ok(_), .. })));
387 }
388
389 #[tokio::test]
390 async fn call_tool_cancellation_cancels_server_task_and_ends_stream() {
391 let seed = task(TaskStatus::Working);
392 let server = FakeMcpServer::new()
393 .with_tool(FakeTool::new("deferred").responds(FakeToolResponse::task(CreateTaskResult::new(seed.clone()))))
394 .with_task("task-1", [DetailedTask::new(seed, TaskPayload::Working)]);
395 let state = server.state();
396 let (event_tx, _event_rx) = mpsc::channel::<McpClientEvent>(4);
397 let client = McpClient::new(
398 ClientInfo::new(client_capabilities(), Implementation::new("test-client", "0.1.0")),
399 "task-server".into(),
400 event_tx,
401 );
402 let (_server, client) = connect(server, client).await.expect("connect task server");
403
404 let cancel = CancellationToken::new();
405 let options = CallToolOptions { timeout: Duration::from_secs(5), meta: None, cancel: cancel.clone() };
406 let mut events = pin!(call_tool(Arc::new(client), CallToolRequestParams::new("deferred"), options));
407
408 assert!(matches!(events.next().await, Some(ToolCallEvent::TaskCreated(_))));
409 cancel.cancel();
410 let mut last = None;
411 while let Some(event) = events.next().await {
412 last = Some(event);
413 }
414
415 assert!(matches!(last, Some(ToolCallEvent::Cancelled { task_id: Some(task_id) }) if task_id == "task-1"));
416 assert_eq!(state.task_cancel_ids(), ["task-1"]);
417 }
418
419 #[tokio::test]
420 async fn call_tool_deadline_includes_task_elicitation() {
421 let result = task_test([input_required_task()]).with_timeout(Duration::from_millis(25)).run().await;
422
423 assert!(matches!(
424 result.events.last(),
425 Some(ToolCallEvent::TaskComplete {
426 result: Err(CallToolError::Task { reason, .. }),
427 ..
428 }) if matches!(reason.as_ref(), TaskErrorReason::TimedOut { .. })
429 ));
430 assert_eq!(result.state.task_cancel_ids(), ["task-1"]);
431 }
432
433 struct TaskTest {
434 seed: Task,
435 states: Vec<DetailedTask>,
436 timeout: Duration,
437 }
438
439 struct TaskTestResult {
440 events: Vec<ToolCallEvent>,
441 state: FakeMcpState,
442 }
443
444 fn task_test(states: impl IntoIterator<Item = DetailedTask>) -> TaskTest {
445 TaskTest {
446 seed: task(TaskStatus::Working),
447 states: states.into_iter().collect(),
448 timeout: Duration::from_secs(1),
449 }
450 }
451
452 impl TaskTest {
453 fn with_task(mut self, seed: Task) -> Self {
454 self.seed = seed;
455 self
456 }
457
458 fn with_timeout(mut self, timeout: Duration) -> Self {
459 self.timeout = timeout;
460 self
461 }
462
463 async fn run(self) -> TaskTestResult {
464 let task_id = self.seed.task_id.clone();
465 let server = FakeMcpServer::new()
466 .with_tool(FakeTool::new("deferred").responds(FakeToolResponse::task(CreateTaskResult::new(self.seed))))
467 .with_task(task_id, self.states);
468 let state = server.state();
469 let (event_tx, _event_rx) = mpsc::channel::<McpClientEvent>(4);
470 let client = McpClient::new(
471 ClientInfo::new(client_capabilities(), Implementation::new("test-client", "0.1.0")),
472 "task-server".into(),
473 event_tx,
474 );
475 let (_server, client) = connect(server, client).await.expect("connect task server");
476 assert_eq!(client.peer_info().expect("peer info").protocol_version, ProtocolVersion::V_2026_07_28);
477
478 let events = call_tool(
479 Arc::new(client),
480 CallToolRequestParams::new("deferred"),
481 CallToolOptions { timeout: self.timeout, ..CallToolOptions::default() },
482 )
483 .collect()
484 .await;
485 TaskTestResult { events, state }
486 }
487 }
488
489 fn completed_task() -> DetailedTask {
490 let result = CallToolResult::success(vec![rmcp::model::ContentBlock::text("finished")]);
491 DetailedTask::new(
492 task(TaskStatus::Completed),
493 TaskPayload::Completed { result: serde_json::from_value(json!(result)).expect("serialize tool result") },
494 )
495 }
496
497 fn input_required_task() -> DetailedTask {
498 let request = ElicitRequest::new(ElicitRequestParams::FormElicitationParams {
499 meta: None,
500 message: "Provide input".to_string(),
501 requested_schema: serde_json::from_value(json!({
502 "type": "object",
503 "properties": {}
504 }))
505 .expect("valid elicitation schema"),
506 });
507 DetailedTask::new(
508 task(TaskStatus::InputRequired),
509 TaskPayload::InputRequired {
510 input_requests: InputRequests::from([("answer".to_string(), InputRequest::Elicitation(request))]),
511 },
512 )
513 }
514
515 fn task(status: TaskStatus) -> Task {
516 let now = chrono::Utc::now().to_rfc3339();
517 Task::new("task-1", status, now.clone(), now).with_poll_interval_ms(10)
518 }
519}