1use std::{collections::HashSet, fmt::Debug, sync::Arc, time::Duration};
8
9use bimap::BiMap;
10use hotshot_types::traits::{
11 network::NetworkError, node_implementation::NodeType, signature_key::SignatureKey,
12};
13use libp2p::{Multiaddr, request_response::ResponseChannel};
14use libp2p_identity::PeerId;
15use parking_lot::Mutex;
16use tokio::{
17 sync::mpsc::{Receiver, UnboundedReceiver, UnboundedSender},
18 time::{sleep, timeout},
19};
20use tracing::{debug, info, instrument};
21
22use crate::network::{
23 ClientRequest, NetworkEvent, NetworkNode, NetworkNodeConfig, SwarmTaskHandle,
24 behaviours::dht::{
25 record::{Namespace, RecordKey, RecordValue},
26 store::persistent::DhtPersistentStorage,
27 },
28 gen_multiaddr, log_summary,
29};
30
31#[derive(Debug, Clone)]
35pub struct NetworkNodeHandle<T: NodeType> {
36 network_config: NetworkNodeConfig,
38
39 send_network: UnboundedSender<ClientRequest>,
41
42 consensus_key_to_pid_map: Arc<Mutex<BiMap<T::SignatureKey, PeerId>>>,
44
45 listen_addr: Multiaddr,
47
48 peer_id: PeerId,
50
51 id: usize,
53
54 swarm_task: Arc<Mutex<Option<SwarmTaskHandle>>>,
56}
57
58#[derive(Debug)]
60pub struct NetworkNodeReceiver {
61 receiver: UnboundedReceiver<NetworkEvent>,
63
64 recv_kill: Option<Receiver<()>>,
66}
67
68impl NetworkNodeReceiver {
69 pub async fn recv(&mut self) -> Result<NetworkEvent, NetworkError> {
73 self.receiver
74 .recv()
75 .await
76 .ok_or(NetworkError::ChannelReceiveError(
77 "Receiver channel closed".to_string(),
78 ))
79 }
80 pub fn set_kill_switch(&mut self, kill_switch: Receiver<()>) {
82 self.recv_kill = Some(kill_switch);
83 }
84
85 pub fn take_kill_switch(&mut self) -> Option<Receiver<()>> {
87 self.recv_kill.take()
88 }
89}
90
91pub async fn spawn_network_node<T: NodeType, D: DhtPersistentStorage>(
95 config: NetworkNodeConfig,
96 dht_persistent_storage: D,
97 consensus_key_to_pid_map: Arc<Mutex<BiMap<T::SignatureKey, PeerId>>>,
98 id: usize,
99) -> Result<(NetworkNodeReceiver, NetworkNodeHandle<T>), NetworkError> {
100 let mut network: NetworkNode<T, _> = NetworkNode::new(
101 config.clone(),
102 dht_persistent_storage,
103 Arc::clone(&consensus_key_to_pid_map),
104 )
105 .await
106 .map_err(|e| NetworkError::ConfigError(format!("failed to create network node: {e}")))?;
107 let listen_addr = config
109 .bind_address
110 .clone()
111 .unwrap_or_else(|| gen_multiaddr(0));
112 let peer_id = network.peer_id();
113 let listen_addr = network.start_listen(listen_addr).await.map_err(|e| {
114 NetworkError::ListenError(format!("failed to start listening on Libp2p: {e}"))
115 })?;
116 let (send_chan, recv_chan, swarm_task) = network.spawn_listeners().map_err(|err| {
119 NetworkError::ListenError(format!("failed to spawn listeners for Libp2p: {err}"))
120 })?;
121 log_summary::spawn_summary_task();
122 let receiver = NetworkNodeReceiver {
123 receiver: recv_chan,
124 recv_kill: None,
125 };
126
127 let handle = NetworkNodeHandle::<T> {
128 network_config: config,
129 send_network: send_chan,
130 consensus_key_to_pid_map,
131 listen_addr,
132 peer_id,
133 id,
134 swarm_task: Arc::new(Mutex::new(Some(swarm_task))),
135 };
136 Ok((receiver, handle))
137}
138
139impl<T: NodeType> NetworkNodeHandle<T> {
140 #[instrument]
144 pub async fn shutdown(&self) -> Result<(), NetworkError> {
145 self.send_request(ClientRequest::Shutdown)?;
146
147 let task = self.swarm_task.lock().take();
150 if let Some(task) = task {
151 match timeout(Duration::from_secs(5), task).await {
152 Ok(Ok(_)) => {},
153 Ok(Err(err)) => debug!(%err, "swarm task ended with error during shutdown"),
154 Err(_) => {
155 debug!("timed out waiting for swarm task to finish during shutdown");
156 },
157 }
158 }
159 Ok(())
160 }
161 pub fn begin_bootstrap(&self) -> Result<(), NetworkError> {
166 let req = ClientRequest::BeginBootstrap;
167 self.send_request(req)
168 }
169
170 #[must_use]
172 pub fn listen_addr(&self) -> Multiaddr {
173 self.listen_addr.clone()
174 }
175
176 pub async fn print_routing_table(&self) -> Result<(), NetworkError> {
181 let (s, r) = futures::channel::oneshot::channel();
182 let req = ClientRequest::GetRoutingTable(s);
183 self.send_request(req)?;
184 r.await
185 .map_err(|e| NetworkError::ChannelReceiveError(e.to_string()))
186 }
187 pub async fn wait_to_connect(
192 &self,
193 num_required_peers: usize,
194 node_id: usize,
195 ) -> Result<(), NetworkError> {
196 loop {
198 let num_connected = self.num_connected().await?;
200 if num_connected >= num_required_peers {
201 break;
202 }
203
204 info!(
206 "Node {} connected to {}/{} peers",
207 node_id, num_connected, num_required_peers
208 );
209
210 sleep(Duration::from_secs(1)).await;
212 }
213
214 Ok(())
215 }
216
217 pub async fn lookup_pid(&self, peer_id: PeerId) -> Result<(), NetworkError> {
222 let (s, r) = futures::channel::oneshot::channel();
223 let req = ClientRequest::LookupPeer(peer_id, s);
224 self.send_request(req)?;
225 r.await
226 .map_err(|err| NetworkError::ChannelReceiveError(err.to_string()))
227 }
228
229 pub async fn lookup_node(
234 &self,
235 consensus_key: &T::SignatureKey,
236 dht_timeout: Duration,
237 ) -> Result<PeerId, NetworkError> {
238 if let Some(pid) = self
240 .consensus_key_to_pid_map
241 .lock()
242 .get_by_left(consensus_key)
243 {
244 return Ok(*pid);
245 }
246
247 let key = RecordKey::new(Namespace::Lookup, consensus_key.to_bytes());
249
250 let pid = self.get_record_timeout(key, dht_timeout).await?;
252
253 PeerId::from_bytes(&pid).map_err(|err| NetworkError::FailedToDeserialize(err.to_string()))
254 }
255
256 pub async fn put_record(
260 &self,
261 key: RecordKey,
262 value: RecordValue<T::SignatureKey>,
263 ) -> Result<(), NetworkError> {
264 let key = key.to_bytes();
266
267 let value = bincode::serialize(&value)
269 .map_err(|e| NetworkError::FailedToSerialize(e.to_string()))?;
270
271 let (s, r) = futures::channel::oneshot::channel();
272 let req = ClientRequest::PutDHT {
273 key: key.clone(),
274 value,
275 notify: s,
276 };
277
278 self.send_request(req)?;
279
280 r.await.map_err(|_| NetworkError::RequestCancelled)
281 }
282
283 pub async fn get_record(
289 &self,
290 key: RecordKey,
291 retry_count: u8,
292 ) -> Result<Vec<u8>, NetworkError> {
293 let serialized_key = key.to_bytes();
295
296 let (s, r) = futures::channel::oneshot::channel();
297 let req = ClientRequest::GetDHT {
298 key: serialized_key.clone(),
299 notify: vec![s],
300 retry_count,
301 };
302 self.send_request(req)?;
303
304 let result = r.await.map_err(|_| NetworkError::RequestCancelled)?;
306
307 let record: RecordValue<T::SignatureKey> = bincode::deserialize(&result)
309 .map_err(|e| NetworkError::FailedToDeserialize(e.to_string()))?;
310
311 Ok(record.value().to_vec())
312 }
313
314 pub async fn get_record_timeout(
320 &self,
321 key: RecordKey,
322 timeout_duration: Duration,
323 ) -> Result<Vec<u8>, NetworkError> {
324 timeout(timeout_duration, self.get_record(key, 3))
325 .await
326 .map_err(|err| NetworkError::Timeout(err.to_string()))?
327 }
328
329 pub async fn put_record_timeout(
335 &self,
336 key: RecordKey,
337 value: RecordValue<T::SignatureKey>,
338 timeout_duration: Duration,
339 ) -> Result<(), NetworkError> {
340 timeout(timeout_duration, self.put_record(key, value))
341 .await
342 .map_err(|err| NetworkError::Timeout(err.to_string()))?
343 }
344
345 pub async fn subscribe(&self, topic: String) -> Result<(), NetworkError> {
349 let (s, r) = futures::channel::oneshot::channel();
350 let req = ClientRequest::Subscribe(topic, Some(s));
351 self.send_request(req)?;
352 r.await
353 .map_err(|err| NetworkError::ChannelReceiveError(err.to_string()))
354 }
355
356 pub async fn unsubscribe(&self, topic: String) -> Result<(), NetworkError> {
360 let (s, r) = futures::channel::oneshot::channel();
361 let req = ClientRequest::Unsubscribe(topic, Some(s));
362 self.send_request(req)?;
363 r.await
364 .map_err(|err| NetworkError::ChannelReceiveError(err.to_string()))
365 }
366
367 pub fn ignore_peers(&self, peers: Vec<PeerId>) -> Result<(), NetworkError> {
372 let req = ClientRequest::IgnorePeers(peers);
373 self.send_request(req)
374 }
375
376 pub fn direct_request(&self, pid: PeerId, msg: &[u8]) -> Result<(), NetworkError> {
381 self.direct_request_no_serialize(pid, msg.to_vec())
382 }
383
384 pub fn direct_request_no_serialize(
389 &self,
390 pid: PeerId,
391 contents: Vec<u8>,
392 ) -> Result<(), NetworkError> {
393 let req = ClientRequest::DirectRequest {
394 pid,
395 contents,
396 retry_count: 1,
397 };
398 self.send_request(req)
399 }
400
401 pub fn direct_response(
406 &self,
407 chan: ResponseChannel<Vec<u8>>,
408 msg: &[u8],
409 ) -> Result<(), NetworkError> {
410 let req = ClientRequest::DirectResponse(chan, msg.to_vec());
411 self.send_request(req)
412 }
413
414 pub fn prune_peer(&self, pid: PeerId) -> Result<(), NetworkError> {
422 let req = ClientRequest::Prune(pid);
423 self.send_request(req)
424 }
425
426 pub fn gossip(&self, topic: String, msg: &[u8]) -> Result<(), NetworkError> {
431 self.gossip_no_serialize(topic, msg.to_vec())
432 }
433
434 pub fn gossip_no_serialize(&self, topic: String, msg: Vec<u8>) -> Result<(), NetworkError> {
439 let req = ClientRequest::GossipMsg(topic, msg);
440 self.send_request(req)
441 }
442
443 pub fn add_known_peers(
447 &self,
448 known_peers: Vec<(PeerId, Multiaddr)>,
449 ) -> Result<(), NetworkError> {
450 debug!("Adding {} known peers", known_peers.len());
451 let req = ClientRequest::AddKnownPeers(known_peers);
452 self.send_request(req)
453 }
454
455 fn send_request(&self, req: ClientRequest) -> Result<(), NetworkError> {
460 self.send_network
461 .send(req)
462 .map_err(|err| NetworkError::ChannelSendError(err.to_string()))
463 }
464
465 pub async fn num_connected(&self) -> Result<usize, NetworkError> {
473 let (s, r) = futures::channel::oneshot::channel();
474 let req = ClientRequest::GetConnectedPeerNum(s);
475 self.send_request(req)?;
476 Ok(r.await.unwrap())
477 }
478
479 pub async fn connected_pids(&self) -> Result<HashSet<PeerId>, NetworkError> {
487 let (s, r) = futures::channel::oneshot::channel();
488 let req = ClientRequest::GetConnectedPeers(s);
489 self.send_request(req)?;
490 Ok(r.await.unwrap())
491 }
492
493 pub async fn kad_routing_peers(&self) -> Result<HashSet<PeerId>, NetworkError> {
497 let (s, r) = futures::channel::oneshot::channel();
498 let req = ClientRequest::GetKadRoutingPeers(s);
499 self.send_request(req)?;
500 Ok(r.await.unwrap())
501 }
502
503 #[must_use]
505 pub fn id(&self) -> usize {
506 self.id
507 }
508
509 #[must_use]
511 pub fn peer_id(&self) -> PeerId {
512 self.peer_id
513 }
514
515 #[must_use]
517 pub fn config(&self) -> &NetworkNodeConfig {
518 &self.network_config
519 }
520}