1use std::sync::atomic::{AtomicUsize, Ordering};
2
3use dashmap::DashMap;
4
5pub 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 pub fn new() -> Self {
28 Self::default()
29 }
30
31 pub fn len(&self) -> usize {
33 self.tail.load(Ordering::Relaxed) - self.head.load(Ordering::Relaxed)
34 }
35
36 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 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}