1use 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;
42const 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
53fn 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
82struct Shared<'a>(&'a SimpleCommunicator);
88unsafe impl Send for Shared<'_> {}
89
90fn 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 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 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 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 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 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 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 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 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 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}