Skip to main content

ntex_service/
ctx.rs

1use std::{cell, fmt, future, marker, ops, pin, rc::Rc, task::Context, task::Poll, task::Waker};
2
3use crate::Service;
4
5/// Context provided to [`Service`] lifecycle methods.
6///
7/// A context gives a service access to pipeline state and coordinates calls to
8/// inner services with the pipeline's readiness and shutdown machinery.
9pub struct Ctx<'a, Svc: ?Sized, St = ()> {
10    idx: u32,
11    st: &'a St,
12    waiters: &'a WaitersRef,
13    _t: marker::PhantomData<Rc<Svc>>,
14}
15
16#[derive(Debug)]
17pub(crate) struct WaitersRef {
18    cur: cell::Cell<u32>,
19    running: cell::Cell<bool>,
20    shutdown: cell::Cell<bool>,
21    ready: cell::Cell<bool>,
22    wakers: cell::UnsafeCell<Vec<u32>>,
23    indexes: cell::UnsafeCell<slab::Slab<Option<Waker>>>,
24}
25
26impl WaitersRef {
27    pub(crate) fn new() -> Self {
28        let mut waiters = slab::Slab::with_capacity(16);
29        waiters.insert(None);
30        WaitersRef {
31            cur: cell::Cell::new(u32::MAX),
32            running: cell::Cell::new(false),
33            shutdown: cell::Cell::new(false),
34            ready: cell::Cell::new(false),
35            indexes: cell::UnsafeCell::new(waiters),
36            wakers: cell::UnsafeCell::new(Vec::default()),
37        }
38    }
39
40    #[allow(clippy::mut_from_ref)]
41    pub(crate) fn get(&self) -> &mut slab::Slab<Option<Waker>> {
42        unsafe { &mut *self.indexes.get() }
43    }
44
45    #[allow(clippy::mut_from_ref)]
46    pub(crate) fn get_wakers(&self) -> &mut Vec<u32> {
47        unsafe { &mut *self.wakers.get() }
48    }
49
50    pub(crate) fn insert(&self) -> u32 {
51        self.get().insert(None) as u32
52    }
53
54    pub(crate) fn remove(&self, idx: u32) {
55        self.get().remove(idx as usize);
56
57        if self.cur.get() == idx {
58            self.notify();
59        }
60    }
61
62    pub(crate) fn notify(&self) {
63        let wakers = self.get_wakers();
64        if !wakers.is_empty() {
65            let indexes = self.get();
66            for idx in wakers.drain(..) {
67                if let Some(item) = indexes.get_mut(idx as usize)
68                    && let Some(waker) = item.take()
69                {
70                    waker.wake();
71                }
72            }
73        }
74
75        self.cur.set(u32::MAX);
76    }
77
78    pub(crate) fn run<F, R>(&self, idx: u32, cx: &mut Context<'_>, f: F) -> Poll<R>
79    where
80        F: FnOnce(&mut Context<'_>) -> Poll<R>,
81    {
82        // Ctx::poll_xxx() methods requires current waker always available,
83        // nested readiness checks poll with the same waker
84        let slot = &mut self.get()[idx as usize];
85        if !slot.as_ref().is_some_and(|w| w.will_wake(cx.waker())) {
86            *slot = Some(cx.waker().clone());
87        }
88
89        // calculate owner for readiness check
90        let cur = self.cur.get();
91        let can_check = if cur == idx {
92            true
93        } else if cur == u32::MAX {
94            self.cur.set(idx);
95            true
96        } else {
97            false
98        };
99
100        if can_check {
101            // only one readiness check can manage waiters
102            let initial_run = !self.running.get();
103            if initial_run {
104                self.running.set(true);
105            }
106
107            let result = f(cx);
108
109            if initial_run {
110                if result.is_pending() {
111                    self.get_wakers().push(idx);
112                } else {
113                    self.notify();
114                }
115                self.running.set(false);
116            }
117            result
118        } else {
119            // Another pipeline binding owns the readiness check.
120            self.get_wakers().push(idx);
121            Poll::Pending
122        }
123    }
124
125    pub(crate) fn shutdown(&self) {
126        self.shutdown.set(true);
127        self.ready.set(false);
128    }
129
130    /// Records the result of a pipeline readiness check
131    pub(crate) fn set_ready(&self, ready: bool) {
132        self.ready.set(ready);
133    }
134
135    /// Consumes the result of the last successful pipeline readiness check
136    pub(crate) fn take_ready(&self) -> bool {
137        self.ready.replace(false)
138    }
139
140    pub(crate) fn is_shutdown(&self) -> bool {
141        self.shutdown.get()
142    }
143}
144
145impl<'a, Svc, St> Ctx<'a, Svc, St> {
146    pub(crate) fn new(idx: u32, waiters: &'a WaitersRef, st: &'a St) -> Self {
147        Self {
148            idx,
149            waiters,
150            st,
151            _t: marker::PhantomData,
152        }
153    }
154
155    pub(crate) fn inner(self) -> (u32, &'a WaitersRef, &'a St) {
156        (self.idx, self.waiters, self.st)
157    }
158
159    #[inline]
160    /// Returns the identifier of the current pipeline binding.
161    pub fn id(&self) -> u32 {
162        self.idx
163    }
164
165    #[inline]
166    /// Returns the pipeline state.
167    pub fn st(&'a self) -> &'a St {
168        self.st
169    }
170
171    /// Waits until `svc` is ready to process a request.
172    pub async fn ready<S, Req>(&self, svc: &'a S) -> Result<(), S::Error>
173    where
174        S: Service<St, Req>,
175    {
176        // check readiness and notify waiters
177        ReadyCall {
178            completed: false,
179            fut: svc.ready(Ctx {
180                st: self.st,
181                idx: self.idx,
182                waiters: self.waiters,
183                _t: marker::PhantomData,
184            }),
185            idx: self.idx,
186            waiters: self.waiters,
187        }
188        .await
189    }
190
191    #[inline]
192    /// Waits for `svc` to become ready, then calls it.
193    pub async fn call<S, Req>(&self, svc: &'a S, req: Req) -> Result<S::Res, S::Error>
194    where
195        S: Service<St, Req>,
196    {
197        self.ready(svc).await?;
198
199        svc.call(
200            req,
201            Ctx {
202                idx: self.idx,
203                st: self.st,
204                waiters: self.waiters,
205                _t: marker::PhantomData,
206            },
207        )
208        .await
209    }
210
211    #[inline]
212    /// Calls `svc` without checking readiness.
213    ///
214    /// The caller must ensure that `svc` is ready.
215    pub async fn call_nowait<S, Req>(&self, svc: &'a S, req: Req) -> Result<S::Res, S::Error>
216    where
217        S: Service<St, Req>,
218    {
219        svc.call(
220            req,
221            Ctx {
222                st: self.st,
223                idx: self.idx,
224                waiters: self.waiters,
225                _t: marker::PhantomData,
226            },
227        )
228        .await
229    }
230
231    #[inline]
232    /// Polls a closure until completion using the dispatcher's task [`Context`].
233    ///
234    /// The [`Context`] uses the waker of the task that currently owns the readiness
235    /// check. Outside of a readiness check a no-op waker is used, so a `Pending`
236    /// result is not woken. If the current owner is not the main runner, the main
237    /// runner is also scheduled for wake-up.
238    pub async fn poll_fn<F, R>(&'a self, f: F) -> R
239    where
240        F: Fn(&mut Context<'a>) -> Poll<R>,
241    {
242        future::poll_fn(move |_| {
243            let wakers = self.waiters.get();
244            let idx = self.waiters.cur.get() as usize;
245            let mut ctx = if let Some(w) = wakers.get(idx).and_then(|w| w.as_ref()) {
246                Context::from_waker(w)
247            } else {
248                Context::from_waker(Waker::noop())
249            };
250
251            // if the current runner is not the main one, wake up the main task
252            // so it re-registers all required wakers
253            if idx != 0 {
254                self.waiters.get_wakers().push(0);
255            }
256
257            f(&mut ctx)
258        })
259        .await
260    }
261
262    #[inline]
263    /// Calls a closure once with the dispatcher's task [`Context`].
264    ///
265    /// The [`Context`] uses the waker of the task that currently owns the readiness
266    /// check. Outside of a readiness check a no-op waker is used, so a `Pending`
267    /// result is not woken. If the current owner is not the main runner, the main
268    /// runner is also scheduled for wake-up.
269    pub fn poll_once<F, R>(&'a self, f: F) -> R
270    where
271        F: FnOnce(&mut Context<'a>) -> R,
272    {
273        let wakers = self.waiters.get();
274        let idx = self.waiters.cur.get() as usize;
275        let mut ctx = if let Some(w) = wakers.get(idx).and_then(|w| w.as_ref()) {
276            Context::from_waker(w)
277        } else {
278            Context::from_waker(Waker::noop())
279        };
280
281        // if the current runner is not the main one, wake up the main task
282        // so it re-registers all required wakers
283        if idx != 0 {
284            self.waiters.get_wakers().push(0);
285        }
286
287        f(&mut ctx)
288    }
289
290    #[inline]
291    /// Shuts down `svc`.
292    pub async fn shutdown<S, Req>(&self, svc: &'a S)
293    where
294        S: Service<St, Req>,
295    {
296        svc.shutdown(Ctx {
297            idx: self.idx,
298            st: self.st,
299            waiters: self.waiters,
300            _t: marker::PhantomData,
301        })
302        .await;
303    }
304
305    #[inline]
306    /// Returns a context that uses `st` as its state.
307    pub fn map_state<NewSt>(&'a self, st: &'a NewSt) -> Ctx<'a, Self, NewSt> {
308        Ctx {
309            st,
310            idx: self.idx,
311            waiters: self.waiters,
312            _t: marker::PhantomData,
313        }
314    }
315}
316
317impl<S, St> Copy for Ctx<'_, S, St> {}
318
319impl<S, St> Clone for Ctx<'_, S, St> {
320    #[inline]
321    fn clone(&self) -> Self {
322        *self
323    }
324}
325
326impl<S, St> ops::Deref for Ctx<'_, S, St> {
327    type Target = St;
328
329    #[inline]
330    fn deref(&self) -> &St {
331        self.st
332    }
333}
334
335impl<S, St> fmt::Debug for Ctx<'_, S, St> {
336    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
337        f.debug_struct("Ctx")
338            .field("idx", &self.idx)
339            .field("waiters", &self.waiters.get().len())
340            .finish()
341    }
342}
343
344struct ReadyCall<'a, F: future::Future> {
345    completed: bool,
346    fut: F,
347    idx: u32,
348    waiters: &'a WaitersRef,
349}
350
351impl<F: future::Future> Drop for ReadyCall<'_, F> {
352    fn drop(&mut self) {
353        if !self.completed && self.waiters.cur.get() == self.idx {
354            self.waiters.notify();
355        }
356    }
357}
358
359impl<F: future::Future> Unpin for ReadyCall<'_, F> {}
360
361impl<F: future::Future> future::Future for ReadyCall<'_, F> {
362    type Output = F::Output;
363
364    fn poll(mut self: pin::Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
365        self.waiters.run(self.idx, cx, |cx| {
366            // SAFETY: `fut` never moves
367            let result = unsafe { pin::Pin::new_unchecked(&mut self.as_mut().fut).poll(cx) };
368            if result.is_ready() {
369                self.completed = true;
370            }
371            result
372        })
373    }
374}
375
376#[cfg(test)]
377mod tests {
378    use std::{cell::Cell, cell::RefCell, future::poll_fn};
379
380    use ntex::channel::{condition, oneshot};
381    use ntex::{rt::spawn, time, util::lazy, util::select};
382
383    use super::*;
384    use crate::Pipeline;
385
386    struct Srv(Rc<Cell<usize>>, condition::Waiter);
387
388    impl Service<(), &'static str> for Srv {
389        type Res = &'static str;
390        type Error = ();
391
392        async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), Self::Error> {
393            self.0.set(self.0.get() + 1);
394            self.1.ready().await;
395            Ok(())
396        }
397
398        async fn call(
399            &self,
400            req: &'static str,
401            ctx: Ctx<'_, Self>,
402        ) -> Result<Self::Res, Self::Error> {
403            let _ = format!("{ctx:?}");
404            let _ = format!("{:?}", ctx.id());
405            let () = *ctx;
406            #[allow(clippy::clone_on_copy)]
407            let _ = ctx.clone();
408            Ok(req)
409        }
410    }
411
412    #[ntex::test]
413    async fn test_ready() {
414        let cnt = Rc::new(Cell::new(0));
415        let con = condition::Condition::new();
416
417        let srv = Pipeline::new((), Srv(cnt.clone(), con.wait()));
418        let res = lazy(|cx| srv.poll_ready(cx)).await;
419        assert_eq!(res, Poll::Pending);
420        assert_eq!(cnt.get(), 1);
421
422        let res = lazy(|cx| srv.poll_ready(cx)).await;
423        assert_eq!(res, Poll::Pending);
424        assert_eq!(cnt.get(), 1);
425
426        con.notify(());
427        let res = lazy(|cx| srv.poll_ready(cx)).await;
428        assert_eq!(res, Poll::Ready(Ok(())));
429        assert_eq!(cnt.get(), 1);
430
431        let res = lazy(|cx| srv.poll_ready(cx)).await;
432        assert_eq!(res, Poll::Pending);
433        assert_eq!(cnt.get(), 2);
434
435        con.notify(());
436        let res = lazy(|cx| srv.poll_ready(cx)).await;
437        assert_eq!(res, Poll::Ready(Ok(())));
438        assert_eq!(cnt.get(), 2);
439
440        let res = lazy(|cx| srv.poll_ready(cx)).await;
441        assert_eq!(res, Poll::Pending);
442        assert_eq!(cnt.get(), 3);
443    }
444
445    struct WakeCounter {
446        clones: Cell<usize>,
447        wakes: Cell<usize>,
448    }
449
450    fn counting_waker() -> (&'static WakeCounter, Waker) {
451        use std::task::{RawWaker, RawWakerVTable};
452
453        static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, wake, wake, drop);
454
455        unsafe fn clone(data: *const ()) -> RawWaker {
456            let cnt = unsafe { &*data.cast::<WakeCounter>() };
457            cnt.clones.set(cnt.clones.get() + 1);
458            RawWaker::new(data, &VTABLE)
459        }
460        unsafe fn wake(data: *const ()) {
461            let cnt = unsafe { &*data.cast::<WakeCounter>() };
462            cnt.wakes.set(cnt.wakes.get() + 1);
463        }
464        unsafe fn drop(_: *const ()) {}
465
466        let cnt: &'static WakeCounter = Box::leak(Box::new(WakeCounter {
467            clones: Cell::new(0),
468            wakes: Cell::new(0),
469        }));
470        let raw = RawWaker::new(std::ptr::from_ref(cnt).cast(), &VTABLE);
471        (cnt, unsafe { Waker::from_raw(raw) })
472    }
473
474    /// Wakes the dispatcher's task context during the readiness check
475    struct WakeSrv;
476
477    impl Service<(), &'static str> for WakeSrv {
478        type Res = &'static str;
479        type Error = ();
480
481        async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
482            ctx.poll_once(|cx| cx.waker().wake_by_ref());
483            Ok(())
484        }
485
486        async fn call(&self, req: &'static str, _: Ctx<'_, Self>) -> Result<&'static str, ()> {
487            Ok(req)
488        }
489    }
490
491    struct Nested<S>(S);
492
493    impl<S: Service<(), &'static str>> Service<(), &'static str> for Nested<S> {
494        type Res = S::Res;
495        type Error = S::Error;
496
497        async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
498            ctx.ready(&self.0).await
499        }
500
501        async fn call(
502            &self,
503            req: &'static str,
504            ctx: Ctx<'_, Self>,
505        ) -> Result<Self::Res, Self::Error> {
506            ctx.call(&self.0, req).await
507        }
508    }
509
510    #[ntex::test]
511    async fn test_ready_waker_not_recloned() {
512        let srv = Pipeline::new((), Nested(Nested(WakeSrv)));
513
514        let (cnt1, waker1) = counting_waker();
515        let mut cx = Context::from_waker(&waker1);
516        for _ in 0..4 {
517            assert_eq!(srv.poll_ready(&mut cx), Poll::Ready(Ok(())));
518        }
519        assert_eq!(cnt1.clones.get(), 1);
520        assert_eq!(cnt1.wakes.get(), 4);
521
522        // a different waker replaces the stored one
523        let (cnt2, waker2) = counting_waker();
524        let mut cx = Context::from_waker(&waker2);
525        assert_eq!(srv.poll_ready(&mut cx), Poll::Ready(Ok(())));
526        assert_eq!(cnt2.clones.get(), 1);
527        assert_eq!(cnt2.wakes.get(), 1);
528        assert_eq!(cnt1.wakes.get(), 4);
529    }
530
531    #[ntex::test]
532    async fn test_ready_on_drop() {
533        let cnt = Rc::new(Cell::new(0));
534        let con = condition::Condition::new();
535        let srv = Pipeline::new((), Srv(cnt.clone(), con.wait()));
536
537        let srv1 = srv.bind();
538        let (tx, rx) = oneshot::channel();
539        spawn(async move {
540            select(rx, srv1.ready()).await;
541            time::sleep(time::Millis(25000)).await;
542        });
543        time::sleep(time::Millis(250)).await;
544
545        let res = lazy(|cx| srv.poll_ready(cx)).await;
546        assert_eq!(res, Poll::Pending);
547
548        let _ = tx.send(());
549        time::sleep(time::Millis(250)).await;
550
551        let res = lazy(|cx| srv.poll_ready(cx)).await;
552        assert_eq!(res, Poll::Pending);
553
554        con.notify(());
555        let res = lazy(|cx| srv.poll_ready(cx)).await;
556        assert_eq!(res, Poll::Ready(Ok(())));
557    }
558
559    #[ntex::test]
560    async fn test_ready_after_shutdown() {
561        let cnt = Rc::new(Cell::new(0));
562        let con = condition::Condition::new();
563        let srv = Pipeline::new((), Srv(cnt.clone(), con.wait()));
564
565        let res = lazy(|cx| srv.poll_ready(cx)).await;
566        assert_eq!(res, Poll::Pending);
567
568        let (tx, rx) = oneshot::channel();
569        let (tx2, rx2) = oneshot::channel();
570        spawn(async move {
571            select(rx, srv.ready()).await;
572            srv.shutdown().await;
573            let _ = tx2.send(srv);
574        });
575        time::sleep(time::Millis(250)).await;
576
577        let _ = tx.send(());
578        let srv = rx2.await.unwrap();
579
580        let res = lazy(|cx| srv.poll_ready(cx)).await;
581        assert_eq!(res, Poll::Ready(Ok(())));
582
583        con.notify(());
584        let res = lazy(|cx| srv.poll_ready(cx)).await;
585        assert_eq!(res, Poll::Ready(Ok(())));
586    }
587
588    #[ntex::test]
589    async fn test_pipeline_binding_after_shutdown() {
590        let cnt = Rc::new(Cell::new(0));
591        let con = condition::Condition::new();
592        let srv = Pipeline::new((), Srv(cnt.clone(), con.wait()));
593        poll_fn(|cx| srv.poll_shutdown(cx)).await;
594        let _ = poll_fn(|cx| srv.poll_ready(cx)).await;
595    }
596
597    #[ntex::test]
598    async fn test_shared_call() {
599        let data = Rc::new(RefCell::new(Vec::new()));
600
601        let cnt = Rc::new(Cell::new(0));
602        let con = condition::Condition::new();
603
604        let srv = Pipeline::new((), Srv(cnt.clone(), con.wait()));
605
606        let srv1 = srv.bind();
607        let data1 = data.clone();
608        ntex::rt::spawn(async move {
609            let _ = srv1.ready().await;
610            let fut = srv1.call_static("srv1");
611            assert!(format!("{fut:?}").contains("PipelineCall"));
612            let i = fut.await.unwrap();
613            data1.borrow_mut().push(i);
614        });
615
616        let srv2 = srv.bind();
617        let data2 = data.clone();
618        ntex::rt::spawn(async move {
619            let i = srv2.call("srv2").await.unwrap();
620            data2.borrow_mut().push(i);
621        });
622        time::sleep(time::Millis(50)).await;
623
624        con.notify(());
625        time::sleep(time::Millis(150)).await;
626
627        assert_eq!(cnt.get(), 2);
628        assert_eq!(&*data.borrow(), &["srv1"]);
629
630        con.notify(());
631        time::sleep(time::Millis(150)).await;
632
633        assert_eq!(cnt.get(), 2);
634        assert_eq!(&*data.borrow(), &["srv1", "srv2"]);
635    }
636}