Skip to main content

subetha_cxc/
blocking_mpsc_ring.rs

1//! `BlockingMpscRing`: composed-SPSC MPSC fan-in with cross-process
2//! futex-shaped `send_blocking` / `recv_blocking`.
3//!
4//! Wraps [`crate::mpsc_ring::SharedRingMpsc`] (N independent Lamport
5//! SPSC rings, one per producer) with one [`CrossProcessWaker`] per
6//! ring on the producer side plus one shared consumer waker.
7//!
8//! Wake routing:
9//! - Each producer parks on its own ring's `producer_waker[i]` when
10//!   the ring is full. The consumer wakes that specific waker after
11//!   popping from ring `i` so only the producer who was actually
12//!   blocked on ring `i` runs.
13//! - The consumer parks on a single shared `consumer_waker` when
14//!   every ring is empty. Any producer that pushes wakes that
15//!   single waker by advancing a shared `total_published` counter.
16//!
17//! See [`crate::cross_process_waker`] for the wake protocol +
18//! storage layout. See [`crate::blocking_spsc_ring::BlockingSpscRing`]
19//! for the simpler 1P/1C shape.
20
21use std::cell::Cell;
22use std::marker::PhantomData;
23use std::path::{Path, PathBuf};
24use std::sync::Arc;
25use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
26use std::time::{Duration, Instant};
27
28use crate::blocking_spsc_ring::BlockingError;
29use crate::cross_process_waker::{
30    CrossProcessWaker, MAX_WAITERS_DEFAULT, WakerError,
31};
32use crate::shared_ring::RingError;
33use crate::spsc_ring::SpscRingCore;
34
35const PRE_PARK_SPIN: u32 = 32;
36
37/// Factory for an MPSC pool of N SPSC rings wired into the
38/// cross-process waker primitive.
39pub struct BlockingMpscRing;
40
41/// A single producer handle. Sole writer to one underlying SPSC
42/// ring. `!Sync + !Clone + Send`: one producer per thread.
43pub struct BlockingMpscProducer {
44    ring: Arc<SpscRingCore>,
45    own_waker: Arc<CrossProcessWaker>,
46    consumer_waker: Arc<CrossProcessWaker>,
47    total_published: Arc<AtomicU64>,
48    _not_sync: PhantomData<Cell<()>>,
49}
50
51/// The single consumer handle. Drains all N producer rings
52/// round-robin. `!Sync + !Clone + Send`.
53pub struct BlockingMpscConsumer {
54    rings: Vec<Arc<SpscRingCore>>,
55    producer_wakers: Vec<Arc<CrossProcessWaker>>,
56    consumer_waker: Arc<CrossProcessWaker>,
57    total_published: Arc<AtomicU64>,
58    next_drain: AtomicUsize,
59    _not_sync: PhantomData<Cell<()>>,
60}
61
62impl BlockingMpscRing {
63    /// In-process pool: N producer rings + per-ring producer wakers
64    /// + one consumer waker, all anon-mapped.
65    pub fn create_anon_pool(
66        n_producers: usize,
67        capacity: usize,
68    ) -> Result<(Vec<BlockingMpscProducer>, BlockingMpscConsumer), BlockingError> {
69        assert!(n_producers >= 1, "n_producers must be >= 1");
70        let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
71        let mut producer_wakers: Vec<Arc<CrossProcessWaker>> =
72            Vec::with_capacity(n_producers);
73        for _ in 0..n_producers {
74            rings.push(Arc::new(
75                SpscRingCore::create_anon(capacity).map_err(BlockingError::from)?,
76            ));
77            producer_wakers.push(Arc::new(
78                CrossProcessWaker::create_anon(MAX_WAITERS_DEFAULT)
79                    .map_err(BlockingError::from)?,
80            ));
81        }
82        let consumer_waker = Arc::new(
83            CrossProcessWaker::create_anon(MAX_WAITERS_DEFAULT)
84                .map_err(BlockingError::from)?,
85        );
86        let total_published = Arc::new(AtomicU64::new(0));
87        Ok(build_pool(
88            rings,
89            producer_wakers,
90            consumer_waker,
91            total_published,
92        ))
93    }
94
95    /// File-backed pool. Path layout:
96    ///   `<prefix>.ring.{i}.bin`  - SPSC ring for producer `i`
97    ///   `<prefix>.pw.{i}.bin`    - producer-side waker for ring `i`
98    ///   `<prefix>.cw.bin`        - shared consumer waker
99    pub fn create_pool(
100        path_prefix: impl AsRef<Path>,
101        n_producers: usize,
102        capacity: usize,
103    ) -> Result<(Vec<BlockingMpscProducer>, BlockingMpscConsumer), BlockingError> {
104        assert!(n_producers >= 1, "n_producers must be >= 1");
105        let base = path_prefix.as_ref().to_path_buf();
106        let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
107        let mut producer_wakers: Vec<Arc<CrossProcessWaker>> =
108            Vec::with_capacity(n_producers);
109        for i in 0..n_producers {
110            rings.push(Arc::new(
111                SpscRingCore::create(ring_path(&base, i), capacity)
112                    .map_err(BlockingError::from)?,
113            ));
114            producer_wakers.push(Arc::new(
115                CrossProcessWaker::create(pw_path(&base, i), MAX_WAITERS_DEFAULT)
116                    .map_err(BlockingError::from)?,
117            ));
118        }
119        let consumer_waker = Arc::new(
120            CrossProcessWaker::create(cw_path(&base), MAX_WAITERS_DEFAULT)
121                .map_err(BlockingError::from)?,
122        );
123        let total_published = Arc::new(AtomicU64::new(0));
124        Ok(build_pool(
125            rings,
126            producer_wakers,
127            consumer_waker,
128            total_published,
129        ))
130    }
131
132    /// Open an existing file-backed pool. Caller passes the same
133    /// `path_prefix`, `n_producers`, and `capacity` the pool was
134    /// created with.
135    pub fn open_pool(
136        path_prefix: impl AsRef<Path>,
137        n_producers: usize,
138        expected_capacity: usize,
139    ) -> Result<(Vec<BlockingMpscProducer>, BlockingMpscConsumer), BlockingError> {
140        assert!(n_producers >= 1, "n_producers must be >= 1");
141        let base = path_prefix.as_ref().to_path_buf();
142        let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
143        let mut producer_wakers: Vec<Arc<CrossProcessWaker>> =
144            Vec::with_capacity(n_producers);
145        for i in 0..n_producers {
146            rings.push(Arc::new(
147                SpscRingCore::open(ring_path(&base, i), expected_capacity)
148                    .map_err(BlockingError::from)?,
149            ));
150            producer_wakers.push(Arc::new(
151                CrossProcessWaker::open(pw_path(&base, i), MAX_WAITERS_DEFAULT)
152                    .map_err(BlockingError::from)?,
153            ));
154        }
155        let consumer_waker = Arc::new(
156            CrossProcessWaker::open(cw_path(&base), MAX_WAITERS_DEFAULT)
157                .map_err(BlockingError::from)?,
158        );
159        let total_published = Arc::new(AtomicU64::new(0));
160        Ok(build_pool(
161            rings,
162            producer_wakers,
163            consumer_waker,
164            total_published,
165        ))
166    }
167}
168
169fn ring_path(base: &Path, i: usize) -> PathBuf {
170    let mut s = base.as_os_str().to_owned();
171    s.push(format!(".ring.{i}.bin"));
172    PathBuf::from(s)
173}
174
175fn pw_path(base: &Path, i: usize) -> PathBuf {
176    let mut s = base.as_os_str().to_owned();
177    s.push(format!(".pw.{i}.bin"));
178    PathBuf::from(s)
179}
180
181fn cw_path(base: &Path) -> PathBuf {
182    let mut s = base.as_os_str().to_owned();
183    s.push(".cw.bin");
184    PathBuf::from(s)
185}
186
187fn build_pool(
188    rings: Vec<Arc<SpscRingCore>>,
189    producer_wakers: Vec<Arc<CrossProcessWaker>>,
190    consumer_waker: Arc<CrossProcessWaker>,
191    total_published: Arc<AtomicU64>,
192) -> (Vec<BlockingMpscProducer>, BlockingMpscConsumer) {
193    let producers: Vec<BlockingMpscProducer> = rings
194        .iter()
195        .zip(producer_wakers.iter())
196        .map(|(r, pw)| BlockingMpscProducer {
197            ring: Arc::clone(r),
198            own_waker: Arc::clone(pw),
199            consumer_waker: Arc::clone(&consumer_waker),
200            total_published: Arc::clone(&total_published),
201            _not_sync: PhantomData,
202        })
203        .collect();
204    let consumer = BlockingMpscConsumer {
205        rings,
206        producer_wakers,
207        consumer_waker,
208        total_published,
209        next_drain: AtomicUsize::new(0),
210        _not_sync: PhantomData,
211    };
212    (producers, consumer)
213}
214
215impl BlockingMpscProducer {
216    /// Non-blocking push. On success, increments the shared
217    /// `total_published` counter and fires a wake at the consumer
218    /// waker (cheap if no parked consumer).
219    #[inline]
220    pub fn try_push(&self, payload: &[u8]) -> Result<(), RingError> {
221        let r = self.ring.try_push(payload);
222        if r.is_ok() {
223            let new_seq = self.total_published.fetch_add(1, Ordering::Release) + 1;
224            self.consumer_waker.wake_up_to(new_seq);
225        }
226        r
227    }
228
229    /// Block until either push succeeds or `timeout` elapses.
230    /// Producer parks on its OWN ring's waker; the consumer's
231    /// pop-side wakes the right ring's waker by `try_pop_ring_blocking`.
232    pub fn send_blocking(
233        &self,
234        payload: &[u8],
235        timeout: Option<Duration>,
236    ) -> Result<(), BlockingError> {
237        let deadline = timeout.map(|d| Instant::now() + d);
238        loop {
239            match self.try_push(payload) {
240                Ok(()) => return Ok(()),
241                Err(RingError::Full) => {}
242                Err(e) => return Err(BlockingError::Ring(e)),
243            }
244            for _ in 0..PRE_PARK_SPIN {
245                if self.ring.try_push(payload).is_ok() {
246                    let new_seq = self.total_published.fetch_add(1, Ordering::Release) + 1;
247                    self.consumer_waker.wake_up_to(new_seq);
248                    return Ok(());
249                }
250                std::hint::spin_loop();
251            }
252            let target = self.ring.tail() + 1;
253            let token = self.own_waker.try_park(target)?;
254            // Wake-before-park recovery.
255            if self.ring.try_push(payload).is_ok() {
256                self.own_waker.release(token);
257                let new_seq = self.total_published.fetch_add(1, Ordering::Release) + 1;
258                self.consumer_waker.wake_up_to(new_seq);
259                return Ok(());
260            }
261            let remaining = match deadline {
262                None => None,
263                Some(d) => {
264                    let now = Instant::now();
265                    if now >= d {
266                        self.own_waker.release(token);
267                        return Err(BlockingError::Timeout);
268                    }
269                    Some(d - now)
270                }
271            };
272            match self.own_waker.wait(token, remaining) {
273                Ok(()) => continue,
274                Err(WakerError::Timeout) => return Err(BlockingError::Timeout),
275                Err(e) => return Err(BlockingError::from(e)),
276            }
277        }
278    }
279
280    /// This producer's ring capacity (constant).
281    pub fn capacity(&self) -> usize { self.ring.capacity() }
282    /// This producer's own publish head.
283    pub fn head(&self) -> u64 { self.ring.head() }
284}
285
286impl BlockingMpscConsumer {
287    /// Non-blocking pop, round-robin across all N rings.
288    pub fn try_pop(&self, out: &mut [u8]) -> Result<usize, RingError> {
289        let n = self.rings.len();
290        let start = self.next_drain.load(Ordering::Relaxed);
291        for i in 0..n {
292            let idx = (start + i) % n;
293            if let Ok(bytes) = self.rings[idx].try_pop(out) {
294                self.next_drain.store((idx + 1) % n, Ordering::Relaxed);
295                let tail = self.rings[idx].tail();
296                self.producer_wakers[idx].wake_up_to(tail);
297                return Ok(bytes);
298            }
299        }
300        Err(RingError::Empty)
301    }
302
303    /// Block until either a pop succeeds or `timeout` elapses.
304    /// Consumer parks on the SHARED consumer waker; any producer's
305    /// push fires it.
306    pub fn recv_blocking(
307        &self,
308        out: &mut [u8],
309        timeout: Option<Duration>,
310    ) -> Result<usize, BlockingError> {
311        let deadline = timeout.map(|d| Instant::now() + d);
312        loop {
313            match self.try_pop(out) {
314                Ok(n) => return Ok(n),
315                Err(RingError::Empty) => {}
316                Err(e) => return Err(BlockingError::Ring(e)),
317            }
318            for _ in 0..PRE_PARK_SPIN {
319                if let Ok(n) = self.try_pop_inner(out) {
320                    return Ok(n);
321                }
322                std::hint::spin_loop();
323            }
324            let target = self.total_published.load(Ordering::Acquire) + 1;
325            let token = self.consumer_waker.try_park(target)?;
326            if let Ok(n) = self.try_pop_inner(out) {
327                self.consumer_waker.release(token);
328                return Ok(n);
329            }
330            let remaining = match deadline {
331                None => None,
332                Some(d) => {
333                    let now = Instant::now();
334                    if now >= d {
335                        self.consumer_waker.release(token);
336                        return Err(BlockingError::Timeout);
337                    }
338                    Some(d - now)
339                }
340            };
341            match self.consumer_waker.wait(token, remaining) {
342                Ok(()) => continue,
343                Err(WakerError::Timeout) => return Err(BlockingError::Timeout),
344                Err(e) => return Err(BlockingError::from(e)),
345            }
346        }
347    }
348
349    #[inline]
350    fn try_pop_inner(&self, out: &mut [u8]) -> Result<usize, RingError> {
351        let n = self.rings.len();
352        let start = self.next_drain.load(Ordering::Relaxed);
353        for i in 0..n {
354            let idx = (start + i) % n;
355            if let Ok(bytes) = self.rings[idx].try_pop(out) {
356                self.next_drain.store((idx + 1) % n, Ordering::Relaxed);
357                let tail = self.rings[idx].tail();
358                self.producer_wakers[idx].wake_up_to(tail);
359                return Ok(bytes);
360            }
361        }
362        Err(RingError::Empty)
363    }
364
365    /// Number of producer rings this consumer drains.
366    pub fn n_producers(&self) -> usize { self.rings.len() }
367
368    /// Approximate total pending items across every ring.
369    pub fn approx_total_len(&self) -> usize {
370        self.rings.iter().map(|r| r.approx_len()).sum()
371    }
372}
373
374#[cfg(test)]
375mod tests {
376    use super::*;
377    use std::thread;
378
379    #[test]
380    fn round_trip_4p_1c_anon() {
381        let (producers, consumer) =
382            BlockingMpscRing::create_anon_pool(4, 8).expect("create");
383        const PER_PROD: u64 = 25;
384        let total: u64 = PER_PROD * 4;
385        let handles: Vec<_> = producers
386            .into_iter()
387            .enumerate()
388            .map(|(pid, p)| {
389                thread::spawn(move || {
390                    for i in 0..PER_PROD {
391                        let val = (pid as u64) * 1_000_000 + i;
392                        let mut payload = [0u8; 56];
393                        payload[..8].copy_from_slice(&val.to_le_bytes());
394                        p.send_blocking(&payload, Some(Duration::from_secs(5)))
395                            .expect("send");
396                    }
397                })
398            })
399            .collect();
400        let mut buf = [0u8; 64];
401        let mut seen: Vec<u64> = Vec::with_capacity(total as usize);
402        for _ in 0..total {
403            consumer
404                .recv_blocking(&mut buf, Some(Duration::from_secs(5)))
405                .expect("recv");
406            seen.push(u64::from_le_bytes(buf[..8].try_into().unwrap()));
407        }
408        for h in handles {
409            h.join().unwrap();
410        }
411        seen.sort_unstable();
412        let mut expected: Vec<u64> = Vec::with_capacity(total as usize);
413        for pid in 0..4u64 {
414            for i in 0..PER_PROD {
415                expected.push(pid * 1_000_000 + i);
416            }
417        }
418        expected.sort_unstable();
419        assert_eq!(seen, expected, "every item delivered exactly once");
420    }
421
422    #[test]
423    fn recv_blocking_returns_timeout() {
424        let (_producers, consumer) =
425            BlockingMpscRing::create_anon_pool(2, 4).expect("create");
426        let mut buf = [0u8; 64];
427        let t0 = Instant::now();
428        let err = consumer.recv_blocking(&mut buf, Some(Duration::from_millis(60)));
429        assert_eq!(err, Err(BlockingError::Timeout));
430        assert!(t0.elapsed() >= Duration::from_millis(50));
431    }
432}