use crossbeam::deque::{Injector, Steal, Stealer, Worker};
use std::fmt::{self, Debug};
use std::iter::repeat_with;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::thread::{self, Builder, JoinHandle};
use std::{cmp, panic};
use crate::executor::strategy::{Signal, Strategy};
use crate::executor::task::Task;
use crate::executor::{Error, Result};
pub struct WorkStealing {
injector: Arc<Injector<Box<dyn Task>>>,
signal: Arc<Signal>,
threads: Vec<JoinHandle<Result>>,
running: Arc<AtomicUsize>,
pending: Arc<AtomicUsize>,
capacity: usize,
}
impl WorkStealing {
#[must_use]
pub fn new(num_workers: usize) -> Self {
Self::with_capacity(num_workers, 8 * num_workers)
}
#[must_use]
pub fn with_capacity(num_workers: usize, capacity: usize) -> Self {
let injector = Arc::new(Injector::new());
let signal = Arc::new(Signal::new());
let mut workers = Vec::with_capacity(num_workers);
for _ in 0..num_workers {
workers.push(Worker::new_fifo());
}
let stealers: Arc<[Stealer<Box<dyn Task>>]> =
Arc::from(workers.iter().map(Worker::stealer).collect::<Vec<_>>());
let running = Arc::new(AtomicUsize::new(0));
let pending = Arc::new(AtomicUsize::new(0));
let iter = workers.into_iter().enumerate().map(|(index, worker)| {
let injector = Arc::clone(&injector);
let stealers = Arc::clone(&stealers);
let signal = Arc::clone(&signal);
let running = Arc::clone(&running);
let pending = Arc::clone(&pending);
let h = move || {
let injector = injector.as_ref();
let stealers = stealers.as_ref();
loop {
let Some(task) = get(&worker, injector, stealers) else {
if signal.should_terminate()? {
break;
}
continue;
};
pending.fetch_sub(1, Ordering::Acquire);
running.fetch_add(1, Ordering::Release);
let subtasks = panic::catch_unwind(|| task.execute())
.unwrap_or_default();
running.fetch_sub(1, Ordering::Acquire);
if !subtasks.is_empty() {
let added = subtasks
.into_iter()
.map(|subtask| worker.push(subtask))
.count();
pending.fetch_add(added, Ordering::Release);
signal.notify();
}
}
Ok(())
};
Builder::new()
.name(format!("zrx/executor/{}", index + 1))
.spawn(h)
.unwrap()
});
let threads = iter.collect();
Self {
injector,
signal,
threads,
running,
pending,
capacity,
}
}
}
impl Strategy for WorkStealing {
fn submit(&self, task: Box<dyn Task>) -> Result {
let pending = self.pending.fetch_add(1, Ordering::Release);
if pending == self.capacity {
self.pending.fetch_sub(1, Ordering::Release);
return Err(Error::Submit(task));
}
self.injector.push(task);
self.signal.notify();
Ok(())
}
#[inline]
fn num_workers(&self) -> usize {
self.threads.len()
}
#[inline]
fn num_tasks_running(&self) -> usize {
self.running.load(Ordering::Relaxed)
}
#[inline]
fn num_tasks_pending(&self) -> usize {
self.pending.load(Ordering::Relaxed)
}
#[inline]
fn capacity(&self) -> usize {
self.capacity
}
}
impl Default for WorkStealing {
#[inline]
fn default() -> Self {
Self::new(cmp::max(
thread::available_parallelism()
.map_or(1, |num| num.get().saturating_sub(1)),
1,
))
}
}
impl Drop for WorkStealing {
fn drop(&mut self) {
let _ = self.signal.terminate();
for handle in self.threads.drain(..) {
let _ = handle.join();
}
}
}
impl Debug for WorkStealing {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("WorkStealing")
.field("workers", &self.num_workers())
.field("running", &self.num_tasks_running())
.field("pending", &self.num_tasks_pending())
.finish()
}
}
fn get<T>(
worker: &Worker<T>, injector: &Injector<T>, stealers: &[Stealer<T>],
) -> Option<T> {
worker
.pop()
.or_else(|| steal_or_retry(worker, injector, stealers))
}
fn steal_or_retry<T>(
worker: &Worker<T>, injector: &Injector<T>, stealers: &[Stealer<T>],
) -> Option<T> {
repeat_with(|| steal(worker, injector, stealers))
.find(|steal| !steal.is_retry())
.and_then(Steal::success)
}
fn steal<T>(
worker: &Worker<T>, injector: &Injector<T>, stealers: &[Stealer<T>],
) -> Steal<T> {
injector
.steal_batch_and_pop(worker)
.or_else(|| stealers.iter().map(Stealer::steal).collect())
}