use std::{
collections::HashMap,
num::NonZeroUsize,
sync::mpsc,
thread::{self, ScopedJoinHandle},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Resolution {
Running,
Completed,
Skipped,
Failed,
}
impl Resolution {
fn satisfies_dependents(self) -> bool {
matches!(self, Self::Running | Self::Completed)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Gate {
DependencySkipped,
DependencyFailed,
}
pub trait Units: Sync {
type Error: Send;
fn static_resolution(&self, service: &str)
-> Result<Option<Resolution>, Self::Error>;
fn start(
&self,
service: &str,
deps: &HashMap<String, Resolution>,
) -> Result<Resolution, Self::Error>;
fn gated(
&self,
service: &str,
dependency: &str,
gate: Gate,
) -> Result<Resolution, Self::Error>;
fn active(&self) -> bool;
}
pub struct Schedule<'a> {
order: &'a [String],
deps: HashMap<&'a str, Vec<&'a str>>,
}
impl<'a> Schedule<'a> {
pub fn new(order: &'a [String], edges: impl Fn(&str) -> Vec<&'a str>) -> Self {
let covered: Vec<&str> = order.iter().map(String::as_str).collect();
let deps = order
.iter()
.map(|service| {
let kept = edges(service)
.into_iter()
.filter(|dep| covered.contains(dep))
.collect();
(service.as_str(), kept)
})
.collect();
Self { order, deps }
}
pub fn run<U: Units>(
&self,
units: &U,
limit: Option<NonZeroUsize>,
) -> Result<HashMap<String, Resolution>, U::Error> {
let limit = limit.map_or(usize::MAX, NonZeroUsize::get);
let mut resolved: HashMap<String, Resolution> = HashMap::new();
let mut pending: Vec<usize> = (0..self.order.len()).collect();
let mut checked = vec![false; self.order.len()];
let mut fatal: Option<(usize, U::Error)> = None;
thread::scope(|scope| {
let (tx, rx) = mpsc::channel();
let mut workers: Vec<ScopedJoinHandle<'_, ()>> = Vec::new();
let mut inflight = 0usize;
loop {
let mut dispatched = false;
let mut index = 0;
while index < pending.len() {
if inflight >= limit {
break;
}
let rank = pending[index];
let service = &self.order[rank];
if !checked[rank] {
checked[rank] = true;
match units.static_resolution(service) {
Ok(Some(state)) => {
pending.remove(index);
resolved.insert(service.clone(), state);
dispatched = true;
continue;
}
Ok(None) => {}
Err(err) => {
pending.remove(index);
Self::keep_first(&mut fatal, rank, err);
pending.clear();
break;
}
}
}
match self.gate_of(service, &resolved) {
DepState::Waiting => {
index += 1;
continue;
}
DepState::Gated(dependency, gate) => {
pending.remove(index);
match units.gated(service, dependency, gate) {
Ok(state) => {
resolved.insert(service.clone(), state);
}
Err(err) => {
Self::keep_first(&mut fatal, rank, err);
pending.clear();
break;
}
}
dispatched = true;
continue;
}
DepState::Ready => {}
}
if fatal.is_some() || !units.active() {
pending.clear();
break;
}
pending.remove(index);
let deps = self.dep_states(service, &resolved);
let name = service.clone();
let mut report = Report {
tx: tx.clone(),
message: Some((rank, name.clone(), Ok(Resolution::Failed))),
};
workers.push(scope.spawn(move || {
let outcome = units.start(&name, &deps);
report.finish((rank, name, outcome));
}));
inflight += 1;
dispatched = true;
}
if inflight == 0 {
if !dispatched || pending.is_empty() {
break;
}
continue;
}
let Ok((rank, service, outcome)) = rx.recv() else {
break;
};
inflight -= 1;
match outcome {
Ok(state) => {
resolved.insert(service, state);
}
Err(err) => {
Self::keep_first(&mut fatal, rank, err);
resolved.insert(service, Resolution::Failed);
pending.clear();
}
}
}
});
match fatal {
Some((_, err)) => Err(err),
None => Ok(resolved),
}
}
fn keep_first<E>(slot: &mut Option<(usize, E)>, rank: usize, err: E) {
match slot {
Some((held, _)) if *held <= rank => {}
_ => *slot = Some((rank, err)),
}
}
fn gate_of(
&self,
service: &str,
resolved: &HashMap<String, Resolution>,
) -> DepState<'_> {
let mut waiting = false;
for dep in self.deps.get(service).into_iter().flatten() {
match resolved.get(*dep) {
None => waiting = true,
Some(Resolution::Skipped) => {
return DepState::Gated(dep, Gate::DependencySkipped);
}
Some(Resolution::Failed) => {
return DepState::Gated(dep, Gate::DependencyFailed);
}
Some(state) if state.satisfies_dependents() => {}
Some(_) => return DepState::Gated(dep, Gate::DependencyFailed),
}
}
if waiting {
DepState::Waiting
} else {
DepState::Ready
}
}
fn dep_states(
&self,
service: &str,
resolved: &HashMap<String, Resolution>,
) -> HashMap<String, Resolution> {
self.deps
.get(service)
.into_iter()
.flatten()
.filter_map(|dep| {
resolved.get(*dep).map(|state| ((*dep).to_string(), *state))
})
.collect()
}
}
struct Report<E> {
tx: mpsc::Sender<(usize, String, Result<Resolution, E>)>,
message: Option<(usize, String, Result<Resolution, E>)>,
}
impl<E> Report<E> {
fn finish(&mut self, message: (usize, String, Result<Resolution, E>)) {
self.message = Some(message);
}
}
impl<E> Drop for Report<E> {
fn drop(&mut self) {
if let Some(message) = self.message.take() {
let _ = self.tx.send(message);
}
}
}
enum DepState<'a> {
Ready,
Waiting,
Gated(&'a str, Gate),
}
#[cfg(test)]
mod tests {
use std::sync::{
Mutex,
atomic::{AtomicUsize, Ordering},
};
use super::*;
struct Recorder {
outcomes: HashMap<String, Resolution>,
statics: HashMap<String, Resolution>,
started: Mutex<Vec<String>>,
gated: Mutex<Vec<(String, String, Gate)>>,
inflight: AtomicUsize,
peak: AtomicUsize,
}
impl Recorder {
fn new() -> Self {
Self {
outcomes: HashMap::new(),
statics: HashMap::new(),
started: Mutex::new(Vec::new()),
gated: Mutex::new(Vec::new()),
inflight: AtomicUsize::new(0),
peak: AtomicUsize::new(0),
}
}
fn outcome(mut self, service: &str, state: Resolution) -> Self {
self.outcomes.insert(service.to_string(), state);
self
}
fn statically(mut self, service: &str, state: Resolution) -> Self {
self.statics.insert(service.to_string(), state);
self
}
}
impl Units for Recorder {
type Error = String;
fn static_resolution(
&self,
service: &str,
) -> Result<Option<Resolution>, Self::Error> {
Ok(self.statics.get(service).copied())
}
fn start(
&self,
service: &str,
_deps: &HashMap<String, Resolution>,
) -> Result<Resolution, Self::Error> {
let now = self.inflight.fetch_add(1, Ordering::SeqCst) + 1;
self.peak.fetch_max(now, Ordering::SeqCst);
self.started.lock().unwrap().push(service.to_string());
thread::sleep(std::time::Duration::from_millis(20));
self.inflight.fetch_sub(1, Ordering::SeqCst);
Ok(self
.outcomes
.get(service)
.copied()
.unwrap_or(Resolution::Running))
}
fn gated(
&self,
service: &str,
dependency: &str,
gate: Gate,
) -> Result<Resolution, Self::Error> {
self.gated.lock().unwrap().push((
service.to_string(),
dependency.to_string(),
gate,
));
Ok(match gate {
Gate::DependencySkipped => Resolution::Skipped,
Gate::DependencyFailed => Resolution::Failed,
})
}
fn active(&self) -> bool {
true
}
}
fn schedule<'a>(
order: &'a [String],
edges: &'a HashMap<&'a str, Vec<&'a str>>,
) -> Schedule<'a> {
Schedule::new(order, |service| {
edges.get(service).cloned().unwrap_or_default()
})
}
fn units(names: &[&str]) -> Vec<String> {
names.iter().map(|name| (*name).to_string()).collect()
}
#[test]
fn independent_units_run_concurrently() {
let order = units(&["a", "b", "c", "d"]);
let edges = HashMap::new();
let recorder = Recorder::new();
let resolved = schedule(&order, &edges).run(&recorder, None).unwrap();
assert_eq!(resolved.len(), 4);
assert_eq!(recorder.peak.load(Ordering::SeqCst), 4);
}
#[test]
fn a_limit_of_one_walks_units_in_order() {
let order = units(&["a", "b", "c"]);
let edges = HashMap::new();
let recorder = Recorder::new();
schedule(&order, &edges)
.run(&recorder, NonZeroUsize::new(1))
.unwrap();
assert_eq!(recorder.peak.load(Ordering::SeqCst), 1);
assert_eq!(*recorder.started.lock().unwrap(), units(&["a", "b", "c"]));
}
#[test]
fn a_cap_bounds_units_in_flight() {
let order = units(&["a", "b", "c", "d", "e"]);
let edges = HashMap::new();
let recorder = Recorder::new();
schedule(&order, &edges)
.run(&recorder, NonZeroUsize::new(2))
.unwrap();
assert_eq!(recorder.peak.load(Ordering::SeqCst), 2);
assert_eq!(recorder.started.lock().unwrap().len(), 5);
}
#[test]
fn a_dependent_waits_only_for_its_dependency() {
let order = units(&["db", "api", "worker"]);
let mut edges = HashMap::new();
edges.insert("api", vec!["db"]);
let recorder = Recorder::new();
schedule(&order, &edges).run(&recorder, None).unwrap();
let started = recorder.started.lock().unwrap().clone();
let db = started.iter().position(|name| name == "db").unwrap();
let api = started.iter().position(|name| name == "api").unwrap();
assert!(db < api, "api started before its dependency: {started:?}");
assert!(started.contains(&"worker".to_string()));
}
#[test]
fn a_failed_dependency_gates_its_dependents() {
let order = units(&["db", "api", "web"]);
let mut edges = HashMap::new();
edges.insert("api", vec!["db"]);
edges.insert("web", vec!["api"]);
let recorder = Recorder::new().outcome("db", Resolution::Failed);
let resolved = schedule(&order, &edges).run(&recorder, None).unwrap();
assert_eq!(resolved.get("api"), Some(&Resolution::Failed));
assert_eq!(resolved.get("web"), Some(&Resolution::Failed));
assert_eq!(*recorder.started.lock().unwrap(), units(&["db"]));
let gated = recorder.gated.lock().unwrap().clone();
assert_eq!(
gated[0],
("api".into(), "db".into(), Gate::DependencyFailed)
);
assert_eq!(
gated[1],
("web".into(), "api".into(), Gate::DependencyFailed)
);
}
#[test]
fn a_skipped_dependency_skips_its_dependents() {
let order = units(&["migrations", "api"]);
let mut edges = HashMap::new();
edges.insert("api", vec!["migrations"]);
let recorder = Recorder::new().statically("migrations", Resolution::Skipped);
let resolved = schedule(&order, &edges).run(&recorder, None).unwrap();
assert_eq!(resolved.get("api"), Some(&Resolution::Skipped));
assert!(recorder.started.lock().unwrap().is_empty());
assert_eq!(
recorder.gated.lock().unwrap()[0],
("api".into(), "migrations".into(), Gate::DependencySkipped)
);
}
#[test]
fn a_skipped_unit_is_not_reported_as_a_dependency_casualty() {
let order = units(&["db", "api"]);
let mut edges = HashMap::new();
edges.insert("api", vec!["db"]);
let recorder = Recorder::new()
.outcome("db", Resolution::Failed)
.statically("api", Resolution::Skipped);
let resolved = schedule(&order, &edges).run(&recorder, None).unwrap();
assert_eq!(resolved.get("api"), Some(&Resolution::Skipped));
assert!(
recorder.gated.lock().unwrap().is_empty(),
"a skipped unit must not be gated by its dependency"
);
}
#[test]
fn a_panicking_unit_does_not_wedge_the_schedule() {
struct Exploding;
impl Units for Exploding {
type Error = String;
fn static_resolution(
&self,
_service: &str,
) -> Result<Option<Resolution>, Self::Error> {
Ok(None)
}
fn start(
&self,
_service: &str,
_deps: &HashMap<String, Resolution>,
) -> Result<Resolution, Self::Error> {
panic!("unit exploded");
}
fn gated(
&self,
_service: &str,
_dependency: &str,
_gate: Gate,
) -> Result<Resolution, Self::Error> {
Ok(Resolution::Failed)
}
fn active(&self) -> bool {
true
}
}
let order = units(&["boom"]);
let edges = HashMap::new();
let previous = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let outcome =
std::panic::catch_unwind(|| schedule(&order, &edges).run(&Exploding, None));
std::panic::set_hook(previous);
assert!(outcome.is_err(), "the unit's panic must not be swallowed");
}
#[test]
fn edges_outside_the_schedule_are_dropped() {
let order = units(&["api"]);
let mut edges = HashMap::new();
edges.insert("api", vec!["db"]);
let recorder = Recorder::new();
let resolved = schedule(&order, &edges).run(&recorder, None).unwrap();
assert_eq!(resolved.get("api"), Some(&Resolution::Running));
}
}