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};
13use super::webhook::parse_webhook_url;
14
15#[derive(Debug, Clone)]
17pub struct TaskManager {
18 state: Arc<RwLock<TaskManagerState>>,
20 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 pub fn new() -> Self {
40 Self {
41 state: Arc::new(RwLock::new(TaskManagerState::default())),
42 max_tasks: 1000,
43 }
44 }
45
46 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 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 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 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 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 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 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 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 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 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 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 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 async fn task_count(&self) -> usize {
320 self.state.read().await.tasks.len()
321 }
322
323 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 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 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 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}