1use std::{
4 fmt::{Debug, Display},
5 marker::PhantomData,
6 pin::Pin,
7 task::{Context, Poll},
8};
9
10use futures::{stream::FusedStream, Sink, Stream};
11use num::{traits::bounds::UpperBounded, One, Zero};
12use pin_project::pin_project;
13use serde::{Deserialize, Serialize};
14
15use super::unwrap::Unwrap;
16
17pub trait NumberMessages<N, Incoming, Outgoing, Error>:
19 dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, Error> + Sized + Unpin
20where
21 Error: std::error::Error,
22{
23 fn number_messages(self) -> Numbered<N, Self, Error, Incoming, Outgoing>
25 where
26 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
27 {
28 Numbered::new(self)
29 }
30}
31
32impl<T, N, Incoming, Outgoing, Error> NumberMessages<N, Incoming, Outgoing, Error> for T
33where
34 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, Error> + Unpin,
35 Error: std::error::Error,
36{
37}
38
39pub trait NumberMessagesUsize<Incoming, Outgoing, Error>:
43 dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, Error> + Sized + Unpin
44where
45 Error: std::error::Error,
46{
47 fn number_messages_u64(self) -> Numbered<usize, Self, Error, Incoming, Outgoing> {
49 Numbered::new(self)
50 }
51}
52
53impl<T, Incoming, Outgoing, Error> NumberMessagesUsize<Incoming, Outgoing, Error> for T
54where
55 T: dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, Error> + Unpin,
56 Error: std::error::Error,
57{
58}
59
60pub trait NumberMessagesU32<Incoming, Outgoing, Error>:
62 dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, Error> + Sized + Unpin
63where
64 Error: std::error::Error,
65{
66 fn number_messages_u32(self) -> Numbered<u32, Self, Error, Incoming, Outgoing> {
68 Numbered::new(self)
69 }
70}
71
72impl<T, Incoming, Outgoing, Error> NumberMessagesU32<Incoming, Outgoing, Error> for T
73where
74 T: dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, Error> + Unpin,
75 Error: std::error::Error,
76{
77}
78
79pub trait NumberMessagesU64<Incoming, Outgoing, Error>:
81 dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, Error> + Sized + Unpin
82where
83 Error: std::error::Error,
84{
85 fn number_messages_u64(self) -> Numbered<u64, Self, Error, Incoming, Outgoing> {
87 Numbered::new(self)
88 }
89}
90
91impl<T, Incoming, Outgoing, Error> NumberMessagesU64<Incoming, Outgoing, Error> for T
92where
93 T: dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, Error> + Unpin,
94 Error: std::error::Error,
95{
96}
97
98pub trait NumberMessagesU128<Incoming, Outgoing, Error>:
100 dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, Error> + Sized + Unpin
101where
102 Error: std::error::Error,
103{
104 fn number_messages_u128(self) -> Numbered<u128, Self, Error, Incoming, Outgoing> {
106 Numbered::new(self)
107 }
108}
109
110impl<T, Incoming, Outgoing, Error> NumberMessagesU128<Incoming, Outgoing, Error> for T
111where
112 T: dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, Error> + Unpin,
113 Error: std::error::Error,
114{
115}
116
117#[derive(Debug)]
119pub enum Error<T> {
120 MaximumNumberReached,
122
123 Transport(T),
125}
126
127impl<T> Display for Error<T>
128where
129 T: Display,
130{
131 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
132 match self {
133 Error::MaximumNumberReached => write!(f, "maximum number reached"),
134 Error::Transport(error) => write!(f, "transport error: {error}"),
135 }
136 }
137}
138
139impl<T> std::error::Error for Error<T> where T: Debug + Display {}
140
141pub trait Number {
143 type Output;
145
146 fn number(&self) -> Self::Output;
148}
149
150#[derive(Debug, PartialEq, Eq, Serialize, Deserialize)]
152pub struct Wrapper<N, T> {
153 pub number: N,
155
156 pub wrapped: T,
158}
159
160impl<N, T> Unwrap for Wrapper<N, T> {
161 type Output = T;
162
163 fn unwrap(self) -> Self::Output {
164 self.wrapped
165 }
166}
167
168impl<N, T> Number for Wrapper<N, T>
169where
170 N: Clone,
171{
172 type Output = N;
173
174 fn number(&self) -> Self::Output {
175 self.number.clone()
176 }
177}
178
179#[pin_project]
191pub struct Numbered<N, T, E, Incoming, Outgoing>
192where
193 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
194 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
195{
196 #[pin]
197 inner: T,
198 current_number: N,
199
200 #[cfg(feature = "logging")]
201 logger: dnet_base::Logger,
202
203 _error: PhantomData<E>,
204 _number: PhantomData<N>,
205 _incoming: PhantomData<Incoming>,
206 _outgoing: PhantomData<Outgoing>,
207}
208
209impl<N, T, E, Incoming, Outgoing> Numbered<N, T, E, Incoming, Outgoing>
210where
211 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
212 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
213 E: std::error::Error,
214{
215 pub fn new(transport: T) -> Self {
219 Numbered {
220 inner: transport,
221 current_number: N::zero(),
222
223 #[cfg(feature = "logging")]
224 logger: dnet_base::Logger::new::<Self>(),
225
226 _error: PhantomData,
227 _number: PhantomData,
228 _incoming: PhantomData,
229 _outgoing: PhantomData,
230 }
231 }
232
233 pub fn current_number(&self) -> N {
235 self.current_number.clone()
236 }
237}
238
239impl<T, E, Incoming, Outgoing> Numbered<usize, T, E, Incoming, Outgoing>
240where
241 T: dnet_base::Transport<Wrapper<usize, Incoming>, Wrapper<usize, Outgoing>, E>,
242 E: std::error::Error,
243{
244 pub fn new_usize(transport: T) -> Self {
246 Numbered::new(transport)
247 }
248}
249
250impl<T, E, Incoming, Outgoing> Numbered<u32, T, E, Incoming, Outgoing>
251where
252 T: dnet_base::Transport<Wrapper<u32, Incoming>, Wrapper<u32, Outgoing>, E>,
253 E: std::error::Error,
254{
255 pub fn new_u32(transport: T) -> Self {
257 Numbered::new(transport)
258 }
259}
260
261impl<T, E, Incoming, Outgoing> Numbered<u64, T, E, Incoming, Outgoing>
262where
263 T: dnet_base::Transport<Wrapper<u64, Incoming>, Wrapper<u64, Outgoing>, E>,
264 E: std::error::Error,
265{
266 pub fn new_u64(transport: T) -> Self {
268 Numbered::new(transport)
269 }
270}
271
272impl<T, E, Incoming, Outgoing> Numbered<u128, T, E, Incoming, Outgoing>
273where
274 T: dnet_base::Transport<Wrapper<u128, Incoming>, Wrapper<u128, Outgoing>, E>,
275 E: std::error::Error,
276{
277 pub fn new_u128(transport: T) -> Self {
279 Numbered::new(transport)
280 }
281}
282
283impl<N, T, E, Incoming, Outgoing> Sink<Outgoing> for Numbered<N, T, E, Incoming, Outgoing>
284where
285 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
286 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
287 E: std::error::Error,
288{
289 type Error = dnet_base::Error<Error<E>>;
290
291 fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
292 let me = self.project();
293 let result = me.inner.poll_ready(cx).map_err(map_error);
294
295 #[cfg(feature = "logging")]
296 me.logger.log_ready(&result);
297
298 result
299 }
300
301 fn start_send(self: Pin<&mut Self>, item: Outgoing) -> Result<(), Self::Error> {
302 let me = self.project();
303 let result = if *me.current_number == N::max_value() {
304 Err(dnet_base::Error::Other(Error::MaximumNumberReached))
305 } else {
306 let item = Wrapper {
307 number: me.current_number.clone(),
308 wrapped: item,
309 };
310 let result = me.inner.start_send(item);
311 if result.is_ok() {
312 *me.current_number = me.current_number.clone().add(One::one());
313 }
314 result.map_err(map_error)
315 };
316
317 #[cfg(feature = "logging")]
318 match &result {
319 Ok(_) => me.logger.log_message_preparation_success::<Outgoing>(None),
320 Err(error) => me.logger.log_sending_failure(error),
321 }
322
323 result
324 }
325
326 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
327 let me = self.project();
328 let result = me.inner.poll_flush(cx).map_err(map_error);
329
330 #[cfg(feature = "logging")]
331 me.logger.log_flush(&result);
332
333 result
334 }
335
336 fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
337 let me = self.project();
338 let result = me.inner.poll_close(cx).map_err(map_error);
339
340 #[cfg(feature = "logging")]
341 me.logger.log_close(&result);
342
343 result
344 }
345}
346
347impl<N, T, E, Incoming, Outgoing> Stream for Numbered<N, T, E, Incoming, Outgoing>
348where
349 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E>,
350 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
351 E: std::error::Error,
352{
353 type Item = Result<Wrapper<N, Incoming>, Error<E>>;
354
355 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
356 let me = self.project();
357 let result = me.inner.poll_next(cx).map_err(Error::Transport);
358
359 #[cfg(feature = "logging")]
360 me.logger.log_receiving(&result, None);
361
362 result
363 }
364}
365
366impl<N, T, E, Incoming, Outgoing> FusedStream for Numbered<N, T, E, Incoming, Outgoing>
367where
368 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E> + FusedStream,
369 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
370 E: std::error::Error,
371{
372 fn is_terminated(&self) -> bool {
373 self.inner.is_terminated()
374 }
375}
376
377#[cfg(feature = "logging")]
378impl<N, T, E, Incoming, Outgoing> dnet_base::Logging for Numbered<N, T, E, Incoming, Outgoing>
379where
380 T: dnet_base::Transport<Wrapper<N, Incoming>, Wrapper<N, Outgoing>, E> + dnet_base::Logging,
381 N: Clone + Zero + One + UpperBounded + PartialEq + Eq,
382 E: std::error::Error,
383{
384 const KIND: &'static str = "Numbered";
385
386 fn with_logger<F, R>(&self, f: F) -> R
387 where
388 F: FnOnce(&dnet_base::Logger) -> R,
389 {
390 f(&self.logger)
391 }
392
393 fn with_logger_mut<F, R>(&mut self, f: F) -> R
394 where
395 F: FnOnce(&mut dnet_base::Logger) -> R,
396 {
397 f(&mut self.logger)
398 }
399}
400
401fn map_error<T>(error: dnet_base::Error<T>) -> dnet_base::Error<Error<T>> {
402 match error {
403 dnet_base::Error::Closed => dnet_base::Error::Closed,
404 dnet_base::Error::Other(error) => dnet_base::Error::Other(Error::Transport(error)),
405 }
406}
407
408#[cfg(test)]
409mod tests {
410 use dnet_base::Receive;
411 use dnet_tests::{dtest, dtest_configure};
412 use futures::SinkExt;
413
414 use crate::{
415 channel::transports,
416 number::{Numbered, Wrapper},
417 };
418
419 dtest_configure!();
420
421 #[dtest]
422 async fn test_transport() {
423 let (left, right) = transports();
424
425 let mut left = Numbered::new_usize(left);
426 let mut right = Numbered::new_usize(right);
427
428 dnet_tests::init_logging(&mut left, &mut right);
429
430 left.send(1).await.unwrap();
431 left.send(2).await.unwrap();
432 left.send(3).await.unwrap();
433
434 assert_eq!(
435 right.receive().await.unwrap(),
436 Wrapper {
437 number: 0,
438 wrapped: 1,
439 }
440 );
441 assert_eq!(
442 right.receive().await.unwrap(),
443 Wrapper {
444 number: 1,
445 wrapped: 2,
446 }
447 );
448 assert_eq!(
449 right.receive().await.unwrap(),
450 Wrapper {
451 number: 2,
452 wrapped: 3,
453 }
454 );
455
456 right.send(1).await.unwrap();
457 right.send(2).await.unwrap();
458 right.send(3).await.unwrap();
459
460 assert_eq!(
461 left.receive().await.unwrap(),
462 Wrapper {
463 number: 0,
464 wrapped: 1,
465 }
466 );
467 assert_eq!(
468 left.receive().await.unwrap(),
469 Wrapper {
470 number: 1,
471 wrapped: 2,
472 }
473 );
474 assert_eq!(
475 left.receive().await.unwrap(),
476 Wrapper {
477 number: 2,
478 wrapped: 3,
479 }
480 );
481 }
482}