Skip to main content

ntex_io/
utils.rs

1use std::{cell::Cell, task::Poll, task::Waker};
2
3use ntex_service::state::{RequestState, State};
4use ntex_util::task::LocalWaker;
5
6use crate::{Filter, Io, IoBoxed, IoCallbacks};
7
8/// Decoded item from buffer
9#[doc(hidden)]
10#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
11pub struct Decoded<T> {
12    pub item: Option<T>,
13    pub remains: usize,
14    pub consumed: usize,
15}
16
17pub(crate) struct Extensions(Cell<Option<Box<ExtensionsInner>>>);
18
19#[derive(Default)]
20pub(crate) struct ExtensionsInner {
21    disconnect: Option<Vec<LocalWaker>>,
22    pub(crate) callbacks: Option<Box<dyn IoCallbacks>>,
23}
24
25impl Default for Extensions {
26    fn default() -> Extensions {
27        Extensions(Cell::new(None))
28    }
29}
30
31impl Extensions {
32    fn with<F, R>(&self, f: F) -> R
33    where
34        F: FnOnce(&mut ExtensionsInner) -> R,
35    {
36        let mut inner = if let Some(inner) = self.0.take() {
37            inner
38        } else {
39            Box::new(ExtensionsInner::default())
40        };
41        let result = f(&mut inner);
42        self.0.set(Some(inner));
43        result
44    }
45
46    fn with_opt<F>(&self, f: F)
47    where
48        F: FnOnce(&mut ExtensionsInner),
49    {
50        if let Some(mut inner) = self.0.take() {
51            f(&mut inner);
52            self.0.set(Some(inner));
53        }
54    }
55
56    pub(super) fn notify_disconnect(&self) {
57        self.with_opt(|inner| {
58            if let Some(disconnect) = inner.disconnect.take() {
59                for item in disconnect {
60                    item.wake();
61                }
62            }
63        });
64    }
65
66    pub(super) fn register_disconnect(&self) -> usize {
67        self.with(|inner| {
68            if let Some(ref mut disconnect) = inner.disconnect {
69                let token = disconnect.len();
70                disconnect.push(LocalWaker::default());
71                token
72            } else {
73                inner.disconnect = Some(vec![LocalWaker::default()]);
74                0
75            }
76        })
77    }
78
79    pub(super) fn poll_disconnect(&self, token: usize, waker: &Waker) -> Poll<()> {
80        self.with(|inner| {
81            if let Some(ref mut disconnect) = inner.disconnect {
82                disconnect[token].register(waker);
83                Poll::Pending
84            } else {
85                Poll::Ready(())
86            }
87        })
88    }
89
90    pub(super) fn register_filter_callbacks<T: IoCallbacks + 'static>(&self, cb: T) {
91        self.with(|inner| {
92            inner.callbacks = Some(Box::new(cb));
93        });
94    }
95
96    pub(crate) fn with_callbacks<F>(&self, f: F)
97    where
98        F: FnOnce(&dyn IoCallbacks),
99    {
100        self.with_opt(|inner| {
101            if let Some(ref cb) = inner.callbacks {
102                f(cb.as_ref());
103            }
104        });
105    }
106}
107
108impl<F> RequestState<Io<F>> for Io<F> {
109    type State = ();
110
111    #[inline]
112    fn unpack(self) -> ((), Io<F>) {
113        ((), self)
114    }
115}
116
117impl<F: Filter> RequestState<IoBoxed> for Io<F> {
118    type State = ();
119
120    #[inline]
121    fn unpack(self) -> ((), IoBoxed) {
122        ((), self.boxed())
123    }
124}
125
126impl RequestState<IoBoxed> for IoBoxed {
127    type State = ();
128
129    #[inline]
130    fn unpack(self) -> ((), IoBoxed) {
131        ((), self)
132    }
133}
134
135impl<F: Filter, St: 'static> RequestState<IoBoxed> for State<St, Io<F>> {
136    type State = St;
137
138    #[inline]
139    fn unpack(self) -> (St, IoBoxed) {
140        let State { req, state } = self;
141        (state, req.boxed())
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use ntex_bytes::BytePageSize;
148    use ntex_service::cfg::SharedCfg;
149
150    use super::*;
151    use crate::{buf::Stack, filter::NullFilter, testing::IoTest};
152
153    #[ntex::test]
154    async fn test_null_filter() {
155        let (_, server) = IoTest::create();
156        let io = Io::new(server, SharedCfg::default());
157        let ioref = io.get_ref();
158        let stack = Stack::new(BytePageSize::Size16);
159        assert!(NullFilter.query(std::any::TypeId::of::<()>()).is_none());
160        assert!(
161            stack
162                .with_filter(&ioref, |ctx| NullFilter.shutdown(ctx))
163                .unwrap()
164                .is_ready()
165        );
166        assert_eq!(
167            std::future::poll_fn(|cx| NullFilter.poll_read_ready(cx)).await,
168            crate::Readiness::Terminate
169        );
170        assert_eq!(
171            std::future::poll_fn(|cx| NullFilter.poll_write_ready(cx)).await,
172            crate::Readiness::Terminate
173        );
174        assert!(
175            stack
176                .with_filter(&ioref, |ctx| NullFilter.process_write_buf(ctx))
177                .is_ok()
178        );
179        assert_eq!(
180            stack.with_filter(&ioref, |ctx| NullFilter.process_read_buf(ctx).unwrap()),
181            ()
182        );
183    }
184}