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#[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}