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