use std::any::Any;
use std::cell::UnsafeCell;
use std::collections::VecDeque;
use std::ops::{Deref, DerefMut, Range};
use std::panic::{AssertUnwindSafe, catch_unwind, resume_unwind};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::{self, JoinHandle};
pub(crate) struct DisjointMut<T> {
inner: UnsafeCell<Vec<T>>,
#[cfg(debug_assertions)]
borrows: Mutex<Vec<Range<usize>>>,
}
unsafe impl<T: Send> Sync for DisjointMut<T> {}
unsafe impl<T: Send> Send for DisjointMut<T> {}
impl<T> DisjointMut<T> {
pub(crate) fn new(v: Vec<T>) -> Self {
DisjointMut {
inner: UnsafeCell::new(v),
#[cfg(debug_assertions)]
borrows: Mutex::new(Vec::new()),
}
}
#[allow(dead_code)]
pub(crate) fn len(&self) -> usize {
unsafe { (*self.inner.get()).len() }
}
pub(crate) fn slice_mut(&self, range: Range<usize>) -> DisjointMutGuard<'_, T> {
let vec = unsafe { &mut *self.inner.get() };
assert!(
range.end <= vec.len() && range.start <= range.end,
"DisjointMut::slice_mut range {range:?} out of bounds (len {})",
vec.len()
);
#[cfg(debug_assertions)]
{
let conflict = {
let mut live = self.borrows.lock().unwrap_or_else(|p| p.into_inner());
let clash = live
.iter()
.find(|other| range.start < other.end && other.start < range.end)
.cloned();
if clash.is_none() {
live.push(range.clone());
}
clash
};
assert!(
conflict.is_none(),
"DisjointMut: overlapping borrow {range:?} vs live {:?}",
conflict.unwrap()
);
}
let ptr = vec.as_mut_ptr();
let slice = unsafe {
std::slice::from_raw_parts_mut(ptr.add(range.start), range.end - range.start)
};
DisjointMutGuard {
slice,
#[cfg(debug_assertions)]
parent: self,
#[cfg(debug_assertions)]
range,
}
}
pub(crate) fn into_inner(self) -> Vec<T> {
self.inner.into_inner()
}
}
pub(crate) struct DisjointMutGuard<'a, T> {
slice: &'a mut [T],
#[cfg(debug_assertions)]
parent: &'a DisjointMut<T>,
#[cfg(debug_assertions)]
range: Range<usize>,
}
impl<T> Deref for DisjointMutGuard<'_, T> {
type Target = [T];
fn deref(&self) -> &[T] {
self.slice
}
}
impl<T> DerefMut for DisjointMutGuard<'_, T> {
fn deref_mut(&mut self) -> &mut [T] {
self.slice
}
}
#[cfg(debug_assertions)]
impl<T> Drop for DisjointMutGuard<'_, T> {
fn drop(&mut self) {
let mut live = self
.parent
.borrows
.lock()
.unwrap_or_else(|p| p.into_inner());
if let Some(pos) = live.iter().position(|r| *r == self.range) {
live.swap_remove(pos);
}
}
}
type Job = Box<dyn FnOnce() + Send + 'static>;
pub(crate) struct ProgressGate {
v: AtomicUsize,
waiters: AtomicUsize,
lock: Mutex<()>,
cvar: Condvar,
}
impl ProgressGate {
pub(crate) fn new() -> Self {
ProgressGate {
v: AtomicUsize::new(0),
waiters: AtomicUsize::new(0),
lock: Mutex::new(()),
cvar: Condvar::new(),
}
}
pub(crate) fn publish(&self, x: usize) {
self.v.store(x, Ordering::SeqCst);
if self.waiters.load(Ordering::SeqCst) > 0 {
let _g = self.lock.lock().unwrap_or_else(|p| p.into_inner());
self.cvar.notify_all();
}
}
pub(crate) fn wait_at_least(&self, x: usize) {
for _ in 0..512 {
if self.v.load(Ordering::Acquire) >= x {
return;
}
std::hint::spin_loop();
}
self.waiters.fetch_add(1, Ordering::SeqCst);
let mut g = self.lock.lock().unwrap_or_else(|p| p.into_inner());
while self.v.load(Ordering::SeqCst) < x {
g = self.cvar.wait(g).unwrap_or_else(|p| p.into_inner());
}
drop(g);
self.waiters.fetch_sub(1, Ordering::SeqCst);
}
}
struct Deque {
jobs: Mutex<VecDeque<Job>>,
}
impl Deque {
fn new() -> Self {
Deque {
jobs: Mutex::new(VecDeque::new()),
}
}
fn push(&self, job: Job) {
self.jobs.lock().unwrap().push_back(job);
}
fn pop(&self) -> Option<Job> {
self.jobs.lock().unwrap().pop_back()
}
fn steal(&self) -> Option<Job> {
self.jobs.lock().unwrap().pop_front()
}
}
struct Shared {
deques: Vec<Deque>,
injector: Mutex<VecDeque<Job>>,
pending: AtomicUsize,
queued: AtomicUsize,
cvar: Condvar,
lock: Mutex<()>,
shutdown: AtomicBool,
}
impl Shared {
fn find_job(&self, me: usize) -> Option<Job> {
let job = self.find_job_inner(me);
if job.is_some() {
self.queued.fetch_sub(1, Ordering::Relaxed);
}
job
}
fn find_job_inner(&self, me: usize) -> Option<Job> {
if let Some(d) = self.deques.get(me)
&& let Some(j) = d.pop()
{
return Some(j);
}
let n = self.deques.len();
for off in 1..=n {
let victim = (me + off) % n;
if victim == me {
continue;
}
if let Some(j) = self.deques[victim].steal() {
return Some(j);
}
}
self.injector.lock().unwrap().pop_front()
}
fn notify_one_job(&self) {
let _g = self.lock.lock().unwrap();
self.cvar.notify_one();
}
fn notify_all_workers(&self) {
let _g = self.lock.lock().unwrap();
self.cvar.notify_all();
}
fn finish_one(&self) {
if self.pending.fetch_sub(1, Ordering::Release) == 1 {
let _g = self.lock.lock().unwrap();
self.cvar.notify_all();
}
}
}
pub(crate) struct ThreadPool {
shared: Arc<Shared>,
workers: Vec<JoinHandle<()>>,
round_robin: AtomicUsize,
}
impl ThreadPool {
pub(crate) fn new(threads: usize) -> Self {
let threads = threads.max(1);
let mut deques = Vec::with_capacity(threads);
for _ in 0..threads {
deques.push(Deque::new());
}
let shared = Arc::new(Shared {
deques,
injector: Mutex::new(VecDeque::new()),
pending: AtomicUsize::new(0),
queued: AtomicUsize::new(0),
cvar: Condvar::new(),
lock: Mutex::new(()),
shutdown: AtomicBool::new(false),
});
let mut workers = Vec::with_capacity(threads);
for id in 0..threads {
let shared = Arc::clone(&shared);
let handle = thread::Builder::new()
.name(format!("hpvcd-worker-{}", id))
.spawn(move || worker_loop(shared, id))
.expect("spawn worker thread");
workers.push(handle);
}
ThreadPool {
shared,
workers,
round_robin: AtomicUsize::new(0),
}
}
pub(crate) fn with_available_parallelism() -> Self {
let n = thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
ThreadPool::new(n)
}
pub(crate) fn threads(&self) -> usize {
self.workers.len()
}
fn submit(&self, job: Job) {
self.shared.pending.fetch_add(1, Ordering::Relaxed);
self.shared.queued.fetch_add(1, Ordering::Relaxed);
let n = self.shared.deques.len();
let idx = self.round_robin.fetch_add(1, Ordering::Relaxed) % n;
self.shared.deques[idx].push(job);
self.shared.notify_one_job();
}
pub(crate) fn scope<'scope, F, R>(&'scope self, f: F) -> R
where
F: FnOnce(&Scope<'scope>) -> R,
{
let scope = Scope {
pool: self,
state: Mutex::new(ScopeState {
outstanding: 0,
panic: None,
}),
done: Condvar::new(),
};
let result = catch_unwind(AssertUnwindSafe(|| f(&scope)));
let worker_panic = scope.wait();
match result {
Ok(value) => {
if let Some(payload) = worker_panic {
resume_unwind(payload);
}
value
}
Err(payload) => resume_unwind(payload),
}
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
self.shared.shutdown.store(true, Ordering::SeqCst);
self.shared.notify_all_workers();
for w in self.workers.drain(..) {
let _ = w.join();
}
}
}
fn worker_loop(shared: Arc<Shared>, id: usize) {
loop {
if let Some(job) = shared.find_job(id) {
job();
shared.finish_one();
continue;
}
if shared.shutdown.load(Ordering::SeqCst) {
if let Some(job) = shared.find_job(id) {
job();
shared.finish_one();
continue;
}
break;
}
let guard = shared.lock.lock().unwrap();
if shared.shutdown.load(Ordering::SeqCst) {
break;
}
let has_work = shared.queued.load(Ordering::Acquire) > 0;
if has_work {
drop(guard);
continue;
}
let _unused = shared.cvar.wait(guard).unwrap();
}
}
pub(crate) struct Scope<'scope> {
pool: &'scope ThreadPool,
state: Mutex<ScopeState>,
done: Condvar,
}
struct ScopeState {
outstanding: usize,
panic: Option<Box<dyn Any + Send + 'static>>,
}
impl<'scope> Scope<'scope> {
pub(crate) fn spawn<F>(&self, f: F)
where
F: FnOnce() + Send + 'scope,
{
self.state
.lock()
.unwrap_or_else(|p| p.into_inner())
.outstanding += 1;
let scope_ptr: *const Scope<'scope> = self;
let scope_addr = scope_ptr as usize;
let job: Box<dyn FnOnce() + Send + 'scope> = Box::new(move || {
let result = catch_unwind(AssertUnwindSafe(f));
let scope = unsafe { &*(scope_addr as *const Scope<'scope>) };
let mut state = scope.state.lock().unwrap_or_else(|p| p.into_inner());
if let Err(payload) = result
&& state.panic.is_none()
{
state.panic = Some(payload);
}
state.outstanding -= 1;
if state.outstanding == 0 {
scope.done.notify_all();
}
});
let job: Job = unsafe {
std::mem::transmute::<
Box<dyn FnOnce() + Send + 'scope>,
Box<dyn FnOnce() + Send + 'static>,
>(job)
};
self.pool.submit(job);
}
fn wait(&self) -> Option<Box<dyn Any + Send + 'static>> {
let external = self.pool.shared.deques.len();
loop {
if let Some(job) = self.pool.shared.find_job(external) {
job();
self.pool.shared.finish_one();
continue;
}
let mut state = self.state.lock().unwrap_or_else(|p| p.into_inner());
if state.outstanding == 0 {
return state.panic.take();
}
let _unused = self.done.wait(state).unwrap_or_else(|p| p.into_inner());
}
}
}
pub(crate) fn parallel_for<F>(pool: &ThreadPool, count: usize, body: F)
where
F: Fn(usize) + Send + Sync,
{
if count == 0 {
return;
}
if count == 1 || pool.threads() == 1 {
for i in 0..count {
body(i);
}
return;
}
let body = &body;
pool.scope(|s| {
for i in 0..count {
s.spawn(move || body(i));
}
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disjoint_regions_write_independently() {
let dm = DisjointMut::new(vec![0u32; 8]);
{
let mut a = dm.slice_mut(0..4);
let mut b = dm.slice_mut(4..8);
for (i, x) in a.iter_mut().enumerate() {
*x = i as u32;
}
for (i, x) in b.iter_mut().enumerate() {
*x = 100 + i as u32;
}
}
assert_eq!(dm.into_inner(), vec![0, 1, 2, 3, 100, 101, 102, 103]);
}
#[test]
#[should_panic(expected = "overlapping borrow")]
#[cfg(debug_assertions)]
fn overlapping_borrow_panics() {
let dm = DisjointMut::new(vec![0u8; 8]);
let _a = dm.slice_mut(0..5);
let _b = dm.slice_mut(3..8); }
#[test]
fn pool_parallel_for_disjoint_write() {
let pool = ThreadPool::new(4);
let dm = DisjointMut::new(vec![0usize; 16]);
parallel_for(&pool, 4, |tile| {
let mut region = dm.slice_mut(tile * 4..tile * 4 + 4);
for (i, x) in region.iter_mut().enumerate() {
*x = tile * 4 + i;
}
});
let out = dm.into_inner();
assert_eq!(out, (0..16).collect::<Vec<_>>());
}
#[test]
fn pool_reused_across_scopes() {
let pool = ThreadPool::new(3);
let sum = AtomicUsize::new(0);
for _ in 0..5 {
parallel_for(&pool, 10, |i| {
sum.fetch_add(i, Ordering::Relaxed);
});
}
assert_eq!(sum.load(Ordering::Relaxed), 45 * 5);
}
#[test]
fn worker_panic_is_propagated_after_all_jobs_finish() {
let pool = ThreadPool::new(3);
let completed = AtomicUsize::new(0);
let panic = catch_unwind(AssertUnwindSafe(|| {
parallel_for(&pool, 16, |i| {
if i == 7 {
panic!("worker boom");
}
completed.fetch_add(1, Ordering::Relaxed);
});
}));
assert!(panic.is_err());
assert_eq!(completed.load(Ordering::Relaxed), 15);
parallel_for(&pool, 8, |_| {
completed.fetch_add(1, Ordering::Relaxed);
});
assert_eq!(completed.load(Ordering::Relaxed), 23);
}
#[test]
fn scope_body_panic_still_joins_borrowing_jobs() {
let pool = ThreadPool::new(3);
let completed = AtomicUsize::new(0);
let panic = catch_unwind(AssertUnwindSafe(|| {
pool.scope(|scope| {
for _ in 0..16 {
scope.spawn(|| {
completed.fetch_add(1, Ordering::Relaxed);
});
}
panic!("scope body boom");
});
}));
assert!(panic.is_err());
assert_eq!(completed.load(Ordering::Relaxed), 16);
}
}