Skip to main content

dnet_utils/
filter.rs

1//! Filtering incoming and/or outgoing messages.
2
3#![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
14/// Trait for filtering messages.
15pub trait Filter<Item> {
16    /// Filter item.
17    ///
18    /// Return [`true`] to keep the item, [`false`] to filter it out.
19    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
31/// Trait for filtering incoming and/or outgoing messages.
32pub trait FilterExt<Incoming, Outgoing, Error>:
33    dnet_base::Transport<Incoming, Outgoing, Error> + Sized + Unpin
34where
35    Error: std::error::Error,
36{
37    /// Filter outgoing massages.
38    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    /// Filter incoming messages.
49    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    /// Filter incoming and outgoing massages.
60    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/// Wrapper transport mapping outgoing and/or incoming messages into other message types.
94#[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}