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 = std::thread::available_parallelism()
475 .map(|count| count.get())
476 .unwrap_or(1)
477 .min(num_chunks)
478 .max(1);
479 let chunks_per_worker = num_chunks.div_ceil(workers);
480 let base = DisjointMutPtr(data.as_mut_ptr());
481 let init = &init;
482 let f = &f;
483 global()
484 .for_each_indexed::<SyncTask, _>(workers, move |worker| {
485 let first_chunk = worker * chunks_per_worker;
486 let last_chunk = ((worker + 1) * chunks_per_worker).min(num_chunks);
487 if first_chunk >= last_chunk {
488 return;
489 }
490 let mut state = init();
491 for chunk_index in first_chunk..last_chunk {
492 let start = chunk_index * chunk_size;
493 let end = (start + chunk_size).min(n);
494 let chunk =
497 unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
498 f(&mut state, chunk);
499 }
500 })
501 .expect("moirai global executor: for_each_chunk_mut_with_state");
502}
503
504pub fn for_each_chunk_pair_mut_enumerated_with<P, A, B, F>(
513 a: &mut [A],
514 b: &mut [B],
515 chunk_size: usize,
516 f: F,
517) where
518 P: ExecutionPolicy,
519 A: Send,
520 B: Send,
521 F: Fn(usize, &mut [A], &mut [B]) + Send + Sync,
522{
523 let na = a.len();
524 let nb = b.len();
525 if chunk_size == 0 || na == 0 {
526 return;
527 }
528 let num_chunks = na.div_ceil(chunk_size);
529 if !P::parallelize(na) || num_chunks <= 1 {
530 a.chunks_mut(chunk_size)
531 .zip(b.chunks_mut(chunk_size))
532 .enumerate()
533 .for_each(|(i, (ca, cb))| f(i, ca, cb));
534 return;
535 }
536 let abase = DisjointMutPtr(a.as_mut_ptr());
537 let bbase = DisjointMutPtr(b.as_mut_ptr());
538 let f = &f;
539 global()
540 .for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
541 let start = c * chunk_size;
542 if start >= na || start >= nb {
543 return;
544 }
545 let ea = (start + chunk_size).min(na);
546 let eb = (start + chunk_size).min(nb);
547 let ca =
551 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
552 let cb =
553 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
554 f(c, ca, cb);
555 })
556 .expect("moirai global executor: for_each_chunk_pair_mut_enumerated_with");
557}
558
559pub fn for_each_chunk_quad_mut_enumerated_with<P, A, B, C, D, F>(
567 a: &mut [A],
568 b: &mut [B],
569 c: &mut [C],
570 d: &mut [D],
571 chunk_size: usize,
572 f: F,
573) where
574 P: ExecutionPolicy,
575 A: Send,
576 B: Send,
577 C: Send,
578 D: Send,
579 F: Fn(usize, &mut [A], &mut [B], &mut [C], &mut [D]) + Send + Sync,
580{
581 let na = a.len();
582 let nb = b.len();
583 let nc = c.len();
584 let nd = d.len();
585 assert_eq!(na, nb, "quad chunk buffers must have equal lengths");
586 assert_eq!(na, nc, "quad chunk buffers must have equal lengths");
587 assert_eq!(na, nd, "quad chunk buffers must have equal lengths");
588 if chunk_size == 0 || na == 0 {
589 return;
590 }
591 let num_chunks = na.div_ceil(chunk_size);
592 if !P::parallelize(na) || num_chunks <= 1 {
593 a.chunks_mut(chunk_size)
594 .zip(b.chunks_mut(chunk_size))
595 .zip(c.chunks_mut(chunk_size))
596 .zip(d.chunks_mut(chunk_size))
597 .enumerate()
598 .for_each(|(i, (((ca, cb), cc), cd))| f(i, ca, cb, cc, cd));
599 return;
600 }
601 let abase = DisjointMutPtr(a.as_mut_ptr());
602 let bbase = DisjointMutPtr(b.as_mut_ptr());
603 let cbase = DisjointMutPtr(c.as_mut_ptr());
604 let dbase = DisjointMutPtr(d.as_mut_ptr());
605 let f = &f;
606 global()
607 .for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
608 let start = chunk_index * chunk_size;
609 if start >= na || start >= nb || start >= nc || start >= nd {
610 return;
611 }
612 let ea = (start + chunk_size).min(na);
613 let eb = (start + chunk_size).min(nb);
614 let ec = (start + chunk_size).min(nc);
615 let ed = (start + chunk_size).min(nd);
616 let ca =
622 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
623 let cb =
624 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
625 let cc =
626 unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
627 let cd =
628 unsafe { core::slice::from_raw_parts_mut(dbase.base().add(start), ed - start) };
629 f(chunk_index, ca, cb, cc, cd);
630 })
631 .expect("moirai global executor: for_each_chunk_quad_mut_enumerated_with");
632}
633
634pub fn for_each_chunk_triple_mut_enumerated_with<P, A, B, C, F>(
640 a: &mut [A],
641 b: &mut [B],
642 c: &mut [C],
643 chunk_size: usize,
644 f: F,
645) where
646 P: ExecutionPolicy,
647 A: Send,
648 B: Send,
649 C: Send,
650 F: Fn(usize, &mut [A], &mut [B], &mut [C]) + Send + Sync,
651{
652 let na = a.len();
653 let nb = b.len();
654 let nc = c.len();
655 assert_eq!(na, nb, "triple chunk buffers must have equal lengths");
656 assert_eq!(na, nc, "triple chunk buffers must have equal lengths");
657 if chunk_size == 0 || na == 0 {
658 return;
659 }
660 let num_chunks = na.div_ceil(chunk_size);
661 if !P::parallelize(na) || num_chunks <= 1 {
662 a.chunks_mut(chunk_size)
663 .zip(b.chunks_mut(chunk_size))
664 .zip(c.chunks_mut(chunk_size))
665 .enumerate()
666 .for_each(|(i, ((ca, cb), cc))| f(i, ca, cb, cc));
667 return;
668 }
669 let abase = DisjointMutPtr(a.as_mut_ptr());
670 let bbase = DisjointMutPtr(b.as_mut_ptr());
671 let cbase = DisjointMutPtr(c.as_mut_ptr());
672 let f = &f;
673 global()
674 .for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
675 let start = chunk_index * chunk_size;
676 if start >= na || start >= nb || start >= nc {
677 return;
678 }
679 let ea = (start + chunk_size).min(na);
680 let eb = (start + chunk_size).min(nb);
681 let ec = (start + chunk_size).min(nc);
682 let ca =
688 unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
689 let cb =
690 unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
691 let cc =
692 unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
693 f(chunk_index, ca, cb, cc);
694 })
695 .expect("moirai global executor: for_each_chunk_triple_mut_enumerated_with");
696}
697
698pub fn for_each_chunk_mut_enumerated_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
702where
703 P: ExecutionPolicy,
704 T: Send,
705 F: Fn(usize, &mut [T]) + Send + Sync,
706{
707 let n = data.len();
708 if n == 0 || chunk_size == 0 {
709 return;
710 }
711 let num_chunks = n.div_ceil(chunk_size);
712 if !P::parallelize(n) || num_chunks <= 1 {
713 data.chunks_mut(chunk_size)
714 .enumerate()
715 .for_each(|(i, c)| f(i, c));
716 return;
717 }
718 let base = DisjointMutPtr(data.as_mut_ptr());
719 let f = &f;
720 global()
721 .for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
722 let start = c * chunk_size;
723 if start >= n {
724 return;
725 }
726 let end = (start + chunk_size).min(n);
727 let chunk =
730 unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
731 f(c, chunk);
732 })
733 .expect("moirai global executor: for_each_chunk_mut_enumerated_with");
734}
735
736pub fn map_collect_with<P, T, R, F>(data: &[T], f: F) -> Vec<R>
739where
740 P: ExecutionPolicy,
741 T: Sync,
742 R: Send,
743 F: Fn(&T) -> R + Send + Sync,
744{
745 let n = data.len();
746 if !P::parallelize(n) {
747 return data.iter().map(f).collect();
748 }
749 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
750 unsafe {
753 out.set_len(n);
754 }
755 enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
756 slot.write(f(&data[i]));
757 });
758 let mut out = core::mem::ManuallyDrop::new(out);
760 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
761}
762
763pub fn map_reduce_with<P, T, R, M, Rd>(data: &[T], identity: R, map: M, reduce: Rd) -> R
768where
769 P: ExecutionPolicy,
770 T: Sync,
771 R: Send + Sync + Clone,
772 M: Fn(&T) -> R + Send + Sync,
773 Rd: Fn(R, R) -> R + Send + Sync,
774{
775 let n = data.len();
776 if n == 0 || !P::parallelize(n) {
777 let mut acc = identity;
778 for item in data {
779 acc = reduce(acc, map(item));
780 }
781 return acc;
782 }
783 let map = ↦
784 let reduce = &reduce;
785 global()
788 .map_reduce_indexed::<SyncTask, _, _, _>(n, identity, move |i| map(&data[i]), reduce)
789 .expect("moirai global executor: map_reduce_with")
790}
791
792pub fn fold_reduce_with<P, A, Init, Fold, Red>(len: usize, init: Init, fold: Fold, reduce: Red) -> A
801where
802 P: ExecutionPolicy,
803 A: Send,
804 Init: Fn() -> A + Send + Sync,
805 Fold: Fn(A, usize) -> A + Send + Sync,
806 Red: Fn(A, A) -> A,
807{
808 if len == 0 {
809 return init();
810 }
811 if !P::parallelize(len) {
812 let mut acc = init();
813 for i in 0..len {
814 acc = fold(acc, i);
815 }
816 return acc;
817 }
818 let workers = std::thread::available_parallelism()
819 .map(|n| n.get())
820 .unwrap_or(1);
821 let chunks = workers.min(len).max(1);
822 let chunk = len.div_ceil(chunks);
823 let mut slots: Vec<Option<A>> = (0..chunks).map(|_| None).collect();
824 let base = DisjointMutPtr(slots.as_mut_ptr());
825 let init_ref = &init;
826 let fold_ref = &fold;
827 global()
828 .for_each_indexed::<SyncTask, _>(chunks, move |ci| {
829 let start = ci * chunk;
830 if start >= len {
831 return;
832 }
833 let end = (start + chunk).min(len);
834 let mut acc = init_ref();
835 for i in start..end {
836 acc = fold_ref(acc, i);
837 }
838 unsafe {
841 *base.get_mut(ci) = Some(acc);
842 }
843 })
844 .expect("moirai global executor: fold_reduce_with");
845 slots
846 .into_iter()
847 .flatten()
848 .reduce(reduce)
849 .unwrap_or_else(init)
850}
851
852pub fn map_collect_index_with<P, R, Map>(len: usize, map: Map) -> Vec<R>
859where
860 P: ExecutionPolicy,
861 R: Send,
862 Map: Fn(usize) -> R + Send + Sync,
863{
864 if !P::parallelize(len) {
865 return (0..len).map(map).collect();
866 }
867 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(len);
868 unsafe {
870 out.set_len(len);
871 }
872 enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
873 slot.write(map(i));
874 });
875 let mut out = core::mem::ManuallyDrop::new(out);
877 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), len, out.capacity()) }
878}
879
880pub fn map_collect_mut_with<P, T, R, F>(data: &mut [T], f: F) -> Vec<R>
887where
888 P: ExecutionPolicy,
889 T: Send,
890 R: Send,
891 F: Fn(usize, &mut T) -> R + Send + Sync,
892{
893 let n = data.len();
894 if !P::parallelize(n) {
895 return data.iter_mut().enumerate().map(|(i, x)| f(i, x)).collect();
896 }
897 let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
898 unsafe {
900 out.set_len(n);
901 }
902 let data_ptr = DisjointMutPtr(data.as_mut_ptr());
903 let out_ptr = DisjointMutPtr(out.as_mut_ptr());
904 let f = &f;
905 global()
906 .for_each_indexed::<SyncTask, _>(n, move |i| {
907 let elem = unsafe { data_ptr.get_mut(i) };
910 let result = f(i, elem);
911 unsafe { out_ptr.get_mut(i).write(result) };
912 })
913 .expect("moirai global executor: map_collect_mut_with");
914 let mut out = core::mem::ManuallyDrop::new(out);
916 unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
917}
918
919pub fn reduce_index_with<P, R, Map, Red>(len: usize, identity: R, map: Map, reduce: Red) -> R
927where
928 P: ExecutionPolicy,
929 R: Send + Sync + Clone,
930 Map: Fn(usize) -> R + Send + Sync,
931 Red: Fn(R, R) -> R + Send + Sync,
932{
933 if len == 0 || !P::parallelize(len) {
934 let mut acc = identity;
935 for i in 0..len {
936 acc = reduce(acc, map(i));
937 }
938 return acc;
939 }
940 global()
941 .map_reduce_indexed::<SyncTask, _, _, _>(len, identity, map, reduce)
942 .expect("moirai global executor: reduce_index_with")
943}