Skip to main content

soak/
soak.rs

1// Ring correctness under sustained all-to-all load.
2//
3//   mpirun -n <2+> soak [mode=safe] [messages=20000] [payload=256] [depth=8]
4//                       [noise_kib=0] [pace_ns=0] [rep=0]  # rep labels output rows
5//
6// Every rank opens a lane to every other rank and drives them all from the main
7// thread while a poller thread drains its own incoming lanes. Payloads are
8// self-describing: origin, sequence, and a sequence-derived fill, so a receiver
9// checks *content*, not just arrival counts.
10//
11// What is asserted, per mode:
12//
13//   safe  every message arrives exactly once, in order, intact. Zero loss.
14//   raw   arrivals are a strictly increasing subsequence of what was sent, every
15//         arrival is intact, and delivered + lost equals what was sent. Loss is
16//         expected; corruption, reordering and duplication are not.
17//
18// `noise_kib > 0` runs background p2p ping-pong between rank pairs on a separate
19// thread, so the ring is measured against a busy progress engine rather than an
20// idle one.
21//
22// `pace_ns > 0` holds each sender to one round every `pace_ns` nanoseconds, one
23// round being one message to every peer. That is the offered-rate axis: sweeping
24// it against raw-mode loss is what locates the point where a lossy lane of a
25// given depth stops keeping up. At 0 the sender runs flat out.
26//
27// Writes one TSV row per rank on stdout, diagnostics on stderr.
28
29use std::collections::HashMap;
30use std::sync::Arc;
31use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
32
33use mpi::Threading;
34use mpi::collective::CommunicatorCollectives;
35use mpi::point_to_point::{Destination, Source};
36use mpi::topology::{Communicator, Rank, SimpleCommunicator};
37
38use mpi_rma::{Message, Ring};
39
40const TAG_NOISE: i32 = 11;
41const TAG_STOP: i32 = 12;
42/// Origin and sequence, ahead of the fill.
43const STAMP: usize = 12;
44
45fn paint(origin: Rank, sequence: u64, buf: &mut [u8]) {
46    buf[..4].copy_from_slice(&origin.to_le_bytes());
47    buf[4..STAMP].copy_from_slice(&sequence.to_le_bytes());
48    for (i, b) in buf[STAMP..].iter_mut().enumerate() {
49        *b = (sequence as u8).wrapping_add(i as u8);
50    }
51}
52
53/// Check a payload against the origin and sequence the ring reported for it.
54fn verify(m: &Message) -> Result<(), String> {
55    if m.data.len() < STAMP {
56        return Err(format!(
57            "payload from {} is {} bytes",
58            m.origin,
59            m.data.len()
60        ));
61    }
62    let origin = i32::from_le_bytes(m.data[..4].try_into().unwrap());
63    let sequence = u64::from_le_bytes(m.data[4..STAMP].try_into().unwrap());
64    if origin != m.origin || sequence != m.sequence {
65        return Err(format!(
66            "payload says ({origin}, {sequence}), ring says ({}, {})",
67            m.origin, m.sequence
68        ));
69    }
70    for (i, &b) in m.data[STAMP..].iter().enumerate() {
71        let want = (sequence as u8).wrapping_add(i as u8);
72        if b != want {
73            return Err(format!(
74                "payload from {origin} seq {sequence} corrupt at byte {}: {b} != {want}",
75                i + STAMP
76            ));
77        }
78    }
79    Ok(())
80}
81
82/// Carries a communicator reference onto the noise thread.
83///
84/// rsmpi wraps a raw MPI handle, so `SimpleCommunicator` is neither `Send` nor
85/// `Sync`. Under `MPI_THREAD_MULTIPLE` concurrent calls are legal, and the noise
86/// thread only ever talks to its own partner on its own tags.
87struct Shared<'a>(&'a SimpleCommunicator);
88unsafe impl Send for Shared<'_> {}
89
90/// Background p2p ping-pong inside rank pairs (0,1), (2,3), ... until `stop`.
91///
92/// The lower rank of each pair owns termination: it stops sending payload and
93/// sends a STOP instead, which the higher rank answers by leaving. An odd rank
94/// at the end has no partner and idles.
95fn noise(world: &SimpleCommunicator, kib: usize, stop: &AtomicBool) {
96    let rank = world.rank();
97    let partner = if rank % 2 == 0 { rank + 1 } else { rank - 1 };
98    if partner >= world.size() {
99        return;
100    }
101    let payload = vec![0xA5u8; kib * 1024];
102    let peer = world.process_at_rank(partner);
103    loop {
104        if rank < partner {
105            if stop.load(Ordering::Relaxed) {
106                peer.send_with_tag(&[0u8], TAG_STOP);
107                return;
108            }
109            peer.send_with_tag(&payload[..], TAG_NOISE);
110            peer.receive_vec_with_tag::<u8>(TAG_NOISE);
111        } else {
112            let (_, status) = peer.receive_vec::<u8>();
113            if status.tag() == TAG_STOP {
114                return;
115            }
116            peer.send_with_tag(&payload[..], TAG_NOISE);
117        }
118    }
119}
120
121fn main() {
122    let mut args = std::env::args().skip(1);
123    let mode = args.next().unwrap_or_else(|| "safe".into());
124    let messages: u64 = args.next().and_then(|s| s.parse().ok()).unwrap_or(20_000);
125    let payload: usize = args.next().and_then(|s| s.parse().ok()).unwrap_or(256);
126    let depth: usize = args.next().and_then(|s| s.parse().ok()).unwrap_or(8);
127    let noise_kib: usize = args.next().and_then(|s| s.parse().ok()).unwrap_or(0);
128    let pace_ns: u64 = args.next().and_then(|s| s.parse().ok()).unwrap_or(0);
129    // Echoed straight into the output so repeats of one point stay distinguishable.
130    let rep: u32 = args.next().and_then(|s| s.parse().ok()).unwrap_or(0);
131    assert!(
132        payload >= STAMP,
133        "payload must hold the origin and sequence"
134    );
135    let safe = match mode.as_str() {
136        "safe" => true,
137        "raw" => false,
138        other => panic!("mode must be safe or raw, got {other:?}"),
139    };
140
141    let (universe, provided) =
142        mpi::initialize_with_threading(Threading::Multiple).expect("MPI must initialize once");
143    assert_eq!(provided, Threading::Multiple);
144    let world = universe.world();
145    let rank = world.rank();
146    let size = world.size();
147    assert!(size >= 2, "soak needs at least 2 ranks");
148
149    // All-to-all: every ordered pair of distinct ranks gets a lane.
150    let mut lanes = Vec::new();
151    for source in 0..size {
152        for destination in 0..size {
153            if source != destination {
154                lanes.push((source, destination, depth, payload));
155            }
156        }
157    }
158    let ring = Arc::new(if safe {
159        Ring::safe(&world, &lanes).unwrap()
160    } else {
161        Ring::raw(&world, &lanes).unwrap()
162    });
163
164    let stop = Arc::new(AtomicBool::new(false));
165    let done = Arc::new(AtomicBool::new(false));
166    let count = Arc::new(AtomicU64::new(0));
167    let peers = (size - 1) as u64;
168    let sent = messages * peers;
169
170    world.barrier();
171    let started = std::time::Instant::now();
172
173    let outcome = std::thread::scope(|scope| {
174        if noise_kib > 0 {
175            let stop = Arc::clone(&stop);
176            let shared = Shared(&world);
177            scope.spawn(move || {
178                // Name the wrapper as a whole first. Closures capture the
179                // narrowest place they use, and `shared.0` alone would be
180                // captured as a bare reference, losing the `Send` impl.
181                let shared = shared;
182                noise(shared.0, noise_kib, &stop);
183            });
184        }
185
186        let poller = {
187            let (ring, done, count) = (Arc::clone(&ring), Arc::clone(&done), Arc::clone(&count));
188            scope.spawn(move || -> Result<Vec<u64>, String> {
189                let mut seen = vec![0u64; size as usize];
190                let mut total = 0u64;
191                loop {
192                    // If true, this poll follows the global send barrier.
193                    let finished = done.load(Ordering::Acquire);
194                    let batch = ring.poll().map_err(|e| format!("poll: {e}"))?;
195                    let empty = batch.is_empty();
196                    let mut furthest: HashMap<Rank, u64> = HashMap::new();
197                    for m in &batch {
198                        verify(m)?;
199                        let previous = seen[m.origin as usize];
200                        if m.sequence <= previous {
201                            return Err(format!(
202                                "sequence from {} went backwards: {} after {previous}",
203                                m.origin, m.sequence
204                            ));
205                        }
206                        if safe && m.sequence != previous + 1 {
207                            return Err(format!(
208                                "safe lane from {} skipped {} to {}",
209                                m.origin,
210                                previous + 1,
211                                m.sequence
212                            ));
213                        }
214                        seen[m.origin as usize] = m.sequence;
215                        let f = furthest.entry(m.origin).or_insert(0);
216                        *f = (*f).max(m.sequence);
217                        total += 1;
218                    }
219                    for (origin, sequence) in furthest {
220                        ring.ack(origin, sequence)
221                            .map_err(|e| format!("ack to {origin}: {e}"))?;
222                    }
223                    count.store(total, Ordering::Relaxed);
224                    if finished && empty {
225                        return Ok(seen);
226                    }
227                    if empty {
228                        std::thread::yield_now();
229                    }
230                }
231            })
232        };
233
234        // Round-robin over peers so every outgoing lane stays hot at once.
235        let mut buf = vec![0u8; payload];
236        let pace = std::time::Duration::from_nanos(pace_ns);
237        let opened = std::time::Instant::now();
238        for sequence in 1..=messages {
239            for destination in (0..size).filter(|&r| r != rank) {
240                paint(rank, sequence, &mut buf);
241                ring.send(destination, &buf)
242                    .unwrap_or_else(|e| panic!("rank {rank}: send to {destination}: {e}"));
243            }
244            // Absolute deadlines, so a slow round is not paid for twice and the
245            // offered rate stays the one that was asked for.
246            if pace_ns > 0 {
247                let due = opened + pace * sequence as u32;
248                while std::time::Instant::now() < due {
249                    std::hint::spin_loop();
250                }
251            }
252        }
253        // Every rank's sends have completed at their targets before this
254        // returns, so after the barrier nothing new can appear in a slot.
255        world.barrier();
256        done.store(true, Ordering::Release);
257
258        let seen = poller.join().expect("poller thread panicked");
259        stop.store(true, Ordering::Relaxed);
260        seen
261    });
262
263    let elapsed = started.elapsed().as_secs_f64();
264    let fail = |why: String| -> ! {
265        eprintln!("[soak] FAIL rank {rank}: {why}");
266        world.abort(1)
267    };
268    let seen = outcome.unwrap_or_else(|why| fail(why));
269
270    let got = count.load(Ordering::Relaxed);
271    let lost = ring.lost();
272    let corrupt = ring.corrupt();
273
274    // The last message on a lane is never overwritten, so every lane must end
275    // exactly at the sender's high-water mark in both modes.
276    for origin in (0..size).filter(|&r| r != rank) {
277        if seen[origin as usize] != messages {
278            fail(format!(
279                "lane from {origin} ended at {} not {messages}",
280                seen[origin as usize]
281            ));
282        }
283    }
284    if safe && lost != 0 {
285        fail(format!("safe ring lost {lost}"));
286    }
287    if got + lost != sent {
288        fail(format!("received {got} + lost {lost} != sent {sent}"));
289    }
290
291    // Gather the counters and let rank 0 write every row. Interleaving the
292    // ranks' own stdout would scatter the header into the middle of the table:
293    // mpirun orders a rank's output against itself and nothing else.
294    let mine = [
295        got,
296        lost,
297        corrupt,
298        ring.max_lag(),
299        ring.waits(),
300        ring.wait_ns(),
301    ];
302    let mut all = vec![0u64; mine.len() * size as usize];
303    world.all_gather_into(&mine[..], &mut all[..]);
304    if rank == 0 {
305        println!(
306            "mode\tranks\tdepth\tpayload\tnoise_kib\tpace_ns\trep\trank\tsent\treceived\tlost\tcorrupt\tmax_lag\twaits\twait_s\tseconds"
307        );
308        for (r, row) in all.chunks_exact(mine.len()).enumerate() {
309            println!(
310                "{mode}\t{size}\t{depth}\t{payload}\t{noise_kib}\t{pace_ns}\t{rep}\t{r}\t{sent}\t{}\t{}\t{}\t{}\t{}\t{:.6}\t{elapsed:.6}",
311                row[0],
312                row[1],
313                row[2],
314                row[3],
315                row[4],
316                row[5] as f64 * 1e-9,
317            );
318        }
319    }
320
321    world.barrier();
322    Arc::try_unwrap(ring)
323        .unwrap_or_else(|_| panic!("ring still shared at close"))
324        .close()
325        .unwrap();
326    if rank == 0 {
327        eprintln!("[soak] ok: {mode}, {size} ranks, {sent} messages/rank in {elapsed:.2}s");
328    }
329}