1use std::{cell, fmt, future, marker, ops, pin, rc::Rc, task::Context, task::Poll, task::Waker};
2
3use crate::Service;
4
5pub 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 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 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 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 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 pub(crate) fn set_ready(&self, ready: bool) {
132 self.ready.set(ready);
133 }
134
135 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 pub fn id(&self) -> u32 {
162 self.idx
163 }
164
165 #[inline]
166 pub fn st(&'a self) -> &'a St {
168 self.st
169 }
170
171 pub async fn ready<S, Req>(&self, svc: &'a S) -> Result<(), S::Error>
173 where
174 S: Service<St, Req>,
175 {
176 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 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 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 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 idx != 0 {
254 self.waiters.get_wakers().push(0);
255 }
256
257 f(&mut ctx)
258 })
259 .await
260 }
261
262 #[inline]
263 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 idx != 0 {
284 self.waiters.get_wakers().push(0);
285 }
286
287 f(&mut ctx)
288 }
289
290 #[inline]
291 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 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 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 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 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}