1use std::cell::Cell;
51use std::marker::PhantomData;
52use std::path::Path;
53use std::sync::Arc;
54use std::sync::atomic::{AtomicUsize, Ordering};
55
56use crate::shared_ring::RingError;
57use crate::spsc_ring::SpscRingCore;
58
59pub struct SharedRingMpmc;
62
63pub struct MpmcProducer {
65 inner: Arc<SpscRingCore>,
66 _not_sync: PhantomData<Cell<()>>,
67}
68
69pub struct MpmcConsumer {
72 rings: Vec<Arc<SpscRingCore>>,
73 next_drain: AtomicUsize,
74 _not_sync: PhantomData<Cell<()>>,
75}
76
77impl SharedRingMpmc {
78 pub fn create_anon_grid(
87 n_producers: usize,
88 n_consumers: usize,
89 capacity: usize,
90 ) -> Result<(Vec<MpmcProducer>, Vec<MpmcConsumer>), RingError> {
91 assert!(n_consumers >= 1, "n_consumers must be >= 1");
92 assert!(
93 n_producers >= n_consumers,
94 "n_producers ({n_producers}) must be >= n_consumers ({n_consumers}); \
95 every consumer needs at least one ring to drain",
96 );
97
98 let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
99 for _ in 0..n_producers {
100 rings.push(Arc::new(SpscRingCore::create_anon(capacity)?));
101 }
102 build_grid(rings, n_consumers)
103 }
104
105 pub fn create_grid_in_region<R: crate::spsc_ring::RegionOwner>(
117 mut region: R,
118 n_producers: usize,
119 n_consumers: usize,
120 capacity: usize,
121 ) -> Result<(Vec<MpmcProducer>, Vec<MpmcConsumer>), RingError> {
122 assert!(n_consumers >= 1, "n_consumers must be >= 1");
123 assert!(n_producers >= n_consumers,
124 "n_producers must be >= n_consumers");
125 let lane_bytes = crate::spsc_ring::spsc_ring_file_size(capacity);
126 let need = lane_bytes
127 .checked_mul(n_producers)
128 .ok_or(RingError::LayoutMismatch)?;
129 if region.region_len() < need {
130 return Err(RingError::LayoutMismatch);
131 }
132 let base = region.region_ptr();
137 let whole: Arc<dyn std::any::Any + Send + Sync> = Arc::new(region);
138
139 let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
140 for i in 0..n_producers {
141 let lane = SubRegion {
142 _whole: Arc::clone(&whole),
143 ptr: unsafe { base.add(i * lane_bytes) },
144 len: lane_bytes,
145 };
146 rings.push(Arc::new(SpscRingCore::create_in_region(lane, capacity)?));
147 }
148 build_grid(rings, n_consumers)
149 }
150
151 pub fn create_grid(
154 path_prefix: impl AsRef<Path>,
155 n_producers: usize,
156 n_consumers: usize,
157 capacity: usize,
158 ) -> Result<(Vec<MpmcProducer>, Vec<MpmcConsumer>), RingError> {
159 assert!(n_consumers >= 1, "n_consumers must be >= 1");
160 assert!(n_producers >= n_consumers,
161 "n_producers must be >= n_consumers");
162 let base = path_prefix.as_ref().to_path_buf();
163 let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
164 for i in 0..n_producers {
165 let path = ring_path(&base, i);
166 rings.push(Arc::new(SpscRingCore::create(&path, capacity)?));
167 }
168 build_grid(rings, n_consumers)
169 }
170
171 pub fn open_grid(
173 path_prefix: impl AsRef<Path>,
174 n_producers: usize,
175 n_consumers: usize,
176 expected_capacity: usize,
177 ) -> Result<(Vec<MpmcProducer>, Vec<MpmcConsumer>), RingError> {
178 assert!(n_consumers >= 1, "n_consumers must be >= 1");
179 assert!(n_producers >= n_consumers,
180 "n_producers must be >= n_consumers");
181 let base = path_prefix.as_ref().to_path_buf();
182 let mut rings: Vec<Arc<SpscRingCore>> = Vec::with_capacity(n_producers);
183 for i in 0..n_producers {
184 let path = ring_path(&base, i);
185 rings.push(Arc::new(SpscRingCore::open(&path, expected_capacity)?));
186 }
187 build_grid(rings, n_consumers)
188 }
189}
190
191fn build_grid(
192 rings: Vec<Arc<SpscRingCore>>,
193 n_consumers: usize,
194) -> Result<(Vec<MpmcProducer>, Vec<MpmcConsumer>), RingError> {
195 let producers: Vec<MpmcProducer> = rings
196 .iter()
197 .map(|r| MpmcProducer {
198 inner: Arc::clone(r),
199 _not_sync: PhantomData,
200 })
201 .collect();
202
203 let mut consumer_rings: Vec<Vec<Arc<SpscRingCore>>> =
205 (0..n_consumers).map(|_| Vec::new()).collect();
206 for (producer_idx, ring) in rings.iter().enumerate() {
207 consumer_rings[producer_idx % n_consumers].push(Arc::clone(ring));
208 }
209
210 let consumers: Vec<MpmcConsumer> = consumer_rings
211 .into_iter()
212 .map(|subset| MpmcConsumer {
213 rings: subset,
214 next_drain: AtomicUsize::new(0),
215 _not_sync: PhantomData,
216 })
217 .collect();
218
219 Ok((producers, consumers))
220}
221
222fn ring_path(prefix: &std::path::Path, i: usize) -> std::path::PathBuf {
223 let mut s = prefix.as_os_str().to_owned();
224 s.push(format!(".{i}.bin"));
225 std::path::PathBuf::from(s)
226}
227
228struct SubRegion {
233 _whole: Arc<dyn std::any::Any + Send + Sync>,
234 ptr: *mut u8,
235 len: usize,
236}
237
238unsafe impl Send for SubRegion {}
243unsafe impl Sync for SubRegion {}
244
245impl crate::spsc_ring::RegionOwner for SubRegion {
246 fn region_ptr(&mut self) -> *mut u8 { self.ptr }
247 fn region_len(&self) -> usize { self.len }
248}
249
250impl MpmcProducer {
251 pub fn try_push(&self, payload: &[u8]) -> Result<(), RingError> {
253 self.inner.try_push(payload)
254 }
255
256 pub fn capacity(&self) -> usize {
258 self.inner.capacity()
259 }
260
261 pub fn head(&self) -> u64 {
263 self.inner.head()
264 }
265}
266
267impl MpmcConsumer {
268 pub fn try_pop(&self, out: &mut [u8]) -> Result<usize, RingError> {
273 let n = self.rings.len();
274 let start = self.next_drain.load(Ordering::Relaxed);
275 for i in 0..n {
276 let idx = (start + i) % n;
277 if let Ok(bytes) = self.rings[idx].try_pop(out) {
278 self.next_drain.store((idx + 1) % n, Ordering::Relaxed);
279 return Ok(bytes);
280 }
281 }
282 Err(RingError::Empty)
283 }
284
285 pub fn n_rings(&self) -> usize {
287 self.rings.len()
288 }
289
290 pub fn approx_subset_len(&self) -> usize {
292 self.rings.iter().map(|r| r.approx_len()).sum()
293 }
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299 use crate::spsc_ring::SPSC_PAYLOAD_BYTES;
300 use std::thread;
301
302 #[test]
303 fn create_anon_grid_round_trip() {
304 let (producers, consumers) =
306 SharedRingMpmc::create_anon_grid(4, 2, 8).unwrap();
307 assert_eq!(producers.len(), 4);
308 assert_eq!(consumers.len(), 2);
309 assert_eq!(consumers[0].n_rings(), 2);
312 assert_eq!(consumers[1].n_rings(), 2);
313
314 for (i, p) in producers.iter().enumerate() {
315 let mut buf = [0u8; SPSC_PAYLOAD_BYTES];
316 buf[..4].copy_from_slice(&(i as u32).to_le_bytes());
317 p.try_push(&buf).unwrap();
318 }
319
320 let mut c0_seen = Vec::new();
322 let mut c1_seen = Vec::new();
323 let mut out = [0u8; SPSC_PAYLOAD_BYTES];
324 while consumers[0].try_pop(&mut out).is_ok() {
325 c0_seen.push(u32::from_le_bytes(out[..4].try_into().unwrap()));
326 }
327 while consumers[1].try_pop(&mut out).is_ok() {
328 c1_seen.push(u32::from_le_bytes(out[..4].try_into().unwrap()));
329 }
330 c0_seen.sort();
331 c1_seen.sort();
332 assert_eq!(c0_seen, vec![0, 2]);
333 assert_eq!(c1_seen, vec![1, 3]);
334 }
335
336 #[test]
337 fn concurrent_mpmc_loses_no_items() {
338 const N_PRODUCERS: usize = 4;
339 const N_CONSUMERS: usize = 2;
340 const PER_PRODUCER: u32 = 10_000;
341
342 let (producers, consumers) =
343 SharedRingMpmc::create_anon_grid(N_PRODUCERS, N_CONSUMERS, 64).unwrap();
344
345 let producer_handles: Vec<_> = producers
346 .into_iter()
347 .enumerate()
348 .map(|(pid, p)| {
349 thread::spawn(move || {
350 for i in 0..PER_PRODUCER {
351 let mut buf = [0u8; SPSC_PAYLOAD_BYTES];
352 buf[..4].copy_from_slice(&(pid as u32).to_le_bytes());
353 buf[4..8].copy_from_slice(&i.to_le_bytes());
354 while p.try_push(&buf).is_err() {
355 std::hint::spin_loop();
356 }
357 }
358 })
359 })
360 .collect();
361
362 let target_per_consumer = (PER_PRODUCER as usize * N_PRODUCERS / N_CONSUMERS) as u32;
363 let consumer_handles: Vec<_> = consumers
364 .into_iter()
365 .map(|c| {
366 thread::spawn(move || -> (u32, std::collections::HashMap<u32, u32>) {
367 let mut next: std::collections::HashMap<u32, u32> = Default::default();
368 let mut total: u32 = 0;
369 let mut out = [0u8; SPSC_PAYLOAD_BYTES];
370 while total < target_per_consumer {
371 if c.try_pop(&mut out).is_ok() {
372 let pid = u32::from_le_bytes(out[..4].try_into().unwrap());
373 let seq = u32::from_le_bytes(out[4..8].try_into().unwrap());
374 let expected = next.entry(pid).or_insert(0);
375 assert_eq!(*expected, seq,
376 "per-producer FIFO violated for producer {pid}: expected {} got {}",
377 expected, seq);
378 *expected += 1;
379 total += 1;
380 } else {
381 std::hint::spin_loop();
382 }
383 }
384 (total, next)
385 })
386 })
387 .collect();
388
389 for h in producer_handles {
390 h.join().unwrap();
391 }
392 let mut grand_total: u32 = 0;
393 for h in consumer_handles {
394 let (t, _next) = h.join().unwrap();
395 grand_total += t;
396 }
397 assert_eq!(grand_total, PER_PRODUCER * N_PRODUCERS as u32);
398 }
399
400 #[test]
401 fn create_grid_in_region_round_trip() {
402 use crate::spsc_ring::{spsc_ring_file_size, RegionOwner};
407 #[repr(C, align(64))]
410 #[derive(Clone, Copy)]
411 struct Block64([u8; 64]);
412 struct HeapRegion(Vec<Block64>);
413 impl RegionOwner for HeapRegion {
414 fn region_ptr(&mut self) -> *mut u8 {
415 self.0.as_mut_ptr() as *mut u8
416 }
417 fn region_len(&self) -> usize { self.0.len() * 64 }
418 }
419
420 let (n_prod, n_cons, cap) = (4usize, 2usize, 8usize);
421 let bytes = spsc_ring_file_size(cap) * n_prod;
422 let region = HeapRegion(vec![Block64([0u8; 64]); bytes.div_ceil(64)]);
423 let (producers, consumers) =
424 SharedRingMpmc::create_grid_in_region(region, n_prod, n_cons, cap)
425 .unwrap();
426 assert_eq!(producers.len(), 4);
427 assert_eq!(consumers.len(), 2);
428
429 for (i, p) in producers.iter().enumerate() {
432 let mut buf = [0u8; SPSC_PAYLOAD_BYTES];
433 buf[..4].copy_from_slice(&(i as u32).to_le_bytes());
434 p.try_push(&buf).unwrap();
435 }
436 let mut seen = Vec::new();
437 let mut out = [0u8; SPSC_PAYLOAD_BYTES];
438 for c in &consumers {
439 while c.try_pop(&mut out).is_ok() {
440 seen.push(u32::from_le_bytes(out[..4].try_into().unwrap()));
441 }
442 }
443 seen.sort();
444 assert_eq!(seen, vec![0, 1, 2, 3]);
445 }
446}