1use 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#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct Message {
24 pub origin: Rank,
26 pub sequence: u64,
28 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
46struct Outgoing {
48 sent: u64,
50 acked: u64,
52 image: Vec<u8>,
57}
58
59pub struct Ring {
65 slots: Window<u8>,
66 acks: Option<Window<u64>>,
67 ranks: usize,
68 to: Vec<Option<Lane>>,
70 from: Vec<Option<Lane>>,
72 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 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 pub fn raw<C: Communicator + ?Sized>(
123 comm: &C,
124 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 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 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 pub fn is_safe(&self) -> bool {
284 self.acks.is_some()
285 }
286
287 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 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 pub fn lost(&self) -> u64 {
305 self.lost.load(Ordering::Relaxed)
306 }
307
308 pub fn corrupt(&self) -> u64 {
310 self.corrupt.load(Ordering::Relaxed)
311 }
312
313 pub fn max_lag(&self) -> u64 {
318 self.max_lag.load(Ordering::Relaxed)
319 }
320
321 pub fn waits(&self) -> u64 {
326 self.waits.load(Ordering::Relaxed)
327 }
328
329 pub fn wait_ns(&self) -> u64 {
331 self.wait_ns.load(Ordering::Relaxed)
332 }
333
334 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 *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 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 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 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 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 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 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}