use std::collections::HashMap;
use std::path::Path;
use chrono::Utc;
use pe_core::PeError;
use pe_core::scope::ExecutionScope;
use surrealdb::Surreal;
use surrealdb::engine::any::Any;
use surrealdb::types::SurrealValue;
use crate::dependency::{self, DependencyType, TaskDependency};
use crate::lifecycle::{DefaultLifecycle, TaskLifecycle};
use crate::registry::{TaskFilter, TaskRegistry};
use crate::task::{Task, TaskId, TaskPriority, TaskStatus, TaskType};
const TASKS_SCHEMA: &str = "
DEFINE TABLE IF NOT EXISTS tasks SCHEMALESS;
DEFINE INDEX IF NOT EXISTS idx_tasks_task_id ON tasks FIELDS task_id UNIQUE;
DEFINE INDEX IF NOT EXISTS idx_tasks_status ON tasks FIELDS status;
DEFINE INDEX IF NOT EXISTS idx_tasks_type ON tasks FIELDS task_type;
DEFINE INDEX IF NOT EXISTS idx_tasks_agent ON tasks FIELDS agent_id;
DEFINE INDEX IF NOT EXISTS idx_tasks_assignee ON tasks FIELDS assignee;
DEFINE INDEX IF NOT EXISTS idx_tasks_parent ON tasks FIELDS parent_task_id;
DEFINE INDEX IF NOT EXISTS idx_tasks_deleted ON tasks FIELDS is_deleted;
";
const DEPENDENCIES_SCHEMA: &str = "
DEFINE TABLE IF NOT EXISTS task_dependencies SCHEMALESS;
DEFINE INDEX IF NOT EXISTS idx_task_deps_task ON task_dependencies FIELDS task_id;
DEFINE INDEX IF NOT EXISTS idx_task_deps_depends_on ON task_dependencies FIELDS depends_on_id;
DEFINE INDEX IF NOT EXISTS idx_task_deps_pair ON task_dependencies FIELDS task_id, depends_on_id;
";
pub struct SurrealTaskRegistry {
db: Surreal<Any>,
namespace: String,
database: String,
lifecycle: Box<dyn TaskLifecycle>,
}
impl SurrealTaskRegistry {
pub async fn connect(url: &str, scope: &ExecutionScope) -> Result<Self, PeError> {
let db = surrealdb::engine::any::connect(url)
.await
.map_err(|error| PeError::Storage {
details: format!("SurrealDB task registry connect to '{url}' failed: {error}"),
})?;
Self::from_existing(db, scope).await
}
pub async fn connect_memory(scope: &ExecutionScope) -> Result<Self, PeError> {
Self::connect("mem://", scope).await
}
pub async fn connect_embedded(
path: impl AsRef<Path>,
scope: &ExecutionScope,
) -> Result<Self, PeError> {
let endpoint = format!("surrealkv://{}", path.as_ref().to_string_lossy());
Self::connect(&endpoint, scope).await
}
pub async fn from_existing(db: Surreal<Any>, scope: &ExecutionScope) -> Result<Self, PeError> {
Self::from_existing_with_lifecycle(db, scope, DefaultLifecycle).await
}
pub async fn from_existing_with_lifecycle(
db: Surreal<Any>,
scope: &ExecutionScope,
lifecycle: impl TaskLifecycle + 'static,
) -> Result<Self, PeError> {
let registry = Self {
db,
namespace: scope.namespace().to_string(),
database: scope.database().to_string(),
lifecycle: Box::new(lifecycle),
};
registry.apply_scope().await?;
registry.init_schema().await?;
Ok(registry)
}
async fn apply_scope(&self) -> Result<(), PeError> {
self.db
.use_ns(&self.namespace)
.use_db(&self.database)
.await
.map_err(|error| PeError::Storage {
details: format!("Task registry scope application failed: {error}"),
})?;
Ok(())
}
async fn init_schema(&self) -> Result<(), PeError> {
self.db
.query(TASKS_SCHEMA)
.await
.map_err(|error| PeError::Storage {
details: format!("Task schema init failed: {error}"),
})?;
self.db
.query(DEPENDENCIES_SCHEMA)
.await
.map_err(|error| PeError::Storage {
details: format!("Task dependency schema init failed: {error}"),
})?;
Ok(())
}
async fn write_task(&self, task: &Task) -> Result<Task, PeError> {
self.apply_scope().await?;
let params = TaskWriteParams::from_task(task)?;
let task_id = task.id.clone();
self.db
.query("DELETE FROM tasks WHERE task_id = $task_id")
.bind(TaskIdParams {
task_id: task_id.clone(),
})
.await
.map_err(|error| PeError::Storage {
details: format!("Task update delete failed: {error}"),
})?;
self.db
.query(
"CREATE tasks SET \
task_id = $task_id, \
task = $task, \
status = $status, \
task_type = $task_type, \
priority = $priority, \
agent_id = $agent_id, \
assignee = $assignee, \
parent_task_id = $parent_task_id, \
tags = $tags, \
is_deleted = $is_deleted, \
created_at = $created_at, \
updated_at = $updated_at, \
completed_at = $completed_at, \
deleted_at = $deleted_at",
)
.bind(params)
.await
.map_err(|error| PeError::Storage {
details: format!("Task update create failed: {error}"),
})?;
Ok(task.clone())
}
async fn fetch_task(
&self,
task_id: &TaskId,
include_deleted: bool,
) -> Result<Option<Task>, PeError> {
self.apply_scope().await?;
let mut result = self
.db
.query("SELECT task FROM tasks WHERE task_id = $task_id LIMIT 1")
.bind(TaskIdParams {
task_id: task_id.clone(),
})
.await
.map_err(|error| PeError::Storage {
details: format!("Task get failed: {error}"),
})?;
let row: Option<TaskRow> = result.take(0).map_err(|error| PeError::Storage {
details: format!("Task get deserialize failed: {error}"),
})?;
let Some(task) = row.map(|row| decode_task(row.task)).transpose()? else {
return Ok(None);
};
if !include_deleted && task.deleted_at.is_some() {
return Ok(None);
}
Ok(Some(task))
}
async fn fetch_all_tasks(&self, include_deleted: bool) -> Result<Vec<Task>, PeError> {
self.apply_scope().await?;
let query = if include_deleted {
"SELECT task FROM tasks"
} else {
"SELECT task FROM tasks WHERE is_deleted = false"
};
let mut result = self
.db
.query(query)
.await
.map_err(|error| PeError::Storage {
details: format!("Task list failed: {error}"),
})?;
let rows: Vec<TaskRow> = result.take(0).map_err(|error| PeError::Storage {
details: format!("Task list deserialize failed: {error}"),
})?;
rows.into_iter().map(|row| decode_task(row.task)).collect()
}
async fn fetch_all_dependencies(&self) -> Result<Vec<TaskDependency>, PeError> {
self.apply_scope().await?;
let mut result = self
.db
.query("SELECT dependency FROM task_dependencies")
.await
.map_err(|error| PeError::Storage {
details: format!("Task dependency list failed: {error}"),
})?;
let rows: Vec<DependencyRow> = result.take(0).map_err(|error| PeError::Storage {
details: format!("Task dependency list deserialize failed: {error}"),
})?;
rows.into_iter()
.map(|row| decode_dependency(row.dependency))
.collect()
}
fn collect_tree(root_id: &str, tasks: &[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 {
if task.parent_task_id.as_deref() == Some(&id) && task.deleted_at.is_none() {
stack.push(task.id.clone());
result.push(task.clone());
}
}
}
result
}
}
#[async_trait::async_trait]
impl TaskRegistry for SurrealTaskRegistry {
async fn create(&self, task: Task) -> Result<Task, PeError> {
let tasks = self.fetch_all_tasks(false).await?;
dependency::validate_parent_link(&task, &tasks)?;
let created = self.write_task(&task).await?;
self.lifecycle.on_create(&created);
Ok(created)
}
async fn get(&self, id: &TaskId) -> Result<Option<Task>, PeError> {
self.fetch_task(id, false).await
}
async fn update(&self, task: &Task) -> Result<Task, PeError> {
let mut updated = task.clone();
updated.updated_at = Some(Utc::now());
let tasks = self.fetch_all_tasks(false).await?;
dependency::validate_parent_link(&updated, &tasks)?;
self.write_task(&updated).await
}
async fn delete(&self, id: &TaskId) -> Result<bool, PeError> {
let Some(mut task) = self.fetch_task(id, true).await? else {
return Ok(false);
};
task.deleted_at = Some(Utc::now());
task.updated_at = Some(Utc::now());
self.write_task(&task).await?;
Ok(true)
}
async fn restore(&self, id: &TaskId) -> Result<bool, PeError> {
let Some(mut task) = self.fetch_task(id, true).await? else {
return Ok(false);
};
if task.deleted_at.is_none() {
return Ok(false);
}
task.deleted_at = None;
task.updated_at = Some(Utc::now());
self.write_task(&task).await?;
Ok(true)
}
async fn update_status(
&self,
id: &TaskId,
status: TaskStatus,
result: Option<serde_json::Value>,
error: Option<String>,
) -> Result<Task, PeError> {
let mut task = self
.fetch_task(id, false)
.await?
.ok_or_else(|| PeError::NodeNotFound { node: id.clone() })?;
self.lifecycle.validate_transition(&task, &status)?;
let old_status = task.status.clone();
task.status = status.clone();
task.updated_at = Some(Utc::now());
if let Some(result) = result {
task.result = Some(result);
}
if let Some(error) = error {
task.error = Some(error);
}
if status == TaskStatus::Completed {
task.completed_at = Some(Utc::now());
}
let updated = self.write_task(&task).await?;
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 limit = if filter.limit == 0 { 100 } else { filter.limit };
let tasks = self.fetch_all_tasks(filter.include_deleted).await?;
Ok(tasks
.into_iter()
.filter(|task| filter.include_deleted || task.deleted_at.is_none())
.filter(|task| filter.status.as_ref().is_none_or(|s| task.status == *s))
.filter(|task| {
filter
.task_type
.as_ref()
.is_none_or(|task_type| task.task_type == *task_type)
})
.filter(|task| {
filter
.priority
.as_ref()
.is_none_or(|priority| task.priority == *priority)
})
.filter(|task| {
filter
.agent_id
.as_ref()
.is_none_or(|agent| task.agent_id.as_ref() == Some(agent))
})
.filter(|task| {
filter
.assignee
.as_ref()
.is_none_or(|assignee| task.assignee == *assignee)
})
.filter(|task| {
filter
.parent_id
.as_ref()
.is_none_or(|parent| task.parent_task_id.as_ref() == Some(parent))
})
.filter(|task| {
filter
.tag
.as_ref()
.is_none_or(|tag| task.tags.contains(tag))
})
.take(limit)
.collect())
}
async fn get_subtasks(&self, parent_id: &TaskId) -> Result<Vec<Task>, PeError> {
Ok(self
.fetch_all_tasks(false)
.await?
.into_iter()
.filter(|task| task.parent_task_id.as_ref() == Some(parent_id))
.collect())
}
async fn get_tree(&self, root_id: &TaskId) -> Result<Vec<Task>, PeError> {
let tasks = self.fetch_all_tasks(false).await?;
Ok(Self::collect_tree(root_id, &tasks))
}
async fn add_dependency(&self, dep: TaskDependency) -> Result<(), PeError> {
let tasks = self.fetch_all_tasks(false).await?;
let deps = self.fetch_all_dependencies().await?;
dependency::validate_dependency_endpoints(&dep, &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
),
});
}
self.apply_scope().await?;
self.db
.query(
"CREATE task_dependencies SET \
task_id = $task_id, \
depends_on_id = $depends_on_id, \
dependency_type = $dependency_type, \
dependency = $dependency, \
created_at = $created_at",
)
.bind(DependencyWriteParams::from_dependency(&dep)?)
.await
.map_err(|error| PeError::Storage {
details: format!("Task dependency create failed: {error}"),
})?;
Ok(())
}
async fn remove_dependency(
&self,
task_id: &TaskId,
depends_on_id: &TaskId,
) -> Result<(), PeError> {
self.apply_scope().await?;
self.db
.query("DELETE FROM task_dependencies WHERE task_id = $task_id AND depends_on_id = $depends_on_id")
.bind(DependencyPairParams {
task_id: task_id.clone(),
depends_on_id: depends_on_id.clone(),
})
.await
.map_err(|error| PeError::Storage {
details: format!("Task dependency delete failed: {error}"),
})?;
Ok(())
}
async fn get_dependencies(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
Ok(self
.fetch_all_dependencies()
.await?
.into_iter()
.filter(|dep| dep.task_id == *task_id)
.collect())
}
async fn get_dependents(&self, task_id: &TaskId) -> Result<Vec<TaskDependency>, PeError> {
Ok(self
.fetch_all_dependencies()
.await?
.into_iter()
.filter(|dep| dep.depends_on_id == *task_id)
.collect())
}
async fn get_ready_tasks(&self) -> Result<Vec<Task>, PeError> {
let tasks = self.fetch_all_tasks(false).await?;
let deps = self.fetch_all_dependencies().await?;
let pending = tasks
.iter()
.filter(|task| task.status == TaskStatus::Pending)
.map(|task| task.id.clone())
.collect::<Vec<_>>();
let statuses: HashMap<TaskId, TaskStatus> = tasks
.iter()
.map(|task| (task.id.clone(), task.status.clone()))
.collect();
let ready_ids = dependency::find_ready_tasks(&pending, &deps, &statuses);
Ok(ready_ids
.iter()
.filter_map(|id| tasks.iter().find(|task| task.id == *id).cloned())
.collect())
}
}
#[derive(Debug, SurrealValue)]
struct TaskRow {
task: serde_json::Value,
}
#[derive(Debug, SurrealValue)]
struct DependencyRow {
dependency: serde_json::Value,
}
#[derive(SurrealValue)]
struct TaskIdParams {
task_id: String,
}
#[derive(SurrealValue)]
struct DependencyPairParams {
task_id: String,
depends_on_id: String,
}
#[derive(SurrealValue)]
struct TaskWriteParams {
task_id: String,
task: serde_json::Value,
status: String,
task_type: String,
priority: String,
agent_id: Option<String>,
assignee: String,
parent_task_id: Option<String>,
tags: Vec<String>,
is_deleted: bool,
created_at: String,
updated_at: Option<String>,
completed_at: Option<String>,
deleted_at: Option<String>,
}
impl TaskWriteParams {
fn from_task(task: &Task) -> Result<Self, PeError> {
Ok(Self {
task_id: task.id.clone(),
task: serde_json::to_value(task).map_err(|error| PeError::Storage {
details: format!("Task serialize failed: {error}"),
})?,
status: task_status_label(&task.status),
task_type: task_type_label(&task.task_type),
priority: task_priority_label(&task.priority),
agent_id: task.agent_id.clone(),
assignee: task.assignee.clone(),
parent_task_id: task.parent_task_id.clone(),
tags: task.tags.clone(),
is_deleted: task.deleted_at.is_some(),
created_at: task.created_at.to_rfc3339(),
updated_at: task.updated_at.map(|dt| dt.to_rfc3339()),
completed_at: task.completed_at.map(|dt| dt.to_rfc3339()),
deleted_at: task.deleted_at.map(|dt| dt.to_rfc3339()),
})
}
}
#[derive(SurrealValue)]
struct DependencyWriteParams {
task_id: String,
depends_on_id: String,
dependency_type: String,
dependency: serde_json::Value,
created_at: String,
}
impl DependencyWriteParams {
fn from_dependency(dep: &TaskDependency) -> Result<Self, PeError> {
Ok(Self {
task_id: dep.task_id.clone(),
depends_on_id: dep.depends_on_id.clone(),
dependency_type: dependency_type_label(&dep.dependency_type),
dependency: serde_json::to_value(dep).map_err(|error| PeError::Storage {
details: format!("Task dependency serialize failed: {error}"),
})?,
created_at: dep.created_at.to_rfc3339(),
})
}
}
fn decode_task(value: serde_json::Value) -> Result<Task, PeError> {
serde_json::from_value(value).map_err(|error| PeError::Storage {
details: format!("Task deserialize failed: {error}"),
})
}
fn decode_dependency(value: serde_json::Value) -> Result<TaskDependency, PeError> {
serde_json::from_value(value).map_err(|error| PeError::Storage {
details: format!("Task dependency deserialize failed: {error}"),
})
}
fn task_status_label(status: &TaskStatus) -> String {
match status {
TaskStatus::Pending => "Pending",
TaskStatus::InProgress => "InProgress",
TaskStatus::Completed => "Completed",
TaskStatus::Failed => "Failed",
TaskStatus::Blocked => "Blocked",
TaskStatus::Cancelled => "Cancelled",
}
.into()
}
fn task_type_label(task_type: &TaskType) -> String {
match task_type {
TaskType::System => "System",
TaskType::Agent => "Agent",
TaskType::Human => "Human",
TaskType::Plan => "Plan",
}
.into()
}
fn task_priority_label(priority: &TaskPriority) -> String {
match priority {
TaskPriority::Urgent => "Urgent",
TaskPriority::High => "High",
TaskPriority::Medium => "Medium",
TaskPriority::Low => "Low",
}
.into()
}
fn dependency_type_label(dependency_type: &DependencyType) -> String {
match dependency_type {
DependencyType::Blocks => "Blocks",
DependencyType::Requires => "Requires",
DependencyType::Related => "Related",
}
.into()
}
#[cfg(test)]
mod tests {
use super::*;
fn scope(tenant: &str) -> ExecutionScope {
ExecutionScope::new(tenant, "user-1", "session-1", "thread-1")
}
#[tokio::test]
async fn creates_lists_and_completes_tasks() {
let registry = SurrealTaskRegistry::connect_memory(&scope("tasks-crud"))
.await
.unwrap();
let task = registry
.create(Task::agent_task("Persist task", "agent-1"))
.await
.unwrap();
let listed = registry
.list(&TaskFilter::default().with_agent("agent-1"))
.await
.unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].title, "Persist task");
registry
.update_status(&task.id, TaskStatus::InProgress, None, None)
.await
.unwrap();
let completed = registry
.update_status(
&task.id,
TaskStatus::Completed,
Some(serde_json::json!({ "ok": true })),
None,
)
.await
.unwrap();
assert_eq!(completed.status, TaskStatus::Completed);
assert!(completed.completed_at.is_some());
}
#[tokio::test]
async fn preserves_tasks_across_registry_instances_on_same_database() {
let db = surrealdb::engine::any::connect("mem://").await.unwrap();
let scope = scope("tasks-same-db");
let first = SurrealTaskRegistry::from_existing(db.clone(), &scope)
.await
.unwrap();
let created = first
.create(Task::new("Stored in SurrealDB"))
.await
.unwrap();
drop(first);
let second = SurrealTaskRegistry::from_existing(db, &scope)
.await
.unwrap();
let fetched = second.get(&created.id).await.unwrap().unwrap();
assert_eq!(fetched.title, "Stored in SurrealDB");
}
#[tokio::test]
async fn isolates_tasks_by_execution_scope_namespace() {
let db = surrealdb::engine::any::connect("mem://").await.unwrap();
let tenant_a = SurrealTaskRegistry::from_existing(db.clone(), &scope("tenant-a"))
.await
.unwrap();
let tenant_b = SurrealTaskRegistry::from_existing(db, &scope("tenant-b"))
.await
.unwrap();
tenant_a
.create(Task::new("Only tenant A can see this"))
.await
.unwrap();
assert_eq!(
tenant_a.list(&TaskFilter::default()).await.unwrap().len(),
1
);
assert_eq!(
tenant_b.list(&TaskFilter::default()).await.unwrap().len(),
0
);
}
#[tokio::test]
async fn tracks_dependencies_and_ready_tasks() {
let registry = SurrealTaskRegistry::connect_memory(&scope("tasks-ready"))
.await
.unwrap();
let a = registry.create(Task::new("A")).await.unwrap();
let b = registry.create(Task::new("B")).await.unwrap();
let c = registry.create(Task::new("C")).await.unwrap();
registry
.add_dependency(TaskDependency::new(&c.id, &a.id, DependencyType::Blocks))
.await
.unwrap();
let ready = registry.get_ready_tasks().await.unwrap();
let ready_ids = ready
.iter()
.map(|task| task.id.as_str())
.collect::<Vec<_>>();
assert!(ready_ids.contains(&a.id.as_str()));
assert!(ready_ids.contains(&b.id.as_str()));
assert!(!ready_ids.contains(&c.id.as_str()));
registry
.update_status(&a.id, TaskStatus::InProgress, None, None)
.await
.unwrap();
registry
.update_status(&a.id, TaskStatus::Completed, None, None)
.await
.unwrap();
let ready = registry.get_ready_tasks().await.unwrap();
let ready_ids = ready
.iter()
.map(|task| task.id.as_str())
.collect::<Vec<_>>();
assert!(ready_ids.contains(&c.id.as_str()));
}
#[tokio::test]
async fn validates_dependency_endpoints_and_duplicates() {
let registry = SurrealTaskRegistry::connect_memory(&scope("tasks-dep-validation"))
.await
.unwrap();
let blocker = registry.create(Task::new("Blocker")).await.unwrap();
let blocked = registry.create(Task::new("Blocked")).await.unwrap();
let err = registry
.add_dependency(TaskDependency::new(
&blocked.id,
"missing",
DependencyType::Blocks,
))
.await
.unwrap_err();
assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing"));
registry
.add_dependency(TaskDependency::new(
&blocked.id,
&blocker.id,
DependencyType::Blocks,
))
.await
.unwrap();
let err = registry
.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 validates_parent_links() {
let registry = SurrealTaskRegistry::connect_memory(&scope("tasks-parent-validation"))
.await
.unwrap();
let err = registry
.create(Task::new("Orphan").with_parent("missing-parent"))
.await
.unwrap_err();
assert!(matches!(err, PeError::NodeNotFound { node } if node == "missing-parent"));
let parent = registry.create(Task::new("Parent")).await.unwrap();
let child = registry
.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 = registry.update(&updated_parent).await.unwrap_err();
assert!(matches!(err, PeError::InvalidUpdate { .. }));
let stored_parent = registry.get(&parent.id).await.unwrap().unwrap();
assert!(stored_parent.parent_task_id.is_none());
}
#[tokio::test]
async fn soft_deletes_and_restores_tasks() {
let registry = SurrealTaskRegistry::connect_memory(&scope("tasks-delete"))
.await
.unwrap();
let task = registry.create(Task::new("Temporary")).await.unwrap();
assert!(registry.delete(&task.id).await.unwrap());
assert!(registry.get(&task.id).await.unwrap().is_none());
assert_eq!(
registry.list(&TaskFilter::default()).await.unwrap().len(),
0
);
assert_eq!(
registry
.list(&TaskFilter::default().including_deleted())
.await
.unwrap()
.len(),
1
);
assert!(registry.restore(&task.id).await.unwrap());
assert!(registry.get(&task.id).await.unwrap().is_some());
}
}