1use crate::execution_api::{ForkJoinPhase, WorkerState as ForkJoinWorkerState};
4use vstd::prelude::*;
5
6verus! {
7
8pub open spec fn worker_at(
10 workers: Seq<ForkJoinWorkerState>,
11 worker: int,
12) -> Option<ForkJoinWorkerState> {
13 if 0 <= worker < workers.len() {
14 Some(workers[worker])
15 } else {
16 None
17 }
18}
19
20pub open spec fn start_worker_transition(
22 before: Option<ForkJoinWorkerState>,
23 after: Option<ForkJoinWorkerState>,
24 phase: ForkJoinPhase,
25 selected: bool,
26 accepted: bool,
27) -> bool {
28 let enabled = selected
29 && phase == ForkJoinPhase::Fork
30 && before == Some(ForkJoinWorkerState::Ready);
31 &&& accepted == enabled
32 &&& after == if accepted {
33 Some(ForkJoinWorkerState::Running)
34 } else {
35 before
36 }
37}
38
39pub open spec fn complete_worker_transition(
41 before: Option<ForkJoinWorkerState>,
42 after: Option<ForkJoinWorkerState>,
43 phase: ForkJoinPhase,
44 selected: bool,
45 value_admitted: bool,
46 accepted: bool,
47) -> bool {
48 let enabled = selected
49 && value_admitted
50 && phase == ForkJoinPhase::Fork
51 && before == Some(ForkJoinWorkerState::Running);
52 &&& accepted == enabled
53 &&& after == if accepted {
54 Some(ForkJoinWorkerState::Complete)
55 } else {
56 before
57 }
58}
59
60pub open spec fn start_worker_action(
62 before: Seq<ForkJoinWorkerState>,
63 after: Seq<ForkJoinWorkerState>,
64 phase: ForkJoinPhase,
65 worker: int,
66 selected: bool,
67 accepted: bool,
68) -> bool {
69 &&& start_worker_transition(
70 worker_at(before, worker),
71 worker_at(after, worker),
72 phase,
73 selected,
74 accepted,
75 )
76 &&& after == if accepted {
77 before.update(worker, ForkJoinWorkerState::Running)
78 } else {
79 before
80 }
81}
82
83pub open spec fn complete_worker_action(
85 before: Seq<ForkJoinWorkerState>,
86 after: Seq<ForkJoinWorkerState>,
87 phase: ForkJoinPhase,
88 worker: int,
89 selected: bool,
90 value_admitted: bool,
91 accepted: bool,
92) -> bool {
93 &&& complete_worker_transition(
94 worker_at(before, worker),
95 worker_at(after, worker),
96 phase,
97 selected,
98 value_admitted,
99 accepted,
100 )
101 &&& after == if accepted {
102 before.update(worker, ForkJoinWorkerState::Complete)
103 } else {
104 before
105 }
106}
107
108pub open spec fn barrier_action(
110 before: ForkJoinPhase,
111 after: ForkJoinPhase,
112 selected: bool,
113 completion_observed: bool,
114 accepted: bool,
115) -> bool {
116 let enabled = selected && before == ForkJoinPhase::Fork && completion_observed;
117 &&& accepted == enabled
118 &&& after == if accepted { ForkJoinPhase::Join } else { before }
119}
120
121pub open spec fn produce_output_action(
123 before: ForkJoinPhase,
124 after: ForkJoinPhase,
125 selected: bool,
126 accepted: bool,
127) -> bool {
128 let enabled = selected && before == ForkJoinPhase::Join;
129 &&& accepted == enabled
130 &&& after == if accepted { ForkJoinPhase::Done } else { before }
131}
132
133pub struct ForkJoin {
135 pub value_domain_size: u64,
137 pub wstate: Vec<ForkJoinWorkerState>,
139 pub wvalue: Vec<u64>,
141 pub phase: ForkJoinPhase,
143 pub output_ready: bool,
145 pub output_snapshot: Vec<u64>,
147}
148
149impl ForkJoin {
150 pub open spec fn all_complete(&self) -> bool {
152 forall|i: int| 0 <= i < self.wstate@.len()
153 ==> #[trigger] self.wstate@[i] == ForkJoinWorkerState::Complete
154 }
155
156 pub open spec fn type_invariant(&self) -> bool {
158 &&& self.value_domain_size > 0
159 &&& self.wstate@.len() == self.wvalue@.len()
160 &&& self.output_snapshot@.len() == self.wvalue@.len()
161 &&& (forall|i: int| 0 <= i < self.wvalue@.len()
162 ==> #[trigger] self.wvalue@[i] < self.value_domain_size)
163 &&& (forall|i: int| 0 <= i < self.output_snapshot@.len()
164 ==> #[trigger] self.output_snapshot@[i] < self.value_domain_size)
165 }
166
167 pub open spec fn barrier_completeness(&self) -> bool {
169 self.output_ready ==> self.all_complete()
170 }
171
172 pub open spec fn output_consistency(&self) -> bool {
174 self.output_ready ==> self.output_snapshot@ == self.wvalue@
175 }
176
177 pub open spec fn phase_ordering(&self) -> bool {
179 &&& (self.phase == ForkJoinPhase::Join ==> self.all_complete())
180 &&& (self.phase == ForkJoinPhase::Done ==> self.output_ready)
181 }
182
183 pub open spec fn ready_only_done(&self) -> bool {
185 self.output_ready ==> self.phase == ForkJoinPhase::Done
186 }
187
188 pub open spec fn inv(&self) -> bool {
190 &&& self.type_invariant()
191 &&& self.barrier_completeness()
192 &&& self.output_consistency()
193 &&& self.phase_ordering()
194 &&& self.ready_only_done()
195 }
196
197 #[expect(clippy::arithmetic_side_effects, reason = "Verus proves the construction cursor remains within the worker bound")]
198 pub fn new(workers: usize, value_domain_size: u64, initial_value: u64) -> (s: ForkJoin)
200 requires
201 value_domain_size > 0,
202 initial_value < value_domain_size,
203 ensures
204 s.value_domain_size == value_domain_size,
205 s.wstate@.len() == workers,
206 s.wvalue@.len() == workers,
207 s.output_snapshot@.len() == workers,
208 forall|i: int| 0 <= i < workers ==>
209 s.wstate@[i] == ForkJoinWorkerState::Ready,
210 forall|i: int| 0 <= i < workers ==> s.wvalue@[i] == initial_value,
211 forall|i: int| 0 <= i < workers ==> s.output_snapshot@[i] == initial_value,
212 s.phase == ForkJoinPhase::Fork,
213 !s.output_ready,
214 s.inv(),
215 {
216 let mut wstate = Vec::new();
217 let mut wvalue = Vec::new();
218 let mut snapshot = Vec::new();
219 let mut i = 0;
220 while i < workers
221 invariant
222 i <= workers,
223 wstate@.len() == i,
224 wvalue@.len() == i,
225 snapshot@.len() == i,
226 forall|k: int| 0 <= k < i ==>
227 wstate@[k] == ForkJoinWorkerState::Ready,
228 forall|k: int| 0 <= k < i ==> wvalue@[k] == initial_value,
229 forall|k: int| 0 <= k < i ==> snapshot@[k] == initial_value,
230 decreases workers - i,
231 {
232 wstate.push(ForkJoinWorkerState::Ready);
233 wvalue.push(initial_value);
234 snapshot.push(initial_value);
235 i += 1;
236 }
237 ForkJoin {
238 value_domain_size,
239 wstate,
240 wvalue,
241 phase: ForkJoinPhase::Fork,
242 output_ready: false,
243 output_snapshot: snapshot,
244 }
245 }
246
247 #[expect(clippy::indexing_slicing, reason = "Verus proves the completeness cursor remains in bounds")]
248 #[expect(clippy::arithmetic_side_effects, reason = "Verus proves the completeness cursor increment remains in bounds")]
249 fn all_complete_exec(&self) -> (b: bool)
250 ensures b == self.all_complete(),
251 {
252 let mut i = 0;
253 while i < self.wstate.len()
254 invariant
255 i <= self.wstate.len(),
256 forall|k: int| 0 <= k < i ==>
257 self.wstate@[k] == ForkJoinWorkerState::Complete,
258 decreases self.wstate.len() - i,
259 {
260 if !matches!(self.wstate[i], ForkJoinWorkerState::Complete) {
261 assert(!self.all_complete()) by {
262 assert(self.wstate@[i as int] != ForkJoinWorkerState::Complete);
263 }
264 return false;
265 }
266 i += 1;
267 }
268 true
269 }
270
271 #[expect(clippy::indexing_slicing, reason = "the action guard and Verus invariant bound the worker index")]
272 pub fn start_worker(&mut self, worker: usize) -> (accepted: bool)
274 requires old(self).inv(),
275 ensures
276 final(self).value_domain_size == old(self).value_domain_size,
277 final(self).wvalue@ == old(self).wvalue@,
278 final(self).phase == old(self).phase,
279 final(self).output_ready == old(self).output_ready,
280 final(self).output_snapshot@ == old(self).output_snapshot@,
281 start_worker_action(
282 old(self).wstate@,
283 final(self).wstate@,
284 old(self).phase,
285 worker as int,
286 true,
287 accepted,
288 ),
289 final(self).inv(),
290 {
291 if worker < self.wstate.len()
292 && matches!(self.phase, ForkJoinPhase::Fork)
293 && matches!(self.wstate[worker], ForkJoinWorkerState::Ready)
294 {
295 assert(!self.output_ready);
296 self.wstate.set(worker, ForkJoinWorkerState::Running);
297 assert(!self.all_complete()) by {
298 assert(self.wstate@[worker as int] == ForkJoinWorkerState::Running);
299 }
300 true
301 } else {
302 false
303 }
304 }
305
306 #[expect(clippy::indexing_slicing, reason = "the action guard and Verus invariant bound the worker index")]
307 pub fn complete_worker(&mut self, worker: usize, value: u64) -> (accepted: bool)
309 requires old(self).inv(),
310 ensures
311 final(self).value_domain_size == old(self).value_domain_size,
312 final(self).phase == old(self).phase,
313 final(self).output_ready == old(self).output_ready,
314 final(self).output_snapshot@ == old(self).output_snapshot@,
315 complete_worker_action(
316 old(self).wstate@,
317 final(self).wstate@,
318 old(self).phase,
319 worker as int,
320 true,
321 value < old(self).value_domain_size,
322 accepted,
323 ),
324 final(self).wvalue@ == if accepted {
325 old(self).wvalue@.update(worker as int, value)
326 } else { old(self).wvalue@ },
327 final(self).inv(),
328 {
329 if worker < self.wstate.len()
330 && matches!(self.phase, ForkJoinPhase::Fork)
331 && matches!(self.wstate[worker], ForkJoinWorkerState::Running)
332 && value < self.value_domain_size
333 {
334 assert(!self.output_ready);
335 self.wvalue.set(worker, value);
336 self.wstate.set(worker, ForkJoinWorkerState::Complete);
337 assert forall|i: int| 0 <= i < self.wvalue@.len()
338 implies #[trigger] self.wvalue@[i] < self.value_domain_size by {
339 if i != worker as int {
340 assert(self.wvalue@[i] == old(self).wvalue@[i]);
341 }
342 }
343 true
344 } else {
345 false
346 }
347 }
348
349 pub fn barrier(&mut self) -> (accepted: bool)
351 requires old(self).inv(),
352 ensures
353 final(self).value_domain_size == old(self).value_domain_size,
354 final(self).wstate@ == old(self).wstate@,
355 final(self).wvalue@ == old(self).wvalue@,
356 final(self).output_ready == old(self).output_ready,
357 final(self).output_snapshot@ == old(self).output_snapshot@,
358 barrier_action(
359 old(self).phase,
360 final(self).phase,
361 true,
362 old(self).all_complete(),
363 accepted,
364 ),
365 final(self).inv(),
366 {
367 if matches!(self.phase, ForkJoinPhase::Fork) && self.all_complete_exec() {
368 assert(!self.output_ready);
369 self.phase = ForkJoinPhase::Join;
370 true
371 } else {
372 false
373 }
374 }
375
376 pub fn produce_output(&mut self) -> (accepted: bool)
378 requires old(self).inv(),
379 ensures
380 final(self).value_domain_size == old(self).value_domain_size,
381 final(self).wstate@ == old(self).wstate@,
382 final(self).wvalue@ == old(self).wvalue@,
383 if accepted {
384 final(self).phase == ForkJoinPhase::Done
385 && final(self).output_ready
386 && final(self).output_snapshot@ == old(self).wvalue@
387 } else {
388 final(self).phase == old(self).phase
389 && final(self).output_ready == old(self).output_ready
390 && final(self).output_snapshot@ == old(self).output_snapshot@
391 },
392 produce_output_action(
393 old(self).phase,
394 final(self).phase,
395 true,
396 accepted,
397 ),
398 final(self).inv(),
399 {
400 if matches!(self.phase, ForkJoinPhase::Join) {
401 self.output_snapshot = self.wvalue.clone();
402 self.output_ready = true;
403 self.phase = ForkJoinPhase::Done;
404 true
405 } else {
406 false
407 }
408 }
409
410 pub fn done_stuttering(&mut self) -> (enabled: bool)
412 requires old(self).inv(),
413 ensures
414 enabled == (old(self).phase == ForkJoinPhase::Done),
415 final(self).value_domain_size == old(self).value_domain_size,
416 final(self).wstate@ == old(self).wstate@,
417 final(self).wvalue@ == old(self).wvalue@,
418 final(self).phase == old(self).phase,
419 final(self).output_ready == old(self).output_ready,
420 final(self).output_snapshot@ == old(self).output_snapshot@,
421 final(self).inv(),
422 {
423 matches!(self.phase, ForkJoinPhase::Done)
424 }
425}
426
427}