Skip to main content

vtcode_a2a/
task_manager.rs

1//! A2A Task Manager
2//!
3//! Manages task lifecycle, storage, and queries for the A2A protocol.
4//! Provides an in-memory store with support for concurrent access.
5
6use 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/// A2A Task Manager - handles task creation, updates, and queries
16#[derive(Debug, Clone)]
17pub struct TaskManager {
18    /// All mutable task manager state lives behind one lock so related indexes stay in sync.
19    state: Arc<RwLock<TaskManagerState>>,
20    /// Maximum tasks to retain (for memory management)
21    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    /// Create a new task manager
39    pub fn new() -> Self {
40        Self {
41            state: Arc::new(RwLock::new(TaskManagerState::default())),
42            max_tasks: 1000,
43        }
44    }
45
46    /// Create a new task manager with custom capacity
47    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    /// Create a new task
59    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    /// Evict oldest completed tasks when at capacity
81    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    /// Get a task by ID
110    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    /// Get a task by ID, returning an error if not found
116    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    /// Get a task while limiting the returned conversation history.
123    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    /// Update task status
133    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    /// Add an artifact to a task
154    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    /// Add a message to task history
166    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    /// Cancel a task
178    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    /// List tasks with optional filtering
217    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    /// Get tasks by context ID
294    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    /// Get the number of tasks
304    async fn task_count(&self) -> usize {
305        self.state.read().await.tasks.len()
306    }
307
308    /// Clear all tasks (for testing)
309    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    /// Set webhook configuration for a task
317    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    /// Get webhook configuration for a task
330    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    /// Remove webhook configuration for a task
336    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}