1use std::{future::Future, sync::Arc, sync::atomic::AtomicUsize, sync::atomic::Ordering};
2
3static mut CBS: Option<Arc<dyn CallbacksApi>> = None;
6
7static STATE: AtomicUsize = AtomicUsize::new(0);
8
9const UNINITIALIZED: usize = 0;
13const INITIALIZING: usize = 1;
14const INITIALIZED: usize = 2;
15
16trait CallbacksApi {
17 fn before(&self) -> Option<*const ()>;
18 fn enter(&self, _: *const ()) -> *const ();
19 fn exit(&self, _: *const ());
20 fn after(&self, _: *const ());
21}
22
23#[allow(clippy::struct_field_names)]
24struct Callbacks<A, B, C, D> {
25 f_before: A,
26 f_enter: B,
27 f_exit: C,
28 f_after: D,
29}
30
31impl<A, B, C, D> CallbacksApi for Callbacks<A, B, C, D>
32where
33 A: Fn() -> Option<*const ()> + 'static,
34 B: Fn(*const ()) -> *const () + 'static,
35 C: Fn(*const ()) + 'static,
36 D: Fn(*const ()) + 'static,
37{
38 fn before(&self) -> Option<*const ()> {
39 (self.f_before)()
40 }
41 fn enter(&self, d: *const ()) -> *const () {
42 (self.f_enter)(d)
43 }
44 fn exit(&self, d: *const ()) {
45 (self.f_exit)(d);
46 }
47 fn after(&self, d: *const ()) {
48 (self.f_after)(d);
49 }
50}
51
52pub(crate) struct Data {
53 cb: &'static dyn CallbacksApi,
54 ptr: *const (),
55}
56
57impl Data {
58 #[allow(clippy::if_not_else)]
59 pub(crate) fn load() -> Option<Data> {
60 let cb = if STATE.load(Ordering::Acquire) != INITIALIZED {
62 None
63 } else {
64 #[allow(static_mut_refs)]
65 unsafe {
66 Some(CBS.as_ref().map(AsRef::as_ref).unwrap())
67 }
68 };
69
70 if let Some(cb) = cb
71 && let Some(ptr) = cb.before()
72 {
73 return Some(Data { cb, ptr });
74 }
75 None
76 }
77
78 pub(crate) fn run<F, R>(&mut self, f: F) -> R
79 where
80 F: FnOnce() -> R,
81 {
82 let ptr = self.cb.enter(self.ptr);
83 let result = f();
84 self.cb.exit(ptr);
85 result
86 }
87}
88
89pub(crate) fn wrap<F: Future>(fut: F) -> impl Future<Output = F::Output> {
94 let mut data = Data::load();
95 async move {
96 let mut f = std::pin::pin!(fut);
97 std::future::poll_fn(|cx| match data.as_mut() {
98 Some(data) => data.run(|| f.as_mut().poll(cx)),
99 None => f.as_mut().poll(cx),
100 })
101 .await
102 }
103}
104
105impl Drop for Data {
106 fn drop(&mut self) {
107 self.cb.after(self.ptr);
108 }
109}
110
111pub unsafe fn task_callbacks<FBefore, FEnter, FExit, FAfter>(
121 f_before: FBefore,
122 f_enter: FEnter,
123 f_exit: FExit,
124 f_after: FAfter,
125) where
126 FBefore: Fn() -> Option<*const ()> + 'static + Sync,
127 FEnter: Fn(*const ()) -> *const () + 'static + Sync,
128 FExit: Fn(*const ()) + 'static + Sync,
129 FAfter: Fn(*const ()) + 'static + Sync,
130{
131 let new = Arc::new(Callbacks {
132 f_before,
133 f_enter,
134 f_exit,
135 f_after,
136 });
137 let _ = set_cbs(new);
138}
139
140pub unsafe fn task_opt_callbacks<FBefore, FEnter, FExit, FAfter>(
149 f_before: FBefore,
150 f_enter: FEnter,
151 f_exit: FExit,
152 f_after: FAfter,
153) -> bool
154where
155 FBefore: Fn() -> Option<*const ()> + Sync + 'static,
156 FEnter: Fn(*const ()) -> *const () + Sync + 'static,
157 FExit: Fn(*const ()) + Sync + 'static,
158 FAfter: Fn(*const ()) + Sync + 'static,
159{
160 let new = Arc::new(Callbacks {
161 f_before,
162 f_enter,
163 f_exit,
164 f_after,
165 });
166 set_cbs(new).is_ok()
167}
168
169fn set_cbs(cbs: Arc<dyn CallbacksApi>) -> Result<(), ()> {
170 match STATE.compare_exchange(
171 UNINITIALIZED,
172 INITIALIZING,
173 Ordering::Acquire,
174 Ordering::Relaxed,
175 ) {
176 Ok(UNINITIALIZED) => {
177 unsafe {
178 CBS = Some(cbs);
179 }
180 STATE.store(INITIALIZED, Ordering::Release);
181 Ok(())
182 }
183 Err(INITIALIZING) => {
184 while STATE.load(Ordering::Relaxed) == INITIALIZING {
185 std::hint::spin_loop();
186 }
187 Err(())
188 }
189 _ => Err(()),
190 }
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196
197 #[test]
198 fn load_callbacks_set_by_other_thread() {
199 let hnd = std::thread::spawn(|| unsafe {
200 task_opt_callbacks(|| Some(std::ptr::null()), |p| p, |_| {}, |_| {})
201 });
202 while Data::load().is_none() {
204 std::hint::spin_loop();
205 }
206 assert!(hnd.join().unwrap());
207 }
208}