1use super::DisjointMutPtr;
29use crate::policy::{ExecutionPolicy, Parallel};
30use moirai_core::error::{ExecutorError, ExecutorResult};
31use moirai_executor::{global, HybridExecutor, SchedulerScope, SyncTask};
32use std::sync::Mutex;
33
34enum Branch<F, R> {
42 Pending(F),
44 Claimed,
47 Done(R),
49}
50
51impl<F, R> Branch<F, R>
52where
53 F: FnOnce() -> R,
54{
55 fn claim(&mut self) -> Option<F> {
57 match std::mem::replace(self, Self::Claimed) {
58 Self::Pending(branch) => Some(branch),
59 other => {
60 *self = other;
61 None
62 }
63 }
64 }
65
66 fn complete(&mut self, result: R) {
67 *self = Self::Done(result);
68 }
69
70 fn run_shared(slot: &Mutex<Self>) {
76 let Some(branch) = lock(slot).claim() else {
77 return;
78 };
79 let result = branch();
80 lock(slot).complete(result);
81 }
82
83 fn run_here(&mut self) {
85 let Some(branch) = self.claim() else {
86 return;
87 };
88 let result = branch();
89 self.complete(result);
90 }
91
92 fn into_result(self) -> R {
94 match self {
95 Self::Done(result) => result,
96 _ => panic!("invariant: a join branch neither ran nor reported failure"),
97 }
98 }
99}
100
101fn lock<T>(slot: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
105 slot.lock()
106 .unwrap_or_else(std::sync::PoisonError::into_inner)
107}
108
109fn lock_owned<T>(slot: Mutex<T>) -> T {
111 slot.into_inner()
112 .unwrap_or_else(std::sync::PoisonError::into_inner)
113}
114
115pub fn join_with<P, A, B, RA, RB>(left: A, right: B) -> (RA, RB)
131where
132 P: ExecutionPolicy,
133 A: FnOnce() -> RA + Send,
134 B: FnOnce() -> RB,
135 RA: Send,
136{
137 if !P::parallelize_pair() {
138 return (left(), right());
139 }
140
141 join_on(global(), left, right)
142}
143
144pub(crate) fn join_on<A, B, RA, RB>(executor: &HybridExecutor, left: A, right: B) -> (RA, RB)
149where
150 A: FnOnce() -> RA + Send,
151 B: FnOnce() -> RB,
152 RA: Send,
153{
154 let left_slot = Mutex::new(Branch::Pending(left));
158 let mut right_slot = Branch::Pending(right);
159
160 let forked = executor.scope::<SyncTask, _>(|scope| {
161 scope.spawn(|_| Branch::run_shared(&left_slot))?;
162 scope.flush()?;
165 right_slot.run_here();
166 Ok(())
167 });
168
169 match forked {
170 Ok(()) => {}
171 Err(ExecutorError::ShuttingDown | ExecutorError::ResourceExhausted(_)) => {
175 Branch::run_shared(&left_slot);
176 right_slot.run_here();
177 }
178 Err(error) => panic!("invariant: scheduled join branch failed ({error})"),
179 }
180
181 (
182 lock_owned(left_slot).into_result(),
183 right_slot.into_result(),
184 )
185}
186
187pub fn join<A, B, RA, RB>(left: A, right: B) -> (RA, RB)
191where
192 A: FnOnce() -> RA + Send,
193 B: FnOnce() -> RB,
194 RA: Send,
195{
196 join_with::<crate::Adaptive, _, _, _, _>(left, right)
197}
198
199pub struct Scope<'scope> {
211 inner: &'scope SchedulerScope<'scope, SyncTask>,
212}
213
214impl<'scope> Scope<'scope> {
215 #[inline]
228 pub fn spawn<F>(&self, task: F)
229 where
230 F: FnOnce() + Send + 'scope,
231 {
232 self.inner
233 .spawn(move |_| task())
234 .expect("moirai global executor: scope spawn");
235 }
236}
237
238#[inline]
268pub fn scope<F, R>(body: F) -> R
269where
270 F: for<'scope> FnOnce(&Scope<'scope>) -> R,
271 R: Send,
272{
273 let mut result = None;
274 global()
275 .scope::<SyncTask, _>(|inner| {
276 let scope = Scope { inner };
277 result = Some(body(&scope));
278 ExecutorResult::Ok(())
279 })
280 .expect("moirai global executor: scope");
281
282 result.expect("scoped body must complete")
283}
284
285pub fn for_each_with<P, T, F>(data: &[T], f: F)
287where
288 P: ExecutionPolicy,
289 T: Sync,
290 F: Fn(&T) + Send + Sync,
291{
292 let n = data.len();
293 if n == 0 {
294 return;
295 }
296 if !P::parallelize(n) {
297 data.iter().for_each(f);
298 return;
299 }
300 let f = &f;
301 global()
302 .for_each_indexed::<SyncTask, _>(n, move |i| f(&data[i]))
303 .expect("moirai global executor: for_each_with");
304}
305
306pub fn for_each_mut_with<P, T, F>(data: &mut [T], f: F)
308where
309 P: ExecutionPolicy,
310 T: Send,
311 F: Fn(&mut T) + Send + Sync,
312{
313 let n = data.len();
314 if n == 0 {
315 return;
316 }
317 if !P::parallelize(n) {
318 data.iter_mut().for_each(f);
319 return;
320 }
321 let base = DisjointMutPtr(data.as_mut_ptr());
322 let f = &f;
323 global()
324 .for_each_indexed::<SyncTask, _>(n, move |i| {
325 f(unsafe { base.get_mut(i) });
329 })
330 .expect("moirai global executor: for_each_mut_with");
331}
332
333pub fn enumerate_with<P, T, F>(data: &[T], f: F)
335where
336 P: ExecutionPolicy,
337 T: Sync,
338 F: Fn(usize, &T) + Send + Sync,
339{
340 let n = data.len();
341 if n == 0 {
342 return;
343 }
344 if !P::parallelize(n) {
345 data.iter().enumerate().for_each(|(i, x)| f(i, x));
346 return;
347 }
348 let f = &f;
349 global()
350 .for_each_indexed::<SyncTask, _>(n, move |i| f(i, &data[i]))
351 .expect("moirai global executor: enumerate_with");
352}
353
354pub fn enumerate_mut_with<P, T, F>(data: &mut [T], f: F)
357where
358 P: ExecutionPolicy,
359 T: Send,
360 F: Fn(usize, &mut T) + Send + Sync,
361{
362 let n = data.len();
363 if n == 0 {
364 return;
365 }
366 if !P::parallelize(n) {
367 data.iter_mut().enumerate().for_each(|(i, x)| f(i, x));
368 return;
369 }
370 let base = DisjointMutPtr(data.as_mut_ptr());
371 let f = &f;
372 global()
373 .for_each_indexed::<SyncTask, _>(n, move |i| {
374 f(i, unsafe { base.get_mut(i) });
377 })
378 .expect("moirai global executor: enumerate_mut_with");
379}
380
381pub fn for_each_index_with<P, F>(len: usize, f: F)
387where
388 P: ExecutionPolicy,
389 F: Fn(usize) + Send + Sync,
390{
391 if len == 0 {
392 return;
393 }
394 if !P::parallelize(len) {
395 (0..len).for_each(f);
396 return;
397 }
398 let f = &f;
399 global()
400 .for_each_indexed::<SyncTask, _>(len, f)
401 .expect("moirai global executor: for_each_index_with");
402}
403
404pub fn for_each_chunk_mut_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
410where
411 P: ExecutionPolicy,
412 T: Send,
413 F: Fn(&mut [T]) + Send + Sync,
414{
415 let n = data.len();
416 if n == 0 || chunk_size == 0 {
417 return;
418 }
419 let num_chunks = n.div_ceil(chunk_size);
420 if !P::parallelize(n) || num_chunks <= 1 {
421 data.chunks_mut(chunk_size).for_each(&f);
422 return;
423 }
424 let base = DisjointMutPtr(data.as_mut_ptr());
425 let f = &f;
426 global()
427 .for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
428 let start = c * chunk_size;
429 if start >= n {
430 return;
431 }
432 let end = (start + chunk_size).min(n);
433 let chunk =
436 unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
437 f(chunk);
438 })
439 .expect("moirai global executor: for_each_chunk_mut_with");
440}
441
442pub fn for_each_chunk_mut_with_state<P, T, S, Init, F>(
450 data: &mut [T],
451 chunk_size: usize,
452 init: Init,
453 f: F,
454) where
455 P: ExecutionPolicy,
456 T: Send,
457 S: Send,
458 Init: Fn() -> S + Send + Sync,
459 F: Fn(&mut S, &mut [T]) + Send + Sync,
460{
461 let n = data.len();
462 if n == 0 || chunk_size == 0 {
463 return;
464 }
465 let num_chunks = n.div_ceil(chunk_size);
466 if !P::parallelize(n) || num_chunks <= 1 {
467 let mut state = init();
468 for chunk in data.chunks_mut(chunk_size) {
469 f(&mut state, chunk);
470 }
471 return;
472 }
473
474 let workers = themis::CpuTopology::detect()
475 .map(|topology| topology.logical_processors())
476 .or_else(|| std::thread::available_parallelism().ok().map(|n| n.get()))
477 .unwrap_or(1)
478 .min(num_chunks)
479 .max(1);
480 let chunks_per_worker = num_chunks.div_ceil(workers);
481 let base = DisjointMutPtr(data.as_mut_ptr());
482 let init = &init;
483 let f = &f;
484 global()
485 .for_each_indexed::<SyncTask, _>(workers, move |worker| {
486 let first_chunk = worker * chunks_per_worker;
487 let last_chunk = ((worker + 1) * chunks_per_worker).min(num_chunks);
488 if first_chunk >= last_chunk {
489 return;
490 }
491 let mut state = init();
492 for chunk_index in first_chunk..last_chunk {
493 let start = chunk_index * chunk_size;
494 let end = (start + chunk_size).min(n);
495 let chunk =
498 unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
499 f(&mut state, chunk);
500 }
501 })
502 .expect("moirai global executor: for_each_chunk_mut_with_state");
503}
504
505pub fn for_each_chunk_pair_mut_enumerated_with<P, A, B, F>(
514 a: &mut [A],
515 b: &mut [B],
516 chunk_size: usize,
517 f: F,
518) where
519 P: ExecutionPolicy,
520 A: Send,
521 B: Send,
522 F: Fn(usize, &mut [A], &mut [B]) + Send + Sync,
523{
524 let na = a.len();
525 let nb = b.len();
526 if chunk_size == 0 || na == 0 {
527 return;
528 }
529 let num_chunks = na.div_ceil(chunk_size);
530 if !P::parallelize(na) || num_chunks <= 1 {
531 a.chunks_mut(chunk_size)
532 .zip(b.chunks_mut(chunk_size))
533 .enumerate()
534 .for_each(|(i, (ca, cb))| f(i, ca, cb));
535 return;
536 }
537 let abase = DisjointMutPtr(a.as_mut_ptr());
538 let bbase = DisjointMutPtr(b.as_mut_ptr());
539 let f = &f;
540 global()
541 .for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
542 let start = c * chunk_size;
543 if start >= na || start >= nb {
544 return;
545 }
546 let ea = (start + chunk_size).min(na);
547 let eb = (start + chunk_size).min(nb);
548 let ca =
552 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
553 let cb =
554 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
555 f(c, ca, cb);
556 })
557 .expect("moirai global executor: for_each_chunk_pair_mut_enumerated_with");
558}
559
560pub fn for_each_chunk_quad_mut_enumerated_with<P, A, B, C, D, F>(
568 a: &mut [A],
569 b: &mut [B],
570 c: &mut [C],
571 d: &mut [D],
572 chunk_size: usize,
573 f: F,
574) where
575 P: ExecutionPolicy,
576 A: Send,
577 B: Send,
578 C: Send,
579 D: Send,
580 F: Fn(usize, &mut [A], &mut [B], &mut [C], &mut [D]) + Send + Sync,
581{
582 let na = a.len();
583 let nb = b.len();
584 let nc = c.len();
585 let nd = d.len();
586 assert_eq!(na, nb, "quad chunk buffers must have equal lengths");
587 assert_eq!(na, nc, "quad chunk buffers must have equal lengths");
588 assert_eq!(na, nd, "quad chunk buffers must have equal lengths");
589 if chunk_size == 0 || na == 0 {
590 return;
591 }
592 let num_chunks = na.div_ceil(chunk_size);
593 if !P::parallelize(na) || num_chunks <= 1 {
594 a.chunks_mut(chunk_size)
595 .zip(b.chunks_mut(chunk_size))
596 .zip(c.chunks_mut(chunk_size))
597 .zip(d.chunks_mut(chunk_size))
598 .enumerate()
599 .for_each(|(i, (((ca, cb), cc), cd))| f(i, ca, cb, cc, cd));
600 return;
601 }
602 let abase = DisjointMutPtr(a.as_mut_ptr());
603 let bbase = DisjointMutPtr(b.as_mut_ptr());
604 let cbase = DisjointMutPtr(c.as_mut_ptr());
605 let dbase = DisjointMutPtr(d.as_mut_ptr());
606 let f = &f;
607 global()
608 .for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
609 let start = chunk_index * chunk_size;
610 if start >= na || start >= nb || start >= nc || start >= nd {
611 return;
612 }
613 let ea = (start + chunk_size).min(na);
614 let eb = (start + chunk_size).min(nb);
615 let ec = (start + chunk_size).min(nc);
616 let ed = (start + chunk_size).min(nd);
617 let ca =
623 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
624 let cb =
625 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
626 let cc =
627 unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
628 let cd =
629 unsafe { core::slice::from_raw_parts_mut(dbase.base().add(start), ed - start) };
630 f(chunk_index, ca, cb, cc, cd);
631 })
632 .expect("moirai global executor: for_each_chunk_quad_mut_enumerated_with");
633}
634
635pub fn for_each_chunk_triple_mut_enumerated_with<P, A, B, C, F>(
641 a: &mut [A],
642 b: &mut [B],
643 c: &mut [C],
644 chunk_size: usize,
645 f: F,
646) where
647 P: ExecutionPolicy,
648 A: Send,
649 B: Send,
650 C: Send,
651 F: Fn(usize, &mut [A], &mut [B], &mut [C]) + Send + Sync,
652{
653 let na = a.len();
654 let nb = b.len();
655 let nc = c.len();
656 assert_eq!(na, nb, "triple chunk buffers must have equal lengths");
657 assert_eq!(na, nc, "triple chunk buffers must have equal lengths");
658 if chunk_size == 0 || na == 0 {
659 return;
660 }
661 let num_chunks = na.div_ceil(chunk_size);
662 if !P::parallelize(na) || num_chunks <= 1 {
663 a.chunks_mut(chunk_size)
664 .zip(b.chunks_mut(chunk_size))
665 .zip(c.chunks_mut(chunk_size))
666 .enumerate()
667 .for_each(|(i, ((ca, cb), cc))| f(i, ca, cb, cc));
668 return;
669 }
670 let abase = DisjointMutPtr(a.as_mut_ptr());
671 let bbase = DisjointMutPtr(b.as_mut_ptr());
672 let cbase = DisjointMutPtr(c.as_mut_ptr());
673 let f = &f;
674 global()
675 .for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
676 let start = chunk_index * chunk_size;
677 if start >= na || start >= nb || start >= nc {
678 return;
679 }
680 let ea = (start + chunk_size).min(na);
681 let eb = (start + chunk_size).min(nb);
682 let ec = (start + chunk_size).min(nc);
683 let ca =
689 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
690 let cb =
691 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
692 let cc =
693 unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
694 f(chunk_index, ca, cb, cc);
695 })
696 .expect("moirai global executor: for_each_chunk_triple_mut_enumerated_with");
697}
698
699pub fn for_each_chunk_mut_enumerated_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
703where
704 P: ExecutionPolicy,
705 T: Send,
706 F: Fn(usize, &mut [T]) + Send + Sync,
707{
708 let n = data.len();
709 if n == 0 || chunk_size == 0 {
710 return;
711 }
712 let num_chunks = n.div_ceil(chunk_size);
713 if !P::parallelize(n) || num_chunks <= 1 {
714 data.chunks_mut(chunk_size)
715 .enumerate()
716 .for_each(|(i, c)| f(i, c));
717 return;
718 }
719 let base = DisjointMutPtr(data.as_mut_ptr());
720 let f = &f;
721 global()
722 .for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
723 let start = c * chunk_size;
724 if start >= n {
725 return;
726 }
727 let end = (start + chunk_size).min(n);
728 let chunk =
731 unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
732 f(c, chunk);
733 })
734 .expect("moirai global executor: for_each_chunk_mut_enumerated_with");
735}
736
737pub fn map_collect_with<P, T, R, F>(data: &[T], f: F) -> Vec<R>
740where
741 P: ExecutionPolicy,
742 T: Sync,
743 R: Send,
744 F: Fn(&T) -> R + Send + Sync,
745{
746 let n = data.len();
747 if !P::parallelize(n) {
748 return data.iter().map(f).collect();
749 }
750 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
751 unsafe {
754 out.set_len(n);
755 }
756 enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
757 slot.write(f(&data[i]));
758 });
759 let mut out = core::mem::ManuallyDrop::new(out);
761 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
762}
763
764pub fn map_reduce_with<P, T, R, M, Rd>(data: &[T], identity: R, map: M, reduce: Rd) -> R
769where
770 P: ExecutionPolicy,
771 T: Sync,
772 R: Send + Sync + Clone,
773 M: Fn(&T) -> R + Send + Sync,
774 Rd: Fn(R, R) -> R + Send + Sync,
775{
776 let n = data.len();
777 if n == 0 || !P::parallelize(n) {
778 let mut acc = identity;
779 for item in data {
780 acc = reduce(acc, map(item));
781 }
782 return acc;
783 }
784 let map = ↦
785 let reduce = &reduce;
786 global()
789 .map_reduce_indexed::<SyncTask, _, _, _>(n, identity, move |i| map(&data[i]), reduce)
790 .expect("moirai global executor: map_reduce_with")
791}
792
793pub fn fold_reduce_with<P, A, Init, Fold, Red>(len: usize, init: Init, fold: Fold, reduce: Red) -> A
802where
803 P: ExecutionPolicy,
804 A: Send,
805 Init: Fn() -> A + Send + Sync,
806 Fold: Fn(A, usize) -> A + Send + Sync,
807 Red: Fn(A, A) -> A,
808{
809 if len == 0 {
810 return init();
811 }
812 if !P::parallelize(len) {
813 let mut acc = init();
814 for i in 0..len {
815 acc = fold(acc, i);
816 }
817 return acc;
818 }
819 let workers = themis::CpuTopology::detect()
820 .map(|topology| topology.logical_processors())
821 .or_else(|| std::thread::available_parallelism().ok().map(|n| n.get()))
822 .unwrap_or(1)
823 .max(1);
824 let chunks = workers.min(len).max(1);
825 let chunk = len.div_ceil(chunks);
826 let mut slots: Vec<Option<A>> = (0..chunks).map(|_| None).collect();
827 let base = DisjointMutPtr(slots.as_mut_ptr());
828 let init_ref = &init;
829 let fold_ref = &fold;
830 global()
831 .for_each_indexed::<SyncTask, _>(chunks, move |ci| {
832 let start = ci * chunk;
833 if start >= len {
834 return;
835 }
836 let end = (start + chunk).min(len);
837 let mut acc = init_ref();
838 for i in start..end {
839 acc = fold_ref(acc, i);
840 }
841 unsafe {
844 *base.get_mut(ci) = Some(acc);
845 }
846 })
847 .expect("moirai global executor: fold_reduce_with");
848 slots
849 .into_iter()
850 .flatten()
851 .reduce(reduce)
852 .unwrap_or_else(init)
853}
854
855pub fn map_collect_index_with<P, R, Map>(len: usize, map: Map) -> Vec<R>
862where
863 P: ExecutionPolicy,
864 R: Send,
865 Map: Fn(usize) -> R + Send + Sync,
866{
867 if !P::parallelize(len) {
868 return (0..len).map(map).collect();
869 }
870 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(len);
871 unsafe {
873 out.set_len(len);
874 }
875 enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
876 slot.write(map(i));
877 });
878 let mut out = core::mem::ManuallyDrop::new(out);
880 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), len, out.capacity()) }
881}
882
883pub fn map_collect_mut_with<P, T, R, F>(data: &mut [T], f: F) -> Vec<R>
890where
891 P: ExecutionPolicy,
892 T: Send,
893 R: Send,
894 F: Fn(usize, &mut T) -> R + Send + Sync,
895{
896 let n = data.len();
897 if !P::parallelize(n) {
898 return data.iter_mut().enumerate().map(|(i, x)| f(i, x)).collect();
899 }
900 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
901 unsafe {
903 out.set_len(n);
904 }
905 let data_ptr = DisjointMutPtr(data.as_mut_ptr());
906 let out_ptr = DisjointMutPtr(out.as_mut_ptr());
907 let f = &f;
908 global()
909 .for_each_indexed::<SyncTask, _>(n, move |i| {
910 let elem = unsafe { data_ptr.get_mut(i) };
913 let result = f(i, elem);
914 unsafe { out_ptr.get_mut(i).write(result) };
915 })
916 .expect("moirai global executor: map_collect_mut_with");
917 let mut out = core::mem::ManuallyDrop::new(out);
919 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
920}
921
922pub fn reduce_index_with<P, R, Map, Red>(len: usize, identity: R, map: Map, reduce: Red) -> R
930where
931 P: ExecutionPolicy,
932 R: Send + Sync + Clone,
933 Map: Fn(usize) -> R + Send + Sync,
934 Red: Fn(R, R) -> R + Send + Sync,
935{
936 if len == 0 || !P::parallelize(len) {
937 let mut acc = identity;
938 for i in 0..len {
939 acc = reduce(acc, map(i));
940 }
941 return acc;
942 }
943 global()
944 .map_reduce_indexed::<SyncTask, _, _, _>(len, identity, map, reduce)
945 .expect("moirai global executor: reduce_index_with")
946}