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};
13use super::webhook::parse_webhook_url;
14
15#[derive(Debug, Clone)]
17pub struct TaskManager {
18 state: Arc<RwLock<TaskManagerState>>,
20 max_tasks: usize,
22}
23
24#[derive(Debug, Default)]
25struct TaskManagerState {
26 tasks: HashMap<String, Task>,
27 contexts: HashMap<String, Vec<String>>,
28 webhook_configs: HashMap<String, TaskPushNotificationConfig>,
29}
30
31impl Default for TaskManager {
32 fn default() -> Self {
33 Self::new()
34 }
35}
36
37impl TaskManager {
38 pub fn new() -> Self {
40 Self {
41 state: Arc::new(RwLock::new(TaskManagerState::default())),
42 max_tasks: 1000,
43 }
44 }
45
46 fn with_capacity(max_tasks: usize) -> Self {
48 Self {
49 state: Arc::new(RwLock::new(TaskManagerState {
50 tasks: HashMap::with_capacity(max_tasks.min(100)),
51 contexts: HashMap::new(),
52 webhook_configs: HashMap::new(),
53 })),
54 max_tasks,
55 }
56 }
57
58 pub(crate) async fn create_task(&self, context_id: Option<String>) -> Task {
60 let mut task = Task::new();
61 if let Some(ref ctx_id) = context_id {
62 task = task.with_context_id(ctx_id);
63 }
64
65 let task_id = task.id.clone();
66 let mut state = self.state.write().await;
67
68 if state.tasks.len() >= self.max_tasks {
69 self.evict_oldest_tasks(&mut state);
70 }
71
72 drop(state.tasks.insert(task_id.clone(), task.clone()));
73 if let Some(ctx_id) = context_id {
74 state.contexts.entry(ctx_id).or_default().push(task_id);
75 }
76
77 task
78 }
79
80 fn evict_oldest_tasks(&self, state: &mut TaskManagerState) {
82 let mut completed_tasks: Vec<_> = state
83 .tasks
84 .iter()
85 .filter(|(_, task)| task.is_terminal())
86 .map(|(id, task)| (id.clone(), task.status.timestamp))
87 .collect();
88
89 completed_tasks.sort_by_key(|a| a.1);
90
91 let evict_count = (self.max_tasks / 10).max(1);
92 let evicted_ids: HashSet<_> = completed_tasks.into_iter().take(evict_count).map(|(id, _)| id).collect();
93
94 if evicted_ids.is_empty() {
95 return;
96 }
97
98 for id in &evicted_ids {
99 drop(state.tasks.remove(id));
100 drop(state.webhook_configs.remove(id));
101 }
102
103 state.contexts.retain(|_, task_ids| {
104 task_ids.retain(|task_id| !evicted_ids.contains(task_id));
105 !task_ids.is_empty()
106 });
107 }
108
109 async fn get_task(&self, task_id: &str) -> Option<Task> {
111 let state = self.state.read().await;
112 state.tasks.get(task_id).cloned()
113 }
114
115 pub(crate) async fn get_task_or_error(&self, task_id: &str) -> A2aResult<Task> {
117 self.get_task(task_id)
118 .await
119 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))
120 }
121
122 pub(crate) async fn get_task_or_error_with_history(&self, task_id: &str, history_length: usize) -> A2aResult<Task> {
124 let state = self.state.read().await;
125 state
126 .tasks
127 .get(task_id)
128 .map(|task| task.clone_for_query(history_length, true))
129 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))
130 }
131
132 pub(crate) async fn update_status(
134 &self,
135 task_id: &str,
136 state: TaskState,
137 message: Option<Message>,
138 ) -> A2aResult<Task> {
139 let mut manager_state = self.state.write().await;
140 let task = manager_state
141 .tasks
142 .get_mut(task_id)
143 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
144
145 task.status = match message {
146 Some(msg) => TaskStatus::with_message(state, msg),
147 None => TaskStatus::new(state),
148 };
149
150 Ok(task.clone())
151 }
152
153 async fn add_artifact(&self, task_id: &str, artifact: Artifact) -> A2aResult<Task> {
155 let mut state = self.state.write().await;
156 let task = state
157 .tasks
158 .get_mut(task_id)
159 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
160
161 task.artifacts.push(artifact);
162 Ok(task.clone())
163 }
164
165 pub(crate) async fn add_message(&self, task_id: &str, message: Message) -> A2aResult<Task> {
167 let mut state = self.state.write().await;
168 let task = state
169 .tasks
170 .get_mut(task_id)
171 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
172
173 task.history.push(message);
174 Ok(task.clone())
175 }
176
177 pub(crate) async fn cancel_task(&self, task_id: &str) -> A2aResult<Task> {
179 let mut state = self.state.write().await;
180 let task = state
181 .tasks
182 .get_mut(task_id)
183 .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
184
185 if !task.is_cancelable() {
186 return Err(A2aError::TaskNotCancelable(format!(
187 "Task {} is in state {:?} and cannot be canceled",
188 task_id, task.status.state
189 )));
190 }
191
192 task.status = TaskStatus::new(TaskState::Canceled);
193 Ok(task.clone())
194 }
195
196 fn matches_list_filters(
197 task: &Task,
198 status: Option<&TaskState>,
199 updated_after: Option<&chrono::DateTime<chrono::Utc>>,
200 ) -> bool {
201 if let Some(status) = status
202 && &task.status.state != status
203 {
204 return false;
205 }
206
207 if let Some(updated_after) = updated_after
208 && task.status.timestamp < *updated_after
209 {
210 return false;
211 }
212
213 true
214 }
215
216 pub(crate) async fn list_tasks(&self, params: ListTasksParams) -> ListTasksResult {
218 let updated_after = params
219 .last_updated_after
220 .as_deref()
221 .and_then(|after| chrono::DateTime::parse_from_rfc3339(after).ok())
222 .map(|after| after.to_utc());
223
224 let mut matching_tasks: Vec<(String, chrono::DateTime<chrono::Utc>)> = {
225 let state = self.state.read().await;
226 if let Some(context_id) = params.context_id.as_deref() {
227 state
228 .contexts
229 .get(context_id)
230 .into_iter()
231 .flat_map(|task_ids| task_ids.iter())
232 .filter_map(|task_id| {
233 let task = state.tasks.get(task_id)?;
234 Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
235 .then(|| (task_id.clone(), task.status.timestamp))
236 })
237 .collect()
238 } else {
239 state
240 .tasks
241 .iter()
242 .filter(|(_, task)| {
243 Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
244 })
245 .map(|(task_id, task)| (task_id.clone(), task.status.timestamp))
246 .collect()
247 }
248 };
249
250 matching_tasks.sort_by_key(|a| std::cmp::Reverse(a.1));
251
252 let total_size = u32::try_from(matching_tasks.len()).unwrap_or(u32::MAX);
253 let page_size = params.page_size.unwrap_or(50).min(100);
254 let start_idx = params
255 .page_token
256 .as_ref()
257 .and_then(|token| token.parse::<usize>().ok())
258 .unwrap_or(0);
259
260 let end_idx = (start_idx + page_size as usize).min(matching_tasks.len());
261 let next_page_token = if end_idx < matching_tasks.len() {
262 Some(end_idx.to_string())
263 } else {
264 None
265 };
266
267 let include_artifacts = params.include_artifacts == Some(true);
268 let history_length = params.history_length.map(|len| len as usize);
269 let page_task_ids: Vec<_> = matching_tasks.into_iter().skip(start_idx).take(page_size as usize).collect();
270 let result = if page_task_ids.is_empty() {
271 Vec::new()
272 } else {
273 let state = self.state.read().await;
274 page_task_ids
275 .into_iter()
276 .filter_map(|(task_id, _)| {
277 state
278 .tasks
279 .get(&task_id)
280 .map(|task| task.clone_for_query(history_length.unwrap_or(0), include_artifacts))
281 })
282 .collect()
283 };
284
285 ListTasksResult {
286 tasks: result,
287 total_size: Some(total_size),
288 page_size: Some(page_size),
289 next_page_token,
290 }
291 }
292
293 async fn get_tasks_by_context(&self, context_id: &str) -> Vec<Task> {
295 let state = self.state.read().await;
296 state
297 .contexts
298 .get(context_id)
299 .map(|task_ids| task_ids.iter().filter_map(|id| state.tasks.get(id).cloned()).collect())
300 .unwrap_or_default()
301 }
302
303 async fn task_count(&self) -> usize {
305 self.state.read().await.tasks.len()
306 }
307
308 pub async fn clear(&self) {
310 let mut state = self.state.write().await;
311 state.tasks.clear();
312 state.contexts.clear();
313 state.webhook_configs.clear();
314 }
315
316 pub(crate) async fn set_webhook_config(&self, config: TaskPushNotificationConfig) -> A2aResult<()> {
318 drop(parse_webhook_url(&config.url).map_err(A2aError::UnsupportedOperation)?);
319
320 let mut state = self.state.write().await;
321 if !state.tasks.contains_key(&config.task_id) {
322 return Err(A2aError::TaskNotFound(config.task_id));
323 }
324
325 drop(state.webhook_configs.insert(config.task_id.clone(), config));
326 Ok(())
327 }
328
329 pub(crate) async fn get_webhook_config(&self, task_id: &str) -> Option<TaskPushNotificationConfig> {
331 let state = self.state.read().await;
332 state.webhook_configs.get(task_id).cloned()
333 }
334
335 pub async fn remove_webhook_config(&self, task_id: &str) {
337 let mut state = self.state.write().await;
338 drop(state.webhook_configs.remove(task_id));
339 }
340}
341
342#[cfg(test)]
343mod tests {
344 use super::*;
345 use crate::types::MessageRole;
346
347 #[tokio::test]
348 async fn test_create_task() {
349 let manager = TaskManager::new();
350 let task = manager.create_task(None).await;
351
352 assert!(!task.id.is_empty());
353 assert_eq!(task.state(), TaskState::Submitted);
354 assert_eq!(manager.task_count().await, 1);
355 }
356
357 #[tokio::test]
358 async fn test_create_task_with_context() {
359 let manager = TaskManager::new();
360 let task = manager.create_task(Some("ctx-1".to_string())).await;
361
362 assert_eq!(task.context_id, Some("ctx-1".to_string()));
363
364 let tasks = manager.get_tasks_by_context("ctx-1").await;
365 assert_eq!(tasks.len(), 1);
366 assert_eq!(tasks[0].id, task.id);
367 }
368
369 #[tokio::test]
370 async fn test_get_task() {
371 let manager = TaskManager::new();
372 let task = manager.create_task(None).await;
373
374 let retrieved = manager.get_task(&task.id).await;
375 assert!(retrieved.is_some());
376 assert_eq!(retrieved.unwrap().id, task.id);
377
378 let missing = manager.get_task("nonexistent").await;
379 assert!(missing.is_none());
380 }
381
382 #[tokio::test]
383 async fn bounded_queries_preserve_task_fields_suffix_and_source() {
384 let manager = TaskManager::new();
385 let mut task = Task::with_id("projection-task");
386 task.context_id = Some("projection-context".to_string());
387 task.status = TaskStatus::with_message(TaskState::Working, Message::agent_text("status payload"));
388 task.history = vec![
389 Message::user_text("first"),
390 Message::agent_text("middle payload"),
391 Message::user_text("last"),
392 ];
393 task.artifacts = vec![
394 Artifact::text("artifact-a", "small"),
395 Artifact::text("artifact-b", "longer output"),
396 ];
397 let mut original = serde_json::to_value(&task).unwrap();
398 original["metadata"] = serde_json::json!({"nested": {"sequence": [2, 7, 3]}, "label": "preserved"});
399 original["kind"] = serde_json::json!("custom-task");
400 let task: Task = serde_json::from_value(original.clone()).unwrap();
401 {
402 let mut state = manager.state.write().await;
403 drop(state.tasks.insert(task.id.clone(), task));
404 drop(
405 state
406 .contexts
407 .insert("projection-context".into(), vec!["projection-task".into()]),
408 );
409 }
410 let first = original["history"][0].clone();
411 let middle = original["history"][1].clone();
412 let last = original["history"][2].clone();
413 for (limit, expected_history) in [
414 (0, vec![]),
415 (1, vec![last.clone()]),
416 (2, vec![middle.clone(), last.clone()]),
417 (3, vec![first.clone(), middle.clone(), last.clone()]),
418 (usize::MAX, vec![first, middle, last]),
419 ] {
420 let mut expected = original.clone();
421 if expected_history.is_empty() {
422 drop(expected.as_object_mut().unwrap().remove("history"));
423 } else {
424 expected["history"] = serde_json::json!(expected_history);
425 }
426 let queried = manager.get_task_or_error_with_history("projection-task", limit).await.unwrap();
427 assert_eq!(serde_json::to_value(queried).unwrap(), expected, "get limit {limit}");
428 for include_artifacts in [false, true] {
429 let mut listed_expected = expected.clone();
430 if !include_artifacts {
431 drop(listed_expected.as_object_mut().unwrap().remove("artifacts"));
432 }
433 let params = ListTasksParams {
434 context_id: Some("projection-context".into()),
435 history_length: Some(u32::try_from(limit).unwrap_or(u32::MAX)),
436 include_artifacts: Some(include_artifacts),
437 ..Default::default()
438 };
439 let listed = manager.list_tasks(params).await;
440 assert_eq!(listed.total_size, Some(1));
441 assert_eq!(listed.tasks.len(), 1);
442 assert_eq!(
443 serde_json::to_value(&listed.tasks[0]).unwrap(),
444 listed_expected,
445 "list limit {limit}, artifacts {include_artifacts}"
446 );
447 }
448 }
449 assert_eq!(
450 serde_json::to_value(manager.get_task_or_error("projection-task").await.unwrap()).unwrap(),
451 original
452 );
453 assert!(matches!(manager.get_task_or_error_with_history("missing-task", 1).await,
454 Err(A2aError::TaskNotFound(id)) if id == "missing-task"));
455 }
456
457 #[tokio::test]
458 async fn test_update_status() {
459 let manager = TaskManager::new();
460 let task = manager.create_task(None).await;
461
462 let updated = manager.update_status(&task.id, TaskState::Working, None).await.expect("update");
463 assert_eq!(updated.state(), TaskState::Working);
464
465 let msg = Message::agent_text("Task completed successfully");
466 let completed = manager
467 .update_status(&task.id, TaskState::Completed, Some(msg))
468 .await
469 .expect("complete");
470 assert_eq!(completed.state(), TaskState::Completed);
471 assert!(completed.status.message.is_some());
472 }
473
474 #[tokio::test]
475 async fn test_add_artifact() {
476 let manager = TaskManager::new();
477 let task = manager.create_task(None).await;
478
479 let artifact = Artifact::text("art-1", "Generated content");
480 let updated = manager.add_artifact(&task.id, artifact).await.expect("add artifact");
481 assert_eq!(updated.artifacts.len(), 1);
482 assert_eq!(updated.artifacts[0].id, "art-1");
483 }
484
485 #[tokio::test]
486 async fn test_cancel_task() {
487 let manager = TaskManager::new();
488 let task = manager.create_task(None).await;
489
490 let canceled = manager.cancel_task(&task.id).await.expect("cancel");
491 assert_eq!(canceled.state(), TaskState::Canceled);
492 }
493
494 #[tokio::test]
495 async fn test_cancel_completed_task_fails() {
496 let manager = TaskManager::new();
497 let task = manager.create_task(None).await;
498
499 drop(
500 manager
501 .update_status(&task.id, TaskState::Completed, None)
502 .await
503 .expect("complete"),
504 );
505
506 let result = manager.cancel_task(&task.id).await;
507 drop(result.unwrap_err());
508 }
509
510 #[tokio::test]
511 async fn test_eviction_cleans_context_and_webhook_indexes() {
512 let manager = TaskManager::with_capacity(1);
513 let task = manager.create_task(Some("ctx-1".to_string())).await;
514
515 drop(
516 manager
517 .update_status(&task.id, TaskState::Completed, None)
518 .await
519 .expect("complete"),
520 );
521 manager
522 .set_webhook_config(TaskPushNotificationConfig {
523 task_id: task.id.clone(),
524 url: "https://example.com/webhook".to_string(),
525 authentication: None,
526 })
527 .await
528 .expect("set webhook");
529
530 let replacement = manager.create_task(None).await;
531
532 assert_eq!(manager.task_count().await, 1);
533 assert!(manager.get_task(&task.id).await.is_none());
534 assert!(manager.get_webhook_config(&task.id).await.is_none());
535 assert!(manager.get_tasks_by_context("ctx-1").await.is_empty());
536 assert_eq!(manager.get_task(&replacement.id).await.unwrap().id, replacement.id);
537 }
538
539 #[tokio::test]
540 async fn test_list_tasks() {
541 let manager = TaskManager::new();
542
543 let task1 = manager.create_task(Some("ctx-1".to_string())).await;
544 let _task2 = manager.create_task(Some("ctx-1".to_string())).await;
545 let _task3 = manager.create_task(Some("ctx-2".to_string())).await;
546 drop(
547 manager
548 .add_message(&task1.id, Message::user_text("private message"))
549 .await
550 .expect("add message"),
551 );
552
553 let all = manager.list_tasks(ListTasksParams::default()).await;
554 assert_eq!(all.tasks.len(), 3);
555 let listed_task = all.tasks.iter().find(|task| task.id == task1.id).expect("listed task");
556 assert!(listed_task.history.is_empty());
557
558 let ctx1_tasks = manager
559 .list_tasks(ListTasksParams {
560 context_id: Some("ctx-1".to_string()),
561 ..Default::default()
562 })
563 .await;
564 assert_eq!(ctx1_tasks.tasks.len(), 2);
565 }
566
567 #[tokio::test]
568 async fn test_list_tasks_paginates_and_trims_after_sorting() {
569 let manager = TaskManager::new();
570
571 let older = manager.create_task(Some("ctx-1".to_string())).await;
572 tokio::time::sleep(std::time::Duration::from_millis(2)).await;
573 let newer = manager.create_task(Some("ctx-1".to_string())).await;
574
575 drop(
576 manager
577 .add_artifact(&newer.id, Artifact::text("art-1", "Generated content"))
578 .await
579 .expect("add artifact"),
580 );
581 drop(
582 manager
583 .add_message(&newer.id, Message::user_text("Hello"))
584 .await
585 .expect("add msg1"),
586 );
587 drop(
588 manager
589 .add_message(&newer.id, Message::agent_text("Hi there"))
590 .await
591 .expect("add msg2"),
592 );
593
594 let first_page = manager
595 .list_tasks(ListTasksParams {
596 context_id: Some("ctx-1".to_string()),
597 page_size: Some(1),
598 history_length: Some(1),
599 include_artifacts: Some(false),
600 ..Default::default()
601 })
602 .await;
603
604 assert_eq!(first_page.total_size, Some(2));
605 assert_eq!(first_page.next_page_token.as_deref(), Some("1"));
606 assert_eq!(first_page.tasks.len(), 1);
607 assert_eq!(first_page.tasks[0].id, newer.id);
608 assert!(first_page.tasks[0].artifacts.is_empty());
609 assert_eq!(first_page.tasks[0].history.len(), 1);
610 assert_eq!(first_page.tasks[0].history[0].role, MessageRole::Agent);
611
612 let second_page = manager
613 .list_tasks(ListTasksParams {
614 context_id: Some("ctx-1".to_string()),
615 page_size: Some(1),
616 page_token: Some("1".to_string()),
617 ..Default::default()
618 })
619 .await;
620
621 assert_eq!(second_page.tasks.len(), 1);
622 assert_eq!(second_page.tasks[0].id, older.id);
623 assert!(second_page.next_page_token.is_none());
624 }
625
626 #[tokio::test]
627 async fn test_add_message_to_history() {
628 let manager = TaskManager::new();
629 let task = manager.create_task(None).await;
630
631 let msg1 = Message::user_text("Hello");
632 let msg2 = Message::agent_text("Hi there!");
633
634 drop(manager.add_message(&task.id, msg1).await.expect("add msg1"));
635 let updated = manager.add_message(&task.id, msg2).await.expect("add msg2");
636
637 assert_eq!(updated.history.len(), 2);
638 assert_eq!(updated.history[0].role, MessageRole::User);
639 assert_eq!(updated.history[1].role, MessageRole::Agent);
640 }
641
642 #[tokio::test]
643 async fn test_get_task_history_length_is_enforced() {
644 let manager = TaskManager::new();
645 let task = manager.create_task(None).await;
646 drop(
647 manager
648 .add_message(&task.id, Message::user_text("first"))
649 .await
650 .expect("add first message"),
651 );
652 drop(
653 manager
654 .add_message(&task.id, Message::agent_text("second"))
655 .await
656 .expect("add second message"),
657 );
658
659 let without_history = manager
660 .get_task_or_error_with_history(&task.id, 0)
661 .await
662 .expect("get task without history");
663 assert!(without_history.history.is_empty());
664
665 let last_message = manager
666 .get_task_or_error_with_history(&task.id, 1)
667 .await
668 .expect("get task with one history item");
669 assert_eq!(last_message.history.len(), 1);
670 assert_eq!(last_message.history[0].role, MessageRole::Agent);
671 }
672
673 #[tokio::test]
674 async fn test_webhook_url_validation_requires_exact_localhost_for_http() {
675 let manager = TaskManager::new();
676 let task = manager.create_task(None).await;
677
678 let invalid_urls = [
679 "http://localhost.evil.example/hook",
680 "http://localhost@evil.example/hook",
681 "http://example.com/hook",
682 "ftp://example.com/hook",
683 "https://user:password@example.com/hook",
684 ];
685
686 for url in invalid_urls {
687 let result = manager
688 .set_webhook_config(TaskPushNotificationConfig {
689 task_id: task.id.clone(),
690 url: url.to_string(),
691 authentication: None,
692 })
693 .await;
694 assert!(result.is_err(), "URL should be rejected: {url}");
695 }
696
697 for url in [
698 "https://example.com/hook",
699 "http://localhost:8080/hook",
700 "http://127.0.0.1:8080/hook",
701 "http://[::1]:8080/hook",
702 ] {
703 manager
704 .set_webhook_config(TaskPushNotificationConfig {
705 task_id: task.id.clone(),
706 url: url.to_string(),
707 authentication: None,
708 })
709 .await
710 .expect("valid webhook URL");
711 }
712 }
713}