use std::panic::{AssertUnwindSafe, catch_unwind, resume_unwind};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::JoinHandle;
type Task<'a> = &'a (dyn Fn() + Sync);
struct Shared {
job: Mutex<State>,
post: Condvar,
done: Condvar,
}
struct State {
job: Option<*const (dyn Fn() + Sync)>,
epoch: u64,
unclaimed: usize,
running: usize,
shutdown: bool,
panic: Option<Box<dyn std::any::Any + Send>>,
}
unsafe impl Send for State {}
pub(crate) struct Pool {
shared: Arc<Shared>,
handles: Vec<JoinHandle<()>>,
width: usize,
}
impl Pool {
pub(crate) fn new(threads: usize) -> Self {
let width = crate::coder::resolve_threads(threads).max(1);
let shared = Arc::new(Shared {
job: Mutex::new(State {
job: None,
epoch: 0,
unclaimed: 0,
running: 0,
shutdown: false,
panic: None,
}),
post: Condvar::new(),
done: Condvar::new(),
});
let handles = (1..width)
.map(|i| {
let sh = Arc::clone(&shared);
std::thread::Builder::new()
.name(format!("maroontree-{i}"))
.spawn(move || worker(&sh))
.expect("spawn pool worker")
})
.collect();
Pool {
shared,
handles,
width,
}
}
pub(crate) fn width(&self) -> usize {
self.width
}
fn run(&self, cap: usize, f: Task) {
let want = cap.min(self.width);
if want <= 1 {
f();
return;
}
let sh = &*self.shared;
let erased: *const (dyn Fn() + Sync) =
unsafe { std::mem::transmute::<Task, *const (dyn Fn() + Sync)>(f) };
{
let mut s = sh.job.lock().expect("pool poisoned");
debug_assert!(s.job.is_none(), "nested pool run");
s.job = Some(erased);
s.epoch += 1;
s.unclaimed = want - 1;
s.panic = None;
sh.post.notify_all();
}
let local = catch_unwind(AssertUnwindSafe(f)).err();
let mut s = sh.job.lock().expect("pool poisoned");
s.unclaimed = 0;
while s.running > 0 {
s = sh.done.wait(s).expect("pool poisoned");
}
s.job = None;
let worker_panic = s.panic.take();
drop(s);
if let Some(p) = local.or(worker_panic) {
resume_unwind(p);
}
}
pub(crate) fn map_indexed<T, F>(&self, cap: usize, n: usize, f: F) -> Vec<T>
where
T: Send,
F: Fn(usize) -> T + Sync,
{
if self.width <= 1 || cap <= 1 || n <= 1 {
return (0..n).map(f).collect();
}
let slots: Vec<Mutex<Option<T>>> = std::iter::repeat_with(|| Mutex::new(None))
.take(n)
.collect();
let next = AtomicUsize::new(0);
self.run(cap.min(n), &|| loop {
let i = next.fetch_add(1, Ordering::Relaxed);
if i >= n {
break;
}
let v = f(i);
*slots[i].lock().expect("slot poisoned") = Some(v);
});
slots
.into_iter()
.map(|c| {
c.into_inner()
.expect("slot poisoned")
.expect("index produced")
})
.collect()
}
pub(crate) fn for_each<T, F>(&self, cap: usize, items: Vec<T>, f: F)
where
T: Send,
F: Fn(T) + Sync,
{
let n = items.len();
if self.width <= 1 || cap <= 1 || n <= 1 {
items.into_iter().for_each(f);
return;
}
let slots: Vec<Mutex<Option<T>>> = items.into_iter().map(|t| Mutex::new(Some(t))).collect();
let next = AtomicUsize::new(0);
self.run(cap.min(n), &|| loop {
let i = next.fetch_add(1, Ordering::Relaxed);
if i >= n {
break;
}
let it = slots[i].lock().expect("slot poisoned").take();
f(it.expect("index claimed once"));
});
}
}
impl Drop for Pool {
fn drop(&mut self) {
{
let mut s = self.shared.job.lock().expect("pool poisoned");
s.shutdown = true;
self.shared.post.notify_all();
}
for h in self.handles.drain(..) {
let _ = h.join();
}
}
}
fn worker(sh: &Shared) {
let mut seen = 0u64;
loop {
let job = {
let mut s = sh.job.lock().expect("pool poisoned");
loop {
if s.shutdown {
return;
}
if s.epoch != seen
&& s.unclaimed > 0
&& let Some(job) = s.job
{
seen = s.epoch;
s.unclaimed -= 1;
s.running += 1;
break job;
}
s = sh.post.wait(s).expect("pool poisoned");
}
};
let r = catch_unwind(AssertUnwindSafe(|| unsafe { (*job)() }));
let mut s = sh.job.lock().expect("pool poisoned");
s.running -= 1;
if let Err(p) = r {
s.panic.get_or_insert(p);
}
if s.running == 0 {
sh.done.notify_all();
}
}
}