use crate::padded_type::PaddedType;
use crate::retired_list::RetiredList;
use crate::task_batch::TaskBatch;
use crate::{TaskFnPointer, TaskFuture, TaskParamPointer};
use std::ptr::NonNull;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicBool, AtomicPtr, AtomicU8, AtomicUsize, Ordering, fence};
use std::thread::{self, Thread};
const NOT_IN_CRITICAL: usize = usize::MAX;
pub const EPOCH_MASK: usize = usize::MAX >> 1; pub const EPOCH_MASK_HALF: usize = EPOCH_MASK / 2;
const STATE_RUNNING: u8 = 0;
const STATE_SLEEPING: u8 = 1;
const STATE_NOTIFIED: u8 = 2;
pub struct Queue {
head: PaddedType<AtomicPtr<TaskBatch>>,
tail: PaddedType<AtomicPtr<TaskBatch>>,
global_epoch: PaddedType<AtomicUsize>,
local_epochs: Box<[PaddedType<AtomicUsize>]>,
worker_states: Box<[PaddedType<AtomicU8>]>,
threads: Box<[OnceLock<Thread>]>,
shutdown: AtomicBool,
}
impl Queue {
pub fn new(worker_count: usize) -> Self {
fn noop(_: TaskParamPointer) {}
let anchor = TaskBatch::new(noop, NonNull::dangling(), 0, 0, TaskFuture::new(0));
let local_epochs = (0..worker_count)
.map(|_| PaddedType::new(AtomicUsize::new(NOT_IN_CRITICAL)))
.collect();
let worker_states = (0..worker_count)
.map(|_| PaddedType::new(AtomicU8::new(STATE_RUNNING)))
.collect();
let threads = (0..worker_count).map(|_| OnceLock::new()).collect();
Queue {
head: PaddedType::new(AtomicPtr::new(anchor)),
tail: PaddedType::new(AtomicPtr::new(anchor)),
global_epoch: PaddedType::new(AtomicUsize::new(0)),
local_epochs,
worker_states,
threads,
shutdown: AtomicBool::new(false),
}
}
pub fn push_task_batch<T>(&self, task_fn: fn(&T), params: &[T]) -> TaskFuture {
if params.is_empty() {
return TaskFuture::new(0);
}
let future = TaskFuture::new(params.len());
let batch = TaskBatch::new(
unsafe { std::mem::transmute::<fn(&T), TaskFnPointer>(task_fn) },
NonNull::from(params).cast(),
std::mem::size_of::<T>(),
std::mem::size_of_val(params),
future.clone(),
);
self.link_and_notify(batch, params.len());
future
}
fn link_and_notify(&self, batch: *mut TaskBatch, count: usize) {
let prev_tail = self.tail.swap(batch, Ordering::Release);
unsafe {
(*prev_tail).next.store(batch, Ordering::Release);
}
let mut remaining = count.min(self.threads.len());
for (state, thread_oncelock) in self.worker_states.iter().zip(self.threads.iter()) {
if let Some(thread) = thread_oncelock.get()
&& state.swap(STATE_NOTIFIED, Ordering::Release) == STATE_SLEEPING
{
thread.unpark();
remaining -= 1;
if remaining == 0 {
break;
}
}
}
}
pub fn get_next_batch(
&self,
worker_id: usize,
retired_list: &mut RetiredList,
) -> Option<(&TaskBatch, TaskParamPointer)> {
let global_epoch = self.global_epoch.load(Ordering::Relaxed) & EPOCH_MASK;
if self.local_epochs[worker_id].load(Ordering::Relaxed) != global_epoch {
self.local_epochs[worker_id].store(global_epoch, Ordering::Relaxed);
fence(Ordering::SeqCst);
}
let mut current = self.head.load(Ordering::Acquire);
loop {
let batch = unsafe { &*current };
if let Some(param) = batch.claim_next_param() {
return Some((batch, param));
}
let next = batch.next.load(Ordering::Acquire);
if next.is_null() {
return None;
}
match self.head.compare_exchange_weak(
current,
next,
Ordering::Release,
Ordering::Acquire,
) {
Ok(_) => {
let fresh_epoch = self.global_epoch.load(Ordering::Relaxed) & EPOCH_MASK;
retired_list.push(current, fresh_epoch);
current = next;
}
Err(new_head) => {
current = new_head;
}
}
}
}
pub fn register_thread(&self, worker_id: usize, thread: Thread) {
let _ = self.threads[worker_id].set(thread);
}
pub fn wait_for_work(&self, worker_id: usize) -> bool {
loop {
if self.shutdown.load(Ordering::Relaxed) {
return false;
}
if self.worker_states[worker_id]
.compare_exchange(
STATE_RUNNING,
STATE_SLEEPING,
Ordering::Relaxed,
Ordering::Acquire,
)
.is_err()
{
self.worker_states[worker_id].store(STATE_RUNNING, Ordering::Relaxed);
return true;
}
self.local_epochs[worker_id].store(NOT_IN_CRITICAL, Ordering::Relaxed);
thread::park();
if self.worker_states[worker_id].swap(STATE_RUNNING, Ordering::Acquire)
== STATE_NOTIFIED
{
return true;
}
}
}
pub fn advance_and_min_epoch(&self) -> usize {
let global_epoch = self
.global_epoch
.fetch_add(1, Ordering::Relaxed)
.wrapping_add(1);
let mut min_epoch = global_epoch & EPOCH_MASK;
fence(Ordering::SeqCst);
for local_epoch in &self.local_epochs {
let e = local_epoch.load(Ordering::Relaxed);
if e != NOT_IN_CRITICAL && (min_epoch.wrapping_sub(e) & EPOCH_MASK < EPOCH_MASK_HALF) {
min_epoch = e;
}
}
min_epoch
}
pub fn shutdown(&self) {
self.shutdown.store(true, Ordering::Relaxed);
self.threads
.iter()
.filter_map(OnceLock::get)
.for_each(Thread::unpark);
}
}
impl Drop for Queue {
fn drop(&mut self) {
let mut current = self.head.load(Ordering::Relaxed);
while !current.is_null() {
let batch = unsafe { Box::from_raw(current) };
current = batch.next.load(Ordering::Relaxed);
}
}
}