1use std::thread;
5use std::time::{Duration, Instant};
6
7use mpi::Threading;
8use mpi::collective::CommunicatorCollectives;
9use mpi::topology::{Communicator, Rank, SimpleCommunicator};
10
11use mpi_rma::{Error, Message, Ring};
12
13fn pair(depth: usize, capacity: usize) -> Vec<(Rank, Rank, usize, usize)> {
14 vec![(0, 1, depth, capacity), (1, 0, depth, capacity)]
15}
16
17fn oneway(depth: usize, capacity: usize) -> Vec<(Rank, Rank, usize, usize)> {
18 vec![(0, 1, depth, capacity)]
19}
20
21fn safe_basics(world: &SimpleCommunicator, rank: Rank, next: Rank, prev: Rank) {
22 let ring = Ring::safe(world, &pair(4, 8)).unwrap();
26 assert!(ring.is_safe());
27 assert_eq!(ring.depth(next), Some(4));
28 assert_eq!(ring.capacity(next), Some(8));
29
30 assert!(matches!(
31 ring.send(next, &[0; 9]),
32 Err(Error::Payload {
33 len: 9,
34 capacity: 8
35 })
36 ));
37 assert!(matches!(
38 ring.ack(prev, 1),
39 Err(Error::Ack {
40 sequence: 1,
41 received: 0,
42 ..
43 })
44 ));
45
46 for s in 1..=3u8 {
47 assert_eq!(ring.send(next, &[rank as u8, s]).unwrap(), u64::from(s));
48 }
49 world.barrier();
50
51 let messages = ring.poll().unwrap();
52 assert_eq!(
53 messages,
54 (1..=3)
55 .map(|s| Message {
56 origin: prev,
57 sequence: s,
58 data: vec![prev as u8, s as u8],
59 })
60 .collect::<Vec<_>>()
61 );
62 assert!(ring.poll().unwrap().is_empty());
63 assert_eq!(ring.lost(), 0);
64 assert_eq!(ring.max_lag(), 3);
65
66 assert!(matches!(
67 ring.ack(prev, 4),
68 Err(Error::Ack {
69 sequence: 4,
70 received: 3,
71 ..
72 })
73 ));
74 ring.ack(prev, 3).unwrap();
75 ring.ack(prev, 3).unwrap();
76 ring.ack(prev, 2).unwrap();
77
78 world.barrier();
79 ring.close().unwrap();
80}
81
82fn backpressure(world: &SimpleCommunicator, rank: Rank) {
89 let ring = Ring::safe(world, &oneway(2, 8)).unwrap();
90 if rank == 0 {
91 assert_eq!(ring.send(1, &[1]).unwrap(), 1);
92 assert_eq!(ring.send(1, &[2]).unwrap(), 2);
93 }
94 world.barrier();
95
96 if rank == 0 {
97 let sent = thread::scope(|scope| {
98 let sender = ˚
99 let handle = scope.spawn(move || sender.send(1, &[3]));
100 let deadline = Instant::now() + Duration::from_secs(5);
103 while ring.waits() == 0 {
104 assert!(Instant::now() < deadline, "safe sender never blocked");
105 thread::yield_now();
106 }
107 world.barrier();
109 handle.join().expect("sender thread panicked").unwrap()
110 });
111 assert_eq!(sent, 3);
112 assert!(ring.waits() > 0);
113 assert!(ring.wait_ns() > 0);
114 } else {
115 world.barrier();
116 let messages = ring.poll().unwrap();
117 assert_eq!(messages.len(), 2);
118 assert_eq!(messages[0].data, vec![1]);
119 assert_eq!(messages[1].data, vec![2]);
120 assert!(matches!(
121 ring.ack(0, 3),
122 Err(Error::Ack {
123 sequence: 3,
124 received: 2,
125 ..
126 })
127 ));
128 ring.ack(0, 2).unwrap();
129 ring.ack(0, 2).unwrap();
130 }
131
132 world.barrier();
133 if rank == 1 {
134 let messages = ring.poll().unwrap();
135 assert_eq!(messages.len(), 1);
136 assert_eq!(messages[0].sequence, 3);
137 assert_eq!(messages[0].data, vec![3]);
138 ring.ack(0, 3).unwrap();
139 }
140 world.barrier();
141 ring.close().unwrap();
142}
143
144fn oneway_basics(world: &SimpleCommunicator, rank: Rank) {
145 let ring = Ring::safe(world, &oneway(2, 8)).unwrap();
146 if rank == 0 {
147 assert_eq!(ring.send(1, &[7]).unwrap(), 1);
148 }
149 world.barrier();
150 if rank == 1 {
151 let messages = ring.poll().unwrap();
152 assert_eq!(messages.len(), 1);
153 assert_eq!(messages[0].origin, 0);
154 assert_eq!(messages[0].data, vec![7]);
155 ring.ack(0, 1).unwrap();
156 }
157 world.barrier();
158 ring.close().unwrap();
159}
160
161fn raw_basics(world: &SimpleCommunicator, rank: Rank, next: Rank, prev: Rank) {
162 let ring = Ring::raw(world, &pair(2, 8)).unwrap();
163 assert!(!ring.is_safe());
164
165 for s in 1..=3u8 {
166 assert_eq!(ring.send(next, &[rank as u8, s]).unwrap(), u64::from(s));
167 }
168 world.barrier();
169
170 let messages = ring.poll().unwrap();
171 assert_eq!(
172 messages,
173 vec![
174 Message {
175 origin: prev,
176 sequence: 2,
177 data: vec![prev as u8, 2]
178 },
179 Message {
180 origin: prev,
181 sequence: 3,
182 data: vec![prev as u8, 3]
183 },
184 ]
185 );
186 assert_eq!(ring.lost(), 1);
187
188 ring.ack(prev, u64::MAX).unwrap();
190 world.barrier();
191 ring.close().unwrap();
192}
193
194fn main() {
195 let (universe, provided) =
196 mpi::initialize_with_threading(Threading::Multiple).expect("MPI must initialize once");
197 assert_eq!(provided, Threading::Multiple);
198 let world = universe.world();
199 let rank = world.rank();
200 let size = world.size();
201 assert_eq!(size, 2);
202 let next = (rank + 1) % size;
203 let prev = (rank + size - 1) % size;
204
205 if rank == 0 {
207 assert!(matches!(
208 Ring::safe(&world, &pair(1, 8)),
209 Err(Error::Ring("configuration differs between ranks"))
210 ));
211 } else {
212 assert!(matches!(
213 Ring::safe(&world, &pair(2, 8)),
214 Err(Error::Ring("configuration differs between ranks"))
215 ));
216 }
217 assert!(matches!(
218 Ring::raw(&world, &pair(0, 8)),
219 Err(Error::Ring("depth must be positive"))
220 ));
221
222 safe_basics(&world, rank, next, prev);
223 oneway_basics(&world, rank);
224 backpressure(&world, rank);
225 raw_basics(&world, rank, next, prev);
226
227 if rank == 0 {
228 println!("test_ring: ok");
229 }
230}