Skip to main content

pe_tasks/
in_memory.rs

1//! # InMemoryTaskRegistry — HashMap-backed implementation for testing.
2//!
3//! Full CRUD + dependency DAG with cycle detection. No persistence.
4//! Use for testing, prototyping, and simple single-process agents.
5
6use 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
17/// In-memory task registry backed by `HashMap` + `Mutex`.
18///
19/// Thread-safe via interior mutability. All tasks and dependencies
20/// are stored in memory — lost on drop.
21///
22/// # Example
23///
24/// ```
25/// use pe_tasks::InMemoryTaskRegistry;
26///
27/// let registry = InMemoryTaskRegistry::new();
28/// ```
29pub struct InMemoryTaskRegistry {
30    tasks: Mutex<HashMap<TaskId, Task>>,
31    deps: Mutex<Vec<TaskDependency>>,
32    lifecycle: Box<dyn TaskLifecycle>,
33}
34
35impl InMemoryTaskRegistry {
36    /// Create with default lifecycle (transition validation only).
37    #[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    /// Create with a custom lifecycle.
47    #[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        // Acquire lock, mutate, clone result, THEN release lock before hooks.
148        // This prevents deadlock if lifecycle hooks call back into the registry.
149        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        }; // Lock released here — safe to call lifecycle hooks
173
174        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()); // hidden
329        assert!(reg.restore(&task.id).await.unwrap());
330        assert!(reg.get(&task.id).await.unwrap().is_some()); // restored
331    }
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        // Pending → Completed is invalid (must go through InProgress first)
349        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        // B depends on A
410        reg.add_dependency(TaskDependency::new(&b.id, &a.id, DependencyType::Blocks))
411            .await
412            .unwrap();
413        // A depends on B → cycle!
414        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        // C depends on A (blocks)
509        reg.add_dependency(TaskDependency::new(&c.id, &a.id, DependencyType::Blocks))
510            .await
511            .unwrap();
512
513        // Initially: A and B are ready, C is not
514        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        // Complete A → C becomes ready
521        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}