use std::cell::UnsafeCell;
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use crate::err::{Error, Result};
const CACHE_LINE: usize = 64;
#[repr(C, align(64))]
struct BroadcastSlot<T> {
seq: AtomicUsize,
data: UnsafeCell<MaybeUninit<T>>,
}
#[repr(C, align(64))]
pub struct BroadcastBuffer<T, const N: usize> {
tail: AtomicUsize,
_pad: [u8; CACHE_LINE - 8],
closed: AtomicBool,
slots: [BroadcastSlot<T>; N],
}
unsafe impl<T: Send, const N: usize> Send for BroadcastBuffer<T, N> {}
unsafe impl<T: Send, const N: usize> Sync for BroadcastBuffer<T, N> {}
impl<T: Clone, const N: usize> BroadcastBuffer<T, N> {
pub fn new() -> Self {
assert!(N.is_power_of_two());
let mut slots: [MaybeUninit<BroadcastSlot<T>>; N] =
unsafe { MaybeUninit::uninit().assume_init() };
for (i, slot) in slots.iter_mut().enumerate() {
slot.write(BroadcastSlot {
seq: AtomicUsize::new(i),
data: UnsafeCell::new(MaybeUninit::uninit()),
});
}
Self {
tail: AtomicUsize::new(0),
_pad: [0; CACHE_LINE - 8],
closed: AtomicBool::new(false),
slots: unsafe { std::mem::transmute_copy(&slots) },
}
}
const fn mask() -> usize {
N - 1
}
}
pub struct BroadcastSender<T: Clone, const N: usize> {
buf: Arc<BroadcastBuffer<T, N>>,
}
impl<T: Clone, const N: usize> BroadcastSender<T, N> {
fn new(buf: Arc<BroadcastBuffer<T, N>>) -> Self {
Self { buf }
}
pub fn send(&self, value: T) -> Result<()> {
if self.buf.closed.load(Ordering::Acquire) {
return Err(Error::ChannelClosed);
}
let tail = self.buf.tail.fetch_add(1, Ordering::AcqRel);
let slot = &self.buf.slots[tail & BroadcastBuffer::<T, N>::mask()];
unsafe {
(*slot.data.get()).write(value);
}
slot.seq.store(tail + 1, Ordering::Release);
Ok(())
}
pub fn close(&self) {
self.buf.closed.store(true, Ordering::Release);
}
pub fn is_closed(&self) -> bool {
self.buf.closed.load(Ordering::Acquire)
}
pub fn receiver_count(&self) -> usize {
Arc::strong_count(&self.buf) - 1
}
pub fn subscribe(&self) -> BroadcastReceiver<T, N> {
BroadcastReceiver::new(self.buf.clone(), self.buf.tail.load(Ordering::Acquire))
}
}
impl<T: Clone, const N: usize> Clone for BroadcastSender<T, N> {
fn clone(&self) -> Self {
Self {
buf: self.buf.clone(),
}
}
}
pub struct BroadcastReceiver<T: Clone, const N: usize> {
buf: Arc<BroadcastBuffer<T, N>>,
cursor: AtomicUsize,
}
impl<T: Clone, const N: usize> BroadcastReceiver<T, N> {
fn new(buf: Arc<BroadcastBuffer<T, N>>, cursor: usize) -> Self {
Self {
buf,
cursor: AtomicUsize::new(cursor),
}
}
pub fn recv(&self) -> Result<T> {
loop {
let cursor = self.cursor.load(Ordering::Relaxed);
let tail = self.buf.tail.load(Ordering::Acquire);
if cursor == tail {
if self.buf.closed.load(Ordering::Acquire) {
return Err(Error::ChannelClosed);
}
return Err(Error::QueueEmpty);
}
if tail.wrapping_sub(cursor) > N {
self.cursor.store(tail - N + 1, Ordering::Release);
return Err(Error::Interrupted);
}
let slot = &self.buf.slots[cursor & BroadcastBuffer::<T, N>::mask()];
let seq = slot.seq.load(Ordering::Acquire);
if seq == cursor + 1 {
let value = unsafe { (*slot.data.get()).assume_init_ref().clone() };
self.cursor.store(cursor + 1, Ordering::Release);
return Ok(value);
}
std::hint::spin_loop();
}
}
pub fn try_recv(&self) -> Result<T> {
self.recv()
}
pub fn recv_spin(&self) -> Result<T> {
loop {
match self.recv() {
Ok(v) => return Ok(v),
Err(Error::QueueEmpty) => std::hint::spin_loop(),
Err(e) => return Err(e),
}
}
}
pub fn is_closed(&self) -> bool {
self.buf.closed.load(Ordering::Acquire)
}
pub fn is_lagged(&self) -> bool {
let cursor = self.cursor.load(Ordering::Relaxed);
let tail = self.buf.tail.load(Ordering::Acquire);
tail.wrapping_sub(cursor) > N
}
pub fn len(&self) -> usize {
let cursor = self.cursor.load(Ordering::Relaxed);
let tail = self.buf.tail.load(Ordering::Acquire);
tail.saturating_sub(cursor)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl<T: Clone, const N: usize> Clone for BroadcastReceiver<T, N> {
fn clone(&self) -> Self {
Self {
buf: self.buf.clone(),
cursor: AtomicUsize::new(self.cursor.load(Ordering::Relaxed)),
}
}
}
pub fn broadcast<T: Clone, const N: usize>() -> (BroadcastSender<T, N>, BroadcastReceiver<T, N>) {
let buf = Arc::new(BroadcastBuffer::<T, N>::new());
let sender = BroadcastSender::new(buf.clone());
let receiver = BroadcastReceiver::new(buf, 0);
(sender, receiver)
}