Skip to main content

ntex_rt/
task.rs

1use std::{future::Future, sync::Arc, sync::atomic::AtomicUsize, sync::atomic::Ordering};
2
3// The Callbacks static holds a pointer to the global callbacks. It is protected by
4// the STATE static which determines whether `CBS` has been initialized yet.
5static mut CBS: Option<Arc<dyn CallbacksApi>> = None;
6
7static STATE: AtomicUsize = AtomicUsize::new(0);
8
9// There are three different states that we care about: the callback's
10// uninitialized, the callback's initializing (set_cbs's been called but
11// CBS hasn't actually been set yet), or the callbacks's active.
12const 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        // `Acquire` pairs with the `Release` store in `set_cbs`, so `CBS` is visible
61        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
89/// Wraps a future so that task callbacks, if registered, run around each poll.
90///
91/// Always returns the same future type, so spawning code is generated once
92/// per future rather than once per callback mode.
93pub(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
111/// # Safety
112///
113/// The user must ensure that the pointer returned by `before` has a `'static` lifetime.
114/// This pointer will be owned by the spawned task for the duration of that task, and
115/// ownership will be returned to the user at the end of the task via `after`.
116/// The pointer remains opaque to the runtime.
117///
118/// Does nothing if task callbacks have already been set, use
119/// [`task_opt_callbacks`] to check whether they were set.
120pub 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
140/// # Safety
141///
142/// The user must ensure that the pointer returned by `before` has a `'static` lifetime.
143/// This pointer will be owned by the spawned task for the duration of that task, and
144/// ownership will be returned to the user at the end of the task via `after`.
145/// The pointer remains opaque to the runtime.
146///
147/// Returns false if task callbacks have already been set.
148pub 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        // observes callbacks set by another thread without synchronization
203        while Data::load().is_none() {
204            std::hint::spin_loop();
205        }
206        assert!(hnd.join().unwrap());
207    }
208}