1use 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
14pub 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
24pub trait Message: Clone + ConditionalSend + 'static {}
26impl<T> Message for T where T: Clone + ConditionalSend + 'static {}
27
28pub trait Error: ConditionalSend + 'static {}
30impl<T> Error for T where T: ConditionalSend + 'static {}
31
32#[derive(Debug, PartialEq, Eq, Clone, Copy, Default)]
34pub enum ErrorHandlingStrategy {
35 #[default]
37 Ignore,
38
39 Retry,
41
42 Close,
44}
45
46pub trait SendErrorCallback<Message, Error>: ConditionalSend + 'static {
53 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
80pub trait ReceiveErrorCallback<Error>: ConditionalSend + 'static {
83 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#[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#[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
128pub struct ErrorHandler<Message, Error> {
130 pub send_error_callback: Box<dyn SendErrorCallback<Message, Error>>,
133
134 pub receive_error_callback: Box<dyn ReceiveErrorCallback<Error>>,
137}
138
139impl<Message, Error> ErrorHandler<Message, Error> {
140 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#[derive(Debug)]
161pub struct Pipe {
162 stop_sender: Option<oneshot::Sender<()>>,
163 keep_open: bool,
164 closed: Notifier,
165}
166
167impl Pipe {
168 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 pub fn open(&self) -> bool {
234 !self.closed.already_notified()
235 }
236
237 pub fn break_pipe(mut self) {
239 self.keep_open = false;
240 }
242
243 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
273pub 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}