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