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};
13
14/// A2A Task Manager - handles task creation, updates, and queries
15#[derive(Debug, Clone)]
16pub struct TaskManager {
17    /// All mutable task manager state lives behind one lock so related indexes stay in sync.
18    state: Arc<RwLock<TaskManagerState>>,
19    /// Maximum tasks to retain (for memory management)
20    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    /// Create a new task manager
38    pub fn new() -> Self {
39        Self {
40            state: Arc::new(RwLock::new(TaskManagerState::default())),
41            max_tasks: 1000,
42        }
43    }
44
45    /// Create a new task manager with custom capacity
46    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    /// Create a new task
58    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    /// Evict oldest completed tasks when at capacity
80    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    /// Get a task by ID
109    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    /// Get a task by ID, returning an error if not found
115    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    /// Update task status
122    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    /// Add an artifact to a task
138    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    /// Add a message to task history
150    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    /// Cancel a task
162    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    /// List tasks with optional filtering
218    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    /// Get tasks by context ID
295    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    /// Get the number of tasks
305    pub async fn task_count(&self) -> usize {
306        self.state.read().await.tasks.len()
307    }
308
309    /// Clear all tasks (for testing)
310    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    /// Set webhook configuration for a task
318    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    /// Get webhook configuration for a task
333    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    /// Remove webhook configuration for a task
339    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}