copy-channels 1.0.0

A collection of cross-thread channels for copyable types
Documentation
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(),
        }
    }
}

/// Sender side of the channel.
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> {
    /// Create a new channel, and return the sender side. The channel is initially set to `value`.
    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),
                }),
            },
        }
    }

    /// Subscribe to the channel, creating a new receiver. The receiver will be up to date
    /// with the current version.
    pub fn subscribe(&self) -> Receiver<T, S> {
        let version = self.inner.counts.sequence.load(Relaxed);
        Receiver {
            inner: self.inner.clone(),
            version,
        }
    }

    /// Update the channel with `value`, and mark as changed.
    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 {
            // We take &mut self, which means a unique reference and we can't be cloning
            // self at the same time. And if tx_count = 1, then this must be the one and
            // only sender, and nowhere throughout this function can there be more. We
            // can skip the CAS to lock the object.
            debug_assert!(
                seq & 1 == 0,
                "locked, multiple senders while there shouldn't be"
            );
            // the Release here could be Relaxed if it was only about this invocation,
            // but it could be ordered before the final release in a previous invocation.
            self.inner.counts.sequence.store(seq + 1, Release);
        } else {
            // Multiple senders are (or may be) around. Use the "write-in-progress" bit as
            // a spinlock.
            loop {
                if seq & 1 != 0 {
                    spin_loop_hint();
                    seq = self.inner.counts.sequence.load(Relaxed);
                    continue;
                }
                // this AcqRel here looks unnecessary, but if Acquire it could be ordered before the final
                // release in a previous invocation
                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);
    }
}

/// Clone this sender. A [`Sender`] may be freely cloned, however if there is more than one sender
/// for a channel, there is a small performance implication (a single extra CAS).
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);
    }
}

/// Receiver end of the channel. A receiver tracks the last-seen version.
pub struct Receiver<T, S: Slotable<T>> {
    inner: Inner<T, S>,
    version: u64,
}

/// Receiver may be freely cloned, without performance impact. The clone will
/// have the same last-seen version, so this is slightly different from a new
/// subscription to the [`Sender`].
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> {
    /// Check if the channel has changed. If it has, update the last-seen version. If all senders
    /// have been dropped, an error will be returned.
    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)
        }
    }

    /// Retrieve the current value and version. This does not update the last-seen version.
    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 {
                    // safety: source is valid, sequence handling ensures validity; value was either
                    // set initially or written
                    let value = unsafe { value.assume_init() };
                    return (value, seq2);
                }
            }
            spin_loop_hint();
        }
    }

    /// Retrieve the current value. This does not update the last-seen version.
    pub fn get(&self) -> T {
        self.get_value_and_sequence().0
    }

    /// Retrieve the current value, and update the last-seen version.
    pub fn get_and_update(&mut self) -> T {
        let (value, seq) = self.get_value_and_sequence();
        self.version = seq;
        value
    }

    /// Receive a new value. If a new value is available, retrieve it and update
    /// the last-seen version. If no new value is available, [`Empty`](TryRecvError::Empty) is
    /// returned. If all senders have been dropped, [`Closed`](TryRecvError::Closed) will be returned.
    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),
        }
    }
}

/// Create a new channel, and set the initial value to `value`. Returns a single sender and receiver.
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 {

            /// Sender side of the channel.
            pub type Sender<T> = crate::channels::watch::Sender<T, $s>;

            /// Receiver side of the channel.
            pub type Receiver<T> = crate::channels::watch::Receiver<T, $s>;

            pub use crate::channels::watch::{RecvError, TryRecvError};

            /// Create a new channel, and set the initial value to `value`.
            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);
            }
        }

        // RUSTFLAGS="--cfg loom" cargo test --release loom

        #[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);