use alloc::{boxed::Box, vec::Vec};
use atomig::Atomic;
use core::{
cell::UnsafeCell,
mem::{MaybeUninit, transmute},
num::NonZeroUsize,
sync::atomic::{AtomicU32, Ordering},
};
use crate::{
Config, Task, TaskPriority,
util::{Backoff, extend_vec_from_slice},
};
#[repr(transparent)]
pub struct Priority<T: Task> {
raw: Atomic<<T::Priority as TaskPriority>::Repr>,
}
impl<T: Task> Default for Priority<T> {
#[inline]
fn default() -> Self {
Self {
raw: Atomic::new(<T::Priority as TaskPriority>::pack(None)),
}
}
}
impl<T: Task> Priority<T> {
#[inline]
pub unsafe fn set(&self, value: T::Priority) {
let pack = <T::Priority as TaskPriority>::pack;
let unpack = <T::Priority as TaskPriority>::unpack;
debug_assert_eq!(unpack(self.raw.load(Ordering::Relaxed)), None);
self.raw.store(pack(Some(value)), Ordering::Release);
}
#[inline]
pub unsafe fn steal_self(&self, expected: T::Priority) -> Result<(), ()> {
let pack = <T::Priority as TaskPriority>::pack;
let unpack = <T::Priority as TaskPriority>::unpack;
let prev = unpack(self.raw.swap(pack(None), Ordering::Acquire));
debug_assert!(
prev.is_none_or(|prev| prev == expected),
"'prev' ({prev:?}) can only be empty or 'expected' ({expected:?})"
);
if prev.is_some() { Ok(()) } else { Err(()) }
}
#[inline]
pub unsafe fn steal_if_above(
&self,
min: Option<T::Priority>,
) -> Option<T::Priority> {
let pack = <T::Priority as TaskPriority>::pack;
let unpack = <T::Priority as TaskPriority>::unpack;
let result = self.raw.fetch_update(
Ordering::Acquire,
Ordering::Relaxed,
|cur| {
if unpack(cur) > min {
Some(pack(None))
} else {
None
}
},
);
match result {
Ok(prev) => Some(unsafe { unpack(prev).unwrap_unchecked() }),
Err(_) => None,
}
}
}
#[repr(transparent)]
pub struct Stealer {
raw: AtomicU32,
}
impl Stealer {
#[inline]
pub const unsafe fn new(id: usize, config: &Config) -> Self {
debug_assert!(id < config.num_workers.get());
let value = unsafe { PubQueue::new(id, config.batch_size, config) };
Self {
raw: AtomicU32::new(value.raw),
}
}
#[inline]
pub unsafe fn set(&self, expected: PubQueue, new: PubQueue) {
debug_assert_eq!(
self.raw.load(Ordering::Relaxed) % (1u32 << 31),
expected.raw
);
self.raw.store(new.raw, Ordering::Relaxed);
}
#[inline]
pub unsafe fn try_dec(
&self,
expected: PubQueue,
config: &Config,
) -> Result<PubQueue, PubQueue> {
debug_assert!(expected.len(config).get() > 1);
let res = self.raw.fetch_sub(1, Ordering::Relaxed);
if res == expected.raw {
let len = expected.len(config).get() - 1;
let len = unsafe { NonZeroUsize::new_unchecked(len) };
Ok(unsafe { expected.with_len(len, config) })
} else {
Err(unsafe { PubQueue::from_raw(res - 1) })
}
}
#[inline]
pub unsafe fn steal(
&self,
replacement: usize,
config: &Config,
) -> PubQueue {
debug_assert!(replacement < config.num_workers.get());
let replacement =
unsafe { PubQueue::new(replacement, config.batch_size, config) };
let mut value = self.raw.swap(replacement.raw, Ordering::Release);
if value >= (1u32 << 31) {
value -= 1u32 << 31;
atomic_wait::wake_one(&self.raw);
}
unsafe { PubQueue::from_raw(value) }
}
pub unsafe fn wait_for_theft(&self, original: PubQueue) -> PubQueue {
let backoff = Backoff::new();
while !backoff.is_completed() {
let raw = self.raw.load(Ordering::Acquire);
if raw != original.raw {
debug_assert!(raw < (1u32 << 31));
return unsafe { PubQueue::from_raw(raw) };
}
backoff.snooze();
}
let raw = self.raw.fetch_or(1u32 << 31, Ordering::Relaxed);
if raw != original.raw {
debug_assert!(raw < (1u32 << 31));
self.raw.store(raw, Ordering::Relaxed);
return unsafe { PubQueue::from_raw(raw) };
}
loop {
let old = original.raw | (1u32 << 31);
atomic_wait::wait(&self.raw, old);
let raw = self.raw.load(Ordering::Relaxed);
if raw != old {
debug_assert!(raw < (1u32 << 31));
return unsafe { PubQueue::from_raw(raw) };
}
}
}
}
#[derive(Copy, Clone)]
#[repr(transparent)]
pub struct PubQueue {
raw: u32,
}
impl PubQueue {
#[inline]
pub const unsafe fn new(
id: usize,
len: NonZeroUsize,
config: &Config,
) -> Self {
debug_assert!(id < config.num_workers.get());
debug_assert!(len.get() <= config.batch_size.get());
Self {
raw: (id * config.batch_size.get() + (len.get() - 1)) as u32,
}
}
#[inline]
pub const unsafe fn from_raw(raw: u32) -> Self {
debug_assert!(raw < (1u32 << 31));
Self { raw }
}
#[inline]
pub const fn id(&self, config: &Config) -> usize {
debug_assert!(self.raw < (1u32 << 31));
self.raw as usize / config.batch_size.get()
}
#[inline]
pub const fn len(&self, config: &Config) -> NonZeroUsize {
let batch_size = config.batch_size.get();
unsafe {
NonZeroUsize::new_unchecked((self.raw as usize % batch_size) + 1)
}
}
#[inline]
pub const unsafe fn with_len(
self,
len: NonZeroUsize,
config: &Config,
) -> Self {
debug_assert!(len.get() <= config.batch_size.get());
unsafe { Self::new(self.id(config), len, config) }
}
}
#[repr(transparent)]
pub struct PQContents<T> {
tasks: [UnsafeCell<MaybeUninit<T>>],
}
impl<T> PQContents<T> {
pub fn new_boxed(config: &Config) -> Box<Self> {
let size: usize = config.batch_size.get();
let ptr = Box::<[MaybeUninit<T>]>::new_uninit_slice(size);
unsafe { transmute(ptr) }
}
#[inline]
pub unsafe fn read(&self, index: usize) -> MaybeUninit<T> {
debug_assert!(index < self.tasks.len());
unsafe { self.tasks.get_unchecked(index).get().read() }
}
pub unsafe fn move_to(&self, len: usize, vec: &mut Vec<T>) {
debug_assert!(len <= self.tasks.len());
let tasks = unsafe { self.tasks.get_unchecked(..len) };
let data = unsafe {
transmute::<&[UnsafeCell<MaybeUninit<T>>], &[MaybeUninit<T>]>(tasks)
};
unsafe { extend_vec_from_slice(vec, data) };
}
pub unsafe fn fill(&self, iter: impl Iterator<Item = T>) {
iter.zip(&self.tasks).for_each(|(elem, slot)| {
unsafe { slot.get().write(MaybeUninit::new(elem)) };
});
}
}
unsafe impl<T> Send for PQContents<T> {}
unsafe impl<T> Sync for PQContents<T> {}
#[cfg(test)]
mod tests {
use std::{
num::{NonZeroU32, NonZeroUsize},
sync::{Barrier, atomic::Ordering},
thread,
vec::Vec,
};
use crate::{Config, Task};
use super::{PQContents, Priority, PubQueue, Stealer};
#[test]
fn priority_theft() {
enum Foo {}
impl Task for Foo {
type Priority = NonZeroU32;
fn priority(&self) -> Self::Priority {
match *self {}
}
}
let priority = Priority::<Foo>::default();
unsafe { priority.set(NonZeroU32::new(42).unwrap()) };
let barrier = Barrier::new(2);
let mut stealer_success = None;
let mut owner_success = None;
thread::scope(|s| {
s.spawn(|| {
barrier.wait();
let min = Some(NonZeroU32::new(30).unwrap());
let res = unsafe { priority.steal_if_above(min) };
stealer_success = Some(res.is_some());
if let Some(res) = res {
assert_eq!(res.get(), 42);
}
});
s.spawn(|| {
barrier.wait();
let orig = NonZeroU32::new(42).unwrap();
let res = unsafe { priority.steal_self(orig) };
owner_success = Some(res.is_ok());
});
});
assert!(stealer_success.unwrap() != owner_success.unwrap());
}
#[test]
fn id_len() {
let num_workers = NonZeroUsize::new(4).unwrap();
let batch_size = NonZeroUsize::new(4).unwrap();
let config = Config::new(num_workers).with_batch_size(batch_size);
let cases = [(0, 1), (3, 4)];
for (id, len) in cases {
let len = NonZeroUsize::new(len).unwrap();
let pq = unsafe { PubQueue::new(id, len, &config) };
assert_eq!(pq.id(&config), id);
assert_eq!(pq.len(&config), len);
}
}
#[test]
fn pq_theft() {
let num_workers = NonZeroUsize::new(2).unwrap();
let batch_size = NonZeroUsize::new(4).unwrap();
let config = Config::new(num_workers).with_batch_size(batch_size);
let stealer = unsafe { Stealer::new(0, &config) };
let barrier = Barrier::new(2);
assert_eq!(stealer.raw.load(Ordering::Relaxed), 3);
thread::scope(|s| {
s.spawn(|| {
barrier.wait();
let pq = unsafe { stealer.steal(1, &config) };
assert_eq!(pq.id(&config), 0);
assert_eq!(pq.len(&config).get(), 4);
});
s.spawn(|| {
barrier.wait();
let orig = unsafe { PubQueue::new(0, batch_size, &config) };
let pq = unsafe { stealer.wait_for_theft(orig) };
assert_eq!(pq.id(&config), 1);
assert_eq!(pq.len(&config).get(), 4);
});
})
}
#[test]
fn contents() {
let num_workers = NonZeroUsize::new(2).unwrap();
let batch_size = NonZeroUsize::new(4).unwrap();
let config = Config::new(num_workers).with_batch_size(batch_size);
let contents = PQContents::<u32>::new_boxed(&config);
unsafe { contents.fill(4..8) };
for index in 0..4 {
assert_eq!(
unsafe { contents.read(index).assume_init() },
index as u32 + 4
);
}
let mut elems = Vec::new();
unsafe { contents.move_to(4, &mut elems) };
assert_eq!(elems, [4, 5, 6, 7]);
}
}