strided_basic/
execution_policy.rs1use core::cell::Cell;
2use core::num::NonZeroUsize;
3
4#[derive(Clone, Copy, Debug, Eq, PartialEq)]
24pub enum ExecutionPolicy {
25 AmbientRayon,
30 Sequential,
32 Rayon { max_threads: NonZeroUsize },
37}
38
39thread_local! {
40 static ACTIVE_POLICY: Cell<ExecutionPolicy> = const { Cell::new(ExecutionPolicy::AmbientRayon) };
41 static ACTIVE_FANOUT: Cell<bool> = const { Cell::new(false) };
42}
43
44fn restrict(outer: ExecutionPolicy, inner: ExecutionPolicy) -> ExecutionPolicy {
45 match (outer, inner) {
46 (ExecutionPolicy::Sequential, _) | (_, ExecutionPolicy::Sequential) => {
47 ExecutionPolicy::Sequential
48 }
49 (ExecutionPolicy::AmbientRayon, policy) | (policy, ExecutionPolicy::AmbientRayon) => policy,
50 (
51 ExecutionPolicy::Rayon { max_threads: outer },
52 ExecutionPolicy::Rayon { max_threads: inner },
53 ) => ExecutionPolicy::Rayon {
54 max_threads: outer.min(inner),
55 },
56 }
57}
58
59#[derive(Clone, Copy)]
60struct ExecutionState {
61 policy: ExecutionPolicy,
62 fanout_active: bool,
63}
64
65struct StateGuard {
66 previous: ExecutionState,
67}
68
69impl Drop for StateGuard {
70 fn drop(&mut self) {
71 set_state(self.previous);
72 }
73}
74
75fn state() -> ExecutionState {
76 ExecutionState {
77 policy: ACTIVE_POLICY.with(Cell::get),
78 fanout_active: ACTIVE_FANOUT.with(Cell::get),
79 }
80}
81
82fn set_state(state: ExecutionState) {
83 ACTIVE_POLICY.with(|active| active.set(state.policy));
84 ACTIVE_FANOUT.with(|active| active.set(state.fanout_active));
85}
86
87#[cfg(feature = "parallel")]
88fn with_state<R>(next: ExecutionState, operation: impl FnOnce() -> R) -> R {
89 let previous = state();
90 set_state(next);
91 let _guard = StateGuard { previous };
92 operation()
93}
94
95#[inline]
122pub fn with_execution_policy<R>(policy: ExecutionPolicy, operation: impl FnOnce() -> R) -> R {
123 let policy = match policy {
124 ExecutionPolicy::AmbientRayon => return operation(),
125 policy => policy,
126 };
127 let previous = state();
128 set_state(ExecutionState {
129 policy: restrict(previous.policy, policy),
130 fanout_active: previous.fanout_active,
131 });
132 let _guard = StateGuard { previous };
133 operation()
134}
135
136#[cfg(feature = "parallel")]
137pub(crate) fn active_policy() -> ExecutionPolicy {
138 ACTIVE_POLICY.with(Cell::get)
139}
140
141#[cfg(feature = "parallel")]
142pub(crate) fn fanout_active() -> bool {
143 ACTIVE_FANOUT.with(Cell::get)
144}
145
146#[cfg(feature = "parallel")]
147#[inline(always)]
148pub(crate) fn with_owned_execution<R>(
149 policy: ExecutionPolicy,
150 fanout_active: bool,
151 operation: impl FnOnce() -> R,
152) -> R {
153 match policy {
154 ExecutionPolicy::AmbientRayon => operation(),
155 ExecutionPolicy::Sequential | ExecutionPolicy::Rayon { .. } => {
156 let previous = state();
157 with_state(
158 ExecutionState {
159 policy: restrict(previous.policy, policy),
160 fanout_active: previous.fanout_active || fanout_active,
161 },
162 operation,
163 )
164 }
165 }
166}
167
168#[cfg(feature = "parallel")]
169pub(crate) fn with_scheduler_suspended<R>(operation: impl FnOnce() -> R) -> R {
170 with_state(
171 ExecutionState {
172 policy: ExecutionPolicy::AmbientRayon,
173 fanout_active: false,
174 },
175 operation,
176 )
177}
178
179#[cfg(feature = "parallel")]
180pub(crate) fn permutation_copy_parallel_eligible(
181 policy: ExecutionPolicy,
182 fanout_active: bool,
183 current_pool_threads: usize,
184) -> bool {
185 if fanout_active || current_pool_threads <= 1 {
186 return false;
187 }
188 match policy {
189 ExecutionPolicy::AmbientRayon => true,
190 ExecutionPolicy::Sequential => false,
191 ExecutionPolicy::Rayon { max_threads } => current_pool_threads <= max_threads.get(),
192 }
193}
194
195#[cfg(feature = "parallel")]
196pub fn rayon_threads() -> usize {
197 if fanout_active() {
198 return 1;
199 }
200 match active_policy() {
201 ExecutionPolicy::Sequential => 1,
202 ExecutionPolicy::AmbientRayon => crate::threading::current_pool_threads(),
203 ExecutionPolicy::Rayon { max_threads } => {
204 crate::threading::current_pool_threads().min(max_threads.get())
205 }
206 }
207}
208
209#[cfg(test)]
210#[path = "execution_policy/tests/default_tests.rs"]
211mod default_tests;
212
213#[cfg(all(test, feature = "parallel"))]
214#[path = "execution_policy/tests/tests.rs"]
215mod tests;