use crate::padded_type::PaddedType;
use crate::retired_list::RetiredList;
use crate::task_batch::TaskBatch;
use crate::{TaskFnPointer, TaskFuture, TaskParamPointer};
use std::cell::UnsafeCell;
use std::ptr::NonNull;
use std::sync::atomic::{AtomicBool, AtomicPtr, AtomicUsize, Ordering, fence};
use std::thread::{self, Thread};
pub const NOT_IN_CRITICAL: usize = usize::MAX;
pub const EPOCH_MASK: usize = usize::MAX >> 1; pub const EPOCH_MASK_HALF: usize = EPOCH_MASK / 2;
pub struct Queue {
head: PaddedType<AtomicPtr<TaskBatch>>,
tail: PaddedType<AtomicPtr<TaskBatch>>,
global_epoch: PaddedType<AtomicUsize>,
local_epochs: Box<[PaddedType<AtomicUsize>]>,
threads: Box<[UnsafeCell<Thread>]>,
shutdown: AtomicBool,
}
unsafe impl Sync for Queue {}
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 threads = (0..worker_count)
.map(|_| UnsafeCell::new(thread::current()))
.collect();
Queue {
head: PaddedType::new(AtomicPtr::new(anchor)),
tail: PaddedType::new(AtomicPtr::new(anchor)),
global_epoch: PaddedType::new(AtomicUsize::new(0)),
local_epochs,
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::AcqRel);
unsafe {
(*prev_tail).next.store(batch, Ordering::Release);
}
let global_epoch = self.global_epoch.load(Ordering::Relaxed) & EPOCH_MASK;
fence(Ordering::SeqCst);
let mut remaining = count.min(self.threads.len());
for (epoch, thread) in self.local_epochs.iter().zip(self.threads.iter()) {
if epoch
.compare_exchange(
NOT_IN_CRITICAL,
global_epoch,
Ordering::Release,
Ordering::Relaxed,
)
.is_ok()
{
unsafe {
(*thread.get()).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_worker_thread(&self, worker_id: usize) {
unsafe {
*self.threads[worker_id].get() = thread::current();
}
}
pub fn wait_for_work(&self, worker_id: usize) -> bool {
loop {
if self.has_tasks() {
return true;
}
if self.is_shutdown() {
return false;
}
self.local_epochs[worker_id].store(NOT_IN_CRITICAL, Ordering::Relaxed);
fence(Ordering::SeqCst);
if self.has_tasks() {
return true;
}
thread::park();
}
}
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.iter() {
let e = local_epoch.load(Ordering::Relaxed);
if e != NOT_IN_CRITICAL {
if min_epoch.wrapping_sub(e) & EPOCH_MASK < EPOCH_MASK_HALF {
min_epoch = e;
}
}
}
min_epoch
}
pub fn is_shutdown(&self) -> bool {
self.shutdown.load(Ordering::Acquire)
}
pub fn has_tasks(&self) -> bool {
let tail = self.tail.load(Ordering::Acquire);
unsafe { (&*tail).has_unclaimed_tasks() }
}
pub fn shutdown(&self) {
self.shutdown.store(true, Ordering::Release);
self.threads.iter().for_each(|t| unsafe {
(*t.get()).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);
}
}
}