mod job;
mod latch;
use std::cell::Cell;
use std::fmt;
use std::panic;
use std::ptr;
use self::job::{HeapJob, JobRef, StackJob};
use self::latch::{LockLatch, SpinLatch};
use crate::deque::{RawStealer, RawWorker, Steal};
use crate::queue::LockFreeQueue;
use crate::sync::{
Arc, AtomicUsize, Backoff,
Ordering::{AcqRel, Acquire},
WaitQueue, thread, thread_local,
};
use crate::utils::CachePadded;
const TERMINATE: usize = 1 << (usize::BITS - 1);
const INJECTOR_CAPACITY: usize = 1024;
struct Registry {
injector: LockFreeQueue<JobRef>,
stealers: Box<[RawStealer<JobRef>]>,
sleep: WaitQueue,
events: CachePadded<AtomicUsize>,
}
impl Registry {
fn inject(&self, job: JobRef) {
if self.injector.push(job).is_err() {
unreachable!("the injector is closed only by Drop, which owns the pool");
}
self.events.fetch_add(1, AcqRel);
self.sleep.notify_one();
}
}
struct WorkerThread {
index: usize,
deque: RawWorker<JobRef>,
rng: Cell<u64>,
registry: Arc<Registry>,
}
thread_local! {
#[allow(clippy::missing_const_for_thread_local)]
static CURRENT: Cell<*const WorkerThread> = Cell::new(ptr::null());
}
fn current_worker() -> Option<&'static WorkerThread> {
let ptr = CURRENT.with(Cell::get);
(!ptr.is_null()).then(|| unsafe { &*ptr })
}
enum Found {
Job(JobRef),
Retry,
Nothing,
}
impl WorkerThread {
fn next_random(&self) -> u64 {
let mut x = self.rng.get();
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.rng.set(x);
x
}
fn push(&self, job: JobRef) {
self.deque.push(job);
let registry = &self.registry;
if registry.sleep.has_waiters() {
registry.events.fetch_add(1, AcqRel);
registry.sleep.notify_one();
}
}
fn find_work(&self) -> Found {
if let Some(job) = self.deque.pop() {
return Found::Job(job);
}
let stealers = &self.registry.stealers;
let n = stealers.len();
let mut retry = false;
if n > 1 {
#[allow(clippy::cast_possible_truncation)]
let start = ((u128::from(self.next_random()) * n as u128) >> 64) as usize;
for k in 0..n {
let victim = (start + k) % n;
if victim == self.index {
continue;
}
match stealers[victim].steal() {
Steal::Success(job) => return Found::Job(job),
Steal::Retry => retry = true,
Steal::Empty => {}
}
}
}
if let Ok(job) = self.registry.injector.try_pop() {
return Found::Job(job);
}
if retry { Found::Retry } else { Found::Nothing }
}
fn join<A, B, RA, RB>(&self, a: A, b: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB + Send,
RA: Send,
RB: Send,
{
let job_b = StackJob::new(b, SpinLatch::new());
let job_b_ref = unsafe { job_b.as_job_ref() };
self.push(job_b_ref);
let result_a = panic::catch_unwind(panic::AssertUnwindSafe(a));
let mut backoff = Backoff::new();
let result_b = loop {
if job_b.latch.probe() {
break job_b.into_result();
}
match self.deque.pop() {
Some(job) if job == job_b_ref => {
break unsafe { job_b.run_inline() };
}
Some(job) => {
unsafe { job.execute() };
backoff.reset();
}
None => match self.find_work() {
Found::Job(job) => unsafe { job.execute() },
Found::Retry | Found::Nothing => backoff.snooze(),
},
}
};
match (result_a, result_b) {
(Ok(ra), Ok(rb)) => (ra, rb),
(Err(payload), _) | (_, Err(payload)) => panic::resume_unwind(payload),
}
}
}
#[allow(clippy::needless_pass_by_value)]
fn worker_main(worker: WorkerThread) {
CURRENT.with(|c| c.set(&raw const worker));
let registry = Arc::clone(&worker.registry);
let mut backoff = Backoff::new();
loop {
let seen = registry.events.load(Acquire);
match worker.find_work() {
Found::Job(job) => {
unsafe { job.execute() };
backoff.reset();
}
Found::Retry => backoff.spin(),
Found::Nothing if seen & TERMINATE != 0 => break,
Found::Nothing if backoff.is_completed() => {
registry
.sleep
.wait_until(|| registry.events.fetch_add(0, AcqRel) != seen, None);
backoff.reset();
}
Found::Nothing => backoff.snooze(),
}
}
CURRENT.with(|c| c.set(ptr::null()));
}
pub struct ThreadPool {
registry: Arc<Registry>,
threads: Vec<thread::JoinHandle<()>>,
}
impl ThreadPool {
pub fn new(threads: usize) -> Self {
assert!(threads > 0, "a pool needs at least one thread");
let deques: Vec<RawWorker<JobRef>> = (0..threads).map(|_| RawWorker::new()).collect();
let registry = Arc::new(Registry {
injector: LockFreeQueue::new(INJECTOR_CAPACITY),
stealers: deques.iter().map(RawWorker::stealer).collect(),
sleep: WaitQueue::new(),
events: CachePadded::new(AtomicUsize::new(0)),
});
let threads = deques
.into_iter()
.enumerate()
.map(|(index, deque)| {
let worker = WorkerThread {
index,
deque,
rng: Cell::new((index as u64 + 1).wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1),
registry: Arc::clone(®istry),
};
thread::spawn(move || worker_main(worker))
})
.collect();
Self { registry, threads }
}
#[must_use]
pub fn threads(&self) -> usize {
self.threads.len()
}
pub fn install<F, R>(&self, f: F) -> R
where
F: FnOnce() -> R + Send,
R: Send,
{
if let Some(worker) = current_worker() {
if ptr::eq(Arc::as_ptr(&worker.registry), Arc::as_ptr(&self.registry)) {
return f();
}
}
let job = StackJob::new(f, LockLatch::new());
self.registry.inject(unsafe { job.as_job_ref() });
job.latch.wait();
match job.into_result() {
Ok(r) => r,
Err(payload) => panic::resume_unwind(payload),
}
}
pub fn spawn<F>(&self, f: F)
where
F: FnOnce() + Send + 'static,
{
let job = HeapJob::new_ref(f);
match current_worker() {
Some(worker) if ptr::eq(Arc::as_ptr(&worker.registry), Arc::as_ptr(&self.registry)) => {
worker.push(job);
}
_ => self.registry.inject(job),
}
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
self.registry.events.fetch_or(TERMINATE, AcqRel);
self.registry.sleep.notify_all();
for handle in self.threads.drain(..) {
handle.join().expect("worker thread panicked");
}
}
}
impl fmt::Debug for ThreadPool {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ThreadPool")
.field("threads", &self.threads())
.finish_non_exhaustive()
}
}
pub fn join<A, B, RA, RB>(a: A, b: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB + Send,
RA: Send,
RB: Send,
{
match current_worker() {
Some(worker) => worker.join(a, b),
None => (a(), b()),
}
}