use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
struct Shared {
job: Mutex<Option<(*const (dyn Fn() + Sync + 'static), usize)>>,
cv: Condvar,
done: AtomicUsize,
done_cv: Condvar,
done_lock: Mutex<()>,
gen0: AtomicUsize,
shutdown: AtomicUsize,
}
unsafe impl Send for Shared {}
unsafe impl Sync for Shared {}
pub struct ThreadPool {
shared: Arc<Shared>,
workers: Vec<std::thread::JoinHandle<()>>,
nthreads: usize,
}
impl ThreadPool {
pub fn new(threads: usize) -> Arc<Self> {
let nthreads = threads.max(1);
let shared = Arc::new(Shared {
job: Mutex::new(None),
cv: Condvar::new(),
done: AtomicUsize::new(0),
done_cv: Condvar::new(),
done_lock: Mutex::new(()),
gen0: AtomicUsize::new(0),
shutdown: AtomicUsize::new(0),
});
let mut workers = Vec::new();
for _ in 0..nthreads.saturating_sub(1) {
let sh = shared.clone();
workers.push(std::thread::spawn(move || worker_loop(sh)));
}
Arc::new(ThreadPool {
shared,
workers,
nthreads,
})
}
pub fn threads(&self) -> usize {
self.nthreads
}
pub fn map_indexed<T, F>(&self, n: usize, f: F) -> Vec<T>
where
T: Send,
F: Fn(usize) -> T + Sync,
{
if self.nthreads <= 1 || n <= 1 {
return (0..n).map(f).collect();
}
let next = AtomicUsize::new(0);
let out: Mutex<Vec<(usize, T)>> = Mutex::new(Vec::with_capacity(n));
let work = || {
let mut got: Vec<(usize, T)> = Vec::new();
loop {
let i = next.fetch_add(1, Ordering::Relaxed);
if i >= n {
break;
}
got.push((i, f(i)));
}
out.lock().unwrap().extend(got);
};
let erased: *const (dyn Fn() + Sync) = &work;
let erased: *const (dyn Fn() + Sync + 'static) = unsafe { std::mem::transmute(erased) };
let spawned = self.workers.len();
{
let mut j = self.shared.job.lock().unwrap();
self.shared.done.store(0, Ordering::Release);
let g = self.shared.gen0.fetch_add(1, Ordering::AcqRel) + 1;
*j = Some((erased, g));
self.shared.cv.notify_all();
}
work(); if spawned > 0 {
let mut guard = self.shared.done_lock.lock().unwrap();
while self.shared.done.load(Ordering::Acquire) < spawned {
guard = self.shared.done_cv.wait(guard).unwrap();
}
}
*self.shared.job.lock().unwrap() = None;
let mut slots: Vec<Option<T>> = (0..n).map(|_| None).collect();
for (i, v) in out.into_inner().unwrap() {
slots[i] = Some(v);
}
slots
.into_iter()
.map(|s| s.expect("index produced"))
.collect()
}
}
fn worker_loop(shared: Arc<Shared>) {
let mut last_gen0 = 0usize;
loop {
let ptr = {
let mut j = shared.job.lock().unwrap();
loop {
if shared.shutdown.load(Ordering::Acquire) != 0 {
return;
}
if let Some((p, g)) = *j
&& g != last_gen0
{
last_gen0 = g;
break p;
}
j = shared.cv.wait(j).unwrap();
}
};
unsafe { (*ptr)() };
{
let _l = shared.done_lock.lock().unwrap();
shared.done.fetch_add(1, Ordering::AcqRel);
shared.done_cv.notify_all();
}
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
self.shared.shutdown.store(1, Ordering::Release);
self.shared.cv.notify_all();
for h in self.workers.drain(..) {
let _ = h.join();
}
}
}