use std::vec::Vec;
use super::executor::registered_parallel_executor;
type PanicPayload = Option<std::boxed::Box<dyn std::any::Any + Send>>;
type PanicPayloadMutex = std::sync::Mutex<PanicPayload>;
fn lock_panic_payload(mutex: &PanicPayloadMutex) -> std::sync::MutexGuard<'_, PanicPayload> {
mutex
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn take_panic_payload(mutex: PanicPayloadMutex) -> PanicPayload {
mutex
.into_inner()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
struct ExecutorDropGuard<R> {
out_ptr: *mut core::mem::MaybeUninit<R>,
capacity: usize,
successful_ptr: *const bool,
num_chunks: usize,
active: bool,
}
impl<R> Drop for ExecutorDropGuard<R> {
fn drop(&mut self) {
if self.active {
for index in 0..self.num_chunks {
unsafe {
if *self.successful_ptr.add(index) {
self.out_ptr.add(index).cast::<R>().drop_in_place();
}
}
}
unsafe {
let _ = Vec::from_raw_parts(self.out_ptr, 0, self.capacity);
}
}
}
}
struct TaskContext<'a, R, Run> {
run: &'a Run,
out_ptr: *mut core::mem::MaybeUninit<R>,
successful_ptr: *mut bool,
panic_payload: &'a PanicPayloadMutex,
}
unsafe fn task_wrapper<R, Run>(index: usize, data: *mut ())
where
R: Send,
Run: Fn(usize) -> R + Sync,
{
let ctx = unsafe { &*(data as *const TaskContext<'_, R, Run>) };
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| (ctx.run)(index)));
match result {
Ok(val) => {
unsafe {
ctx.out_ptr
.add(index)
.write(core::mem::MaybeUninit::new(val));
*ctx.successful_ptr.add(index) = true;
}
}
Err(payload) => {
let mut g = lock_panic_payload(ctx.panic_payload);
if g.is_none() {
*g = Some(payload);
}
}
}
}
pub(super) fn drive<R, Run>(num_chunks: usize, run: Run) -> Vec<R>
where
R: Send,
Run: Fn(usize) -> R + Sync,
{
debug_assert!(num_chunks >= 1, "drive requires at least one chunk");
if let Some(executor) = registered_parallel_executor() {
let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(num_chunks);
unsafe {
out.set_len(num_chunks);
}
let out_ptr = out.as_mut_ptr();
let capacity = out.capacity();
core::mem::forget(out);
let mut successful = std::vec![false; num_chunks];
let panic_payload = PanicPayloadMutex::new(None);
let mut guard = ExecutorDropGuard {
out_ptr,
capacity,
successful_ptr: successful.as_ptr(),
num_chunks,
active: true,
};
let mut ctx = TaskContext {
run: &run,
out_ptr,
successful_ptr: successful.as_mut_ptr(),
panic_payload: &panic_payload,
};
unsafe {
executor.execute(
num_chunks,
task_wrapper::<R, Run>,
core::ptr::addr_of_mut!(ctx).cast::<()>(),
);
}
guard.active = false;
if let Some(payload) = take_panic_payload(panic_payload) {
for (index, &success) in successful.iter().enumerate() {
if success {
unsafe {
out_ptr.add(index).cast::<R>().drop_in_place();
}
}
}
unsafe {
let _ = Vec::from_raw_parts(out_ptr, 0, capacity);
}
std::panic::resume_unwind(payload);
}
return unsafe { Vec::from_raw_parts(out_ptr.cast::<R>(), num_chunks, capacity) };
}
if num_chunks == 1 {
return std::vec![run(0)];
}
std::thread::scope(|scope| {
let run = &run;
let mut handles = Vec::with_capacity(num_chunks - 1);
for index in 0..(num_chunks - 1) {
handles.push(scope.spawn(move || run(index)));
}
let last = run(num_chunks - 1);
let mut results = Vec::with_capacity(num_chunks);
for h in handles {
match h.join() {
Ok(value) => results.push(value),
Err(payload) => std::panic::resume_unwind(payload),
}
}
results.push(last);
results
})
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use std::boxed::Box;
use super::{lock_panic_payload, take_panic_payload, PanicPayloadMutex, TaskContext};
#[test]
fn poisoned_payload_mutex_preserves_the_first_panic() {
let mutex = PanicPayloadMutex::new(None);
let poison = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let mut payload = mutex.lock().unwrap();
*payload = Some(Box::new("first panic"));
panic!("poison payload mutex");
}));
let Err(poison_payload) = poison else {
panic!("poisoning the payload mutex must unwind");
};
assert_eq!(
poison_payload.downcast_ref::<&'static str>(),
Some(&"poison payload mutex")
);
let run: fn(usize) = |_| panic!("second panic");
let mut successful = false;
let mut context = TaskContext::<(), fn(usize)> {
run: &run,
out_ptr: core::ptr::null_mut(),
successful_ptr: &mut successful,
panic_payload: &mutex,
};
unsafe {
super::task_wrapper::<(), fn(usize)>(0, core::ptr::addr_of_mut!(context).cast::<()>());
}
let payload = take_panic_payload(mutex).expect("first panic payload survives poisoning");
assert_eq!(payload.downcast_ref::<&'static str>(), Some(&"first panic"));
let mutex = PanicPayloadMutex::new(None);
let poison = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _guard = mutex.lock().unwrap();
panic!("poison empty payload mutex");
}));
let Err(poison_payload) = poison else {
panic!("poisoning the empty payload mutex must unwind");
};
assert_eq!(
poison_payload.downcast_ref::<&'static str>(),
Some(&"poison empty payload mutex")
);
let payload = lock_panic_payload(&mutex);
assert_eq!(payload.as_ref().map(|_| ()), None);
}
}