Skip to main content

prism_mcp_rs/server/
tasks.rs

1//! Production runtime for the MCP Tasks extension.
2
3use std::{collections::HashMap, sync::Arc, time::Duration};
4
5use async_trait::async_trait;
6use serde_json::Value;
7use tokio::sync::{broadcast, Mutex, Notify, RwLock};
8use tokio_util::sync::CancellationToken;
9
10use crate::{
11    core::error::{McpError, McpResult},
12    core::MultiRoundToolCall,
13    protocol::{error_codes, JsonRpcNotification, Task, TaskStatus, ToolResult, TASKS_STATUS},
14};
15
16/// Application handler for a tool that executes as a durable MCP task.
17#[async_trait]
18pub trait TaskToolHandler: Send + Sync {
19    async fn call(
20        &self,
21        arguments: HashMap<String, Value>,
22        context: TaskContext,
23    ) -> McpResult<ToolResult>;
24}
25
26#[async_trait]
27impl<F, Fut> TaskToolHandler for F
28where
29    F: Fn(HashMap<String, Value>, TaskContext) -> Fut + Send + Sync,
30    Fut: std::future::Future<Output = McpResult<ToolResult>> + Send,
31{
32    async fn call(
33        &self,
34        arguments: HashMap<String, Value>,
35        context: TaskContext,
36    ) -> McpResult<ToolResult> {
37        self(arguments, context).await
38    }
39}
40
41/// Application handler for the durable phase of a tool that first completes
42/// an MCP multi-round input exchange and then escalates to a Task.
43#[async_trait]
44pub trait ComposedTaskToolHandler: Send + Sync {
45    async fn call(&self, call: MultiRoundToolCall, context: TaskContext) -> McpResult<ToolResult>;
46}
47
48#[async_trait]
49impl<F, Fut> ComposedTaskToolHandler for F
50where
51    F: Fn(MultiRoundToolCall, TaskContext) -> Fut + Send + Sync,
52    Fut: std::future::Future<Output = McpResult<ToolResult>> + Send,
53{
54    async fn call(&self, call: MultiRoundToolCall, context: TaskContext) -> McpResult<ToolResult> {
55        self(call, context).await
56    }
57}
58
59struct TaskRecord {
60    owner: String,
61    task: RwLock<Task>,
62    responses: Mutex<HashMap<String, Value>>,
63    response_ready: Notify,
64    cancellation: CancellationToken,
65    expires_at: Option<tokio::time::Instant>,
66}
67
68/// Context supplied to a task tool. It supports cooperative cancellation and
69/// task-native multi-round input without leaking `requestState` into Tasks.
70#[derive(Clone)]
71pub struct TaskContext {
72    record: Arc<TaskRecord>,
73    status_sender: broadcast::Sender<JsonRpcNotification>,
74}
75
76impl TaskContext {
77    pub async fn task_id(&self) -> String {
78        self.record.task.read().await.task_id.clone()
79    }
80
81    pub fn is_cancelled(&self) -> bool {
82        self.record.cancellation.is_cancelled()
83    }
84
85    pub async fn cancelled(&self) {
86        self.record.cancellation.cancelled().await;
87    }
88
89    /// Update the human-readable working status.
90    pub async fn report_progress(&self, message: impl Into<String>) -> McpResult<()> {
91        {
92            let mut task = self.record.task.write().await;
93            if task.status.is_terminal() {
94                return Err(McpError::Cancelled("task is already terminal".to_string()));
95            }
96            task.status = TaskStatus::Working;
97            task.status_message = Some(message.into());
98            task.input_requests.clear();
99            task.last_updated_at = chrono::Utc::now().to_rfc3339();
100        }
101        self.publish().await
102    }
103
104    /// Publish input requests and wait until all of them are answered or the
105    /// task is cancelled. Keys remain unique for the task lifetime.
106    pub async fn require_input(
107        &self,
108        requests: HashMap<String, Value>,
109        message: Option<String>,
110    ) -> McpResult<HashMap<String, Value>> {
111        if requests.is_empty() {
112            return Err(McpError::InvalidParams(
113                "task inputRequests cannot be empty".to_string(),
114            ));
115        }
116        {
117            let mut task = self.record.task.write().await;
118            if task.status.is_terminal() {
119                return Err(McpError::Cancelled("task is already terminal".to_string()));
120            }
121            task.status = TaskStatus::InputRequired;
122            task.status_message = message;
123            task.input_requests = requests.clone();
124            task.last_updated_at = chrono::Utc::now().to_rfc3339();
125        }
126        self.publish().await?;
127
128        loop {
129            if self.is_cancelled() {
130                return Err(McpError::Cancelled(
131                    "task cancellation requested".to_string(),
132                ));
133            }
134            let mut responses = self.record.responses.lock().await;
135            if requests.keys().all(|key| responses.contains_key(key)) {
136                let values = requests
137                    .keys()
138                    .filter_map(|key| responses.remove(key).map(|value| (key.clone(), value)))
139                    .collect();
140                drop(responses);
141                {
142                    let mut task = self.record.task.write().await;
143                    task.status = TaskStatus::Working;
144                    task.input_requests.clear();
145                    task.last_updated_at = chrono::Utc::now().to_rfc3339();
146                }
147                self.publish().await?;
148                return Ok(values);
149            }
150            drop(responses);
151            tokio::select! {
152                _ = self.record.response_ready.notified() => {},
153                _ = self.record.cancellation.cancelled() => {
154                    return Err(McpError::Cancelled("task cancellation requested".to_string()));
155                }
156            }
157        }
158    }
159
160    async fn publish(&self) -> McpResult<()> {
161        let task = self.record.task.read().await;
162        publish_task(&self.status_sender, &task).await
163    }
164}
165
166/// In-memory task store with caller binding, TTL enforcement, status
167/// notifications, input delivery, and cooperative cancellation.
168#[derive(Clone)]
169pub(crate) struct TaskRegistry {
170    records: Arc<RwLock<HashMap<String, Arc<TaskRecord>>>>,
171    status_sender: broadcast::Sender<JsonRpcNotification>,
172    default_ttl: Option<Duration>,
173    poll_interval_ms: u64,
174}
175
176impl Default for TaskRegistry {
177    fn default() -> Self {
178        let (status_sender, _) = broadcast::channel(1024);
179        Self {
180            records: Arc::new(RwLock::new(HashMap::new())),
181            status_sender,
182            default_ttl: Some(Duration::from_secs(60 * 60)),
183            poll_interval_ms: 1_000,
184        }
185    }
186}
187
188impl TaskRegistry {
189    pub fn subscribe(&self) -> broadcast::Receiver<JsonRpcNotification> {
190        self.status_sender.subscribe()
191    }
192
193    pub async fn create(
194        &self,
195        owner: String,
196        arguments: HashMap<String, Value>,
197        handler: Arc<dyn TaskToolHandler>,
198    ) -> McpResult<Task> {
199        self.create_with(owner, move |context| async move {
200            handler.call(arguments, context).await
201        })
202        .await
203    }
204
205    pub async fn create_composed(
206        &self,
207        owner: String,
208        call: MultiRoundToolCall,
209        handler: Arc<dyn ComposedTaskToolHandler>,
210    ) -> McpResult<Task> {
211        self.create_with(owner, move |context| async move {
212            handler.call(call, context).await
213        })
214        .await
215    }
216
217    async fn create_with<F, Fut>(&self, owner: String, run: F) -> McpResult<Task>
218    where
219        F: FnOnce(TaskContext) -> Fut + Send + 'static,
220        Fut: std::future::Future<Output = McpResult<ToolResult>> + Send + 'static,
221    {
222        let now = chrono::Utc::now().to_rfc3339();
223        let task = Task {
224            task_id: uuid::Uuid::new_v4().to_string(),
225            status: TaskStatus::Working,
226            status_message: Some("Task accepted".to_string()),
227            created_at: now.clone(),
228            last_updated_at: now,
229            ttl_ms: self.default_ttl.map(|value| value.as_millis() as u64),
230            poll_interval_ms: Some(self.poll_interval_ms),
231            input_requests: HashMap::new(),
232            result: None,
233            error: None,
234        };
235        let record = Arc::new(TaskRecord {
236            owner,
237            task: RwLock::new(task.clone()),
238            responses: Mutex::new(HashMap::new()),
239            response_ready: Notify::new(),
240            cancellation: CancellationToken::new(),
241            expires_at: self
242                .default_ttl
243                .map(|ttl| tokio::time::Instant::now() + ttl),
244        });
245        // Strong creation consistency: insert before returning the handle.
246        self.records
247            .write()
248            .await
249            .insert(task.task_id.clone(), record.clone());
250
251        let sender = self.status_sender.clone();
252        tokio::spawn(async move {
253            let context = TaskContext {
254                record: record.clone(),
255                status_sender: sender.clone(),
256            };
257            let outcome = run(context).await;
258            let mut state = record.task.write().await;
259            if record.cancellation.is_cancelled() && !state.status.is_terminal() {
260                state.status = TaskStatus::Cancelled;
261                state.status_message = Some("Task cancellation accepted".to_string());
262                state.input_requests.clear();
263            } else {
264                match outcome {
265                    Ok(result) => {
266                        state.status = TaskStatus::Completed;
267                        state.status_message = Some("Task completed".to_string());
268                        state.input_requests.clear();
269                        state.result = serde_json::to_value(result).ok();
270                    }
271                    Err(error) => {
272                        state.status = TaskStatus::Failed;
273                        state.status_message = Some(error.to_string());
274                        state.input_requests.clear();
275                        state.error = Some(serde_json::json!({
276                            "code": error_codes::INTERNAL_ERROR,
277                            "message": error.to_string()
278                        }));
279                    }
280                }
281            }
282            state.last_updated_at = chrono::Utc::now().to_rfc3339();
283            let _ = publish_task(&sender, &state).await;
284        });
285        Ok(task)
286    }
287
288    async fn record(&self, task_id: &str, owner: &str) -> McpResult<Arc<TaskRecord>> {
289        let record = self
290            .records
291            .read()
292            .await
293            .get(task_id)
294            .cloned()
295            .ok_or_else(|| McpError::InvalidParams("Task not found".to_string()))?;
296        if record.owner != owner {
297            // Do not disclose whether a handle belongs to another caller.
298            return Err(McpError::InvalidParams("Task not found".to_string()));
299        }
300        if record
301            .expires_at
302            .is_some_and(|deadline| tokio::time::Instant::now() >= deadline)
303        {
304            self.records.write().await.remove(task_id);
305            return Err(McpError::InvalidParams("Task has expired".to_string()));
306        }
307        Ok(record)
308    }
309
310    pub async fn get(&self, task_id: &str, owner: &str) -> McpResult<Task> {
311        let record = self.record(task_id, owner).await?;
312        let task = record.task.read().await.clone();
313        task.validate().map_err(McpError::Internal)?;
314        Ok(task)
315    }
316
317    pub async fn update(
318        &self,
319        task_id: &str,
320        owner: &str,
321        responses: HashMap<String, Value>,
322    ) -> McpResult<()> {
323        let record = self.record(task_id, owner).await?;
324        let mut task = record.task.write().await;
325        let mut accepted = record.responses.lock().await;
326        for (key, value) in responses {
327            if task.input_requests.contains_key(&key) && !accepted.contains_key(&key) {
328                accepted.insert(key.clone(), value);
329                task.input_requests.remove(&key);
330            }
331        }
332        drop(accepted);
333        if task.status == TaskStatus::InputRequired && task.input_requests.is_empty() {
334            task.status = TaskStatus::Working;
335        }
336        task.last_updated_at = chrono::Utc::now().to_rfc3339();
337        let snapshot = task.clone();
338        drop(task);
339        publish_task(&self.status_sender, &snapshot).await?;
340        record.response_ready.notify_waiters();
341        Ok(())
342    }
343
344    pub async fn cancel(&self, task_id: &str, owner: &str) -> McpResult<()> {
345        let record = self.record(task_id, owner).await?;
346        record.cancellation.cancel();
347        Ok(())
348    }
349}
350
351async fn publish_task(
352    sender: &broadcast::Sender<JsonRpcNotification>,
353    task: &Task,
354) -> McpResult<()> {
355    task.validate().map_err(McpError::Internal)?;
356    let notification =
357        JsonRpcNotification::new(TASKS_STATUS.to_string(), Some(serde_json::to_value(task)?))?;
358    let _ = sender.send(notification);
359    Ok(())
360}