1use 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#[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#[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#[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 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 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#[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 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 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}