1#![allow(clippy::type_complexity)]
4
5use std::{
6 marker::PhantomData,
7 pin::Pin,
8 task::{Context, Poll},
9};
10
11use futures::{ready, stream::FusedStream, Sink, SinkExt, Stream, StreamExt};
12use pin_project::pin_project;
13
14pub trait Filter<Item> {
16 fn filter(&mut self, item: &Item) -> bool;
20}
21
22impl<T, I> Filter<I> for T
23where
24 T: FnMut(&I) -> bool,
25{
26 fn filter(&mut self, item: &I) -> bool {
27 (self)(item)
28 }
29}
30
31pub trait FilterExt<Incoming, Outgoing, Error>:
33 dnet_base::Transport<Incoming, Outgoing, Error> + Sized + Unpin
34where
35 Error: std::error::Error,
36{
37 fn filter_outgoing<Filter>(
39 self,
40 filter: Filter,
41 ) -> Filtering<Self, Incoming, Outgoing, fn(&Incoming) -> bool, Filter, Error>
42 where
43 Filter: self::Filter<Outgoing>,
44 {
45 self.filter_incoming_and_outgoing(|_| true, filter)
46 }
47
48 fn filter_incoming<Filter>(
50 self,
51 filter: Filter,
52 ) -> Filtering<Self, Incoming, Outgoing, Filter, fn(&Outgoing) -> bool, Error>
53 where
54 Filter: self::Filter<Incoming>,
55 {
56 self.filter_incoming_and_outgoing(filter, |_| true)
57 }
58
59 fn filter_incoming_and_outgoing<IncomingFilter, OutgoingFilter>(
61 self,
62 incoming_filter: IncomingFilter,
63 outgoing_filter: OutgoingFilter,
64 ) -> Filtering<Self, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
65 where
66 IncomingFilter: self::Filter<Incoming>,
67 OutgoingFilter: self::Filter<Outgoing>,
68 {
69 Filtering {
70 inner: self,
71 incoming_filter,
72 outgoing_filter,
73
74 #[cfg(feature = "logging")]
75 logger: dnet_base::Logger::new::<
76 Filtering<Self, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>,
77 >(),
78
79 _incoming: PhantomData,
80 _outgoing: PhantomData,
81 _error: PhantomData,
82 }
83 }
84}
85
86impl<T, Incoming, Outgoing, Error> FilterExt<Incoming, Outgoing, Error> for T
87where
88 T: dnet_base::Transport<Incoming, Outgoing, Error> + Unpin,
89 Error: std::error::Error,
90{
91}
92
93#[pin_project]
95pub struct Filtering<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
96where
97 Transport: dnet_base::Transport<Incoming, Outgoing, Error> + Unpin,
98 IncomingFilter: self::Filter<Incoming>,
99 OutgoingFilter: self::Filter<Outgoing>,
100 Error: std::error::Error,
101{
102 inner: Transport,
103 incoming_filter: IncomingFilter,
104 outgoing_filter: OutgoingFilter,
105
106 #[cfg(feature = "logging")]
107 logger: dnet_base::Logger,
108
109 _incoming: PhantomData<Incoming>,
110 _outgoing: PhantomData<Outgoing>,
111 _error: PhantomData<Error>,
112}
113
114impl<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error> Sink<Outgoing>
115 for Filtering<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
116where
117 Transport: dnet_base::Transport<Incoming, Outgoing, Error> + Unpin,
118 IncomingFilter: self::Filter<Incoming>,
119 OutgoingFilter: self::Filter<Outgoing>,
120 Error: std::error::Error,
121{
122 type Error = dnet_base::Error<Error>;
123
124 fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
125 let me = self.project();
126 let result = me.inner.poll_ready_unpin(cx);
127
128 #[cfg(feature = "logging")]
129 me.logger.log_ready(&result);
130
131 result
132 }
133
134 fn start_send(self: Pin<&mut Self>, item: Outgoing) -> Result<(), Self::Error> {
135 let me = self.project();
136 if me.outgoing_filter.filter(&item) {
137 let result = me.inner.start_send_unpin(item);
138
139 #[cfg(feature = "logging")]
140 match &result {
141 Ok(_) => me.logger.log_message_preparation_success::<Outgoing>(None),
142 Err(error) => me.logger.log_sending_failure(error),
143 }
144
145 result
146 } else {
147 #[cfg(feature = "logging")]
148 me.logger.log_outgoing_filtered_out::<Outgoing>();
149
150 Ok(())
151 }
152 }
153
154 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
155 let me = self.project();
156 let result = me.inner.poll_flush_unpin(cx);
157
158 #[cfg(feature = "logging")]
159 me.logger.log_flush(&result);
160
161 result
162 }
163
164 fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
165 let me = self.project();
166 let result = me.inner.poll_close_unpin(cx);
167
168 #[cfg(feature = "logging")]
169 me.logger.log_close(&result);
170
171 result
172 }
173}
174
175impl<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error> Stream
176 for Filtering<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
177where
178 Transport: dnet_base::Transport<Incoming, Outgoing, Error> + Unpin,
179 IncomingFilter: self::Filter<Incoming>,
180 OutgoingFilter: self::Filter<Outgoing>,
181 Error: std::error::Error,
182{
183 type Item = Result<Incoming, Error>;
184
185 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
186 let me = self.project();
187 let result = loop {
188 let result = ready!(me.inner.poll_next_unpin(cx));
189 match result {
190 Some(Ok(item)) => {
191 if me.incoming_filter.filter(&item) {
192 break Poll::Ready(Some(Ok(item)));
193 } else {
194 #[cfg(feature = "logging")]
195 me.logger.log_incoming_filtered_out::<Incoming>();
196
197 continue;
198 }
199 }
200 Some(Err(error)) => break Poll::Ready(Some(Err(error))),
201 None => break Poll::Ready(None),
202 }
203 };
204
205 #[cfg(feature = "logging")]
206 me.logger.log_receiving(&result, None);
207
208 result
209 }
210}
211
212impl<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error> FusedStream
213 for Filtering<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
214where
215 Transport: dnet_base::Transport<Incoming, Outgoing, Error> + FusedStream + Unpin,
216 IncomingFilter: self::Filter<Incoming>,
217 OutgoingFilter: self::Filter<Outgoing>,
218 Error: std::error::Error,
219{
220 fn is_terminated(&self) -> bool {
221 self.inner.is_terminated()
222 }
223}
224
225#[cfg(feature = "logging")]
226impl<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error> dnet_base::Logging
227 for Filtering<Transport, Incoming, Outgoing, IncomingFilter, OutgoingFilter, Error>
228where
229 Transport: dnet_base::Transport<Incoming, Outgoing, Error> + dnet_base::Logging + Unpin,
230 IncomingFilter: self::Filter<Incoming>,
231 OutgoingFilter: self::Filter<Outgoing>,
232 Error: std::error::Error,
233{
234 const KIND: &'static str = "Filtering";
235
236 fn with_logger<F, R>(&self, f: F) -> R
237 where
238 F: FnOnce(&dnet_base::Logger) -> R,
239 {
240 f(&self.logger)
241 }
242
243 fn with_logger_mut<F, R>(&mut self, f: F) -> R
244 where
245 F: FnOnce(&mut dnet_base::Logger) -> R,
246 {
247 f(&mut self.logger)
248 }
249}
250
251#[cfg(test)]
252mod tests {
253 use dnet_base::Messages;
254 use dnet_tests::{dtest, dtest_configure};
255 use futures::{stream, SinkExt, StreamExt};
256
257 use crate::channel::transports;
258
259 use super::FilterExt;
260
261 dtest_configure!();
262
263 #[dtest]
264 async fn test_filter() {
265 let (mut left, right) = transports();
266 let mut right = right.filter_incoming_and_outgoing(
267 |integer: &u32| integer % 2 == 0,
268 |string: &String| !string.starts_with("A"),
269 );
270
271 dnet_tests::init_logging(&mut left, &mut right);
272
273 left.send(1).await.unwrap();
274 left.send(2).await.unwrap();
275 left.send(3).await.unwrap();
276 left.send(4).await.unwrap();
277 left.send(5).await.unwrap();
278 left.close().await.unwrap();
279
280 right
281 .send_all(&mut stream::iter(
282 vec!["Anna", "Tom", "Albert", "Robert"]
283 .into_iter()
284 .map(String::from)
285 .map(Ok),
286 ))
287 .await
288 .unwrap();
289 right.close().await.unwrap();
290
291 assert_eq!(
292 vec!["Tom", "Robert"],
293 left.messages().collect::<Vec<String>>().await
294 );
295
296 assert_eq!(vec![2, 4], right.messages().collect::<Vec<u32>>().await);
297 }
298}