1use crate::{Channel, CheckedSender, LimitedSender, Message, Receiver, Recipients, Sender};
12use commonware_actor::{Feedback, Unreliable};
13use commonware_codec::{varint::UInt, Encode, Error as CodecError, ReadExt};
14use commonware_macros::select_loop;
15use commonware_runtime::{spawn_cell, ContextCell, Handle, IoBuf, IoBufs, Spawner};
16use commonware_utils::channel::{
17 fallible::FallibleExt,
18 mpsc::{self, error::TrySendError},
19 oneshot,
20};
21use std::{collections::HashMap, fmt::Debug, time::SystemTime};
22use thiserror::Error;
23use tracing::debug;
24
25#[derive(Error, Debug)]
27pub enum Error {
28 #[error("subchannel already registered: {0}")]
29 AlreadyRegistered(Channel),
30 #[error("muxer is closed")]
31 Closed,
32 #[error("recv failed")]
33 RecvFailed,
34}
35
36pub fn parse(mut buf: IoBuf) -> Result<(Channel, IoBuf), CodecError> {
38 let subchannel: Channel = UInt::read(&mut buf)?.into();
39 Ok((subchannel, buf))
40}
41
42enum Control<R: Receiver> {
44 Register {
45 subchannel: Channel,
46 sender: oneshot::Sender<mpsc::Receiver<Message<R::PublicKey>>>,
47 },
48 Deregister {
49 subchannel: Channel,
50 },
51}
52
53type Routes<P> = HashMap<Channel, mpsc::Sender<Message<P>>>;
55
56type BackupResponse<P> = (Channel, Message<P>);
59
60pub struct Muxer<E: Spawner, S: Sender, R: Receiver> {
62 context: ContextCell<E>,
63 sender: S,
64 receiver: R,
65 mailbox_size: usize,
66 control_rx: mpsc::UnboundedReceiver<Control<R>>,
67 routes: Routes<R::PublicKey>,
68 backup: Option<mpsc::Sender<BackupResponse<R::PublicKey>>>,
69}
70
71impl<E: Spawner, S: Sender, R: Receiver> Muxer<E, S, R> {
72 pub fn new(context: E, sender: S, receiver: R, mailbox_size: usize) -> (Self, MuxHandle<S, R>) {
75 Self::builder(context, sender, receiver, mailbox_size).build()
76 }
77
78 pub fn builder(
80 context: E,
81 sender: S,
82 receiver: R,
83 mailbox_size: usize,
84 ) -> MuxerBuilder<E, S, R> {
85 let (control_tx, control_rx) = mpsc::unbounded_channel();
86 let mux = Self {
87 context: ContextCell::new(context),
88 sender,
89 receiver,
90 mailbox_size,
91 control_rx,
92 routes: HashMap::new(),
93 backup: None,
94 };
95
96 let mux_handle = MuxHandle {
97 sender: mux.sender.clone(),
98 control_tx,
99 };
100
101 MuxerBuilder { mux, mux_handle }
102 }
103
104 pub fn start(mut self) -> Handle<Result<(), R::Error>> {
106 spawn_cell!(self.context, self.run())
107 }
108
109 pub async fn run(mut self) -> Result<(), R::Error> {
114 select_loop! {
115 self.context,
116 on_stopped => {
117 debug!("context shutdown, stopping muxer");
118 },
119 Some(control) = self.control_rx.recv() else {
122 return Ok(());
125 } => match control {
126 Control::Register { subchannel, sender } => {
127 if self.routes.contains_key(&subchannel) {
129 continue;
130 }
131
132 let (tx, rx) = mpsc::channel(self.mailbox_size);
134 self.routes.insert(subchannel, tx);
135 let _ = sender.send(rx);
136 }
137 Control::Deregister { subchannel } => {
138 self.routes.remove(&subchannel);
140 }
141 },
142 message = self.receiver.recv() => {
144 let (pk, bytes) = message?;
146 let (subchannel, bytes) = match parse(bytes) {
147 Ok(parsed) => parsed,
148 Err(_) => {
149 debug!(?pk, "invalid message: missing subchannel");
150 continue;
151 }
152 };
153
154 let Some(sender) = self.routes.get_mut(&subchannel) else {
156 if let Some(backup) = &mut self.backup {
158 if let Err(e) = backup.try_send((subchannel, (pk, bytes))) {
159 debug!(?subchannel, ?e, "failed to send message to backup channel");
160 }
161 }
162
163 continue;
166 };
167
168 if let Err(e) = sender.try_send((pk, bytes)) {
171 if matches!(e, TrySendError::Closed(_)) {
173 self.routes.remove(&subchannel);
175 debug!(?subchannel, "subchannel receiver dropped, removing route");
176 } else {
177 debug!(?subchannel, "subchannel full, dropping message");
179 }
180 }
181 },
182 }
183
184 Ok(())
185 }
186}
187
188#[derive(Clone)]
190pub struct MuxHandle<S: Sender, R: Receiver> {
191 sender: S,
192 control_tx: mpsc::UnboundedSender<Control<R>>,
193}
194
195impl<S: Sender, R: Receiver> MuxHandle<S, R> {
196 pub async fn register(
201 &mut self,
202 subchannel: Channel,
203 ) -> Result<(SubSender<S>, SubReceiver<R>), Error> {
204 let (tx, rx) = oneshot::channel();
205 self.control_tx
206 .send(Control::Register {
207 subchannel,
208 sender: tx,
209 })
210 .map_err(|_| Error::Closed)?;
211 let receiver = rx.await.map_err(|_| Error::AlreadyRegistered(subchannel))?;
212
213 Ok((
214 SubSender {
215 subchannel,
216 inner: GlobalSender::new(self.sender.clone()),
217 },
218 SubReceiver {
219 receiver,
220 control_tx: Some(self.control_tx.clone()),
221 subchannel,
222 },
223 ))
224 }
225}
226
227#[derive(Clone, Debug)]
229pub struct SubSender<S: Sender> {
230 inner: GlobalSender<S>,
231 subchannel: Channel,
232}
233
234impl<S: Sender> LimitedSender for SubSender<S> {
235 type PublicKey = S::PublicKey;
236 type Checked<'a> = CheckedGlobalSender<'a, S>;
237
238 fn check(
239 &mut self,
240 recipients: Recipients<Self::PublicKey>,
241 ) -> Result<Self::Checked<'_>, SystemTime> {
242 self.inner
243 .check(recipients)
244 .map(|checked| checked.with_subchannel(self.subchannel))
245 }
246}
247
248pub struct SubReceiver<R: Receiver> {
250 receiver: mpsc::Receiver<Message<R::PublicKey>>,
251 control_tx: Option<mpsc::UnboundedSender<Control<R>>>,
252 subchannel: Channel,
253}
254
255impl<R: Receiver> Receiver for SubReceiver<R> {
256 type Error = Error;
257 type PublicKey = R::PublicKey;
258
259 async fn recv(&mut self) -> Result<Message<Self::PublicKey>, Self::Error> {
260 self.receiver.recv().await.ok_or(Error::RecvFailed)
261 }
262}
263
264impl<R: Receiver> Debug for SubReceiver<R> {
265 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
266 write!(f, "SubReceiver({})", self.subchannel)
267 }
268}
269
270impl<R: Receiver> Drop for SubReceiver<R> {
271 fn drop(&mut self) {
272 let control_tx = self
274 .control_tx
275 .take()
276 .expect("SubReceiver::drop called twice");
277
278 control_tx.send_lossy(Control::Deregister {
280 subchannel: self.subchannel,
281 });
282 }
283}
284
285#[derive(Clone, Debug)]
287pub struct GlobalSender<S: Sender> {
288 inner: S,
289}
290
291impl<S: Sender> GlobalSender<S> {
292 pub const fn new(inner: S) -> Self {
294 Self { inner }
295 }
296
297 pub fn send(
299 &mut self,
300 subchannel: Channel,
301 recipients: Recipients<S::PublicKey>,
302 payload: impl Into<IoBufs> + Send,
303 priority: bool,
304 ) -> Unreliable<Feedback> {
305 self.check(recipients).map_or_else(
306 |_| Unreliable::Rejected,
307 |checked| checked.with_subchannel(subchannel).send(payload, priority),
308 )
309 }
310}
311
312impl<S: Sender> LimitedSender for GlobalSender<S> {
313 type PublicKey = S::PublicKey;
314 type Checked<'a> = CheckedGlobalSender<'a, S>;
315
316 fn check(
317 &mut self,
318 recipients: Recipients<Self::PublicKey>,
319 ) -> Result<Self::Checked<'_>, SystemTime> {
320 self.inner
321 .check(recipients)
322 .map(|checked| CheckedGlobalSender {
323 subchannel: None,
324 inner: checked,
325 })
326 }
327}
328
329pub struct CheckedGlobalSender<'a, S: Sender> {
331 subchannel: Option<Channel>,
332 inner: S::Checked<'a>,
333}
334
335impl<'a, S: Sender> CheckedGlobalSender<'a, S> {
336 pub const fn with_subchannel(mut self, subchannel: Channel) -> Self {
338 self.subchannel = Some(subchannel);
339 self
340 }
341}
342
343impl<'a, S: Sender> CheckedSender for CheckedGlobalSender<'a, S> {
344 type PublicKey = S::PublicKey;
345
346 fn recipients(&self) -> Vec<Self::PublicKey> {
347 self.inner.recipients()
348 }
349
350 fn send(self, message: impl Into<IoBufs> + Send, priority: bool) -> Unreliable<Feedback> {
351 let subchannel = UInt(self.subchannel.expect("subchannel not set"));
352 let mut message = message.into();
353 message.prepend(subchannel.encode().into());
354 self.inner.send(message, priority)
355 }
356}
357
358pub trait Builder {
360 type Output;
362
363 fn build(self) -> Self::Output;
365}
366
367pub struct MuxerBuilder<E: Spawner, S: Sender, R: Receiver> {
369 mux: Muxer<E, S, R>,
370 mux_handle: MuxHandle<S, R>,
371}
372
373impl<E: Spawner, S: Sender, R: Receiver> Builder for MuxerBuilder<E, S, R> {
374 type Output = (Muxer<E, S, R>, MuxHandle<S, R>);
375
376 fn build(self) -> Self::Output {
377 (self.mux, self.mux_handle)
378 }
379}
380
381impl<E: Spawner, S: Sender, R: Receiver> MuxerBuilder<E, S, R> {
382 pub fn with_backup(mut self) -> MuxerBuilderWithBackup<E, S, R> {
384 let (tx, rx) = mpsc::channel(self.mux.mailbox_size);
385 self.mux.backup = Some(tx);
386
387 MuxerBuilderWithBackup {
388 mux: self.mux,
389 mux_handle: self.mux_handle,
390 backup_rx: rx,
391 }
392 }
393
394 pub fn with_global_sender(self) -> MuxerBuilderWithGlobalSender<E, S, R> {
396 let global_sender = GlobalSender::new(self.mux.sender.clone());
397
398 MuxerBuilderWithGlobalSender {
399 mux: self.mux,
400 mux_handle: self.mux_handle,
401 global_sender,
402 }
403 }
404}
405
406pub struct MuxerBuilderWithBackup<E: Spawner, S: Sender, R: Receiver> {
408 mux: Muxer<E, S, R>,
409 mux_handle: MuxHandle<S, R>,
410 backup_rx: mpsc::Receiver<BackupResponse<R::PublicKey>>,
411}
412
413impl<E: Spawner, S: Sender, R: Receiver> MuxerBuilderWithBackup<E, S, R> {
414 pub fn with_global_sender(self) -> MuxerBuilderAllOpts<E, S, R> {
416 let global_sender = GlobalSender::new(self.mux.sender.clone());
417
418 MuxerBuilderAllOpts {
419 mux: self.mux,
420 mux_handle: self.mux_handle,
421 backup_rx: self.backup_rx,
422 global_sender,
423 }
424 }
425}
426
427impl<E: Spawner, S: Sender, R: Receiver> Builder for MuxerBuilderWithBackup<E, S, R> {
428 type Output = (
429 Muxer<E, S, R>,
430 MuxHandle<S, R>,
431 mpsc::Receiver<BackupResponse<R::PublicKey>>,
432 );
433
434 fn build(self) -> Self::Output {
435 (self.mux, self.mux_handle, self.backup_rx)
436 }
437}
438
439pub struct MuxerBuilderWithGlobalSender<E: Spawner, S: Sender, R: Receiver> {
441 mux: Muxer<E, S, R>,
442 mux_handle: MuxHandle<S, R>,
443 global_sender: GlobalSender<S>,
444}
445
446impl<E: Spawner, S: Sender, R: Receiver> MuxerBuilderWithGlobalSender<E, S, R> {
447 pub fn with_backup(mut self) -> MuxerBuilderAllOpts<E, S, R> {
449 let (tx, rx) = mpsc::channel(self.mux.mailbox_size);
450 self.mux.backup = Some(tx);
451
452 MuxerBuilderAllOpts {
453 mux: self.mux,
454 mux_handle: self.mux_handle,
455 backup_rx: rx,
456 global_sender: self.global_sender,
457 }
458 }
459}
460
461impl<E: Spawner, S: Sender, R: Receiver> Builder for MuxerBuilderWithGlobalSender<E, S, R> {
462 type Output = (Muxer<E, S, R>, MuxHandle<S, R>, GlobalSender<S>);
463
464 fn build(self) -> Self::Output {
465 (self.mux, self.mux_handle, self.global_sender)
466 }
467}
468
469pub struct MuxerBuilderAllOpts<E: Spawner, S: Sender, R: Receiver> {
471 mux: Muxer<E, S, R>,
472 mux_handle: MuxHandle<S, R>,
473 backup_rx: mpsc::Receiver<BackupResponse<R::PublicKey>>,
474 global_sender: GlobalSender<S>,
475}
476
477impl<E: Spawner, S: Sender, R: Receiver> Builder for MuxerBuilderAllOpts<E, S, R> {
478 type Output = (
479 Muxer<E, S, R>,
480 MuxHandle<S, R>,
481 mpsc::Receiver<BackupResponse<R::PublicKey>>,
482 GlobalSender<S>,
483 );
484
485 fn build(self) -> Self::Output {
486 (
487 self.mux,
488 self.mux_handle,
489 self.backup_rx,
490 self.global_sender,
491 )
492 }
493}
494
495#[cfg(test)]
496mod tests {
497 use super::*;
498 use crate::{
499 simulated::{self, Link, Network, Oracle},
500 Manager as _, Provider as _, Recipients,
501 };
502 use commonware_cryptography::{
503 ed25519::{PrivateKey, PublicKey},
504 Signer,
505 };
506 use commonware_macros::{select, test_traced};
507 use commonware_runtime::{deterministic, IoBuf, Quota, Runner, Supervisor as _};
508 use commonware_utils::{ordered::Set, NZUsize};
509 use std::{
510 num::NonZeroU32,
511 time::{Duration, SystemTime},
512 };
513
514 const LINK: Link = Link {
515 latency: Duration::from_millis(0),
516 jitter: Duration::from_millis(0),
517 success_rate: 1.0,
518 };
519 const CAPACITY: usize = 5usize;
520
521 const TEST_QUOTA: Quota = Quota::per_second(NonZeroU32::MAX);
523
524 fn start_network(context: deterministic::Context) -> Oracle<PublicKey, deterministic::Context> {
526 let (network, oracle) = Network::new(
527 context.child("network"),
528 simulated::Config {
529 max_size: 1024 * 1024,
530 disconnect_on_block: true,
531 tracked_peer_sets: NZUsize!(1),
532 },
533 );
534 network.start();
535 oracle
536 }
537
538 fn pk(seed: u64) -> PublicKey {
540 PrivateKey::from_seed(seed).public_key()
541 }
542
543 async fn link_bidirectional(
545 oracle: &mut Oracle<PublicKey, deterministic::Context>,
546 a: PublicKey,
547 b: PublicKey,
548 ) {
549 let mut manager = oracle.manager();
550 let peers = manager.peer_set(0).await.unwrap_or_default();
551 manager.track(
552 0,
553 Set::from_iter_dedup(peers.primary.iter().cloned().chain([a.clone(), b.clone()])),
554 );
555 oracle.add_link(a.clone(), b.clone(), LINK).await.unwrap();
556 oracle.add_link(b, a, LINK).await.unwrap();
557 }
558
559 async fn create_peer(
561 context: &deterministic::Context,
562 oracle: &mut Oracle<PublicKey, deterministic::Context>,
563 seed: u64,
564 ) -> (
565 PublicKey,
566 MuxHandle<impl Sender<PublicKey = PublicKey>, impl Receiver<PublicKey = PublicKey>>,
567 ) {
568 let pubkey = pk(seed);
569 let (sender, receiver) = oracle
570 .control(pubkey.clone())
571 .register(0, TEST_QUOTA)
572 .await
573 .unwrap();
574 let (mux, handle) = Muxer::new(context.child("mux"), sender, receiver, CAPACITY);
575 mux.start();
576 (pubkey, handle)
577 }
578
579 async fn create_peer_with_backup_and_global_sender(
581 context: &deterministic::Context,
582 oracle: &mut Oracle<PublicKey, deterministic::Context>,
583 seed: u64,
584 ) -> (
585 PublicKey,
586 MuxHandle<impl Sender<PublicKey = PublicKey>, impl Receiver<PublicKey = PublicKey>>,
587 mpsc::Receiver<BackupResponse<PublicKey>>,
588 GlobalSender<simulated::Sender<PublicKey, deterministic::Context>>,
589 ) {
590 let pubkey = pk(seed);
591 let (sender, receiver) = oracle
592 .control(pubkey.clone())
593 .register(0, TEST_QUOTA)
594 .await
595 .unwrap();
596 let (mux, handle, backup, global_sender) =
597 Muxer::builder(context.child("mux"), sender, receiver, CAPACITY)
598 .with_backup()
599 .with_global_sender()
600 .build();
601 mux.start();
602 (pubkey, handle, backup, global_sender)
603 }
604
605 fn send_burst<S: Sender>(txs: &mut [SubSender<S>], count: usize) {
607 for i in 0..count {
608 let payload = IoBuf::from(vec![i as u8]);
609 for tx in txs.iter_mut() {
610 tx.send(Recipients::All, payload.clone(), false);
611 }
612 }
613 }
614
615 async fn expect_n_messages(
617 rx: &mut SubReceiver<impl Receiver<PublicKey = PublicKey>>,
618 n: usize,
619 ) {
620 let mut count = 0;
621 loop {
622 select! {
623 res = rx.recv() => {
624 res.expect("should have received message");
625 count += 1;
626 },
627 }
628
629 if count >= n {
630 break;
631 }
632 }
633 assert_eq!(n, count);
634 }
635
636 async fn expect_n_messages_with_backup(
638 rx: &mut SubReceiver<impl Receiver<PublicKey = PublicKey>>,
639 backup_rx: &mut mpsc::Receiver<BackupResponse<PublicKey>>,
640 n: usize,
641 n_backup: usize,
642 ) {
643 let mut count_std = 0;
644 let mut count_backup = 0;
645 loop {
646 select! {
647 res = rx.recv() => {
648 res.expect("should have received message");
649 count_std += 1;
650 },
651 res = backup_rx.recv() => {
652 res.expect("should have received message");
653 count_backup += 1;
654 },
655 }
656
657 if count_std >= n && count_backup >= n_backup {
658 break;
659 }
660 }
661 assert_eq!(n, count_std);
662 assert_eq!(n_backup, count_backup);
663 }
664
665 #[derive(Clone)]
666 struct RateLimitedSender;
667
668 struct UnusedCheckedSender;
669
670 impl CheckedSender for UnusedCheckedSender {
671 type PublicKey = PublicKey;
672
673 fn recipients(&self) -> Vec<Self::PublicKey> {
674 unreachable!("rate-limited sender should not produce a checked sender");
675 }
676
677 fn send(self, _: impl Into<IoBufs> + Send, _: bool) -> Unreliable<Feedback> {
678 unreachable!("rate-limited sender should not send");
679 }
680 }
681
682 impl LimitedSender for RateLimitedSender {
683 type PublicKey = PublicKey;
684 type Checked<'a> = UnusedCheckedSender;
685
686 fn check(
687 &mut self,
688 _: Recipients<Self::PublicKey>,
689 ) -> Result<Self::Checked<'_>, SystemTime> {
690 Err(SystemTime::UNIX_EPOCH)
691 }
692 }
693
694 #[test]
695 fn test_global_sender_rate_limited_send_rejected() {
696 let mut sender = GlobalSender::new(RateLimitedSender);
697 let feedback = sender.send(0, Recipients::One(pk(0)), b"rate-limited", false);
698 assert_eq!(feedback, Unreliable::Rejected);
699 assert!(!feedback.accepted());
700 }
701
702 #[test]
703 fn test_basic_routing() {
704 let executor = deterministic::Runner::default();
706 executor.start(|context| async move {
707 let mut oracle = start_network(context.child("network"));
708
709 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
710 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
711 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
712
713 let (_, mut sub_rx1) = handle1.register(7).await.unwrap();
714 let (mut sub_tx2, _) = handle2.register(7).await.unwrap();
715
716 let payload = IoBuf::from(b"hello");
718 let _ = sub_tx2.send(Recipients::One(pk1.clone()), payload.clone(), false);
719 let (from, bytes) = sub_rx1.recv().await.unwrap();
720 assert_eq!(from, pk2);
721 assert_eq!(bytes, payload);
722 });
723 }
724
725 #[test]
726 fn test_multiple_routes() {
727 let executor = deterministic::Runner::default();
729 executor.start(|context| async move {
730 let mut oracle = start_network(context.child("network"));
731
732 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
733 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
734 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
735
736 let (_, mut rx_a) = handle1.register(10).await.unwrap();
737 let (_, mut rx_b) = handle1.register(20).await.unwrap();
738
739 let (mut tx2_a, _) = handle2.register(10).await.unwrap();
740 let (mut tx2_b, _) = handle2.register(20).await.unwrap();
741
742 let payload_a = IoBuf::from(b"A");
743 let payload_b = IoBuf::from(b"B");
744 let _ = tx2_a.send(Recipients::One(pk1.clone()), payload_a.clone(), false);
745 let _ = tx2_b.send(Recipients::One(pk1.clone()), payload_b.clone(), false);
746
747 let (from_a, bytes_a) = rx_a.recv().await.unwrap();
748 assert_eq!(from_a, pk2);
749 assert_eq!(bytes_a, payload_a);
750
751 let (from_b, bytes_b) = rx_b.recv().await.unwrap();
752 assert_eq!(from_b, pk2);
753 assert_eq!(bytes_b, payload_b);
754 });
755 }
756
757 #[test_traced]
758 fn test_mailbox_capacity_drops_when_full() {
759 let executor = deterministic::Runner::default();
762 executor.start(|context| async move {
763 let mut oracle = start_network(context.child("network"));
764
765 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
766 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
767 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
768
769 let (tx1, _) = handle1.register(99).await.unwrap();
771 let (tx2, _) = handle1.register(100).await.unwrap();
772 let (_, mut rx1) = handle2.register(99).await.unwrap();
773 let (_, mut rx2) = handle2.register(100).await.unwrap();
774
775 send_burst(&mut [tx1, tx2], CAPACITY * 2);
778
779 expect_n_messages(&mut rx1, CAPACITY).await;
781 expect_n_messages(&mut rx2, CAPACITY).await;
782 });
783 }
784
785 #[test]
786 fn test_drop_subchannel_receiver_deregisters_route() {
787 let executor = deterministic::Runner::default();
790 executor.start(|context| async move {
791 let mut oracle = start_network(context.child("network"));
792
793 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
794 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
795 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
796
797 let (tx1, _) = handle1.register(99).await.unwrap();
799 let (tx2, _) = handle1.register(100).await.unwrap();
800 let (_, rx1) = handle2.register(99).await.unwrap();
801 let (_, mut rx2) = handle2.register(100).await.unwrap();
802
803 drop(rx1);
805
806 send_burst(&mut [tx1, tx2], CAPACITY);
809
810 expect_n_messages(&mut rx2, CAPACITY).await;
812 });
813 }
814
815 #[test]
816 fn test_drop_messages_for_unregistered_subchannel() {
817 let executor = deterministic::Runner::default();
820 executor.start(|context| async move {
821 let mut oracle = start_network(context.child("network"));
822
823 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
824 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
825 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
826
827 let (tx1, _) = handle1.register(1).await.unwrap();
829 let (tx2, _) = handle1.register(2).await.unwrap();
830 let (_, mut rx2) = handle2.register(2).await.unwrap();
832
833 send_burst(&mut [tx1, tx2], CAPACITY);
837
838 expect_n_messages(&mut rx2, CAPACITY).await;
840 });
841 }
842
843 #[test]
844 fn test_backup_for_unregistered_subchannel() {
845 let executor = deterministic::Runner::default();
848 executor.start(|context| async move {
849 let mut oracle = start_network(context.child("network"));
850
851 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
852 let (pk2, mut handle2, mut backup2, _) =
853 create_peer_with_backup_and_global_sender(&context, &mut oracle, 1).await;
854 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
855
856 let (tx1, _) = handle1.register(1).await.unwrap();
858 let (tx2, _) = handle1.register(2).await.unwrap();
859 let (_, mut rx2) = handle2.register(2).await.unwrap();
861
862 send_burst(&mut [tx1, tx2], CAPACITY);
865
866 expect_n_messages_with_backup(&mut rx2, &mut backup2, CAPACITY, CAPACITY).await;
868 });
869 }
870
871 #[test]
872 fn test_backup_for_unregistered_subchannel_response() {
873 let executor = deterministic::Runner::default();
876 executor.start(|context| async move {
877 let mut oracle = start_network(context.child("network"));
878
879 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
880 let (pk2, _handle2, mut backup2, mut global_sender2) =
881 create_peer_with_backup_and_global_sender(&context, &mut oracle, 1).await;
882 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
883
884 let (mut tx1, mut rx1) = handle1.register(1).await.unwrap();
886 tx1.send(Recipients::One(pk2.clone()), b"REQUEST", false);
890
891 let (subchannel, (from, _)) = backup2.recv().await.unwrap();
893 assert_eq!(subchannel, 1);
894 assert_eq!(from, pk1);
895 let checked = global_sender2.check(Recipients::One(pk1.clone())).unwrap();
896 assert_eq!(checked.recipients(), vec![pk1]);
897 checked.with_subchannel(subchannel).send(b"TEST", true);
898
899 let (from, bytes) = rx1.recv().await.unwrap();
901 assert_eq!(from, pk2);
902 assert_eq!(bytes, b"TEST");
903 });
904 }
905
906 #[test]
907 fn test_message_dropped_for_closed_subchannel() {
908 let executor = deterministic::Runner::default();
913 executor.start(|context| async move {
914 let mut oracle = start_network(context.child("network"));
915
916 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
917 let (pk2, mut handle2) = create_peer(&context, &mut oracle, 1).await;
918 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
919
920 let (mut tx1, _) = handle1.register(1).await.unwrap();
922 let (mut tx2, _) = handle1.register(2).await.unwrap();
923 let (_, rx1) = handle2.register(1).await.unwrap();
924 let (_, mut rx2) = handle2.register(2).await.unwrap();
925
926 drop(rx1);
928
929 tx1.send(Recipients::One(pk2.clone()), b"closed", false);
932 tx2.send(Recipients::One(pk2.clone()), b"open", false);
933
934 expect_n_messages(&mut rx2, 1).await;
936 });
937 }
938
939 #[test]
940 fn test_dropped_backup_channel_doesnt_block() {
941 let executor = deterministic::Runner::default();
944 executor.start(|context| async move {
945 let mut oracle = start_network(context.child("network"));
946
947 let (pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
948 let (pk2, mut handle2, backup2, _) =
949 create_peer_with_backup_and_global_sender(&context, &mut oracle, 1).await;
950 link_bidirectional(&mut oracle, pk1.clone(), pk2.clone()).await;
951
952 drop(backup2);
954
955 let (tx1, _) = handle1.register(1).await.unwrap();
957 let (tx2, _) = handle1.register(2).await.unwrap();
958 let (_, mut rx2) = handle2.register(2).await.unwrap();
960
961 send_burst(&mut [tx1, tx2], CAPACITY);
965
966 expect_n_messages(&mut rx2, CAPACITY).await;
968 });
969 }
970
971 #[test]
972 fn test_duplicate_registration() {
973 let executor = deterministic::Runner::default();
975 executor.start(|context| async move {
976 let mut oracle = start_network(context.child("network"));
977
978 let (_pk1, mut handle1) = create_peer(&context, &mut oracle, 0).await;
979
980 let (_, _rx) = handle1.register(7).await.unwrap();
982
983 assert!(matches!(
985 handle1.register(7).await,
986 Err(Error::AlreadyRegistered(_))
987 ));
988 });
989 }
990
991 #[test]
992 fn test_register_after_deregister() {
993 let executor = deterministic::Runner::default();
995 executor.start(|context| async move {
996 let mut oracle = start_network(context.child("network"));
997
998 let (_, mut handle) = create_peer(&context, &mut oracle, 0).await;
999 let (_, rx) = handle.register(7).await.unwrap();
1000 drop(rx);
1001
1002 handle.register(7).await.unwrap();
1004 });
1005 }
1006}