1use std::{cell, fmt, future, pin::Pin, ptr, rc::Rc, task::Context, task::Poll};
2
3use crate::{Ctx, IntoService, Service, ctx::WaitersRef, util::BoxFuture};
4
5use crate::pipeline::PipelineBinding;
6use crate::pl_inner::{PipelineApi, PipelineInternalApi};
7
8pub struct PipelineState<St, Req, Res, Err> {
14 api: Rc<dyn PipelineStateApi<St, Req, Res, Err>>,
15}
16
17impl<St, Req, Res, Err> PipelineState<St, Req, Res, Err>
18where
19 St: 'static,
20 Req: 'static,
21 Res: 'static,
22 Err: 'static,
23{
24 #[inline]
25 pub fn new<S>(service: impl IntoService<S, St, Req>) -> Self
27 where
28 S: Service<St, Req, Res = Res, Error = Err> + 'static,
29 St: 'static,
30 {
31 PipelineState {
32 api: Rc::new(PipelineInner {
33 s: service.into_service(),
34 waiters: WaitersRef::new(),
35 st_runtime: cell::UnsafeCell::new(RuntimeState::New),
36 }),
37 }
38 }
39
40 #[inline]
41 pub async fn ready(&self, st: &St) -> Result<(), Err> {
46 self.api.ready(0, st).await
47 }
48
49 #[inline]
50 pub async fn call(&self, req: Req, st: &St) -> Result<Res, Err> {
55 let pl = self.binding();
56 self.api.call(pl.idx, req, st).await
57 }
58
59 #[inline]
60 pub async fn shutdown(&self, st: &St) {
62 self.api.shutdown(0, st).await;
63 }
64
65 #[inline]
66 pub fn poll_ready(&self, cx: &mut Context<'_>, st: &St) -> Poll<Result<(), Err>>
76 where
77 St: Clone,
78 {
79 self.api.poll_ready(cx, st)
80 }
81
82 fn binding(&self) -> Binding<'_, St, Req, Res, Err> {
83 Binding {
84 idx: self.api.reg(),
85 api: self.api.as_ref(),
86 }
87 }
88
89 #[inline]
90 pub fn bind(&self) -> PipelineStateBinding<St, Req, Res, Err> {
94 PipelineStateBinding {
95 idx: self.api.reg(),
96 api: self.api.clone(),
97 }
98 }
99
100 #[inline]
101 pub fn bind_state(&self, st: St) -> PipelineBinding<Req, Res, Err>
105 where
106 St: Clone,
107 {
108 let internal = PipelineInternal {
109 st,
110 api: self.api.clone(),
111 };
112
113 PipelineBinding::with(self.api.reg(), PipelineApi::with(internal))
114 }
115}
116
117impl<St, Req, Res, Err> Drop for PipelineState<St, Req, Res, Err> {
118 #[inline]
119 fn drop(&mut self) {
120 self.api.unreg(0);
121 }
122}
123
124struct Binding<'a, St, Req, Res, Err> {
125 idx: u32,
126 api: &'a dyn PipelineStateApi<St, Req, Res, Err>,
127}
128
129impl<St, Req, Res, Err> Drop for Binding<'_, St, Req, Res, Err> {
130 #[inline]
131 fn drop(&mut self) {
132 self.api.unreg(self.idx);
133 }
134}
135
136pub struct PipelineStateBinding<St, Req, Res, Err> {
140 idx: u32,
141 api: Rc<dyn PipelineStateApi<St, Req, Res, Err>>,
142}
143
144impl<St, Req, Res, Err> Drop for PipelineStateBinding<St, Req, Res, Err> {
145 #[inline]
146 fn drop(&mut self) {
147 self.api.unreg(self.idx);
148 }
149}
150
151impl<St, Req, Res, Err> Clone for PipelineStateBinding<St, Req, Res, Err> {
152 #[inline]
153 fn clone(&self) -> Self {
154 PipelineStateBinding {
155 idx: self.api.reg(),
156 api: self.api.clone(),
157 }
158 }
159}
160
161impl<St, Req, Res, Err> PipelineStateBinding<St, Req, Res, Err>
162where
163 St: 'static,
164 Req: 'static,
165 Res: 'static,
166 Err: 'static,
167{
168 #[inline]
169 pub async fn call(&self, req: Req, st: &St) -> Result<Res, Err> {
174 let pl = Binding {
175 idx: self.api.reg(),
176 api: self.api.as_ref(),
177 };
178 pl.api.call(pl.idx, req, st).await
179 }
180}
181
182struct PipelineInternal<St, Req, Res, Err> {
185 st: St,
186 api: Rc<dyn PipelineStateApi<St, Req, Res, Err>>,
187}
188
189impl<St, Req, Res, Err> PipelineInternalApi<Req, Res, Err> for PipelineInternal<St, Req, Res, Err> {
190 fn reg(&self) -> u32 {
191 self.api.reg()
192 }
193
194 fn unreg(&self, idx: u32) {
195 self.api.unreg(idx);
196 }
197
198 fn ready(&self, idx: u32) -> BoxFuture<'_, Result<(), Err>> {
199 self.api.ready(idx, &self.st)
200 }
201
202 fn call(&self, idx: u32, req: Req) -> BoxFuture<'_, Result<Res, Err>> {
203 self.api.call(idx, req, &self.st)
204 }
205
206 fn poll_ready(&self, _: &mut Context<'_>) -> Poll<Result<(), Err>> {
207 unreachable!()
208 }
209
210 fn poll_shutdown(&self, _: &mut Context<'_>) -> Poll<()> {
211 unreachable!()
212 }
213
214 fn is_shutdown(&self) -> bool {
215 self.api.is_shutdown()
216 }
217}
218
219struct PipelineInner<S, St, E> {
222 s: S,
223 waiters: WaitersRef,
224 st_runtime: cell::UnsafeCell<RuntimeState<St, E>>,
225}
226
227impl<S, St, E> Drop for PipelineInner<S, St, E> {
228 fn drop(&mut self) {
229 *self.st_runtime.get_mut() = RuntimeState::New;
232 }
233}
234
235enum RuntimeState<St, E> {
236 New,
237 Readiness(Box<dyn CheckReadiness<St, E>>),
238 Shutdown,
239}
240
241trait PipelineStateApi<St, Req, Res, Err> {
242 fn reg(&self) -> u32;
243 fn unreg(&self, idx: u32);
244
245 fn call<'a>(&'a self, idx: u32, req: Req, st: &'a St) -> BoxFuture<'a, Result<Res, Err>>
246 where
247 Req: 'a;
248
249 fn ready<'a>(&'a self, idx: u32, st: &'a St) -> BoxFuture<'a, Result<(), Err>>
250 where
251 Req: 'a;
252
253 fn poll_ready(&self, cx: &mut Context<'_>, st: &St) -> Poll<Result<(), Err>>
254 where
255 St: Clone;
256
257 fn shutdown<'a>(&'a self, idx: u32, st: &'a St) -> BoxFuture<'a, ()>;
258
259 fn is_shutdown(&self) -> bool;
260}
261
262impl<S, St, Req, E> PipelineStateApi<St, Req, S::Res, S::Error> for PipelineInner<S, St, E>
263where
264 S: Service<St, Req, Error = E> + 'static,
265 St: 'static,
266 Req: 'static,
267 E: 'static,
268{
269 fn reg(&self) -> u32 {
270 self.waiters.insert()
271 }
272
273 fn unreg(&self, idx: u32) {
274 self.waiters.remove(idx);
275 }
276
277 fn ready<'a>(&'a self, idx: u32, st: &'a St) -> BoxFuture<'a, Result<(), S::Error>>
278 where
279 Req: 'a,
280 {
281 Box::pin(async move {
282 self.waiters.set_ready(false);
283 let result = Ctx::<'_, S, St>::new(idx, &self.waiters, st)
284 .ready(&self.s)
285 .await;
286 self.waiters.set_ready(result.is_ok());
287 result
288 })
289 }
290
291 fn shutdown<'a>(&'a self, idx: u32, st: &'a St) -> BoxFuture<'a, ()> {
292 Box::pin(async move {
293 let pl_state = unsafe { &mut *self.st_runtime.get() };
294 *pl_state = RuntimeState::Shutdown;
295 self.waiters.set_ready(false);
296
297 Ctx::<'_, S, St>::new(idx, &self.waiters, st)
298 .shutdown(&self.s)
299 .await;
300 })
301 }
302
303 fn call<'a>(&'a self, idx: u32, req: Req, st: &'a St) -> BoxFuture<'a, Result<S::Res, S::Error>>
304 where
305 Req: 'a,
306 {
307 Box::pin(async move {
308 let ctx = Ctx::<'_, S, St>::new(idx, &self.waiters, st);
309 if !self.waiters.take_ready() {
310 let result = ctx.ready(&self.s).await;
311 self.waiters.set_ready(false);
313 result?;
314 }
315 ctx.call_nowait(&self.s, req).await
316 })
317 }
318
319 fn poll_ready(&self, cx: &mut Context<'_>, st: &St) -> Poll<Result<(), S::Error>>
320 where
321 St: Clone,
322 {
323 let pl_state = unsafe { &mut *self.st_runtime.get() };
324 match pl_state {
325 RuntimeState::New => {
326 let pl = unsafe { &*(ptr::from_ref(self)) };
330 let fut = Box::new(CheckReadinessFut {
331 pl,
332 f: ready,
333 st: st.clone(),
334 fut: None,
335 });
336 *pl_state = RuntimeState::Readiness(fut);
337 self.poll_ready(cx, st)
338 }
339 RuntimeState::Readiness(fut) => fut.poll(cx, st),
340 RuntimeState::Shutdown => panic!("Pipeline is shutting down"),
341 }
342 }
343
344 fn is_shutdown(&self) -> bool {
345 self.waiters.is_shutdown()
346 }
347}
348
349trait CheckReadiness<St, E> {
350 fn poll(&mut self, cx: &mut Context<'_>, st: &St) -> Poll<Result<(), E>>;
351}
352
353struct CheckReadinessFut<S, St, Req, F, Fut>
354where
355 S: Service<St, Req> + 'static,
356 St: 'static,
357 Req: 'static,
358{
359 f: F,
360 st: St,
361 fut: Option<Fut>,
362 pl: &'static PipelineInner<S, St, S::Error>,
363}
364
365fn ready<S, St, Req>(
366 st: &'static St,
367 pl: &'static PipelineInner<S, St, S::Error>,
368) -> impl future::Future<Output = Result<(), S::Error>>
369where
370 S: Service<St, Req>,
371{
372 pl.s.ready(Ctx::<'_, S, St>::new(0, &pl.waiters, st))
373}
374
375impl<S: Service<St, Req>, St, Req, F, Fut> Drop for CheckReadinessFut<S, St, Req, F, Fut> {
376 fn drop(&mut self) {
377 if self.fut.is_some() {
379 self.pl.waiters.notify();
380 }
381 }
382}
383
384impl<S, St, Req, F, Fut> CheckReadiness<St, S::Error> for CheckReadinessFut<S, St, Req, F, Fut>
385where
386 St: Clone,
387 S: Service<St, Req>,
388 F: Fn(&'static St, &'static PipelineInner<S, St, S::Error>) -> Fut,
389 Fut: Future<Output = Result<(), S::Error>>,
390{
391 fn poll(&mut self, cx: &mut Context<'_>, st: &St) -> Poll<Result<(), S::Error>> {
392 let result = self.pl.waiters.run(0, cx, |cx| {
393 if self.fut.is_none() {
394 self.st = st.clone();
395 let st: &'static St = unsafe { std::mem::transmute(&self.st) };
396 self.fut = Some((self.f)(st, self.pl));
397 }
398 let fut = self.fut.as_mut().unwrap();
399 let result = unsafe { Pin::new_unchecked(fut) }.poll(cx);
400 if result.is_ready() {
401 let _ = self.fut.take();
402 }
403 result
404 });
405 self.pl
406 .waiters
407 .set_ready(matches!(result, Poll::Ready(Ok(()))));
408 result
409 }
410}
411
412impl<St, Req, Res, Err> fmt::Debug for PipelineState<St, Req, Res, Err> {
413 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
414 f.debug_struct("PipelineState").finish()
415 }
416}
417
418impl<St, Req, Res, Err> fmt::Debug for PipelineStateBinding<St, Req, Res, Err> {
419 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
420 f.debug_struct("PipelineStateBinding").finish()
421 }
422}
423
424#[cfg(test)]
425mod tests {
426 use std::{cell::Cell, future::pending, task::Waker};
427
428 use ntex::{channel::condition, util::lazy};
429
430 use super::*;
431
432 struct Srv(Rc<Cell<usize>>, condition::Waiter);
433
434 impl Service<usize, usize> for Srv {
435 type Res = usize;
436 type Error = ();
437
438 async fn ready(&self, _: Ctx<'_, Self, usize>) -> Result<(), ()> {
439 self.0.set(self.0.get() + 1);
440 self.1.ready().await;
441 Ok(())
442 }
443
444 async fn call(&self, req: usize, ctx: Ctx<'_, Self, usize>) -> Result<usize, ()> {
445 if req == 0 { Err(()) } else { Ok(req + *ctx.st()) }
446 }
447
448 async fn shutdown(&self, ctx: Ctx<'_, Self, usize>) {
449 self.0.set(self.0.get() + 100 * *ctx);
450 }
451 }
452
453 #[ntex::test]
454 async fn pipeline_state() {
455 let cnt = Rc::new(Cell::new(0));
456 let cond = condition::Condition::new();
457 let pl = PipelineState::new(Srv(cnt.clone(), cond.wait()));
458 assert!(format!("{pl:?}").contains("PipelineState"));
459
460 cond.notify_and_lock(());
461 assert_eq!(pl.ready(&1).await, Ok(()));
462 assert_eq!(cnt.get(), 1);
463 assert_eq!(pl.call(1, &2).await, Ok(3));
465 assert_eq!(cnt.get(), 1);
466 assert_eq!(pl.call(0, &2).await, Err(()));
467 assert_eq!(pl.call(2, &3).await, Ok(5));
468 assert_eq!(cnt.get(), 3);
469
470 let b = pl.bind();
471 assert!(format!("{b:?}").contains("PipelineStateBinding"));
472 let b2 = b.clone();
473 drop(b);
474 assert_eq!(b2.call(1, &10).await, Ok(11));
475 assert_eq!(cnt.get(), 4);
476 assert_eq!(pl.ready(&1).await, Ok(()));
477 assert_eq!(b2.call(1, &20).await, Ok(21));
478 assert_eq!(cnt.get(), 5);
479
480 let b = pl.bind_state(7);
481 assert_eq!(b.ready().await, Ok(()));
482 assert_eq!(b.call(1).await, Ok(8));
483 assert_eq!(cnt.get(), 6);
484 assert_eq!(b.call(2).await, Ok(9));
485 assert_eq!(b.clone().call_static(3).await, Ok(10));
486 assert_eq!(cnt.get(), 8);
487 drop(b);
488
489 assert_eq!(lazy(|cx| pl.poll_ready(cx, &1)).await, Poll::Ready(Ok(())));
490 assert_eq!(cnt.get(), 9);
491 assert_eq!(pl.call(1, &2).await, Ok(3));
492 assert_eq!(cnt.get(), 9);
493
494 assert_eq!(pl.ready(&1).await, Ok(()));
496 pl.shutdown(&2).await;
497 assert_eq!(cnt.get(), 210);
498 assert_eq!(pl.call(1, &2).await, Ok(3));
499 assert_eq!(cnt.get(), 211);
500 }
501
502 #[ntex::test]
503 async fn pipeline_state_poll_ready() {
504 let cnt = Rc::new(Cell::new(0));
505 let cond = condition::Condition::new();
506 let pl = PipelineState::new(Srv(cnt.clone(), cond.wait()));
507
508 assert!(lazy(|cx| pl.poll_ready(cx, &1)).await.is_pending());
509 assert!(lazy(|cx| pl.poll_ready(cx, &1)).await.is_pending());
510 assert_eq!(cnt.get(), 1);
511
512 let b = pl.bind_state(1);
514 let mut fut = Box::pin(b.ready());
515 assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
516 assert_eq!(cnt.get(), 1);
517
518 cond.notify(());
519 assert_eq!(lazy(|cx| pl.poll_ready(cx, &1)).await, Poll::Ready(Ok(())));
520 assert_eq!(cnt.get(), 1);
521
522 assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
524 assert_eq!(cnt.get(), 2);
525 assert!(lazy(|cx| pl.poll_ready(cx, &1)).await.is_pending());
526 assert_eq!(cnt.get(), 2);
527
528 drop(fut);
530 assert!(lazy(|cx| pl.poll_ready(cx, &1)).await.is_pending());
531 assert_eq!(cnt.get(), 3);
532 }
533
534 #[ntex::test]
535 async fn pipeline_state_ready_flag() {
536 let cnt = Rc::new(Cell::new(0));
537 let cond = condition::Condition::new();
538 let pl = PipelineState::new(Srv(cnt.clone(), cond.wait()));
539
540 let mut fut = Box::pin(pl.ready(&1));
541 assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
542 cond.notify(());
543 assert_eq!(fut.await, Ok(()));
544 assert_eq!(cnt.get(), 1);
545
546 assert!(lazy(|cx| pl.poll_ready(cx, &1)).await.is_pending());
548 assert_eq!(cnt.get(), 2);
549 let mut call = Box::pin(pl.call(1, &2));
550 assert!(lazy(|cx| call.as_mut().poll(cx)).await.is_pending());
551 assert_eq!(cnt.get(), 2);
552
553 cond.notify(());
554 assert_eq!(lazy(|cx| pl.poll_ready(cx, &1)).await, Poll::Ready(Ok(())));
555
556 assert!(lazy(|cx| call.as_mut().poll(cx)).await.is_pending());
558 assert_eq!(cnt.get(), 3);
559 cond.notify(());
560 assert_eq!(lazy(|cx| call.as_mut().poll(cx)).await, Poll::Ready(Ok(3)));
561 assert_eq!(cnt.get(), 3);
562
563 let mut call = Box::pin(pl.call(1, &2));
565 assert!(lazy(|cx| call.as_mut().poll(cx)).await.is_pending());
566 assert_eq!(cnt.get(), 4);
567 }
568
569 #[ntex::test]
570 #[should_panic(expected = "Pipeline is shutting down")]
571 async fn pipeline_state_poll_ready_after_shutdown() {
572 let cond = condition::Condition::new();
573 let pl = PipelineState::new(Srv(Rc::default(), cond.wait()));
574 pl.shutdown(&1).await;
575 let _ = lazy(|cx| pl.poll_ready(cx, &1)).await;
576 }
577
578 struct Guard<'a>(&'a [usize]);
579
580 impl Drop for Guard<'_> {
581 fn drop(&mut self) {
582 assert_eq!(self.0.iter().sum::<usize>(), 3);
583 }
584 }
585
586 struct Pending(Vec<usize>);
587
588 impl Service<usize, ()> for Pending {
589 type Res = ();
590 type Error = ();
591
592 async fn ready(&self, _: Ctx<'_, Self, usize>) -> Result<(), ()> {
593 let _g = Guard(&self.0);
594 pending().await
595 }
596
597 async fn call(&self, (): (), _: Ctx<'_, Self, usize>) -> Result<(), ()> {
598 Ok(())
599 }
600 }
601
602 #[test]
603 fn miri_drop_with_pending_readiness() {
604 let mut cx = Context::from_waker(Waker::noop());
605
606 let pl = PipelineState::new(Pending(vec![1, 2]));
607 assert!(pl.poll_ready(&mut cx, &1).is_pending());
608 drop(pl);
609 }
610}