Skip to main content

shuttle_std/sync/
barrier.rs

1use crate::sync::{ResourceSignature, ResourceType};
2use shuttle_engine::runtime::execution::ExecutionState;
3use shuttle_engine::runtime::task::clock::VectorClock;
4use shuttle_engine::runtime::task::TaskId;
5use shuttle_engine::runtime::thread;
6use std::cell::RefCell;
7use std::collections::HashSet;
8use std::fmt;
9use std::rc::Rc;
10use tracing::trace;
11
12#[derive(Clone, Copy, Debug)]
13/// A `BarrierWaitResult` is returned by `Barrier::wait()` when all threads in the `Barrier` have rendezvoused.
14pub struct BarrierWaitResult {
15    is_leader: bool,
16}
17
18impl BarrierWaitResult {
19    /// Returns true if this thread is the "leader thread" for the call to `Barrier::wait()`.
20    pub fn is_leader(&self) -> bool {
21        self.is_leader
22    }
23}
24
25/// We implement [Barrier] by keeping track of a list of [TaskId]s of threads that have called
26/// `wait`. When the numbers of waiters gets to the barrier's `bound`, they all become unblocked.
27/// Whichever task wakes up first after becoming unblocked will be designated as the "leader" via
28/// the return value of `wait`.
29///
30/// Because barriers can be reused, it does not suffice to designate a single, permanent thread as
31/// the leader in the [BarrierState], since if a different batch of threads waits on the same
32/// barrier, one of those new threads should be the leader of that batch.
33///
34/// For example, if there's a barrier where the bound is 2 and there are 6 threads all concurrently
35/// calling `wait`, then all threads will become unblocked in 3 batches of 2 threads each, with
36/// each batch having its own leader, resulting in 3 (unique) leaders.
37///
38/// We implement this by tracking an `epoch` counter that gets incremented whenever a batch of
39/// threads is released by `wait`. When that happens, we make a "leader token_" available for that
40/// batch. The first thread of a batch to get scheduled after becoming unblocked takes the leader
41/// token (without affecting any threads that are part of a different batch).
42struct BarrierState {
43    /// The number of tasks that must call `wait` before they all get unblocked (and one gets
44    /// chosen as the leader).
45    bound: usize,
46    /// A counter of the number of "batches" of threads that have been released by the barrier,
47    /// needed in order to keep track of the leaders of each batch separately.
48    epoch: u64,
49    /// The set of waiting tasks for the current epoch. Then the size of this set becomes equal to
50    /// the bound, all waiters are unblocked and one of them becomes the leader.
51    waiters: HashSet<TaskId>,
52    /// The set of epochs of this [Barrier] that have reached the number of waiters to be
53    /// unblocked, but haven't had a leader selected yet. The first thread to wake up after wait
54    /// will take the token (by removing the epoch from the set) and become the leader.
55    leader_tokens: HashSet<u64>,
56    clock: VectorClock,
57}
58
59// Implement debug in order to not output the `VectorClock`
60impl fmt::Debug for BarrierState {
61    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
62        f.debug_struct("BarrierState")
63            .field("bound", &self.bound)
64            .field("epoch", &self.epoch)
65            .field("waiters", &self.waiters)
66            .field("leader_tokens", &self.leader_tokens)
67            .finish()
68    }
69}
70
71#[derive(Debug)]
72/// A barrier enables multiple threads to synchronize the beginning of some computation.
73pub struct Barrier {
74    state: Rc<RefCell<BarrierState>>,
75    #[allow(unused)]
76    signature: ResourceSignature,
77}
78
79impl Barrier {
80    /// Creates a new barrier that can block a given number of threads.
81    /// A barrier will block n-1 threads which call `wait()` and then wake up all threads
82    /// at once when the nth thread calls `wait()`.
83    #[track_caller]
84    pub fn new(n: usize) -> Self {
85        let state = BarrierState {
86            bound: n,
87            epoch: 0,
88            waiters: HashSet::new(),
89            leader_tokens: HashSet::new(),
90            clock: VectorClock::new(),
91        };
92
93        Self {
94            state: Rc::new(RefCell::new(state)),
95            signature: ExecutionState::new_resource_signature(ResourceType::Barrier),
96        }
97    }
98
99    /// Blocks the current thread until all threads have rendezvoused here.
100    pub fn wait(&self) -> BarrierWaitResult {
101        let state = self.state.borrow_mut();
102        // The barrier will block if the number of current waiters *plus* an additional waiter
103        // for this thread is less than the bound
104        let will_block = state.waiters.len() + 1 < state.bound;
105        drop(state);
106
107        // If all tasks have already rendezvoused, we need to context switch once to allow the
108        // previous event to become visible before the epoch changes. Otherwise, we can omit the
109        // scheduling point if the wait commutes with other blocking waits (double-yield optimization,
110        // reasoning below).
111        //
112        // Blocking waits Y1 and Z on threads T1 and T2 always commute with each other because
113        // the waiters for a barrier are represented by an unordered set. Thus for both orderings
114        // `Y1 Z` and `Z Y1`, the state of the barrier is {T1, T2}. As a result, we never need to
115        // switch before blocking on a barrier wait.
116        if !will_block {
117            thread::switch();
118        }
119        let mut state = self.state.borrow_mut();
120        let my_epoch = state.epoch;
121
122        trace!(waiters=?state.waiters, epoch=my_epoch, "waiting on barrier {:p}", self);
123
124        // Update the barrier's clock with the clock of this thread
125        ExecutionState::with(|s| {
126            let clock = s.increment_clock();
127            state.clock.update(clock);
128        });
129
130        // Add the current thread to `waiters`. It shouldn't already be present.
131        let me = ExecutionState::me();
132        assert!(state.waiters.insert(me));
133
134        if state.waiters.len() < state.bound {
135            trace!(waiters=?state.waiters, epoch=my_epoch, "blocked on barrier {:?}", self);
136            drop(state);
137
138            // If execution teardown unwinds the task from the switch below (see
139            // `ExecutionState::tear_down`), the task is no longer waiting. Its destructors run as the
140            // task, and may wait on this barrier.
141            struct StopWaitingOnUnwind<'a>(&'a Barrier, TaskId);
142            impl Drop for StopWaitingOnUnwind<'_> {
143                fn drop(&mut self) {
144                    self.0.state.borrow_mut().waiters.remove(&self.1);
145                }
146            }
147            let stop_waiting_on_unwind = StopWaitingOnUnwind(self, me);
148
149            ExecutionState::with(|s| s.current_mut().block(false));
150            thread::switch();
151            std::mem::forget(stop_waiting_on_unwind);
152        } else {
153            trace!(waiters=?state.waiters, epoch=my_epoch, "releasing waiters on barrier {:?}", self);
154
155            debug_assert!(state.waiters.len() == state.bound || state.bound == 0);
156
157            // Make the leader token available for this epoch. The first task to wake up will
158            // take it and become the leader. The token shouldn't already be available.
159            assert!(state.leader_tokens.insert(my_epoch));
160
161            // Drain the set of waiters and increment the barrier's epoch, so any other task that
162            // calls `wait` from now on becomes part of a separate group with its own leader.
163            let waiters = state.waiters.drain().collect::<Vec<_>>();
164            state.epoch += 1;
165
166            trace!(
167                waiters=?state.waiters,
168                epoch=state.epoch,
169                "releasing waiters on barrier {:?}",
170                self,
171            );
172
173            let clock = state.clock.clone();
174            ExecutionState::with(|s| {
175                // `waiters` includes the current task.
176                for tid in waiters {
177                    let t = s.get_mut(tid);
178                    t.clock.increment(tid);
179                    t.clock.update(&clock);
180                    t.unblock();
181                }
182            });
183            drop(state);
184        };
185
186        // Try to remove the leader token for this epoch. If true, then the token was present and
187        // we are the leader. Any future attempts to remove the token will return false.
188        let is_leader = self.state.borrow_mut().leader_tokens.remove(&my_epoch);
189
190        trace!(epoch=?my_epoch, is_leader, "returning from barrier {:?}", self);
191
192        BarrierWaitResult { is_leader }
193    }
194}
195
196// Safety: Barrier is never actually passed across threads, only across continuations. The
197// Rc<RefCell<_>> type therefore can't be preempted mid-bookkeeping-operation.
198// TODO we shouldn't need to do this, but RefCell is not Send, and Barrier needs to be Send.
199unsafe impl Send for Barrier {}
200unsafe impl Sync for Barrier {}
201
202#[cfg(test)]
203mod tests {
204    use super::*;
205
206    #[test]
207    fn unique_resource_signature_barrier() {
208        shuttle_schedulers::check_random(
209            || {
210                let barrier1 = Barrier::new(2);
211                let barrier2 = Barrier::new(2);
212                assert_ne!(barrier1.signature, barrier2.signature);
213            },
214            1,
215        );
216    }
217}