Skip to main content

dnet_utils/
number.rs

1//! Wrapper transport attaching a number to messages.
2
3use std::{
4    fmt::{Debug, Display},
5    marker::PhantomData,
6    pin::Pin,
7    task::{Context, Poll},
8};
9
10use futures::{stream::FusedStream, Sink, Stream};
11use num::{traits::bounds::UpperBounded, One, Zero};
12use pin_project::pin_project;
13use serde::{Deserialize, Serialize};
14
15use super::unwrap::Unwrap;
16
17/// Trait for adding message number to transport messages.
18pub trait NumberMessages<N, Incoming, Outgoing, Error>:
19    dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, Error> + Sized + Unpin
20where
21    Error: std::error::Error,
22{
23    /// Wrap transport with [Numbered] transport adding message number of type `N`.
24    fn number_messages(self) -> Numbered<N, Self, Error, Incoming, Outgoing>
25    where
26        N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
27    {
28        Numbered::new(self)
29    }
30}
31
32impl<T, N, Incoming, Outgoing, Error> NumberMessages<N, Incoming, Outgoing, Error> for T
33where
34    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, Error> + Unpin,
35    Error: std::error::Error,
36{
37}
38
39/// Trait for adding [usize] message number to transport messages.
40///
41/// **NOTE**: [usize] size may differ between platforms.
42pub trait NumberMessagesUsize<Incoming, Outgoing, Error>:
43    dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, Error> + Sized + Unpin
44where
45    Error: std::error::Error,
46{
47    /// Wrap transport with [Numbered] transport adding message number of type [`usize`].
48    fn number_messages_u64(self) -> Numbered<usize, Self, Error, Incoming, Outgoing> {
49        Numbered::new(self)
50    }
51}
52
53impl<T, Incoming, Outgoing, Error> NumberMessagesUsize<Incoming, Outgoing, Error> for T
54where
55    T: dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, Error> + Unpin,
56    Error: std::error::Error,
57{
58}
59
60/// Trait for adding [u32] message number to transport messages.
61pub trait NumberMessagesU32<Incoming, Outgoing, Error>:
62    dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, Error> + Sized + Unpin
63where
64    Error: std::error::Error,
65{
66    /// Wrap transport with [Numbered] transport adding message number of type [`u32`].
67    fn number_messages_u32(self) -> Numbered<u32, Self, Error, Incoming, Outgoing> {
68        Numbered::new(self)
69    }
70}
71
72impl<T, Incoming, Outgoing, Error> NumberMessagesU32<Incoming, Outgoing, Error> for T
73where
74    T: dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, Error> + Unpin,
75    Error: std::error::Error,
76{
77}
78
79/// Trait for adding [u64] message number to transport messages.
80pub trait NumberMessagesU64<Incoming, Outgoing, Error>:
81    dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, Error> + Sized + Unpin
82where
83    Error: std::error::Error,
84{
85    /// Wrap transport with [Numbered] transport adding message number of type [`u64`].
86    fn number_messages_u64(self) -> Numbered<u64, Self, Error, Incoming, Outgoing> {
87        Numbered::new(self)
88    }
89}
90
91impl<T, Incoming, Outgoing, Error> NumberMessagesU64<Incoming, Outgoing, Error> for T
92where
93    T: dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, Error> + Unpin,
94    Error: std::error::Error,
95{
96}
97
98/// Trait for adding [u128] message number to transport messages.
99pub trait NumberMessagesU128<Incoming, Outgoing, Error>:
100    dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, Error> + Sized + Unpin
101where
102    Error: std::error::Error,
103{
104    /// Wrap transport with [Numbered] transport adding message number of type [`u128`].
105    fn number_messages_u128(self) -> Numbered<u128, Self, Error, Incoming, Outgoing> {
106        Numbered::new(self)
107    }
108}
109
110impl<T, Incoming, Outgoing, Error> NumberMessagesU128<Incoming, Outgoing, Error> for T
111where
112    T: dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, Error> + Unpin,
113    Error: std::error::Error,
114{
115}
116
117/// [Numbered] transport error.
118#[derive(Debug)]
119pub enum Error<T> {
120    /// Maximum message number reached.
121    MaximumNumberReached,
122
123    /// Wrapped transport error.
124    Transport(T),
125}
126
127impl<T> Display for Error<T>
128where
129    T: Display,
130{
131    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
132        match self {
133            Error::MaximumNumberReached => write!(f, "maximum number reached"),
134            Error::Transport(error) => write!(f, "transport error: {error}"),
135        }
136    }
137}
138
139impl<T> std::error::Error for Error<T> where T: Debug + Display {}
140
141/// Trait implemented by messages with attached number.
142pub trait Number {
143    /// Message number type.
144    type Output;
145
146    /// Message number.
147    fn number(&self) -> Self::Output;
148}
149
150/// Message wrapper used by [`Numbered`] transport.
151#[derive(Debug, PartialEq, Eq, Serialize, Deserialize)]
152pub struct Wrapper<N, T> {
153    /// Message number.
154    pub number: N,
155
156    /// Wrapped message.
157    pub wrapped: T,
158}
159
160impl<N, T> Unwrap for Wrapper<N, T> {
161    type Output = T;
162
163    fn unwrap(self) -> Self::Output {
164        self.wrapped
165    }
166}
167
168impl<N, T> Number for Wrapper<N, T>
169where
170    N: Clone,
171{
172    type Output = N;
173
174    fn number(&self) -> Self::Output {
175        self.number.clone()
176    }
177}
178
179/// Wrapper transport attaching number to sent messages.
180///
181/// Received messages are of [`Wrapper`] type.
182///
183/// First message number is [zero].<br>
184/// Next message number = previous message number + [one].
185///
186/// Sending will result in an [Error::MaximumNumberReached] error when reaching maximum value.
187///
188/// [zero]: num::Zero
189/// [one]: num::One
190#[pin_project]
191pub struct Numbered<N, T, E, Incoming, Outgoing>
192where
193    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
194    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
195{
196    #[pin]
197    inner: T,
198    current_number: N,
199
200    #[cfg(feature = "logging")]
201    logger: dnet_base::Logger,
202
203    _error: PhantomData<E>,
204    _number: PhantomData<N>,
205    _incoming: PhantomData<Incoming>,
206    _outgoing: PhantomData<Outgoing>,
207}
208
209impl<N, T, E, Incoming, Outgoing> Numbered<N, T, E, Incoming, Outgoing>
210where
211    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
212    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
213    E: std::error::Error,
214{
215    /// Create new [numbered] transport wrapping provided transport.
216    ///
217    /// [numbered]: self::Number
218    pub fn new(transport: T) -> Self {
219        Numbered {
220            inner: transport,
221            current_number: N::zero(),
222
223            #[cfg(feature = "logging")]
224            logger: dnet_base::Logger::new::<Self>(),
225
226            _error: PhantomData,
227            _number: PhantomData,
228            _incoming: PhantomData,
229            _outgoing: PhantomData,
230        }
231    }
232
233    /// Number that will be attached to the next sent message.
234    pub fn current_number(&self) -> N {
235        self.current_number.clone()
236    }
237}
238
239impl<T, E, Incoming, Outgoing> Numbered<usize, T, E, Incoming, Outgoing>
240where
241    T: dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, E>,
242    E: std::error::Error,
243{
244    /// Create new [`Numbered`] transport using [`usize`] type as  message number.
245    pub fn new_usize(transport: T) -> Self {
246        Numbered::new(transport)
247    }
248}
249
250impl<T, E, Incoming, Outgoing> Numbered<u32, T, E, Incoming, Outgoing>
251where
252    T: dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, E>,
253    E: std::error::Error,
254{
255    /// Create new [`Numbered`] transport using [`u32`] type as  message number.
256    pub fn new_u32(transport: T) -> Self {
257        Numbered::new(transport)
258    }
259}
260
261impl<T, E, Incoming, Outgoing> Numbered<u64, T, E, Incoming, Outgoing>
262where
263    T: dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, E>,
264    E: std::error::Error,
265{
266    /// Create new [`Numbered`] transport using [`u64`] type as  message number.
267    pub fn new_u64(transport: T) -> Self {
268        Numbered::new(transport)
269    }
270}
271
272impl<T, E, Incoming, Outgoing> Numbered<u128, T, E, Incoming, Outgoing>
273where
274    T: dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, E>,
275    E: std::error::Error,
276{
277    /// Create new [`Numbered`] transport using [`u128`] type as  message number.
278    pub fn new_u128(transport: T) -> Self {
279        Numbered::new(transport)
280    }
281}
282
283impl<N, T, E, Incoming, Outgoing> Sink<Outgoing> for Numbered<N, T, E, Incoming, Outgoing>
284where
285    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
286    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
287    E: std::error::Error,
288{
289    type Error = dnet_base::Error<Error<E>>;
290
291    fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
292        let me = self.project();
293        let result = me.inner.poll_ready(cx).map_err(map_error);
294
295        #[cfg(feature = "logging")]
296        me.logger.log_ready(&result);
297
298        result
299    }
300
301    fn start_send(self: Pin<&mut Self>, item: Outgoing) -> Result<(), Self::Error> {
302        let me = self.project();
303        let result = if *me.current_number == N::max_value() {
304            Err(dnet_base::Error::Other(Error::MaximumNumberReached))
305        } else {
306            let item = Wrapper {
307                number: me.current_number.clone(),
308                wrapped: item,
309            };
310            let result = me.inner.start_send(item);
311            if result.is_ok() {
312                *me.current_number = me.current_number.clone().add(One::one());
313            }
314            result.map_err(map_error)
315        };
316
317        #[cfg(feature = "logging")]
318        match &result {
319            Ok(_) => me.logger.log_message_preparation_success::<Outgoing>(None),
320            Err(error) => me.logger.log_sending_failure(error),
321        }
322
323        result
324    }
325
326    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
327        let me = self.project();
328        let result = me.inner.poll_flush(cx).map_err(map_error);
329
330        #[cfg(feature = "logging")]
331        me.logger.log_flush(&result);
332
333        result
334    }
335
336    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
337        let me = self.project();
338        let result = me.inner.poll_close(cx).map_err(map_error);
339
340        #[cfg(feature = "logging")]
341        me.logger.log_close(&result);
342
343        result
344    }
345}
346
347impl<N, T, E, Incoming, Outgoing> Stream for Numbered<N, T, E, Incoming, Outgoing>
348where
349    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
350    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
351    E: std::error::Error,
352{
353    type Item = Result<Wrapper<N, Incoming>, Error<E>>;
354
355    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
356        let me = self.project();
357        let result = me.inner.poll_next(cx).map_err(Error::Transport);
358
359        #[cfg(feature = "logging")]
360        me.logger.log_receiving(&result, None);
361
362        result
363    }
364}
365
366impl<N, T, E, Incoming, Outgoing> FusedStream for Numbered<N, T, E, Incoming, Outgoing>
367where
368    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E> + FusedStream,
369    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
370    E: std::error::Error,
371{
372    fn is_terminated(&self) -> bool {
373        self.inner.is_terminated()
374    }
375}
376
377#[cfg(feature = "logging")]
378impl<N, T, E, Incoming, Outgoing> dnet_base::Logging for Numbered<N, T, E, Incoming, Outgoing>
379where
380    T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E> + dnet_base::Logging,
381    N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
382    E: std::error::Error,
383{
384    const KIND: &'static str = "Numbered";
385
386    fn with_logger<F, R>(&self, f: F) -> R
387    where
388        F: FnOnce(&dnet_base::Logger) -> R,
389    {
390        f(&self.logger)
391    }
392
393    fn with_logger_mut<F, R>(&mut self, f: F) -> R
394    where
395        F: FnOnce(&mut dnet_base::Logger) -> R,
396    {
397        f(&mut self.logger)
398    }
399}
400
401fn map_error<T>(error: dnet_base::Error<T>) -> dnet_base::Error<Error<T>> {
402    match error {
403        dnet_base::Error::Closed => dnet_base::Error::Closed,
404        dnet_base::Error::Other(error) => dnet_base::Error::Other(Error::Transport(error)),
405    }
406}
407
408#[cfg(test)]
409mod tests {
410    use dnet_base::Receive;
411    use dnet_tests::{dtest, dtest_configure};
412    use futures::SinkExt;
413
414    use crate::{
415        channel::transports,
416        number::{Numbered, Wrapper},
417    };
418
419    dtest_configure!();
420
421    #[dtest]
422    async fn test_transport() {
423        let (left, right) = transports();
424
425        let mut left = Numbered::new_usize(left);
426        let mut right = Numbered::new_usize(right);
427
428        dnet_tests::init_logging(&mut left, &mut right);
429
430        left.send(1).await.unwrap();
431        left.send(2).await.unwrap();
432        left.send(3).await.unwrap();
433
434        assert_eq!(
435            right.receive().await.unwrap(),
436            Wrapper {
437                number: 0,
438                wrapped: 1,
439            }
440        );
441        assert_eq!(
442            right.receive().await.unwrap(),
443            Wrapper {
444                number: 1,
445                wrapped: 2,
446            }
447        );
448        assert_eq!(
449            right.receive().await.unwrap(),
450            Wrapper {
451                number: 2,
452                wrapped: 3,
453            }
454        );
455
456        right.send(1).await.unwrap();
457        right.send(2).await.unwrap();
458        right.send(3).await.unwrap();
459
460        assert_eq!(
461            left.receive().await.unwrap(),
462            Wrapper {
463                number: 0,
464                wrapped: 1,
465            }
466        );
467        assert_eq!(
468            left.receive().await.unwrap(),
469            Wrapper {
470                number: 1,
471                wrapped: 2,
472            }
473        );
474        assert_eq!(
475            left.receive().await.unwrap(),
476            Wrapper {
477                number: 2,
478                wrapped: 3,
479            }
480        );
481    }
482}