use std::collections::HashMap;
use std::sync::{Mutex, MutexGuard};
use chrono::Utc;
use pe_core::PeError;
use crate::dependency::{self, TaskDependency};
use crate::lifecycle::{DefaultLifecycle, TaskLifecycle};
use crate::registry::{TaskFilter, TaskRegistry};
use crate::task::{Task, TaskId, TaskStatus};
pub struct InMemoryTaskRegistry {
tasks: Mutex<HashMap<TaskId, Task>>,
deps: Mutex<Vec<TaskDependency>>,
lifecycle: Box<dyn TaskLifecycle>,
}
impl InMemoryTaskRegistry {
#[must_use]
pub fn new() -> Self {
Self {
tasks: Mutex::new(HashMap::new()),
deps: Mutex::new(Vec::new()),
lifecycle: Box::new(DefaultLifecycle),
}
}
#[must_use]
pub fn with_lifecycle(lifecycle: impl TaskLifecycle + 'static) -> Self {
Self {
tasks: Mutex::new(HashMap::new()),
deps: Mutex::new(Vec::new()),
lifecycle: Box::new(lifecycle),
}
}
fn tasks_guard(&self) -> MutexGuard<'_, HashMap<TaskId, Task>> {
match self.tasks.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
}
fn deps_guard(&self) -> MutexGuard<'_, Vec<TaskDependency>> {
match self.deps.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
}
fn collect_tree(&self, root_id: &str, tasks: &HashMap<TaskId, Task>) -> Vec<Task> {
let mut result = Vec::new();
let mut stack = vec![root_id.to_string()];
while let Some(id) = stack.pop() {
for task in tasks.values() {
if task.parent_task_id.as_deref() == Some(&id) && task.deleted_at.is_none() {
stack.push(task.id.clone());
result.push(task.clone());
}
}
}
result
}
}
impl Default for InMemoryTaskRegistry {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
impl TaskRegistry for InMemoryTaskRegistry {
async fn create(&self, task: Task) -> Result<Task, PeError> {
{
let mut tasks = self.tasks_guard();
let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
dependency::validate_parent_link(&task, &existing_tasks)?;
tasks.insert(task.id.clone(), task.clone());
}
self.lifecycle.on_create(&task);
Ok(task)
}
async fn get(&self, id: &TaskId) -> Result<Option<Task>, PeError> {
let tasks = self.tasks_guard();
Ok(tasks.get(id).filter(|t| t.deleted_at.is_none()).cloned())
}
async fn update(&self, task: &Task) -> Result<Task, PeError> {
let mut tasks = self.tasks_guard();
let mut updated = task.clone();
updated.updated_at = Some(Utc::now());
let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
dependency::validate_parent_link(&updated, &existing_tasks)?;
tasks.insert(task.id.clone(), updated.clone());
Ok(updated)
}
async fn delete(&self, id: &TaskId) -> Result<bool, PeError> {
let mut tasks = self.tasks_guard();
if let Some(task) = tasks.get_mut(id) {
task.deleted_at = Some(Utc::now());
Ok(true)
} else {
Ok(false)
}
}
async fn restore(&self, id: &TaskId) -> Result<bool, PeError> {
let mut tasks = self.tasks_guard();
if let Some(task) = tasks.get_mut(id) {
if task.deleted_at.is_some() {
task.deleted_at = None;
return Ok(true);
}
}
Ok(false)
}
async fn update_status(
&self,
id: &TaskId,
status: TaskStatus,
result: Option<serde_json::Value>,
error: Option<String>,
) -> Result<Task, PeError> {
let (updated, old_status) = {
let mut tasks = self.tasks_guard();
let task = tasks
.get(id)
.ok_or_else(|| PeError::NodeNotFound { node: id.clone() })?;
self.lifecycle.validate_transition(task, &status)?;
let old_status = task.status.clone();
let task = tasks.get_mut(id).unwrap();
task.status = status.clone();
task.updated_at = Some(Utc::now());
if let Some(r) = result {
task.result = Some(r);
}
if let Some(e) = error {
task.error = Some(e);
}
if status == TaskStatus::Completed {
task.completed_at = Some(Utc::now());
}
(task.clone(), old_status)
};
self.lifecycle.on_transition(&updated, &old_status);
match &status {
TaskStatus::Completed => self.lifecycle.on_complete(&updated),
TaskStatus::Failed => self
.lifecycle
.on_fail(&updated, updated.error.as_deref().unwrap_or("")),
TaskStatus::Cancelled => self.lifecycle.on_cancel(&updated),
_ => {}
}
Ok(updated)
}
async fn list(&self, filter: &TaskFilter) -> Result<Vec<Task>, PeError> {
let tasks = self.tasks_guard();
let limit = if filter.limit == 0 { 100 } else { filter.limit };
let results: Vec<Task> = tasks
.values()
.filter(|t| filter.include_deleted || t.deleted_at.is_none())
.filter(|t| filter.status.as_ref().is_none_or(|s| t.status == *s))
.filter(|t| {
filter
.task_type
.as_ref()
.is_none_or(|tt| t.task_type == *tt)
})
.filter(|t| filter.priority.as_ref().is_none_or(|p| t.priority == *p))
.filter(|t| {
filter
.agent_id
.as_ref()
.is_none_or(|a| t.agent_id.as_ref() == Some(a))
})
.filter(|t| filter.assignee.as_ref().is_none_or(|a| t.assignee == *a))
.filter(|t| {
filter
.parent_id
.as_ref()
.is_none_or(|p| t.parent_task_id.as_ref() == Some(p))
})
.filter(|t| filter.tag.as_ref().is_none_or(|tag| t.tags.contains(tag)))
.take(limit)
.cloned()
.collect();
Ok(results)
}
async fn get_subtasks(&self, parent_id: &TaskId) -> Result<Vec<Task>, PeError> {
let tasks = self.tasks_guard();
Ok(tasks
.values()
.filter(|t| t.parent_task_id.as_ref() == Some(parent_id) && t.deleted_at.is_none())
.cloned()
.collect())
}
async fn get_tree(&self, root_id: &TaskId) -> Result<Vec<Task>, PeError> {
let tasks = self.tasks_guard();
Ok(self.collect_tree(root_id, &tasks))
}
async fn add_dependency(&self, dep: TaskDependency) -> Result<(), PeError> {
let tasks = self.tasks_guard();
let existing_tasks = tasks.values().cloned().collect::<Vec<_>>();
drop(tasks);
let mut deps = self.deps_guard();
dependency::validate_dependency_endpoints(&dep, &existing_tasks)?;
dependency::validate_dependency_unique(&dep, &deps)?;
if dependency::would_create_cycle(&dep.task_id, &dep.depends_on_id, &deps) {
return Err(PeError::InvalidUpdate {
details: format!(
"Adding dependency {} → {} would create a cycle",
dep.task_id, dep.depends_on_id
),
});
}
deps.push(dep);
Ok(())
}
async fn remove_dependency(
&self,
task_id: &TaskId,
depends_on_id: &TaskId,
) -> Result<(), PeError> {
let mut deps = self.deps_guard();
deps.retain(|d| !(d.task_id == *task_id && d.depends_on_id == *depends_on_id));
Ok(())
}
async fn get_dependencies(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
let deps = self.deps_guard();
Ok(deps
.iter()
.filter(|d| d.task_id == *task_id)
.cloned()
.collect())
}
async fn get_dependents(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
let deps = self.deps_guard();
Ok(deps
.iter()
.filter(|d| d.depends_on_id == *task_id)
.cloned()
.collect())
}
async fn get_ready_tasks(&self) -> Result<Vec<Task>, PeError> {
let tasks = self.tasks_guard();
let deps = self.deps_guard();
let pending: Vec<TaskId> = tasks
.values()
.filter(|t| t.status == TaskStatus::Pending && t.deleted_at.is_none())
.map(|t| t.id.clone())
.collect();
let statuses: HashMap<TaskId, TaskStatus> = tasks
.values()
.map(|t| (t.id.clone(), t.status.clone()))
.collect();
let ready_ids = dependency::find_ready_tasks(&pending, &deps, &statuses);
Ok(ready_ids
.iter()
.filter_map(|id| tasks.get(id).cloned())
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dependency::DependencyType;
use std::panic::{AssertUnwindSafe, catch_unwind};
#[tokio::test]
async fn test_create_and_get() {
let reg = InMemoryTaskRegistry::new();
let task = Task::new("Build feature");
let created = reg.create(task).await.unwrap();
let fetched = reg.get(&created.id).await.unwrap().unwrap();
assert_eq!(fetched.title, "Build feature");
}
#[tokio::test]
async fn test_soft_delete_and_restore() {
let reg = InMemoryTaskRegistry::new();
let task = reg.create(Task::new("Deletable")).await.unwrap();
assert!(reg.delete(&task.id).await.unwrap());
assert!(reg.get(&task.id).await.unwrap().is_none()); assert!(reg.restore(&task.id).await.unwrap());
assert!(reg.get(&task.id).await.unwrap().is_some()); }
#[tokio::test]
async fn test_status_transition() {
let reg = InMemoryTaskRegistry::new();
let task = reg.create(Task::new("Work item")).await.unwrap();
let updated = reg
.update_status(&task.id, TaskStatus::InProgress, None, None)
.await
.unwrap();
assert_eq!(updated.status, TaskStatus::InProgress);
}
#[tokio::test]
async fn test_invalid_transition_rejected() {
let reg = InMemoryTaskRegistry::new();
let task = reg.create(Task::new("Pending task")).await.unwrap();
let err = reg
.update_status(&task.id, TaskStatus::Completed, None, None)
.await;
assert!(err.is_err());
}
#[tokio::test]
async fn test_completion_sets_timestamp() {
let reg = InMemoryTaskRegistry::new();
let task = reg.create(Task::new("Finish me")).await.unwrap();
reg.update_status(&task.id, TaskStatus::InProgress, None, None)
.await
.unwrap();
let done = reg
.update_status(&task.id, TaskStatus::Completed, None, None)
.await
.unwrap();
assert!(done.completed_at.is_some());
}
#[tokio::test]
async fn test_list_with_filter() {
let reg = InMemoryTaskRegistry::new();
reg.create(Task::agent_task("Agent work", "a1"))
.await
.unwrap();
reg.create(Task::new("Human work")).await.unwrap();
let agent_tasks = reg
.list(&TaskFilter::default().with_agent("a1"))
.await
.unwrap();
assert_eq!(agent_tasks.len(), 1);
assert_eq!(agent_tasks[0].title, "Agent work");
}
#[tokio::test]
async fn test_subtasks_and_tree() {
let reg = InMemoryTaskRegistry::new();
let parent = reg.create(Task::plan("Project")).await.unwrap();
reg.create(Task::new("Step 1").with_parent(&parent.id))
.await
.unwrap();
reg.create(Task::new("Step 2").with_parent(&parent.id))
.await
.unwrap();
let subs = reg.get_subtasks(&parent.id).await.unwrap();
assert_eq!(subs.len(), 2);
let tree = reg.get_tree(&parent.id).await.unwrap();
assert_eq!(tree.len(), 2);
}
#[tokio::test]
async fn test_dependency_cycle_rejected() {
let reg = InMemoryTaskRegistry::new();
let a = reg.create(Task::new("A")).await.unwrap();
let b = reg.create(Task::new("B")).await.unwrap();
reg.add_dependency(TaskDependency::new(&b.id, &a.id, DependencyType::Blocks))
.await
.unwrap();
let err = reg
.add_dependency(TaskDependency::new(&a.id, &b.id, DependencyType::Blocks))
.await;
assert!(err.is_err());
}
#[tokio::test]
async fn test_dependency_requires_existing_tasks() {
let reg = InMemoryTaskRegistry::new();
let task = reg.create(Task::new("Known")).await.unwrap();
let err = reg
.add_dependency(TaskDependency::new(
&task.id,
"missing",
DependencyType::Blocks,
))
.await
.unwrap_err();
assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing"));
let err = reg
.add_dependency(TaskDependency::new(
"missing",
&task.id,
DependencyType::Blocks,
))
.await
.unwrap_err();
assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing"));
}
#[tokio::test]
async fn test_duplicate_dependency_rejected() {
let reg = InMemoryTaskRegistry::new();
let blocker = reg.create(Task::new("Blocker")).await.unwrap();
let blocked = reg.create(Task::new("Blocked")).await.unwrap();
reg.add_dependency(TaskDependency::new(
&blocked.id,
&blocker.id,
DependencyType::Blocks,
))
.await
.unwrap();
let err = reg
.add_dependency(TaskDependency::new(
&blocked.id,
&blocker.id,
DependencyType::Related,
))
.await
.unwrap_err();
assert!(
matches!(err, PeError::InvalidUpdate { details } if details.contains("already exists"))
);
}
#[tokio::test]
async fn test_parent_task_must_exist() {
let reg = InMemoryTaskRegistry::new();
let err = reg
.create(Task::new("Orphan").with_parent("missing-parent"))
.await
.unwrap_err();
assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing-parent"));
}
#[tokio::test]
async fn test_parent_cycle_rejected_on_update() {
let reg = InMemoryTaskRegistry::new();
let parent = reg.create(Task::new("Parent")).await.unwrap();
let child = reg
.create(Task::new("Child").with_parent(&parent.id))
.await
.unwrap();
let mut updated_parent = parent.clone();
updated_parent.parent_task_id = Some(child.id.clone());
let err = reg.update(&updated_parent).await.unwrap_err();
assert!(matches!(err, PeError::InvalidUpdate { .. }));
let stored_parent = reg.get(&parent.id).await.unwrap().unwrap();
assert!(stored_parent.parent_task_id.is_none());
}
#[tokio::test]
async fn test_ready_tasks() {
let reg = InMemoryTaskRegistry::new();
let a = reg.create(Task::new("A")).await.unwrap();
let b = reg.create(Task::new("B")).await.unwrap();
let c = reg.create(Task::new("C")).await.unwrap();
reg.add_dependency(TaskDependency::new(&c.id, &a.id, DependencyType::Blocks))
.await
.unwrap();
let ready = reg.get_ready_tasks().await.unwrap();
let ready_ids: Vec<&str> = ready.iter().map(|t| t.id.as_str()).collect();
assert!(ready_ids.contains(&a.id.as_str()));
assert!(ready_ids.contains(&b.id.as_str()));
assert!(!ready_ids.contains(&c.id.as_str()));
reg.update_status(&a.id, TaskStatus::InProgress, None, None)
.await
.unwrap();
reg.update_status(&a.id, TaskStatus::Completed, None, None)
.await
.unwrap();
let ready2 = reg.get_ready_tasks().await.unwrap();
let ready_ids2: Vec<&str> = ready2.iter().map(|t| t.id.as_str()).collect();
assert!(ready_ids2.contains(&c.id.as_str()));
}
#[tokio::test]
async fn test_poisoned_task_lock_is_recovered() {
let reg = InMemoryTaskRegistry::new();
let _ = catch_unwind(AssertUnwindSafe(|| {
let _guard = reg.tasks.lock().unwrap();
panic!("poison tasks");
}));
let created = reg.create(Task::new("Recovered task")).await.unwrap();
assert_eq!(created.title, "Recovered task");
}
#[tokio::test]
async fn test_poisoned_dependency_lock_is_recovered() {
let reg = InMemoryTaskRegistry::new();
let a = reg.create(Task::new("A")).await.unwrap();
let b = reg.create(Task::new("B")).await.unwrap();
let _ = catch_unwind(AssertUnwindSafe(|| {
let _guard = reg.deps.lock().unwrap();
panic!("poison deps");
}));
reg.add_dependency(TaskDependency::new(&b.id, &a.id, DependencyType::Blocks))
.await
.unwrap();
let deps = reg.get_dependencies(&b.id).await.unwrap();
assert_eq!(deps.len(), 1);
}
}