Skip to main content

dnet_utils/
pipe.rs

1//! Pipe messages from one transport into another and vice versa.
2
3use std::{
4    future::Future,
5    pin::Pin,
6    task::{Context, Poll},
7};
8
9use dportable::{create_non_sync_send_variant_for_wasm, spawn, value::Notifier};
10use futures::{channel::oneshot, future::FusedFuture, select, FutureExt, Sink, SinkExt, StreamExt};
11
12create_non_sync_send_variant_for_wasm! {
13    /// Helper trait for transports that can be piped.
14    pub trait Transport<Incoming, Outgoing, Error>:
15        dnet_base::Transport<Incoming, Outgoing, Error> + Send + Unpin + 'static {}
16    impl<T, Incoming, Outgoing, Error>  Transport<Incoming, Outgoing, Error> for T
17        where T: dnet_base::Transport<Incoming, Outgoing, Error> + Send + Unpin + 'static {}
18}
19
20create_non_sync_send_variant_for_wasm! {
21    /// Helper trait for pipe-able transport message.
22    pub trait Message: Clone + Send + 'static {}
23    impl<T> Message for T where T: Clone + Send + 'static {}
24}
25
26create_non_sync_send_variant_for_wasm! {
27    /// Helper trait for pipe-able transport error.
28    pub trait Error: Send + 'static {}
29    impl<T> Error for T where T: Send + 'static {}
30}
31
32/// Strategy of handling message passing errors.
33#[derive(Debug, PartialEq, Eq, Clone, Copy, Default)]
34pub enum ErrorHandlingStrategy {
35    /// Ignore the error.
36    #[default]
37    Ignore,
38
39    /// Try again to send/receive message.
40    Retry,
41
42    /// Close pipe.
43    Close,
44}
45
46create_non_sync_send_variant_for_wasm! {
47    /// Callback called when an error is encountered while trying to send a message
48    /// through the transport.
49    ///
50    /// **NOTE**: it also receives [dnet_base::Error::Closed] errors and
51    /// they need to be handled properly (most likely by returning
52    /// [ErrorHandlingStrategy::Close]).
53    pub trait SendErrorCallback<Message, Error>: Send + 'static {
54        /// Handle sending error.
55        ///
56        /// **NOTE**: it also receives [dnet_base::Error::Closed] errors and
57        /// they need to be handled properly (most likely by returning
58        /// [ErrorHandlingStrategy::Close]).
59        fn on_send_error(
60            &mut self,
61            message: &Message,
62            error: dnet_base::Error<Error>,
63        ) -> ErrorHandlingStrategy;
64    }
65
66    impl<T, Message, Error> SendErrorCallback<Message, Error> for T
67    where
68        T: FnMut(&Message, dnet_base::Error<Error>) -> ErrorHandlingStrategy + Send + 'static,
69    {
70        fn on_send_error(
71            &mut self,
72            message: &Message,
73            error: dnet_base::Error<Error>,
74        ) -> ErrorHandlingStrategy {
75            (self)(message, error)
76        }
77    }
78}
79
80create_non_sync_send_variant_for_wasm! {
81    /// Callback called when an error is encountered while
82    /// trying to receive a message from the transport.
83    pub trait ReceiveErrorCallback<Error>: Send + 'static {
84        /// Handle receiving error.
85        fn on_receive_error(&mut self, error: Error) -> ErrorHandlingStrategy;
86    }
87
88    impl<T, Error> ReceiveErrorCallback<Error> for T
89    where
90        T: FnMut(Error) -> ErrorHandlingStrategy + Send + 'static,
91    {
92        fn on_receive_error(&mut self, error: Error) -> ErrorHandlingStrategy {
93            (self)(error)
94        }
95    }
96}
97
98/// Default [SendErrorCallback].
99///
100/// It ignores errors (except [dnet_base::Error::Closed] error - which results in
101/// closing pipe).
102#[derive(Debug)]
103pub struct DefaultSendErrorCallback;
104
105impl<Message, Error> SendErrorCallback<Message, Error> for DefaultSendErrorCallback {
106    fn on_send_error(
107        &mut self,
108        _message: &Message,
109        error: dnet_base::Error<Error>,
110    ) -> ErrorHandlingStrategy {
111        match error {
112            dnet_base::Error::Closed => ErrorHandlingStrategy::Close,
113            dnet_base::Error::Other(_) => ErrorHandlingStrategy::Ignore,
114        }
115    }
116}
117
118/// Default [ReceiveErrorCallback].
119///
120/// It ignores errors.
121#[derive(Debug)]
122pub struct DefaultReceiveErrorCallback;
123
124impl<Error> ReceiveErrorCallback<Error> for DefaultReceiveErrorCallback {
125    fn on_receive_error(&mut self, _error: Error) -> ErrorHandlingStrategy {
126        ErrorHandlingStrategy::Ignore
127    }
128}
129
130/// Pipe message passing error handler.
131pub struct ErrorHandler<Message, Error> {
132    /// Callback called when an error is encountered while trying to send a message
133    /// through the transport.
134    pub send_error_callback: Box<dyn SendErrorCallback<Message, Error>>,
135
136    /// Callback called when an error is encountered while trying to receive a message
137    /// from the transport.
138    pub receive_error_callback: Box<dyn ReceiveErrorCallback<Error>>,
139}
140
141impl<Message, Error> ErrorHandler<Message, Error> {
142    /// Create new error handler.
143    pub fn new<S, R>(send_error_callback: S, receive_error_callback: R) -> Self
144    where
145        S: SendErrorCallback<Message, Error>,
146        R: ReceiveErrorCallback<Error>,
147    {
148        ErrorHandler {
149            send_error_callback: Box::new(send_error_callback),
150            receive_error_callback: Box::new(receive_error_callback),
151        }
152    }
153}
154
155impl<Message, Error> Default for ErrorHandler<Message, Error> {
156    fn default() -> Self {
157        ErrorHandler::new(DefaultSendErrorCallback, DefaultReceiveErrorCallback)
158    }
159}
160
161/// Pipe sending messages from one transport into another and vice versa.
162#[derive(Debug)]
163pub struct Pipe {
164    stop_sender: Option<oneshot::Sender<()>>,
165    keep_open: bool,
166    closed: Notifier,
167}
168
169impl Pipe {
170    /// Create new pipe sending messages from one transport into another and vice versa.
171    pub fn new<A, B, M1, M2, E1, E2>(
172        a: A,
173        b: B,
174        mut a_error_handler: ErrorHandler<M2, E1>,
175        mut b_error_handler: ErrorHandler<M1, E2>,
176    ) -> Self
177    where
178        A: Transport<M1, M2, E1> + Unpin,
179        B: Transport<M2, M1, E2> + Unpin,
180        M1: Message,
181        M2: Message,
182        E1: Error,
183        E2: Error,
184    {
185        let (stop_sender, mut stop_receiver) = oneshot::channel();
186        let stop_sender = Some(stop_sender);
187        let closed = Notifier::new();
188        let closed_clone = closed.clone();
189        spawn(async move {
190            let (mut sender_a, receiver_a) = a.split();
191            let mut receiver_a = receiver_a.fuse();
192            let (mut sender_b, receiver_b) = b.split();
193            let mut receiver_b = receiver_b.fuse();
194            let mut should_close = false;
195            loop {
196                select! {
197                    a = receiver_a.next() => {
198                        handle_receive_result(
199                            &mut sender_b,
200                            a,
201                            &mut a_error_handler.receive_error_callback,
202                            &mut b_error_handler.send_error_callback,
203                            &mut should_close
204                        ).await;
205                    }
206                    b = receiver_b.next() => {
207                        handle_receive_result(
208                            &mut sender_a,
209                            b,
210                            &mut b_error_handler.receive_error_callback,
211                            &mut a_error_handler.send_error_callback,
212                            &mut should_close
213                        ).await;
214                    }
215                    result = stop_receiver => {
216                        if result.is_ok() {
217                            should_close = true
218                        }
219                    }
220                }
221                if should_close {
222                    break;
223                }
224            }
225            closed_clone.notify();
226        });
227        Pipe {
228            stop_sender,
229            keep_open: false,
230            closed,
231        }
232    }
233
234    /// Is pipe still open.
235    pub fn open(&self) -> bool {
236        !self.closed.already_notified()
237    }
238
239    /// Stop message interchange.
240    pub fn break_pipe(mut self) {
241        self.keep_open = false;
242        // drop(self)
243    }
244
245    /// Keep pipe open after drop.
246    pub fn keep_open(&mut self) {
247        self.keep_open = true;
248    }
249}
250
251impl Drop for Pipe {
252    fn drop(&mut self) {
253        if !self.keep_open {
254            if let Some(sender) = self.stop_sender.take() {
255                let _ = sender.send(());
256            }
257        }
258    }
259}
260
261impl Future for Pipe {
262    type Output = ();
263
264    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
265        self.closed.poll_unpin(cx)
266    }
267}
268
269impl FusedFuture for Pipe {
270    fn is_terminated(&self) -> bool {
271        self.closed.is_terminated()
272    }
273}
274
275/// Pipe messages from one transport into another and vice versa.
276///
277/// Pipe created with this function ignores errors.
278pub fn pipe<A, B, M1, M2, E1, E2>(a: A, b: B) -> Pipe
279where
280    A: Transport<M1, M2, E1> + Unpin,
281    B: Transport<M2, M1, E2> + Unpin,
282    M1: Message,
283    M2: Message,
284    E1: Error,
285    E2: Error,
286{
287    Pipe::new(a, b, Default::default(), Default::default())
288}
289
290async fn handle_receive_result<S, M, ER, ES>(
291    sender: &mut S,
292    result: Option<Result<M, ER>>,
293    receive_error_callback: &mut Box<dyn ReceiveErrorCallback<ER>>,
294    send_error_callback: &mut Box<dyn SendErrorCallback<M, ES>>,
295    should_close: &mut bool,
296) where
297    S: Sink<M, Error = dnet_base::Error<ES>> + Unpin,
298    M: Message,
299    ES: Error,
300    ER: Error,
301{
302    if let Some(result) = result {
303        match result {
304            Ok(message) => {
305                send(sender, message, send_error_callback, should_close).await;
306            }
307            Err(error) => {
308                let strategy = receive_error_callback.on_receive_error(error);
309                if matches!(strategy, ErrorHandlingStrategy::Close) {
310                    *should_close = true;
311                }
312            }
313        }
314    } else {
315        *should_close = true;
316    }
317}
318
319async fn send<S, M, E>(
320    sender: &mut S,
321    message: M,
322    send_error_callback: &mut Box<dyn SendErrorCallback<M, E>>,
323    should_close: &mut bool,
324) where
325    S: Sink<M, Error = dnet_base::Error<E>> + Unpin,
326    M: Message,
327    E: Error,
328{
329    while let Err(error) = sender.send(message.clone()).await {
330        match send_error_callback.on_send_error(&message, error) {
331            ErrorHandlingStrategy::Ignore => {
332                break;
333            }
334            ErrorHandlingStrategy::Retry => {
335                continue;
336            }
337            ErrorHandlingStrategy::Close => {
338                *should_close = true;
339                return;
340            }
341        }
342    }
343    *should_close = false;
344}
345
346#[cfg(test)]
347mod tests {
348    use dnet_base::Receive;
349    use dnet_tests::{dtest, dtest_configure};
350    use futures::SinkExt;
351
352    use crate::channel::{transports, ChannelTransport};
353
354    use super::{pipe, Message, Pipe};
355
356    dtest_configure!();
357
358    fn create_transports<A, B>() -> (ChannelTransport<A, B>, ChannelTransport<B, A>, Pipe)
359    where
360        A: Message,
361        B: Message,
362    {
363        let (out_a, to_pipe_a) = transports();
364        let (out_b, to_pipe_b) = transports();
365        let pipe = pipe(to_pipe_a, to_pipe_b);
366        (out_a, out_b, pipe)
367    }
368
369    #[dtest]
370    async fn test_transport() {
371        let (left, right, _pipe) = create_transports();
372        dnet_tests::test_transport(left, right).await;
373    }
374
375    #[dtest]
376    async fn test_unit_message() {
377        let (left, right, _pipe) = create_transports();
378        dnet_tests::test_unit_message(left, right).await;
379    }
380
381    #[dtest]
382    async fn test_stream() {
383        let (left, right, _pipe) = create_transports();
384        dnet_tests::test_stream(left, right).await;
385    }
386
387    #[dtest]
388    async fn test_pipe_drop() {
389        let (mut left, mut right, pipe) = create_transports();
390
391        dnet_tests::init_logging(&mut left, &mut right);
392
393        left.send(1).await.unwrap();
394        right.send(1).await.unwrap();
395        assert_eq!(left.receive().await.unwrap(), 1);
396        assert_eq!(right.receive().await.unwrap(), 1);
397        drop(pipe);
398        assert!(matches!(
399            left.receive().await,
400            Err(dnet_base::Error::Closed)
401        ));
402        assert!(matches!(
403            right.receive().await,
404            Err(dnet_base::Error::Closed)
405        ));
406    }
407}