1use std::collections::HashMap;
7use std::sync::{Mutex, MutexGuard};
8
9use chrono::Utc;
10use pe_core::PeError;
11
12use crate::dependency::{self, TaskDependency};
13use crate::lifecycle::{DefaultLifecycle, TaskLifecycle};
14use crate::registry::{TaskFilter, TaskRegistry};
15use crate::task::{Task, TaskId, TaskStatus};
16
17pub struct InMemoryTaskRegistry {
30 tasks: Mutex<HashMap<TaskId, Task>>,
31 deps: Mutex<Vec<TaskDependency>>,
32 lifecycle: Box<dyn TaskLifecycle>,
33}
34
35impl InMemoryTaskRegistry {
36 #[must_use]
38 pub fn new() -> Self {
39 Self {
40 tasks: Mutex::new(HashMap::new()),
41 deps: Mutex::new(Vec::new()),
42 lifecycle: Box::new(DefaultLifecycle),
43 }
44 }
45
46 #[must_use]
48 pub fn with_lifecycle(lifecycle: impl TaskLifecycle + 'static) -> Self {
49 Self {
50 tasks: Mutex::new(HashMap::new()),
51 deps: Mutex::new(Vec::new()),
52 lifecycle: Box::new(lifecycle),
53 }
54 }
55
56 fn tasks_guard(&self) -> MutexGuard<'_, HashMap<TaskId, Task>> {
57 match self.tasks.lock() {
58 Ok(guard) => guard,
59 Err(poisoned) => poisoned.into_inner(),
60 }
61 }
62
63 fn deps_guard(&self) -> MutexGuard<'_, Vec<TaskDependency>> {
64 match self.deps.lock() {
65 Ok(guard) => guard,
66 Err(poisoned) => poisoned.into_inner(),
67 }
68 }
69
70 fn collect_tree(&self, root_id: &str, tasks: &HashMap<TaskId, Task>) -> Vec<Task> {
71 let mut result = Vec::new();
72 let mut stack = vec![root_id.to_string()];
73 while let Some(id) = stack.pop() {
74 for task in tasks.values() {
75 if task.parent_task_id.as_deref() == Some(&id) && task.deleted_at.is_none() {
76 stack.push(task.id.clone());
77 result.push(task.clone());
78 }
79 }
80 }
81 result
82 }
83}
84
85impl Default for InMemoryTaskRegistry {
86 fn default() -> Self {
87 Self::new()
88 }
89}
90
91#[async_trait::async_trait]
92impl TaskRegistry for InMemoryTaskRegistry {
93 async fn create(&self, task: Task) -> Result<Task, PeError> {
94 {
95 let mut tasks = self.tasks_guard();
96 let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
97 dependency::validate_parent_link(&task, &existing_tasks)?;
98 tasks.insert(task.id.clone(), task.clone());
99 }
100 self.lifecycle.on_create(&task);
101 Ok(task)
102 }
103
104 async fn get(&self, id: &TaskId) -> Result<Option<Task>, PeError> {
105 let tasks = self.tasks_guard();
106 Ok(tasks.get(id).filter(|t| t.deleted_at.is_none()).cloned())
107 }
108
109 async fn update(&self, task: &Task) -> Result<Task, PeError> {
110 let mut tasks = self.tasks_guard();
111 let mut updated = task.clone();
112 updated.updated_at = Some(Utc::now());
113 let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
114 dependency::validate_parent_link(&updated, &existing_tasks)?;
115 tasks.insert(task.id.clone(), updated.clone());
116 Ok(updated)
117 }
118
119 async fn delete(&self, id: &TaskId) -> Result<bool, PeError> {
120 let mut tasks = self.tasks_guard();
121 if let Some(task) = tasks.get_mut(id) {
122 task.deleted_at = Some(Utc::now());
123 Ok(true)
124 } else {
125 Ok(false)
126 }
127 }
128
129 async fn restore(&self, id: &TaskId) -> Result<bool, PeError> {
130 let mut tasks = self.tasks_guard();
131 if let Some(task) = tasks.get_mut(id) {
132 if task.deleted_at.is_some() {
133 task.deleted_at = None;
134 return Ok(true);
135 }
136 }
137 Ok(false)
138 }
139
140 async fn update_status(
141 &self,
142 id: &TaskId,
143 status: TaskStatus,
144 result: Option<serde_json::Value>,
145 error: Option<String>,
146 ) -> Result<Task, PeError> {
147 let (updated, old_status) = {
150 let mut tasks = self.tasks_guard();
151 let task = tasks
152 .get(id)
153 .ok_or_else(|| PeError::NodeNotFound { node: id.clone() })?;
154
155 self.lifecycle.validate_transition(task, &status)?;
156
157 let old_status = task.status.clone();
158 let task = tasks.get_mut(id).unwrap();
159 task.status = status.clone();
160 task.updated_at = Some(Utc::now());
161 if let Some(r) = result {
162 task.result = Some(r);
163 }
164 if let Some(e) = error {
165 task.error = Some(e);
166 }
167 if status == TaskStatus::Completed {
168 task.completed_at = Some(Utc::now());
169 }
170
171 (task.clone(), old_status)
172 }; self.lifecycle.on_transition(&updated, &old_status);
175 match &status {
176 TaskStatus::Completed => self.lifecycle.on_complete(&updated),
177 TaskStatus::Failed => self
178 .lifecycle
179 .on_fail(&updated, updated.error.as_deref().unwrap_or("")),
180 TaskStatus::Cancelled => self.lifecycle.on_cancel(&updated),
181 _ => {}
182 }
183
184 Ok(updated)
185 }
186
187 async fn list(&self, filter: &TaskFilter) -> Result<Vec<Task>, PeError> {
188 let tasks = self.tasks_guard();
189 let limit = if filter.limit == 0 { 100 } else { filter.limit };
190
191 let results: Vec<Task> = tasks
192 .values()
193 .filter(|t| filter.include_deleted || t.deleted_at.is_none())
194 .filter(|t| filter.status.as_ref().is_none_or(|s| t.status == *s))
195 .filter(|t| {
196 filter
197 .task_type
198 .as_ref()
199 .is_none_or(|tt| t.task_type == *tt)
200 })
201 .filter(|t| filter.priority.as_ref().is_none_or(|p| t.priority == *p))
202 .filter(|t| {
203 filter
204 .agent_id
205 .as_ref()
206 .is_none_or(|a| t.agent_id.as_ref() == Some(a))
207 })
208 .filter(|t| filter.assignee.as_ref().is_none_or(|a| t.assignee == *a))
209 .filter(|t| {
210 filter
211 .parent_id
212 .as_ref()
213 .is_none_or(|p| t.parent_task_id.as_ref() == Some(p))
214 })
215 .filter(|t| filter.tag.as_ref().is_none_or(|tag| t.tags.contains(tag)))
216 .take(limit)
217 .cloned()
218 .collect();
219
220 Ok(results)
221 }
222
223 async fn get_subtasks(&self, parent_id: &TaskId) -> Result<Vec<Task>, PeError> {
224 let tasks = self.tasks_guard();
225 Ok(tasks
226 .values()
227 .filter(|t| t.parent_task_id.as_ref() == Some(parent_id) && t.deleted_at.is_none())
228 .cloned()
229 .collect())
230 }
231
232 async fn get_tree(&self, root_id: &TaskId) -> Result<Vec<Task>, PeError> {
233 let tasks = self.tasks_guard();
234 Ok(self.collect_tree(root_id, &tasks))
235 }
236
237 async fn add_dependency(&self, dep: TaskDependency) -> Result<(), PeError> {
238 let tasks = self.tasks_guard();
239 let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
240 drop(tasks);
241
242 let mut deps = self.deps_guard();
243 dependency::validate_dependency_endpoints(&dep, &existing_tasks)?;
244 dependency::validate_dependency_unique(&dep, &deps)?;
245 if dependency::would_create_cycle(&dep.task_id, &dep.depends_on_id, &deps) {
246 return Err(PeError::InvalidUpdate {
247 details: format!(
248 "Adding dependency {} → {} would create a cycle",
249 dep.task_id, dep.depends_on_id
250 ),
251 });
252 }
253 deps.push(dep);
254 Ok(())
255 }
256
257 async fn remove_dependency(
258 &self,
259 task_id: &TaskId,
260 depends_on_id: &TaskId,
261 ) -> Result<(), PeError> {
262 let mut deps = self.deps_guard();
263 deps.retain(|d| !(d.task_id == *task_id && d.depends_on_id == *depends_on_id));
264 Ok(())
265 }
266
267 async fn get_dependencies(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
268 let deps = self.deps_guard();
269 Ok(deps
270 .iter()
271 .filter(|d| d.task_id == *task_id)
272 .cloned()
273 .collect())
274 }
275
276 async fn get_dependents(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
277 let deps = self.deps_guard();
278 Ok(deps
279 .iter()
280 .filter(|d| d.depends_on_id == *task_id)
281 .cloned()
282 .collect())
283 }
284
285 async fn get_ready_tasks(&self) -> Result<Vec<Task>, PeError> {
286 let tasks = self.tasks_guard();
287 let deps = self.deps_guard();
288
289 let pending: Vec<TaskId> = tasks
290 .values()
291 .filter(|t| t.status == TaskStatus::Pending && t.deleted_at.is_none())
292 .map(|t| t.id.clone())
293 .collect();
294
295 let statuses: HashMap<TaskId, TaskStatus> = tasks
296 .values()
297 .map(|t| (t.id.clone(), t.status.clone()))
298 .collect();
299
300 let ready_ids = dependency::find_ready_tasks(&pending, &deps, &statuses);
301 Ok(ready_ids
302 .iter()
303 .filter_map(|id| tasks.get(id).cloned())
304 .collect())
305 }
306}
307
308#[cfg(test)]
309mod tests {
310 use super::*;
311 use crate::dependency::DependencyType;
312 use std::panic::{AssertUnwindSafe, catch_unwind};
313
314 #[tokio::test]
315 async fn test_create_and_get() {
316 let reg = InMemoryTaskRegistry::new();
317 let task = Task::new("Build feature");
318 let created = reg.create(task).await.unwrap();
319 let fetched = reg.get(&created.id).await.unwrap().unwrap();
320 assert_eq!(fetched.title, "Build feature");
321 }
322
323 #[tokio::test]
324 async fn test_soft_delete_and_restore() {
325 let reg = InMemoryTaskRegistry::new();
326 let task = reg.create(Task::new("Deletable")).await.unwrap();
327 assert!(reg.delete(&task.id).await.unwrap());
328 assert!(reg.get(&task.id).await.unwrap().is_none()); assert!(reg.restore(&task.id).await.unwrap());
330 assert!(reg.get(&task.id).await.unwrap().is_some()); }
332
333 #[tokio::test]
334 async fn test_status_transition() {
335 let reg = InMemoryTaskRegistry::new();
336 let task = reg.create(Task::new("Work item")).await.unwrap();
337 let updated = reg
338 .update_status(&task.id, TaskStatus::InProgress, None, None)
339 .await
340 .unwrap();
341 assert_eq!(updated.status, TaskStatus::InProgress);
342 }
343
344 #[tokio::test]
345 async fn test_invalid_transition_rejected() {
346 let reg = InMemoryTaskRegistry::new();
347 let task = reg.create(Task::new("Pending task")).await.unwrap();
348 let err = reg
350 .update_status(&task.id, TaskStatus::Completed, None, None)
351 .await;
352 assert!(err.is_err());
353 }
354
355 #[tokio::test]
356 async fn test_completion_sets_timestamp() {
357 let reg = InMemoryTaskRegistry::new();
358 let task = reg.create(Task::new("Finish me")).await.unwrap();
359 reg.update_status(&task.id, TaskStatus::InProgress, None, None)
360 .await
361 .unwrap();
362 let done = reg
363 .update_status(&task.id, TaskStatus::Completed, None, None)
364 .await
365 .unwrap();
366 assert!(done.completed_at.is_some());
367 }
368
369 #[tokio::test]
370 async fn test_list_with_filter() {
371 let reg = InMemoryTaskRegistry::new();
372 reg.create(Task::agent_task("Agent work", "a1"))
373 .await
374 .unwrap();
375 reg.create(Task::new("Human work")).await.unwrap();
376
377 let agent_tasks = reg
378 .list(&TaskFilter::default().with_agent("a1"))
379 .await
380 .unwrap();
381 assert_eq!(agent_tasks.len(), 1);
382 assert_eq!(agent_tasks[0].title, "Agent work");
383 }
384
385 #[tokio::test]
386 async fn test_subtasks_and_tree() {
387 let reg = InMemoryTaskRegistry::new();
388 let parent = reg.create(Task::plan("Project")).await.unwrap();
389 reg.create(Task::new("Step 1").with_parent(&parent.id))
390 .await
391 .unwrap();
392 reg.create(Task::new("Step 2").with_parent(&parent.id))
393 .await
394 .unwrap();
395
396 let subs = reg.get_subtasks(&parent.id).await.unwrap();
397 assert_eq!(subs.len(), 2);
398
399 let tree = reg.get_tree(&parent.id).await.unwrap();
400 assert_eq!(tree.len(), 2);
401 }
402
403 #[tokio::test]
404 async fn test_dependency_cycle_rejected() {
405 let reg = InMemoryTaskRegistry::new();
406 let a = reg.create(Task::new("A")).await.unwrap();
407 let b = reg.create(Task::new("B")).await.unwrap();
408
409 reg.add_dependency(TaskDependency::new(&b.id, &a.id, DependencyType::Blocks))
411 .await
412 .unwrap();
413 let err = reg
415 .add_dependency(TaskDependency::new(&a.id, &b.id, DependencyType::Blocks))
416 .await;
417 assert!(err.is_err());
418 }
419
420 #[tokio::test]
421 async fn test_dependency_requires_existing_tasks() {
422 let reg = InMemoryTaskRegistry::new();
423 let task = reg.create(Task::new("Known")).await.unwrap();
424
425 let err = reg
426 .add_dependency(TaskDependency::new(
427 &task.id,
428 "missing",
429 DependencyType::Blocks,
430 ))
431 .await
432 .unwrap_err();
433 assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing"));
434
435 let err = reg
436 .add_dependency(TaskDependency::new(
437 "missing",
438 &task.id,
439 DependencyType::Blocks,
440 ))
441 .await
442 .unwrap_err();
443 assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing"));
444 }
445
446 #[tokio::test]
447 async fn test_duplicate_dependency_rejected() {
448 let reg = InMemoryTaskRegistry::new();
449 let blocker = reg.create(Task::new("Blocker")).await.unwrap();
450 let blocked = reg.create(Task::new("Blocked")).await.unwrap();
451
452 reg.add_dependency(TaskDependency::new(
453 &blocked.id,
454 &blocker.id,
455 DependencyType::Blocks,
456 ))
457 .await
458 .unwrap();
459
460 let err = reg
461 .add_dependency(TaskDependency::new(
462 &blocked.id,
463 &blocker.id,
464 DependencyType::Related,
465 ))
466 .await
467 .unwrap_err();
468 assert!(
469 matches!(err, PeError::InvalidUpdate { details } if details.contains("already exists"))
470 );
471 }
472
473 #[tokio::test]
474 async fn test_parent_task_must_exist() {
475 let reg = InMemoryTaskRegistry::new();
476 let err = reg
477 .create(Task::new("Orphan").with_parent("missing-parent"))
478 .await
479 .unwrap_err();
480 assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing-parent"));
481 }
482
483 #[tokio::test]
484 async fn test_parent_cycle_rejected_on_update() {
485 let reg = InMemoryTaskRegistry::new();
486 let parent = reg.create(Task::new("Parent")).await.unwrap();
487 let child = reg
488 .create(Task::new("Child").with_parent(&parent.id))
489 .await
490 .unwrap();
491
492 let mut updated_parent = parent.clone();
493 updated_parent.parent_task_id = Some(child.id.clone());
494 let err = reg.update(&updated_parent).await.unwrap_err();
495 assert!(matches!(err, PeError::InvalidUpdate { .. }));
496
497 let stored_parent = reg.get(&parent.id).await.unwrap().unwrap();
498 assert!(stored_parent.parent_task_id.is_none());
499 }
500
501 #[tokio::test]
502 async fn test_ready_tasks() {
503 let reg = InMemoryTaskRegistry::new();
504 let a = reg.create(Task::new("A")).await.unwrap();
505 let b = reg.create(Task::new("B")).await.unwrap();
506 let c = reg.create(Task::new("C")).await.unwrap();
507
508 reg.add_dependency(TaskDependency::new(&c.id, &a.id, DependencyType::Blocks))
510 .await
511 .unwrap();
512
513 let ready = reg.get_ready_tasks().await.unwrap();
515 let ready_ids: Vec<&str> = ready.iter().map(|t| t.id.as_str()).collect();
516 assert!(ready_ids.contains(&a.id.as_str()));
517 assert!(ready_ids.contains(&b.id.as_str()));
518 assert!(!ready_ids.contains(&c.id.as_str()));
519
520 reg.update_status(&a.id, TaskStatus::InProgress, None, None)
522 .await
523 .unwrap();
524 reg.update_status(&a.id, TaskStatus::Completed, None, None)
525 .await
526 .unwrap();
527 let ready2 = reg.get_ready_tasks().await.unwrap();
528 let ready_ids2: Vec<&str> = ready2.iter().map(|t| t.id.as_str()).collect();
529 assert!(ready_ids2.contains(&c.id.as_str()));
530 }
531
532 #[tokio::test]
533 async fn test_poisoned_task_lock_is_recovered() {
534 let reg = InMemoryTaskRegistry::new();
535
536 let _ = catch_unwind(AssertUnwindSafe(|| {
537 let _guard = reg.tasks.lock().unwrap();
538 panic!("poison tasks");
539 }));
540
541 let created = reg.create(Task::new("Recovered task")).await.unwrap();
542 assert_eq!(created.title, "Recovered task");
543 }
544
545 #[tokio::test]
546 async fn test_poisoned_dependency_lock_is_recovered() {
547 let reg = InMemoryTaskRegistry::new();
548 let a = reg.create(Task::new("A")).await.unwrap();
549 let b = reg.create(Task::new("B")).await.unwrap();
550
551 let _ = catch_unwind(AssertUnwindSafe(|| {
552 let _guard = reg.deps.lock().unwrap();
553 panic!("poison deps");
554 }));
555
556 reg.add_dependency(TaskDependency::new(&b.id, &a.id, DependencyType::Blocks))
557 .await
558 .unwrap();
559 let deps = reg.get_dependencies(&b.id).await.unwrap();
560 assert_eq!(deps.len(), 1);
561 }
562}