Skip to main content

automation_structures/modalities/
fork_join.rs

1//! Finite-index ForkJoin execution carrier.
2
3use crate::execution_api::{ForkJoinPhase, WorkerState as ForkJoinWorkerState};
4use vstd::prelude::*;
5
6verus! {
7
8/// Observe one worker without inventing a value for an invalid ordinal.
9pub 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
20/// Transition for starting one observed ForkJoin worker.
21pub 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
39/// Transition for completing one observed ForkJoin worker.
40pub 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
60/// Worker-start action over any faithful ForkJoin state carrier.
61pub 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
83/// Worker-completion action over any faithful ForkJoin state carrier.
84pub 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
108/// Barrier phase transition.
109pub 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
121/// Output-publication phase transition.
122pub 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
133/// Fork-join execution owner.
134pub struct ForkJoin {
135    /// Exclusive upper bound of worker values.
136    pub value_domain_size: u64,
137    /// Worker lifecycle states by worker index.
138    pub wstate: Vec<ForkJoinWorkerState>,
139    /// Current values by worker index.
140    pub wvalue: Vec<u64>,
141    /// Global fork-join phase.
142    pub phase: ForkJoinPhase,
143    /// Whether a stable output snapshot has been produced.
144    pub output_ready: bool,
145    /// Stable joined output values.
146    pub output_snapshot: Vec<u64>,
147}
148
149impl ForkJoin {
150    /// Whether every worker has completed.
151    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    /// Whether worker, phase, and output storage has valid shape and values.
157    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    /// Whether entering the join phase requires every worker to be complete.
168    pub open spec fn barrier_completeness(&self) -> bool {
169        self.output_ready ==> self.all_complete()
170    }
171
172    /// Whether a published output exactly reflects completed worker values.
173    pub open spec fn output_consistency(&self) -> bool {
174        self.output_ready ==> self.output_snapshot@ == self.wvalue@
175    }
176
177    /// Whether phase progress follows fork, join, and completion order.
178    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    /// Whether output readiness occurs only in the terminal phase.
184    pub open spec fn ready_only_done(&self) -> bool {
185        self.output_ready ==> self.phase == ForkJoinPhase::Done
186    }
187
188    /// Whether all fork-join obligations hold.
189    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    /// Construct a fork-phase execution with ready workers.
199    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    /// Start one ready worker, returning false when the transition is disabled.
273    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    /// Complete one running worker with an in-domain value.
308    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    /// Commit the join barrier after every worker completes.
350    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    /// Produce the stable output snapshot from joined worker values.
377    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    /// Execute the terminal stutter when output production is complete.
411    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}