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        assert!(state.waiters.insert(ExecutionState::me()));
132
133        if state.waiters.len() < state.bound {
134            trace!(waiters=?state.waiters, epoch=my_epoch, "blocked on barrier {:?}", self);
135            drop(state);
136            ExecutionState::with(|s| s.current_mut().block(false));
137            thread::switch();
138        } else {
139            trace!(waiters=?state.waiters, epoch=my_epoch, "releasing waiters on barrier {:?}", self);
140
141            debug_assert!(state.waiters.len() == state.bound || state.bound == 0);
142
143            // Make the leader token available for this epoch. The first task to wake up will
144            // take it and become the leader. The token shouldn't already be available.
145            assert!(state.leader_tokens.insert(my_epoch));
146
147            // Drain the set of waiters and increment the barrier's epoch, so any other task that
148            // calls `wait` from now on becomes part of a separate group with its own leader.
149            let waiters = state.waiters.drain().collect::<Vec<_>>();
150            state.epoch += 1;
151
152            trace!(
153                waiters=?state.waiters,
154                epoch=state.epoch,
155                "releasing waiters on barrier {:?}",
156                self,
157            );
158
159            let clock = state.clock.clone();
160            ExecutionState::with(|s| {
161                // `waiters` includes the current task.
162                for tid in waiters {
163                    let t = s.get_mut(tid);
164                    t.clock.increment(tid);
165                    t.clock.update(&clock);
166                    t.unblock();
167                }
168            });
169            drop(state);
170        };
171
172        // Try to remove the leader token for this epoch. If true, then the token was present and
173        // we are the leader. Any future attempts to remove the token will return false.
174        let is_leader = self.state.borrow_mut().leader_tokens.remove(&my_epoch);
175
176        trace!(epoch=?my_epoch, is_leader, "returning from barrier {:?}", self);
177
178        BarrierWaitResult { is_leader }
179    }
180}
181
182// Safety: Barrier is never actually passed across threads, only across continuations. The
183// Rc<RefCell<_>> type therefore can't be preempted mid-bookkeeping-operation.
184// TODO we shouldn't need to do this, but RefCell is not Send, and Barrier needs to be Send.
185unsafe impl Send for Barrier {}
186unsafe impl Sync for Barrier {}
187
188#[cfg(test)]
189mod tests {
190    use super::*;
191
192    #[test]
193    fn unique_resource_signature_barrier() {
194        shuttle_schedulers::check_random(
195            || {
196                let barrier1 = Barrier::new(2);
197                let barrier2 = Barrier::new(2);
198                assert_ne!(barrier1.signature, barrier2.signature);
199            },
200            1,
201        );
202    }
203}