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#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
13pub struct Decoded<T> {
14 pub item: Option<T>,
16 pub remains: usize,
18 pub consumed: usize,
20}
21
22pub(crate) struct Extensions(Cell<Option<Box<ExtensionsInner>>>);
23
24#[derive(Default)]
25pub(crate) struct ExtensionsInner {
26 waiters: Waiters,
28 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 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 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 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 pub(super) fn wake(&self, tag: usize) {
102 self.with_opt(|inner| inner.waiters.wake(tag));
103 }
104
105 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
178pub(crate) struct WriteDeadline {
180 timeout: Seconds,
181 sleep: Option<Sleep>,
182}
183
184impl WriteDeadline {
185 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 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}