Skip to main content

shuttle_std/sync/
condvar.rs

1use crate::sync::{MutexGuard, ResourceSignature, ResourceType};
2use assoc::AssocExt;
3use shuttle_engine::current;
4use shuttle_engine::runtime::execution::ExecutionState;
5use shuttle_engine::runtime::task::clock::VectorClock;
6use shuttle_engine::runtime::task::TaskId;
7use shuttle_engine::runtime::thread;
8use std::cell::RefCell;
9use std::collections::VecDeque;
10use std::sync::{LockResult, PoisonError};
11use std::time::Duration;
12use tracing::trace;
13
14/// A `Condvar` represents the ability to block a thread such that it consumes no CPU time while
15/// waiting for an event to occur.
16#[derive(Debug)]
17pub struct Condvar {
18    state: RefCell<CondvarState>,
19    #[allow(unused)]
20    signature: ResourceSignature,
21}
22
23#[derive(Debug)]
24struct CondvarState {
25    // TODO: this should be a HashMap but [HashMap::new] is not const
26    waiters: Vec<(TaskId, CondvarWaitStatus)>,
27    next_epoch: usize,
28}
29
30// For tracking causal dependencies, we record the clock C of the thread that does the notify.
31// When a thread is unblocked, its clock is updated by C.
32#[derive(PartialEq, Eq, Debug)]
33enum CondvarWaitStatus {
34    Waiting,
35    // invariant: VecDeque is non-empty (if it's empty, we should be Waiting instead)
36    Signal(VecDeque<(usize, VectorClock)>),
37    Broadcast(VectorClock),
38}
39
40// TODO Check if we can avoid using epochs now that we have vector clocks.
41// TODO See [Issue 39](https://github.com/awslabs/shuttle/issues/39)
42
43// We implement `Condvar` by tracking the `CondvarWaitStatus` of each thread currently waiting on
44// the `Condvar`.
45//
46// ## Terminology
47//
48// There's competing notions of "unblocked" here -- unblocked *from the condition variable*, which
49// is the API-level notion of blocked, and unblocked *within Shuttle's scheduler*, which is an
50// implementation detail. To disambiguate, we'll call the latter "runnable" or "unrunnable", even
51// though that's a little awkward.
52//
53// ## `notify_one`
54//
55// A `notify_one` unblocks *one* currently blocked thread. We want the scheduler to be able to
56// choose which thread that is, and so we implement `notify_one` by marking all waiters as runnable.
57// Whichever waiter wins the race by running first will mark all other waiters as unrunnable again,
58//
59// This gets a little hairy if there are racing `wait`ers and `notify_one`s. The scenario we're
60// concerned about is this:
61//
62//          Thread 1       | Thread 2       | Thread 3       | Thread 4       | Thread 5
63//          ---------------|----------------|----------------|----------------| ----------------
64//     (1)   wait()        |                |                |                |
65//     (2)                 |  wait()        |                |                |
66//     (3)                 |                |  notify_one()  |                |
67//     (4)                 |                |                |  wait()        |
68//     (5)                 |                |                |                |  wait()
69//     (6)                 |                |  notify_one()  |                |
70//     (7)                 |                |                |  wake          |
71//
72// Here, after (6), all 4 waiter threads are runnable. Thread 4 wins the race to run first at (7),
73// and so is chosen as the unblocked thread. After (7), Thread 5 needs to be made unrunnable,
74// because the only signal it can see was (6), which has already unblocked Thread 4. However,
75// Threads 1 and 2 need to remain runnable, because they are eligible to be unblocked by (3), which
76// has not yet been consumed.
77//
78// We solve this problem with "epochs". Each `notify_one` is associated with a unique epoch, and
79// each waiter in the `CondvarState` remembers a list of the epochs that occurred while it was
80// blocked. When a waiter runs after being made runnable by a `notify_one`, it checks to see which
81// epoch that notify was associated with, and removes that epoch from the lists of every other
82// waiter. Any waiter that still has a non-empty list of epochs should remain runnable, because
83// there are still signals it's eligible to receive. Any waiter with an empty list of epochs is made
84// unrunnable, because all the signals it was present for have been consumed.
85//
86// In the scenario above, there are two epochs (3) and (6). At (7), Threads 1 and 2 have the same
87// epoch list [0, 1], and Threads 4 and 5 have the same epoch list [1]. When Thread 4 wins the race,
88// it sees that it was woken by epoch 1, and so removes that epoch from all other waiter lists.
89// Thread 1 and 2 still have epoch 0 in their list, so they remain runnable. Thread 5 has no more
90// epochs in its list, so it is made unrunnable. Whichever of Threads 1 and 2 wins the subsequent
91// race (not shown) will observe that it was woken by epoch 0, remove epoch 0 from whichever of the
92// two threads lost the race, and then make that thread unrunnable again because there are no
93// signals remaining for it to observe.
94//
95// ## `notify_all`
96//
97// `notify_all` is a broadcast that unblocks *all* currently blocked threads. Once a `notify_all`
98// occurs, all the state discussed above is irrelevant -- every thread should be unblocked, so we
99// don't need to remember whether there is also a `notify_one` that could have unblocked them.
100// Waiters that arrive after the `notify_all` will be blocked as usual.
101//
102// `notify_all` atomically unblocks all currently blocked threads. For example, consider:
103//
104//          Thread 1       | Thread 2       | Thread 3       | Thread 4
105//          ---------------|----------------|----------------|----------------
106//     (1)   wait()        |                |                |
107//     (2)                 |  wait()        |                |
108//     (3)                 |                |  notify_all()  |
109//     (4)                 |                |                |  wait()
110//     (5)                 |                |  notify_one()  |
111//
112// After (3), Threads 1 and 2 are considered unblocked, even though the scheduler has not yet run
113// them again, Thread 4 becomes blocked at (4). At (5), the only blocked thread is Thread 4, because
114// the others were unblocked by (3), even though they have not yet woken up to discover this fact.
115// In other words, this execution cannot deadlock -- if (4) happens-before (5), then Thread 4 is
116// guaranteed to be the thread unblocked by (5). After (5), Threads 1, 2, and 4 are all runnable,
117// and can run in any order (because they are all contending on the same mutex).
118impl Condvar {
119    /// Creates a new condition variable which is ready to be waited on and notified.
120    #[track_caller]
121    pub const fn new() -> Self {
122        let state = CondvarState {
123            waiters: Vec::new(),
124            next_epoch: 0,
125        };
126
127        Self {
128            state: RefCell::new(state),
129            signature: ResourceSignature::new_const(ResourceType::Condvar),
130        }
131    }
132
133    /// Blocks the current thread until this condition variable receives a notification.
134    pub fn wait<'a, T>(&self, guard: MutexGuard<'a, T>) -> LockResult<MutexGuard<'a, T>> {
135        let me = ExecutionState::me();
136
137        // Release the lock, which allows for a switch *before* unlocking, but not after
138        // This is because the MutexGuard internally calls `batch_semaphore::release` when it
139        // unlocks via it's drop handler. `release` itself is a visible operation, so it provides
140        // it's own scheduling point prior to releasing the mutex. As all scheduling points are
141        // *only before* visible operations, anything done after this point is not visible until
142        // we switch ourselves
143        let mutex = guard.unlock();
144        // Unlocked, but no other task has run yet. We thus block ourselves and switch
145        let mut state = self.state.borrow_mut();
146
147        trace!(waiters=?state.waiters, next_epoch=state.next_epoch, "waiting on condvar {:p}", self);
148
149        debug_assert!(<_ as AssocExt<_, _>>::get(&state.waiters, &me).is_none());
150        state.waiters.push((me, CondvarWaitStatus::Waiting));
151        drop(state);
152
153        // TODO: Condvar::wait should allow for spurious wakeups.
154        ExecutionState::with(|s| s.current_mut().block(false));
155        thread::switch();
156
157        // After the context switch, consume whichever signal that woke this thread
158        let mut state = self.state.borrow_mut();
159        trace!(waiters=?state.waiters, next_epoch=state.next_epoch, "woken from condvar {:p}", self);
160        let my_status = <_ as AssocExt<_, _>>::remove(&mut state.waiters, &me).expect("should be waiting");
161        match my_status {
162            CondvarWaitStatus::Broadcast(clock) => {
163                // Woken by a broadcast, so nothing to do except update the clock
164                ExecutionState::with(|s| s.update_clock(&clock));
165            }
166            CondvarWaitStatus::Signal(mut epochs) => {
167                let (epoch, clock) = epochs.pop_front().expect("should be a pending signal");
168                // No other waiter is allowed to be unblocked by the epoch that woke us
169                for (tid, status) in state.waiters.iter_mut() {
170                    if let CondvarWaitStatus::Signal(epochs) = status {
171                        if let Some(i) = epochs.iter().position(|e| epoch == e.0) {
172                            epochs.remove(i);
173                            if epochs.is_empty() {
174                                *status = CondvarWaitStatus::Waiting;
175                                // Make the task unrunnable if there are no pending signals that
176                                // could unblock it
177                                // TODO: Condvar::wait should allow for spurious wakeups.
178                                ExecutionState::with(|s| s.get_mut(*tid).block(false));
179                            }
180                        }
181                    }
182                }
183                // Update the thread's clock with the clock from the notifier
184                ExecutionState::with(|s| s.update_clock(&clock));
185            }
186            CondvarWaitStatus::Waiting => panic!("should not have been woken while in Waiting status"),
187        }
188        drop(state);
189
190        // Reacquire the lock
191        // TODO The context switch involved here might be redundant? The scheduler implicitly chose
192        // TODO this thread to win the lock when it ran us after the context switch above.
193        mutex.lock()
194    }
195
196    /// Blocks the current thread until this condition variable receives a notification and the
197    /// provided condition is false.
198    pub fn wait_while<'a, T, F>(&self, mut guard: MutexGuard<'a, T>, mut condition: F) -> LockResult<MutexGuard<'a, T>>
199    where
200        F: FnMut(&mut T) -> bool,
201    {
202        while condition(&mut *guard) {
203            guard = self.wait(guard)?;
204        }
205        Ok(guard)
206    }
207
208    /// Waits on this condition variable for a notification, timing out after a specified duration.
209    pub fn wait_timeout<'a, T>(
210        &self,
211        guard: MutexGuard<'a, T>,
212        _dur: Duration,
213    ) -> LockResult<(MutexGuard<'a, T>, WaitTimeoutResult)> {
214        // TODO support the timeout case -- this method never times out
215        self.wait(guard)
216            .map(|guard| (guard, WaitTimeoutResult(false)))
217            .map_err(|e| PoisonError::new((e.into_inner(), WaitTimeoutResult(false))))
218    }
219
220    /// Waits on this condition variable for a notification, timing out after a specified duration.
221    ///
222    /// The semantics of this function are equivalent to [`wait_while`](Self::wait_while) except
223    /// that the thread will be blocked for roughly no longer than `dur`.
224    pub fn wait_timeout_while<'a, T, F>(
225        &self,
226        guard: MutexGuard<'a, T>,
227        _dur: Duration,
228        condition: F,
229    ) -> LockResult<(MutexGuard<'a, T>, WaitTimeoutResult)>
230    where
231        F: FnMut(&mut T) -> bool,
232    {
233        // TODO support the timeout case -- this method never times out
234        self.wait_while(guard, condition)
235            .map(|guard| (guard, WaitTimeoutResult(false)))
236            .map_err(|e| PoisonError::new((e.into_inner(), WaitTimeoutResult(false))))
237    }
238
239    /// Wakes up one blocked thread on this condvar.
240    ///
241    /// If there is a blocked thread on this condition variable, then it will be woken up from its
242    /// call to wait or wait_timeout. Calls to notify_one are not buffered in any way.
243    pub fn notify_one(&self) {
244        thread::switch();
245
246        let me = ExecutionState::me();
247
248        let mut state = self.state.borrow_mut();
249
250        trace!(waiters=?state.waiters, next_epoch=state.next_epoch, "notifying one on condvar {:p}", self);
251
252        let epoch = state.next_epoch;
253        for (tid, status) in state.waiters.iter_mut() {
254            assert_ne!(*tid, me);
255
256            let clock = current::clock();
257            match status {
258                CondvarWaitStatus::Waiting => {
259                    let mut epochs = VecDeque::new();
260                    epochs.push_back((epoch, clock));
261                    *status = CondvarWaitStatus::Signal(epochs);
262                }
263                CondvarWaitStatus::Signal(epochs) => {
264                    epochs.push_back((epoch, clock));
265                }
266                CondvarWaitStatus::Broadcast(_) => {
267                    // no-op, broadcast will already unblock this task
268                }
269            }
270
271            // Note: the task might have been unblocked by a previous signal
272            ExecutionState::with(|s| s.get_mut(*tid).unblock());
273        }
274        state.next_epoch += 1;
275
276        drop(state);
277    }
278
279    /// Wakes up all blocked threads on this condvar.
280    pub fn notify_all(&self) {
281        thread::switch();
282
283        let me = ExecutionState::me();
284
285        let mut state = self.state.borrow_mut();
286
287        trace!(waiters=?state.waiters, next_epoch=state.next_epoch, "notifying all on condvar {:p}", self);
288
289        for (tid, status) in state.waiters.iter_mut() {
290            assert_ne!(*tid, me);
291            *status = CondvarWaitStatus::Broadcast(current::clock());
292            // Note: the task might have been unblocked by a previous signal
293            ExecutionState::with(|s| s.get_mut(*tid).unblock());
294        }
295
296        drop(state);
297    }
298}
299
300// Safety: Condvar is never actually passed across true threads, only across continuations. The
301// Rc<RefCell<_>> type therefore can't be preempted mid-bookkeeping-operation.
302// TODO we shouldn't need to do this, but RefCell is not Send
303unsafe impl Send for Condvar {}
304unsafe impl Sync for Condvar {}
305
306impl Default for Condvar {
307    fn default() -> Self {
308        Self::new()
309    }
310}
311
312/// A type indicating whether a timed wait on a condition variable returned due to a time out or not.
313#[derive(Debug, PartialEq, Eq, Copy, Clone)]
314pub struct WaitTimeoutResult(bool);
315
316impl WaitTimeoutResult {
317    /// Returns `true` if the wait was known to have timed out.
318    pub fn timed_out(&self) -> bool {
319        self.0
320    }
321}
322
323#[cfg(test)]
324mod tests {
325    use super::*;
326
327    #[test]
328    fn unique_resource_signature_condvar() {
329        shuttle_schedulers::check_random(
330            || {
331                let condvar1 = Condvar::new();
332                let condvar2 = Condvar::new();
333                assert_ne!(condvar1.signature, condvar2.signature);
334            },
335            1,
336        );
337    }
338}