use std::{
panic::{catch_unwind, AssertUnwindSafe},
sync::{atomic::Ordering, Arc},
};
use moirai_core::{
error::{ExecutorError, ExecutorResult},
Priority,
};
use super::super::super::{class::WorkClass, reduce::ReduceSlots};
use super::super::types::{
get_current_worker_id, is_in_indexed_region, IndexedRegionGuard, SchedulerScopeState,
SharedScopedTaskCompletion, ThreadScheduler,
};
use super::super::worker::{
indexed_chunk_bounds, indexed_chunk_count, inline_map_reduce, map_reduce_range,
};
fn execute_catching_panic<T>(operation: impl FnOnce() -> T) -> ExecutorResult<T> {
catch_unwind(AssertUnwindSafe(operation))
.map_err(|_| ExecutorError::SpawnFailed(moirai_core::error::TaskError::Panicked))
}
impl<const QUEUE_CAPACITY: usize, const SPIN_LIMIT: usize>
ThreadScheduler<QUEUE_CAPACITY, SPIN_LIMIT>
{
pub fn for_each_indexed<C, F>(
&self,
priority: Priority,
locality_hint: Option<usize>,
count: usize,
task: F,
) -> ExecutorResult<()>
where
C: WorkClass,
F: Fn(usize) + Send + Sync,
{
if self.inner.shutdown.load(Ordering::Acquire) {
return Err(ExecutorError::ShuttingDown);
}
if count == 0 {
return Ok(());
}
if get_current_worker_id().is_some() || is_in_indexed_region() {
return execute_catching_panic(|| {
for index in 0..count {
task(index);
}
});
}
let chunk_count = indexed_chunk_count(count, self.worker_count());
let (_, caller_end) = indexed_chunk_bounds(count, chunk_count, 0);
if chunk_count == 1 {
return execute_catching_panic(|| {
let _region = IndexedRegionGuard::enter();
for index in 0..caller_end {
task(index);
}
});
}
let state = Arc::new(SchedulerScopeState::new());
let task = &task;
let mut schedule_result = Ok(());
let mut inline_result = Ok(());
for chunk_index in 1..chunk_count {
let (start, end) = indexed_chunk_bounds(count, chunk_count, chunk_index);
state.register_task();
let completion = SharedScopedTaskCompletion {
state: Arc::clone(&state),
};
let scoped_job = move |_| {
let completion = completion;
let result = execute_catching_panic(|| {
for index in start..end {
task(index);
}
});
if result.is_err() {
completion.mark_failed();
}
};
if let Err(error) =
self.schedule_scoped_job::<C, _>(priority, locality_hint, scoped_job)
{
match error {
ExecutorError::ResourceExhausted(_) => {
self.record_admission_caller_run();
inline_result = execute_catching_panic(|| {
for index in start..end {
task(index);
}
});
if inline_result.is_err() {
break;
}
}
other => {
schedule_result = Err(other);
break;
}
}
}
}
let caller_result = if schedule_result.is_ok() && inline_result.is_ok() {
execute_catching_panic(|| {
let _region = IndexedRegionGuard::enter();
for index in 0..caller_end {
task(index);
}
})
} else {
Ok(())
};
self.drain_scope(&state);
if state.failed_tasks.load(Ordering::Acquire)
|| inline_result.is_err()
|| caller_result.is_err()
{
Err(ExecutorError::SpawnFailed(
moirai_core::error::TaskError::Panicked,
))
} else {
schedule_result
}
}
pub fn map_reduce_indexed<C, T, Map, Reduce>(
&self,
priority: Priority,
locality_hint: Option<usize>,
count: usize,
identity: T,
map: Map,
reduce: Reduce,
) -> ExecutorResult<T>
where
C: WorkClass,
T: Send + Clone,
Map: Fn(usize) -> T + Send + Sync,
Reduce: Fn(T, T) -> T + Send + Sync,
{
if self.inner.shutdown.load(Ordering::Acquire) {
return Err(ExecutorError::ShuttingDown);
}
if count == 0 {
return Ok(identity);
}
if get_current_worker_id().is_some() || is_in_indexed_region() {
return inline_map_reduce(count, identity, map, reduce);
}
let chunk_count = indexed_chunk_count(count, self.worker_count());
let (_, caller_end) = indexed_chunk_bounds(count, chunk_count, 0);
if chunk_count == 1 {
let _region = IndexedRegionGuard::enter();
return inline_map_reduce(count, identity, map, reduce);
}
let state = Arc::new(SchedulerScopeState::new());
let slots = Arc::new(ReduceSlots::new(chunk_count - 1));
let map = ↦
let reduce = &reduce;
let mut schedule_result = Ok(());
let mut inline_result = Ok(());
for chunk_index in 1..chunk_count {
let (start, end) = indexed_chunk_bounds(count, chunk_count, chunk_index);
state.register_task();
let completion = SharedScopedTaskCompletion {
state: Arc::clone(&state),
};
let slots_chunk = Arc::clone(&slots);
let identity_chunk = identity.clone();
let scoped_job = move |_| {
let completion = completion;
let result = execute_catching_panic(|| {
let accumulator = map_reduce_range(start, end, identity_chunk, map, reduce);
slots_chunk.write(chunk_index - 1, accumulator);
});
if result.is_err() {
completion.mark_failed();
}
};
if let Err(error) =
self.schedule_scoped_job::<C, _>(priority, locality_hint, scoped_job)
{
match error {
ExecutorError::ResourceExhausted(_) => {
self.record_admission_caller_run();
inline_result = execute_catching_panic(|| {
let accumulator =
map_reduce_range(start, end, identity.clone(), map, reduce);
slots.write(chunk_index - 1, accumulator);
});
if inline_result.is_err() {
break;
}
}
other => {
schedule_result = Err(other);
break;
}
}
}
}
let caller_result = if schedule_result.is_ok() && inline_result.is_ok() {
execute_catching_panic(|| {
let _region = IndexedRegionGuard::enter();
map_reduce_range(0, caller_end, identity.clone(), map, reduce)
})
} else {
Ok(identity.clone())
};
self.drain_scope(&state);
if state.failed_tasks.load(Ordering::Acquire) || inline_result.is_err() {
Err(ExecutorError::SpawnFailed(
moirai_core::error::TaskError::Panicked,
))
} else {
schedule_result?;
Ok(slots.reduce(caller_result?, reduce))
}
}
}