use std::vec::Vec;
use super::executor::registered_parallel_executor;
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 std::sync::Mutex<Option<std::boxed::Box<dyn std::any::Any + Send>>>,
}
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) => {
if let Ok(mut g) = ctx.panic_payload.lock() {
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 = std::sync::Mutex::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>,
&mut ctx as *mut TaskContext<'_, R, Run> as *mut (),
);
}
guard.active = false;
if let Some(payload) = panic_payload.into_inner().unwrap() {
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
})
}