use arc_swap::ArcSwapOption;
use std::{
fmt::{Debug, Display, Write},
sync::Arc,
};
type CellInner<T> = Option<Arc<TreiberCell<T>>>;
pub struct TreiberStack<T: Send + Sync> {
head: ArcSwapOption<TreiberCell<T>>,
}
struct TreiberCell<T: Send + Sync> {
value: Arc<T>,
next: CellInner<T>,
}
impl<T: Send + Sync> Drop for TreiberCell<T> {
fn drop(&mut self) {
let mut next = self.next.take();
while let Some(n) = next.take() {
if let Some(mut n) = Arc::into_inner(n) {
next = n.next.take();
} else {
break;
}
}
}
}
impl<T: Send + Sync> Default for TreiberStack<T> {
fn default() -> Self {
Self {
head: ArcSwapOption::empty(),
}
}
}
impl<T: Send + Sync> Into<Vec<Arc<T>>> for TreiberStack<T> {
fn into(self) -> Vec<Arc<T>> {
self.drain()
}
}
impl<T: Send + Sync, I: IntoIterator<Item = J>, J: Into<T>> From<I> for TreiberStack<T> {
fn from(value: I) -> Self {
let result = Self::default();
for node in value.into_iter() {
result.push(node)
}
result
}
}
impl<T: Send + Sync> TryInto<Vec<T>> for TreiberStack<T> {
type Error = IntoInnerError<T>;
fn try_into(mut self) -> Result<Vec<T>, Self::Error> {
let mut result = Vec::with_capacity(self.len());
let mut old_head = ArcSwapOption::empty();
std::mem::swap(&mut old_head, &mut self.head);
let mut node = old_head.into_inner();
while let Some(inner) = node.take() {
match Arc::try_unwrap(inner) {
Ok(nd) => {
let (value, next) = nd.into_parts();
match Arc::try_unwrap(value) {
Ok(val) => {
result.push(val);
node = next;
}
Err(e) => {
let new_cell = TreiberCell { value: e, next };
let remainder: TreiberStack<T> = TreiberStack {
head: ArcSwapOption::new(Some(Arc::new(new_cell))),
};
return Err(IntoInnerError {
drained_elements: result,
remainder,
});
}
}
}
Err(e) => {
let head = ArcSwapOption::from(Some(e));
let unremoved: TreiberStack<T> = TreiberStack { head };
return Err(IntoInnerError {
drained_elements: result,
remainder: unremoved,
});
}
}
}
Ok(result)
}
}
pub struct IntoInnerError<T: Send + Sync> {
pub drained_elements: Vec<T>,
pub remainder: TreiberStack<T>,
}
impl<T: Send + Sync> Display for IntoInnerError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_fmt(format_args!(
"drained-elements: {}, unremovable-elements: {}",
self.drained_elements.len(),
self.remainder.len()
))
}
}
impl<T: Send + Sync> std::fmt::Debug for IntoInnerError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("IntoInnerError")
.field("drained_elements", &self.drained_elements.len())
.field("remainder", &self.remainder.len())
.finish()
}
}
impl<T: Send + Sync> std::error::Error for IntoInnerError<T> {}
impl<T: Send + Sync> TreiberStack<T> {
pub fn initialized_with(item: T) -> Self {
Self {
head: ArcSwapOption::new(Some(Arc::new(TreiberCell {
value: Arc::new(item),
next: None,
}))),
}
}
pub fn push<I: Into<T>>(&self, val: I) {
let a = Arc::new(val.into());
self.push_arc(a);
}
pub fn push_arc(&self, a: Arc<T>) {
self.head.rcu(|old| {
if let Some(curr_head) = old {
let new_head = prepend(curr_head.clone(), a.clone());
Some(Arc::new(new_head))
} else {
Some(Arc::new(TreiberCell {
value: a.clone(),
next: None,
}))
}
});
}
pub fn pop(&self) -> Option<Arc<T>> {
let popped = self.head.rcu(|old| match old {
Some(curr_head) => curr_head.next.clone(),
None => None,
});
popped.map(|v| v.value.clone())
}
pub fn pop_raw(&self) -> Option<T>
where
T: Copy,
{
self.pop().map(|result| *result)
}
pub fn drain_all_into<F: FnMut(Arc<T>)>(&self, mut f: F) -> usize {
let mut head = self.head.swap(None);
let mut processed = 0_usize;
while let Some(curr) = head {
f(curr.value.clone());
processed += 1;
head = curr.next.clone()
}
processed
}
pub fn drain_all_copy<F: FnMut(T)>(&self, mut f: F) -> usize
where
T: Copy,
{
let mut head = self.head.swap(None);
let mut processed = 0_usize;
while let Some(curr) = head {
let val = *curr.value;
f(val);
processed += 1;
head = curr.next.clone()
}
processed
}
pub fn drain_into<F: FnMut(Arc<T>) -> bool>(&self, mut f: F) -> usize {
let mut processed = 0_usize;
while let Some(item) = self.pop() {
processed += 1;
if !f(item) {
break;
}
}
processed
}
#[inline]
pub fn is_empty(&self) -> bool {
self.head.load().is_none()
}
pub fn len(&self) -> usize {
if let Some(head) = self.head.load().as_ref() {
head.len()
} else {
0
}
}
pub fn clear(&self) {
self.head.store(None);
}
pub fn contains<F: FnMut(&T) -> bool>(&self, mut predicate: F) -> bool {
if let Some(head) = self.head.load().as_ref() {
predicate(&head.value) || {
let mut maybe_next = &head.next;
while let Some(next) = maybe_next {
if predicate(&next.value) {
return true;
}
maybe_next = &next.next
}
false
}
} else {
false
}
}
pub fn drain(&self) -> Vec<Arc<T>> {
if let Some(head) = self.head.swap(None) {
let mut result = Vec::new();
result.push(head.value.clone());
let mut maybe_next = &head.next;
while let Some(next) = maybe_next {
result.push(next.value.clone());
maybe_next = &next.next
}
result
} else {
vec![]
}
}
pub fn drain_replace(&self, new_head: T) -> Vec<Arc<T>> {
if let Some(head) = self.head.swap(Some(Arc::new(TreiberCell {
value: Arc::new(new_head),
next: None,
}))) {
let mut result = Vec::new();
result.push(head.value.clone());
let mut maybe_next = &head.next;
while let Some(next) = maybe_next {
result.push(next.value.clone());
maybe_next = &next.next
}
result
} else {
vec![]
}
}
pub fn drain_transforming<R, F: FnMut(Arc<T>) -> R>(&self, mut transform: F) -> Vec<R> {
if let Some(head) = self.head.swap(None) {
let mut result = Vec::new();
result.push(transform(head.value.clone()));
let mut nxt = &head.next;
while let Some(next) = nxt {
result.push(transform(next.value.clone()));
nxt = &next.next
}
result
} else {
vec![]
}
}
pub fn snapshot(&self) -> Vec<Arc<T>> {
if let Some(head) = self.head.load().as_ref() {
let mut result = Vec::new();
head.copy_into(&mut result);
result
} else {
vec![]
}
}
pub fn peek(&self) -> Option<Arc<T>> {
self.head.load().as_ref().map(|head| head.value.clone())
}
pub fn iter(&self) -> TreiberStackIterator<T> {
TreiberStackIterator {
curr: self.head.load().clone(),
}
}
pub fn exchange_contents(&self, other: &Self) {
let _ = self.head.rcu(|my_head| other.head.rcu(|_| my_head.clone()));
}
}
impl<T: Send + Sync> TreiberCell<T> {
fn len(&self) -> usize {
let mut result = 1;
let mut nxt = &self.next;
while let Some(next) = nxt {
result += 1;
nxt = &next.next;
}
result
}
fn copy_into(&self, into: &mut Vec<Arc<T>>) {
into.push(self.value.clone());
let mut nxt = &self.next;
while let Some(next) = nxt {
into.push(next.value.clone());
nxt = &next.next;
}
}
fn into_parts(mut self) -> (Arc<T>, Option<Arc<TreiberCell<T>>>) {
let next = self.next.take();
(self.value.clone(), next)
}
}
pub struct TreiberStackIterator<T: Send + Sync> {
curr: Option<Arc<TreiberCell<T>>>,
}
impl<T: Send + Sync> Iterator for TreiberStackIterator<T> {
type Item = Arc<T>;
fn next(&mut self) -> Option<Self::Item> {
if let Some(curr) = self.curr.take() {
self.curr = curr.next.to_owned();
Some(curr.value.clone())
} else {
None
}
}
}
fn prepend<T: Send + Sync>(old: Arc<TreiberCell<T>>, t: Arc<T>) -> TreiberCell<T> {
let op: CellInner<T> = Some(old);
TreiberCell { value: t, next: op }
}
impl<T: Send + Sync> TreiberCell<T>
where
T: Display,
{
fn stringify(&self, into: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
into.write_fmt(format_args!("{}", self.value))?;
let mut nxt = &self.next;
while let Some(next) = nxt {
into.write_char(',')?;
into.write_fmt(format_args!("{}", next.value))?;
nxt = &next.next;
}
Ok(())
}
}
impl<T: Send + Sync> TreiberCell<T>
where
T: Debug,
{
fn debugify(&self, into: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
into.write_fmt(format_args!("{:?}", self.value))?;
let mut nxt = &self.next;
while let Some(next) = nxt {
into.write_char(',')?;
into.write_fmt(format_args!("{:?}", next.value))?;
nxt = &next.next;
}
Ok(())
}
}
impl<T: Send + Sync> Display for TreiberCell<T>
where
T: Display,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.stringify(f)
}
}
impl<T: Send + Sync + Debug> Debug for TreiberCell<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.debugify(f)
}
}
impl<T: Send + Sync> Display for TreiberStack<T>
where
T: Display,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_char('(')?;
if let Some(head) = self.head.load().as_ref() {
head.stringify(f)?;
}
f.write_char(')')
}
}
impl<T: Send + Sync> Debug for TreiberStack<T>
where
T: Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_char('(')?;
if let Some(head) = self.head.load().as_ref() {
head.debugify(f)?;
}
f.write_char(')')
}
}
#[cfg(test)]
#[allow(unused_imports, dead_code, clippy::vec_init_then_push)]
mod treiber_stack_tests {
use std::{
fmt::Display,
ops::Range,
sync::{Arc, atomic::AtomicUsize},
thread::{self, JoinHandle},
};
use crate::IntoInnerError;
use super::TreiberStack;
#[test]
fn test_simple() {
let ts: TreiberStack<&str> = TreiberStack::default();
assert!(ts.is_empty());
assert!(ts.peek().is_none());
ts.push("one");
assert!(!ts.is_empty());
assert_eq!(1_usize, ts.len());
assert!(ts.peek().is_some());
assert_eq!(Some(Arc::new("one")), ts.peek());
assert!(
ts.contains(|o| "one".eq(*o)),
"Not present or unequal: 'one'"
);
assert_eq!(1_usize, ts.len());
ts.push("two");
assert!(!ts.is_empty());
assert_eq!(2_usize, ts.len());
ts.push("three");
assert!(!ts.is_empty());
assert_eq!(3_usize, ts.len());
ts.push("four");
assert!(!ts.is_empty());
assert_eq!(4_usize, ts.len());
assert!(ts.peek().is_some());
assert_eq!(Some(Arc::new("four")), ts.peek());
let text = ts.to_string();
assert_eq!("(four,three,two,one)", &text);
let dbg_text = format!("{:?}", ts);
assert_eq!("(\"four\",\"three\",\"two\",\"one\")", &dbg_text);
let a = ts.pop();
assert!(a.is_some());
assert_eq!(&"four", a.as_ref().unwrap().as_ref());
let a = ts.pop();
assert!(a.is_some());
assert_eq!(&"three", a.as_ref().unwrap().as_ref());
let b = ts.pop();
assert_eq!(&"two", b.as_ref().unwrap().as_ref());
let c = ts.pop();
assert_eq!(&"one", c.as_ref().unwrap().as_ref());
assert_eq!(None, ts.pop());
ts.clear();
ts.push("five");
assert_eq!(1, ts.len());
assert!(!ts.is_empty());
ts.clear();
assert_eq!(0, ts.len());
assert!(ts.is_empty());
}
#[test]
fn test_from_and_into() {
let v: Vec<usize> = vec![6, 5, 4, 3, 2, 1];
let stack: TreiberStack<usize> = TreiberStack::from(v);
assert!(!stack.is_empty());
assert_eq!(6, stack.len());
let v: Vec<Arc<usize>> = stack.into();
assert_eq!(6, v.len());
println!("{:?}", v);
let mut vv: Vec<usize> = Vec::with_capacity(v.len());
for item in v.into_iter() {
vv.push(*item);
}
assert_eq!(vec![1, 2, 3, 4, 5, 6], vv);
}
#[test]
fn test_pop_fn() {
let stack: TreiberStack<usize> = TreiberStack::from(vec![6_usize, 5, 4, 3, 2, 1]);
let mut v = Vec::with_capacity(6);
stack.drain_into(|item| {
v.push(*item);
true
});
assert_eq!(6, v.len());
assert_eq!(vec![1, 2, 3, 4, 5, 6], v);
}
#[test]
fn test_pop_fn_filter() {
let stack: TreiberStack<usize> = TreiberStack::from(vec![6_usize, 5, 4, 3, 2, 1]);
let mut v = Vec::with_capacity(6);
stack.drain_into(|item| {
v.push(*item);
*item < 3
});
assert_eq!(vec![1, 2, 3], v);
assert_eq!(3, stack.len());
assert_eq!(vec![4_usize, 5, 6], stack.drain_transforming(|item| *item));
}
#[test]
fn test_threaded() {
const THREADS: usize = 8;
const MAX: usize = 1000;
let ts: TreiberStack<Thing> = TreiberStack::default();
let counter = AtomicUsize::new(0);
let thread_id = AtomicUsize::new(0);
thread::scope(|scope| {
for _ in 0..THREADS {
scope.spawn(|| {
let id = thread_id.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let mut count: usize = 0;
loop {
let next = counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if next > MAX {
break;
}
thread::yield_now();
ts.push(Thing { value: next });
count += 1;
}
assert!(count > 0, "No items added by thread {}", id);
});
}
});
let mut from_iter = Vec::with_capacity(ts.len());
for item in ts.iter() {
from_iter.push(item.value);
}
let copy = ts.snapshot();
let mut from_copy = Vec::with_capacity(copy.len());
for t in copy {
from_copy.push(t.value);
}
from_copy.sort();
let mut expected = Vec::with_capacity(MAX + 1);
for i in 0_usize..(MAX + 1) {
expected.push(i);
}
let mut got = ts.drain_transforming(|t| t.value);
got.sort();
from_iter.sort();
assert_eq!(expected, got, "Contents do not match");
assert_eq!(expected, from_copy, "Contents from copy do not match");
assert_eq!(expected, from_iter, "Contents from iterator do not match");
assert!(ts.is_empty(), "Should be empty");
}
#[derive(Debug)]
struct Thing {
value: usize,
}
impl Display for Thing {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.value.to_string().as_str())
}
}
#[test]
fn test_batches() {
unsafe {
backtrace_on_stack_overflow::enable();
}
const THREADS: usize = 12;
const ITEMS: usize = 15000;
let ts: Arc<TreiberStack<usize>> = Default::default();
fn run_range(rng: Range<usize>, stack: Arc<TreiberStack<usize>>) -> JoinHandle<()> {
std::thread::spawn(move || {
for i in rng {
stack.push(i);
}
})
}
fn ranges() -> Vec<Range<usize>> {
let mut result = Vec::with_capacity(THREADS);
for i in 0..THREADS {
let start = i * ITEMS;
let end = start + ITEMS;
result.push(start..end);
}
result
}
let mut handles = Vec::new();
for range in ranges() {
handles.push(run_range(range, ts.clone()));
}
for h in handles {
h.join().unwrap();
}
let mut all = ts.drain();
all.sort();
let mut prev: Option<usize> = None;
for item in all.iter() {
let item = **item;
if let Some(p) = prev {
if p != item - 1 {
println!("Discontinuity: {} - {}", p, item);
}
}
prev = Some(item);
}
assert_eq!(THREADS * ITEMS, all.len(), "Size mismatch");
}
#[test]
fn test_into_item_vec_clean() {
#[derive(Debug, Eq, PartialEq, Copy, Clone)]
struct Thing {
index: usize,
}
let orig: [Thing; 5] = std::array::from_fn(|index| Thing { index });
let stack: TreiberStack<Thing> = TreiberStack::from(orig);
for (ix, v) in stack.iter().enumerate() {
let inv_index = 4 - ix;
assert_eq!(inv_index, v.index, "Misordered at {} / {}", ix, inv_index);
}
let mut v: Vec<Thing> = stack.try_into().expect("Should not be any Arc clones here");
v.reverse();
assert_eq!(orig.to_vec(), v);
}
#[test]
fn test_into_item_vec_dirty() {
#[derive(Debug, Eq, PartialEq, Copy, Clone)]
struct Thing {
index: usize,
}
let orig: [Thing; 10] = std::array::from_fn(|index| Thing { index });
let stack: TreiberStack<Thing> = TreiberStack::from(orig);
let mut hold: Option<Arc<Thing>> = None;
for (ix, v) in stack.iter().enumerate() {
let inv_index = 9 - ix;
assert_eq!(inv_index, v.index, "Misordered at {} / {}", ix, inv_index);
if ix == 5 {
hold = Some(v);
}
}
let res: Result<Vec<Thing>, IntoInnerError<Thing>> = stack.try_into();
assert!(hold.take().is_some());
let err = match res {
Ok(stuff) => {
panic!(
"Conversion with an outstanding Arc clone should not have succeeded, but got {:?}",
stuff
)
}
Err(e) => e,
};
assert_eq!(5, err.drained_elements.len());
assert_eq!(5, err.remainder.len());
let mut last_index = 1000;
for e in err.drained_elements {
assert!(e.index >= 5);
assert!(e.index <= 9);
assert!(last_index > e.index, "Drain order should be LIFO");
last_index = e.index;
}
last_index = 1000;
for e in err.remainder.iter() {
assert!(e.index <= 4);
assert!(last_index > e.index, "Drain order should be LIFO");
last_index = e.index;
}
let mut rem_vec: Vec<Thing> = err
.remainder
.try_into()
.expect("Should now be able to drain the rest");
rem_vec.reverse();
let portion = (&orig[0..5]).to_vec();
assert_eq!(portion, rem_vec);
}
}