use crate::Slotable;
use crate::loom::sync::atomic::{AtomicU64, Ordering::*};
use crate::loom::sync::spin_loop_hint;
use std::sync::Arc;
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub enum RecvError {
#[error("channel was closed")]
Closed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub enum TryRecvError {
#[error("channel was closed")]
Closed,
#[error("no new value available")]
Empty,
}
struct Counts {
sequence: AtomicU64,
tx_count: AtomicU64,
}
struct Inner<T, S: Slotable<T>> {
value: Arc<S::Slot>,
counts: Arc<Counts>,
}
impl<T, S: Slotable<T>> Clone for Inner<T, S> {
fn clone(&self) -> Self {
Inner {
value: self.value.clone(),
counts: self.counts.clone(),
}
}
}
pub struct Sender<T, S: Slotable<T>> {
inner: Inner<T, S>,
}
impl<T: Default, S: Slotable<T>> Default for Sender<T, S> {
fn default() -> Self {
Sender::new(Default::default())
}
}
impl<T, S: Slotable<T>> Sender<T, S> {
pub fn new(value: T) -> Self {
Sender {
inner: Inner {
value: S::create_boxed(value).into(),
counts: Arc::new(Counts {
sequence: AtomicU64::new(0),
tx_count: AtomicU64::new(1),
}),
},
}
}
pub fn subscribe(&self) -> Receiver<T, S> {
let version = self.inner.counts.sequence.load(Relaxed);
Receiver {
inner: self.inner.clone(),
version,
}
}
pub fn send_replace(&mut self, value: T) {
let tx_count = self.inner.counts.tx_count.load(Acquire);
let mut seq = self.inner.counts.sequence.load(Relaxed);
debug_assert!(tx_count > 0, "bad tx_count");
if tx_count == 1 {
debug_assert!(
seq & 1 == 0,
"locked, multiple senders while there shouldn't be"
);
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();
}
}
S::write(&self.inner.value, 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>,
version: u64,
}
impl<T, S: Slotable<T>> Clone for Receiver<T, S> {
fn clone(&self) -> Self {
Receiver {
inner: self.inner.clone(),
version: self.version,
}
}
}
impl<T, S: Slotable<T>> Receiver<T, S> {
pub fn has_changed(&mut self) -> Result<bool, RecvError> {
let seq = self.inner.counts.sequence.load(Relaxed);
if seq > self.version {
self.version = seq;
Ok(true)
} else if self.inner.counts.tx_count.load(Relaxed) == 0 {
Err(RecvError::Closed)
} else {
Ok(false)
}
}
fn get_value_and_sequence(&self) -> (T, u64) {
loop {
let seq1 = self.inner.counts.sequence.load(Acquire);
if seq1 & 1 == 0 {
let value = S::read(&self.inner.value, Acquire);
let seq2 = self.inner.counts.sequence.load(Relaxed);
if seq2 == seq1 {
let value = unsafe { value.assume_init() };
return (value, seq2);
}
}
spin_loop_hint();
}
}
pub fn get(&self) -> T {
self.get_value_and_sequence().0
}
pub fn get_and_update(&mut self) -> T {
let (value, seq) = self.get_value_and_sequence();
self.version = seq;
value
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
match self.has_changed() {
Ok(true) => Ok(self.get_and_update()),
Ok(false) => Err(TryRecvError::Empty),
Err(RecvError::Closed) => Err(TryRecvError::Closed),
}
}
}
pub fn channel<T, S: Slotable<T>>(value: T) -> (Sender<T, S>, Receiver<T, S>) {
let sender = Sender::new(value);
let receiver = sender.subscribe();
(sender, receiver)
}
macro_rules! watch_impl {
($s:ty, $bound:path) => {
pub mod watch {
pub type Sender<T> = crate::channels::watch::Sender<T, $s>;
pub type Receiver<T> = crate::channels::watch::Receiver<T, $s>;
pub use crate::channels::watch::{RecvError, TryRecvError};
pub fn channel<T: $bound>(value: T) -> (Sender<T>, Receiver<T>) {
crate::channels::watch::channel::<T, $s>(value)
}
}
};
}
pub(crate) use watch_impl;
macro_rules! def_tests {
($modname:ident,$loommodname:ident,$s:ty) => {
#[cfg(test)]
mod $modname {
use super::*;
use crate::loom::thread;
type S = $s;
#[test]
fn get_latest() {
let (mut sender, receiver) = channel::<u32, S>(0u32);
let th = thread::spawn(move || {
for i in 1..=10 {
sender.send_replace(i);
}
});
let mut prev = 0;
for _ in 0..25 {
let value = receiver.get();
assert!(value >= prev);
prev = value;
}
th.join().unwrap();
}
#[test]
fn wait_for_change() {
let (mut sender, mut receiver) = channel::<u32, S>(1u32);
let th = thread::spawn(move || {
sender.send_replace(2);
sender.send_replace(3);
});
let value = loop {
if let Ok(value) = receiver.try_recv() {
break value;
}
spin_loop_hint();
};
th.join().unwrap();
dbg!(value);
assert!(value == 2 || value == 3);
}
#[test]
fn multi_sender_lock() {
let (mut sender1, receiver) = channel::<u32, S>(1u32);
let mut sender2 = sender1.clone();
let th1 = thread::spawn(move || {
sender1.send_replace(2);
});
let th2 = thread::spawn(move || {
sender2.send_replace(3);
});
let value = receiver.get();
th1.join().unwrap();
th2.join().unwrap();
dbg!(value);
assert!(value == 1 || value == 2 || value == 3);
}
}
#[cfg(all(loom, test))]
pub mod $loommodname {
use super::*;
use crate::loom::thread;
type S = $s;
#[test]
fn get_latest() {
loom::model(|| {
let (mut sender, receiver) = channel::<u32, S>(1u32);
sender.send_replace(1);
assert_eq!(receiver.get(), 1);
let th = thread::spawn(move || {
sender.send_replace(2);
});
let value = receiver.get();
dbg!(value);
assert!(value == 1 || value == 2);
th.join().unwrap();
assert_eq!(receiver.get(), 2);
});
}
#[test]
fn wait_for_change() {
loom::model(|| {
let (mut sender, mut receiver) = channel::<u32, S>(1u32);
let th = thread::spawn(move || {
sender.send_replace(2);
sender.send_replace(3);
});
let value = loop {
if let Ok(value) = receiver.try_recv() {
break value;
}
spin_loop_hint();
};
th.join().unwrap();
dbg!(value);
assert!(value == 2 || value == 3);
});
}
#[test]
fn multi_sender_lock() {
loom::model(|| {
let (mut sender1, receiver) = channel::<u32, S>(1u32);
let mut sender2 = sender1.clone();
let th1 = thread::spawn(move || {
sender1.send_replace(2);
});
sender2.send_replace(3);
th1.join().unwrap();
let value = receiver.get();
dbg!(value);
assert!(value == 2 || value == 3);
});
}
}
};
}
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);