rama_http/layer/har/
toggle.rs1use 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}