Skip to main content

mcp_utils/client/
task.rs

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}