diff --git a/Cargo.lock b/Cargo.lock index 7a29e5a..fab1d45 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -582,6 +582,10 @@ dependencies = [ "which", ] +[[package]] +name = "rb-task" +version = "0.3.1" + [[package]] name = "rb-tests" version = "0.3.1" diff --git a/crates/rb-task/Cargo.toml b/crates/rb-task/Cargo.toml new file mode 100644 index 0000000..8dd9749 --- /dev/null +++ b/crates/rb-task/Cargo.toml @@ -0,0 +1,8 @@ +[package] +name = "rb-task" +version.workspace = true +edition.workspace = true +license.workspace = true +description = "Dependency graphs and bounded task execution" + +[dependencies] diff --git a/crates/rb-task/README.md b/crates/rb-task/README.md new file mode 100644 index 0000000..daf7969 --- /dev/null +++ b/crates/rb-task/README.md @@ -0,0 +1,45 @@ +# rb-task + +Dependency graphs, bounded workers, and task events using only the standard library. + +```rust +use rb_task::{TaskEvents, Executor, TaskAction, TaskGraph}; +use std::{collections::BTreeMap, sync::Arc}; + +let mut graph = TaskGraph::new(); +let prepare = graph.add("prepare", []); +let finish = graph.add("finish", [prepare]); +let actions = BTreeMap::from([ + (prepare, Arc::new(|context: rb_task::TaskContext| { + context.output("preparing input"); + Ok(()) + }) as TaskAction), + (finish, Arc::new(|_| Ok(())) as TaskAction), +]); +Executor::new(2, TaskEvents::default()).run(graph, actions)?; +# Ok::<(), rb_task::ExecutionError>(()) +``` + +Subscribe a `TaskEventSink` to the `TaskEvents` to receive lifecycle, progress, +output, and duration events. Callbacks run synchronously on workers and may +run concurrently; slow callbacks delay their worker, and panics fail the run. +The crate does not format output. + +Graphs are fixed per run. Task IDs are graph-local indexes; callers must not mix +IDs from different graphs. Ready tasks run within +the worker limit as dependencies finish. On an observed action failure or panic, +or an observer panic, execution stops dispatching and waits for running work. +Cancellation, retries, and graph expansion belong to the caller. + +Run `cargo test -p rb-task` for the example and regression suite. + +Runnable examples: + +- [Dependencies](examples/dependencies.rs): two tasks run independently between + shared preparation and completion steps. +- [Reporting](examples/reporting.rs): observe lifecycle, output, and progress events. + +```sh +cargo run -p rb-task --example dependencies +cargo run -p rb-task --example reporting +``` diff --git a/crates/rb-task/examples/dependencies.rs b/crates/rb-task/examples/dependencies.rs new file mode 100644 index 0000000..dbb1de3 --- /dev/null +++ b/crates/rb-task/examples/dependencies.rs @@ -0,0 +1,28 @@ +use rb_task::{ExecutionError, Executor, TaskAction, TaskContext, TaskEvents, TaskGraph}; +use std::{collections::BTreeMap, sync::Arc, thread, time::Duration}; + +fn main() -> Result<(), ExecutionError> { + let mut graph = TaskGraph::new(); + let prepare = graph.add("prepare", []); + let left = graph.add("process left", [prepare]); + let right = graph.add("process right", [prepare]); + let finish = graph.add("finish", [left, right]); + + let actions = BTreeMap::from([ + (prepare, action("prepare")), + (left, action("process left")), + (right, action("process right")), + (finish, action("finish")), + ]); + + Executor::new(2, TaskEvents::default()).run(graph, actions) +} + +fn action(label: &'static str) -> TaskAction { + Arc::new(move |context: TaskContext| { + println!("worker {}: starting {label}", context.worker); + thread::sleep(Duration::from_millis(50)); + println!("worker {}: finished {label}", context.worker); + Ok(()) + }) +} diff --git a/crates/rb-task/examples/reporting.rs b/crates/rb-task/examples/reporting.rs new file mode 100644 index 0000000..afc5b64 --- /dev/null +++ b/crates/rb-task/examples/reporting.rs @@ -0,0 +1,45 @@ +use rb_task::{ + ExecutionError, Executor, TaskAction, TaskContext, TaskEvent, TaskEventSink, TaskEvents, + TaskGraph, +}; +use std::{collections::BTreeMap, sync::Arc}; + +struct Console; + +impl TaskEventSink for Console { + fn event(&self, event: TaskEvent) { + match event { + TaskEvent::Started { worker, label, .. } => { + println!("worker {worker}: {label}"); + } + TaskEvent::Output { line, .. } => println!(" {line}"), + TaskEvent::Progress { + phase, done, total, .. + } => { + println!(" {phase}: {done}/{total}"); + } + TaskEvent::Finished { elapsed, .. } => println!(" finished in {elapsed:?}"), + TaskEvent::Failed { error, .. } => println!(" failed: {error}"), + } + } +} + +fn main() -> Result<(), ExecutionError> { + let events = TaskEvents::default(); + events.subscribe(Arc::new(Console)); + + let mut graph = TaskGraph::new(); + let process = graph.add("process documents", []); + let actions = BTreeMap::from([( + process, + Arc::new(|context: TaskContext| { + for done in 1..=3 { + context.output(format!("Processed document {done}")); + context.progress("processing", done, 3, None); + } + Ok(()) + }) as TaskAction, + )]); + + Executor::new(1, events).run(graph, actions) +} diff --git a/crates/rb-task/src/error.rs b/crates/rb-task/src/error.rs new file mode 100644 index 0000000..a8af1d5 --- /dev/null +++ b/crates/rb-task/src/error.rs @@ -0,0 +1,38 @@ +use crate::{CompletionError, PlanError, TaskId}; +use std::{error::Error, fmt}; + +/// Identifies a validation, task, or panic failure encountered during execution. +#[derive(Debug, Eq, PartialEq)] +pub enum ExecutionError { + InvalidActions, + InvalidGraph(PlanError), + InvalidCompletion(CompletionError), + TaskFailed { task: TaskId, message: String }, + Panicked { task: TaskId }, +} + +impl fmt::Display for ExecutionError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::InvalidActions => formatter.write_str("missing or invalid task action"), + Self::InvalidGraph(error) => write!(formatter, "invalid task graph: {error}"), + Self::InvalidCompletion(error) => write!(formatter, "invalid task completion: {error}"), + Self::TaskFailed { task, message } => { + write!(formatter, "task {} failed: {message}", task.index()) + } + Self::Panicked { task } => { + write!(formatter, "task {} or its observer panicked", task.index()) + } + } + } +} + +impl Error for ExecutionError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + Self::InvalidGraph(error) => Some(error), + Self::InvalidCompletion(error) => Some(error), + _ => None, + } + } +} diff --git a/crates/rb-task/src/events.rs b/crates/rb-task/src/events.rs new file mode 100644 index 0000000..0fede64 --- /dev/null +++ b/crates/rb-task/src/events.rs @@ -0,0 +1,67 @@ +use crate::TaskId; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +/// Reports task lifecycle changes, output, and progress to observers. +#[derive(Clone, Debug, Eq, PartialEq)] +pub enum TaskEvent { + Started { + task: TaskId, + worker: usize, + label: String, + elapsed: Duration, + }, + Output { + task: TaskId, + worker: usize, + line: String, + elapsed: Duration, + }, + Progress { + task: TaskId, + worker: usize, + phase: &'static str, + done: usize, + total: usize, + detail: Option, + elapsed: Duration, + }, + Finished { + task: TaskId, + worker: usize, + elapsed: Duration, + }, + Failed { + task: TaskId, + worker: usize, + error: String, + elapsed: Duration, + }, +} + +/// Receives events synchronously on workers; callbacks may run concurrently. +pub trait TaskEventSink: Send + Sync { + fn event(&self, event: TaskEvent); +} + +/// Shares subscriptions and delivers each event to registered observers. +#[derive(Default, Clone)] +pub struct TaskEvents { + sinks: Arc>>>, +} + +impl TaskEvents { + pub fn subscribe(&self, sink: Arc) { + self.sinks.lock().unwrap().push(sink); + } + + pub fn emit(&self, event: TaskEvent) { + let sinks = self.sinks.lock().unwrap().clone(); + for sink in sinks { + sink.event(event.clone()); + } + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/rb-task/src/events/tests.rs b/crates/rb-task/src/events/tests.rs new file mode 100644 index 0000000..90037f1 --- /dev/null +++ b/crates/rb-task/src/events/tests.rs @@ -0,0 +1,59 @@ +use super::*; +use std::sync::{ + atomic::{AtomicUsize, Ordering}, + mpsc, +}; +use std::thread; + +struct CountingSink(AtomicUsize); + +impl TaskEventSink for CountingSink { + fn event(&self, _event: TaskEvent) { + self.0.fetch_add(1, Ordering::Relaxed); + } +} + +#[test] +fn event_bus_fans_out_without_owning_execution() { + let bus = TaskEvents::default(); + let sink = Arc::new(CountingSink(AtomicUsize::new(0))); + let other = Arc::new(CountingSink(AtomicUsize::new(0))); + bus.subscribe(sink.clone()); + bus.subscribe(other.clone()); + bus.emit(TaskEvent::Finished { + task: TaskId(0), + worker: 0, + elapsed: Duration::ZERO, + }); + assert_eq!(sink.0.load(Ordering::Relaxed), 1); + assert_eq!(other.0.load(Ordering::Relaxed), 1); +} + +#[test] +fn event_callbacks_can_subscribe_without_locking_the_bus() { + struct Subscriber(TaskEvents); + struct Ignore; + impl TaskEventSink for Ignore { + fn event(&self, _: TaskEvent) {} + } + impl TaskEventSink for Subscriber { + fn event(&self, _: TaskEvent) { + self.0.subscribe(Arc::new(Ignore)); + } + } + let bus = TaskEvents::default(); + bus.subscribe(Arc::new(Subscriber(bus.clone()))); + let emitting = bus.clone(); + let (send, receive) = mpsc::channel(); + thread::spawn(move || { + emitting.emit(TaskEvent::Finished { + task: TaskId(0), + worker: 0, + elapsed: Duration::ZERO, + }); + send.send(()).unwrap(); + }); + receive.recv_timeout(Duration::from_secs(2)).unwrap(); + assert_eq!(bus.sinks.lock().unwrap().len(), 2); + bus.sinks.lock().unwrap().clear(); +} diff --git a/crates/rb-task/src/executor.rs b/crates/rb-task/src/executor.rs new file mode 100644 index 0000000..f6cc23c --- /dev/null +++ b/crates/rb-task/src/executor.rs @@ -0,0 +1,166 @@ +use crate::{ExecutionError, TaskEvent, TaskEvents, TaskGraph, TaskId}; +use std::collections::{BTreeMap, VecDeque}; +use std::sync::Arc; +use std::thread; +use std::time::Duration; +use std::time::Instant; + +/// Gives an action its task and worker identity and methods to report progress. +pub struct TaskContext { + pub task: TaskId, + pub worker: usize, + pub events: TaskEvents, + started: Instant, +} + +impl TaskContext { + pub fn output(&self, line: impl Into) { + self.events.emit(TaskEvent::Output { + task: self.task, + worker: self.worker, + line: line.into(), + elapsed: self.started.elapsed(), + }); + } + + pub fn progress(&self, phase: &'static str, done: usize, total: usize, detail: Option) { + self.events.emit(TaskEvent::Progress { + task: self.task, + worker: self.worker, + phase, + done, + total, + detail, + elapsed: self.started.elapsed(), + }); + } +} + +/// Holds the caller's executable work, returning an error message on failure. +pub type TaskAction = Arc Result<(), String> + Send + Sync + 'static>; + +/// Runs ready tasks within a worker limit and waits for active work before returning. +pub struct Executor { + workers: usize, + events: TaskEvents, +} + +impl Executor { + pub fn new(workers: usize, events: TaskEvents) -> Self { + Self { + workers: workers.max(1), + events, + } + } + + pub fn run( + &self, + graph: TaskGraph, + actions: BTreeMap, + ) -> Result<(), ExecutionError> { + if actions.len() != graph.len() || actions.keys().any(|id| graph.task(*id).is_none()) { + return Err(ExecutionError::InvalidActions); + } + let mut schedule = graph.schedule().map_err(ExecutionError::InvalidGraph)?; + let count = self.workers.min(graph.len()); + thread::scope(|scope| { + let (completed, completions) = std::sync::mpsc::channel(); + let mut senders = Vec::new(); + for worker in 0..count { + let (sender, receiver) = std::sync::mpsc::channel::<(TaskId, String, TaskAction)>(); + senders.push(sender); + let completed = completed.clone(); + let events = self.events.clone(); + scope.spawn(move || { + while let Ok((task, label, action)) = receiver.recv() { + let started = Instant::now(); + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + events.emit(TaskEvent::Started { + task, + worker, + label, + elapsed: Duration::ZERO, + }); + action(TaskContext { + task, + worker, + events: events.clone(), + started, + }) + })) + .map_err(|_| ExecutionError::Panicked { task }) + .and_then(|result| { + result.map_err(|message| ExecutionError::TaskFailed { task, message }) + }); + let terminal = match &result { + Ok(()) => TaskEvent::Finished { + task, + worker, + elapsed: started.elapsed(), + }, + Err(error) => TaskEvent::Failed { + task, + worker, + error: error.to_string(), + elapsed: started.elapsed(), + }, + }; + let notification = + std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + events.emit(terminal) + })); + let result = if notification.is_err() { + Err(ExecutionError::Panicked { task }) + } else { + result + }; + if completed.send((task, worker, result)).is_err() { + break; + } + } + }); + } + drop(completed); + let mut available: VecDeque<_> = (0..count).collect(); + let mut active = 0; + let mut failure = None; + loop { + while failure.is_none() && !available.is_empty() { + let Some(task) = schedule.take_ready() else { + break; + }; + let worker = available.pop_front().unwrap(); + senders[worker] + .send(( + task, + schedule.task(task).name.clone(), + actions[&task].clone(), + )) + .unwrap(); + active += 1; + } + if active == 0 { + break; + } + let (task, worker, result) = completions.recv().unwrap(); + active -= 1; + available.push_back(worker); + match result { + Ok(()) => { + if let Err(error) = schedule.complete(task) { + failure.get_or_insert(ExecutionError::InvalidCompletion(error)); + } + } + Err(error) => { + failure.get_or_insert(error); + } + } + } + drop(senders); + failure.map_or(Ok(()), Err) + }) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/rb-task/src/executor/tests.rs b/crates/rb-task/src/executor/tests.rs new file mode 100644 index 0000000..3383cfb --- /dev/null +++ b/crates/rb-task/src/executor/tests.rs @@ -0,0 +1,377 @@ +use super::*; +use crate::TaskEventSink; +use std::collections::BTreeSet; +use std::sync::{ + Condvar, Mutex, + atomic::{AtomicBool, AtomicUsize, Ordering}, + mpsc, +}; + +#[derive(Default)] +struct Rendezvous { + arrived: Mutex, + changed: Condvar, +} + +impl Rendezvous { + fn wait(&self) { + let mut arrived = self.arrived.lock().unwrap(); + *arrived += 1; + self.changed.notify_all(); + let (arrived, _) = self + .changed + .wait_timeout_while(arrived, Duration::from_secs(2), |count| *count < 2) + .unwrap(); + assert_eq!(*arrived, 2, "two actions must overlap"); + } +} + +struct FinishedTimingSink(Mutex>); + +impl TaskEventSink for FinishedTimingSink { + fn event(&self, event: TaskEvent) { + if let TaskEvent::Finished { elapsed, .. } = event { + *self.0.lock().unwrap() = Some(elapsed); + } + } +} + +#[test] +fn finished_task_reports_elapsed_time() { + let mut graph = TaskGraph::new(); + let task = graph.add("timed", []); + let events = TaskEvents::default(); + let timing = Arc::new(FinishedTimingSink(Mutex::new(None))); + events.subscribe(timing.clone()); + let actions = BTreeMap::from([( + task, + Arc::new(|_| { + thread::sleep(std::time::Duration::from_millis(2)); + Ok(()) + }) as TaskAction, + )]); + Executor::new(1, events).run(graph, actions).unwrap(); + assert!(timing.0.lock().unwrap().unwrap() >= Duration::from_millis(2)); +} + +#[test] +fn executor_runs_independent_tasks_in_parallel() { + let mut graph = TaskGraph::new(); + let first = graph.add("first", []); + let second = graph.add("second", []); + let final_task = graph.add("final", [first, second]); + let rendezvous = Arc::new(Rendezvous::default()); + let counts = Arc::new([ + AtomicUsize::new(0), + AtomicUsize::new(0), + AtomicUsize::new(0), + ]); + let mut actions = BTreeMap::new(); + for task in [first, second] { + let rendezvous = rendezvous.clone(); + let counts = counts.clone(); + actions.insert( + task, + Arc::new(move |context: TaskContext| { + assert_eq!(context.task, task); + rendezvous.wait(); + assert_eq!(counts[task.index()].fetch_add(1, Ordering::SeqCst), 0); + Ok(()) + }) as TaskAction, + ); + } + let finished = counts.clone(); + actions.insert( + final_task, + Arc::new(move |_| { + assert_eq!(finished[first.index()].load(Ordering::SeqCst), 1); + assert_eq!(finished[second.index()].load(Ordering::SeqCst), 1); + assert_eq!( + finished[final_task.index()].fetch_add(1, Ordering::SeqCst), + 0 + ); + Ok(()) + }) as TaskAction, + ); + Executor::new(2, TaskEvents::default()) + .run(graph, actions) + .unwrap(); + assert!(counts.iter().all(|count| count.load(Ordering::SeqCst) == 1)); +} + +#[test] +fn limits_active_tasks_and_never_shares_a_busy_worker() { + let active = Arc::new(AtomicUsize::new(0)); + let completed = Arc::new(AtomicUsize::new(0)); + let busy = Arc::new(Mutex::new(BTreeSet::new())); + let mut graph = TaskGraph::new(); + let mut actions = BTreeMap::new(); + for _ in 0..24 { + let task = graph.add("work", []); + let (active, completed, busy) = (active.clone(), completed.clone(), busy.clone()); + actions.insert( + task, + Arc::new(move |context: TaskContext| { + assert!(busy.lock().unwrap().insert(context.worker)); + let count = active.fetch_add(1, Ordering::SeqCst) + 1; + assert!(count <= 2); + assert!(context.worker < 2); + thread::yield_now(); + active.fetch_sub(1, Ordering::SeqCst); + assert!(busy.lock().unwrap().remove(&context.worker)); + completed.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) as TaskAction, + ); + } + Executor::new(2, TaskEvents::default()) + .run(graph, actions) + .unwrap(); + assert_eq!(completed.load(Ordering::SeqCst), 24); + assert!(busy.lock().unwrap().is_empty()); + assert_eq!(active.load(Ordering::SeqCst), 0); +} + +#[test] +fn releases_dependents_before_unrelated_running_tasks_finish() { + let mut graph = TaskGraph::new(); + let slow = graph.add("waiting for dependent", []); + let fast = graph.add("prerequisite", []); + let dependent = graph.add("dependent", [fast]); + let ran = Arc::new(AtomicBool::new(false)); + let marker = ran.clone(); + let (send, receive) = mpsc::channel(); + let receive = Mutex::new(receive); + let actions = BTreeMap::from([ + ( + slow, + Arc::new(move |_| { + receive + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(2)) + .map_err(|e| e.to_string()) + }) as TaskAction, + ), + (fast, Arc::new(|_| Ok(())) as TaskAction), + ( + dependent, + Arc::new(move |_| { + marker.store(true, Ordering::SeqCst); + send.send(()).map_err(|e| e.to_string()) + }) as TaskAction, + ), + ]); + Executor::new(2, TaskEvents::default()) + .run(graph, actions) + .unwrap(); + assert!(ran.load(Ordering::SeqCst)); +} + +#[test] +fn failures_and_panics_join_running_work_and_block_dependents() { + for panic in [false, true] { + let mut graph = TaskGraph::new(); + let failed = graph.add("failure", []); + let running = graph.add("running", []); + let blocked = graph.add("blocked", [failed]); + let rendezvous = Arc::new(Rendezvous::default()); + let done = Arc::new(AtomicBool::new(false)); + let blocked_ran = Arc::new(AtomicBool::new(false)); + let blocked_marker = blocked_ran.clone(); + let marker = done.clone(); + let other = rendezvous.clone(); + let actions = BTreeMap::from([ + ( + failed, + Arc::new(move |_| { + rendezvous.wait(); + if panic { + panic!("fixture"); + } + Err("fixture".into()) + }) as TaskAction, + ), + ( + running, + Arc::new(move |_| { + other.wait(); + thread::sleep(Duration::from_millis(30)); + marker.store(true, Ordering::SeqCst); + Ok(()) + }) as TaskAction, + ), + ( + blocked, + Arc::new(move |_| { + blocked_marker.store(true, Ordering::SeqCst); + Ok(()) + }) as TaskAction, + ), + ]); + let error = Executor::new(2, TaskEvents::default()) + .run(graph, actions) + .unwrap_err(); + assert_eq!( + error, + if panic { + ExecutionError::Panicked { task: failed } + } else { + ExecutionError::TaskFailed { + task: failed, + message: "fixture".into(), + } + } + ); + assert!(done.load(Ordering::SeqCst)); + assert!(!blocked_ran.load(Ordering::SeqCst)); + } +} + +#[test] +fn rejects_incorrect_action_ids_and_accepts_empty_graphs() { + let mut graph = TaskGraph::new(); + graph.add("missing", []); + let actions = BTreeMap::from([(TaskId(99), Arc::new(|_| Ok(())) as TaskAction)]); + assert_eq!( + Executor::new(1, TaskEvents::default()).run(graph, actions), + Err(ExecutionError::InvalidActions) + ); + Executor::new(0, TaskEvents::default()) + .run(TaskGraph::new(), BTreeMap::new()) + .unwrap(); +} + +#[test] +fn varied_dags_execute_every_task_once_after_its_dependencies() { + for workers in [1, 2, 4, 8] { + for seed in 0..12 { + let mut graph = TaskGraph::new(); + let mut actions = BTreeMap::new(); + let counts = Arc::new((0..32).map(|_| AtomicUsize::new(0)).collect::>()); + for index in 0..32 { + let deps = (0..index) + .filter(|dep| (dep * 7 + index * 11 + seed) % 9 == 0) + .map(TaskId) + .collect::>(); + let task = graph.add(format!("work {index}"), deps.clone()); + let counts = counts.clone(); + actions.insert( + task, + Arc::new(move |context: TaskContext| { + assert!(context.worker < workers); + for dep in &deps { + assert_eq!(counts[dep.index()].load(Ordering::SeqCst), 1); + } + context.output("processing"); + context.progress("work", 1, 1, None); + thread::yield_now(); + assert_eq!(counts[index].fetch_add(1, Ordering::SeqCst), 0); + Ok(()) + }) as TaskAction, + ); + } + Executor::new(workers, TaskEvents::default()) + .run(graph, actions) + .unwrap(); + assert!(counts.iter().all(|count| count.load(Ordering::SeqCst) == 1)); + } + } +} + +#[test] +fn observer_panics_return_an_error_without_running_dependents() { + struct PanickingSink(usize); + impl TaskEventSink for PanickingSink { + fn event(&self, event: TaskEvent) { + let kind = match event { + TaskEvent::Started { .. } => 0, + TaskEvent::Progress { .. } => 1, + TaskEvent::Output { .. } => 2, + TaskEvent::Finished { .. } => 3, + TaskEvent::Failed { .. } => 4, + }; + assert_ne!(kind, self.0, "fixture observer panic"); + } + } + for kind in 0..5 { + let mut graph = TaskGraph::new(); + let task = graph.add("work", []); + let dependent = graph.add("dependent", [task]); + let ran = Arc::new(AtomicUsize::new(0)); + let marker = ran.clone(); + let events = TaskEvents::default(); + events.subscribe(Arc::new(PanickingSink(kind))); + let actions = BTreeMap::from([ + ( + task, + Arc::new(move |context: TaskContext| { + context.progress("working", 0, 1, None); + context.output("output"); + if kind == 4 { + Err("fixture action failure".into()) + } else { + Ok(()) + } + }) as TaskAction, + ), + ( + dependent, + Arc::new(move |_| { + marker.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) as TaskAction, + ), + ]); + assert_eq!( + Executor::new(2, events).run(graph, actions), + Err(ExecutionError::Panicked { task }) + ); + assert_eq!(ran.load(Ordering::SeqCst), 0); + } +} + +#[test] +fn invalid_graph_errors_preserve_their_source() { + use std::error::Error; + + for (dependency, expected) in [ + (TaskId(0), crate::PlanError::Cycle), + ( + TaskId(99), + crate::PlanError::UnknownDependency(TaskId(0), TaskId(99)), + ), + ] { + let mut graph = TaskGraph::new(); + let task = graph.add("invalid", [dependency]); + let actions = BTreeMap::from([( + task, + Arc::new(|_| -> Result<(), String> { panic!("invalid graph executed") }) as TaskAction, + )]); + let error = Executor::new(1, TaskEvents::default()) + .run(graph, actions) + .unwrap_err(); + assert_eq!(error.source().unwrap().to_string(), expected.to_string()); + assert_eq!(error, ExecutionError::InvalidGraph(expected)); + } +} + +#[test] +fn zero_workers_still_executes_work() { + let mut graph = TaskGraph::new(); + let task = graph.add("work", []); + let ran = Arc::new(AtomicUsize::new(0)); + let marker = ran.clone(); + let actions = BTreeMap::from([( + task, + Arc::new(move |context: TaskContext| { + assert_eq!(context.worker, 0); + marker.fetch_add(1, Ordering::SeqCst); + Ok(()) + }) as TaskAction, + )]); + Executor::new(0, TaskEvents::default()) + .run(graph, actions) + .unwrap(); + assert_eq!(ran.load(Ordering::SeqCst), 1); +} diff --git a/crates/rb-task/src/graph.rs b/crates/rb-task/src/graph.rs new file mode 100644 index 0000000..56f7d2e --- /dev/null +++ b/crates/rb-task/src/graph.rs @@ -0,0 +1,188 @@ +use crate::TaskId; +use std::collections::{BTreeSet, VecDeque}; + +/// Describes a unit of work and the tasks that must finish before it can run. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct Task { + pub name: String, + pub deps: Vec, +} + +/// Collects tasks and dependencies for validation and scheduling. +#[derive(Clone, Debug, Default)] +pub struct TaskGraph { + tasks: Vec, +} + +impl TaskGraph { + pub fn new() -> Self { + Self::default() + } + + pub fn add( + &mut self, + name: impl Into, + deps: impl IntoIterator, + ) -> TaskId { + let id = TaskId(self.tasks.len()); + self.tasks.push(Task { + name: name.into(), + deps: deps.into_iter().collect(), + }); + id + } + + pub fn task(&self, id: TaskId) -> Option<&Task> { + self.tasks.get(id.0) + } + + pub fn len(&self) -> usize { + self.tasks.len() + } + + pub fn is_empty(&self) -> bool { + self.tasks.is_empty() + } + + pub fn schedule(&self) -> Result { + let mut indegree = vec![0usize; self.tasks.len()]; + let mut dependents = vec![Vec::new(); self.tasks.len()]; + for (index, task) in self.tasks.iter().enumerate() { + for &dependency in &task.deps { + if dependency.0 >= self.tasks.len() { + return Err(PlanError::UnknownDependency(TaskId(index), dependency)); + } + indegree[index] += 1; + dependents[dependency.0].push(TaskId(index)); + } + } + let ready = indegree + .iter() + .enumerate() + .filter_map(|(index, count)| (*count == 0).then_some(TaskId(index))) + .collect(); + let schedule = Schedule { + graph: self.clone(), + indegree, + dependents, + ready, + completed: BTreeSet::new(), + running: BTreeSet::new(), + }; + if schedule.topological_count() != self.tasks.len() { + return Err(PlanError::Cycle); + } + Ok(schedule) + } +} + +/// Explains why a graph cannot produce a valid schedule. +#[derive(Debug, Eq, PartialEq)] +pub enum PlanError { + UnknownDependency(TaskId, TaskId), + Cycle, +} + +impl std::fmt::Display for PlanError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::UnknownDependency(task, dependency) => write!( + formatter, + "task {} has unknown dependency {}", + task.index(), + dependency.index() + ), + Self::Cycle => formatter.write_str("task graph contains a cycle"), + } + } +} + +impl std::error::Error for PlanError {} + +/// Explains why a task cannot be marked complete in a schedule. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum CompletionError { + UnknownTask(TaskId), + NotRunning(TaskId), +} + +impl std::fmt::Display for CompletionError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::UnknownTask(id) => write!(formatter, "unknown task {}", id.index()), + Self::NotRunning(id) => write!(formatter, "task {} is not running", id.index()), + } + } +} + +impl std::error::Error for CompletionError {} + +/// Tracks ready, running, and completed tasks in a validated dependency graph. +#[derive(Debug)] +pub struct Schedule { + graph: TaskGraph, + indegree: Vec, + dependents: Vec>, + ready: VecDeque, + completed: BTreeSet, + running: BTreeSet, +} + +impl Schedule { + pub fn take_ready(&mut self) -> Option { + let id = self.ready.pop_front()?; + self.running.insert(id); + Some(id) + } + + pub fn task(&self, id: TaskId) -> &Task { + &self.graph.tasks[id.0] + } + + pub fn complete(&mut self, id: TaskId) -> Result { + if self.graph.task(id).is_none() { + return Err(CompletionError::UnknownTask(id)); + } + if self.completed.contains(&id) { + return Ok(false); + } + if !self.running.remove(&id) { + return Err(CompletionError::NotRunning(id)); + } + self.completed.insert(id); + for &dependent in &self.dependents[id.0] { + self.indegree[dependent.0] -= 1; + if self.indegree[dependent.0] == 0 { + self.ready.push_back(dependent); + } + } + Ok(true) + } + + pub fn is_complete(&self) -> bool { + self.completed.len() == self.graph.len() + } + + fn topological_count(&self) -> usize { + let mut indegree = self.indegree.clone(); + let mut queue: VecDeque<_> = indegree + .iter() + .enumerate() + .filter_map(|(index, count)| (*count == 0).then_some(TaskId(index))) + .collect(); + let mut count = 0; + while let Some(id) = queue.pop_front() { + count += 1; + for &dependent in &self.dependents[id.0] { + indegree[dependent.0] -= 1; + if indegree[dependent.0] == 0 { + queue.push_back(dependent); + } + } + } + count + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/rb-task/src/graph/tests.rs b/crates/rb-task/src/graph/tests.rs new file mode 100644 index 0000000..ac7811d --- /dev/null +++ b/crates/rb-task/src/graph/tests.rs @@ -0,0 +1,64 @@ +use super::*; + +#[test] +fn releases_independent_tasks_and_then_dependents() { + let mut graph = TaskGraph::new(); + let resolve = graph.add("resolve", []); + let loader = graph.add("loader", []); + let pack = graph.add("pack", [resolve]); + let snapshot = graph.add("snapshot", [pack, loader]); + let mut schedule = graph.schedule().unwrap(); + assert_eq!(schedule.take_ready(), Some(resolve)); + assert_eq!(schedule.take_ready(), Some(loader)); + assert_eq!(schedule.take_ready(), None); + schedule.complete(resolve).unwrap(); + assert_eq!(schedule.take_ready(), Some(pack)); + schedule.complete(loader).unwrap(); + assert_eq!(schedule.take_ready(), None); + assert!(!schedule.is_complete()); + schedule.complete(pack).unwrap(); + assert_eq!(schedule.take_ready(), Some(snapshot)); + assert!(!schedule.is_complete()); + assert_eq!(schedule.complete(snapshot), Ok(true)); + assert!(schedule.is_complete()); + assert_eq!(schedule.take_ready(), None); +} + +#[test] +fn rejects_cycles() { + let mut graph = TaskGraph::new(); + let first = graph.add("first", [TaskId(1)]); + let second = graph.add("second", [first]); + assert_eq!(graph.schedule().unwrap_err(), PlanError::Cycle); + assert_eq!(second.index(), 1); +} + +#[test] +fn cannot_release_dependencies_by_completing_unscheduled_work() { + let mut graph = TaskGraph::new(); + let first = graph.add("first", []); + let second = graph.add("second", [first]); + let mut schedule = graph.schedule().unwrap(); + assert_eq!( + schedule.complete(first), + Err(CompletionError::NotRunning(first)) + ); + assert_eq!(schedule.take_ready(), Some(first)); + assert_eq!(schedule.take_ready(), None); + assert_eq!(schedule.complete(first), Ok(true)); + assert_eq!(schedule.complete(first), Ok(false)); + assert_eq!(schedule.take_ready(), Some(second)); +} + +#[test] +fn unknown_completion_does_not_panic_or_change_the_plan() { + let mut graph = TaskGraph::new(); + let task = graph.add("work", []); + let mut schedule = graph.schedule().unwrap(); + assert_eq!( + schedule.complete(TaskId(99)), + Err(CompletionError::UnknownTask(TaskId(99))) + ); + assert_eq!(schedule.take_ready(), Some(task)); + assert!(!schedule.is_complete()); +} diff --git a/crates/rb-task/src/lib.rs b/crates/rb-task/src/lib.rs new file mode 100644 index 0000000..8af6b84 --- /dev/null +++ b/crates/rb-task/src/lib.rs @@ -0,0 +1,21 @@ +#![doc = include_str!("../README.md")] + +mod error; +mod events; +mod executor; +mod graph; + +pub use error::ExecutionError; +pub use events::{TaskEvent, TaskEventSink, TaskEvents}; +pub use executor::{Executor, TaskAction, TaskContext}; +pub use graph::{CompletionError, PlanError, Schedule, Task, TaskGraph}; + +/// Identifies a task within its graph; IDs must not be mixed between graphs. +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct TaskId(usize); + +impl TaskId { + pub fn index(self) -> usize { + self.0 + } +}