1use crate::{Blocker, CheckedSender, Receiver, Recipients, Sender};
4use commonware_actor::{mailbox, Feedback, Unreliable};
5use commonware_codec::{Codec, Error};
6use commonware_cryptography::PublicKey;
7use commonware_macros::select_loop;
8use commonware_parallel::Strategy;
9use commonware_runtime::{
10 iobuf::EncodeExt, spawn_cell, BufferPool, ContextCell, Handle, Metrics, Spawner,
11};
12use commonware_utils::futures::Pool;
13use std::{collections::VecDeque, num::NonZeroUsize, time::SystemTime};
14
15pub const fn wrap<S: Sender, R: Receiver, V: Codec>(
17 config: V::Cfg,
18 pool: BufferPool,
19 sender: S,
20 receiver: R,
21) -> (WrappedSender<S, V>, WrappedReceiver<R, V>) {
22 (
23 WrappedSender::new(pool, sender),
24 WrappedReceiver::new(config, receiver),
25 )
26}
27
28pub type WrappedMessage<P, V> = (P, Result<V, Error>);
30
31#[derive(Clone)]
33pub struct WrappedSender<S: Sender, V: Codec> {
34 pool: BufferPool,
35 sender: S,
36 _phantom_v: std::marker::PhantomData<V>,
37}
38
39impl<S: Sender, V: Codec> WrappedSender<S, V> {
40 pub const fn new(pool: BufferPool, sender: S) -> Self {
42 Self {
43 pool,
44 sender,
45 _phantom_v: std::marker::PhantomData,
46 }
47 }
48
49 pub fn send(
51 &mut self,
52 recipients: Recipients<S::PublicKey>,
53 message: V,
54 priority: bool,
55 ) -> Vec<S::PublicKey> {
56 self.send_ref(recipients, &message, priority)
57 }
58
59 pub fn send_ref(
61 &mut self,
62 recipients: Recipients<S::PublicKey>,
63 message: &V,
64 priority: bool,
65 ) -> Vec<S::PublicKey> {
66 let encoded = message.encode_with_pool(&self.pool);
67 self.sender.send(recipients, encoded, priority)
68 }
69
70 pub fn check(
73 &mut self,
74 recipients: Recipients<S::PublicKey>,
75 ) -> Result<CheckedWrappedSender<'_, S, V>, SystemTime> {
76 self.sender
77 .check(recipients)
78 .map(|checked| CheckedWrappedSender {
79 pool: &self.pool,
80 sender: checked,
81 _phantom_v: std::marker::PhantomData,
82 })
83 }
84}
85
86#[derive(Debug)]
88pub struct CheckedWrappedSender<'a, S: Sender, V: Codec> {
89 pool: &'a BufferPool,
90 sender: S::Checked<'a>,
91 _phantom_v: std::marker::PhantomData<V>,
92}
93
94impl<'a, S: Sender, V: Codec> CheckedWrappedSender<'a, S, V> {
95 pub fn recipients(&self) -> Vec<S::PublicKey> {
96 self.sender.recipients()
97 }
98
99 pub fn send(self, message: V, priority: bool) -> Unreliable<Feedback> {
100 self.send_ref(&message, priority)
101 }
102
103 pub fn send_ref(self, message: &V, priority: bool) -> Unreliable<Feedback> {
104 let encoded = message.encode_with_pool(self.pool);
105 self.sender.send(encoded, priority)
106 }
107}
108
109pub struct WrappedReceiver<R: Receiver, V: Codec> {
111 config: V::Cfg,
112 receiver: R,
113}
114
115impl<R: Receiver, V: Codec> WrappedReceiver<R, V> {
116 pub const fn new(config: V::Cfg, receiver: R) -> Self {
118 Self { config, receiver }
119 }
120
121 pub async fn recv(&mut self) -> Result<WrappedMessage<R::PublicKey, V>, R::Error> {
123 let (pk, bytes) = self.receiver.recv().await?;
124 let decoded = match V::decode_cfg(bytes.as_ref(), &self.config) {
125 Ok(decoded) => decoded,
126 Err(e) => {
127 return Ok((pk, Err(e)));
128 }
129 };
130 Ok((pk, Ok(decoded)))
131 }
132}
133
134struct Decoded<P: PublicKey, V>(P, V);
145
146impl<P: PublicKey, V> mailbox::UnreliablePolicy for Decoded<P, V> {
147 type Overflow = VecDeque<Self>;
148
149 fn handle(_overflow: &mut Self::Overflow, _message: Self) -> bool {
150 false
151 }
152}
153
154pub struct BackgroundReceiver<P: PublicKey, V> {
156 receiver: mailbox::UnreliableReceiver<Decoded<P, V>>,
157}
158
159impl<P: PublicKey, V> BackgroundReceiver<P, V> {
160 pub async fn recv(&mut self) -> Option<(P, V)> {
162 self.receiver
163 .recv()
164 .await
165 .map(|Decoded(peer, value)| (peer, value))
166 }
167}
168
169pub struct WrappedBackgroundReceiver<E, P, B, R, V, T>
170where
171 E: Spawner,
172 P: PublicKey,
173 B: Blocker<PublicKey = P>,
174 R: Receiver<PublicKey = P>,
175 V: Codec + Send,
176 T: Strategy,
177{
178 context: ContextCell<E>,
179 receiver: R,
180 codec_config: V::Cfg,
181 blocker: B,
182 sender: mailbox::UnreliableSender<Decoded<P, V>>,
183 strategy: T,
184}
185
186impl<E, P, B, R, V, T> WrappedBackgroundReceiver<E, P, B, R, V, T>
187where
188 E: Spawner + Metrics,
189 P: PublicKey,
190 B: Blocker<PublicKey = P>,
191 R: Receiver<PublicKey = P>,
192 V: Codec + Send + 'static,
193 T: Strategy,
194{
195 pub fn new(
199 context: E,
200 receiver: R,
201 codec_config: V::Cfg,
202 blocker: B,
203 channel_capacity: NonZeroUsize,
204 strategy: T,
205 ) -> (Self, BackgroundReceiver<P, V>) {
206 let (tx, rx) = mailbox::new_unreliable(context.child("mailbox"), channel_capacity);
207 (
208 Self {
209 context: ContextCell::new(context),
210 receiver,
211 codec_config,
212 blocker,
213 sender: tx,
214 strategy,
215 },
216 BackgroundReceiver { receiver: rx },
217 )
218 }
219
220 pub fn start(mut self) -> Handle<()> {
225 spawn_cell!(self.context, self.run())
226 }
227
228 async fn run(mut self) {
234 let decode_queue_capacity = self.strategy.manual().parallelism();
235 let mut decode_pool = Pool::default();
236 let mut receiver_closed = false;
237
238 select_loop! {
239 self.context,
240 on_start => {
241 while decode_pool.len() >= decode_queue_capacity
242 || (receiver_closed && !decode_pool.is_empty())
243 {
244 let result = decode_pool.next_completed().await;
245 Self::handle_decode_result(&mut self.blocker, &mut self.sender, result);
246 }
247 if receiver_closed && decode_pool.is_empty() {
248 break;
249 }
250 },
251 on_stopped => {},
252 result = decode_pool.next_completed() => {
254 Self::handle_decode_result(&mut self.blocker, &mut self.sender, result);
255 },
256 Ok((peer, bytes)) = self.receiver.recv() else {
258 receiver_closed = true;
259 continue;
260 } => {
261 let config = self.codec_config.clone();
262 let handle = self.strategy.spawn(move |_| {
263 let result = V::decode_cfg(bytes.as_ref(), &config);
264 (peer, result)
265 });
266 decode_pool.push(handle);
267 },
268 }
269 }
270
271 fn handle_decode_result(
272 blocker: &mut B,
273 sender: &mut mailbox::UnreliableSender<Decoded<P, V>>,
274 result: (P, Result<V, commonware_codec::Error>),
275 ) {
276 let (peer, decode_result) = result;
277 match decode_result {
278 Ok(value) => {
279 let _ = sender.enqueue(Decoded(peer, value));
280 }
281 Err(err) => {
282 crate::block!(blocker, peer, ?err, "received invalid message");
283 }
284 }
285 }
286}
287
288#[cfg(test)]
289mod tests {
290 use super::*;
291 use crate::{
292 simulated::{self, Link, Network, Oracle},
293 Manager as _, Recipients,
294 };
295 use commonware_actor::Feedback;
296 use commonware_codec::Encode;
297 use commonware_cryptography::{
298 ed25519::{PrivateKey, PublicKey},
299 Signer,
300 };
301 use commonware_macros::test_traced;
302 use commonware_parallel::{Manual, Sequential, Strategy};
303 use commonware_runtime::{deterministic, Clock as _, IoBuf, Quota, Runner, Supervisor as _};
304 use commonware_utils::{channel::mpsc, ordered::Set, NZUsize};
305 use std::{
306 io,
307 num::{NonZeroU32, NonZeroUsize},
308 sync::{
309 atomic::{AtomicUsize, Ordering},
310 Arc,
311 },
312 time::Duration,
313 };
314
315 const LINK: Link = Link {
316 latency: Duration::from_millis(0),
317 jitter: Duration::from_millis(0),
318 success_rate: 1.0,
319 };
320
321 const TEST_QUOTA: Quota = Quota::per_second(NonZeroU32::MAX);
322
323 fn start_network(context: deterministic::Context) -> Oracle<PublicKey, deterministic::Context> {
324 let (network, oracle) = Network::new(
325 context.child("network"),
326 simulated::Config {
327 max_size: 1024 * 1024,
328 disconnect_on_block: true,
329 tracked_peer_sets: NZUsize!(1),
330 },
331 );
332 network.start();
333 oracle
334 }
335
336 fn pk(seed: u64) -> PublicKey {
337 PrivateKey::from_seed(seed).public_key()
338 }
339
340 fn track_peers<I>(oracle: &Oracle<PublicKey, deterministic::Context>, index: u64, peers: I)
341 where
342 I: IntoIterator<Item = PublicKey>,
343 {
344 oracle.manager().track(index, Set::from_iter_dedup(peers));
345 }
346
347 async fn link_bidirectional(
348 oracle: &mut Oracle<PublicKey, deterministic::Context>,
349 a: PublicKey,
350 b: PublicKey,
351 ) {
352 oracle.add_link(a.clone(), b.clone(), LINK).await.unwrap();
353 oracle.add_link(b, a, LINK).await.unwrap();
354 }
355
356 #[derive(Debug)]
357 struct MockReceiver<P: commonware_cryptography::PublicKey> {
358 receiver: mpsc::UnboundedReceiver<crate::Message<P>>,
359 }
360
361 impl<P: commonware_cryptography::PublicKey> crate::Receiver for MockReceiver<P> {
362 type Error = io::Error;
363 type PublicKey = P;
364
365 async fn recv(&mut self) -> Result<crate::Message<Self::PublicKey>, Self::Error> {
366 self.receiver
367 .recv()
368 .await
369 .ok_or_else(|| io::Error::from(io::ErrorKind::BrokenPipe))
370 }
371 }
372
373 #[derive(Debug)]
374 struct CountingReceiver<P: commonware_cryptography::PublicKey> {
375 receiver: mpsc::UnboundedReceiver<crate::Message<P>>,
376 received: Arc<AtomicUsize>,
377 }
378
379 impl<P: commonware_cryptography::PublicKey> crate::Receiver for CountingReceiver<P> {
380 type Error = io::Error;
381 type PublicKey = P;
382
383 async fn recv(&mut self) -> Result<crate::Message<Self::PublicKey>, Self::Error> {
384 self.received.fetch_add(1, Ordering::SeqCst);
385 self.receiver
386 .recv()
387 .await
388 .ok_or_else(|| io::Error::from(io::ErrorKind::BrokenPipe))
389 }
390 }
391
392 #[derive(Clone, Default)]
393 struct NoopBlocker;
394
395 impl crate::Blocker for NoopBlocker {
396 type PublicKey = PublicKey;
397
398 fn block(&mut self, _peer: Self::PublicKey) -> Feedback {
399 Feedback::Ok
400 }
401 }
402
403 #[derive(Clone, Debug)]
404 struct TestStrategy {
405 parallelism: NonZeroUsize,
406 pending: bool,
407 }
408
409 impl TestStrategy {
410 const fn complete(parallelism: NonZeroUsize) -> Self {
411 Self {
412 parallelism,
413 pending: false,
414 }
415 }
416
417 const fn pending(parallelism: NonZeroUsize) -> Self {
418 Self {
419 parallelism,
420 pending: true,
421 }
422 }
423 }
424
425 impl Strategy for TestStrategy {
426 fn manual(&self) -> Manual<Self> {
427 Manual::new(self.clone(), self.parallelism)
428 }
429
430 fn spawn<F, T>(&self, f: F) -> impl core::future::Future<Output = T> + Send + 'static
431 where
432 F: FnOnce(Self) -> T + Send + 'static,
433 T: Send + 'static,
434 {
435 let pending = self.pending;
436 let s = self.clone();
437 async move {
438 if pending {
439 futures::future::pending::<()>().await;
440 }
441 f(s)
442 }
443 }
444
445 fn fold_init<I, INIT, T, R, ID, F, RD>(
446 &self,
447 iter: I,
448 init: INIT,
449 identity: ID,
450 fold_op: F,
451 reduce_op: RD,
452 ) -> R
453 where
454 I: IntoIterator<IntoIter: Send, Item: Send> + Send,
455 INIT: Fn() -> T + Send + Sync,
456 T: Send,
457 R: Send,
458 ID: Fn() -> R + Send + Sync,
459 F: Fn(R, &mut T, I::Item) -> R + Send + Sync,
460 RD: Fn(R, R) -> R + Send + Sync,
461 {
462 Sequential.fold_init(iter, init, identity, fold_op, reduce_op)
463 }
464
465 fn try_fold<I, R, E, ID, F, RD>(
466 &self,
467 iter: I,
468 identity: ID,
469 fold_op: F,
470 reduce_op: RD,
471 ) -> Result<R, E>
472 where
473 I: IntoIterator<IntoIter: Send, Item: Send> + Send,
474 R: Send,
475 E: Send,
476 ID: Fn() -> R + Send + Sync,
477 F: Fn(R, I::Item) -> Result<R, E> + Send + Sync,
478 RD: Fn(R, R) -> R + Send + Sync,
479 {
480 Sequential.try_fold(iter, identity, fold_op, reduce_op)
481 }
482
483 fn run<R, SEQ, PAR>(&self, len: usize, serial: SEQ, parallel: PAR) -> R
484 where
485 R: Send,
486 SEQ: FnOnce() -> R + Send,
487 PAR: FnOnce() -> R + Send,
488 {
489 Sequential.run(len, serial, parallel)
490 }
491
492 fn try_run<R, E, SEQ, PAR>(&self, len: usize, serial: SEQ, parallel: PAR) -> Result<R, E>
493 where
494 R: Send,
495 E: Send,
496 SEQ: FnOnce() -> Result<R, E> + Send,
497 PAR: FnOnce() -> Result<R, E> + Send,
498 {
499 Sequential.try_run(len, serial, parallel)
500 }
501
502 fn join<A, B, RA, RB>(&self, a: A, b: B) -> (RA, RB)
503 where
504 A: FnOnce() -> RA + Send,
505 B: FnOnce() -> RB + Send,
506 RA: Send,
507 RB: Send,
508 {
509 Sequential.join(a, b)
510 }
511
512 fn sort_by<T, C>(&self, items: &mut [T], compare: C)
513 where
514 T: Send,
515 C: Fn(&T, &T) -> std::cmp::Ordering + Send + Sync,
516 {
517 Sequential.sort_by(items, compare);
518 }
519 }
520
521 #[test_traced]
522 fn test_valid_messages_forwarded() {
523 let executor = deterministic::Runner::default();
524 executor.start(|context| async move {
525 let mut oracle = start_network(context.child("network"));
526
527 let pk1 = pk(0);
528 let pk2 = pk(1);
529 let control1 = oracle.control(pk1.clone());
530 let control2 = oracle.control(pk2.clone());
531 track_peers(&oracle, 0, [pk1.clone(), pk2.clone()]);
532 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
533
534 let (mut sender1, _) = control1.register(0, TEST_QUOTA).await.unwrap();
535 let (_, receiver2) = control2.register(0, TEST_QUOTA).await.unwrap();
536
537 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
538 context.child("bg"),
539 receiver2,
540 (),
541 control2.clone(),
542 NZUsize!(16),
543 Sequential,
544 );
545 let _handle = bg.start();
546
547 let msg: u32 = 42;
548 let _ = sender1.send(Recipients::One(pk2.clone()), msg.encode(), true);
549
550 let (from, value) = rx.recv().await.unwrap();
551 assert_eq!(from, pk1);
552 assert_eq!(value, 42u32);
553 });
554 }
555
556 #[test_traced]
557 fn test_invalid_codec_blocks_peer() {
558 let executor = deterministic::Runner::default();
559 executor.start(|context| async move {
560 let mut oracle = start_network(context.child("network"));
561
562 let pk1 = pk(0);
563 let pk2 = pk(1);
564 let pk3 = pk(2);
565 let control1 = oracle.control(pk1.clone());
566 let control2 = oracle.control(pk2.clone());
567 track_peers(&oracle, 0, [pk1.clone(), pk2.clone(), pk3.clone()]);
568 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
569
570 let (mut sender1, _) = control1.register(0, TEST_QUOTA).await.unwrap();
571 let (_, receiver2) = control2.register(0, TEST_QUOTA).await.unwrap();
572
573 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
574 context.child("bg"),
575 receiver2,
576 (),
577 control2.clone(),
578 NZUsize!(16),
579 Sequential,
580 );
581 let _handle = bg.start();
582
583 let invalid = IoBuf::from(vec![0xFFu8]);
585 let _ = sender1.send(Recipients::One(pk2.clone()), invalid, true);
586
587 let control3 = oracle.control(pk3.clone());
590 link_bidirectional(&mut oracle, pk3.clone(), pk2.clone()).await;
591 let (mut sender3, _) = control3.register(0, TEST_QUOTA).await.unwrap();
592
593 let msg: u32 = 99;
594 let _ = sender3.send(Recipients::One(pk2.clone()), msg.encode(), true);
595
596 let (from, value) = rx.recv().await.unwrap();
597 assert_eq!(from, pk3);
598 assert_eq!(value, 99u32);
599
600 loop {
602 let blocked = oracle.blocked().await.unwrap();
603 if blocked.contains(&(pk2.clone(), pk1.clone())) {
604 break;
605 }
606
607 context.sleep(Duration::from_millis(1)).await;
608 }
609 });
610 }
611
612 #[test_traced]
613 fn test_multiple_valid_messages() {
614 let executor = deterministic::Runner::default();
615 executor.start(|context| async move {
616 let mut oracle = start_network(context.child("network"));
617
618 let pk1 = pk(0);
619 let pk2 = pk(1);
620 let control1 = oracle.control(pk1.clone());
621 let control2 = oracle.control(pk2.clone());
622 track_peers(&oracle, 0, [pk1.clone(), pk2.clone()]);
623 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
624
625 let (mut sender1, _) = control1.register(0, TEST_QUOTA).await.unwrap();
626 let (_, receiver2) = control2.register(0, TEST_QUOTA).await.unwrap();
627
628 let count = 20;
629 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
630 context.child("bg"),
631 receiver2,
632 (),
633 control2.clone(),
634 NZUsize!(20),
635 Sequential,
636 );
637 let _handle = bg.start();
638
639 for i in 0..count {
640 let msg: u32 = i;
641 let _ = sender1.send(Recipients::One(pk2.clone()), msg.encode(), true);
642 }
643
644 let mut received = Vec::new();
645 for _ in 0..count {
646 let (from, value) = rx.recv().await.unwrap();
647 assert_eq!(from, pk1);
648 received.push(value);
649 }
650 received.sort();
651 assert_eq!(received, (0..count).collect::<Vec<u32>>());
652 });
653 }
654
655 #[test_traced]
656 fn test_decode_with_strategy() {
657 let executor = deterministic::Runner::default();
658 executor.start(|context| async move {
659 let mut oracle = start_network(context.child("network"));
660
661 let pk1 = pk(0);
662 let pk2 = pk(1);
663 let control1 = oracle.control(pk1.clone());
664 let control2 = oracle.control(pk2.clone());
665 track_peers(&oracle, 0, [pk1.clone(), pk2.clone()]);
666 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
667
668 let (mut sender1, _) = control1.register(0, TEST_QUOTA).await.unwrap();
669 let (_, receiver2) = control2.register(0, TEST_QUOTA).await.unwrap();
670
671 let count = 50u32;
674 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
675 context.child("bg"),
676 receiver2,
677 (),
678 control2.clone(),
679 NZUsize!(50),
680 TestStrategy::complete(NZUsize!(4)),
681 );
682 let _handle = bg.start();
683
684 for i in 0..count {
685 let _ = sender1.send(Recipients::One(pk2.clone()), i.encode(), true);
686 }
687
688 let mut received = Vec::new();
689 for _ in 0..count {
690 let (from, value) = rx.recv().await.unwrap();
691 assert_eq!(from, pk1);
692 received.push(value);
693 }
694 received.sort();
695 assert_eq!(received, (0..count).collect::<Vec<u32>>());
696 });
697 }
698
699 #[test_traced]
700 fn test_invalid_among_valid_only_blocks_offender() {
701 let executor = deterministic::Runner::default();
702 executor.start(|context| async move {
703 let mut oracle = start_network(context.child("network"));
704
705 let pk1 = pk(0);
706 let pk2 = pk(1);
707 let pk3 = pk(2);
708 let control1 = oracle.control(pk1.clone());
709 let control2 = oracle.control(pk2.clone());
710 let control3 = oracle.control(pk3.clone());
711 track_peers(&oracle, 0, [pk1.clone(), pk2.clone(), pk3.clone()]);
712 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
713 link_bidirectional(&mut oracle, pk3.clone(), pk2.clone()).await;
714
715 let (mut sender1, _) = control1.register(0, TEST_QUOTA).await.unwrap();
716 let (_, receiver2) = control2.register(0, TEST_QUOTA).await.unwrap();
717 let (mut sender3, _) = control3.register(0, TEST_QUOTA).await.unwrap();
718
719 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
720 context.child("bg"),
721 receiver2,
722 (),
723 control2.clone(),
724 NZUsize!(16),
725 Sequential,
726 );
727 let _handle = bg.start();
728
729 let _ = sender3.send(Recipients::One(pk2.clone()), 10u32.encode(), true);
731
732 let _ = sender1.send(Recipients::One(pk2.clone()), IoBuf::from(vec![0xFF]), true);
734
735 let _ = sender3.send(Recipients::One(pk2.clone()), 20u32.encode(), true);
737
738 let mut values = Vec::new();
740 for _ in 0..2 {
741 let (from, value) = rx.recv().await.unwrap();
742 assert_eq!(from, pk3);
743 values.push(value);
744 }
745 values.sort();
746 assert_eq!(values, vec![10u32, 20]);
747
748 loop {
750 let blocked = oracle.blocked().await.unwrap();
751 assert!(!blocked.contains(&(pk2.clone(), pk3.clone())));
752 if blocked.contains(&(pk2.clone(), pk1.clone())) {
753 break;
754 }
755
756 context.sleep(Duration::from_millis(1)).await;
757 }
758 });
759 }
760
761 #[test_traced]
762 fn test_decoded_messages_drop_when_receiver_full() {
763 let executor = deterministic::Runner::default();
764 executor.start(|context| async move {
765 let sender = pk(0);
766 let (tx, receiver) = mpsc::unbounded_channel();
767
768 for i in 0..2u32 {
769 tx.send((sender.clone(), IoBuf::from(i.encode())))
770 .expect("mock receiver should be open");
771 }
772 drop(tx);
773
774 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
775 context.child("bg"),
776 MockReceiver { receiver },
777 (),
778 NoopBlocker,
779 NZUsize!(1),
780 Sequential,
781 );
782 let handle = bg.start();
783 handle.await.expect("background receiver should complete");
784
785 let (from, value) = rx.recv().await.unwrap();
786 assert_eq!(from, sender);
787 assert_eq!(value, 0);
788 assert!(rx.recv().await.is_none());
789 });
790 }
791
792 #[test_traced]
793 fn test_decode_backpressure_limits_raw_receives() {
794 let executor = deterministic::Runner::default();
795 executor.start(|context| async move {
796 let sender = pk(0);
797 let (tx, receiver) = mpsc::unbounded_channel();
798 let received = Arc::new(AtomicUsize::new(0));
799
800 for i in 0..10u32 {
801 tx.send((sender.clone(), IoBuf::from(i.encode())))
802 .expect("mock receiver should be open");
803 }
804
805 let (bg, _rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
806 context.child("bg"),
807 CountingReceiver {
808 receiver,
809 received: received.clone(),
810 },
811 (),
812 NoopBlocker,
813 NZUsize!(16),
814 TestStrategy::pending(NZUsize!(2)),
815 );
816 let handle = bg.start();
817
818 while received.load(Ordering::SeqCst) < 2 {
819 context.sleep(Duration::from_millis(1)).await;
820 }
821 for _ in 0..10 {
822 context.sleep(Duration::from_millis(1)).await;
823 assert_eq!(received.load(Ordering::SeqCst), 2);
824 }
825
826 drop(handle);
827 });
828 }
829
830 #[test_traced]
831 fn test_drain_decode_pool_after_receiver_closure() {
832 let executor = deterministic::Runner::default();
833 executor.start(|context| async move {
834 let sender = pk(0);
835 let (tx, receiver) = mpsc::unbounded_channel();
836 let count = 64u32;
837
838 for i in 0..count {
839 tx.send((sender.clone(), IoBuf::from(i.encode())))
840 .expect("mock receiver should be open");
841 }
842 drop(tx);
843
844 let (bg, mut rx) = WrappedBackgroundReceiver::<_, _, _, _, u32, _>::new(
845 context.child("bg"),
846 MockReceiver { receiver },
847 (),
848 NoopBlocker,
849 NZUsize!(64),
850 Sequential,
851 );
852 let _handle = bg.start();
853
854 let mut values = Vec::new();
855 while let Some((from, value)) = rx.recv().await {
856 assert_eq!(from, sender);
857 values.push(value);
858 }
859 values.sort_unstable();
860
861 assert_eq!(values, (0..count).collect::<Vec<u32>>());
862 });
863 }
864}