use alloc::{boxed::Box, vec::Vec};
use core::{
alloc::Layout,
cell::UnsafeCell,
mem::MaybeUninit,
ptr::{self, NonNull},
slice,
task::Waker,
};
use crate::Task;
pub struct Enqueuer<T: Task> {
inner: UnsafeCell<Inner<T>>,
}
struct Inner<T: Task> {
counter: isize,
end: NonNull<T>,
capacity: isize,
waker: Option<Waker>,
}
impl<T: Task> Enqueuer<T> {
pub(crate) const fn new() -> Self {
let capacity = if Self::ZST { isize::MAX } else { 0 };
let counter = -capacity - Self::ESZ;
Self {
inner: UnsafeCell::new(Inner {
counter,
end: NonNull::dangling(),
capacity,
waker: None,
}),
}
}
pub(crate) fn set_waker(&self, waker: Waker) {
match self.waker() {
Some(slot) => {
*slot = waker;
}
slot @ None => {
assert!(!self.is_locked(), "the enqueuer is locked");
let counter = self.counter();
let empty = -self.byte_capacity() - Self::ESZ;
assert!(*counter == empty, "the enqueuer is not empty");
*slot = Some(waker);
*counter = -Self::ESZ;
}
}
}
pub(crate) fn take_waker(&self) -> Option<Waker> {
let waker = self.waker().take()?;
*self.counter() = -self.byte_capacity() - Self::ESZ;
Some(waker)
}
pub fn add(&self, task: T) {
let inner = self.inner.get();
let counter = unsafe { (*inner).counter };
let next = counter + Self::ESZ;
if next < 0 {
let end = unsafe { (*inner).end };
if !Self::ZST {
let ptr = unsafe { end.byte_offset(next) };
unsafe { ptr.write(task) };
}
unsafe { (*inner).counter = next };
return;
}
self.add_slow(task, next)
}
#[cold]
fn add_slow(&self, task: T, next: isize) {
let inner = self.inner.get();
assert!(next == 0, "the enqueuer is already borrowed");
let capacity = unsafe { (*inner).capacity };
if capacity == 0 {
cold_path();
debug_assert!(!Self::ZST);
let mut buffer = Box::new_uninit_slice(16);
buffer[0].write(task);
let ptr = Box::into_raw(buffer);
let end = unsafe { ptr.cast::<T>().add(16) };
let end = unsafe { NonNull::new_unchecked(end) };
unsafe { (*inner).end = end };
unsafe { (*inner).capacity = 16 * Self::ESZ };
unsafe { (*inner).counter = -16 * Self::ESZ };
if let Some(waker) = unsafe { (*inner).waker.take() } {
waker.wake();
}
return;
}
let end = unsafe { (*inner).end };
let ptr = unsafe { end.byte_offset(-capacity) };
if unsafe { (*inner).waker.is_some() } {
let waker = unsafe { (*inner).waker.take() }.unwrap();
unsafe { ptr.write(task) };
unsafe { (*inner).counter = -capacity };
waker.wake();
} else {
let mut vec = unsafe {
Vec::from_raw_parts(
ptr.as_ptr(),
(capacity / Self::ESZ) as usize,
(capacity / Self::ESZ) as usize,
)
};
vec.push(task);
let old_cap = capacity;
let capacity = vec.capacity() as isize * Self::ESZ;
let ptr = unsafe { NonNull::new_unchecked(vec.as_mut_ptr()) };
let end = unsafe { ptr.byte_offset(capacity) };
core::mem::forget(vec);
unsafe { (*inner).end = end };
unsafe { (*inner).capacity = capacity };
unsafe { (*inner).counter = -capacity + old_cap };
}
}
pub fn extend(&self, tasks: impl IntoIterator<Item = T>) {
for task in tasks {
self.add(task);
}
}
pub(crate) fn drain(&self, tasks: &mut Vec<T>) -> usize {
let inner = self.inner.get();
let counter = unsafe { (*inner).counter };
let capacity = unsafe { (*inner).capacity };
assert!(counter < 0, "the enqueuer is already borrowed");
let end = unsafe { (*inner).end };
let ptr = unsafe { end.byte_offset(-capacity) };
let ptr = ptr.cast::<MaybeUninit<T>>().as_ptr();
let len = ((capacity + counter) / Self::ESZ + 1) as usize;
tasks.reserve(len);
let space = tasks.spare_capacity_mut().as_mut_ptr();
unsafe { ptr::copy_nonoverlapping(ptr, space, len) };
unsafe { tasks.set_len(tasks.len() + len) };
unsafe { (*inner).counter = -capacity - Self::ESZ };
len
}
}
impl<T: Task> Extend<T> for &Enqueuer<T> {
fn extend<I: IntoIterator<Item = T>>(&mut self, tasks: I) {
(*self).extend(tasks)
}
}
impl<T: Task> Drop for Enqueuer<T> {
fn drop(&mut self) {
if self.is_locked() {
return;
}
if self.waker().is_none() {
for task in unsafe { self.buffer_init() } {
unsafe { task.assume_init_drop() };
}
}
if !Self::ZST && self.byte_capacity() != 0 {
let buffer = unsafe { self.buffer() };
let layout =
unsafe { Layout::array::<T>(buffer.len()).unwrap_unchecked() };
unsafe {
alloc::alloc::dealloc(buffer.as_mut_ptr().cast(), layout)
};
}
}
}
impl<T: Task> Enqueuer<T> {
const ZST: bool = size_of::<T>() == 0;
const ESZ: isize = if !Self::ZST {
size_of::<T>() as isize
} else {
1
};
fn is_locked(&self) -> bool {
unsafe { (*self.inner.get()).counter == 0 }
}
fn byte_capacity(&self) -> isize {
unsafe { (*self.inner.get()).capacity }
}
fn byte_len(&self) -> isize {
let inner = self.inner.get();
unsafe { (*inner).capacity + (*inner).counter + Self::ESZ }
}
fn len(&self) -> usize {
(self.byte_len() / Self::ESZ) as usize
}
fn capacity(&self) -> usize {
(self.byte_capacity() / Self::ESZ) as usize
}
fn ptr(&self) -> *mut T {
let inner = self.inner.get();
let off = if Self::ZST {
0
} else {
unsafe { (*inner).capacity }
};
unsafe { (*inner).end.as_ptr().byte_offset(-off) }
}
#[allow(clippy::mut_from_ref)] unsafe fn buffer(&self) -> &mut [MaybeUninit<T>] {
unsafe { slice::from_raw_parts_mut(self.ptr().cast(), self.capacity()) }
}
#[allow(clippy::mut_from_ref)] unsafe fn buffer_init(&self) -> &mut [MaybeUninit<T>] {
unsafe { slice::from_raw_parts_mut(self.ptr().cast(), self.len()) }
}
#[allow(clippy::mut_from_ref)] fn waker(&self) -> &mut Option<Waker> {
unsafe { &mut (*self.inner.get()).waker }
}
#[allow(clippy::mut_from_ref)] fn counter(&self) -> &mut isize {
unsafe { &mut (*self.inner.get()).counter }
}
}
#[cold]
fn cold_path() {}
#[cfg(test)]
mod test {
use std::vec::Vec;
use crate::Task;
use super::Enqueuer;
#[derive(Copy, Clone, PartialEq, Eq, Debug)]
struct Foo(u32);
impl Task for Foo {
type Priority = ();
fn priority(&self) -> Self::Priority {}
}
#[test]
fn new() {
let _ = Enqueuer::<()>::new();
let _ = Enqueuer::<Foo>::new();
}
#[test]
fn add() {
let enqueuer = Enqueuer::<()>::new();
enqueuer.add(());
enqueuer.add(());
enqueuer.add(());
enqueuer.add(());
assert_eq!(enqueuer.len(), 4);
let enqueuer = Enqueuer::<Foo>::new();
enqueuer.add(Foo(0));
enqueuer.add(Foo(1));
enqueuer.add(Foo(2));
enqueuer.add(Foo(3));
assert_eq!(enqueuer.len(), 4);
assert_eq!(enqueuer.byte_len(), 16);
let mut tasks = Vec::new();
enqueuer.drain(&mut tasks);
assert_eq!(tasks, [Foo(0), Foo(1), Foo(2), Foo(3)]);
}
#[test]
fn realloc() {
let enqueuer = Enqueuer::<Foo>::new();
for n in 1..4 {
let mut expected = Vec::new();
for i in 0..32 * n {
let task = Foo(n * 32 + i);
enqueuer.add(task);
expected.push(task);
}
assert_eq!(enqueuer.len(), 32 * n as usize);
assert_eq!(enqueuer.byte_len(), 32 * n as isize * 4);
let mut tasks = Vec::new();
enqueuer.drain(&mut tasks);
assert_eq!(tasks, expected);
}
}
}