Skip to main content

rama_http/layer/har/
toggle.rs

1use std::future::ready;
2use std::sync::Arc;
3use std::sync::atomic::{AtomicBool, Ordering};
4use tokio::sync::mpsc;
5
6use rama_core::telemetry::tracing;
7
8pub trait Toggle: Send + Sync + 'static {
9    fn status(&self) -> impl Future<Output = bool> + Send + '_;
10}
11
12impl Toggle for bool {
13    fn status(&self) -> impl Future<Output = Self> + Send + '_ {
14        ready(*self)
15    }
16}
17
18impl<T: Toggle> Toggle for Option<T> {
19    async fn status(&self) -> bool {
20        if let Some(toggle) = self {
21            toggle.status().await
22        } else {
23            false
24        }
25    }
26}
27
28impl Toggle for AtomicBool {
29    fn status(&self) -> impl Future<Output = bool> + Send + '_ {
30        ready(self.load(Ordering::Acquire))
31    }
32}
33
34impl<T: Toggle> Toggle for Arc<T> {
35    fn status(&self) -> impl Future<Output = bool> + Send + '_ {
36        (**self).status()
37    }
38}
39
40impl<F, Fut> Toggle for F
41where
42    F: Fn() -> Fut + Send + Sync + 'static,
43    Fut: Future<Output: Toggle> + Send + 'static,
44{
45    async fn status(&self) -> bool {
46        (self)().await.status().await
47    }
48}
49
50macro_rules! impl_toggle_either {
51    ($id:ident, $($variant:ident),+ $(,)?) => {
52        impl<$($variant),+> Toggle for rama_core::combinators::$id<$($variant),+>
53        where
54            $($variant: Toggle),+
55        {
56            async fn status(&self) -> bool {
57                match self {
58                    $(
59                        rama_core::combinators::$id::$variant(inner) => inner.status().await,
60                    )+
61                }
62            }
63        }
64    };
65}
66
67rama_core::combinators::impl_either!(impl_toggle_either);
68
69pub fn toggle_from_mpsc_recv<T, C>(mut rx: mpsc::Receiver<T>, cancel: C) -> Arc<AtomicBool>
70where
71    T: Send + 'static,
72    C: Future + Send + 'static,
73{
74    let toggle: Arc<AtomicBool> = Default::default();
75    let flag = toggle.clone();
76    tokio::spawn(async move {
77        let mut cancel = std::pin::pin!(cancel);
78        loop {
79            tokio::select! {
80                _ = cancel.as_mut() => {
81                    tracing::trace!("MPSC Toggle cancelled via cancel future; exit");
82                    return;
83                }
84                res = rx.recv() => {
85                    if res.is_some() {
86                        let state = !flag.fetch_xor(true, Ordering::AcqRel);
87                        tracing::trace!("MPSC Toggle received trigger via receiver, new state: {state}");
88                    } else {
89                        tracing::trace!("MPSC Toggle cancelled via closed channel; exit");
90                        return;
91                    }
92                }
93            }
94        }
95    });
96    toggle
97}
98
99pub fn mpsc_toggle<C>(buffer: usize, cancel: C) -> (Arc<AtomicBool>, mpsc::Sender<()>)
100where
101    C: Future + Send + 'static,
102{
103    let (tx, rx) = mpsc::channel(buffer);
104    let toggle = toggle_from_mpsc_recv(rx, cancel);
105    (toggle, tx)
106}
107
108pub fn toggle_from_mpsc_unbounded_recv<T, C>(
109    mut rx: mpsc::UnboundedReceiver<T>,
110    cancel: C,
111) -> Arc<AtomicBool>
112where
113    T: Send + 'static,
114    C: Future + Send + 'static,
115{
116    let toggle: Arc<AtomicBool> = Default::default();
117    let flag = toggle.clone();
118    tokio::spawn(async move {
119        let mut cancel = std::pin::pin!(cancel);
120        loop {
121            tokio::select! {
122                _ = cancel.as_mut() => {
123                    tracing::trace!("uMPSC Toggle cancelled via cancel future; exit");
124                    return;
125                }
126                res = rx.recv() => {
127                    if res.is_some() {
128                        let state = flag.fetch_xor(true, Ordering::AcqRel);
129                        tracing::trace!("uMPSC Toggle received trigger via receiver, new state: {state}");
130                    } else {
131                        tracing::trace!("uMPSC Toggle cancelled via closed channel; exit");
132                        return;
133                    }
134                }
135            }
136        }
137    });
138    toggle
139}
140
141pub fn mpsc_unbounded_toggle<C>(cancel: C) -> (Arc<AtomicBool>, mpsc::UnboundedSender<()>)
142where
143    C: Future + Send + 'static,
144{
145    let (tx, rx) = mpsc::unbounded_channel();
146    let toggle = toggle_from_mpsc_unbounded_recv(rx, cancel);
147    (toggle, tx)
148}
149
150#[cfg(target_family = "unix")]
151#[cfg_attr(docsrs, doc(cfg(target_family = "unix")))]
152pub fn mpsc_toggle_for_unix_signal<C>(
153    signal: tokio::signal::unix::SignalKind,
154    cancel: C,
155) -> Result<Arc<AtomicBool>, std::io::Error>
156where
157    C: Future + Send + 'static,
158{
159    let mut signal = tokio::signal::unix::signal(signal)?;
160
161    let toggle: Arc<AtomicBool> = Default::default();
162    let flag = toggle.clone();
163
164    tokio::spawn(async move {
165        let mut cancel = std::pin::pin!(cancel);
166        loop {
167            tokio::select! {
168                _ = cancel.as_mut() => {
169                    tracing::trace!("unix signal trigger cancelled via cancel future; exit");
170                    return;
171                }
172                res = signal.recv() => {
173                    if res.is_some() {
174                        let state = flag.fetch_xor(true, Ordering::AcqRel);
175                        tracing::trace!("unix signal triggered, new state: {state}");
176                    } else {
177                        tracing::trace!("unix signal closed: cancel and exit");
178                        return;
179                    }
180                }
181            }
182        }
183    });
184
185    Ok(toggle)
186}