use std::{
sync::{Mutex, MutexGuard},
thread,
};
use moirai_core::error::{ExecutorError, ExecutorResult};
use super::super::job::ScheduledJob;
use super::super::queue::WorkerQueueOwner;
use super::types::{set_current_worker_id, ContendedWakePolicy, SchedulerInner, WorkerState};
#[cfg(feature = "scheduler-diagnostics")]
use super::types::BoundedContendedWake;
pub(super) const WORKER_IDLE_SPIN_ATTEMPTS: usize = 256;
pub(super) const JOIN_FAST_SPIN_ATTEMPTS: usize = WORKER_IDLE_SPIN_ATTEMPTS;
pub(super) fn worker_loop<const QUEUE_CAPACITY: usize, const SPIN_LIMIT: usize>(
inner: std::sync::Arc<SchedulerInner<QUEUE_CAPACITY>>,
worker_id: usize,
mut owner: WorkerQueueOwner<QUEUE_CAPACITY>,
) {
set_current_worker_id(Some(worker_id));
let _ = inner.workers[worker_id].thread.set(thread::current());
loop {
if let Some(job) = next_job(&inner, worker_id, &mut owner) {
execute_job(&inner, worker_id, job);
continue;
}
if should_stop(&inner) {
break;
}
if spin_for_work::<QUEUE_CAPACITY, SPIN_LIMIT>(&inner, worker_id) {
continue;
}
run_idle_memory_maintenance();
wait_for_work(&inner, worker_id);
}
}
#[cfg(feature = "mnemosyne")]
melinoe::thread_cached! {
mod last_maintenance_time: std::time::Instant;
}
#[inline]
fn run_idle_memory_maintenance() {
#[cfg(feature = "mnemosyne")]
{
use mnemosyne::{LocalAllocatorSelector, MemoryBackendWrapper, StandardPolicy};
if <MemoryBackendWrapper as LocalAllocatorSelector<MemoryBackendWrapper>>::get_allocator_ptr_raw().is_null() {
return;
}
let now = std::time::Instant::now();
let should_run = if let Some(last) = last_maintenance_time::get() {
if now.duration_since(last) >= std::time::Duration::from_millis(500) {
last_maintenance_time::set(now);
true
} else {
false
}
} else {
last_maintenance_time::set(now);
true
};
if should_run {
let _ =
<MemoryBackendWrapper as LocalAllocatorSelector<MemoryBackendWrapper>>::with_allocator(
|alloc| unsafe {
alloc.periodic_defragmentation_sweep::<StandardPolicy>();
},
);
}
}
}
pub(super) fn next_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
owner: &mut WorkerQueueOwner<QUEUE_CAPACITY>,
) -> Option<ScheduledJob> {
let local = &inner.workers[worker_id];
local
.lifo_slot
.pop()
.or_else(|| owner.pop_local())
.or_else(|| steal_job(inner, worker_id, owner))
}
pub(super) fn next_shared_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
) -> Option<ScheduledJob> {
let local = &inner.workers[worker_id];
local
.lifo_slot
.pop()
.or_else(|| local.queues.steal_one())
.or_else(|| steal_shared_job(inner, worker_id))
}
fn steal_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
owner: &mut WorkerQueueOwner<QUEUE_CAPACITY>,
) -> Option<ScheduledJob> {
let worker_count = inner.workers.len();
let my_node = inner.worker_numa_nodes.get(worker_id).copied().flatten();
if let Some(node) = my_node {
let start = next_steal_start();
for offset in 0..worker_count {
let victim_index = (start.wrapping_add(offset)) % worker_count;
if victim_index == worker_id {
continue;
}
if inner.worker_numa_nodes.get(victim_index).copied().flatten() != Some(node) {
continue;
}
let victim = &inner.workers[victim_index];
if let Some(job) = owner.steal_batch(&victim.queues) {
return Some(job);
}
if let Some(job) = victim.lifo_slot.steal() {
return Some(job);
}
}
}
let start = next_steal_start();
for offset in 0..worker_count {
let victim_index = (start.wrapping_add(offset)) % worker_count;
if victim_index == worker_id {
continue;
}
let victim = &inner.workers[victim_index];
if let Some(job) = owner.steal_batch(&victim.queues) {
return Some(job);
}
if let Some(job) = victim.lifo_slot.steal() {
return Some(job);
}
}
None
}
fn steal_shared_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
) -> Option<ScheduledJob> {
let worker_count = inner.workers.len();
let start = next_steal_start();
for offset in 0..worker_count {
let victim_index = start.wrapping_add(offset) % worker_count;
if victim_index == worker_id {
continue;
}
let victim = &inner.workers[victim_index];
if let Some(job) = victim.queues.steal_one() {
return Some(job);
}
if let Some(job) = victim.lifo_slot.steal() {
return Some(job);
}
}
None
}
fn next_steal_start() -> usize {
use std::cell::Cell;
thread_local!(static RNG: Cell<u64> = const { Cell::new(0) });
RNG.with(|cell| {
let mut x = cell.get();
if x == 0 {
x = (cell as *const Cell<u64> as u64) | 1;
}
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
cell.set(x);
x as usize
})
}
pub(super) fn execute_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
job: ScheduledJob,
) {
execute_job_with_counters(
inner,
worker_id,
job,
&inner.pending_tasks,
&inner.active_workers,
);
}
pub(super) fn execute_blocking_job<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
job: ScheduledJob,
) {
execute_job_with_counters(
inner,
worker_id,
job,
&inner.blocking_pending_tasks,
&inner.blocking_active_workers,
);
}
fn execute_job_with_counters<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
job: ScheduledJob,
pending_tasks: &std::sync::atomic::AtomicUsize,
active_workers: &std::sync::atomic::AtomicUsize,
) {
use std::sync::atomic::Ordering;
active_workers.fetch_add(1, Ordering::Release);
pending_tasks.fetch_sub(1, Ordering::Release);
if job.execute(worker_id) {
inner.completed_tasks.fetch_add(1, Ordering::Relaxed);
} else {
inner.failed_tasks.fetch_add(1, Ordering::Relaxed);
}
if active_workers.fetch_sub(1, Ordering::SeqCst) == 1 {
notify_quiescent(inner);
}
}
fn should_stop<const QUEUE_CAPACITY: usize>(inner: &SchedulerInner<QUEUE_CAPACITY>) -> bool {
use std::sync::atomic::Ordering;
inner.shutdown.load(Ordering::Acquire) && inner.pending_tasks.load(Ordering::Acquire) == 0
}
fn spin_for_work<const QUEUE_CAPACITY: usize, const SPIN_LIMIT: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
) -> bool {
use std::sync::atomic::Ordering;
for attempt in 0..SPIN_LIMIT {
core::hint::spin_loop();
let local = &inner.workers[worker_id];
if !local.queues.is_empty()
|| local.lifo_slot.state.load(Ordering::Relaxed) == 2
|| should_stop(inner)
{
return true;
}
if attempt % 32 == 0 && (has_stealable_work(inner, worker_id) || should_stop(inner)) {
return true;
}
}
false
}
fn has_stealable_work<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
) -> bool {
use std::sync::atomic::Ordering;
let worker_count = inner.workers.len();
const STEAL_PROBE_LIMIT: usize = 8;
let start = next_steal_start();
let probe_count = worker_count.min(STEAL_PROBE_LIMIT);
for offset in 0..probe_count {
let victim_index = (start.wrapping_add(offset)) % worker_count;
if victim_index == worker_id {
continue;
}
let victim = &inner.workers[victim_index];
if !victim.queues.is_empty() || victim.lifo_slot.state.load(Ordering::Relaxed) == 2 {
return true;
}
}
false
}
fn wait_for_work<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
worker_id: usize,
) {
use std::sync::atomic::Ordering;
inner.idle_workers.set(worker_id);
while inner.pending_tasks.load(Ordering::SeqCst) == 0 && !inner.shutdown.load(Ordering::SeqCst)
{
thread::park();
}
inner.idle_workers.clear(worker_id);
}
pub(super) fn wake_worker<const QUEUE_CAPACITY: usize>(worker: &WorkerState<QUEUE_CAPACITY>) {
if let Some(thread) = worker.thread.get() {
thread.unpark();
}
}
pub(super) fn wake_all_workers<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
) {
for worker in inner.workers.iter() {
wake_worker(worker);
}
}
pub(super) trait ContendedWakable {
#[cfg(feature = "scheduler-diagnostics")]
fn worker_count(&self) -> usize;
#[cfg(feature = "scheduler-diagnostics")]
fn wake_worker(&self, worker_index: usize);
fn wake_contended<P>(&self, worker_index: usize, previous_pending: usize) -> usize
where
P: ContendedWakePolicy;
}
impl<const QUEUE_CAPACITY: usize> ContendedWakable for SchedulerInner<QUEUE_CAPACITY> {
#[cfg(feature = "scheduler-diagnostics")]
fn worker_count(&self) -> usize {
self.workers.len()
}
#[cfg(feature = "scheduler-diagnostics")]
fn wake_worker(&self, worker_index: usize) {
wake_worker(&self.workers[worker_index]);
}
fn wake_contended<P>(&self, worker_index: usize, previous_pending: usize) -> usize
where
P: ContendedWakePolicy,
{
let worker_count = self.workers.len();
wake_worker(&self.workers[worker_index]);
if P::WAKE_LIMIT < 2 || worker_count < 2 {
return 1;
}
let peer_index = worker_index.wrapping_add(previous_pending) % worker_count;
wake_worker(&self.workers[peer_index]);
2
}
}
#[cold]
#[inline(never)]
pub(super) fn wake_contended_workers<P>(
inner: &impl ContendedWakable,
worker_index: usize,
previous_pending: usize,
) -> usize
where
P: ContendedWakePolicy,
{
inner.wake_contended::<P>(worker_index, previous_pending)
}
#[cfg(feature = "scheduler-diagnostics")]
#[inline]
pub(super) fn diagnostic_publish_work_available(
inner: &impl ContendedWakable,
worker_index: usize,
previous_pending: usize,
) -> usize {
let worker_count = inner.worker_count();
if previous_pending == 0 {
inner.wake_worker(worker_index);
1
} else if previous_pending < worker_count {
wake_contended_workers::<BoundedContendedWake>(inner, worker_index, previous_pending)
} else {
0
}
}
pub(super) fn notify_quiescent<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
) {
use std::sync::atomic::Ordering;
if inner.join_waiters.load(Ordering::SeqCst) != 0 && is_quiescent(inner) {
let _guard = lock_mutex(&inner.wait_lock);
if is_quiescent(inner) {
inner.wait_signal.notify_all();
}
}
}
pub(super) fn is_quiescent<const QUEUE_CAPACITY: usize>(
inner: &SchedulerInner<QUEUE_CAPACITY>,
) -> bool {
use std::sync::atomic::Ordering;
inner.pending_tasks.load(Ordering::SeqCst) == 0
&& inner.active_workers.load(Ordering::SeqCst) == 0
&& inner.blocking_pending_tasks.load(Ordering::SeqCst) == 0
&& inner.blocking_active_workers.load(Ordering::SeqCst) == 0
}
pub(super) fn inline_map_reduce<T, Map, Reduce>(
count: usize,
identity: T,
map: Map,
reduce: Reduce,
) -> ExecutorResult<T>
where
Map: Fn(usize) -> T,
Reduce: Fn(T, T) -> T,
{
use std::panic::{catch_unwind, AssertUnwindSafe};
catch_unwind(AssertUnwindSafe(|| {
let mut accumulator = identity;
for index in 0..count {
accumulator = reduce(accumulator, map(index));
}
accumulator
}))
.map_err(|_| ExecutorError::SpawnFailed(moirai_core::error::TaskError::Panicked))
}
pub(super) fn map_reduce_range<T, Map, Reduce>(
start: usize,
end: usize,
identity: T,
map: &Map,
reduce: &Reduce,
) -> T
where
Map: Fn(usize) -> T,
Reduce: Fn(T, T) -> T,
{
let mut accumulator = identity;
for index in start..end {
accumulator = reduce(accumulator, map(index));
}
accumulator
}
pub(super) fn indexed_chunk_count(count: usize, worker_count: usize) -> usize {
count.min(worker_count.max(1).saturating_add(1))
}
pub(super) fn indexed_chunk_bounds(
count: usize,
chunk_count: usize,
chunk_index: usize,
) -> (usize, usize) {
let base = count / chunk_count;
let remainder = count % chunk_count;
let start = chunk_index * base + chunk_index.min(remainder);
let len = base + usize::from(chunk_index < remainder);
(start, start + len)
}
pub(super) fn lock_mutex<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
#[cfg(test)]
mod indexed_chunk_count_tests {
use super::{indexed_chunk_bounds, indexed_chunk_count};
#[test]
fn assigns_small_domains_across_available_lanes() {
assert_eq!(indexed_chunk_count(9, 8), 9);
assert_eq!(indexed_chunk_count(2, 8), 2);
assert_eq!(indexed_chunk_count(1, 8), 1);
assert_eq!(indexed_chunk_count(0, 8), 0);
}
#[test]
fn caps_large_domains_at_workers_plus_caller() {
assert_eq!(indexed_chunk_count(1_000_000, 8), 9);
}
#[test]
fn single_worker_uses_worker_plus_caller() {
assert_eq!(indexed_chunk_count(1024, 1), 2);
assert_eq!(indexed_chunk_count(2, 1), 2);
}
#[test]
fn balances_remainder_across_every_chunk() {
let bounds: Vec<_> = (0..9)
.map(|chunk_index| indexed_chunk_bounds(10, 9, chunk_index))
.collect();
assert_eq!(
bounds,
vec![
(0, 2),
(2, 3),
(3, 4),
(4, 5),
(5, 6),
(6, 7),
(7, 8),
(8, 9),
(9, 10)
]
);
}
}