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 task = self.get_task_or_error(task_id).await?;
125        Ok(Self::clone_task_for_history(task, history_length))
126    }
127
128    /// Update task status
129    pub(crate) async fn update_status(
130        &self,
131        task_id: &str,
132        state: TaskState,
133        message: Option<Message>,
134    ) -> A2aResult<Task> {
135        let mut manager_state = self.state.write().await;
136        let task = manager_state
137            .tasks
138            .get_mut(task_id)
139            .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
140
141        task.status = match message {
142            Some(msg) => TaskStatus::with_message(state, msg),
143            None => TaskStatus::new(state),
144        };
145
146        Ok(task.clone())
147    }
148
149    /// Add an artifact to a task
150    async fn add_artifact(&self, task_id: &str, artifact: Artifact) -> 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.artifacts.push(artifact);
158        Ok(task.clone())
159    }
160
161    /// Add a message to task history
162    pub(crate) async fn add_message(&self, task_id: &str, message: Message) -> 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        task.history.push(message);
170        Ok(task.clone())
171    }
172
173    /// Cancel a task
174    pub(crate) async fn cancel_task(&self, task_id: &str) -> A2aResult<Task> {
175        let mut state = self.state.write().await;
176        let task = state
177            .tasks
178            .get_mut(task_id)
179            .ok_or_else(|| A2aError::TaskNotFound(task_id.to_string()))?;
180
181        if !task.is_cancelable() {
182            return Err(A2aError::TaskNotCancelable(format!(
183                "Task {} is in state {:?} and cannot be canceled",
184                task_id, task.status.state
185            )));
186        }
187
188        task.status = TaskStatus::new(TaskState::Canceled);
189        Ok(task.clone())
190    }
191
192    fn matches_list_filters(
193        task: &Task,
194        status: Option<&TaskState>,
195        updated_after: Option<&chrono::DateTime<chrono::Utc>>,
196    ) -> bool {
197        if let Some(status) = status
198            && &task.status.state != status
199        {
200            return false;
201        }
202
203        if let Some(updated_after) = updated_after
204            && task.status.timestamp < *updated_after
205        {
206            return false;
207        }
208
209        true
210    }
211
212    fn clone_task_for_history(mut task: Task, history_length: usize) -> Task {
213        if task.history.len() > history_length {
214            let trim_count = task.history.len() - history_length;
215            drop(task.history.drain(..trim_count));
216        }
217
218        task
219    }
220
221    fn clone_task_for_listing(task: &Task, include_artifacts: bool, history_length: Option<usize>) -> Task {
222        let mut task = Self::clone_task_for_history(task.clone(), history_length.unwrap_or(0));
223
224        if !include_artifacts {
225            task.artifacts.clear();
226        }
227
228        task
229    }
230
231    /// List tasks with optional filtering
232    pub(crate) async fn list_tasks(&self, params: ListTasksParams) -> ListTasksResult {
233        let updated_after = params
234            .last_updated_after
235            .as_deref()
236            .and_then(|after| chrono::DateTime::parse_from_rfc3339(after).ok())
237            .map(|after| after.to_utc());
238
239        let mut matching_tasks: Vec<(String, chrono::DateTime<chrono::Utc>)> = {
240            let state = self.state.read().await;
241            if let Some(context_id) = params.context_id.as_deref() {
242                state
243                    .contexts
244                    .get(context_id)
245                    .into_iter()
246                    .flat_map(|task_ids| task_ids.iter())
247                    .filter_map(|task_id| {
248                        let task = state.tasks.get(task_id)?;
249                        Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
250                            .then(|| (task_id.clone(), task.status.timestamp))
251                    })
252                    .collect()
253            } else {
254                state
255                    .tasks
256                    .iter()
257                    .filter(|(_, task)| {
258                        Self::matches_list_filters(task, params.status.as_ref(), updated_after.as_ref())
259                    })
260                    .map(|(task_id, task)| (task_id.clone(), task.status.timestamp))
261                    .collect()
262            }
263        };
264
265        matching_tasks.sort_by_key(|a| std::cmp::Reverse(a.1));
266
267        let total_size = u32::try_from(matching_tasks.len()).unwrap_or(u32::MAX);
268        let page_size = params.page_size.unwrap_or(50).min(100);
269        let start_idx = params
270            .page_token
271            .as_ref()
272            .and_then(|token| token.parse::<usize>().ok())
273            .unwrap_or(0);
274
275        let end_idx = (start_idx + page_size as usize).min(matching_tasks.len());
276        let next_page_token = if end_idx < matching_tasks.len() {
277            Some(end_idx.to_string())
278        } else {
279            None
280        };
281
282        let include_artifacts = params.include_artifacts == Some(true);
283        let history_length = params.history_length.map(|len| len as usize);
284        let page_task_ids: Vec<_> = matching_tasks.into_iter().skip(start_idx).take(page_size as usize).collect();
285        let result = if page_task_ids.is_empty() {
286            Vec::new()
287        } else {
288            let state = self.state.read().await;
289            page_task_ids
290                .into_iter()
291                .filter_map(|(task_id, _)| {
292                    state
293                        .tasks
294                        .get(&task_id)
295                        .map(|task| Self::clone_task_for_listing(task, include_artifacts, history_length))
296                })
297                .collect()
298        };
299
300        ListTasksResult {
301            tasks: result,
302            total_size: Some(total_size),
303            page_size: Some(page_size),
304            next_page_token,
305        }
306    }
307
308    /// Get tasks by context ID
309    async fn get_tasks_by_context(&self, context_id: &str) -> Vec<Task> {
310        let state = self.state.read().await;
311        state
312            .contexts
313            .get(context_id)
314            .map(|task_ids| task_ids.iter().filter_map(|id| state.tasks.get(id).cloned()).collect())
315            .unwrap_or_default()
316    }
317
318    /// Get the number of tasks
319    async fn task_count(&self) -> usize {
320        self.state.read().await.tasks.len()
321    }
322
323    /// Clear all tasks (for testing)
324    pub async fn clear(&self) {
325        let mut state = self.state.write().await;
326        state.tasks.clear();
327        state.contexts.clear();
328        state.webhook_configs.clear();
329    }
330
331    /// Set webhook configuration for a task
332    pub(crate) async fn set_webhook_config(&self, config: TaskPushNotificationConfig) -> A2aResult<()> {
333        drop(parse_webhook_url(&config.url).map_err(A2aError::UnsupportedOperation)?);
334
335        let mut state = self.state.write().await;
336        if !state.tasks.contains_key(&config.task_id) {
337            return Err(A2aError::TaskNotFound(config.task_id));
338        }
339
340        drop(state.webhook_configs.insert(config.task_id.clone(), config));
341        Ok(())
342    }
343
344    /// Get webhook configuration for a task
345    pub(crate) async fn get_webhook_config(&self, task_id: &str) -> Option<TaskPushNotificationConfig> {
346        let state = self.state.read().await;
347        state.webhook_configs.get(task_id).cloned()
348    }
349
350    /// Remove webhook configuration for a task
351    pub async fn remove_webhook_config(&self, task_id: &str) {
352        let mut state = self.state.write().await;
353        drop(state.webhook_configs.remove(task_id));
354    }
355}
356
357#[cfg(test)]
358mod tests {
359    use super::*;
360    use crate::types::MessageRole;
361
362    #[tokio::test]
363    async fn test_create_task() {
364        let manager = TaskManager::new();
365        let task = manager.create_task(None).await;
366
367        assert!(!task.id.is_empty());
368        assert_eq!(task.state(), TaskState::Submitted);
369        assert_eq!(manager.task_count().await, 1);
370    }
371
372    #[tokio::test]
373    async fn test_create_task_with_context() {
374        let manager = TaskManager::new();
375        let task = manager.create_task(Some("ctx-1".to_string())).await;
376
377        assert_eq!(task.context_id, Some("ctx-1".to_string()));
378
379        let tasks = manager.get_tasks_by_context("ctx-1").await;
380        assert_eq!(tasks.len(), 1);
381        assert_eq!(tasks[0].id, task.id);
382    }
383
384    #[tokio::test]
385    async fn test_get_task() {
386        let manager = TaskManager::new();
387        let task = manager.create_task(None).await;
388
389        let retrieved = manager.get_task(&task.id).await;
390        assert!(retrieved.is_some());
391        assert_eq!(retrieved.unwrap().id, task.id);
392
393        let missing = manager.get_task("nonexistent").await;
394        assert!(missing.is_none());
395    }
396
397    #[tokio::test]
398    async fn test_update_status() {
399        let manager = TaskManager::new();
400        let task = manager.create_task(None).await;
401
402        let updated = manager.update_status(&task.id, TaskState::Working, None).await.expect("update");
403        assert_eq!(updated.state(), TaskState::Working);
404
405        let msg = Message::agent_text("Task completed successfully");
406        let completed = manager
407            .update_status(&task.id, TaskState::Completed, Some(msg))
408            .await
409            .expect("complete");
410        assert_eq!(completed.state(), TaskState::Completed);
411        assert!(completed.status.message.is_some());
412    }
413
414    #[tokio::test]
415    async fn test_add_artifact() {
416        let manager = TaskManager::new();
417        let task = manager.create_task(None).await;
418
419        let artifact = Artifact::text("art-1", "Generated content");
420        let updated = manager.add_artifact(&task.id, artifact).await.expect("add artifact");
421        assert_eq!(updated.artifacts.len(), 1);
422        assert_eq!(updated.artifacts[0].id, "art-1");
423    }
424
425    #[tokio::test]
426    async fn test_cancel_task() {
427        let manager = TaskManager::new();
428        let task = manager.create_task(None).await;
429
430        let canceled = manager.cancel_task(&task.id).await.expect("cancel");
431        assert_eq!(canceled.state(), TaskState::Canceled);
432    }
433
434    #[tokio::test]
435    async fn test_cancel_completed_task_fails() {
436        let manager = TaskManager::new();
437        let task = manager.create_task(None).await;
438
439        drop(
440            manager
441                .update_status(&task.id, TaskState::Completed, None)
442                .await
443                .expect("complete"),
444        );
445
446        let result = manager.cancel_task(&task.id).await;
447        drop(result.unwrap_err());
448    }
449
450    #[tokio::test]
451    async fn test_eviction_cleans_context_and_webhook_indexes() {
452        let manager = TaskManager::with_capacity(1);
453        let task = manager.create_task(Some("ctx-1".to_string())).await;
454
455        drop(
456            manager
457                .update_status(&task.id, TaskState::Completed, None)
458                .await
459                .expect("complete"),
460        );
461        manager
462            .set_webhook_config(TaskPushNotificationConfig {
463                task_id: task.id.clone(),
464                url: "https://example.com/webhook".to_string(),
465                authentication: None,
466            })
467            .await
468            .expect("set webhook");
469
470        let replacement = manager.create_task(None).await;
471
472        assert_eq!(manager.task_count().await, 1);
473        assert!(manager.get_task(&task.id).await.is_none());
474        assert!(manager.get_webhook_config(&task.id).await.is_none());
475        assert!(manager.get_tasks_by_context("ctx-1").await.is_empty());
476        assert_eq!(manager.get_task(&replacement.id).await.unwrap().id, replacement.id);
477    }
478
479    #[tokio::test]
480    async fn test_list_tasks() {
481        let manager = TaskManager::new();
482
483        let task1 = manager.create_task(Some("ctx-1".to_string())).await;
484        let _task2 = manager.create_task(Some("ctx-1".to_string())).await;
485        let _task3 = manager.create_task(Some("ctx-2".to_string())).await;
486        drop(
487            manager
488                .add_message(&task1.id, Message::user_text("private message"))
489                .await
490                .expect("add message"),
491        );
492
493        let all = manager.list_tasks(ListTasksParams::default()).await;
494        assert_eq!(all.tasks.len(), 3);
495        let listed_task = all.tasks.iter().find(|task| task.id == task1.id).expect("listed task");
496        assert!(listed_task.history.is_empty());
497
498        let ctx1_tasks = manager
499            .list_tasks(ListTasksParams {
500                context_id: Some("ctx-1".to_string()),
501                ..Default::default()
502            })
503            .await;
504        assert_eq!(ctx1_tasks.tasks.len(), 2);
505    }
506
507    #[tokio::test]
508    async fn test_list_tasks_paginates_and_trims_after_sorting() {
509        let manager = TaskManager::new();
510
511        let older = manager.create_task(Some("ctx-1".to_string())).await;
512        tokio::time::sleep(std::time::Duration::from_millis(2)).await;
513        let newer = manager.create_task(Some("ctx-1".to_string())).await;
514
515        drop(
516            manager
517                .add_artifact(&newer.id, Artifact::text("art-1", "Generated content"))
518                .await
519                .expect("add artifact"),
520        );
521        drop(
522            manager
523                .add_message(&newer.id, Message::user_text("Hello"))
524                .await
525                .expect("add msg1"),
526        );
527        drop(
528            manager
529                .add_message(&newer.id, Message::agent_text("Hi there"))
530                .await
531                .expect("add msg2"),
532        );
533
534        let first_page = manager
535            .list_tasks(ListTasksParams {
536                context_id: Some("ctx-1".to_string()),
537                page_size: Some(1),
538                history_length: Some(1),
539                include_artifacts: Some(false),
540                ..Default::default()
541            })
542            .await;
543
544        assert_eq!(first_page.total_size, Some(2));
545        assert_eq!(first_page.next_page_token.as_deref(), Some("1"));
546        assert_eq!(first_page.tasks.len(), 1);
547        assert_eq!(first_page.tasks[0].id, newer.id);
548        assert!(first_page.tasks[0].artifacts.is_empty());
549        assert_eq!(first_page.tasks[0].history.len(), 1);
550        assert_eq!(first_page.tasks[0].history[0].role, MessageRole::Agent);
551
552        let second_page = manager
553            .list_tasks(ListTasksParams {
554                context_id: Some("ctx-1".to_string()),
555                page_size: Some(1),
556                page_token: Some("1".to_string()),
557                ..Default::default()
558            })
559            .await;
560
561        assert_eq!(second_page.tasks.len(), 1);
562        assert_eq!(second_page.tasks[0].id, older.id);
563        assert!(second_page.next_page_token.is_none());
564    }
565
566    #[tokio::test]
567    async fn test_add_message_to_history() {
568        let manager = TaskManager::new();
569        let task = manager.create_task(None).await;
570
571        let msg1 = Message::user_text("Hello");
572        let msg2 = Message::agent_text("Hi there!");
573
574        drop(manager.add_message(&task.id, msg1).await.expect("add msg1"));
575        let updated = manager.add_message(&task.id, msg2).await.expect("add msg2");
576
577        assert_eq!(updated.history.len(), 2);
578        assert_eq!(updated.history[0].role, MessageRole::User);
579        assert_eq!(updated.history[1].role, MessageRole::Agent);
580    }
581
582    #[tokio::test]
583    async fn test_get_task_history_length_is_enforced() {
584        let manager = TaskManager::new();
585        let task = manager.create_task(None).await;
586        drop(
587            manager
588                .add_message(&task.id, Message::user_text("first"))
589                .await
590                .expect("add first message"),
591        );
592        drop(
593            manager
594                .add_message(&task.id, Message::agent_text("second"))
595                .await
596                .expect("add second message"),
597        );
598
599        let without_history = manager
600            .get_task_or_error_with_history(&task.id, 0)
601            .await
602            .expect("get task without history");
603        assert!(without_history.history.is_empty());
604
605        let last_message = manager
606            .get_task_or_error_with_history(&task.id, 1)
607            .await
608            .expect("get task with one history item");
609        assert_eq!(last_message.history.len(), 1);
610        assert_eq!(last_message.history[0].role, MessageRole::Agent);
611    }
612
613    #[tokio::test]
614    async fn test_webhook_url_validation_requires_exact_localhost_for_http() {
615        let manager = TaskManager::new();
616        let task = manager.create_task(None).await;
617
618        let invalid_urls = [
619            "http://localhost.evil.example/hook",
620            "http://localhost@evil.example/hook",
621            "http://example.com/hook",
622            "ftp://example.com/hook",
623            "https://user:password@example.com/hook",
624        ];
625
626        for url in invalid_urls {
627            let result = manager
628                .set_webhook_config(TaskPushNotificationConfig {
629                    task_id: task.id.clone(),
630                    url: url.to_string(),
631                    authentication: None,
632                })
633                .await;
634            assert!(result.is_err(), "URL should be rejected: {url}");
635        }
636
637        for url in [
638            "https://example.com/hook",
639            "http://localhost:8080/hook",
640            "http://127.0.0.1:8080/hook",
641            "http://[::1]:8080/hook",
642        ] {
643            manager
644                .set_webhook_config(TaskPushNotificationConfig {
645                    task_id: task.id.clone(),
646                    url: url.to_string(),
647                    authentication: None,
648                })
649                .await
650                .expect("valid webhook URL");
651        }
652    }
653}