1use hashbrown::{HashMap, HashSet};
7use std::sync::Arc;
8use tokio::sync::RwLock;
9
10use super::errors::{A2aError, A2aResult};
11use super::rpc::{ListTasksParams, ListTasksResult, TaskPushNotificationConfig};
12use super::types::{Artifact, Message, Task, TaskState, TaskStatus};
13
14#[derive(Debug, Clone)]
16pub struct TaskManager {
17 state: Arc<RwLock<TaskManagerState>>,
19 max_tasks: usize,
21}
22
23#[derive(Debug, Default)]
24struct TaskManagerState {
25 tasks: HashMap<String, Task>,
26 contexts: HashMap<String, Vec<String>>,
27 webhook_configs: HashMap<String, TaskPushNotificationConfig>,
28}
29
30impl Default for TaskManager {
31 fn default() -> Self {
32 Self::new()
33 }
34}
35
36impl TaskManager {
37 pub fn new() -> Self {
39 Self {
40 state: Arc::new(RwLock::new(TaskManagerState::default())),
41 max_tasks: 1000,
42 }
43 }
44
45 pub fn with_capacity(max_tasks: usize) -> Self {
47 Self {
48 state: Arc::new(RwLock::new(TaskManagerState {
49 tasks: HashMap::with_capacity(max_tasks.min(100)),
50 contexts: HashMap::new(),
51 webhook_configs: HashMap::new(),
52 })),
53 max_tasks,
54 }
55 }
56
57 pub async fn create_task(&self, context_id: Option<String>) -> Task {
59 let mut task = Task::new();
60 if let Some(ref ctx_id) = context_id {
61 task = task.with_context_id(ctx_id);
62 }
63
64 let task_id = task.id.clone();
65 let mut state = self.state.write().await;
66
67 if state.tasks.len() >= self.max_tasks {
68 self.evict_oldest_tasks(&mut state);
69 }
70
71 state.tasks.insert(task_id.clone(), task.clone());
72 if let Some(ctx_id) = context_id {
73 state.contexts.entry(ctx_id).or_default().push(task_id);
74 }
75
76 task
77 }
78
79 fn evict_oldest_tasks(&self, state: &mut TaskManagerState) {
81 let mut completed_tasks: Vec<_> = state
82 .tasks
83 .iter()
84 .filter(|(_, task)| task.is_terminal())
85 .map(|(id, task)| (id.clone(), task.status.timestamp))
86 .collect();
87
88 completed_tasks.sort_by(|a, b| a.1.cmp(&b.1));
89
90 let evict_count = (self.max_tasks / 10).max(1);
91 let evicted_ids: HashSet<_> = completed_tasks.into_iter().take(evict_count).map(|(id, _)| id).collect();
92
93 if evicted_ids.is_empty() {
94 return;
95 }
96
97 for id in &evicted_ids {
98 state.tasks.remove(id);
99 state.webhook_configs.remove(id);
100 }
101
102 state.contexts.retain(|_, task_ids| {
103 task_ids.retain(|task_id| !evicted_ids.contains(task_id));
104 !task_ids.is_empty()
105 });
106 }
107
108 pub async fn get_task(&self, task_id: &str) -> Option<Task> {
110 let state = self.state.read().await;
111 state.tasks.get(task_id).cloned()
112 }
113
114 pub async fn get_task_or_error(&self, task_id: &str) -> A2aResult<Task> {
116 self.get_task(task_id)
117 .await
118 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))
119 }
120
121 pub async fn update_status(&self, task_id: &str, state: TaskState, message: Option<Message>) -> A2aResult<Task> {
123 let mut manager_state = self.state.write().await;
124 let task = manager_state
125 .tasks
126 .get_mut(task_id)
127 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
128
129 task.status = match message {
130 Some(msg) => TaskStatus::with_message(state, msg),
131 None => TaskStatus::new(state),
132 };
133
134 Ok(task.clone())
135 }
136
137 pub async fn add_artifact(&self, task_id: &str, artifact: Artifact) -> A2aResult<Task> {
139 let mut state = self.state.write().await;
140 let task = state
141 .tasks
142 .get_mut(task_id)
143 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
144
145 task.artifacts.push(artifact);
146 Ok(task.clone())
147 }
148
149 pub async fn add_message(&self, task_id: &str, message: Message) -> A2aResult<Task> {
151 let mut state = self.state.write().await;
152 let task = state
153 .tasks
154 .get_mut(task_id)
155 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
156
157 task.history.push(message);
158 Ok(task.clone())
159 }
160
161 pub async fn cancel_task(&self, task_id: &str) -> A2aResult<Task> {
163 let mut state = self.state.write().await;
164 let task = state
165 .tasks
166 .get_mut(task_id)
167 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
168
169 if !task.is_cancelable() {
170 return Err(A2aError::TaskNotCancelable(format!(
171 "Task {} is in state {:?} and cannot be canceled",
172 task_id, task.status.state
173 )));
174 }
175
176 task.status = TaskStatus::new(TaskState::Canceled);
177 Ok(task.clone())
178 }
179
180 fn matches_list_filters(
181 task: &Task,
182 status: Option<&TaskState>,
183 updated_after: Option<&chrono::DateTime<chrono::Utc>>,
184 ) -> bool {
185 if let Some(status) = status
186 && &task.status.state != status
187 {
188 return false;
189 }
190
191 if let Some(updated_after) = updated_after
192 && task.status.timestamp < *updated_after
193 {
194 return false;
195 }
196
197 true
198 }
199
200 fn clone_task_for_listing(task: &Task, include_artifacts: bool, history_length: Option<usize>) -> Task {
201 let mut task = task.clone();
202
203 if !include_artifacts {
204 task.artifacts.clear();
205 }
206
207 if let Some(history_length) = history_length
208 && task.history.len() > history_length
209 {
210 let trim_count = task.history.len() - history_length;
211 task.history.drain(..trim_count);
212 }
213
214 task
215 }
216
217 pub async fn list_tasks(&self, params: ListTasksParams) -> ListTasksResult {
219 let updated_after = params
220 .last_updated_after
221 .as_deref()
222 .and_then(|after| chrono::DateTime::parse_from_rfc3339(after).ok())
223 .map(|after| after.to_utc());
224
225 let mut matching_tasks: Vec<(String, chrono::DateTime<chrono::Utc>)> = {
226 let state = self.state.read().await;
227 if let Some(context_id) = params.context_id.as_deref() {
228 state
229 .contexts
230 .get(context_id)
231 .into_iter()
232 .flat_map(|task_ids| task_ids.iter())
233 .filter_map(|task_id| {
234 let task = state.tasks.get(task_id)?;
235 Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
236 .then(|| (task_id.clone(), task.status.timestamp))
237 })
238 .collect()
239 } else {
240 state
241 .tasks
242 .iter()
243 .filter(|(_, task)| {
244 Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
245 })
246 .map(|(task_id, task)| (task_id.clone(), task.status.timestamp))
247 .collect()
248 }
249 };
250
251 matching_tasks.sort_by(|a, b| b.1.cmp(&a.1));
252
253 let total_size = matching_tasks.len() as u32;
254 let page_size = params.page_size.unwrap_or(50).min(100);
255 let start_idx = params
256 .page_token
257 .as_ref()
258 .and_then(|token| token.parse::<usize>().ok())
259 .unwrap_or(0);
260
261 let end_idx = (start_idx + page_size as usize).min(matching_tasks.len());
262 let next_page_token = if end_idx < matching_tasks.len() {
263 Some(end_idx.to_string())
264 } else {
265 None
266 };
267
268 let include_artifacts = params.include_artifacts == Some(true);
269 let history_length = params.history_length.map(|len| len as usize);
270 let page_task_ids: Vec<_> = matching_tasks.into_iter().skip(start_idx).take(page_size as usize).collect();
271 let result = if page_task_ids.is_empty() {
272 Vec::new()
273 } else {
274 let state = self.state.read().await;
275 page_task_ids
276 .into_iter()
277 .filter_map(|(task_id, _)| {
278 state
279 .tasks
280 .get(&task_id)
281 .map(|task| Self::clone_task_for_listing(task, include_artifacts, history_length))
282 })
283 .collect()
284 };
285
286 ListTasksResult {
287 tasks: result,
288 total_size: Some(total_size),
289 page_size: Some(page_size),
290 next_page_token,
291 }
292 }
293
294 pub async fn get_tasks_by_context(&self, context_id: &str) -> Vec<Task> {
296 let state = self.state.read().await;
297 state
298 .contexts
299 .get(context_id)
300 .map(|task_ids| task_ids.iter().filter_map(|id| state.tasks.get(id).cloned()).collect())
301 .unwrap_or_default()
302 }
303
304 pub async fn task_count(&self) -> usize {
306 self.state.read().await.tasks.len()
307 }
308
309 pub async fn clear(&self) {
311 let mut state = self.state.write().await;
312 state.tasks.clear();
313 state.contexts.clear();
314 state.webhook_configs.clear();
315 }
316
317 pub async fn set_webhook_config(&self, config: TaskPushNotificationConfig) -> A2aResult<()> {
319 if !config.url.starts_with("https://") && !config.url.starts_with("http://localhost") {
320 return Err(A2aError::UnsupportedOperation("Webhook URL must use HTTPS or be localhost".to_string()));
321 }
322
323 let mut state = self.state.write().await;
324 if !state.tasks.contains_key(&config.task_id) {
325 return Err(A2aError::TaskNotFound(config.task_id));
326 }
327
328 state.webhook_configs.insert(config.task_id.clone(), config);
329 Ok(())
330 }
331
332 pub async fn get_webhook_config(&self, task_id: &str) -> Option<TaskPushNotificationConfig> {
334 let state = self.state.read().await;
335 state.webhook_configs.get(task_id).cloned()
336 }
337
338 pub async fn remove_webhook_config(&self, task_id: &str) {
340 let mut state = self.state.write().await;
341 state.webhook_configs.remove(task_id);
342 }
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348 use crate::types::MessageRole;
349
350 #[tokio::test]
351 async fn test_create_task() {
352 let manager = TaskManager::new();
353 let task = manager.create_task(None).await;
354
355 assert!(!task.id.is_empty());
356 assert_eq!(task.state(), TaskState::Submitted);
357 assert_eq!(manager.task_count().await, 1);
358 }
359
360 #[tokio::test]
361 async fn test_create_task_with_context() {
362 let manager = TaskManager::new();
363 let task = manager.create_task(Some("ctx-1".to_string())).await;
364
365 assert_eq!(task.context_id, Some("ctx-1".to_string()));
366
367 let tasks = manager.get_tasks_by_context("ctx-1").await;
368 assert_eq!(tasks.len(), 1);
369 assert_eq!(tasks[0].id, task.id);
370 }
371
372 #[tokio::test]
373 async fn test_get_task() {
374 let manager = TaskManager::new();
375 let task = manager.create_task(None).await;
376
377 let retrieved = manager.get_task(&task.id).await;
378 assert!(retrieved.is_some());
379 assert_eq!(retrieved.unwrap().id, task.id);
380
381 let missing = manager.get_task("nonexistent").await;
382 assert!(missing.is_none());
383 }
384
385 #[tokio::test]
386 async fn test_update_status() {
387 let manager = TaskManager::new();
388 let task = manager.create_task(None).await;
389
390 let updated = manager.update_status(&task.id, TaskState::Working, None).await.expect("update");
391 assert_eq!(updated.state(), TaskState::Working);
392
393 let msg = Message::agent_text("Task completed successfully");
394 let completed = manager
395 .update_status(&task.id, TaskState::Completed, Some(msg))
396 .await
397 .expect("complete");
398 assert_eq!(completed.state(), TaskState::Completed);
399 assert!(completed.status.message.is_some());
400 }
401
402 #[tokio::test]
403 async fn test_add_artifact() {
404 let manager = TaskManager::new();
405 let task = manager.create_task(None).await;
406
407 let artifact = Artifact::text("art-1", "Generated content");
408 let updated = manager.add_artifact(&task.id, artifact).await.expect("add artifact");
409 assert_eq!(updated.artifacts.len(), 1);
410 assert_eq!(updated.artifacts[0].id, "art-1");
411 }
412
413 #[tokio::test]
414 async fn test_cancel_task() {
415 let manager = TaskManager::new();
416 let task = manager.create_task(None).await;
417
418 let canceled = manager.cancel_task(&task.id).await.expect("cancel");
419 assert_eq!(canceled.state(), TaskState::Canceled);
420 }
421
422 #[tokio::test]
423 async fn test_cancel_completed_task_fails() {
424 let manager = TaskManager::new();
425 let task = manager.create_task(None).await;
426
427 manager
428 .update_status(&task.id, TaskState::Completed, None)
429 .await
430 .expect("complete");
431
432 let result = manager.cancel_task(&task.id).await;
433 result.unwrap_err();
434 }
435
436 #[tokio::test]
437 async fn test_eviction_cleans_context_and_webhook_indexes() {
438 let manager = TaskManager::with_capacity(1);
439 let task = manager.create_task(Some("ctx-1".to_string())).await;
440
441 manager
442 .update_status(&task.id, TaskState::Completed, None)
443 .await
444 .expect("complete");
445 manager
446 .set_webhook_config(TaskPushNotificationConfig {
447 task_id: task.id.clone(),
448 url: "https://example.com/webhook".to_string(),
449 authentication: None,
450 })
451 .await
452 .expect("set webhook");
453
454 let replacement = manager.create_task(None).await;
455
456 assert_eq!(manager.task_count().await, 1);
457 assert!(manager.get_task(&task.id).await.is_none());
458 assert!(manager.get_webhook_config(&task.id).await.is_none());
459 assert!(manager.get_tasks_by_context("ctx-1").await.is_empty());
460 assert_eq!(manager.get_task(&replacement.id).await.unwrap().id, replacement.id);
461 }
462
463 #[tokio::test]
464 async fn test_list_tasks() {
465 let manager = TaskManager::new();
466
467 let _task1 = manager.create_task(Some("ctx-1".to_string())).await;
468 let _task2 = manager.create_task(Some("ctx-1".to_string())).await;
469 let _task3 = manager.create_task(Some("ctx-2".to_string())).await;
470
471 let all = manager.list_tasks(ListTasksParams::default()).await;
472 assert_eq!(all.tasks.len(), 3);
473
474 let ctx1_tasks = manager
475 .list_tasks(ListTasksParams {
476 context_id: Some("ctx-1".to_string()),
477 ..Default::default()
478 })
479 .await;
480 assert_eq!(ctx1_tasks.tasks.len(), 2);
481 }
482
483 #[tokio::test]
484 async fn test_list_tasks_paginates_and_trims_after_sorting() {
485 let manager = TaskManager::new();
486
487 let older = manager.create_task(Some("ctx-1".to_string())).await;
488 tokio::time::sleep(std::time::Duration::from_millis(2)).await;
489 let newer = manager.create_task(Some("ctx-1".to_string())).await;
490
491 manager
492 .add_artifact(&newer.id, Artifact::text("art-1", "Generated content"))
493 .await
494 .expect("add artifact");
495 manager
496 .add_message(&newer.id, Message::user_text("Hello"))
497 .await
498 .expect("add msg1");
499 manager
500 .add_message(&newer.id, Message::agent_text("Hi there"))
501 .await
502 .expect("add msg2");
503
504 let first_page = manager
505 .list_tasks(ListTasksParams {
506 context_id: Some("ctx-1".to_string()),
507 page_size: Some(1),
508 history_length: Some(1),
509 include_artifacts: Some(false),
510 ..Default::default()
511 })
512 .await;
513
514 assert_eq!(first_page.total_size, Some(2));
515 assert_eq!(first_page.next_page_token.as_deref(), Some("1"));
516 assert_eq!(first_page.tasks.len(), 1);
517 assert_eq!(first_page.tasks[0].id, newer.id);
518 assert!(first_page.tasks[0].artifacts.is_empty());
519 assert_eq!(first_page.tasks[0].history.len(), 1);
520 assert_eq!(first_page.tasks[0].history[0].role, MessageRole::Agent);
521
522 let second_page = manager
523 .list_tasks(ListTasksParams {
524 context_id: Some("ctx-1".to_string()),
525 page_size: Some(1),
526 page_token: Some("1".to_string()),
527 ..Default::default()
528 })
529 .await;
530
531 assert_eq!(second_page.tasks.len(), 1);
532 assert_eq!(second_page.tasks[0].id, older.id);
533 assert!(second_page.next_page_token.is_none());
534 }
535
536 #[tokio::test]
537 async fn test_add_message_to_history() {
538 let manager = TaskManager::new();
539 let task = manager.create_task(None).await;
540
541 let msg1 = Message::user_text("Hello");
542 let msg2 = Message::agent_text("Hi there!");
543
544 manager.add_message(&task.id, msg1).await.expect("add msg1");
545 let updated = manager.add_message(&task.id, msg2).await.expect("add msg2");
546
547 assert_eq!(updated.history.len(), 2);
548 assert_eq!(updated.history[0].role, MessageRole::User);
549 assert_eq!(updated.history[1].role, MessageRole::Agent);
550 }
551}