Skip to main content

dnet_utils/
split.rs

1//! Splitting transport into multiple transports of different message types.
2
3#![allow(clippy::type_complexity)]
4
5use std::{
6    collections::VecDeque,
7    fmt::{Debug, Display},
8    marker::PhantomData,
9    pin::Pin,
10    sync::{Arc, Mutex},
11    task::{Context, Poll, Waker},
12};
13
14use futures::{stream::FusedStream, Sink, SinkExt, Stream, StreamExt};
15use pin_project::pin_project;
16use serde::{Deserialize, Serialize};
17
18use super::map::Mapper;
19
20/// Trait allowing transports to be split into two of different message types.
21pub trait Split2<I1, O1, I2, O2, E>:
22    dnet_base::Transport<Message<I1, I2, (), (), ()>, Message<O1, O2, (), (), ()>, E> + Sized + Unpin
23where
24    E: std::error::Error,
25{
26    /// Split into two transports.
27    fn split_into_2(
28        self,
29    ) -> (
30        Part<
31            I1,
32            O1,
33            fn(O1) -> Message<O1, O2, (), (), ()>,
34            fn(Message<I1, I2, (), (), ()>) -> Result<I1, Message<I1, I2, (), (), ()>>,
35            Self,
36            E,
37            I1,
38            I2,
39            (),
40            (),
41            (),
42        >,
43        Part<
44            I2,
45            O2,
46            fn(O2) -> Message<O1, O2, (), (), ()>,
47            fn(Message<I1, I2, (), (), ()>) -> Result<I2, Message<I1, I2, (), (), ()>>,
48            Self,
49            E,
50            I1,
51            I2,
52            (),
53            (),
54            (),
55        >,
56    ) {
57        let state = State::new(2, self);
58        (
59            Part::new(0, &state, Message::Variant1, Message::unwrap1),
60            Part::new(1, &state, Message::Variant2, Message::unwrap2),
61        )
62    }
63}
64
65impl<T, I1, O1, I2, O2, E> Split2<I1, O1, I2, O2, E> for T
66where
67    T: dnet_base::Transport<Message<I1, I2, (), (), ()>, Message<O1, O2, (), (), ()>, E>
68        + Sized
69        + Unpin,
70    E: std::error::Error,
71{
72}
73
74/// Trait allowing transports to be split into three of different message types.
75pub trait Split3<I1, O1, I2, O2, I3, O3, E>:
76    dnet_base::Transport<Message<I1, I2, I3, (), ()>, Message<O1, O2, O3, (), ()>, E> + Sized + Unpin
77where
78    E: std::error::Error,
79{
80    /// Split into three transports.
81    fn split_into_3(
82        self,
83    ) -> (
84        Part<
85            I1,
86            O1,
87            fn(O1) -> Message<O1, O2, O3, (), ()>,
88            fn(Message<I1, I2, I3, (), ()>) -> Result<I1, Message<I1, I2, I3, (), ()>>,
89            Self,
90            E,
91            I1,
92            I2,
93            I3,
94            (),
95            (),
96        >,
97        Part<
98            I2,
99            O2,
100            fn(O2) -> Message<O1, O2, O3, (), ()>,
101            fn(Message<I1, I2, I3, (), ()>) -> Result<I2, Message<I1, I2, I3, (), ()>>,
102            Self,
103            E,
104            I1,
105            I2,
106            I3,
107            (),
108            (),
109        >,
110        Part<
111            I3,
112            O3,
113            fn(O3) -> Message<O1, O2, O3, (), ()>,
114            fn(Message<I1, I2, I3, (), ()>) -> Result<I3, Message<I1, I2, I3, (), ()>>,
115            Self,
116            E,
117            I1,
118            I2,
119            I3,
120            (),
121            (),
122        >,
123    ) {
124        let state = State::new(3, self);
125        (
126            Part::new(0, &state, Message::Variant1, Message::unwrap1),
127            Part::new(1, &state, Message::Variant2, Message::unwrap2),
128            Part::new(2, &state, Message::Variant3, Message::unwrap3),
129        )
130    }
131}
132
133impl<T, I1, O1, I2, O2, I3, O3, E> Split3<I1, O1, I2, O2, I3, O3, E> for T
134where
135    T: dnet_base::Transport<Message<I1, I2, I3, (), ()>, Message<O1, O2, O3, (), ()>, E> + Unpin,
136    E: std::error::Error,
137{
138}
139
140/// Trait allowing transports to be split into four of different message types.
141pub trait Split4<I1, O1, I2, O2, I3, O3, I4, O4, E>:
142    dnet_base::Transport<Message<I1, I2, I3, I4, ()>, Message<O1, O2, O3, O4, ()>, E> + Sized + Unpin
143where
144    E: std::error::Error,
145{
146    /// Split into four transports.
147    fn split_into_4(
148        self,
149    ) -> (
150        Part<
151            I1,
152            O1,
153            fn(O1) -> Message<O1, O2, O3, O4, ()>,
154            fn(Message<I1, I2, I3, I4, ()>) -> Result<I1, Message<I1, I2, I3, I4, ()>>,
155            Self,
156            E,
157            I1,
158            I2,
159            I3,
160            I4,
161            (),
162        >,
163        Part<
164            I2,
165            O2,
166            fn(O2) -> Message<O1, O2, O3, O4, ()>,
167            fn(Message<I1, I2, I3, I4, ()>) -> Result<I2, Message<I1, I2, I3, I4, ()>>,
168            Self,
169            E,
170            I1,
171            I2,
172            I3,
173            I4,
174            (),
175        >,
176        Part<
177            I3,
178            O3,
179            fn(O3) -> Message<O1, O2, O3, O4, ()>,
180            fn(Message<I1, I2, I3, I4, ()>) -> Result<I3, Message<I1, I2, I3, I4, ()>>,
181            Self,
182            E,
183            I1,
184            I2,
185            I3,
186            I4,
187            (),
188        >,
189        Part<
190            I4,
191            O4,
192            fn(O4) -> Message<O1, O2, O3, O4, ()>,
193            fn(Message<I1, I2, I3, I4, ()>) -> Result<I4, Message<I1, I2, I3, I4, ()>>,
194            Self,
195            E,
196            I1,
197            I2,
198            I3,
199            I4,
200            (),
201        >,
202    ) {
203        let state = State::new(4, self);
204        (
205            Part::new(0, &state, Message::Variant1, Message::unwrap1),
206            Part::new(1, &state, Message::Variant2, Message::unwrap2),
207            Part::new(2, &state, Message::Variant3, Message::unwrap3),
208            Part::new(3, &state, Message::Variant4, Message::unwrap4),
209        )
210    }
211}
212
213impl<T, I1, O1, I2, O2, I3, O3, I4, O4, E> Split4<I1, O1, I2, O2, I3, O3, I4, O4, E> for T
214where
215    T: dnet_base::Transport<Message<I1, I2, I3, I4, ()>, Message<O1, O2, O3, O4, ()>, E> + Unpin,
216    E: std::error::Error,
217{
218}
219
220/// Trait allowing transports to be split into five of different message types.
221pub trait Split5<I1, O1, I2, O2, I3, O3, I4, O4, I5, O5, E>:
222    dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Message<O1, O2, O3, O4, O5>, E> + Sized + Unpin
223where
224    E: std::error::Error,
225{
226    /// Split into five transports.
227    fn split_into_5(
228        self,
229    ) -> (
230        Part<
231            I1,
232            O1,
233            fn(O1) -> Message<O1, O2, O3, O4, O5>,
234            fn(Message<I1, I2, I3, I4, I5>) -> Result<I1, Message<I1, I2, I3, I4, I5>>,
235            Self,
236            E,
237            I1,
238            I2,
239            I3,
240            I4,
241            I5,
242        >,
243        Part<
244            I2,
245            O2,
246            fn(O2) -> Message<O1, O2, O3, O4, O5>,
247            fn(Message<I1, I2, I3, I4, I5>) -> Result<I2, Message<I1, I2, I3, I4, I5>>,
248            Self,
249            E,
250            I1,
251            I2,
252            I3,
253            I4,
254            I5,
255        >,
256        Part<
257            I3,
258            O3,
259            fn(O3) -> Message<O1, O2, O3, O4, O5>,
260            fn(Message<I1, I2, I3, I4, I5>) -> Result<I3, Message<I1, I2, I3, I4, I5>>,
261            Self,
262            E,
263            I1,
264            I2,
265            I3,
266            I4,
267            I5,
268        >,
269        Part<
270            I4,
271            O4,
272            fn(O4) -> Message<O1, O2, O3, O4, O5>,
273            fn(Message<I1, I2, I3, I4, I5>) -> Result<I4, Message<I1, I2, I3, I4, I5>>,
274            Self,
275            E,
276            I1,
277            I2,
278            I3,
279            I4,
280            I5,
281        >,
282        Part<
283            I5,
284            O5,
285            fn(O5) -> Message<O1, O2, O3, O4, O5>,
286            fn(Message<I1, I2, I3, I4, I5>) -> Result<I5, Message<I1, I2, I3, I4, I5>>,
287            Self,
288            E,
289            I1,
290            I2,
291            I3,
292            I4,
293            I5,
294        >,
295    ) {
296        let state = State::new(5, self);
297        (
298            Part::new(0, &state, Message::Variant1, Message::unwrap1),
299            Part::new(1, &state, Message::Variant2, Message::unwrap2),
300            Part::new(2, &state, Message::Variant3, Message::unwrap3),
301            Part::new(3, &state, Message::Variant4, Message::unwrap4),
302            Part::new(4, &state, Message::Variant5, Message::unwrap5),
303        )
304    }
305}
306
307impl<T, I1, O1, I2, O2, I3, O3, I4, O4, I5, O5, E> Split5<I1, O1, I2, O2, I3, O3, I4, O4, I5, O5, E>
308    for T
309where
310    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Message<O1, O2, O3, O4, O5>, E> + Unpin,
311    E: std::error::Error,
312{
313}
314
315/// [Part] transport error.
316#[derive(Debug)]
317pub enum Error<T> {
318    /// Unexpected variant received.
319    UnexpectedVariantReceived(usize),
320
321    /// Wrapped transport error.
322    Transport(T),
323}
324
325impl<T> Display for Error<T>
326where
327    T: Display,
328{
329    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
330        match self {
331            Self::UnexpectedVariantReceived(variant) => {
332                write!(f, "unexpected message variant received: {variant}")
333            }
334            Self::Transport(error) => write!(f, "transport error: {}", error),
335        }
336    }
337}
338
339impl<T> std::error::Error for Error<T> where T: Debug + Display {}
340
341/// Wrapping message of split transport.
342#[derive(Debug, Serialize, Deserialize)]
343pub enum Message<T1, T2, T3, T4, T5> {
344    /// Message of the first transport.
345    Variant1(T1),
346
347    /// Message of the second transport.
348    Variant2(T2),
349
350    /// Message of the third transport.
351    Variant3(T3),
352
353    /// Message of the fourth transport.
354    Variant4(T4),
355
356    /// Message of the fifth transport.
357    Variant5(T5),
358}
359
360impl<T1, T2, T3, T4, T5> Message<T1, T2, T3, T4, T5> {
361    fn unwrap1(self) -> Result<T1, Self> {
362        if let Message::Variant1(message) = self {
363            Ok(message)
364        } else {
365            Err(self)
366        }
367    }
368
369    fn unwrap2(self) -> Result<T2, Self> {
370        if let Message::Variant2(message) = self {
371            Ok(message)
372        } else {
373            Err(self)
374        }
375    }
376
377    fn unwrap3(self) -> Result<T3, Self> {
378        if let Message::Variant3(message) = self {
379            Ok(message)
380        } else {
381            Err(self)
382        }
383    }
384
385    fn unwrap4(self) -> Result<T4, Self> {
386        if let Message::Variant4(message) = self {
387            Ok(message)
388        } else {
389            Err(self)
390        }
391    }
392
393    fn unwrap5(self) -> Result<T5, Self> {
394        if let Message::Variant5(message) = self {
395            Ok(message)
396        } else {
397            Err(self)
398        }
399    }
400}
401
402/// One of the transports after split.
403#[pin_project]
404pub struct Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5> {
405    variant: usize,
406    state: Arc<Mutex<State<T, I1, I2, I3, I4, I5>>>,
407    wrapper: Wrapper,
408    unwrapper: Unwrapper,
409
410    #[cfg(feature = "logging")]
411    logger: dnet_base::Logger,
412
413    _incoming: PhantomData<Incoming>,
414    _outgoing: PhantomData<Outgoing>,
415    _error: PhantomData<E>,
416}
417
418impl<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
419    Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
420where
421    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Wrapper::Output, E> + Unpin,
422    Wrapper: Mapper<Outgoing>,
423    Unwrapper:
424        Mapper<Message<I1, I2, I3, I4, I5>, Output = Result<Incoming, Message<I1, I2, I3, I4, I5>>>,
425    E: std::error::Error,
426{
427    fn new(
428        variant: usize,
429        state: &Arc<Mutex<State<T, I1, I2, I3, I4, I5>>>,
430        wrapper: Wrapper,
431        unwrapper: Unwrapper,
432    ) -> Self {
433        Part {
434            variant,
435            state: state.clone(),
436            wrapper,
437            unwrapper,
438
439            #[cfg(feature = "logging")]
440            logger: {
441                let mut logger = dnet_base::Logger::new::<Self>();
442                logger.override_kind_part::<Self>(variant);
443                logger
444            },
445
446            _incoming: PhantomData,
447            _outgoing: PhantomData,
448            _error: PhantomData,
449        }
450    }
451}
452
453impl<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5> Sink<Outgoing>
454    for Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
455where
456    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Wrapper::Output, E> + Unpin,
457    Wrapper: Mapper<Outgoing>,
458    Unwrapper:
459        Mapper<Message<I1, I2, I3, I4, I5>, Output = Result<Incoming, Message<I1, I2, I3, I4, I5>>>,
460    E: std::error::Error,
461{
462    type Error = dnet_base::Error<Error<E>>;
463
464    fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
465        let result = self
466            .state
467            .lock()
468            .unwrap()
469            .inner
470            .poll_ready_unpin(cx)
471            .map_err(map_error);
472
473        #[cfg(feature = "logging")]
474        self.logger.log_ready(&result);
475
476        result
477    }
478
479    fn start_send(self: Pin<&mut Self>, item: Outgoing) -> Result<(), Self::Error> {
480        let me = self.project();
481        let item = me.wrapper.map(item);
482        let result = me
483            .state
484            .lock()
485            .unwrap()
486            .inner
487            .start_send_unpin(item)
488            .map_err(map_error);
489
490        #[cfg(feature = "logging")]
491        match &result {
492            Ok(_) => me.logger.log_message_preparation_success::<Outgoing>(None),
493            Err(error) => me.logger.log_sending_failure(error),
494        }
495
496        result
497    }
498
499    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
500        let result = self
501            .state
502            .lock()
503            .unwrap()
504            .inner
505            .poll_flush_unpin(cx)
506            .map_err(map_error);
507
508        #[cfg(feature = "logging")]
509        self.logger.log_flush(&result);
510
511        result
512    }
513
514    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
515        let result = self
516            .state
517            .lock()
518            .unwrap()
519            .inner
520            .poll_close_unpin(cx)
521            .map_err(map_error);
522
523        #[cfg(feature = "logging")]
524        self.logger.log_close(&result);
525
526        result
527    }
528}
529
530impl<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5> Stream
531    for Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
532where
533    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Wrapper::Output, E> + Unpin,
534    Wrapper: Mapper<Outgoing>,
535    Unwrapper:
536        Mapper<Message<I1, I2, I3, I4, I5>, Output = Result<Incoming, Message<I1, I2, I3, I4, I5>>>,
537    E: std::error::Error,
538{
539    type Item = Result<Incoming, Error<E>>;
540
541    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
542        let me = self.project();
543        let mut lock = me.state.lock().unwrap();
544        let result = if let Some(message) = lock.buffers[*me.variant].pop() {
545            Poll::Ready(Some(Ok(me.unwrapper.map(message).ok().unwrap())))
546        } else {
547            loop {
548                let Poll::Ready(item) = lock.inner.poll_next_unpin(cx) else {
549                    lock.buffers[*me.variant].update_waker_with(cx.waker());
550                    break Poll::Pending;
551                };
552                match item {
553                    Some(Ok(item)) => {
554                        let variant = match item {
555                            Message::Variant1(_) => 0,
556                            Message::Variant2(_) => 1,
557                            Message::Variant3(_) => 2,
558                            Message::Variant4(_) => 3,
559                            Message::Variant5(_) => 4,
560                        };
561                        match me.unwrapper.map(item) {
562                            Ok(item) => break Poll::Ready(Some(Ok(item))),
563                            Err(message) => {
564                                if let Some(buffer) = lock.buffers.get_mut(variant) {
565                                    buffer.push(message);
566                                } else {
567                                    break Poll::Ready(Some(Err(
568                                        Error::UnexpectedVariantReceived(variant),
569                                    )));
570                                }
571                            }
572                        }
573                    }
574                    Some(Err(error)) => break Poll::Ready(Some(Err(Error::Transport(error)))),
575                    None => break Poll::Ready(None),
576                }
577            }
578        };
579
580        #[cfg(feature = "logging")]
581        me.logger.log_receiving(&result, None);
582
583        result
584    }
585}
586
587impl<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5> FusedStream
588    for Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
589where
590    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Wrapper::Output, E> + FusedStream + Unpin,
591    Wrapper: Mapper<Outgoing>,
592    Unwrapper:
593        Mapper<Message<I1, I2, I3, I4, I5>, Output = Result<Incoming, Message<I1, I2, I3, I4, I5>>>,
594    E: std::error::Error,
595{
596    fn is_terminated(&self) -> bool {
597        self.state.lock().unwrap().inner.is_terminated()
598    }
599}
600
601#[cfg(feature = "logging")]
602impl<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5> dnet_base::Logging
603    for Part<Incoming, Outgoing, Wrapper, Unwrapper, T, E, I1, I2, I3, I4, I5>
604where
605    T: dnet_base::Transport<Message<I1, I2, I3, I4, I5>, Wrapper::Output, E> + Unpin,
606    Wrapper: Mapper<Outgoing>,
607    Unwrapper:
608        Mapper<Message<I1, I2, I3, I4, I5>, Output = Result<Incoming, Message<I1, I2, I3, I4, I5>>>,
609    E: std::error::Error,
610{
611    const KIND: &'static str = "Part";
612
613    fn with_logger<F, R>(&self, f: F) -> R
614    where
615        F: FnOnce(&dnet_base::Logger) -> R,
616    {
617        f(&self.logger)
618    }
619
620    fn with_logger_mut<F, R>(&mut self, f: F) -> R
621    where
622        F: FnOnce(&mut dnet_base::Logger) -> R,
623    {
624        f(&mut self.logger)
625    }
626}
627
628struct State<T, I1, I2, I3, I4, I5> {
629    inner: T,
630    buffers: Vec<Buffer<I1, I2, I3, I4, I5>>,
631}
632
633impl<T, I1, I2, I3, I4, I5> State<T, I1, I2, I3, I4, I5> {
634    fn new(size: usize, inner: T) -> Arc<Mutex<Self>> {
635        let buffers = (0..size).map(|_| Buffer::new()).collect();
636        let state = State { inner, buffers };
637        Arc::new(Mutex::new(state))
638    }
639}
640
641struct Buffer<I1, I2, I3, I4, I5> {
642    inner: VecDeque<Message<I1, I2, I3, I4, I5>>,
643    waker: Option<Waker>,
644}
645
646impl<I1, I2, I3, I4, I5> Buffer<I1, I2, I3, I4, I5> {
647    fn new() -> Self {
648        Buffer {
649            inner: VecDeque::new(),
650            waker: None,
651        }
652    }
653
654    fn pop(&mut self) -> Option<Message<I1, I2, I3, I4, I5>> {
655        self.inner.pop_front()
656    }
657
658    fn push(&mut self, message: Message<I1, I2, I3, I4, I5>) {
659        self.inner.push_back(message);
660        self.wake();
661    }
662
663    fn update_waker_with(&mut self, other: &Waker) {
664        if let Some(waker) = &self.waker {
665            if !waker.will_wake(other) {
666                self.waker = Some(other.clone());
667            }
668        } else {
669            self.waker = Some(other.clone());
670        }
671    }
672
673    fn wake(&mut self) {
674        if let Some(waker) = self.waker.take() {
675            waker.wake();
676        }
677    }
678}
679
680fn map_error<T>(error: dnet_base::Error<T>) -> dnet_base::Error<Error<T>> {
681    match error {
682        dnet_base::Error::Closed => dnet_base::Error::Closed,
683        dnet_base::Error::Other(error) => dnet_base::Error::Other(Error::Transport(error)),
684    }
685}
686
687#[cfg(test)]
688mod tests {
689    use dnet_base::Receive;
690    use dnet_tests::{dtest, dtest_configure};
691    use futures::SinkExt;
692
693    use crate::{channel::transports, split::Split2};
694
695    dtest_configure!();
696
697    #[dtest]
698    async fn test_split() {
699        let (left, right) = transports();
700
701        let (mut left_string_i32, mut left_u32_f64) = left.split_into_2();
702        let (mut right_i32_string, mut right_f64_u32) = right.split_into_2();
703
704        dnet_tests::init_logging(&mut left_string_i32, &mut right_i32_string);
705        dnet_tests::init_logging(&mut left_u32_f64, &mut right_f64_u32);
706
707        left_string_i32.send(-50).await.unwrap();
708        left_u32_f64.send(770.0).await.unwrap();
709
710        right_i32_string.send("Hello".to_string()).await.unwrap();
711        right_f64_u32.send(66).await.unwrap();
712
713        assert_eq!(left_u32_f64.receive().await.unwrap(), 66);
714        assert_eq!(left_string_i32.receive().await.unwrap(), "Hello");
715        assert_eq!(right_f64_u32.receive().await.unwrap(), 770.0);
716        assert_eq!(right_i32_string.receive().await.unwrap(), -50);
717    }
718}