1use 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 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 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 pub trait Error: Send + 'static {}
29 impl<T> Error for T where T: Send + 'static {}
30}
31
32#[derive(Debug, PartialEq, Eq, Clone, Copy, Default)]
34pub enum ErrorHandlingStrategy {
35 #[default]
37 Ignore,
38
39 Retry,
41
42 Close,
44}
45
46create_non_sync_send_variant_for_wasm! {
47 pub trait SendErrorCallback<Message, Error>: Send + 'static {
54 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 pub trait ReceiveErrorCallback<Error>: Send + 'static {
84 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#[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#[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
130pub struct ErrorHandler<Message, Error> {
132 pub send_error_callback: Box<dyn SendErrorCallback<Message, Error>>,
135
136 pub receive_error_callback: Box<dyn ReceiveErrorCallback<Error>>,
139}
140
141impl<Message, Error> ErrorHandler<Message, Error> {
142 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#[derive(Debug)]
163pub struct Pipe {
164 stop_sender: Option<oneshot::Sender<()>>,
165 keep_open: bool,
166 closed: Notifier,
167}
168
169impl Pipe {
170 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 pub fn open(&self) -> bool {
236 !self.closed.already_notified()
237 }
238
239 pub fn break_pipe(mut self) {
241 self.keep_open = false;
242 }
244
245 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
275pub 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}