Skip to main content

state_m/
barrier.rs

1use std::{
2    fmt::Debug,
3    sync::{
4        Arc,
5        atomic::{AtomicBool, AtomicUsize, Ordering},
6    },
7};
8use tokio::sync::{Notify, futures::Notified};
9
10pub trait AsPassCheck: Debug {
11    fn is_open(&self) -> bool;
12
13    fn notified(&self) -> Notified<'_>;
14}
15
16impl<T> AsPassCheck for Arc<T>
17where
18    T: AsPassCheck,
19{
20    fn is_open(&self) -> bool {
21        self.as_ref().is_open()
22    }
23
24    fn notified(&self) -> Notified<'_> {
25        self.as_ref().notified()
26    }
27}
28
29#[derive(Debug, Default)]
30pub struct Barrier(Arc<AtomicUsize>, Arc<Notify>);
31
32impl Drop for Barrier {
33    fn drop(&mut self) {
34        let counter = self.0.update(Ordering::Release, Ordering::Acquire, |c| {
35            if c == usize::MIN { c } else { c - 1 }
36        });
37        if counter <= 1 {
38            self.1.notify_waiters();
39        }
40    }
41}
42
43#[derive(Clone, Debug, Default)]
44pub struct Barriers(Arc<AtomicUsize>, Arc<Notify>);
45
46impl AsPassCheck for Barriers {
47    fn is_open(&self) -> bool {
48        self.0.load(Ordering::Acquire) == 0
49    }
50
51    fn notified(&self) -> Notified<'_> {
52        self.1.notified()
53    }
54}
55
56impl Barriers {
57    pub fn new() -> Self {
58        Self::default()
59    }
60
61    pub fn add_barrier(&self) -> Option<Barrier> {
62        let counter = self.0.update(Ordering::Release, Ordering::Acquire, |c| {
63            if c < usize::MAX { c + 1 } else { c }
64        });
65        if counter == usize::MAX {
66            return None;
67        } else {
68            Some(Barrier(self.0.clone(), self.1.clone()))
69        }
70    }
71}
72
73#[derive(Clone, Debug, Default)]
74pub struct Door(Arc<AtomicBool>, Arc<Notify>);
75
76impl AsPassCheck for Door {
77    fn is_open(&self) -> bool {
78        self.0.load(Ordering::Acquire) == false
79    }
80
81    fn notified(&self) -> Notified<'_> {
82        self.1.notified()
83    }
84}
85
86impl Door {
87    pub fn new() -> Self {
88        Self::default()
89    }
90
91    pub fn open(&self) {
92        self.0.store(false, Ordering::Release);
93        self.1.notify_waiters();
94    }
95
96    pub fn close(&self) {
97        self.0.store(true, Ordering::Release);
98    }
99}
100
101#[cfg(test)]
102mod tests {
103    use super::*;
104    use tokio::task::JoinSet;
105
106    #[tokio::test]
107    async fn test_door() {
108        let res: Arc<AtomicUsize> = Default::default();
109        let door = Door::new();
110        let func = |join_set: &mut JoinSet<_>, (door, res): (Door, Arc<AtomicUsize>)| {
111            join_set.spawn(async move {
112                if door.is_open() {
113                    println!("door is open");
114                    res.fetch_add(1, Ordering::AcqRel);
115                } else {
116                    println!("door is closed");
117                    door.notified().await;
118                    println!("door is open now");
119                    res.fetch_add(10, Ordering::AcqRel);
120                }
121            })
122        };
123        let mut join_set = JoinSet::new();
124        for _ in 0..10 {
125            let vars = (door.clone(), res.clone());
126            func(&mut join_set, vars);
127        }
128        join_set.join_all().await;
129        assert_eq!(10, res.load(Ordering::Acquire));
130
131        door.close();
132        res.store(0, Ordering::Release);
133        let mut join_set = JoinSet::new();
134        for _ in 0..10 {
135            let vars = (door.clone(), res.clone());
136            func(&mut join_set, vars);
137        }
138        let door_c = door.clone();
139        tokio::spawn(async move {
140            tokio::time::sleep(std::time::Duration::from_secs(1)).await;
141            door_c.open();
142        });
143        join_set.join_all().await;
144        assert_eq!(100, res.load(Ordering::Acquire));
145    }
146
147    #[tokio::test]
148    async fn test_barrier() {
149        let res: Arc<AtomicUsize> = Default::default();
150        let barriers = Barriers::new();
151        let func = |join_set: &mut JoinSet<_>, (barriers, res): (Barriers, Arc<AtomicUsize>)| {
152            join_set.spawn(async move {
153                if barriers.is_open() {
154                    println!("door is open");
155                    res.fetch_add(1, Ordering::AcqRel);
156                } else {
157                    println!("door is closed");
158                    barriers.notified().await;
159                    println!("door is open now");
160                    res.fetch_add(10, Ordering::AcqRel);
161                }
162            })
163        };
164        let mut join_set = JoinSet::new();
165        for _ in 0..10 {
166            let vars = (barriers.clone(), res.clone());
167            func(&mut join_set, vars);
168        }
169        join_set.join_all().await;
170        assert_eq!(10, res.load(Ordering::Acquire));
171
172        res.store(0, Ordering::Release);
173        let mut join_set = JoinSet::new();
174        let barriers_c = barriers.clone();
175        tokio::spawn(async move {
176            let barrier = barriers_c.add_barrier();
177            assert_eq!(true, barrier.is_some());
178            tokio::time::sleep(std::time::Duration::from_secs(2)).await;
179        });
180        tokio::time::sleep(std::time::Duration::from_secs(1)).await;
181        for _ in 0..10 {
182            let vars = (barriers.clone(), res.clone());
183            func(&mut join_set, vars);
184        }
185        join_set.join_all().await;
186        assert_eq!(100, res.load(Ordering::Acquire));
187    }
188}