Skip to main content

ferrijs_std/utils/
mc_oneshot.rs

1// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::sync::{
5    atomic::{AtomicBool, Ordering},
6    Arc, RwLock,
7};
8
9use rquickjs::{
10    class::{Trace, Tracer},
11    Value,
12};
13use std::ops::Deref;
14use tokio::sync::Notify;
15
16#[derive(Debug)]
17pub struct Shared<T> {
18    is_sent: AtomicBool,
19    value: RwLock<Option<T>>,
20    notify: Notify,
21}
22
23#[derive(Clone, Debug)]
24pub struct Sender<T: Clone>(Arc<Shared<T>>);
25
26impl<T: Clone> Deref for Sender<T> {
27    type Target = Arc<Shared<T>>;
28    fn deref(&self) -> &Self::Target {
29        &self.0
30    }
31}
32
33impl<'js> Trace<'js> for Sender<Value<'js>> {
34    fn trace<'a>(&self, tracer: Tracer<'a, 'js>) {
35        if let Ok(v) = self.value.read() {
36            if let Some(v) = v.as_ref() {
37                tracer.mark(v)
38            }
39        }
40    }
41}
42
43impl<T: Clone> Sender<T> {
44    pub fn send(&self, value: T) {
45        if !self.is_sent.load(Ordering::Relaxed) {
46            self.value.write().unwrap().replace(value);
47            self.is_sent.store(true, Ordering::Release);
48            self.notify.notify_waiters();
49        }
50    }
51
52    pub fn subscribe(&self) -> Receiver<T> {
53        Receiver(Arc::clone(&self.0))
54    }
55}
56
57#[derive(Clone, Debug)]
58pub struct Receiver<T: Clone>(Arc<Shared<T>>);
59
60impl<T: Clone> Deref for Receiver<T> {
61    type Target = Arc<Shared<T>>;
62    fn deref(&self) -> &Self::Target {
63        &self.0
64    }
65}
66
67impl<T: Clone> Receiver<T> {
68    pub async fn recv(&self) -> T {
69        if !self.is_sent.load(Ordering::Acquire) {
70            self.notify.notified().await;
71        }
72        self.value.read().unwrap().clone().unwrap()
73    }
74}
75
76pub fn channel<T: Clone>() -> (Sender<T>, Receiver<T>) {
77    let shared = Arc::new(Shared {
78        is_sent: AtomicBool::new(false),
79        value: RwLock::new(None),
80        notify: Notify::new(),
81    });
82
83    (Sender(Arc::clone(&shared)), Receiver(shared))
84}
85
86#[cfg(test)]
87mod tests {
88    use tokio::join;
89
90    #[tokio::test]
91    async fn test() {
92        let (tx, rx1) = super::channel::<bool>();
93
94        let rx2 = tx.subscribe();
95        let rx3 = tx.subscribe();
96
97        let a = tokio::spawn(async move {
98            let val = rx1.recv().await; //wait for value to become false
99            assert!(val)
100        });
101
102        let b = tokio::spawn(async move {
103            let val = rx2.recv().await; //wait for value to become false
104            assert!(val)
105        });
106
107        tokio::time::sleep(std::time::Duration::from_millis(10)).await;
108
109        tx.send(true);
110
111        let val = rx3.recv().await;
112        assert!(val);
113
114        let (a, b) = join!(a, b);
115        a.unwrap();
116        b.unwrap();
117    }
118}