1#![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
15pub trait Mapper<Item> {
17 type Output;
19
20 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
35pub trait Map<Incoming, Outgoing, Error>:
37 dnet_base::Transport<Incoming, Outgoing, Error> + Sized + Unpin
38where
39 Error: std::error::Error,
40{
41 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 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 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#[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}