ghassmo 4 лет назад
Родитель
Сommit
0ffc5f473f
1 измененных файлов с 101 добавлено и 76 удалено
  1. 101 76
      src/raft/raft.rs

+ 101 - 76
src/raft/raft.rs

@@ -27,6 +27,49 @@ const TIMEOUT_NODES: u64 = 300;
 pub type Broadcast<T> = (async_channel::Sender<T>, async_channel::Receiver<T>);
 type Sender = (async_channel::Sender<NetMsg>, async_channel::Receiver<NetMsg>);
 
+async fn load_node_ids_loop(
+    nodes: Arc<Mutex<HashMap<NodeId, SocketAddr>>>,
+    p2p: net::P2pPtr,
+    self_id: Option<NodeId>,
+) -> Result<()> {
+    if self_id.is_none() {
+        return Ok(())
+    }
+    loop {
+        debug!(target: "raft", "load node ids from p2p hosts ips");
+        task::sleep(Duration::from_millis(TIMEOUT_NODES * 10)).await;
+        let hosts = p2p.hosts().clone();
+        let nodes_ip = hosts.load_all().await.clone();
+        let mut nodes = nodes.lock().await;
+        for ip in nodes_ip.iter() {
+            nodes.insert(NodeId::from(*ip), *ip);
+        }
+        drop(nodes);
+    }
+}
+
+async fn p2p_recv_send_loop(
+    p2p_recv: async_channel::Receiver<NetMsg>,
+    p2p: net::P2pPtr,
+) -> Result<()> {
+    loop {
+        let msg: NetMsg = match p2p_recv.recv().await {
+            Ok(m) => m,
+            Err(e) => {
+                error!(target: "raft", "error occurred while receiving a msg: {}", e);
+                continue
+            }
+        };
+        match p2p.broadcast(msg).await {
+            Ok(_) => {}
+            Err(e) => {
+                error!(target: "raft", "error occurred during broadcasting a msg: {}", e);
+                continue
+            }
+        }
+    }
+}
+
 pub struct Raft<T> {
     // this will be derived from the ip
     // if the node doesn't have an id then will become a listener and doesn't have the right
@@ -117,6 +160,7 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
         executor: Arc<Executor<'_>>,
         stop_signal: async_channel::Receiver<()>,
     ) -> Result<()> {
+        // P2p setup
         let (p2p_snd, receive_queues) = async_channel::unbounded::<NetMsg>();
 
         let p2p = net::P2p::new(net_settings).await;
@@ -137,52 +181,19 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
             })
             .await;
 
-        // P2p performs seed session
+        // P2p start
         p2p.clone().start(executor.clone()).await?;
 
         let executor_cloned = executor.clone();
         let p2p_task = executor_cloned.spawn(p2p.clone().run(executor.clone()));
 
-        let p2p_cloned = p2p.clone();
         let p2p_recv = self.sender.1.clone();
-        let p2p_recv_task = executor.spawn(async move {
-            loop {
-                let msg: NetMsg = match p2p_recv.recv().await {
-                    Ok(m) => m,
-                    Err(e) => {
-                        error!(target: "raft", "error occurred while receiving a msg: {}", e);
-                        continue
-                    }
-                };
-                match p2p_cloned.broadcast(msg).await {
-                    Ok(_) => {}
-                    Err(e) => {
-                        error!(target: "raft", "error occurred during broadcasting a msg: {}", e);
-                        continue
-                    }
-                }
-            }
-        });
+        let p2p_recv_task = executor.spawn(p2p_recv_send_loop(p2p_recv.clone(), p2p.clone()));
 
-        let self_nodes = self.nodes.clone();
-        let p2p_cloned = p2p.clone();
-        let self_id = self.id.clone();
-        let load_ips_task = executor.spawn(async move {
-            if self_id.is_none() {
-                return
-            }
-            loop {
-                debug!(target: "raft", "load node ids from p2p hosts ips");
-                task::sleep(Duration::from_millis(TIMEOUT_NODES * 10)).await;
-                let hosts = p2p_cloned.hosts().clone();
-                let nodes_ip = hosts.load_all().await.clone();
-                let mut nodes = self_nodes.lock().await;
-                for ip in nodes_ip.iter() {
-                    nodes.insert(NodeId::from(*ip), *ip);
-                }
-            }
-        });
+        let load_ips_task =
+            executor.spawn(load_node_ids_loop(self.nodes.clone(), p2p.clone(), self.id.clone()));
 
