Skip to main content

waitfree_sync/
spsc.rs

1//! A wait-free single-producer single-consumer (SPSC) queue to send data to another thread.
2//! It is based on the improved FastForward queue.
3//!
4//! This is similar to [`std::sync::mpsc`], but restricted to a single producer and a single
5//! consumer and backed by a fixed-size ring buffer instead of an unbounded linked list. Because
6//! of this, [`Sender::try_send`] and [`Receiver::try_recv`] never block and never allocate: they
7//! return immediately instead of parking the calling thread the way `std`'s blocking `send`/`recv`
8//! do.
9//!
10//! # Example
11//! ```rust
12//! use waitfree_sync::spsc;
13//!
14//! //                            Type ──╮   ╭─ Capacity
15//! let (mut tx, mut rx) = spsc::spsc::<u64>(8);
16//! tx.try_send(234);
17//! assert_eq!(rx.try_recv(),Ok(234u64));
18//! ```
19//!
20//! # Behavior for full and empty queue.
21//! If the queue is full, [`Sender::try_send`] returns [`SendError::NoSpaceLeft`].
22//! If the queue is empty, [`Receiver::try_recv`] returns [`TryRecvError::Empty`].
23//!
24use crate::import::{Arc, AtomicBool, Ordering, UnsafeCell};
25use core::error::Error;
26use crossbeam_utils::CachePadded;
27use std::{fmt::Debug, sync::atomic::AtomicUsize};
28
29/// Create a new wait-free SPSC queue. The `capacity` must be a power of two, which is validate during runtime.
30/// # Panic
31/// Panics if the `capacity` is not a power of two.
32/// # Example
33/// ```rust
34/// use waitfree_sync::spsc;
35///
36/// //               Data type ──╮   ╭─ Capacity
37/// let (tx, rx) = spsc::spsc::<u64>(8);
38/// ```
39pub fn spsc<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
40    if !is_power_of_two(capacity) {
41        panic!("The SIZE must be a power of 2")
42    }
43
44    let chan = Arc::new(Spsc::new(capacity));
45
46    let r = Receiver::new(chan.clone());
47    let w = Sender::new(chan);
48
49    (w, r)
50}
51
52const fn is_power_of_two(x: usize) -> bool {
53    let c = x.wrapping_sub(1);
54    (x != 0) && (x != 1) && ((x & c) == 0)
55}
56
57/// An error returned from the [`Sender::try_send`] function on a [`Sender`].
58///
59/// The error contains the data being sent as a payload so it can be recovered.
60#[derive(Clone, Debug, PartialEq)]
61pub enum SendError<T> {
62    /// The queue is full. The receiving side of the queue must collect items.
63    NoSpaceLeft(T),
64    /// The receiving end of a channel is disconnected, implying that the data could never be received.
65    ReceiverSideDropped(T),
66}
67impl<T: Debug> Error for SendError<T> {}
68impl<T> core::fmt::Display for SendError<T> {
69    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
70        match self {
71            SendError::NoSpaceLeft(_) => write!(f, "No space left in the SPSC queue."),
72            SendError::ReceiverSideDropped(_) => {
73                write!(f, "Receiver side of the SPSC queue dropped.")
74            }
75        }
76    }
77}
78impl<T> SendError<T> {
79    /// Returns the value which was tried to be sent to the queue.
80    pub fn into_value(self) -> T {
81        match self {
82            SendError::NoSpaceLeft(val) => val,
83            SendError::ReceiverSideDropped(val) => val,
84        }
85    }
86}
87
88/// This enumeration is the list of the possible reasons that [`Receiver::try_recv`] could not return data when called.
89#[derive(Clone, Debug, PartialEq)]
90pub enum TryRecvError {
91    /// This queue is currently empty, but the Sender have not yet disconnected, so data may yet become available.
92    Empty,
93    /// The queues sending half has become disconnected, and there will never be any more data received on it.
94    Disconnected,
95}
96impl Error for TryRecvError {}
97impl core::fmt::Display for TryRecvError {
98    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
99        match self {
100            TryRecvError::Empty => write!(f, "No data available in the SPSC queue."),
101            TryRecvError::Disconnected => {
102                write!(f, "Sender side of the SPSC queue dropped.")
103            }
104        }
105    }
106}
107
108#[derive(Debug)]
109struct Slot<T> {
110    value: UnsafeCell<Option<T>>,
111    occupied: CachePadded<AtomicBool>,
112}
113impl<T> Slot<T> {
114    fn new() -> Self {
115        Self {
116            value: UnsafeCell::new(None),
117            occupied: CachePadded::new(false.into()),
118        }
119    }
120}
121
122#[derive(Debug)]
123struct Spsc<T> {
124    mem: Box<[Slot<T>]>,
125    // The mask is written when this structure is created and is then only read.
126    // Therefore, we do not need Atomic here.
127    mask: usize,
128    read: CachePadded<AtomicUsize>,
129    write: CachePadded<AtomicUsize>,
130}
131
132impl<T> Spsc<T> {
133    fn new(size: usize) -> Self {
134        let mut buffer = Vec::with_capacity(size);
135        for _ in 0..size {
136            buffer.push(Slot::new());
137        }
138        let buffer: Box<[Slot<T>]> = buffer.into_boxed_slice();
139        Spsc {
140            mem: buffer,
141            mask: size - 1,
142            read: CachePadded::new(0.into()),
143            write: CachePadded::new(0.into()),
144        }
145    }
146
147    #[inline]
148    fn capacity(&self) -> usize {
149        self.mask + 1
150    }
151
152    #[inline]
153    fn len(&self) -> usize {
154        self.write
155            .load(Ordering::Relaxed)
156            .saturating_sub(self.read.load(Ordering::Relaxed))
157    }
158}
159
160/// The receiving side of the [spsc] queue.
161#[derive(Debug)]
162pub struct Receiver<T> {
163    spsc: Arc<Spsc<T>>,
164}
165unsafe impl<T: Send> Send for Receiver<T> {}
166unsafe impl<T: Send> Sync for Receiver<T> {}
167
168impl<T> Receiver<T> {
169    fn new(spsc: Arc<Spsc<T>>) -> Self {
170        Receiver { spsc }
171    }
172}
173
174impl<T> Receiver<T> {
175    /// Retrieve the next available element from the queue without blocking.
176    ///
177    /// Returns [`TryRecvError::Empty`] if the queue is currently empty,
178    /// or [`TryRecvError::Disconnected`] if the [Sender] has been dropped and
179    /// no further items can arrive.
180    pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
181        let read = self.spsc.read.load(Ordering::Relaxed);
182        let rpos = read & self.spsc.mask;
183        let slot = unsafe { self.spsc.mem.get_unchecked(rpos) };
184        if !slot.occupied.load(Ordering::Acquire) {
185            if self.is_disconnected() {
186                Err(TryRecvError::Disconnected)
187            } else {
188                Err(TryRecvError::Empty)
189            }
190        } else {
191            #[cfg(not(loom))]
192            let val = unsafe { slot.value.get().replace(None) };
193            #[cfg(loom)]
194            let val = unsafe { slot.value.get_mut().with(|ptr| ptr.replace(None)) };
195
196            slot.occupied.store(false, Ordering::Release);
197            // self.read = self.read.wrapping_add(1);
198            self.spsc
199                .read
200                .store(read.wrapping_add(1), Ordering::Relaxed);
201            Ok(val.ok_or(TryRecvError::Empty)?)
202        }
203    }
204    /// Peeks the next element in the queue without removing it.
205    #[cfg(not(loom))] // We can't return a reference to an UnsafeCell of loom.
206    pub fn peek(&self) -> Option<&T> {
207        let rpos = self.spsc.read.load(Ordering::Relaxed) & self.spsc.mask;
208        let slot = unsafe { self.spsc.mem.get_unchecked(rpos) };
209        if !slot.occupied.load(Ordering::Acquire) {
210            None
211        } else {
212            let val = unsafe { &*slot.value.get() };
213            val.as_ref()
214        }
215    }
216
217    /// Returns the total number of items that the queue can hold at most.
218    #[inline]
219    pub fn capacity(&self) -> usize {
220        // SAFETY: This is safe because we only read size which is never written.
221        self.spsc.capacity()
222    }
223
224    /// Returns the number of items in the queue.
225    /// # WARNING
226    /// This length is only a best-effort estimate.
227    /// It is computed from relaxed atomic and is NOT a linearizable value.
228    /// It may be temporarily incorrect (including over/under-counting) due to
229    /// reordering and visibility delays across threads.
230    #[inline]
231    pub fn len(&self) -> usize {
232        self.spsc.len()
233    }
234
235    /// Returns true if the queue is empty.
236    /// # WARNING
237    /// This length is only a best-effort estimate.
238    /// It is computed from relaxed atomic and is NOT a linearizable value.
239    /// It may be temporarily incorrect (including over/under-counting) due to
240    /// reordering and visibility delays across threads.
241    #[inline]
242    pub fn is_empty(&self) -> bool {
243        self.spsc.len() == 0
244    }
245
246    /// Returns true if the channel is disconnected because the sender was dropped.
247    ///
248    /// Note that a return value of false does not guarantee the channel will remain connected.
249    /// The channel may be disconnected immediately after this method returns, so a subsequent [Sender::try_send] may still fail with SendError.
250    #[inline]
251    pub fn is_disconnected(&self) -> bool {
252        Arc::strong_count(&self.spsc) < 2
253    }
254}
255
256/// The sending side of the [spsc] queue.
257#[derive(Debug)]
258pub struct Sender<T> {
259    spsc: Arc<Spsc<T>>,
260}
261unsafe impl<T: Send> Send for Sender<T> {}
262unsafe impl<T: Send> Sync for Sender<T> {}
263impl<T> Sender<T> {
264    fn new(spsc: Arc<Spsc<T>>) -> Self {
265        Sender { spsc }
266    }
267}
268
269impl<T> Sender<T> {
270    /// Attempts to send a value to the queue without blocking.
271    ///
272    /// Because this queue has a fixed capacity, sending returns [`SendError::NoSpaceLeft`]
273    /// instead of blocking or growing the buffer when the queue is full.
274    /// Returns [`SendError::ReceiverSideDropped`] if the [Receiver] has been dropped.
275    pub fn try_send(&mut self, data: T) -> Result<(), SendError<T>> {
276        let write = self.spsc.write.load(Ordering::Relaxed);
277        let wpos = write & self.spsc.mask;
278
279        if self.is_disconnected() {
280            return Err(SendError::ReceiverSideDropped(data));
281        }
282
283        let slot = unsafe { self.spsc.mem.get_unchecked(wpos) };
284        if slot.occupied.load(Ordering::Acquire) {
285            Err(SendError::NoSpaceLeft(data))
286        } else {
287            #[cfg(not(loom))]
288            unsafe {
289                slot.value.get().write(Some(data))
290            };
291            #[cfg(loom)]
292            unsafe {
293                slot.value.get_mut().with(|ptr| ptr.write(Some(data)))
294            };
295            slot.occupied.store(true, Ordering::Release);
296            self.spsc
297                .write
298                .store(write.wrapping_add(1), Ordering::Relaxed);
299            Ok(())
300        }
301    }
302
303    /// Returns the total number of items that the queue can hold at most.
304    #[inline]
305    pub fn capacity(&self) -> usize {
306        // SAFETY: This is safe because we only read size which is never written.
307        self.spsc.capacity()
308    }
309
310    /// Returns the number of items in the queue.
311    /// # WARNING
312    /// This length is only a best-effort estimate.
313    /// It is computed from relaxed atomic and is NOT a linearizable value.
314    /// It may be temporarily incorrect (including over/under-counting) due to
315    /// reordering and visibility delays across threads.
316    #[inline]
317    pub fn len(&self) -> usize {
318        self.spsc.len()
319    }
320
321    /// Returns true if the queue is empty.
322    /// # WARNING
323    /// This length is only a best-effort estimate.
324    /// It is computed from relaxed atomic and is NOT a linearizable value.
325    /// It may be temporarily incorrect (including over/under-counting) due to
326    /// reordering and visibility delays across threads.
327    #[inline]
328    pub fn is_empty(&self) -> bool {
329        self.spsc.len() == 0
330    }
331
332    /// Returns true if the channel is disconnected because the receiver was dropped.
333    ///
334    /// Note that a return value of false does not guarantee the channel will remain connected.
335    /// The channel may be disconnected immediately after this method returns, so a subsequent [Sender::try_send] may still fail with SendError.
336    #[inline]
337    pub fn is_disconnected(&self) -> bool {
338        Arc::strong_count(&self.spsc) < 2
339    }
340}
341
342#[cfg(not(loom))]
343#[cfg(test)]
344mod test {
345    #[cfg(loom)]
346    use loom::thread;
347    #[cfg(not(loom))]
348    use std::thread;
349
350    use super::*;
351
352    #[test]
353    fn smoke() {
354        let (mut w, mut r) = spsc(4);
355        w.try_send(vec![0; 15]).unwrap();
356        w.try_send(vec![0; 16]).unwrap();
357        w.try_send(vec![0; 17]).unwrap();
358        w.try_send(vec![0; 18]).unwrap();
359
360        assert_eq!(r.try_recv(), Ok(vec![0; 15]));
361        assert_eq!(r.try_recv(), Ok(vec![0; 16]));
362        assert_eq!(r.try_recv(), Ok(vec![0; 17]));
363        assert_eq!(r.try_recv(), Ok(vec![0; 18]));
364    }
365
366    #[test]
367    fn test_is_power_of_two() {
368        assert!(!is_power_of_two(0));
369        assert!(!is_power_of_two(1));
370        assert!(is_power_of_two(2));
371        assert!(!is_power_of_two(3));
372        assert!(is_power_of_two(4));
373        assert!(!is_power_of_two(5));
374        assert!(!is_power_of_two(6));
375        assert!(!is_power_of_two(7));
376        assert!(is_power_of_two(8));
377        assert!(!is_power_of_two(9));
378
379        assert!(!is_power_of_two(15));
380        assert!(is_power_of_two(16));
381        assert!(!is_power_of_two(17));
382
383        assert!(!is_power_of_two(31));
384        assert!(is_power_of_two(32));
385        assert!(!is_power_of_two(33));
386    }
387
388    #[test]
389    fn test_drop_read_side() {
390        let (mut write, read) = spsc::<i32>(4);
391
392        assert_eq!(write.try_send(1), Ok(()));
393        assert_eq!(write.len(), 1);
394        assert_eq!(write.try_send(2), Ok(()));
395        assert_eq!(write.len(), 2);
396        drop(read);
397        assert_eq!(write.try_send(3), Err(SendError::ReceiverSideDropped(3)));
398        assert_eq!(write.len(), 2);
399        assert_eq!(write.try_send(4), Err(SendError::ReceiverSideDropped(4)));
400        assert_eq!(write.len(), 2);
401        assert_eq!(write.try_send(5), Err(SendError::ReceiverSideDropped(5)));
402        assert_eq!(write.len(), 2);
403    }
404
405    #[test]
406    fn test_drop_write_side() {
407        let (mut write, mut read) = spsc::<i32>(4);
408
409        write.try_send(0).unwrap();
410        write.try_send(1).unwrap();
411        assert_eq!(read.try_recv(), Ok(0));
412        drop(write);
413        assert_eq!(read.try_recv(), Ok(1));
414    }
415
416    #[test]
417    fn test_full_empty() {
418        let (mut write, mut read) = spsc::<i32>(4);
419        assert_eq!(write.try_send(1), Ok(()));
420        assert_eq!(write.len(), 1);
421        assert_eq!(write.try_send(2), Ok(()));
422        assert_eq!(write.len(), 2);
423        assert_eq!(write.try_send(3), Ok(()));
424        assert_eq!(write.len(), 3);
425        assert_eq!(write.try_send(4), Ok(()));
426        assert_eq!(write.len(), 4);
427        assert_eq!(write.try_send(5), Err(SendError::NoSpaceLeft(5)));
428        assert_eq!(write.len(), 4);
429
430        assert_eq!(read.try_recv(), Ok(1));
431        assert_eq!(write.len(), 3);
432        assert_eq!(write.try_send(6), Ok(()));
433        assert_eq!(write.len(), 4);
434        assert_eq!(read.try_recv(), Ok(2));
435        assert_eq!(write.len(), 3);
436        assert_eq!(read.try_recv(), Ok(3));
437        assert_eq!(write.len(), 2);
438        assert_eq!(read.try_recv(), Ok(4));
439        assert_eq!(write.len(), 1);
440        assert_eq!(read.try_recv(), Ok(6));
441        assert_eq!(read.try_recv(), Err(TryRecvError::Empty));
442    }
443
444    #[test]
445    fn test_drop_one_side() {
446        let (mut write, read) = spsc::<i32>(4);
447        assert_eq!(write.try_send(1), Ok(()));
448        assert_eq!(write.len(), 1);
449        assert_eq!(write.try_send(2), Ok(()));
450        assert_eq!(write.len(), 2);
451        drop(read);
452        assert_eq!(write.try_send(3), Err(SendError::ReceiverSideDropped(3)));
453        assert_eq!(write.len(), 2);
454        assert_eq!(write.try_send(4), Err(SendError::ReceiverSideDropped(4)));
455        assert_eq!(write.len(), 2);
456        assert_eq!(write.try_send(5), Err(SendError::ReceiverSideDropped(5)));
457        assert_eq!(write.len(), 2);
458    }
459
460    #[test]
461    fn test_peek() {
462        let (mut w, mut r) = spsc(4);
463        w.try_send(vec![0; 15]).unwrap();
464        w.try_send(vec![0; 16]).unwrap();
465        w.try_send(vec![0; 17]).unwrap();
466        w.try_send(vec![0; 18]).unwrap();
467
468        assert_eq!(r.peek(), Some(&vec![0; 15]));
469        assert_eq!(r.try_recv(), Ok(vec![0; 15]));
470        assert_eq!(r.peek(), Some(&vec![0; 16]));
471        assert_eq!(r.try_recv(), Ok(vec![0; 16]));
472        assert_eq!(r.peek(), Some(&vec![0; 17]));
473        assert_eq!(r.try_recv(), Ok(vec![0; 17]));
474        assert_eq!(r.peek(), Some(&vec![0; 18]));
475        assert_eq!(r.peek(), Some(&vec![0; 18]));
476        assert_eq!(r.peek(), Some(&vec![0; 18]));
477        assert_eq!(r.try_recv(), Ok(vec![0; 18]));
478        assert_eq!(r.peek(), None);
479    }
480
481    #[test]
482    fn test_peek_threaded() {
483        let (mut sender, mut receiver) = spsc(4);
484
485        let writer_thread = thread::spawn(move || {
486            thread::park();
487            for i in 0..4 {
488                assert_eq!(sender.try_send([i; 50]), Ok(()));
489            }
490        });
491        let reader_thread = thread::spawn(move || {
492            thread::park();
493            let mut i = 0;
494            while i < 4 {
495                if let Some(val) = receiver.peek() {
496                    let first_entry = val[0];
497                    for entry in val {
498                        assert_eq!(*entry, first_entry);
499                    }
500                    let val = receiver.try_recv().unwrap();
501                    let first_entry = val[0];
502                    for entry in val {
503                        assert_eq!(entry, first_entry);
504                    }
505                    i += 1;
506                }
507            }
508        });
509        writer_thread.thread().unpark();
510        reader_thread.thread().unpark();
511        assert!(writer_thread.join().is_ok());
512        assert!(reader_thread.join().is_ok());
513    }
514
515    #[test]
516    fn test_dissconnect() {
517        let (tx, rx) = spsc::<u32>(4);
518        assert!(!tx.is_disconnected());
519        assert!(!rx.is_disconnected());
520        drop(tx);
521        assert!(rx.is_disconnected());
522
523        let (tx, rx) = spsc::<u32>(4);
524        assert!(!tx.is_disconnected());
525        assert!(!rx.is_disconnected());
526        drop(rx);
527        assert!(tx.is_disconnected());
528    }
529}