use std::fmt;
use std::mem;
use std::ptr;
use std::sync::atomic::{Ordering, spin_loop_hint};
pub struct AtomicRingBuffer<T: Sized> {
read_counters: CounterStore,
write_counters: CounterStore,
cap_mask: usize,
ptr: *mut [T],
}
unsafe impl<T: Send> Send for AtomicRingBuffer<T> {}
unsafe impl<T: Send> Sync for AtomicRingBuffer<T> {}
impl<T: Sized> AtomicRingBuffer<T> {
pub fn with_capacity(capacity: usize) -> AtomicRingBuffer<T> {
if capacity > (::std::usize::MAX >> 16) + 1 {
panic!("too large!");
}
let cap = capacity.next_power_of_two();
let cap_mask = cap - 1;
let mut content: Vec<T> = Vec::with_capacity(cap);
unsafe { content.set_len(cap); }
let ptr = Box::into_raw(content.into_boxed_slice());
AtomicRingBuffer {
cap_mask,
ptr,
read_counters: CounterStore::new(),
write_counters: CounterStore::new(),
}
}
#[inline(always)]
pub fn try_push(&self, content: T) -> Result<(), T> {
let mut to_write_index;
let mut write_counters = self.write_counters.load(Ordering::Acquire);
loop {
let write_in_process_count = write_counters.in_process_count();
if write_in_process_count == 255 {
spin_loop_hint();
write_counters = self.write_counters.load(Ordering::Acquire);
continue;
}
let write_idx = write_counters.index();
to_write_index = write_idx.wrapping_add(write_in_process_count as usize) & self.cap_mask;
if to_write_index.wrapping_add(1) & self.cap_mask == self.read_counters.load(Ordering::SeqCst).index() {
return Err(content);
}
let new_counters = write_counters.increment_in_process();
match self.write_counters.compare_and_exchange_weak(write_counters, new_counters, Ordering::Acquire, Ordering::Relaxed) {
Ok(_) => {
write_counters = new_counters;
break;
}
Err(n) => write_counters = n
};
}
unsafe {
ptr::write(&mut (*self.ptr)[to_write_index], content);
}
loop {
let new_counters = write_counters.increment_done(self.cap_mask);
match self.write_counters.compare_and_exchange_weak(write_counters, new_counters, Ordering::Release, Ordering::Relaxed) {
Ok(_) => return Ok(()),
Err(previous) => write_counters = previous
};
}
}
#[inline]
pub fn push_overwrite(&self, content: T) {
let mut cont = content;
loop {
let option = self.try_push(cont);
if option.is_ok() {
return;
}
self.remove_if_full();
cont = option.err().unwrap();
}
}
#[inline(always)]
pub fn try_pop(&self) -> Option<T> {
let mut read_counters = self.read_counters.load(Ordering::Acquire);
let mut to_read_index;
loop {
let read_in_process_count = read_counters.in_process_count();
if read_in_process_count == 255 {
spin_loop_hint();
read_counters = self.read_counters.load(Ordering::Acquire);
continue;
}
to_read_index = read_counters.index().wrapping_add(read_in_process_count as usize) & self.cap_mask;
if to_read_index == self.write_counters.load(Ordering::SeqCst).index() {
return None;
}
let new_counters = read_counters.increment_in_process();
match self.read_counters.compare_and_exchange_weak(read_counters, new_counters, Ordering::Acquire, Ordering::Relaxed) {
Ok(_) => {
read_counters = new_counters;
break;
}
Err(n) => read_counters = n
};
}
let popped = unsafe {
ptr::read(&mut (*self.ptr)[to_read_index])
};
loop {
let new_counters = read_counters.increment_done(self.cap_mask);
match self.read_counters.compare_and_exchange_weak(read_counters, new_counters, Ordering::Release, Ordering::Relaxed) {
Ok(_) => {
break;
}
Err(n) => read_counters = n
};
}
Some(popped)
}
#[inline]
pub fn size(&self) -> usize {
let read_counters = self.read_counters.load(Ordering::SeqCst);
let write_counters = self.write_counters.load(Ordering::SeqCst);
counter_size(read_counters, write_counters, self.cap_mask + 1)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.size() == 0
}
#[inline]
pub fn cap(&self) -> usize {
return self.cap_mask + 1;
}
#[inline]
pub fn remaining_cap(&self) -> usize {
let read_counters = self.read_counters.load(Ordering::SeqCst);
let write_counters = self.write_counters.load(Ordering::SeqCst);
let cap = self.cap_mask + 1;
let read_index = read_counters.index();
let write_index = write_counters.index();
let size = if read_index <= write_index { write_index - read_index } else { write_index + cap - read_index };
cap - 1 - size - write_counters.in_process_count() as usize
}
#[inline]
pub fn clear(&self) {
while let Some(_) = self.try_pop() {}
}
pub fn memory_usage(&self) -> usize {
unsafe { mem::size_of_val(&(*self.ptr)) }
}
fn remove_if_full(&self) -> Option<T> {
let mut read_counters = self.read_counters.load(Ordering::Acquire);
let mut to_read_index;
loop {
let read_in_process_count = read_counters.in_process_count();
if read_in_process_count == 255 {
spin_loop_hint();
read_counters = self.read_counters.load(Ordering::Acquire);
continue;
}
if read_in_process_count > 0 {
return None;
}
to_read_index = read_counters.index();
if to_read_index.wrapping_add(1) & self.cap_mask == self.write_counters.load(Ordering::Acquire).index() {
return None;
}
let new_counters = read_counters.increment_in_process();
let existing = self.read_counters.compare_and_swap(read_counters, new_counters, Ordering::Acquire);
if existing == read_counters {
read_counters = new_counters;
break;
}
read_counters = existing;
}
let popped = unsafe {
ptr::read(&mut (*self.ptr)[to_read_index])
};
loop {
let new_counters = read_counters.increment_done(self.cap_mask);
match self.read_counters.compare_and_exchange_weak(read_counters, new_counters, Ordering::Release, Ordering::Relaxed) {
Ok(_) => {
break;
}
Err(n) => read_counters = n
};
}
Some(popped)
}
}
impl<T> fmt::Debug for AtomicRingBuffer<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
if f.alternate() {
let cap = self.cap_mask + 1;
let read_counters = self.read_counters.load(Ordering::Relaxed);
let write_counters = self.write_counters.load(Ordering::Relaxed);
write!(f, "AtomicRingBuffer cap: {} size: {} read_index: {}, read_in_process_count: {}, read_done_count: {}, write_index: {}, write_in_process_count: {}, write_done_count: {}", cap, self.size(),
read_counters.index(), read_counters.in_process_count(), read_counters.done_count(),
write_counters.index(), write_counters.in_process_count(), write_counters.done_count())
} else {
write!(f, "AtomicRingBuffer cap: {} size: {}", self.cap_mask + 1, self.size())
}
}
}
#[derive(Eq, PartialEq, Copy, Clone)]
pub struct Counters(usize);
fn counter_size(read_counters: Counters, write_counters: Counters, cap: usize) -> usize {
let read_index = read_counters.index();
let write_index = write_counters.index();
let size = if read_index <= write_index { write_index - read_index } else { write_index + cap - read_index };
size - (read_counters.in_process_count() as usize)
}
impl Counters {
#[inline(always)]
fn index(&self) -> usize {
self.0 >> 16
}
#[inline(always)]
fn in_process_count(&self) -> u8 {
self.0 as u8
}
#[inline(always)]
fn done_count(&self) -> u8 {
(self.0 >> 8) as u8
}
#[inline(always)]
fn increment_done(&self, cap_mask: usize) -> Counters {
let in_process_count = self.in_process_count();
Counters(if self.done_count() + 1 == in_process_count {
((self.index().wrapping_add(in_process_count as usize) & cap_mask) << 16)
} else {
self.0 + (1 << 8)
})
}
#[inline(always)]
fn increment_in_process(&self) -> Counters {
Counters(self.0 + 1)
}
}
struct CounterStore {
counters: ::std::sync::atomic::AtomicUsize,
}
impl CounterStore {
pub fn new() -> CounterStore {
CounterStore { counters: ::std::sync::atomic::AtomicUsize::new(0) }
}
#[inline(always)]
pub fn load(&self, ordering: Ordering) -> Counters {
Counters(self.counters.load(ordering))
}
#[inline(always)]
pub fn compare_and_swap(&self, old: Counters, new: Counters, ordering: Ordering) -> Counters {
Counters(self.counters.compare_and_swap(old.0, new.0, ordering))
}
#[inline(always)]
pub fn compare_and_exchange_weak(&self, old: Counters, new: Counters, success: Ordering, failure: Ordering) -> Result<Counters, Counters> {
match self.counters.compare_exchange_weak(old.0, new.0, success, failure) {
Ok(_) => Ok(old),
Err(previous) => Err(Counters(previous))
}
}
}
impl<T> Drop for AtomicRingBuffer<T> {
fn drop(&mut self) {
self.clear();
unsafe { Box::from_raw(self.ptr).into_vec().set_len(0); }
}
}
#[cfg(test)]
mod tests {
#[test]
pub fn test_increments() {
let mut read_counters = super::Counters(0);
let mut write_counters = super::Counters(0);
let cap = 16;
let cap_mask = 0xf;
for i in 0..8 {
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_in_process();
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 1, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_in_process();
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 2, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_in_process();
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 3, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_done(cap_mask);
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 3, 1), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_done(cap_mask);
assert_eq!((0, (0 + i * 3) % 16, 0, 0, (0 + i * 3) % 16, 3, 2), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
write_counters = write_counters.increment_done(cap_mask);
assert_eq!((3, (0 + i * 3) % 16, 0, 0, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_in_process();
assert_eq!((2, (0 + i * 3) % 16, 1, 0, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_in_process();
assert_eq!((1, (0 + i * 3) % 16, 2, 0, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_in_process();
assert_eq!((0, (0 + i * 3) % 16, 3, 0, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_done(cap_mask);
assert_eq!((0, (0 + i * 3) % 16, 3, 1, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_done(cap_mask);
assert_eq!((0, (0 + i * 3) % 16, 3, 2, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
read_counters = read_counters.increment_done(cap_mask);
assert_eq!((0, (3 + i * 3) % 16, 0, 0, (3 + i * 3) % 16, 0, 0), (super::counter_size(read_counters, write_counters, cap), read_counters.index(), read_counters.in_process_count(), read_counters.done_count(), write_counters.index(), write_counters.in_process_count(), write_counters.done_count()));
}
}
#[test]
pub fn test_pushpop() {
let ring = super::AtomicRingBuffer::with_capacity(900);
assert_eq!(None, ring.try_pop());
ring.push_overwrite(1);
assert_eq!(Some(1), ring.try_pop());
assert_eq!(None, ring.try_pop());
for i in 0..5000 {
ring.push_overwrite(i);
assert_eq!(Some(i), ring.try_pop());
assert_eq!(None, ring.try_pop());
}
for i in 0..199999 {
ring.push_overwrite(i);
}
assert_eq!(ring.cap(), ring.size() + 1);
assert_eq!(199999 - (ring.cap() - 1), ring.try_pop().unwrap());
assert_eq!(Ok(()), ring.try_push(199999));
for i in 200000 - (ring.cap() - 1)..200000 {
assert_eq!(i, ring.try_pop().unwrap());
}
}
#[test]
pub fn test_pushpop_large() {
let ring = super::AtomicRingBuffer::with_capacity(65535);
assert_eq!(None, ring.try_pop());
ring.push_overwrite(1);
assert_eq!(Some(1), ring.try_pop());
for i in 0..200000 {
ring.push_overwrite(i);
assert_eq!(Some(i), ring.try_pop());
}
for i in 0..200000 {
ring.push_overwrite(i);
}
assert_eq!(ring.cap(), ring.size() + 1);
for i in 200000 - (ring.cap() - 1)..200000 {
assert_eq!(i, ring.try_pop().unwrap());
}
}
#[test]
pub fn test_pushpop_large2() {
let ring = super::AtomicRingBuffer::with_capacity(65536);
assert_eq!(None, ring.try_pop());
ring.push_overwrite(1);
assert_eq!(Some(1), ring.try_pop());
for i in 0..200000 {
ring.push_overwrite(i);
assert_eq!(Some(i), ring.try_pop());
}
for i in 0..200000 {
ring.push_overwrite(i);
}
assert_eq!(ring.cap(), ring.size() + 1);
for i in 200000 - (ring.cap() - 1)..200000 {
assert_eq!(i, ring.try_pop().unwrap());
}
}
#[test]
pub fn test_pushpop_large2_zerotype() {
#[derive(Eq, PartialEq, Debug)]
struct ZeroType {}
let ring = super::AtomicRingBuffer::with_capacity(65536);
assert_eq!(0, ring.memory_usage());
assert_eq!(None, ring.try_pop());
ring.push_overwrite(ZeroType {});
assert_eq!(Some(ZeroType {}), ring.try_pop());
for _i in 0..200000 {
ring.push_overwrite(ZeroType {});
assert_eq!(Some(ZeroType {}), ring.try_pop());
}
for _i in 0..200000 {
ring.push_overwrite(ZeroType {});
}
assert_eq!(ring.cap(), ring.size() + 1);
for _i in 200000 - (ring.cap() - 1)..200000 {
assert_eq!(ZeroType {}, ring.try_pop().unwrap());
}
}
#[test]
pub fn test_threaded() {
let cap = 65535;
let buf: super::AtomicRingBuffer<usize> = super::AtomicRingBuffer::with_capacity(cap);
for i in 0..cap {
buf.try_push(i).expect("init");
}
let arc = ::std::sync::Arc::new(buf);
let mut handles = Vec::new();
let end = ::std::time::Instant::now() + ::std::time::Duration::from_millis(10000);
for _thread_num in 0..100 {
let buf = ::std::sync::Arc::clone(&arc);
handles.push(::std::thread::spawn(move || {
while ::std::time::Instant::now() < end {
let a = pop_wait(&buf);
let b = pop_wait(&buf);
while let Err(_) = buf.try_push(a) {};
while let Err(_) = buf.try_push(b) {};
}
}));
}
for (_idx, handle) in handles.into_iter().enumerate() {
handle.join().expect("join");
}
assert_eq!(arc.size(), cap);
let mut expected: Vec<usize> = Vec::new();
let mut actual: Vec<usize> = Vec::new();
for i in 0..cap {
expected.push(i);
actual.push(arc.try_pop().expect("check"));
}
actual.sort_by(|&a, b| a.partial_cmp(b).unwrap());
assert_eq!(actual, expected);
}
fn pop_wait(buf: &::std::sync::Arc<super::AtomicRingBuffer<usize>>) -> usize {
loop {
match buf.try_pop() {
None => continue,
Some(v) => return v,
}
}
}
}