use std::cmp::Ordering as CmpOrdering;
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicUsize, Ordering};
use crate::err::{Error, Result};
use crate::sync::asm::fence::pause;
const CACHE_LINE: usize = 64;
#[repr(C, align(64))]
pub struct PriorityQueue<T: Ord, const N: usize> {
len: AtomicUsize,
lock: AtomicUsize,
_pad: [u8; CACHE_LINE - 16],
heap: [MaybeUninit<T>; N],
}
unsafe impl<T: Ord + Send, const N: usize> Send for PriorityQueue<T, N> {}
unsafe impl<T: Ord + Send, const N: usize> Sync for PriorityQueue<T, N> {}
impl<T: Ord, const N: usize> PriorityQueue<T, N> {
pub fn new() -> Self {
Self {
len: AtomicUsize::new(0),
lock: AtomicUsize::new(0),
_pad: [0; CACHE_LINE - 16],
heap: unsafe { MaybeUninit::uninit().assume_init() },
}
}
fn acquire(&self) {
loop {
if self
.lock
.compare_exchange_weak(0, 1, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
return;
}
pause();
}
}
fn release(&self) {
self.lock.store(0, Ordering::Release);
}
pub fn push(&self, value: T) -> Result<()> {
self.acquire();
let len = self.len.load(Ordering::Relaxed);
if len >= N {
self.release();
return Err(Error::QueueFull);
}
unsafe {
let heap = self.heap.as_ptr().cast_mut();
(*heap.add(len)).write(value);
}
self.sift_up(len);
self.len.store(len + 1, Ordering::Release);
self.release();
Ok(())
}
pub fn pop(&self) -> Result<T> {
self.acquire();
let len = self.len.load(Ordering::Relaxed);
if len == 0 {
self.release();
return Err(Error::QueueEmpty);
}
let value = unsafe {
let heap = self.heap.as_ptr().cast_mut();
let result = (*heap).assume_init_read();
if len > 1 {
std::ptr::copy_nonoverlapping(heap.add(len - 1), heap, 1);
}
result
};
let new_len = len - 1;
self.len.store(new_len, Ordering::Release);
if new_len > 0 {
self.sift_down(0, new_len);
}
self.release();
Ok(value)
}
pub fn peek(&self) -> Option<&T> {
self.acquire();
let len = self.len.load(Ordering::Relaxed);
if len == 0 {
self.release();
return None;
}
let value = unsafe { self.heap[0].assume_init_ref() };
self.release();
Some(value)
}
fn sift_up(&self, mut idx: usize) {
let heap = self.heap.as_ptr().cast_mut();
while idx > 0 {
let parent = (idx - 1) / 2;
unsafe {
let current = (*heap.add(idx)).assume_init_ref();
let parent_val = (*heap.add(parent)).assume_init_ref();
if current.cmp(parent_val) != CmpOrdering::Less {
break;
}
std::ptr::swap(heap.add(idx), heap.add(parent));
}
idx = parent;
}
}
fn sift_down(&self, mut idx: usize, len: usize) {
let heap = self.heap.as_ptr().cast_mut();
loop {
let left = 2 * idx + 1;
let right = 2 * idx + 2;
let mut smallest = idx;
unsafe {
if left < len {
let current = (*heap.add(smallest)).assume_init_ref();
let left_val = (*heap.add(left)).assume_init_ref();
if left_val.cmp(current) == CmpOrdering::Less {
smallest = left;
}
}
if right < len {
let current = (*heap.add(smallest)).assume_init_ref();
let right_val = (*heap.add(right)).assume_init_ref();
if right_val.cmp(current) == CmpOrdering::Less {
smallest = right;
}
}
if smallest == idx {
break;
}
std::ptr::swap(heap.add(idx), heap.add(smallest));
}
idx = smallest;
}
}
#[inline]
pub fn len(&self) -> usize {
self.len.load(Ordering::Relaxed)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn capacity(&self) -> usize {
N
}
pub fn is_full(&self) -> bool {
self.len() >= N
}
}
impl<T: Ord, const N: usize> Default for PriorityQueue<T, N> {
fn default() -> Self {
Self::new()
}
}