Skip to main content

ntex_io/
utils.rs

1use std::cell::Cell;
2use std::io;
3use std::task::{Context, Poll, Waker};
4
5use ntex_service::state::{RequestState, State};
6use ntex_util::time::{Seconds, Sleep};
7
8use crate::waiters::{WaiterEntry, Waiters};
9use crate::{Filter, Io, IoBoxed, IoCallbacks};
10
11/// Result of a single decode attempt.
12#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
13pub struct Decoded<T> {
14    /// The decoded item, or `None` when the codec needs more input.
15    pub item: Option<T>,
16    /// Bytes left in the application-facing read buffer after the attempt.
17    pub remains: usize,
18    /// Bytes consumed from the read buffer by the attempt.
19    pub consumed: usize,
20}
21
22pub(crate) struct Extensions(Cell<Option<Box<ExtensionsInner>>>);
23
24#[derive(Default)]
25pub(crate) struct ExtensionsInner {
26    // tasks waiting for io events, by tag
27    waiters: Waiters,
28    // filter callbacks registered for io events
29    pub(crate) callbacks: Option<Box<dyn IoCallbacks>>,
30}
31
32impl Default for Extensions {
33    fn default() -> Extensions {
34        Extensions(Cell::new(None))
35    }
36}
37
38impl Extensions {
39    fn with<F, R>(&self, f: F) -> R
40    where
41        F: FnOnce(&mut ExtensionsInner) -> R,
42    {
43        let mut inner = if let Some(inner) = self.0.take() {
44            inner
45        } else {
46            Box::new(ExtensionsInner::default())
47        };
48        let result = f(&mut inner);
49        self.0.set(Some(inner));
50        result
51    }
52
53    fn with_opt<F>(&self, f: F)
54    where
55        F: FnOnce(&mut ExtensionsInner),
56    {
57        if let Some(mut inner) = self.0.take() {
58            f(&mut inner);
59            self.0.set(Some(inner));
60        }
61    }
62
63    /// Registers the waker of the waiter, a woken waiter gets a new entry.
64    pub(super) fn register_waker(&self, waiter: &WaiterEntry, waker: &Waker) {
65        self.with(|inner| {
66            let waiters = &mut inner.waiters;
67            if !waiter.id.get().is_some_and(|id| waiters.update(id, waker)) {
68                waiter.id.set(Some(waiters.register(waiter.tag, waker)));
69            }
70        });
71    }
72
73    /// Registers the waiter, or reports that its registration was woken.
74    ///
75    /// A reported wake clears the registration, the next poll registers again.
76    pub(super) fn poll_waker(&self, waiter: &WaiterEntry, waker: &Waker) -> Poll<()> {
77        self.with(|inner| {
78            let waiters = &mut inner.waiters;
79            match waiter.id.get() {
80                None => {
81                    waiter.id.set(Some(waiters.register(waiter.tag, waker)));
82                    Poll::Pending
83                }
84                Some(id) if waiters.update(id, waker) => Poll::Pending,
85                Some(_) => {
86                    waiter.id.set(None);
87                    Poll::Ready(())
88                }
89            }
90        })
91    }
92
93    /// Removes the waiter entry unless it is woken.
94    pub(super) fn remove_waker(&self, waiter: &WaiterEntry) {
95        if let Some(id) = waiter.id.take() {
96            self.with_opt(|inner| inner.waiters.remove(id, waiter.tag));
97        }
98    }
99
100    /// Wakes and removes all wakers of the tag.
101    pub(super) fn wake(&self, tag: usize) {
102        self.with_opt(|inner| inner.waiters.wake(tag));
103    }
104
105    /// Wakes and removes all wakers.
106    pub(super) fn wake_all(&self) {
107        self.with_opt(|inner| inner.waiters.wake_all());
108    }
109
110    #[cfg(test)]
111    pub(super) fn wakers_len(&self) -> usize {
112        let mut len = 0;
113        self.with_opt(|inner| len = inner.waiters.len());
114        len
115    }
116
117    pub(super) fn register_filter_callbacks<T: IoCallbacks + 'static>(&self, cb: T) {
118        self.with(|inner| {
119            inner.callbacks = Some(Box::new(cb));
120        });
121    }
122
123    pub(super) fn take_callbacks(&self) -> Option<Box<dyn IoCallbacks>> {
124        let mut callbacks = None;
125        self.with_opt(|inner| callbacks = inner.callbacks.take());
126        callbacks
127    }
128
129    pub(crate) fn with_callbacks<F>(&self, f: F)
130    where
131        F: FnOnce(&dyn IoCallbacks),
132    {
133        self.with_opt(|inner| {
134            if let Some(ref cb) = inner.callbacks {
135                f(cb.as_ref());
136            }
137        });
138    }
139}
140
141impl<F> RequestState<Io<F>> for Io<F> {
142    type State = ();
143
144    #[inline]
145    fn unpack(self) -> ((), Io<F>) {
146        ((), self)
147    }
148}
149
150impl<F: Filter> RequestState<IoBoxed> for Io<F> {
151    type State = ();
152
153    #[inline]
154    fn unpack(self) -> ((), IoBoxed) {
155        ((), self.boxed())
156    }
157}
158
159impl RequestState<IoBoxed> for IoBoxed {
160    type State = ();
161
162    #[inline]
163    fn unpack(self) -> ((), IoBoxed) {
164        ((), self)
165    }
166}
167
168impl<F: Filter, St: 'static> RequestState<IoBoxed> for State<St, Io<F>> {
169    type State = St;
170
171    #[inline]
172    fn unpack(self) -> (St, IoBoxed) {
173        let State { req, state } = self;
174        (state, req.boxed())
175    }
176}
177
178/// Deadline for a wait on output, started lazily on the first wait.
179pub(crate) struct WriteDeadline {
180    timeout: Seconds,
181    sleep: Option<Sleep>,
182}
183
184impl WriteDeadline {
185    /// Creates a deadline, a zero timeout never expires.
186    pub(crate) fn new(timeout: Seconds) -> Self {
187        Self {
188            timeout,
189            sleep: None,
190        }
191    }
192
193    pub(crate) fn poll_expired(&mut self, cx: &mut Context<'_>) -> bool {
194        if self.timeout.is_zero() {
195            false
196        } else {
197            let timeout = self.timeout;
198            self.sleep
199                .get_or_insert_with(|| Sleep::new(timeout.into()))
200                .poll_elapsed(cx)
201                .is_ready()
202        }
203    }
204}
205
206pub(crate) fn write_timed_out() -> io::Error {
207    io::Error::new(io::ErrorKind::TimedOut, "Write timeout")
208}
209
210#[cfg(test)]
211mod tests {
212    use ntex_bytes::BytePageSize;
213    use ntex_service::cfg::SharedCfg;
214
215    use super::*;
216    use crate::{Sealed, buf::Stack, filter::NullFilter, testing::IoTest};
217
218    #[ntex::test]
219    async fn test_null_filter() {
220        let (_, server) = IoTest::create();
221        let io = Io::new(server, SharedCfg::default());
222        let ioref = io.get_ref();
223        let stack = Stack::new(BytePageSize::Size16);
224        assert!(NullFilter.query(std::any::TypeId::of::<()>()).is_none());
225        assert!(
226            stack
227                .with_filter(&ioref, |ctx| NullFilter.shutdown(ctx))
228                .unwrap()
229                .is_ready()
230        );
231        // The chain is gone, so the transport closes the connection
232        // gracefully. `IoContext` escalates to `Terminate` when the connection
233        // was force-closed; `NullFilter` itself cannot see that state.
234        assert_eq!(
235            std::future::poll_fn(|cx| NullFilter.poll_read_ready(cx)).await,
236            crate::Readiness::Close
237        );
238        assert_eq!(
239            std::future::poll_fn(|cx| NullFilter.poll_write_ready(cx)).await,
240            crate::Readiness::Close
241        );
242        assert!(
243            stack
244                .with_filter(&ioref, |ctx| NullFilter.process_write_buf(ctx))
245                .is_ok()
246        );
247        assert_eq!(
248            stack.with_filter(&ioref, |ctx| NullFilter.process_read_buf(ctx).unwrap()),
249            ()
250        );
251    }
252
253    #[ntex::test]
254    async fn request_state_unpack() {
255        use ntex_service::state::{RequestState, State};
256
257        let (_, server) = IoTest::create();
258        let io = Io::from(server);
259        let id = io.id();
260
261        let ((), io) = <Io as RequestState<Io>>::unpack(io);
262        assert_eq!(io.id(), id);
263        let ((), io) = <Io as RequestState<IoBoxed>>::unpack(io);
264        assert_eq!(io.id(), id);
265        let ((), mut io) = <IoBoxed as RequestState<IoBoxed>>::unpack(io);
266        assert_eq!(io.id(), id);
267
268        let st = State {
269            req: Io::<Sealed>::from(io.take()),
270            state: 10u32,
271        };
272        let (state, io) = <_ as RequestState<IoBoxed>>::unpack(st);
273        assert_eq!(state, 10);
274        assert_eq!(io.id(), id);
275    }
276}