Skip to main content

hala_lockfree/
queue.rs

1use std::sync::atomic::{AtomicUsize, Ordering};
2
3use dashmap::DashMap;
4
5/// A lokfree `FIFO` queue implementation.
6pub struct Queue<T> {
7    head: AtomicUsize,
8    tail: AtomicUsize,
9    slots: DashMap<usize, T>,
10}
11
12unsafe impl<T: Send> Send for Queue<T> {}
13unsafe impl<T: Send> Sync for Queue<T> {}
14
15impl<T> Default for Queue<T> {
16    fn default() -> Self {
17        Self {
18            head: Default::default(),
19            tail: Default::default(),
20            slots: DashMap::new(),
21        }
22    }
23}
24
25impl<T> Queue<T> {
26    /// Create new instance of type [`Queue`]
27    pub fn new() -> Self {
28        Self::default()
29    }
30
31    /// Get queue len.
32    pub fn len(&self) -> usize {
33        self.tail.load(Ordering::Relaxed) - self.head.load(Ordering::Relaxed)
34    }
35
36    /// Push one value into queue tail
37    pub fn push(&self, value: T) {
38        let offset = self.tail.fetch_add(1, Ordering::AcqRel);
39
40        self.slots.insert(offset, value);
41    }
42
43    /// Pop one value from the queue's header. returns [`None`] if this queue is empty.
44    pub fn pop(&self) -> Option<T> {
45        loop {
46            let header = self.head.load(Ordering::Acquire);
47            let tail = self.tail.load(Ordering::Acquire);
48
49            if header == tail {
50                return None;
51            }
52
53            if self
54                .head
55                .compare_exchange(header, header + 1, Ordering::AcqRel, Ordering::Relaxed)
56                .is_ok()
57            {
58                return Some(self.spin_pop_by_offset(header));
59            }
60        }
61    }
62
63    #[cold]
64    fn spin_pop_by_offset(&self, index: usize) -> T {
65        loop {
66            if let Some((_, value)) = self.slots.remove(&index) {
67                return value;
68            }
69        }
70    }
71}
72
73#[cfg(test)]
74mod tests {
75    use std::{
76        collections::HashSet,
77        sync::{atomic::AtomicI32, Arc},
78    };
79
80    use super::*;
81
82    #[test]
83    fn test_mulit_push() {
84        let queue = Arc::new(Queue::<i32>::new());
85        let counter = Arc::new(AtomicI32::new(0));
86
87        let push_threads = 10;
88        let loops = 10000;
89
90        for _ in 0..push_threads {
91            let queue = queue.clone();
92            let counter = counter.clone();
93            std::thread::spawn(move || {
94                for _ in 0..loops {
95                    queue.push(counter.fetch_add(1, Ordering::AcqRel));
96                }
97            });
98        }
99
100        let mut set = HashSet::new();
101
102        for _ in 0..push_threads * loops {
103            'inner: loop {
104                if let Some(value) = queue.pop() {
105                    set.insert(value);
106                    break 'inner;
107                }
108            }
109        }
110
111        assert_eq!(set.len() as i32, push_threads * loops);
112
113        for i in 0..push_threads * loops {
114            assert!(set.contains(&i), "Not found {i}");
115        }
116    }
117
118    #[test]
119    fn test_mulit_push_pop() {
120        let queue = Arc::new(Queue::<i32>::new());
121        let counter = Arc::new(AtomicI32::new(0));
122        let pop_counter = Arc::new(AtomicI32::new(0));
123
124        let push_threads = 5;
125        let pop_threads = 5;
126        let loops = 100;
127
128        for _ in 0..push_threads {
129            let queue = queue.clone();
130            let counter = counter.clone();
131            std::thread::spawn(move || {
132                for _ in 0..loops {
133                    queue.push(counter.fetch_add(1, Ordering::AcqRel));
134                }
135            });
136        }
137
138        let mut handles = vec![];
139
140        for _ in 0..pop_threads {
141            let queue = queue.clone();
142            let pop_counter = pop_counter.clone();
143            handles.push(std::thread::spawn(move || {
144                while queue.head.load(Ordering::Acquire) < push_threads * loops {
145                    if let Some(_) = queue.pop() {
146                        pop_counter.fetch_add(1, Ordering::AcqRel);
147                    }
148                }
149            }));
150        }
151
152        for handle in handles {
153            handle.join().unwrap();
154        }
155
156        assert_eq!(
157            pop_counter.load(Ordering::Relaxed),
158            counter.load(Ordering::Relaxed)
159        );
160    }
161}