Skip to main content

mai_sdk_core/
network.rs

1use anyhow::{bail, Result};
2use async_channel::{Receiver, Sender};
3use base64::prelude::BASE64_STANDARD;
4use base64::Engine;
5use either::Either;
6use libp2p::pnet::PreSharedKey;
7use libp2p::Transport;
8use libp2p::{
9    autonat,
10    futures::StreamExt,
11    gossipsub, identify,
12    kad::{self, store::MemoryStore, QueryId},
13    mdns, noise, ping, pnet, relay,
14    swarm::{NetworkBehaviour, SwarmEvent},
15    tcp, yamux, Multiaddr,
16};
17use slog::{debug, error, info, warn, Logger};
18use std::{
19    collections::HashMap,
20    hash::{DefaultHasher, Hash, Hasher},
21    sync::Arc,
22};
23use tokio::{select, sync::RwLock};
24
25use crate::{
26    bridge::{EventBridge, PublishEvents},
27    handler::Startable,
28    storage::{GetEvent, SetEvent},
29};
30
31use serde::{Deserialize, Serialize};
32
33pub type PeerId = String;
34
35pub type Topic = String;
36
37pub struct PeerInquiryResponse {
38    peer_ids: Vec<PeerId>,
39}
40
41pub struct PeerInquiry {
42    response_tx: Sender<PeerInquiryResponse>,
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
46pub struct NetworkMessage {
47    /// Type of message
48    pub message_type: String,
49
50    /// Payload of the message
51    pub payload: Vec<u8>,
52}
53
54impl NetworkMessage {
55    pub fn new(message_type: String, payload: Vec<u8>) -> Self {
56        Self {
57            message_type,
58            payload,
59        }
60    }
61}
62
63pub trait Network {
64    /// Returns the local peer id
65    fn peer_id(&self) -> String;
66
67    /// Returns a list of peers connected to the network
68    fn peers(&self) -> impl std::future::Future<Output = Vec<PeerId>> + Send;
69}
70
71/// Contains information about a message received from the network
72#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
73pub struct HandlerEvent {
74    message: NetworkMessage,
75    peer_id: Option<PeerId>,
76    topic: Option<Topic>,
77}
78
79impl HandlerEvent {
80    pub fn new(message: NetworkMessage, peer_id: Option<PeerId>, topic: Option<Topic>) -> Self {
81        Self {
82            peer_id,
83            message,
84            topic,
85        }
86    }
87
88    pub fn peer_id(&self) -> Option<String> {
89        self.peer_id.clone()
90    }
91
92    pub fn message(&self) -> NetworkMessage {
93        self.message.clone()
94    }
95
96    pub fn topic(&self) -> Option<String> {
97        self.topic.as_ref().map(|topic| topic.to_string())
98    }
99}
100
101#[derive(NetworkBehaviour)]
102struct P2PNetworkBehaviour {
103    gossipsub: gossipsub::Behaviour,
104
105    ping: ping::Behaviour,
106
107    mdns: mdns::tokio::Behaviour,
108
109    autonat: autonat::Behaviour,
110
111    kad: kad::Behaviour<MemoryStore>,
112
113    relay: relay::Behaviour,
114
115    identify: identify::Behaviour,
116}
117
118#[derive(Debug, Clone)]
119pub struct P2PNetwork {
120    /// module logger
121    logger: Logger,
122
123    /// Listener addresses for the libp2p network
124    listen_addrs: Vec<Multiaddr>,
125
126    /// Bootstrap addresses for the libp2p network
127    bootstrap_addrs: Vec<Multiaddr>,
128
129    /// Ping interval
130    ping_interval: std::time::Duration,
131
132    /// Gossipsub heartbeat interval
133    gossipsub_heartbeat_interval: std::time::Duration,
134
135    /// Peer inquiry channel
136    peer_inquiry_tx: Sender<PeerInquiry>,
137
138    /// Peer inquiry channel
139    peer_inquiry_rx: Receiver<PeerInquiry>,
140
141    // Bridge
142    bridge: EventBridge,
143
144    /// Keypair for the broker network
145    keypair: libp2p::identity::Keypair,
146
147    /// KV get requests
148    kv_get_store: Arc<RwLock<HashMap<QueryId, GetEvent>>>,
149
150    /// Pre shared key for the network
151    psk: Option<pnet::PreSharedKey>,
152}
153
154pub struct P2PNetworkConfig {
155    pub listen_addrs: Vec<Multiaddr>,
156    pub bootstrap_addrs: Vec<Multiaddr>,
157    pub ping_interval: std::time::Duration,
158    pub gossipsub_heartbeat_interval: std::time::Duration,
159    pub logger: Logger,
160    pub bridge: EventBridge,
161    pub psk: Option<String>,
162}
163
164impl P2PNetwork {
165    pub fn new(cfg: P2PNetworkConfig) -> Self {
166        let keypair = libp2p::identity::Keypair::generate_ed25519();
167        let (peer_inquiry_tx, peer_inquiry_rx) = async_channel::unbounded();
168        let kv_get_store = Arc::new(RwLock::new(HashMap::new()));
169        let psk = if let Some(psk) = cfg.psk {
170            let psk_bytes: [u8; 32] = BASE64_STANDARD
171                .decode(psk.as_bytes())
172                .unwrap()
173                .try_into()
174                .unwrap();
175            let psk = PreSharedKey::new(psk_bytes);
176            Some(psk)
177        } else {
178            None
179        };
180        Self {
181            logger: cfg.logger,
182            kv_get_store,
183            listen_addrs: cfg.listen_addrs,
184            ping_interval: cfg.ping_interval,
185            gossipsub_heartbeat_interval: cfg.gossipsub_heartbeat_interval,
186            keypair,
187            peer_inquiry_rx,
188            peer_inquiry_tx,
189            bridge: cfg.bridge,
190            bootstrap_addrs: cfg.bootstrap_addrs,
191            psk,
192        }
193    }
194}
195
196async fn handle_set_event(kad: &mut kad::Behaviour<MemoryStore>, event: SetEvent) -> Result<()> {
197    let key = kad::RecordKey::new(&event.key.as_bytes());
198    kad.start_providing(key.clone())?;
199    let value = event.value.clone();
200    let record = kad::Record {
201        key,
202        value,
203        publisher: None,
204        expires: None,
205    };
206    if let Err(e) = kad.put_record(record, kad::Quorum::One) {
207        event.result.send(Err(e.into())).await?;
208    } else {
209        event.result.send(Ok(())).await?;
210    };
211    Ok(())
212}
213
214impl Startable for P2PNetwork {
215    async fn start(&self) -> Result<()> {
216        info!(self.logger, "building p2p network");
217
218        let mut swarm = libp2p::SwarmBuilder::with_existing_identity(self.keypair.clone())
219            .with_tokio()
220            .with_other_transport(|key| {
221                let noise_config = noise::Config::new(key).unwrap();
222                let yamux_config = yamux::Config::default();
223
224                let base_transport =
225                    tcp::tokio::Transport::new(tcp::Config::default().nodelay(true));
226                let maybe_encrypted =
227                    match self.psk {
228                        Some(psk) => Either::Left(base_transport.and_then(move |socket, _| {
229                            pnet::PnetConfig::new(psk).handshake(socket)
230                        })),
231                        None => Either::Right(base_transport),
232                    };
233                maybe_encrypted
234                    .upgrade(libp2p::core::transport::upgrade::Version::V1Lazy)
235                    .authenticate(noise_config)
236                    .multiplex(yamux_config)
237            })?
238            .with_behaviour(|key| {
239                let local_peer_id = key.public().to_peer_id();
240                let message_id_fn = |message: &gossipsub::Message| {
241                    let mut s = DefaultHasher::new();
242                    message.data.hash(&mut s);
243                    gossipsub::MessageId::from(s.finish().to_string())
244                };
245                let gossipsub_config = gossipsub::ConfigBuilder::default()
246                    .heartbeat_interval(self.gossipsub_heartbeat_interval)
247                    .validation_mode(gossipsub::ValidationMode::Strict)
248                    .message_id_fn(message_id_fn)
249                    .build()?;
250                let gossipsub = gossipsub::Behaviour::new(
251                    gossipsub::MessageAuthenticity::Signed(key.clone()),
252                    gossipsub_config,
253                )?;
254                let ping =
255                    ping::Behaviour::new(ping::Config::new().with_interval(self.ping_interval));
256                let mdns = mdns::tokio::Behaviour::new(mdns::Config::default(), local_peer_id)?;
257                let autonat = autonat::Behaviour::new(local_peer_id, Default::default());
258                let kad = kad::Behaviour::new(local_peer_id, MemoryStore::new(local_peer_id));
259                let relay = relay::Behaviour::new(local_peer_id, Default::default());
260                let identify = identify::Behaviour::new(identify::Config::new(
261                    "/TODO/0.0.1".to_string(),
262                    key.public(),
263                ));
264                Ok(P2PNetworkBehaviour {
265                    gossipsub,
266                    ping,
267                    mdns,
268                    autonat,
269                    kad,
270                    relay,
271                    identify,
272                })
273            })?
274            .build();
275
276        swarm.behaviour_mut().kad.set_mode(Some(kad::Mode::Server));
277
278        // Dial bootstrap nodes
279        for bootstrap_addr in self.bootstrap_addrs.iter() {
280            swarm.dial(bootstrap_addr.clone())?;
281        }
282
283        // Setup local listeners
284        for addr in self.listen_addrs.iter() {
285            swarm.listen_on(addr.clone())?;
286        }
287
288        // Subscribe to a common topic where broadcast messages will be sent
289        swarm
290            .behaviour_mut()
291            .gossipsub
292            .subscribe(&gossipsub::IdentTopic::new("broadcast".to_string()))?;
293
294        // Setup channel subscriptions
295        let network_rx: Receiver<NetworkMessage> = self.bridge.subscribe_to_network().await;
296        let kv_get_rx: Receiver<GetEvent> = self.bridge.subscribe_to_kv_get().await;
297        let kv_set_rx: Receiver<SetEvent> = self.bridge.subscribe_to_kv_set().await;
298
299        // Start event loop
300        info!(self.logger, "starting broker network");
301        loop {
302            select! {
303                // handle new kv get events
304                get_event = kv_get_rx.recv() => match get_event {
305                    Ok(event) =>{
306                        let key = kad::RecordKey::new(&event.key.as_bytes());
307                        let result = swarm.behaviour_mut().kad.get_record(key);
308                        self.kv_get_store.write().await.insert(result, event);
309                    },
310                    Err(e) => {
311                        error!(self.logger, "kv get rx channel closed: {e}");
312                        bail!("internal channel error")
313                    }
314                },
315                // handle new kv set events
316                set_event = kv_set_rx.recv() => match set_event {
317                    Ok(event) => {
318                        if let Err(e) = handle_set_event(&mut swarm.behaviour_mut().kad, event).await {
319                            error!(self.logger, "failed to handle set event: {e}");
320                        }
321                    },
322                    Err(e) => {
323                        error!(self.logger, "kv set rx channel closed: {e}");
324                        bail!("internal channel error")
325                    }
326                },
327                // handle new peer inquiries
328                peer_inquiry = self.peer_inquiry_rx.recv() => match peer_inquiry {
329                    Ok(inquiry) => {
330                        info!(self.logger, "received peer inquiry");
331                        let peer_ids = swarm.connected_peers().map(|peer_id| peer_id.to_string()).collect();
332                        let response = PeerInquiryResponse { peer_ids };
333                        if let Err(e) = inquiry.response_tx.send(response).await {
334                            error!(self.logger, "failed to send peer inquiry response: {e}");
335                        }
336                    },
337                    Err(e) => {
338                        error!(self.logger, "peer inquiry rx channel closed: {e}");
339                        bail!("internal channel error")
340                    }
341                },
342                // handle new events from the publish rx channel
343                publish_event = network_rx.recv() => match publish_event {
344                    Ok(event) => {
345                        let message = bincode::serialize(&event)?;
346
347                        // Publish the network to the gossipsub network
348                        if let Err(e) = swarm.behaviour_mut().gossipsub.publish(
349                            gossipsub::IdentTopic::new("broadcast"),
350                            message,
351                        ) {
352                            // If the error is due to insufficient peers, we can safely ignore since we propagate the message to the local handler
353                            match e {
354                                gossipsub::PublishError::InsufficientPeers => {
355                                    info!(self.logger, "no peers to publish message to, message will only propagate locally");
356                                },
357                                e => {
358                                    warn!(self.logger, "failed to publish message: {e}");
359                                },
360                            }
361                        };
362
363                        // Publish the network to the local handler
364                        // NOTE: we do this to allow the local system to handle any jobs it has capacity for
365                        let local_handler_event = HandlerEvent {
366                            peer_id: Some(self.peer_id()),
367                            topic: None,
368                            message: event,
369                        };
370                        if let Err(e) = self.bridge.publish(crate::bridge::PublishEvents::HandlerEvent(local_handler_event)).await {
371                            error!(self.logger, "failed to send message to handler: {e}");
372                        } else {
373                            info!(self.logger, "notified local");
374                        };
375                    },
376                    Err(e) => {
377                        error!(self.logger, "publish rx channel closed: {e}");
378                        bail!("internal channel error")
379                    }
380                },
381                // handle new events from the swarm network
382                event = swarm.select_next_some() => match event {
383                    SwarmEvent::Behaviour(P2PNetworkBehaviourEvent::Mdns(mdns::Event::Discovered(list))) => {
384                        for (peer_id, multiaddr) in list {
385                            info!(self.logger, "mDNS discovered a new peer: {peer_id}");
386                            swarm.behaviour_mut().gossipsub.add_explicit_peer(&peer_id);
387                            swarm.behaviour_mut().kad.add_address(&peer_id, multiaddr);
388                        }
389                    },
390                    SwarmEvent::Behaviour(P2PNetworkBehaviourEvent::Identify(identify::Event::Received {
391                        peer_id,
392                        info: identify::Info { observed_addr, .. },
393                    })) => {
394                        info!(self.logger, "received identify info from {peer_id}");
395                        swarm.add_external_address(observed_addr.clone());
396                    },
397                    SwarmEvent::Behaviour(P2PNetworkBehaviourEvent::Mdns(mdns::Event::Expired(list))) => {
398                        for (peer_id, _multiaddr) in list {
399                            info!(self.logger, "mDNS discover peer has expired: {peer_id}");
400                            swarm.behaviour_mut().gossipsub.remove_explicit_peer(&peer_id);
401                        }
402                    },
403                    SwarmEvent::Behaviour(P2PNetworkBehaviourEvent::Gossipsub(gossipsub::Event::Message {
404                        propagation_source: peer_id,
405                        message_id: id,
406                        message,
407                    })) => {
408                        let topic = message.topic.clone();
409                        info!(self.logger, "received message {id} from {peer_id} on topic {topic}");
410                        let message = bincode::deserialize(&message.data)?;
411                        if let Err(e) = self.bridge.publish(PublishEvents::HandlerEvent(message)).await {
412                            error!(self.logger, "failed to send message to handler: {e}");
413                        };
414                    },
415                    SwarmEvent::NewExternalAddrOfPeer { peer_id, .. }  => {
416                        info!(
417                            self.logger,
418                            "discovered new address for peer";
419                            "peer_id" => peer_id.to_string()
420                        );
421                    },
422                    SwarmEvent::NewListenAddr { address, .. } => {
423                        info!(self.logger, "local node is listening on {address}");
424                    }
425                    SwarmEvent::Behaviour(P2PNetworkBehaviourEvent::Kad(kad::Event::OutboundQueryProgressed {id: query_id, result, ..})) => {
426                        match result {
427                            kad::QueryResult::GetProviders(Ok(kad::GetProvidersOk::FoundProviders { key, providers, .. })) => {
428                                for peer in providers {
429                                    info!(
430                                        self.logger,
431                                        "Peer {peer:?} provides key {:?}",
432                                        std::str::from_utf8(key.as_ref()).unwrap()
433                                    );
434                                }
435                            }
436                            kad::QueryResult::GetProviders(Err(err)) => {
437                                error!(self.logger, "Failed to get providers: {err:?}");
438                            }
439                            kad::QueryResult::GetRecord(Ok(
440                                kad::GetRecordOk::FoundRecord(kad::PeerRecord {
441                                    record: kad::Record { key, value, .. },
442                                    ..
443                                })
444                            )) => {
445                                info!(
446                                    self.logger,
447                                    "got record";
448                                    "key" => std::str::from_utf8(key.as_ref()).unwrap(),
449                                );
450                                let event = self.kv_get_store.write().await.remove(&query_id);
451                                if let Some(event) = event {
452                                    event.result.send(Ok(Some(value))).await?;
453                                }
454                            }
455                            kad::QueryResult::GetRecord(Ok(_)) => {}
456                            kad::QueryResult::GetRecord(Err(err)) => {
457                                error!(self.logger, "Failed to get record: {err:?}");
458                                let event = self.kv_get_store.write().await.remove(&query_id);
459                                if let Some(event) = event {
460                                    event.result.send(Err(err.into())).await?;
461                                }
462                            }
463                            kad::QueryResult::PutRecord(Ok(kad::PutRecordOk { key })) => {
464                                info!(
465                                    self.logger,
466                                    "Successfully put record {:?}",
467                                    std::str::from_utf8(key.as_ref()).unwrap()
468                                );
469                            }
470                            kad::QueryResult::PutRecord(Err(err)) => {
471                                error!(self.logger, "Failed to put record: {err:?}");
472                            }
473                            kad::QueryResult::StartProviding(Ok(kad::AddProviderOk { key })) => {
474                                info!(
475                                    self.logger,
476                                    "Successfully put provider record {:?}",
477                                    std::str::from_utf8(key.as_ref()).unwrap()
478                                );
479                            }
480                            kad::QueryResult::StartProviding(Err(err)) => {
481                                eprintln!("Failed to put provider record: {err:?}");
482                            }
483                            _ => {}
484                        }
485                    }
486                    _ => {
487                        debug!(self.logger, "unhandled event {:?}", event);
488                    }
489                }
490            }
491        }
492    }
493}
494
495impl Network for P2PNetwork {
496    fn peer_id(&self) -> String {
497        self.keypair.public().to_peer_id().to_string()
498    }
499
500    async fn peers(&self) -> Vec<PeerId> {
501        let (tx, rx) = async_channel::bounded(1);
502        let inquiry = PeerInquiry { response_tx: tx };
503        if let Err(e) = self.peer_inquiry_tx.send(inquiry).await {
504            error!(self.logger, "failed to send peer inquiry: {e}");
505            return vec![];
506        }
507        match rx.recv().await {
508            Ok(response) => response.peer_ids,
509            Err(e) => {
510                error!(self.logger, "failed to receive peer inquiry response: {e}");
511                vec![]
512            }
513        }
514    }
515}