Skip to main content

dnet_utils/
map.rs

1//! Mapping transports into other message types.
2
3#![allow(clippy::type_complexity)]
4
5use std::{
6    convert::identity,
7    marker::PhantomData,
8    pin::Pin,
9    task::{Context, Poll},
10};
11
12use futures::{stream::FusedStream, Sink, SinkExt, Stream, StreamExt};
13use pin_project::pin_project;
14
15/// Trait for mapping input into output.
16pub trait Mapper<Item> {
17    /// Output type.
18    type Output;
19
20    /// Map item into output.
21    fn map(&mut self, item: Item) -> Self::Output;
22}
23
24impl<T, I, O> Mapper<I> for T
25where
26    T: FnMut(I) -> O,
27{
28    type Output = O;
29
30    fn map(&mut self, item: I) -> Self::Output {
31        (self)(item)
32    }
33}
34
35/// Trait allowing transports to be mapped into other message types.
36pub trait Map<Incoming, Outgoing, Error>:
37    dnet_base::Transport<Incoming, Outgoing, Error> + Sized + Unpin
38where
39    Error: std::error::Error,
40{
41    /// Map outgoing massages.
42    fn map<O, Mapper>(
43        self,
44        mapper: Mapper,
45    ) -> Mapping<Self, Incoming, O, fn(Incoming) -> Incoming, Mapper, Error>
46    where
47        Mapper: self::Mapper<O, Output = Outgoing>,
48    {
49        self.map_and_unmap(mapper, identity)
50    }
51
52    /// Map incoming messages.
53    fn unmap<I, Unmapper>(
54        self,
55        unmapper: Unmapper,
56    ) -> Mapping<Self, Incoming, Outgoing, Unmapper, fn(Outgoing) -> Outgoing, Error>
57    where
58        Unmapper: self::Mapper<Incoming, Output = I>,
59    {
60        self.map_and_unmap(identity, unmapper)
61    }
62
63    /// Map outgoing messages and unmap incoming messages.
64    fn map_and_unmap<I, O, Unmapper, Mapper>(
65        self,
66        mapper: Mapper,
67        unmapper: Unmapper,
68    ) -> Mapping<Self, Incoming, O, Unmapper, Mapper, Error>
69    where
70        Mapper: self::Mapper<O, Output = Outgoing>,
71        Unmapper: self::Mapper<Incoming, Output = I>,
72    {
73        Mapping {
74            inner: self,
75            mapper,
76            unmapper,
77
78            #[cfg(feature = "logging")]
79            logger: dnet_base::Logger::new::<Self>(),
80
81            _incoming: PhantomData,
82            _outgoing: PhantomData,
83            _error: PhantomData,
84        }
85    }
86}
87
88impl<T, Incoming, Outgoing, Error> Map<Incoming, Outgoing, Error> for T
89where
90    T: dnet_base::Transport<Incoming, Outgoing, Error> + Unpin,
91    Error: std::error::Error,
92{
93}
94
95/// Wrapper transport mapping outgoing and/or incoming messages into other message types.
96#[pin_project]
97pub struct Mapping<Transport, Incoming, Outgoing, Unmapper, Mapper, Error>
98where
99    Transport: dnet_base::Transport<Incoming, Mapper::Output, Error> + Unpin,
100    Mapper: self::Mapper<Outgoing>,
101    Unmapper: self::Mapper<Incoming>,
102{
103    inner: Transport,
104    mapper: Mapper,
105    unmapper: Unmapper,
106
107    #[cfg(feature = "logging")]
108    logger: dnet_base::Logger,
109
110    _incoming: PhantomData<Incoming>,
111    _outgoing: PhantomData<Outgoing>,
112    _error: PhantomData<Error>,
113}
114
115impl<Transport, Incoming, Outgoing, Unmapper, Mapper, Error> Sink<Outgoing>
116    for Mapping<Transport, Incoming, Outgoing, Unmapper, Mapper, Error>
117where
118    Transport: dnet_base::Transport<Incoming, Mapper::Output, Error> + Unpin,
119    Mapper: self::Mapper<Outgoing>,
120    Unmapper: self::Mapper<Incoming>,
121    Error: std::error::Error,
122{
123    type Error = dnet_base::Error<Error>;
124
125    fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
126        let me = self.project();
127        let result = me.inner.poll_ready_unpin(cx);
128
129        #[cfg(feature = "logging")]
130        me.logger.log_ready(&result);
131
132        result
133    }
134
135    fn start_send(self: Pin<&mut Self>, item: Outgoing) -> Result<(), Self::Error> {
136        let me = self.project();
137        let item = me.mapper.map(item);
138        let result = me.inner.start_send_unpin(item);
139
140        #[cfg(feature = "logging")]
141        match &result {
142            Ok(_) => me.logger.log_message_preparation_success::<Outgoing>(None),
143            Err(error) => me.logger.log_sending_failure(error),
144        }
145
146        result
147    }
148
149    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
150        let me = self.project();
151        let result = me.inner.poll_flush_unpin(cx);
152
153        #[cfg(feature = "logging")]
154        me.logger.log_flush(&result);
155
156        result
157    }
158
159    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
160        let me = self.project();
161        let result = me.inner.poll_close_unpin(cx);
162
163        #[cfg(feature = "logging")]
164        me.logger.log_close(&result);
165
166        result
167    }
168}
169
170impl<Transport, Incoming, Outgoing, Unmapper, Mapper, Error> Stream
171    for Mapping<Transport, Incoming, Outgoing, Unmapper, Mapper, Error>
172where
173    Transport: dnet_base::Transport<Incoming, Mapper::Output, Error> + Unpin,
174    Mapper: self::Mapper<Outgoing>,
175    Unmapper: self::Mapper<Incoming>,
176    Error: std::error::Error,
177{
178    type Item = Result<Unmapper::Output, Error>;
179
180    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
181        let me = self.project();
182        let result = me
183            .inner
184            .poll_next_unpin(cx)
185            .map_ok(|item| me.unmapper.map(item));
186
187        #[cfg(feature = "logging")]
188        me.logger.log_receiving(&result, None);
189
190        result
191    }
192}
193
194impl<Transport, Incoming, Outgoing, Unmapper, Mapper, Error> FusedStream
195    for Mapping<Transport, Incoming, Outgoing, Unmapper, Mapper, Error>
196where
197    Transport: dnet_base::Transport<Incoming, Mapper::Output, Error> + FusedStream + Unpin,
198    Mapper: self::Mapper<Outgoing>,
199    Unmapper: self::Mapper<Incoming>,
200    Error: std::error::Error,
201{
202    fn is_terminated(&self) -> bool {
203        self.inner.is_terminated()
204    }
205}
206
207#[cfg(feature = "logging")]
208impl<Transport, Incoming, Outgoing, Unmapper, Mapper, Error> dnet_base::Logging
209    for Mapping<Transport, Incoming, Outgoing, Unmapper, Mapper, Error>
210where
211    Transport: dnet_base::Transport<Incoming, Mapper::Output, Error> + dnet_base::Logging + Unpin,
212    Mapper: self::Mapper<Outgoing>,
213    Unmapper: self::Mapper<Incoming>,
214    Error: std::error::Error,
215{
216    const KIND: &'static str = "Map";
217
218    fn with_logger<F, R>(&self, f: F) -> R
219    where
220        F: FnOnce(&dnet_base::Logger) -> R,
221    {
222        f(&self.logger)
223    }
224
225    fn with_logger_mut<F, R>(&mut self, f: F) -> R
226    where
227        F: FnOnce(&mut dnet_base::Logger) -> R,
228    {
229        f(&mut self.logger)
230    }
231}
232
233#[cfg(test)]
234mod tests {
235    use dnet_base::Receive;
236    use dnet_tests::{dtest, dtest_configure};
237    use futures::SinkExt;
238
239    use crate::channel::transports;
240
241    dtest_configure!();
242
243    use super::Map;
244    #[derive(Debug, PartialEq, Eq)]
245    struct Wrapper<T>(pub T);
246
247    impl<T> Wrapper<T> {
248        fn unwrap(self) -> T {
249            self.0
250        }
251    }
252
253    #[dtest]
254    async fn test_map() {
255        let (left, right) = transports();
256        let mut left = left.map(Wrapper);
257        let mut right = right.map_and_unmap(Wrapper, Wrapper::unwrap);
258
259        dnet_tests::init_logging(&mut left, &mut right);
260
261        left.send(30).await.unwrap();
262        right.send("Hello".to_string()).await.unwrap();
263
264        assert_eq!(left.receive().await.unwrap(), Wrapper("Hello".to_string()));
265        assert_eq!(right.receive().await.unwrap(), 30);
266    }
267
268    #[dtest]
269    async fn test_map_and_unmap() {
270        let (left, right) = transports();
271        let left = left.map_and_unmap(Wrapper, Wrapper::unwrap);
272        let right = right.map_and_unmap(Wrapper, Wrapper::unwrap);
273        dnet_tests::test_transport(left, right).await;
274    }
275
276    #[dtest]
277    async fn test_map_and_unmap_unit_message() {
278        let (left, right) = transports();
279        let left = left.map_and_unmap(Wrapper, Wrapper::unwrap);
280        let right = right.map_and_unmap(Wrapper, Wrapper::unwrap);
281        dnet_tests::test_unit_message(left, right).await;
282    }
283
284    #[dtest]
285    async fn test_map_and_unmap_stream() {
286        let (left, right) = transports();
287        let left = left.map_and_unmap(Wrapper, Wrapper::unwrap);
288        let right = right.map_and_unmap(Wrapper, Wrapper::unwrap);
289        dnet_tests::test_stream(left, right).await;
290    }
291}