use core::cell::Cell;
use core::num::NonZeroUsize;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ExecutionPolicy {
AmbientRayon,
Sequential,
Rayon { max_threads: NonZeroUsize },
}
thread_local! {
static ACTIVE_POLICY: Cell<ExecutionPolicy> = const { Cell::new(ExecutionPolicy::AmbientRayon) };
static ACTIVE_FANOUT: Cell<bool> = const { Cell::new(false) };
}
fn restrict(outer: ExecutionPolicy, inner: ExecutionPolicy) -> ExecutionPolicy {
match (outer, inner) {
(ExecutionPolicy::Sequential, _) | (_, ExecutionPolicy::Sequential) => {
ExecutionPolicy::Sequential
}
(ExecutionPolicy::AmbientRayon, policy) | (policy, ExecutionPolicy::AmbientRayon) => policy,
(
ExecutionPolicy::Rayon { max_threads: outer },
ExecutionPolicy::Rayon { max_threads: inner },
) => ExecutionPolicy::Rayon {
max_threads: outer.min(inner),
},
}
}
#[derive(Clone, Copy)]
struct ExecutionState {
policy: ExecutionPolicy,
fanout_active: bool,
}
struct StateGuard {
previous: ExecutionState,
}
impl Drop for StateGuard {
fn drop(&mut self) {
set_state(self.previous);
}
}
fn state() -> ExecutionState {
ExecutionState {
policy: ACTIVE_POLICY.with(Cell::get),
fanout_active: ACTIVE_FANOUT.with(Cell::get),
}
}
fn set_state(state: ExecutionState) {
ACTIVE_POLICY.with(|active| active.set(state.policy));
ACTIVE_FANOUT.with(|active| active.set(state.fanout_active));
}
#[cfg(feature = "parallel")]
fn with_state<R>(next: ExecutionState, operation: impl FnOnce() -> R) -> R {
let previous = state();
set_state(next);
let _guard = StateGuard { previous };
operation()
}
#[inline]
pub fn with_execution_policy<R>(policy: ExecutionPolicy, operation: impl FnOnce() -> R) -> R {
let policy = match policy {
ExecutionPolicy::AmbientRayon => return operation(),
policy => policy,
};
let previous = state();
set_state(ExecutionState {
policy: restrict(previous.policy, policy),
fanout_active: previous.fanout_active,
});
let _guard = StateGuard { previous };
operation()
}
#[cfg(feature = "parallel")]
pub(crate) fn active_policy() -> ExecutionPolicy {
ACTIVE_POLICY.with(Cell::get)
}
#[cfg(feature = "parallel")]
pub(crate) fn fanout_active() -> bool {
ACTIVE_FANOUT.with(Cell::get)
}
#[cfg(feature = "parallel")]
#[inline(always)]
pub(crate) fn with_owned_execution<R>(
policy: ExecutionPolicy,
fanout_active: bool,
operation: impl FnOnce() -> R,
) -> R {
match policy {
ExecutionPolicy::AmbientRayon => operation(),
ExecutionPolicy::Sequential | ExecutionPolicy::Rayon { .. } => {
let previous = state();
with_state(
ExecutionState {
policy: restrict(previous.policy, policy),
fanout_active: previous.fanout_active || fanout_active,
},
operation,
)
}
}
}
#[cfg(feature = "parallel")]
pub(crate) fn with_scheduler_suspended<R>(operation: impl FnOnce() -> R) -> R {
with_state(
ExecutionState {
policy: ExecutionPolicy::AmbientRayon,
fanout_active: false,
},
operation,
)
}
#[cfg(feature = "parallel")]
pub(crate) fn permutation_copy_parallel_eligible(
policy: ExecutionPolicy,
fanout_active: bool,
current_pool_threads: usize,
) -> bool {
if fanout_active || current_pool_threads <= 1 {
return false;
}
match policy {
ExecutionPolicy::AmbientRayon => true,
ExecutionPolicy::Sequential => false,
ExecutionPolicy::Rayon { max_threads } => current_pool_threads <= max_threads.get(),
}
}
#[cfg(feature = "parallel")]
pub fn rayon_threads() -> usize {
if fanout_active() {
return 1;
}
match active_policy() {
ExecutionPolicy::Sequential => 1,
ExecutionPolicy::AmbientRayon => crate::threading::current_pool_threads(),
ExecutionPolicy::Rayon { max_threads } => {
crate::threading::current_pool_threads().min(max_threads.get())
}
}
}
#[cfg(test)]
#[path = "execution_policy/tests/default_tests.rs"]
mod default_tests;
#[cfg(all(test, feature = "parallel"))]
#[path = "execution_policy/tests/tests.rs"]
mod tests;