use crate::Slotable;
use crate::loom::sync::atomic::{AtomicU64, Ordering::*};
use crate::loom::sync::spin_loop_hint;
use std::sync::Arc;
#[inline]
fn version_to_slot_index(version: u64, queue_len: usize) -> usize {
debug_assert!(queue_len.is_power_of_two());
(version & (queue_len as u64 - 1)) as usize
}
struct Counts {
sequence: AtomicU64,
tx_count: AtomicU64,
}
struct Inner<T, S: Slotable<T>> {
queue: Arc<[S::SlotArrayItem]>,
counts: Arc<Counts>,
queue_len: usize,
}
impl<T, S: Slotable<T>> Clone for Inner<T, S> {
fn clone(&self) -> Self {
Inner {
queue: self.queue.clone(),
counts: self.counts.clone(),
queue_len: self.queue_len,
}
}
}
pub struct Sender<T, S: Slotable<T>> {
inner: Inner<T, S>,
}
impl<T, S: Slotable<T>> Sender<T, S> {
pub fn new(capacity: usize) -> Self {
let queue_len = capacity.next_power_of_two();
let inner = Inner {
queue: S::boxed_uninit_multiple(queue_len).into(),
counts: Arc::new(Counts {
sequence: AtomicU64::new(0),
tx_count: AtomicU64::new(1),
}),
queue_len,
};
Sender { inner }
}
pub fn subscribe(&self) -> Receiver<T, S> {
let next_version = self.inner.counts.sequence.load(Relaxed) >> 1;
Receiver {
inner: self.inner.clone(),
next_version,
}
}
pub fn send(&mut self, value: T) {
let tx_count = self.inner.counts.tx_count.load(Acquire);
let mut seq = self.inner.counts.sequence.load(Relaxed);
if tx_count == 1 {
debug_assert!(seq & 1 == 0, "single sender finds queue locked");
self.inner.counts.sequence.store(seq + 1, Release);
} else {
loop {
if seq & 1 != 0 {
spin_loop_hint();
seq = self.inner.counts.sequence.load(Relaxed);
continue;
}
match self
.inner
.counts
.sequence
.compare_exchange(seq, seq + 1, AcqRel, Relaxed)
{
Ok(_) => break,
Err(c) => seq = c,
}
spin_loop_hint();
}
}
let version = seq >> 1;
let slot_index = version_to_slot_index(version, self.inner.queue_len);
let slot = S::index_in_array(&self.inner.queue, slot_index);
S::write(slot, value, Release);
self.inner.counts.sequence.store(seq + 2, Release);
}
}
impl<T, S: Slotable<T>> Clone for Sender<T, S> {
fn clone(&self) -> Self {
self.inner.counts.tx_count.fetch_add(1, Relaxed);
Sender {
inner: self.inner.clone(),
}
}
}
impl<T, S: Slotable<T>> Drop for Sender<T, S> {
fn drop(&mut self) {
self.inner.counts.tx_count.fetch_sub(1, Release);
}
}
pub struct Receiver<T, S: Slotable<T>> {
inner: Inner<T, S>,
next_version: u64,
}
impl<T, S: Slotable<T>> Clone for Receiver<T, S> {
fn clone(&self) -> Self {
Receiver {
inner: self.inner.clone(),
next_version: self.next_version,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub enum TryRecvError {
#[error("no new value available")]
Empty,
#[error("channel was closed")]
Closed,
#[error("lagged by {0} messages")]
Lagged(u64),
}
impl<T, S: Slotable<T>> Receiver<T, S> {
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
let version = self.next_version;
let seq1 = self.inner.counts.sequence.load(Acquire);
if seq1 < (version + 1) << 1 {
let err = if self.inner.counts.tx_count.load(Relaxed) == 0 {
TryRecvError::Closed
} else {
TryRecvError::Empty
};
return Err(err);
}
let slot_index = version_to_slot_index(version, self.inner.queue_len);
let slot = S::index_in_array(&self.inner.queue, slot_index);
let value = S::read(slot, Acquire);
let seq2 = self.inner.counts.sequence.load(Relaxed);
if seq2 > (version + self.inner.queue_len as u64) << 1 {
let next_version = seq2 >> 1;
self.next_version = next_version;
return Err(TryRecvError::Lagged(next_version - version));
}
self.next_version += 1;
Ok(unsafe { value.assume_init() })
}
}
pub fn channel<T, S: Slotable<T>>(capacity: usize) -> (Sender<T, S>, Receiver<T, S>) {
let sender = Sender::new(capacity);
let receiver = sender.subscribe();
(sender, receiver)
}
macro_rules! broadcast_impl {
($s:ty, $bound:path) => {
pub mod broadcast {
pub type Sender<T> = crate::channels::broadcast::Sender<T, $s>;
pub type Receiver<T> = crate::channels::broadcast::Receiver<T, $s>;
pub use crate::channels::broadcast::TryRecvError;
pub fn channel<T: $bound>(capacity: usize) -> (Sender<T>, Receiver<T>) {
crate::channels::broadcast::channel::<T, $s>(capacity)
}
}
};
}
pub(crate) use broadcast_impl;
macro_rules! def_tests {
($modname:ident,$loommodname:ident,$s:ty) => {
#[cfg(test)]
mod $modname {
use super::*;
use crate::loom::thread;
type S = $s;
fn busy_read(receiver: &mut Receiver<u32, S>) -> Result<u32, &'static str> {
loop {
match receiver.try_recv() {
Ok(x) => return Ok(x),
Err(TryRecvError::Lagged(_)) => return Err("lagged"),
Err(TryRecvError::Empty) => (),
Err(TryRecvError::Closed) => panic!("closed"),
}
spin_loop_hint();
}
}
#[test]
fn rw() {
let mut sender = Sender::<u32, S>::new(3);
let mut receiver = sender.subscribe();
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
sender.send(1);
sender.send(2);
assert_eq!(receiver.try_recv(), Ok(1));
assert_eq!(receiver.try_recv(), Ok(2));
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
sender.send(3);
sender.send(4);
assert_eq!(receiver.try_recv(), Ok(3));
assert_eq!(receiver.try_recv(), Ok(4));
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
}
#[test]
fn rw_multi() {
let mut sender1 = Sender::<u32, S>::new(3);
let mut sender2 = sender1.clone();
let mut receiver = sender1.subscribe();
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
sender1.send(1);
sender2.send(2);
assert_eq!(receiver.try_recv(), Ok(1));
assert_eq!(receiver.try_recv(), Ok(2));
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
sender1.send(3);
sender2.send(4);
assert_eq!(receiver.try_recv(), Ok(3));
assert_eq!(receiver.try_recv(), Ok(4));
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
}
#[test]
fn lagging() {
let mut sender = Sender::<u32, S>::new(3); let mut receiver = sender.subscribe();
sender.send(1); sender.send(2); sender.send(3); assert_eq!(receiver.try_recv(), Ok(1)); sender.send(4); sender.send(5); sender.send(6); assert_eq!(receiver.try_recv(), Err(TryRecvError::Lagged(5)));
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
sender.send(7);
assert_eq!(receiver.try_recv(), Ok(7));
}
#[test]
fn multi_wait() {
let mut sender1 = Sender::<u32, S>::new(3);
let mut sender2 = sender1.clone();
let _no_close = sender1.clone();
let mut receiver = sender1.subscribe();
let th1 = thread::spawn(move || sender1.send(1));
let th2 = thread::spawn(move || sender2.send(2));
let values: [u32; 2] = std::array::from_fn(|_| busy_read(&mut receiver).unwrap());
th1.join().unwrap();
th2.join().unwrap();
assert!(values == [1, 2] || values == [2, 1]);
assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty));
}
#[test]
fn single_wait() {
let mut sender = Sender::<u32, S>::new(3);
let mut receiver = sender.subscribe();
let th = thread::spawn(move || {
for i in 1..=3 {
sender.send(i);
}
});
for i in 1..=3 {
let read = busy_read(&mut receiver).unwrap();
assert_eq!(read, i);
}
th.join().unwrap();
}
}
#[cfg(all(loom, test))]
mod $loommodname {
use super::*;
use crate::loom::thread;
type S = $s;
fn busy_read(receiver: &mut Receiver<u32, S>) -> Result<u32, &'static str> {
loop {
match receiver.try_recv() {
Ok(x) => return Ok(x),
Err(TryRecvError::Lagged(_)) => return Err("lagged"),
Err(TryRecvError::Closed) => return Err("closed"),
Err(TryRecvError::Empty) => (),
}
spin_loop_hint();
}
}
#[test]
fn single_wait() {
loom::model(|| {
let mut sender = Sender::<u32, S>::new(3);
let mut receiver = sender.subscribe();
let th = thread::spawn(move || {
for i in 1..=3 {
sender.send(i);
}
sender });
for i in 1..=3 {
let read = busy_read(&mut receiver).unwrap();
assert_eq!(read, i);
}
th.join().unwrap();
});
}
#[test]
fn multi_wait() {
loom::model(|| {
let mut sender = Sender::<u32, S>::new(3);
let _other = sender.clone();
let mut receiver = sender.subscribe();
let th = thread::spawn(move || {
sender.send(1);
sender.send(2);
});
for i in 1..=2 {
let read = busy_read(&mut receiver).unwrap();
assert_eq!(read, i);
}
th.join().unwrap();
});
}
#[test]
fn multi_sender() {
loom::model(|| {
println!("-- start");
let mut sender1 = Sender::<u32, S>::new(3);
let mut sender2 = sender1.clone();
let mut receiver = sender1.subscribe();
let th1 = thread::spawn(move || {
println!("sending 1");
sender1.send(1);
println!("sending 1 done");
});
let th2 = thread::spawn(move || {
println!("sending 2");
sender2.send(2);
println!("sending 2 done");
});
th1.join().unwrap();
th2.join().unwrap();
let values: [u32; 2] = std::array::from_fn(|_| receiver.try_recv().unwrap());
assert!(values == [1, 2] || values == [2, 1]);
assert_eq!(receiver.try_recv(), Err(TryRecvError::Closed));
});
}
}
};
}
def_tests!(
native_tests,
native_loom_tests,
crate::native::NativeSlotable
);
#[cfg(feature = "atomic")]
def_tests!(
atomic_tests,
atomic_loom_tests,
crate::atomic::AtomicSlotable
);
#[cfg(not(miri))]
def_tests!(fast_tests, fast_loom_tests, crate::fast::Slotable);