1use 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
37pub struct BlockingMpscRing;
40
41pub 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
51pub 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 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 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 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 #[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 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 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 pub fn capacity(&self) -> usize { self.ring.capacity() }
282 pub fn head(&self) -> u64 { self.ring.head() }
284}
285
286impl BlockingMpscConsumer {
287 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 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 pub fn n_producers(&self) -> usize { self.rings.len() }
367
368 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}