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};
13
14#[derive(Debug, Clone)]
16pub struct TaskManager {
17 state: Arc<RwLock<TaskManagerState>>,
19 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 pub fn new() -> Self {
39 Self {
40 state: Arc::new(RwLock::new(TaskManagerState::default())),
41 max_tasks: 1000,
42 }
43 }
44
45 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 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 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 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 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 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 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 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 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 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 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 async fn task_count(&self) -> usize {
311 self.state.read().await.tasks.len()
312 }
313
314 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 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 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 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}