+        // Sync listener node
         if self.id.is_none() {
             let last_term =
                 if !self.logs.0.is_empty() { self.logs.0.last().unwrap().term } else { 0 };
@@ -192,42 +203,7 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
             info!("send sync request");
             self.send(None, &serialize(&sync_request), NetMsgMethod::SyncRequest, None).await?;
 
-            loop {
-                select! {
-                    msg =  receive_queues.recv().fuse() => {
-                        let msg = msg?;
-                        if msg.method == NetMsgMethod::SyncResponse {
-                            info!("receive sync response");
-                            let sr: SyncResponse = deserialize(&msg.payload)?;
-                            if sr.wipe {
-                                self.set_commit_length(&0)?;
-                                self.push_logs(&sr.logs)?;
-                            } else {
-                                for log in sr.logs.0.iter() {
-                                    self.push_log(log)?;
-                                }
-                            }
-
-                            if !self.logs.is_empty() {
-                                self.set_current_term(&self.logs.0.last().unwrap().term.clone())?;
-                            }
-
-                            if self.commit_length > sr.commit_length {
-                                self.set_commit_length(&0)?;
-                            }
-
-                            for i in self.commit_length..sr.commit_length {
-                                self.push_commit(&self.logs.get(i)?.msg).await?;
-                            }
-                            self.set_commit_length(&sr.commit_length)?;
-
-                            self.current_leader = Some(sr.leader_id);
-
-                            break
-                        }},
-                        _ = stop_signal.recv().fuse() => break,
-                }
-            }
+            self.waiting_for_sync(receive_queues.clone(), stop_signal.clone()).await?;
         }
 
         let mut rng = rand::thread_rng();
@@ -338,7 +314,7 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
 
         debug!(
         target: "raft",
-        "{} {:?}  receive msg id: {}  recipient_id: {:?} method: {:?} ",
+        "Node has id: {} {:?}  receive msg id: {}  recipient_id: {:?} method: {:?} ",
         self.id.is_some(), self.role, msg.id, &msg.recipient_id.is_some(), &msg.method
         );
         Ok(())
@@ -393,6 +369,35 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
         Ok(())
     }
 
+    async fn receive_sync_response(&mut self, sr: &SyncResponse) -> Result<()> {
+        info!("receive sync response");
+        if sr.wipe {
+            self.set_commit_length(&0)?;
+            self.push_logs(&sr.logs)?;
+        } else {
+            for log in sr.logs.0.iter() {
+                self.push_log(log)?;
+            }
+        }
+
+        if !self.logs.is_empty() {
+            self.set_current_term(&self.logs.0.last().unwrap().term.clone())?;
+        }
+
+        if self.commit_length > sr.commit_length {
+            self.set_commit_length(&0)?;
+        }
+
+        for i in self.commit_length..sr.commit_length {
+            self.push_commit(&self.logs.get(i)?.msg).await?;
+        }
+        self.set_commit_length(&sr.commit_length)?;
+
+        self.current_leader = Some(sr.leader_id.clone());
+
+        Ok(())
+    }
+
     async fn send(
         &self,
         recipient_id: Option<NodeId>,
@@ -404,7 +409,7 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
 
         debug!(
         target: "raft",
-        "{} {:?}  send a msg id: {}  recipient_id: {:?} method: {:?} ",
+        "Node has id: {} {:?}  send a msg id: {}  recipient_id: {:?} method: {:?} ",
         self.id.is_some(), self.role, random_id, &recipient_id.is_some(), &method
         );
 
@@ -414,6 +419,26 @@ impl<T: Decodable + Encodable + Clone> Raft<T> {
         Ok(())
     }
 
+    async fn waiting_for_sync(
+        &mut self,
+        receive_queues: async_channel::Receiver<NetMsg>,
+        stop_signal: async_channel::Receiver<()>,
+    ) -> Result<()> {
+        loop {
+            select! {
+                msg =  receive_queues.recv().fuse() => {
+                    let msg = msg?;
+                    if msg.method == NetMsgMethod::SyncResponse {
+                        let sr: SyncResponse = deserialize(&msg.payload)?;
+                        self.receive_sync_response(&sr).await?;
+                        break
+                    }},
+                    _ = stop_signal.recv().fuse() => break,
+            }
+        }
+        Ok(())
+    }
+
     async fn send_heartbeat(&self) -> Result<()> {
         if self.role == Role::Leader {
             let nodes = self.nodes.lock().await;