use super::ring::{blocking, SpscChannel};
use crate::channel::error::{Channel, Result};
use crate::channel::roles::{Consumer, Producer};
use std::cell::Cell;
use std::sync::atomic::Ordering;
pub struct SpscRing<T> {
channel: SpscChannel<T>,
}
impl<T: Send> SpscRing<T> {
#[must_use]
pub fn new(capacity: usize) -> Self {
Self {
channel: SpscChannel::new(capacity),
}
}
pub fn split(&mut self) -> (SpscProducer<'_, T>, SpscConsumer<'_, T>) {
let (head, tail) = self.channel.indices();
self.channel.closed.store(false, Ordering::Release);
(
SpscProducer {
channel: &self.channel,
cached_tail: Cell::new(tail),
},
SpscConsumer {
channel: &self.channel,
cached_head: Cell::new(head),
},
)
}
#[must_use]
pub fn len(&self) -> usize {
let (head, tail) = self.channel.indices();
head.wrapping_sub(tail)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn capacity(&self) -> usize {
Channel::capacity(&self.channel).unwrap_or(0)
}
}
pub struct SpscProducer<'ring, T> {
channel: &'ring SpscChannel<T>,
cached_tail: Cell<usize>,
}
impl<T: Send> SpscProducer<'_, T> {
pub fn send(&self, value: T) -> Result<()> {
self.channel.send_cached(value, &self.cached_tail)
}
pub fn try_send(&self, value: T) -> Result<()> {
self.channel.try_send_cached(value, &self.cached_tail)
}
}
pub struct SpscConsumer<'ring, T> {
channel: &'ring SpscChannel<T>,
cached_head: Cell<usize>,
}
impl<T: Send> SpscConsumer<'_, T> {
pub fn recv(&self) -> Result<T> {
blocking(|| self.channel.try_recv_cached(&self.cached_head))
}
pub fn try_recv(&self) -> Result<T> {
self.channel.try_recv_cached(&self.cached_head)
}
}
impl<T: Send> Producer<T> for SpscProducer<'_, T> {
#[inline]
fn send(&self, value: T) -> Result<()> {
SpscProducer::send(self, value)
}
#[inline]
fn try_send(&self, value: T) -> Result<()> {
SpscProducer::try_send(self, value)
}
#[inline]
fn is_full(&self) -> bool {
Channel::is_full(self.channel)
}
#[inline]
fn capacity(&self) -> Option<usize> {
Channel::capacity(self.channel)
}
}
impl<T: Send> Consumer<T> for SpscConsumer<'_, T> {
#[inline]
fn recv(&self) -> Result<T> {
SpscConsumer::recv(self)
}
#[inline]
fn try_recv(&self) -> Result<T> {
SpscConsumer::try_recv(self)
}
#[inline]
fn is_empty(&self) -> bool {
Channel::is_empty(self.channel)
}
}
impl<T> Drop for SpscProducer<'_, T> {
fn drop(&mut self) {
self.channel.closed.store(true, Ordering::Release);
}
}
impl<T> Drop for SpscConsumer<'_, T> {
fn drop(&mut self) {
self.channel.closed.store(true, Ordering::Release);
}
}
#[cfg(test)]
mod auto_traits {
use super::{SpscConsumer, SpscProducer, SpscRing};
use static_assertions::{assert_impl_all, assert_not_impl_any};
assert_impl_all!(SpscRing<u64>: Send, Sync);
assert_impl_all!(SpscProducer<'static, u64>: Send);
assert_impl_all!(SpscConsumer<'static, u64>: Send);
#[allow(dead_code)]
fn halves_are_not_sync() {
assert_not_impl_any!(SpscProducer<'static, u64>: Sync);
assert_not_impl_any!(SpscConsumer<'static, u64>: Sync);
}
}
#[cfg(test)]
mod tests {
use super::SpscRing;
use crate::channel::roles::{Consumer, Producer};
use crate::channel::ChannelError;
#[test]
fn split_halves_transfer_values_in_order() {
let mut ring = SpscRing::<u64>::new(4);
let (tx, rx) = ring.split();
assert!(tx.try_send(1).is_ok());
assert!(tx.try_send(2).is_ok());
assert_eq!(rx.try_recv().expect("a value was sent"), 1);
assert_eq!(rx.try_recv().expect("a value was sent"), 2);
assert!(matches!(rx.try_recv(), Err(ChannelError::Empty)));
}
#[test]
fn capacity_is_exact_with_no_sacrificed_slot() {
let mut ring = SpscRing::<u64>::new(4);
assert_eq!(ring.capacity(), 4);
let (tx, _rx) = ring.split();
for value in 0..4 {
assert!(tx.try_send(value).is_ok(), "slot {value} must be usable");
}
assert!(matches!(tx.try_send(4), Err(ChannelError::Full)));
}
#[test]
fn resplitting_a_drained_ring_reports_empty() {
let mut ring = SpscRing::<u64>::new(4);
{
let (tx, rx) = ring.split();
for value in 0..3 {
tx.try_send(value).expect("capacity is 4");
}
for expected in 0..3 {
assert_eq!(rx.try_recv().expect("three were sent"), expected);
}
}
assert!(ring.is_empty());
let (_tx, rx) = ring.split();
assert!(
matches!(rx.try_recv(), Err(ChannelError::Empty)),
"a drained ring must report empty, not read an unwritten slot"
);
}
#[test]
fn resplitting_preserves_queued_values() {
let mut ring = SpscRing::<u64>::new(8);
{
let (tx, _rx) = ring.split();
tx.try_send(7).expect("capacity is 8");
tx.try_send(8).expect("capacity is 8");
}
assert_eq!(ring.len(), 2);
let (_tx, rx) = ring.split();
assert_eq!(rx.try_recv().expect("queued across the split"), 7);
assert_eq!(rx.try_recv().expect("queued across the split"), 8);
}
#[test]
fn dropping_the_producer_closes_the_consumer() {
let mut ring = SpscRing::<u64>::new(4);
let (tx, rx) = ring.split();
tx.try_send(1).expect("capacity is 4");
drop(tx);
assert_eq!(rx.recv().expect("the queued value survives closure"), 1);
assert!(matches!(rx.recv(), Err(ChannelError::Closed)));
}
#[test]
fn borrowed_halves_satisfy_the_roles() {
fn drain_into<P: Producer<u64>>(producer: &P, values: &[u64]) -> usize {
values
.iter()
.filter(|v| producer.try_send(**v).is_ok())
.count()
}
fn take_all<C: Consumer<u64>>(consumer: &C, limit: usize) -> Vec<u64> {
(0..limit)
.filter_map(|_| consumer.try_recv().ok())
.collect()
}
let mut ring = SpscRing::<u64>::new(8);
let (tx, rx) = ring.split();
assert_eq!(drain_into(&tx, &[1, 2, 3]), 3);
assert_eq!(Producer::capacity(&tx), Some(8));
assert!(!Consumer::is_empty(&rx));
assert_eq!(take_all(&rx, 8), vec![1, 2, 3]);
}
#[test]
fn dropping_a_loaded_ring_drops_each_value_once() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
struct Counted(Arc<AtomicUsize>);
impl Drop for Counted {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
let drops = Arc::new(AtomicUsize::new(0));
{
let mut ring = SpscRing::<Counted>::new(4);
let (tx, _rx) = ring.split();
for _ in 0..3 {
tx.try_send(Counted(Arc::clone(&drops)))
.expect("capacity is 4");
}
}
assert_eq!(
drops.load(Ordering::Relaxed),
3,
"every queued value is dropped exactly once with the ring"
);
}
#[test]
fn scoped_threads_move_the_halves_across_the_boundary() {
const COUNT: u64 = 10_000;
let mut ring = SpscRing::<u64>::new(64);
let (tx, rx) = ring.split();
let sum = std::thread::scope(|scope| {
scope.spawn(move || {
for value in 0..COUNT {
if tx.send(value).is_err() {
break;
}
}
});
let mut sum = 0_u64;
for _ in 0..COUNT {
match rx.recv() {
Ok(value) => sum = sum.wrapping_add(value),
Err(_) => break,
}
}
sum
});
assert_eq!(sum, (0..COUNT).sum::<u64>());
}
}