Skip to main content

moirai_sync/sync/
wait_group.rs

1use std::fmt;
2use std::sync::{Condvar, Mutex};
3
4/// A wait group for synchronizing multiple threads (Go-inspired).
5/// This provides value beyond standard library primitives.
6pub struct WaitGroup {
7    state: Mutex<WaitGroupState>,
8    cond: Condvar,
9}
10
11struct WaitGroupState {
12    counter: u64,
13}
14
15impl fmt::Debug for WaitGroup {
16    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
17        let state = self.state.lock().unwrap();
18        f.debug_struct("WaitGroup")
19            .field("counter", &state.counter)
20            .finish()
21    }
22}
23
24impl Default for WaitGroup {
25    fn default() -> Self {
26        Self::new()
27    }
28}
29
30impl WaitGroup {
31    /// Create a new wait group.
32    pub fn new() -> Self {
33        Self {
34            state: Mutex::new(WaitGroupState { counter: 0 }),
35            cond: Condvar::new(),
36        }
37    }
38
39    /// Add to the wait group counter.
40    pub fn add(&self, delta: u64) {
41        if delta == 0 {
42            return;
43        }
44        let mut state = self.state.lock().unwrap();
45        state.counter = state
46            .counter
47            .checked_add(delta)
48            .expect("WaitGroup counter overflow");
49    }
50
51    /// Decrement the wait group counter.
52    pub fn done(&self) {
53        let mut state = self.state.lock().unwrap();
54        if state.counter == 0 {
55            panic!("WaitGroup counter decremented below zero");
56        }
57        state.counter -= 1;
58        if state.counter == 0 {
59            self.cond.notify_all();
60        }
61    }
62
63    /// Wait for the counter to reach zero.
64    pub fn wait(&self) {
65        let mut state = self.state.lock().unwrap();
66        while state.counter > 0 {
67            state = self.cond.wait(state).unwrap();
68        }
69    }
70}