Skip to main content

mpi_rma/
ring.rs

1//! Fixed-slot message transport over RMA windows.
2//!
3//! See `design.md` for the slot layout, sequencing rules, and acknowledge
4//! gating. Construction is collective over the configured communicator and
5//! requires `MPI_THREAD_MULTIPLE` plus a unified window memory model.
6
7use std::sync::Mutex;
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::time::Instant;
10
11use crc32fast::Hasher;
12use mpi::collective::CommunicatorCollectives;
13use mpi::topology::{Communicator, Rank};
14
15use crate::{CommunicatorRmaExt, Error, MemoryModel, Window};
16
17const WORD: usize = size_of::<u64>();
18const CHECKSUM: usize = size_of::<u32>();
19const HEADER: usize = 2 * WORD + CHECKSUM;
20
21/// One message observed in a ring.
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct Message {
24    /// Rank that sent the message.
25    pub origin: Rank,
26    /// Monotonic sequence number within the directed lane.
27    pub sequence: u64,
28    /// Message payload copied from the slot.
29    pub data: Vec<u8>,
30}
31
32#[derive(Clone, Copy)]
33struct Lane {
34    offset: usize,
35    depth: usize,
36    capacity: usize,
37    slot: usize,
38    ack: usize,
39}
40
41struct Frame {
42    sequence: u64,
43    data: Vec<u8>,
44}
45
46/// Per-lane sender state
47struct Outgoing {
48    /// Highest sequence put on the wire.
49    sent: u64,
50    /// Cumulative acknowledgement, as last read from the counter window.
51    acked: u64,
52    /// Slot-sized scratch for the image, reused across sends.
53    /// Only the header, the payload and the guard are written.
54    /// The padding between payload and guard is never zeroed because nothing reads it.
55    /// The checksum covers the declared length only, and `length` bounds what a receiver copies out.
56    image: Vec<u8>,
57}
58
59/// Sparse directed message rings packed into collective RMA windows.
60///
61/// Safe to share between threads on the same process.
62/// Lifetime is the underlying MPI window
63/// To drop the ring use [`Self::close`] (collective) or let the destructor run symmetrically on every rank.
64pub struct Ring {
65    slots: Window<u8>,
66    acks: Option<Window<u64>>,
67    ranks: usize,
68    /// Outgoing lanes, indexed by destination
69    to: Vec<Option<Lane>>,
70    /// Incoming lanes, indexed by origin
71    from: Vec<Option<Lane>>,
72    // Ranks with incoming lanes
73    incoming: Vec<Rank>,
74    sent: Vec<Mutex<Outgoing>>,
75    seen: Mutex<Vec<u64>>,
76    acked: Mutex<Vec<u64>>,
77    lost: AtomicU64,
78    corrupt: AtomicU64,
79    max_lag: AtomicU64,
80    waits: AtomicU64,
81    wait_ns: AtomicU64,
82}
83
84impl Ring {
85    /// Construct a safe (overwrite-gated) ring. Collective over `comm`.
86    ///
87    /// `rings` names the active directed lanes, one `(source, destination,
88    /// depth, capacity)` tuple per lane: `depth` is the number of fixed slots
89    /// the lane holds, `capacity` the payload bytes one slot fits. Every rank
90    /// passes the same list.
91    ///
92    /// Each source lane gets a cumulative-acknowledgement counter at the source
93    /// The senders spin on `yield_now` until the receiver has acked
94    /// enough earlier messages to make room in the slot ring. Requires a
95    /// unified window memory model otherwise contruction fails
96    ///
97    /// # Errors
98    /// - [`Error::Intercommunicator`] if `comm` is an intercommunicator.
99    /// - [`Error::Ring`] for configuration disagreement, invalid or repeated
100    ///   lanes, or zero depth or capacity.
101    /// - [`Error::Window`] if the window uses a separate memory model, which
102    ///   cannot support the local polling the ring relies on.
103    /// - Plus whatever [`CommunicatorRmaExt::allocate_window`] returns for
104    ///   the slot and counter windows.
105    pub fn safe<C: Communicator + ?Sized>(
106        comm: &C,
107        rings: &[(Rank, Rank, usize, usize)],
108    ) -> Result<Self, Error> {
109        Self::new(comm, rings, true)
110    }
111
112    /// Construct a raw (overwrite-capable) ring. Collective over `comm`.
113    ///
114    /// `rings` has the same shape as for [`Self::safe`]. No acknowledgement
115    /// counter: unread slots are silently overwritten when the depth is
116    /// exhausted, and the receiver reports the gap via [`Self::lost`]. The
117    /// sender never blocks on a slow peer. Requires a unified window memory
118    /// model.
119    ///
120    /// # Errors
121    /// Same conditions as [`Self::safe`].
122    pub fn raw<C: Communicator + ?Sized>(
123        comm: &C,
124        // (Source: Rank, Destination: Rank, Depth: usize, Capacity: usize)
125        rings: &[(Rank, Rank, usize, usize)],
126    ) -> Result<Self, Error> {
127        Self::new(comm, rings, false)
128    }
129
130    fn new<C: Communicator + ?Sized>(
131        comm: &C,
132        rings: &[(Rank, Rank, usize, usize)],
133        safe: bool,
134    ) -> Result<Self, Error> {
135        if comm.test_inter() {
136            return Err(Error::Intercommunicator);
137        }
138        let ranks = usize::try_from(comm.size()).map_err(|_| Error::SizeOverflow)?;
139        let me = usize::try_from(comm.rank()).map_err(|_| Error::SizeOverflow)?;
140
141        let mut config: Vec<_> = rings.to_vec();
142        config.sort_unstable();
143        Self::agree(comm, safe, &config)?;
144        for pair in config.windows(2) {
145            if pair[0].0 == pair[1].0 && pair[0].1 == pair[1].1 {
146                return Err(Error::Ring("a lane is configured twice"));
147            }
148        }
149        for &(source, destination, depth, capacity) in &config {
150            if source < 0
151                || destination < 0
152                || source as usize >= ranks
153                || destination as usize >= ranks
154            {
155                return Err(Error::Ring("lane rank is outside the communicator"));
156            }
157            if source == destination {
158                return Err(Error::Ring("self lanes are not transport"));
159            }
160            if depth == 0 {
161                return Err(Error::Ring("depth must be positive"));
162            }
163            if capacity == 0 {
164                return Err(Error::Ring("capacity must be positive"));
165            }
166        }
167
168        let mut lengths = vec![0usize; ranks];
169        let mut ack_lengths = vec![0usize; ranks];
170        let mut to = vec![None; ranks];
171        let mut from = vec![None; ranks];
172        for &(source, destination, depth, capacity) in &config {
173            let slot = capacity
174                .checked_add(HEADER + WORD)
175                .ok_or(Error::SizeOverflow)?;
176            if slot > i32::MAX as usize {
177                return Err(Error::CountOverflow);
178            }
179            let bytes = depth.checked_mul(slot).ok_or(Error::SizeOverflow)?;
180            let target = destination as usize;
181            let offset = lengths[target];
182            lengths[target] = offset.checked_add(bytes).ok_or(Error::SizeOverflow)?;
183            // The ack counter index counts the source's outgoing lanes in
184            // configuration order. Every rank derives it the same way, even
185            // for lanes it does not own: the `from` table needs it to target
186            // the counter on the origin's window.
187            let ack = ack_lengths[source as usize];
188            ack_lengths[source as usize] = ack.checked_add(1).ok_or(Error::SizeOverflow)?;
189            let lane = Lane {
190                offset,
191                depth,
192                capacity,
193                slot,
194                ack,
195            };
196            // Only lanes where this rank is an endpoint belong to its tables.
197            if source as usize == me {
198                to[target] = Some(lane);
199            }
200            if target == me {
201                from[source as usize] = Some(lane);
202            }
203        }
204
205        let slots = comm.allocate_window::<u8>(lengths[me])?;
206        let acks = safe
207            .then(|| comm.allocate_window::<u64>(ack_lengths[me]))
208            .transpose()?;
209        let unified = slots.memory_model() == MemoryModel::Unified
210            && acks
211                .as_ref()
212                .is_none_or(|window| window.memory_model() == MemoryModel::Unified);
213        let mut models = vec![0u8; ranks];
214        comm.all_gather_into(&(unified as u8), &mut models[..]);
215        if models.contains(&0) {
216            return Err(Error::Window("local ring access requires unified memory"));
217        }
218
219        let mut incoming: Vec<_> = config
220            .iter()
221            .filter_map(|&(source, destination, _, _)| {
222                (destination == comm.rank()).then_some(source)
223            })
224            .collect();
225        incoming.sort_unstable();
226
227        Ok(Ring {
228            slots,
229            acks,
230            ranks,
231            to,
232            from,
233            incoming,
234            sent: (0..ranks)
235                .map(|_| {
236                    Mutex::new(Outgoing {
237                        sent: 0,
238                        acked: 0,
239                        image: Vec::new(),
240                    })
241                })
242                .collect(),
243            seen: Mutex::new(vec![0; ranks]),
244            acked: Mutex::new(vec![0; ranks]),
245            lost: AtomicU64::new(0),
246            corrupt: AtomicU64::new(0),
247            max_lag: AtomicU64::new(0),
248            waits: AtomicU64::new(0),
249            wait_ns: AtomicU64::new(0),
250        })
251    }
252
253    fn agree<C: Communicator + ?Sized>(
254        comm: &C,
255        safe: bool,
256        config: &[(Rank, Rank, usize, usize)],
257    ) -> Result<(), Error> {
258        let ranks = usize::try_from(comm.size()).map_err(|_| Error::SizeOverflow)?;
259        let mut counts = vec![0u64; ranks];
260        comm.all_gather_into(&(config.len() as u64), &mut counts[..]);
261        if counts.iter().any(|&count| count != config.len() as u64) {
262            return Err(Error::Ring("configuration differs between ranks"));
263        }
264        let mut local = Vec::with_capacity(1 + config.len() * 4);
265        local.push(safe as u64);
266        for &(source, destination, depth, capacity) in config {
267            local.extend_from_slice(&[
268                source as u64,
269                destination as u64,
270                depth as u64,
271                capacity as u64,
272            ]);
273        }
274        let mut all = vec![0u64; local.len() * ranks];
275        comm.all_gather_into(&local[..], &mut all[..]);
276        if all.chunks_exact(local.len()).any(|row| row != local) {
277            return Err(Error::Ring("configuration differs between ranks"));
278        }
279        Ok(())
280    }
281
282    /// Whether this ring gates overwrites with cumulative acknowledgements.
283    pub fn is_safe(&self) -> bool {
284        self.acks.is_some()
285    }
286
287    /// Slot count configured for the outgoing lane to `destination`, if any.
288    pub fn depth(&self, destination: Rank) -> Option<usize> {
289        self.peer(destination)
290            .ok()
291            .and_then(|i| self.to[i])
292            .map(|lane| lane.depth)
293    }
294
295    /// Per-slot payload capacity configured for the outgoing lane to `destination`, if any.
296    pub fn capacity(&self, destination: Rank) -> Option<usize> {
297        self.peer(destination)
298            .ok()
299            .and_then(|i| self.to[i])
300            .map(|lane| lane.capacity)
301    }
302
303    /// Total messages observed as lost (raw mode only).
304    pub fn lost(&self) -> u64 {
305        self.lost.load(Ordering::Relaxed)
306    }
307
308    /// Total corrupt slot reads: torn header/footer, bad length, or CRC mismatch.
309    pub fn corrupt(&self) -> u64 {
310        self.corrupt.load(Ordering::Relaxed)
311    }
312
313    /// Greatest sequence advance one [`Self::poll`] made on one lane.
314    ///
315    /// In raw mode this includes gaps counted as lost. In safe mode it is the
316    /// number of messages drained.
317    pub fn max_lag(&self) -> u64 {
318        self.max_lag.load(Ordering::Relaxed)
319    }
320
321    /// Number of times a safe-mode sender actually blocked on acknowledgements.
322    ///
323    /// Refreshing the cached counter does not count; only a sender that
324    /// found no headroom after refreshing and had to spin does.
325    pub fn waits(&self) -> u64 {
326        self.waits.load(Ordering::Relaxed)
327    }
328
329    /// Cumulative nanoseconds spent in safe-mode wait spins.
330    pub fn wait_ns(&self) -> u64 {
331        self.wait_ns.load(Ordering::Relaxed)
332    }
333
334    /// Send one message and return its sequence number.
335    ///
336    /// Completes the put at the target before returning. In safe mode,
337    /// spins on `yield_now` if the destination has not yet acknowledged
338    /// enough earlier messages to make room in the slot ring; the wait is
339    /// recorded in [`Self::waits`] and [`Self::wait_ns`].
340    ///
341    /// # Errors
342    /// - [`Error::Rank`] if `destination` is outside the communicator.
343    /// - [`Error::Ring`] if no lane to `destination` is configured, the
344    ///   sequence number would overflow, the acknowledgement counter
345    ///   regresses or exceeds what was sent, or the send state is poisoned.
346    /// - [`Error::Payload`] if `data` does not fit the lane's capacity.
347    /// - Whatever the underlying `MPI_Put` fails with (wrapped), e.g.
348    ///   [`Error::Mpi`] or [`Error::Range`].
349    pub fn send(&self, destination: Rank, data: &[u8]) -> Result<u64, Error> {
350        let destination_index = self.peer(destination)?;
351        let lane =
352            self.to[destination_index].ok_or(Error::Ring("directed lane is not configured"))?;
353        if data.len() > lane.capacity {
354            return Err(Error::Payload {
355                len: data.len(),
356                capacity: lane.capacity,
357            });
358        }
359
360        let mut guard = self.sent[destination_index]
361            .lock()
362            .map_err(|_| Error::Ring("send state poisoned"))?;
363        let Outgoing { sent, acked, image } = &mut *guard;
364        let sequence = sent
365            .checked_add(1)
366            .ok_or(Error::Ring("sequence number exhausted"))?;
367        if self.is_safe() && sequence - *acked > lane.depth as u64 {
368            // The cached counter only moves when read here, so it is stale
369            // roughly once per `depth` sends. Refresh before calling this a
370            // wait: otherwise `waits` counts cache misses, not blocked senders.
371            *acked = self.acknowledged(lane, *sent, *acked)?;
372            if sequence - *acked > lane.depth as u64 {
373                self.waits.fetch_add(1, Ordering::Relaxed);
374                let started = Instant::now();
375                while sequence - *acked > lane.depth as u64 {
376                    std::thread::yield_now();
377                    *acked = self.acknowledged(lane, *sent, *acked)?;
378                }
379                self.wait_ns.fetch_add(
380                    started.elapsed().as_nanos().min(u128::from(u64::MAX)) as u64,
381                    Ordering::Relaxed,
382                );
383            }
384        }
385
386        // One buffer per lane, grown once. A slot-sized allocation and memset
387        // per send costs more than the transfer itself on a wide lane.
388        if image.len() != lane.slot {
389            image.resize(lane.slot, 0);
390        }
391        let len = data.len() as u64;
392        image[..WORD].copy_from_slice(&sequence.to_le_bytes());
393        image[WORD..2 * WORD].copy_from_slice(&len.to_le_bytes());
394        image[HEADER..HEADER + data.len()].copy_from_slice(data);
395        let checksum = Self::checksum(sequence, len, data);
396        image[2 * WORD..HEADER].copy_from_slice(&checksum.to_le_bytes());
397        image[lane.slot - WORD..].copy_from_slice(&(!sequence).to_le_bytes());
398
399        let position = ((sequence - 1) % lane.depth as u64) as usize;
400        self.slots
401            .put(destination, lane.offset + position * lane.slot, image)?;
402        *sent = sequence;
403        Ok(sequence)
404    }
405
406    /// Drain every incoming lane and return newly observed messages.
407    ///
408    /// Does not call MPI. Safe mode delivers in-order; raw mode scans all
409    /// slots when one is lapped, sorts the survivors, and counts the gaps
410    /// into [`Self::lost`]. The per-origin `seen` cursor only advances.
411    ///
412    /// # Errors
413    /// - [`Error::Lapped`] if a safe lane was overwritten before consumption,
414    ///   which the acknowledge gate makes impossible and so indicates a
415    ///   corrupted counter or window rather than backpressure.
416    /// - Any window read error from the local slot memory.
417    pub fn poll(&self) -> Result<Vec<Message>, Error> {
418        let mut seen = self
419            .seen
420            .lock()
421            .map_err(|_| Error::Ring("receive state poisoned"))?;
422        let mut messages = Vec::new();
423
424        for &origin in &self.incoming {
425            let origin_index = origin as usize;
426            let lane =
427                self.from[origin_index].ok_or(Error::Ring("directed lane is not configured"))?;
428            let behind = seen[origin_index];
429            let mut read = 0;
430            while read < lane.depth {
431                let Some(expected) = seen[origin_index].checked_add(1) else {
432                    break;
433                };
434                let position = ((expected - 1) % lane.depth as u64) as usize;
435                let Some(found) = self.probe(lane, position)? else {
436                    break;
437                };
438                if found < expected {
439                    break;
440                }
441                if found > expected {
442                    // A safe lane cannot lap: the sender cannot reach this slot
443                    // again until `expected` has been acknowledged, and an ack
444                    // only follows consumption. Reaching here means the counter
445                    // or the window has been corrupted, and continuing would
446                    // silently turn a guaranteed lane into a lossy one.
447                    if self.is_safe() {
448                        return Err(Error::Lapped {
449                            origin,
450                            expected,
451                            found,
452                        });
453                    }
454                    let frames = self.recover(lane, seen[origin_index])?;
455                    let mut next = expected;
456                    for frame in frames {
457                        if frame.sequence < next {
458                            continue;
459                        }
460                        self.lost
461                            .fetch_add(frame.sequence - next, Ordering::Relaxed);
462                        next = frame.sequence.saturating_add(1);
463                        seen[origin_index] = frame.sequence;
464                        messages.push(Message {
465                            origin,
466                            sequence: frame.sequence,
467                            data: frame.data,
468                        });
469                    }
470                    break;
471                }
472
473                let Some(frame) = self.frame(lane, position)? else {
474                    break;
475                };
476                if frame.sequence != expected {
477                    break;
478                }
479                seen[origin_index] = expected;
480                messages.push(Message {
481                    origin,
482                    sequence: expected,
483                    data: frame.data,
484                });
485                read += 1;
486            }
487            self.max_lag
488                .fetch_max(seen[origin_index] - behind, Ordering::Relaxed);
489        }
490        Ok(messages)
491    }
492
493    /// Acknowledge all messages from `origin` through `sequence`.
494    ///
495    /// Cumulative and idempotent. In safe mode performs one atomic
496    /// `MPI_Fetch_and_op` against `origin`'s counter window. In raw mode
497    /// this is a no-op (returns `Ok`).
498    ///
499    /// # Errors
500    /// - [`Error::Rank`] if `origin` is outside the communicator.
501    /// - [`Error::Ring`] if no incoming lane from `origin` is configured,
502    ///   or the counter diverges from the locally tracked value.
503    /// - [`Error::Ack`] if `sequence` is past the last received message.
504    pub fn ack(&self, origin: Rank, sequence: u64) -> Result<(), Error> {
505        let origin_index = self.peer(origin)?;
506        let Some(acks) = &self.acks else {
507            return Ok(());
508        };
509        let lane = self.from[origin_index].ok_or(Error::Ring("directed lane is not configured"))?;
510
511        let mut acked = self
512            .acked
513            .lock()
514            .map_err(|_| Error::Ring("acknowledgement state poisoned"))?;
515        if sequence <= acked[origin_index] {
516            return Ok(());
517        }
518        let seen = self
519            .seen
520            .lock()
521            .map_err(|_| Error::Ring("receive state poisoned"))?;
522        if sequence > seen[origin_index] {
523            return Err(Error::Ack {
524                origin,
525                sequence,
526                received: seen[origin_index],
527            });
528        }
529        drop(seen);
530
531        let delta = sequence - acked[origin_index];
532        let previous = acks.fetch_add(origin, lane.ack, delta)?;
533        if previous != acked[origin_index] {
534            return Err(Error::Ring("acknowledgement counter diverged"));
535        }
536        acked[origin_index] = sequence;
537        Ok(())
538    }
539
540    /// Close the underlying windows. Collective over the ring communicator.
541    ///
542    /// Prefer this over relying on `Drop`: the destructor's shutdown runs
543    /// the same steps but the collective boundary is implicit.
544    pub fn close(self) -> Result<(), Error> {
545        let slots = self.slots.close();
546        let acks = self.acks.map_or(Ok(()), Window::close);
547        slots.and(acks)
548    }
549
550    fn probe(&self, lane: Lane, position: usize) -> Result<Option<u64>, Error> {
551        let offset = lane.offset + position * lane.slot;
552        let mut head = [0; WORD];
553        let mut tail = [0; WORD];
554        self.slots.read_local_volatile(offset, &mut head)?;
555        self.slots
556            .read_local_volatile(offset + lane.slot - WORD, &mut tail)?;
557        let sequence = Self::word(&head);
558        if sequence == 0 {
559            return Ok(None);
560        }
561        if Self::word(&tail) != !sequence {
562            self.corrupt.fetch_add(1, Ordering::Relaxed);
563            return Ok(None);
564        }
565        Ok(Some(sequence))
566    }
567
568    fn frame(&self, lane: Lane, position: usize) -> Result<Option<Frame>, Error> {
569        let offset = lane.offset + position * lane.slot;
570        let mut image = vec![0; lane.slot];
571        self.slots.read_local(offset, &mut image)?;
572
573        let sequence = Self::word(&image[..WORD]);
574        let guard = Self::word(&image[lane.slot - WORD..]);
575        if sequence == 0 {
576            return Ok(None);
577        }
578        if guard != !sequence {
579            self.corrupt.fetch_add(1, Ordering::Relaxed);
580            return Ok(None);
581        }
582        let len = Self::word(&image[WORD..2 * WORD]);
583        let Ok(len) = usize::try_from(len) else {
584            self.corrupt.fetch_add(1, Ordering::Relaxed);
585            return Ok(None);
586        };
587        if len > lane.capacity {
588            self.corrupt.fetch_add(1, Ordering::Relaxed);
589            return Ok(None);
590        }
591        let checksum = u32::from_le_bytes(image[2 * WORD..HEADER].try_into().unwrap());
592        if checksum != Self::checksum(sequence, len as u64, &image[HEADER..HEADER + len]) {
593            self.corrupt.fetch_add(1, Ordering::Relaxed);
594            return Ok(None);
595        }
596        Ok(Some(Frame {
597            sequence,
598            data: image[HEADER..HEADER + len].to_vec(),
599        }))
600    }
601
602    fn recover(&self, lane: Lane, seen: u64) -> Result<Vec<Frame>, Error> {
603        let mut frames = Vec::with_capacity(lane.depth);
604        for position in 0..lane.depth {
605            if self
606                .probe(lane, position)?
607                .is_some_and(|sequence| sequence > seen)
608                && let Some(frame) = self.frame(lane, position)?
609                && frame.sequence > seen
610            {
611                frames.push(frame);
612            }
613        }
614        frames.sort_unstable_by_key(|frame| frame.sequence);
615        frames.dedup_by_key(|frame| frame.sequence);
616        Ok(frames)
617    }
618
619    /// Re-read `lane`'s cumulative acknowledgement, rejecting a counter that
620    /// went backwards or ran past what this rank has sent.
621    fn acknowledged(&self, lane: Lane, sent: u64, acked: u64) -> Result<u64, Error> {
622        let mut value = [0];
623        self.acks
624            .as_ref()
625            .expect("safe ring has acknowledgements")
626            .read_local_volatile(lane.ack, &mut value)?;
627        if value[0] < acked {
628            return Err(Error::Ring("acknowledgement counter regressed"));
629        }
630        if value[0] > sent {
631            return Err(Error::Ring("acknowledgement exceeds sent sequence"));
632        }
633        Ok(value[0])
634    }
635
636    fn checksum(sequence: u64, len: u64, data: &[u8]) -> u32 {
637        let mut checksum = Hasher::new();
638        checksum.update(&sequence.to_le_bytes());
639        checksum.update(&len.to_le_bytes());
640        checksum.update(data);
641        checksum.finalize()
642    }
643
644    fn word(bytes: &[u8]) -> u64 {
645        u64::from_le_bytes(bytes.try_into().unwrap())
646    }
647
648    fn peer(&self, rank: Rank) -> Result<usize, Error> {
649        if rank < 0 || rank as usize >= self.ranks {
650            Err(Error::Rank(rank))
651        } else {
652            Ok(rank as usize)
653        }
654    }
655}