use crate::ringbuffer::RingBuffer;
use crate::wait::{TryWaitStrategy, WaitStrategy};
use std::collections::Bound;
use std::ops::RangeBounds;
use std::sync::atomic::{fence, AtomicI64, Ordering};
use std::sync::Arc;
#[derive(Debug)]
pub struct MultiProducer<E, W, const LEAD: bool> {
inner: HandleInner<E, W, LEAD>,
claim: Arc<Cursor>, }
impl<E, W, const LEAD: bool> Clone for MultiProducer<E, W, LEAD>
where
W: Clone,
{
fn clone(&self) -> Self {
MultiProducer {
inner: HandleInner {
cursor: Arc::clone(&self.inner.cursor),
barrier: self.inner.barrier.clone(),
buffer: Arc::clone(&self.inner.buffer),
wait_strategy: self.inner.wait_strategy.clone(),
available: self.inner.available,
},
claim: Arc::clone(&self.claim),
}
}
}
impl<E, W, const LEAD: bool> MultiProducer<E, W, LEAD> {
pub const fn is_lead(&self) -> bool {
LEAD
}
#[inline]
pub fn count(&self) -> usize {
Arc::strong_count(&self.claim)
}
#[inline]
pub fn sequence(&self) -> i64 {
self.inner.cursor.sequence.load(Ordering::Relaxed)
}
#[inline]
pub fn buffer_size(&self) -> usize {
self.inner.buffer_size()
}
#[inline]
pub fn set_wait_strategy<W2>(self, wait_strategy: W2) -> MultiProducer<E, W2, LEAD> {
MultiProducer {
inner: self.inner.set_wait_strategy(wait_strategy),
claim: self.claim,
}
}
#[inline]
pub fn into_producer(self) -> Option<Producer<E, W, LEAD>> {
Arc::into_inner(self.claim).map(|_| self.inner.into_producer())
}
#[inline]
fn wait_bounds(&self, size: i64) -> (i64, i64) {
let mut current_claim = self.claim.sequence.load(Ordering::Relaxed);
let mut claim_end = current_claim + size;
while let Err(new_current) = self.claim.sequence.compare_exchange(
current_claim,
claim_end,
Ordering::AcqRel,
Ordering::Relaxed,
) {
current_claim = new_current;
claim_end = new_current + size;
}
let desired_seq = if LEAD {
claim_end - saturate_i64(self.inner.buffer.size())
} else {
claim_end
};
(current_claim, desired_seq)
}
#[inline]
fn update_cursor(&self, start: i64, end: i64) {
while self.inner.cursor.sequence.load(Ordering::Acquire) != start {
std::hint::spin_loop();
}
self.inner.cursor.sequence.store(end, Ordering::Release)
}
}
impl<E, W, const LEAD: bool> MultiProducer<E, W, LEAD>
where
W: WaitStrategy,
{
#[inline]
pub fn wait_for_each<F>(&mut self, size: usize, mut f: F)
where
F: FnMut(&mut E, i64, bool),
{
debug_assert!(size <= self.inner.buffer.size());
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(size));
if self.inner.available < till_seq {
self.inner.available = self.inner.wait_strategy.wait(till_seq, &self.inner.barrier);
}
debug_assert!(self.inner.available >= till_seq);
fence(Ordering::Acquire);
unsafe {
self.inner.buffer.apply(from_seq + 1, size, |ptr, seq, end| {
let event: &mut E = &mut *ptr;
f(event, seq, end)
})
};
self.update_cursor(from_seq, from_seq + saturate_i64(size))
}
}
impl<E, W, const LEAD: bool> MultiProducer<E, W, LEAD>
where
W: TryWaitStrategy,
{
#[inline]
pub fn try_wait_for_each<F, Err>(&mut self, size: usize, mut f: F) -> Result<(), Err>
where
F: FnMut(&mut E, i64, bool) -> Result<(), Err>,
Err: From<W::Error>,
{
debug_assert!(size <= self.inner.buffer.size());
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(size));
fence(Ordering::Acquire);
if self.inner.available < till_seq {
self.inner.available = self
.inner
.wait_strategy
.try_wait(till_seq, &self.inner.barrier)
.inspect_err(|_| self.update_cursor(from_seq, from_seq + saturate_i64(size)))?;
}
debug_assert!(self.inner.available >= till_seq);
let result = unsafe {
self.inner.buffer.try_apply(from_seq + 1, size, |ptr, seq, end| {
let event: &mut E = &mut *ptr;
f(event, seq, end)
})
};
self.update_cursor(from_seq, from_seq + saturate_i64(size));
result
}
}
#[derive(Debug)]
#[repr(transparent)]
pub struct Producer<E, W, const LEAD: bool> {
inner: HandleInner<E, W, LEAD>,
}
impl<E, W, const LEAD: bool> Producer<E, W, LEAD> {
pub const fn is_lead(&self) -> bool {
LEAD
}
#[inline]
pub fn sequence(&self) -> i64 {
self.inner.cursor.sequence.load(Ordering::Relaxed)
}
#[inline]
pub fn buffer_size(&self) -> usize {
self.inner.buffer_size()
}
#[inline]
pub fn set_wait_strategy<W2>(self, wait_strategy: W2) -> Producer<E, W2, LEAD> {
self.inner.set_wait_strategy(wait_strategy).into_producer()
}
#[inline]
pub fn into_multi(self) -> MultiProducer<E, W, LEAD> {
let producer_seq = self.sequence();
MultiProducer {
inner: self.inner,
claim: Arc::new(Cursor::new(producer_seq)),
}
}
}
impl<E, W, const LEAD: bool> Producer<E, W, LEAD>
where
W: WaitStrategy,
{
#[inline]
pub fn wait(&mut self, size: usize) -> EventsMut<'_, E> {
EventsMut(self.inner.wait(size))
}
#[inline]
pub fn wait_range<R>(&mut self, range: R) -> EventsMut<'_, E>
where
R: RangeBounds<usize>,
{
EventsMut(self.inner.wait_range(range))
}
}
impl<E, W, const LEAD: bool> Producer<E, W, LEAD>
where
W: TryWaitStrategy,
{
#[inline]
pub fn try_wait(&mut self, size: usize) -> Result<EventsMut<'_, E>, W::Error> {
self.inner.try_wait(size).map(EventsMut)
}
#[inline]
pub fn try_wait_range<R>(&mut self, range: R) -> Result<EventsMut<'_, E>, W::Error>
where
R: RangeBounds<usize>,
{
self.inner.try_wait_range(range).map(EventsMut)
}
}
#[derive(Debug)]
#[repr(transparent)]
pub struct Consumer<E, W> {
inner: HandleInner<E, W, false>,
}
impl<E, W> Consumer<E, W> {
#[inline]
pub fn sequence(&self) -> i64 {
self.inner.cursor.sequence.load(Ordering::Relaxed)
}
#[inline]
pub fn buffer_size(&self) -> usize {
self.inner.buffer_size()
}
#[inline]
pub fn set_wait_strategy<W2>(self, wait_strategy: W2) -> Consumer<E, W2> {
self.inner.set_wait_strategy(wait_strategy).into_consumer()
}
}
impl<E, W> Consumer<E, W>
where
W: WaitStrategy,
{
#[inline]
pub fn wait(&mut self, size: usize) -> Events<'_, E> {
Events(self.inner.wait(size))
}
#[inline]
pub fn wait_range<R>(&mut self, range: R) -> Events<'_, E>
where
R: RangeBounds<usize>,
{
Events(self.inner.wait_range(range))
}
}
impl<E, W> Consumer<E, W>
where
W: TryWaitStrategy,
{
#[inline]
pub fn try_wait(&mut self, size: usize) -> Result<Events<'_, E>, W::Error> {
self.inner.try_wait(size).map(Events)
}
#[inline]
pub fn try_wait_range<R>(&mut self, range: R) -> Result<Events<'_, E>, W::Error>
where
R: RangeBounds<usize>,
{
self.inner.try_wait_range(range).map(Events)
}
}
#[derive(Debug)]
pub(crate) struct HandleInner<E, W, const LEAD: bool> {
pub(crate) cursor: Arc<Cursor>,
pub(crate) barrier: Barrier,
pub(crate) buffer: Arc<RingBuffer<E>>,
pub(crate) wait_strategy: W,
pub(crate) available: i64,
}
impl<E, W> HandleInner<E, W, false> {
#[inline]
pub(crate) fn into_consumer(self) -> Consumer<E, W> {
Consumer { inner: self }
}
}
impl<E, W, const LEAD: bool> HandleInner<E, W, LEAD> {
pub(crate) fn new(
cursor: Arc<Cursor>,
barrier: Barrier,
buffer: Arc<RingBuffer<E>>,
wait_strategy: W,
) -> Self {
let available = if LEAD {
CURSOR_START - (buffer.size() as i64)
} else {
CURSOR_START
};
HandleInner {
cursor,
barrier,
buffer,
wait_strategy,
available,
}
}
#[inline]
pub(crate) fn into_producer(self) -> Producer<E, W, LEAD> {
Producer { inner: self }
}
#[inline]
fn buffer_size(&self) -> usize {
self.buffer.size()
}
#[inline]
fn set_wait_strategy<W2>(self, wait_strategy: W2) -> HandleInner<E, W2, LEAD> {
HandleInner {
cursor: self.cursor,
barrier: self.barrier,
buffer: self.buffer,
wait_strategy,
available: self.available,
}
}
#[inline]
fn wait_bounds(&self, size: i64) -> (i64, i64) {
let from_sequence = self.cursor.sequence.load(Ordering::Relaxed);
let batch_end = from_sequence + size;
let till_sequence = if LEAD {
batch_end - saturate_i64(self.buffer.size())
} else {
batch_end
};
(from_sequence, till_sequence)
}
#[inline]
fn as_batch(&mut self, from_seq: i64, batch_size: usize) -> Batch<'_, E> {
Batch {
cursor: &mut self.cursor,
buffer: &self.buffer,
current: from_seq,
size: batch_size,
}
}
#[inline]
fn range_batch_size(&self, from_seq: i64, end_bound: Bound<&usize>) -> usize {
let from = if LEAD {
from_seq - saturate_i64(self.buffer.size())
} else {
from_seq
};
let available_batch = (self.available - from).unsigned_abs() as usize;
match end_bound {
Bound::Included(b) => available_batch.min(*b),
Bound::Excluded(b) => available_batch.min(b.saturating_sub(1)),
Bound::Unbounded => available_batch,
}
}
}
impl<E, W, const LEAD: bool> HandleInner<E, W, LEAD>
where
W: WaitStrategy,
{
#[inline]
fn wait(&mut self, size: usize) -> Batch<'_, E> {
debug_assert!(self.buffer.size() >= size);
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(size));
if self.available < till_seq {
self.available = self.wait_strategy.wait(till_seq, &self.barrier);
}
debug_assert!(self.available >= till_seq);
self.as_batch(from_seq, size)
}
#[inline]
fn wait_range<R>(&mut self, range: R) -> Batch<'_, E>
where
R: RangeBounds<usize>,
{
let batch_min = match range.start_bound() {
Bound::Included(b) => *b,
Bound::Excluded(b) => b.saturating_add(1),
Bound::Unbounded => 0,
};
debug_assert!(self.buffer.size() >= batch_min);
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(batch_min));
if self.available < till_seq.max(1) {
self.available = self.wait_strategy.wait(till_seq, &self.barrier);
}
debug_assert!(self.available >= till_seq);
let batch_max = self.range_batch_size(from_seq, range.end_bound());
self.as_batch(from_seq, batch_max)
}
}
impl<E, W, const LEAD: bool> HandleInner<E, W, LEAD>
where
W: TryWaitStrategy,
{
#[inline]
fn try_wait(&mut self, size: usize) -> Result<Batch<'_, E>, W::Error> {
debug_assert!(self.buffer.size() >= size);
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(size));
if self.available < till_seq {
self.available = self.wait_strategy.try_wait(till_seq, &self.barrier)?;
}
debug_assert!(self.available >= till_seq);
Ok(self.as_batch(from_seq, size))
}
#[inline]
fn try_wait_range<R>(&mut self, range: R) -> Result<Batch<'_, E>, W::Error>
where
R: RangeBounds<usize>,
{
let batch_min = match range.start_bound() {
Bound::Included(b) => *b,
Bound::Excluded(b) => b.saturating_add(1),
Bound::Unbounded => 0,
};
debug_assert!(self.buffer.size() >= batch_min);
let (from_seq, till_seq) = self.wait_bounds(saturate_i64(batch_min));
if self.available < till_seq.max(1) {
self.available = self.wait_strategy.try_wait(till_seq, &self.barrier)?;
}
debug_assert!(self.available >= till_seq);
let batch_max = self.range_batch_size(from_seq, range.end_bound());
Ok(self.as_batch(from_seq, batch_max))
}
}
struct Batch<'a, E> {
cursor: &'a mut Arc<Cursor>, buffer: &'a Arc<RingBuffer<E>>,
current: i64,
size: usize,
}
impl<E> Batch<'_, E> {
#[inline]
fn apply<F>(self, f: F)
where
F: FnMut(*mut E, i64, bool),
{
fence(Ordering::Acquire);
unsafe { self.buffer.apply(self.current + 1, self.size, f) };
let seq_end = self.current + saturate_i64(self.size);
self.cursor.sequence.store(seq_end, Ordering::Release);
}
#[inline]
fn try_apply<F, Err>(self, f: F) -> Result<(), Err>
where
F: FnMut(*mut E, i64, bool) -> Result<(), Err>,
{
fence(Ordering::Acquire);
unsafe { self.buffer.try_apply(self.current + 1, self.size, f)? };
let seq_end = self.current + saturate_i64(self.size);
self.cursor.sequence.store(seq_end, Ordering::Release);
Ok(())
}
#[inline]
fn try_commit<F, Err>(self, mut f: F) -> Result<(), Err>
where
F: FnMut(*mut E, i64, bool) -> Result<(), Err>,
{
fence(Ordering::Acquire);
unsafe {
self.buffer.try_apply(self.current + 1, self.size, |ptr, seq, end| {
f(ptr, seq, end)
.inspect_err(|_| self.cursor.sequence.store(seq - 1, Ordering::Release))
})?;
}
let seq_end = self.current + saturate_i64(self.size);
self.cursor.sequence.store(seq_end, Ordering::Release);
Ok(())
}
}
#[repr(transparent)]
pub struct EventsMut<'a, E>(Batch<'a, E>);
impl<E> EventsMut<'_, E> {
#[inline]
pub fn size(&self) -> usize {
self.0.size
}
#[inline]
pub fn for_each<F>(self, mut f: F)
where
F: FnMut(&mut E, i64, bool),
{
self.0.apply(|ptr, seq, end| {
let event: &mut E = unsafe { &mut *ptr };
f(event, seq, end)
})
}
#[inline]
pub fn try_for_each<F, Err>(self, mut f: F) -> Result<(), Err>
where
F: FnMut(&mut E, i64, bool) -> Result<(), Err>,
{
self.0.try_apply(|ptr, seq, end| {
let event: &mut E = unsafe { &mut *ptr };
f(event, seq, end)
})
}
#[inline]
pub fn try_commit_each<F, Err>(self, mut f: F) -> Result<(), Err>
where
F: FnMut(&mut E, i64, bool) -> Result<(), Err>,
{
self.0.try_commit(|ptr, seq, end| {
let event: &mut E = unsafe { &mut *ptr };
f(event, seq, end)
})
}
}
#[repr(transparent)]
pub struct Events<'a, E>(Batch<'a, E>);
impl<E> Events<'_, E> {
#[inline]
pub fn size(&self) -> usize {
self.0.size
}
#[inline]
pub fn for_each<F>(self, mut f: F)
where
F: FnMut(&E, i64, bool),
{
self.0.apply(|ptr, seq, end| {
let event: &E = unsafe { &*ptr };
f(event, seq, end)
})
}
#[inline]
pub fn try_for_each<F, Err>(self, mut f: F) -> Result<(), Err>
where
F: FnMut(&E, i64, bool) -> Result<(), Err>,
{
self.0.try_apply(|ptr, seq, end| {
let event: &E = unsafe { &*ptr };
f(event, seq, end)
})
}
#[inline]
pub fn try_commit_each<F, Err>(self, mut f: F) -> Result<(), Err>
where
F: FnMut(&E, i64, bool) -> Result<(), Err>,
{
self.0.try_commit(|ptr, seq, end| {
let event: &E = unsafe { &*ptr };
f(event, seq, end)
})
}
}
#[derive(Debug)]
#[repr(transparent)]
pub(crate) struct Cursor {
#[cfg(not(feature = "cache-padded"))]
sequence: AtomicI64,
#[cfg(feature = "cache-padded")]
sequence: crossbeam_utils::CachePadded<AtomicI64>,
}
const CURSOR_START: i64 = -1;
impl Cursor {
pub(crate) const fn new(seq: i64) -> Self {
Cursor {
#[cfg(not(feature = "cache-padded"))]
sequence: AtomicI64::new(seq),
#[cfg(feature = "cache-padded")]
sequence: crossbeam_utils::CachePadded::new(AtomicI64::new(seq)),
}
}
pub(crate) const fn start() -> Self {
Cursor::new(CURSOR_START)
}
}
#[derive(Clone, Debug)]
#[repr(transparent)]
pub struct Barrier(Barrier_);
#[derive(Clone, Debug)]
enum Barrier_ {
One(Arc<Cursor>),
Many(Box<[Arc<Cursor>]>),
}
impl Barrier {
pub(crate) fn one(cursor: Arc<Cursor>) -> Self {
Barrier(Barrier_::One(cursor))
}
pub(crate) fn many(cursors: Box<[Arc<Cursor>]>) -> Self {
Barrier(Barrier_::Many(cursors))
}
#[inline]
pub fn sequence(&self) -> i64 {
match &self.0 {
Barrier_::One(cursor) => cursor.sequence.load(Ordering::Relaxed),
Barrier_::Many(cursors) => cursors.iter().fold(i64::MAX, |seq, cursor| {
seq.min(cursor.sequence.load(Ordering::Relaxed))
}),
}
}
}
impl Drop for Barrier {
fn drop(&mut self) {
self.0 = Barrier_::Many(Box::new([]));
}
}
#[allow(clippy::cast_possible_truncation)]
#[allow(clippy::cast_possible_wrap)]
#[inline]
const fn saturate_i64(u: usize) -> i64 {
if const { size_of::<usize>() >= 8 } {
(u & i64::MAX as usize) as i64
} else {
u as i64
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::wait::WaitBusy;
#[test]
fn sizes() {
assert_eq!(size_of::<Consumer<u8, WaitBusy>>(), 40);
assert_eq!(size_of::<Producer<u8, WaitBusy, true>>(), 40);
assert_eq!(size_of::<MultiProducer<u8, WaitBusy, true>>(), 48);
}
#[test]
fn test_wait_range() {
let buffer = Arc::new(RingBuffer::from_factory(16, || 0));
let lead_cursor = Arc::new(Cursor::new(8));
let follows_cursor = Arc::new(Cursor::new(4));
let mut lead_handle = HandleInner::<_, _, true>::new(
Arc::clone(&lead_cursor),
Barrier::one(Arc::clone(&follows_cursor)),
Arc::clone(&buffer),
WaitBusy,
);
let mut follows_handle = HandleInner::<_, _, false>::new(
Arc::clone(&follows_cursor),
Barrier::one(Arc::clone(&lead_cursor)),
Arc::clone(&buffer),
WaitBusy,
);
let lead_batch = lead_handle.wait_range(1..);
assert_eq!(lead_batch.current, 8);
assert_eq!(lead_batch.size, 12);
let follows_batch = follows_handle.wait_range(1..);
assert_eq!(follows_batch.current, 4);
assert_eq!(follows_batch.size, 4);
}
#[test]
fn test_size_zero_apply() {
let buffer = Arc::new(RingBuffer::from_factory(16, || 0));
let mut lead_handle = HandleInner::<_, _, true>::new(
Arc::new(Cursor::new(8)),
Barrier::one(Arc::new(Cursor::new(4))),
Arc::clone(&buffer),
WaitBusy,
);
let sequence_before_apply = lead_handle.cursor.sequence.load(Ordering::Relaxed);
let batch = lead_handle.wait(0);
assert_eq!(batch.size, 0);
batch.apply(|_, _, _| assert!(false));
let sequence_after_apply = lead_handle.cursor.sequence.load(Ordering::Relaxed);
assert_eq!(sequence_before_apply, sequence_after_apply);
}
}