1#![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
20pub 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 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
74pub 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 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
140pub 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 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
220pub 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 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#[derive(Debug)]
317pub enum Error<T> {
318 UnexpectedVariantReceived(usize),
320
321 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#[derive(Debug, Serialize, Deserialize)]
343pub enum Message<T1, T2, T3, T4, T5> {
344 Variant1(T1),
346
347 Variant2(T2),
349
350 Variant3(T3),
352
353 Variant4(T4),
355
356 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#[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}