1use std::{
10 collections::{BTreeMap, HashMap},
11 future::Future,
12 num::NonZeroUsize,
13 sync::{
14 Arc,
15 atomic::{AtomicBool, AtomicU64, Ordering},
16 },
17 time::Duration,
18};
19
20use async_broadcast::{InactiveReceiver, Sender, broadcast};
21use async_lock::RwLock;
22use async_trait::async_trait;
23use futures::{FutureExt, join, select};
24use hotshot_types::{
25 BoxSyncFuture, boxed_sync,
26 constants::{
27 COMBINED_NETWORK_CACHE_SIZE, COMBINED_NETWORK_DELAY_DURATION,
28 COMBINED_NETWORK_MIN_PRIMARY_FAILURES, COMBINED_NETWORK_PRIMARY_CHECK_INTERVAL,
29 },
30 data::{EpochNumber, ViewNumber},
31 epoch_membership::EpochMembershipCoordinator,
32 traits::{
33 network::{BroadcastDelay, ConnectedNetwork, Topic},
34 node_implementation::NodeType,
35 },
36};
37#[cfg(feature = "hotshot-testing")]
38use hotshot_types::{
39 PeerConnectInfo,
40 traits::network::{AsyncGenerator, NetworkReliability, TestableNetworkingImplementation},
41};
42use lru::LruCache;
43use parking_lot::RwLock as PlRwLock;
44use tokio::{spawn, sync::mpsc::error::TrySendError, time::sleep};
45use tracing::{debug, info, warn};
46
47use super::{NetworkError, push_cdn_network::PushCdnNetwork};
48use crate::traits::implementations::Libp2pNetwork;
49
50type DelayedTasksChannelsMap = Arc<RwLock<BTreeMap<u64, (Sender<()>, InactiveReceiver<()>)>>>;
52
53#[derive(Clone)]
56pub struct CombinedNetworks<TYPES: NodeType> {
57 networks: Arc<UnderlyingCombinedNetworks<TYPES>>,
59
60 message_deduplication_cache: Arc<PlRwLock<MessageDeduplicationCache>>,
63
64 primary_fail_counter: Arc<AtomicU64>,
66
67 primary_down: Arc<AtomicBool>,
69
70 delay_duration: Arc<RwLock<Duration>>,
72
73 delayed_tasks_channels: DelayedTasksChannelsMap,
75
76 no_delay_counter: Arc<AtomicU64>,
78}
79
80impl<TYPES: NodeType> CombinedNetworks<TYPES> {
81 #[must_use]
87 pub fn new(
88 primary_network: PushCdnNetwork<TYPES::SignatureKey>,
89 secondary_network: Libp2pNetwork<TYPES>,
90 delay_duration: Option<Duration>,
91 ) -> Self {
92 let networks = Arc::from(UnderlyingCombinedNetworks(
94 primary_network,
95 secondary_network,
96 ));
97
98 Self {
99 networks,
100 message_deduplication_cache: Arc::new(PlRwLock::new(MessageDeduplicationCache::new())),
101 primary_fail_counter: Arc::new(AtomicU64::new(0)),
102 primary_down: Arc::new(AtomicBool::new(false)),
103 delay_duration: Arc::new(RwLock::new(
104 delay_duration.unwrap_or(Duration::from_millis(COMBINED_NETWORK_DELAY_DURATION)),
105 )),
106 delayed_tasks_channels: Arc::default(),
107 no_delay_counter: Arc::new(AtomicU64::new(0)),
108 }
109 }
110
111 #[must_use]
113 pub fn primary(&self) -> &PushCdnNetwork<TYPES::SignatureKey> {
114 &self.networks.0
115 }
116
117 #[must_use]
119 pub fn secondary(&self) -> &Libp2pNetwork<TYPES> {
120 &self.networks.1
121 }
122
123 async fn send_both_networks(
125 &self,
126 _message: Vec<u8>,
127 primary_future: impl Future<Output = Result<(), NetworkError>> + Send + 'static,
128 secondary_future: impl Future<Output = Result<(), NetworkError>> + Send + 'static,
129 broadcast_delay: BroadcastDelay,
130 ) -> Result<(), NetworkError> {
131 let mut primary_failed = false;
133 if self.primary_down.load(Ordering::Relaxed) {
134 primary_failed = true;
136 } else if self.primary_fail_counter.load(Ordering::Relaxed)
137 > COMBINED_NETWORK_MIN_PRIMARY_FAILURES
138 {
139 info!(
142 "View progression is slower than normally, stop delaying messages on the secondary"
143 );
144 self.primary_down.store(true, Ordering::Relaxed);
145 primary_failed = true;
146 }
147
148 if let Err(e) = primary_future.await {
150 warn!("Error on primary network: {}", e);
152 self.primary_fail_counter.fetch_add(1, Ordering::Relaxed);
153 primary_failed = true;
154 };
155
156 if let (BroadcastDelay::View(view), false) = (broadcast_delay, primary_failed) {
157 let duration = *self.delay_duration.read().await;
159 let primary_down = Arc::clone(&self.primary_down);
160 let primary_fail_counter = Arc::clone(&self.primary_fail_counter);
161 let mut receiver = self
164 .delayed_tasks_channels
165 .write()
166 .await
167 .entry(view)
168 .or_insert_with(|| {
169 let (s, r) = broadcast(1);
170 (s, r.deactivate())
171 })
172 .1
173 .activate_cloned();
174 spawn(async move {
176 sleep(duration).await;
177 if receiver.try_recv().is_ok() {
178 debug!(
180 "Not sending on secondary after delay, task was canceled in view update"
181 );
182 match primary_fail_counter.load(Ordering::Relaxed) {
183 0u64 => {
184 primary_down.store(false, Ordering::Relaxed);
186 debug!("primary_fail_counter reached zero, primary_down set to false");
187 },
188 c => {
189 primary_fail_counter.store(c - 1, Ordering::Relaxed);
191 debug!("primary_fail_counter set to {:?}", c - 1);
192 },
193 }
194 return Ok(());
195 }
196 debug!(
199 "Sending on secondary after delay, message possibly has not reached recipient \
200 on primary"
201 );
202 primary_fail_counter.fetch_add(1, Ordering::Relaxed);
203 secondary_future.await
204 });
205 Ok(())
206 } else {
207 if self.primary_down.load(Ordering::Relaxed) {
209 match self.no_delay_counter.load(Ordering::Relaxed) {
213 c if c < COMBINED_NETWORK_PRIMARY_CHECK_INTERVAL => {
214 self.no_delay_counter.store(c + 1, Ordering::Relaxed);
216 },
217 _ => {
218 debug!(
220 "Sent on secondary without delay more than {} times,try delaying to \
221 check primary",
222 COMBINED_NETWORK_PRIMARY_CHECK_INTERVAL
223 );
224 self.no_delay_counter.store(0u64, Ordering::Relaxed);
226 self.primary_down.store(false, Ordering::Relaxed);
228 self.primary_fail_counter
230 .store(COMBINED_NETWORK_MIN_PRIMARY_FAILURES, Ordering::Relaxed);
231 },
232 }
233 }
234 tokio::time::timeout(Duration::from_secs(2), secondary_future)
236 .await
237 .map_err(|e| NetworkError::Timeout(e.to_string()))?
238 }
239 }
240}
241
242#[derive(Clone)]
246pub struct UnderlyingCombinedNetworks<TYPES: NodeType>(
247 pub PushCdnNetwork<TYPES::SignatureKey>,
248 pub Libp2pNetwork<TYPES>,
249);
250
251#[cfg(feature = "hotshot-testing")]
252impl<TYPES: NodeType> TestableNetworkingImplementation<TYPES> for CombinedNetworks<TYPES> {
253 fn generator(
254 expected_node_count: usize,
255 num_bootstrap: usize,
256 network_id: usize,
257 da_committee_size: usize,
258 reliability_config: Option<Box<dyn NetworkReliability>>,
259 secondary_network_delay: Duration,
260 connect_infos: &mut HashMap<TYPES::SignatureKey, PeerConnectInfo>,
261 ) -> AsyncGenerator<Arc<Self>> {
262 let generators = (
263 <PushCdnNetwork<TYPES::SignatureKey> as TestableNetworkingImplementation<TYPES>>::generator(
264 expected_node_count,
265 num_bootstrap,
266 network_id,
267 da_committee_size,
268 None,
269 Duration::default(),
270 connect_infos
271 ),
272 <Libp2pNetwork<TYPES> as TestableNetworkingImplementation<TYPES>>::generator(
273 expected_node_count,
274 num_bootstrap,
275 network_id,
276 da_committee_size,
277 reliability_config,
278 Duration::default(),
279 connect_infos
280 )
281 );
282 Box::pin(move |node_id| {
283 let gen0 = generators.0(node_id);
284 let gen1 = generators.1(node_id);
285
286 Box::pin(async move {
287 let cdn = gen0.await;
289 let cdn = Arc::<PushCdnNetwork<TYPES::SignatureKey>>::into_inner(cdn).unwrap();
290
291 let p2p = gen1.await;
293
294 let underlying_combined = UnderlyingCombinedNetworks(
296 cdn.clone(),
297 Arc::<Libp2pNetwork<TYPES>>::unwrap_or_clone(p2p),
298 );
299
300 let message_deduplication_cache =
302 Arc::new(PlRwLock::new(MessageDeduplicationCache::new()));
303
304 let combined_network = Self {
306 networks: Arc::new(underlying_combined),
307 primary_fail_counter: Arc::new(AtomicU64::new(0)),
308 primary_down: Arc::new(AtomicBool::new(false)),
309 message_deduplication_cache: Arc::clone(&message_deduplication_cache),
310 delay_duration: Arc::new(RwLock::new(secondary_network_delay)),
311 delayed_tasks_channels: Arc::default(),
312 no_delay_counter: Arc::new(AtomicU64::new(0)),
313 };
314
315 Arc::new(combined_network)
316 })
317 })
318 }
319
320 fn in_flight_message_count(&self) -> Option<usize> {
324 None
325 }
326}
327
328#[async_trait]
329impl<TYPES: NodeType> ConnectedNetwork<TYPES::SignatureKey> for CombinedNetworks<TYPES> {
330 fn pause(&self) {
331 self.networks.0.pause();
332 }
333
334 fn resume(&self) {
335 self.networks.0.resume();
336 }
337
338 async fn wait_for_ready(&self) {
339 join!(
340 self.primary().wait_for_ready(),
341 self.secondary().wait_for_ready()
342 );
343 }
344
345 fn shut_down<'a, 'b>(&'a self) -> BoxSyncFuture<'b, ()>
346 where
347 'a: 'b,
348 Self: 'b,
349 {
350 let closure = async move {
351 join!(self.primary().shut_down(), self.secondary().shut_down());
352 self.message_deduplication_cache.write().drop_contents();
356 self.delayed_tasks_channels.write().await.clear();
357 };
358 boxed_sync(closure)
359 }
360
361 async fn broadcast_message(
362 &self,
363 view: ViewNumber,
364 message: Vec<u8>,
365 topic: Topic,
366 broadcast_delay: BroadcastDelay,
367 ) -> Result<(), NetworkError> {
368 let primary = self.primary().clone();
369 let secondary = self.secondary().clone();
370 let primary_message = message.clone();
371 let secondary_message = message.clone();
372 self.send_both_networks(
373 message,
374 async move {
375 primary
376 .broadcast_message(view, primary_message, topic, BroadcastDelay::None)
377 .await
378 },
379 async move {
380 secondary
381 .broadcast_message(view, secondary_message, topic, BroadcastDelay::None)
382 .await
383 },
384 broadcast_delay,
385 )
386 .await
387 }
388
389 async fn da_broadcast_message(
390 &self,
391 view: ViewNumber,
392 message: Vec<u8>,
393 recipients: Vec<TYPES::SignatureKey>,
394 broadcast_delay: BroadcastDelay,
395 ) -> Result<(), NetworkError> {
396 let primary = self.primary().clone();
397 let secondary = self.secondary().clone();
398 let primary_message = message.clone();
399 let secondary_message = message.clone();
400 let primary_recipients = recipients.clone();
401 self.send_both_networks(
402 message,
403 async move {
404 primary
405 .da_broadcast_message(
406 view,
407 primary_message,
408 primary_recipients,
409 BroadcastDelay::None,
410 )
411 .await
412 },
413 async move {
414 secondary
415 .da_broadcast_message(view, secondary_message, recipients, BroadcastDelay::None)
416 .await
417 },
418 broadcast_delay,
419 )
420 .await
421 }
422
423 async fn direct_message(
424 &self,
425 view: ViewNumber,
426 message: Vec<u8>,
427 recipient: TYPES::SignatureKey,
428 ) -> Result<(), NetworkError> {
429 let primary = self.primary().clone();
430 let secondary = self.secondary().clone();
431 let primary_message = message.clone();
432 let secondary_message = message.clone();
433 let primary_recipient = recipient.clone();
434 self.send_both_networks(
435 message,
436 async move {
437 primary
438 .direct_message(view, primary_message, primary_recipient)
439 .await
440 },
441 async move {
442 secondary
443 .direct_message(view, secondary_message, recipient)
444 .await
445 },
446 BroadcastDelay::None,
447 )
448 .await
449 }
450
451 async fn vid_broadcast_message(
452 &self,
453 messages: HashMap<TYPES::SignatureKey, (ViewNumber, Vec<u8>)>,
454 ) -> Result<(), NetworkError> {
455 self.networks.0.vid_broadcast_message(messages).await
456 }
457
458 async fn recv_message(&self) -> Result<Vec<u8>, NetworkError> {
463 loop {
464 let mut primary_fut = self.primary().recv_message().fuse();
466 let mut secondary_fut = self.secondary().recv_message().fuse();
467
468 let (message, from_primary) = select! {
470 p = primary_fut => (p?, true),
471 s = secondary_fut => (s?, false),
472 };
473
474 if self
476 .message_deduplication_cache
477 .write()
478 .is_unique(&message, from_primary)
479 {
480 break Ok(message);
481 }
482 }
483 }
484
485 fn queue_node_lookup(
486 &self,
487 view_number: ViewNumber,
488 pk: TYPES::SignatureKey,
489 ) -> Result<(), TrySendError<Option<(ViewNumber, TYPES::SignatureKey)>>> {
490 self.primary().queue_node_lookup(view_number, pk.clone())?;
491 self.secondary().queue_node_lookup(view_number, pk)
492 }
493
494 async fn update_view<T>(
495 &self,
496 view: ViewNumber,
497 epoch: Option<EpochNumber>,
498 membership: EpochMembershipCoordinator<T>,
499 ) where
500 T: NodeType<SignatureKey = TYPES::SignatureKey>,
501 {
502 let delayed_tasks_channels = Arc::clone(&self.delayed_tasks_channels);
503 spawn(async move {
504 let mut map_lock = delayed_tasks_channels.write().await;
505 while let Some((first_view, _)) = map_lock.first_key_value() {
506 if *first_view < *view {
508 if let Some((_, (sender, _))) = map_lock.pop_first() {
509 let _ = sender.try_broadcast(());
510 } else {
511 break;
512 }
513 } else {
514 break;
515 }
516 }
517 });
518 self.networks
520 .1
521 .update_view::<T>(view, epoch, membership)
522 .await;
523 }
524
525 fn is_primary_down(&self) -> bool {
526 self.primary_down.load(Ordering::Relaxed)
527 }
528}
529
530struct MessageDeduplicationCache {
531 primary_message_cache: LruCache<blake3::Hash, ()>,
534
535 secondary_message_cache: LruCache<blake3::Hash, ()>,
538}
539
540impl MessageDeduplicationCache {
541 fn new() -> Self {
543 Self {
544 primary_message_cache: LruCache::new(
545 NonZeroUsize::new(COMBINED_NETWORK_CACHE_SIZE).unwrap(),
546 ),
547 secondary_message_cache: LruCache::new(
548 NonZeroUsize::new(COMBINED_NETWORK_CACHE_SIZE).unwrap(),
549 ),
550 }
551 }
552
553 fn drop_contents(&mut self) {
556 self.primary_message_cache = LruCache::new(NonZeroUsize::MIN);
557 self.secondary_message_cache = LruCache::new(NonZeroUsize::MIN);
558 }
559
560 fn is_unique(&mut self, message: &[u8], from_primary: bool) -> bool {
562 let message_hash = blake3::hash(message);
564
565 let (this_cache, other_cache) = if from_primary {
567 (
568 &mut self.primary_message_cache,
569 &mut self.secondary_message_cache,
570 )
571 } else {
572 (
573 &mut self.secondary_message_cache,
574 &mut self.primary_message_cache,
575 )
576 };
577
578 if other_cache.pop(&message_hash).is_some() {
581 false
583 } else {
584 this_cache.put(message_hash, ());
586 true
587 }
588 }
589}
590
591#[cfg(test)]
592mod tests {
593 use super::*;
594
595 #[test]
596 fn test_message_deduplication() {
597 let message = b"hello";
598
599 let mut cache = MessageDeduplicationCache::new();
601 assert!(cache.is_unique(message, true));
602 assert!(!cache.is_unique(message, false));
603
604 assert!(cache.is_unique(message, true));
607 assert!(!cache.is_unique(message, false));
608
609 assert!(cache.is_unique(message, false));
611 assert!(!cache.is_unique(message, true));
612 assert!(cache.is_unique(message, false));
613 assert!(!cache.is_unique(message, true));
614
615 assert!(cache.is_unique(message, true));
617 assert!(cache.is_unique(message, true));
618 assert!(cache.is_unique(message, true));
619 assert!(!cache.is_unique(message, false));
620 assert!(cache.is_unique(message, false));
621 }
622}