consensus.rs 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375
  1. use async_std::{
  2. sync::{Arc, Mutex},
  3. task,
  4. };
  5. use std::time::Duration;
  6. use async_executor::Executor;
  7. use chrono::Utc;
  8. use futures::{select, FutureExt};
  9. use fxhash::FxHashMap;
  10. use log::{debug, error, warn};
  11. use rand::{rngs::OsRng, Rng, RngCore};
  12. use crate::{
  13. net,
  14. util::{
  15. self, gen_id,
  16. serial::{deserialize, serialize, Decodable, Encodable},
  17. },
  18. Error, Result,
  19. };
  20. use super::{
  21. p2p_send_loop,
  22. primitives::{
  23. BroadcastMsgRequest, Channel, Log, LogRequest, LogResponse, Logs, MapLength, NetMsg,
  24. NetMsgMethod, NodeId, NodeIdMsg, Role, Sender, VoteRequest, VoteResponse,
  25. },
  26. prune_map, DataStore, RaftSettings,
  27. };
  28. async fn send_node_id_loop(sender: async_channel::Sender<()>, timeout: i64) -> Result<()> {
  29. loop {
  30. util::sleep(timeout as u64).await;
  31. sender.send(()).await?;
  32. }
  33. }
  34. pub struct Raft<T> {
  35. id: NodeId,
  36. pub(super) role: Role,
  37. pub(super) current_leader: NodeId,
  38. pub(super) votes_received: Vec<NodeId>,
  39. pub(super) sent_length: MapLength,
  40. pub(super) acked_length: MapLength,
  41. pub(super) nodes: Arc<Mutex<FxHashMap<NodeId, i64>>>,
  42. pub(super) last_term: u64,
  43. p2p_sender: Sender,
  44. msgs_channel: Channel<T>,
  45. commits_channel: Channel<T>,
  46. datastore: DataStore<T>,
  47. seen_msgs: Arc<Mutex<FxHashMap<String, i64>>>,
  48. settings: RaftSettings,
  49. pending_msgs: Vec<T>,
  50. }
  51. impl<T: Decodable + Encodable + Clone> Raft<T> {
  52. pub fn new(
  53. settings: RaftSettings,
  54. seen_msgs: Arc<Mutex<FxHashMap<String, i64>>>,
  55. ) -> Result<Self> {
  56. if settings.datastore_path.to_str().is_none() {
  57. error!(target: "raft", "datastore path is incorrect");
  58. return Err(Error::ParseFailed("unable to parse pathbuf to str"))
  59. };
  60. let datastore = DataStore::new(settings.datastore_path.to_str().unwrap())?;
  61. // broadcasting channels
  62. let msgs_channel = async_channel::unbounded::<T>();
  63. let commits_channel = async_channel::unbounded::<T>();
  64. let p2p_sender = async_channel::unbounded::<NetMsg>();
  65. let id = match datastore.id.get_last()? {
  66. Some(_id) => _id,
  67. None => {
  68. let id = NodeId(gen_id(30));
  69. datastore.id.insert(&id)?;
  70. id
  71. }
  72. };
  73. let role = Role::Follower;
  74. Ok(Self {
  75. id,
  76. role,
  77. current_leader: NodeId("".into()),
  78. votes_received: vec![],
  79. sent_length: MapLength(FxHashMap::default()),
  80. acked_length: MapLength(FxHashMap::default()),
  81. nodes: Arc::new(Mutex::new(FxHashMap::default())),
  82. last_term: 0,
  83. p2p_sender,
  84. msgs_channel,
  85. commits_channel,
  86. datastore,
  87. seen_msgs,
  88. settings,
  89. pending_msgs: vec![],
  90. })
  91. }
  92. ///
  93. /// Run raft consensus and wait stop_signal channel to terminate
  94. ///
  95. pub async fn run(
  96. &mut self,
  97. p2p: net::P2pPtr,
  98. p2p_recv_channel: async_channel::Receiver<NetMsg>,
  99. executor: Arc<Executor<'_>>,
  100. stop_signal: async_channel::Receiver<()>,
  101. ) -> Result<()> {
  102. let p2p_send_task = executor.spawn(p2p_send_loop(self.p2p_sender.1.clone(), p2p.clone()));
  103. let prune_seen_messages_task = executor.spawn(prune_map::<String>(
  104. self.seen_msgs.clone(),
  105. self.settings.prun_messages_duration,
  106. ));
  107. let prune_nodes_id_task = executor
  108. .spawn(prune_map::<NodeId>(self.nodes.clone(), self.settings.prun_nodes_ids_duration));
  109. let (node_id_sx, node_id_rv) = async_channel::unbounded::<()>();
  110. let send_node_id_loop_task =
  111. executor.spawn(send_node_id_loop(node_id_sx, self.settings.node_id_timeout));
  112. let mut rng = rand::thread_rng();
  113. let broadcast_msg_rv = self.msgs_channel.1.clone();
  114. loop {
  115. let timeout = if self.role == Role::Leader {
  116. self.settings.heartbeat_timeout
  117. } else {
  118. rng.gen_range(0..self.settings.timeout) + self.settings.timeout
  119. };
  120. let timeout = Duration::from_millis(timeout);
  121. let mut result: Result<()>;
  122. select! {
  123. m = p2p_recv_channel.recv().fuse() => result = self.handle_method(m?).await,
  124. m = broadcast_msg_rv.recv().fuse() => result = self.broadcast_msg(&m?,None).await,
  125. _ = node_id_rv.recv().fuse() => result = self.send_node_id_msg().await,
  126. _ = task::sleep(timeout).fuse() => {
  127. result = if self.role == Role::Leader {
  128. self.send_heartbeat().await
  129. }else {
  130. self.send_vote_request().await
  131. };
  132. },
  133. _ = stop_signal.recv().fuse() => break,
  134. }
  135. // send pending messages
  136. if !self.pending_msgs.is_empty() && self.role != Role::Candidate {
  137. let pending_msgs = self.pending_msgs.clone();
  138. for m in &pending_msgs {
  139. result = self.broadcast_msg(m, None).await;
  140. }
  141. self.pending_msgs = vec![];
  142. }
  143. match result {
  144. Ok(_) => {}
  145. Err(e) => warn!(target: "raft", "warn: {}", e),
  146. }
  147. }
  148. warn!(target: "raft", "Raft Terminating...");
  149. p2p_send_task.cancel().await;
  150. prune_seen_messages_task.cancel().await;
  151. prune_nodes_id_task.cancel().await;
  152. send_node_id_loop_task.cancel().await;
  153. self.datastore.flush().await?;
  154. Ok(())
  155. }
  156. ///
  157. /// Return async receiver channel which can be used to receive T Messages
  158. /// from raft consensus
  159. ///
  160. pub fn receiver(&self) -> async_channel::Receiver<T> {
  161. self.commits_channel.1.clone()
  162. }
  163. ///
  164. /// Return async sender channel which can be used to broadcast T Messages
  165. /// to raft consensus
  166. ///
  167. pub fn sender(&self) -> async_channel::Sender<T> {
  168. self.msgs_channel.0.clone()
  169. }
  170. ///
  171. /// Return the raft node id
  172. ///
  173. pub fn id(&self) -> NodeId {
  174. self.id.clone()
  175. }
  176. async fn send_node_id_msg(&self) -> Result<()> {
  177. let node_id_msg = serialize(&NodeIdMsg { id: self.id.clone() });
  178. self.send(None, &node_id_msg, NetMsgMethod::NodeIdMsg, None).await?;
  179. Ok(())
  180. }
  181. async fn broadcast_msg(&mut self, msg: &T, msg_id: Option<u64>) -> Result<()> {
  182. match self.role {
  183. Role::Leader => {
  184. let msg = serialize(msg);
  185. let log = Log { msg, term: self.current_term()? };
  186. self.push_log(&log)?;
  187. self.acked_length.insert(&self.id, self.logs_len());
  188. }
  189. Role::Follower => {
  190. let b_msg = BroadcastMsgRequest(serialize(msg));
  191. self.send(
  192. Some(self.current_leader.clone()),
  193. &serialize(&b_msg),
  194. NetMsgMethod::BroadcastRequest,
  195. msg_id,
  196. )
  197. .await?;
  198. }
  199. Role::Candidate => {
  200. warn!("The role is Candidate, add the msg to pending_msgs");
  201. self.pending_msgs.push(msg.clone());
  202. }
  203. }
  204. debug!(target: "raft", "Role: {:?} Id: {:?}, broadcast a msg id: {:?} ", self.role, self.id, msg_id);
  205. Ok(())
  206. }
  207. async fn handle_method(&mut self, msg: NetMsg) -> Result<()> {
  208. match msg.method {
  209. NetMsgMethod::LogResponse => {
  210. let lr: LogResponse = deserialize(&msg.payload)?;
  211. self.receive_log_response(lr).await?;
  212. }
  213. NetMsgMethod::LogRequest => {
  214. let lr: LogRequest = deserialize(&msg.payload)?;
  215. self.receive_log_request(lr).await?;
  216. }
  217. NetMsgMethod::VoteResponse => {
  218. let vr: VoteResponse = deserialize(&msg.payload)?;
  219. self.receive_vote_response(vr).await?;
  220. }
  221. NetMsgMethod::VoteRequest => {
  222. let vr: VoteRequest = deserialize(&msg.payload)?;
  223. self.receive_vote_request(vr).await?;
  224. }
  225. NetMsgMethod::BroadcastRequest => {
  226. let vr: BroadcastMsgRequest = deserialize(&msg.payload)?;
  227. let d: T = deserialize(&vr.0)?;
  228. self.broadcast_msg(&d, Some(msg.id)).await?;
  229. }
  230. NetMsgMethod::NodeIdMsg => {
  231. let node_id_msg: NodeIdMsg = deserialize(&msg.payload)?;
  232. if node_id_msg.id != self.id {
  233. self.nodes.lock().await.insert(node_id_msg.id, Utc::now().timestamp());
  234. }
  235. }
  236. }
  237. debug!(target: "raft", "Role: {:?} Id: {:?}, receive a msg with id: {} recipient_id: {:?} method: {:?} ",
  238. self.role, self.id, msg.id, &msg.recipient_id, &msg.method);
  239. Ok(())
  240. }
  241. pub(super) async fn send(
  242. &self,
  243. recipient_id: Option<NodeId>,
  244. payload: &[u8],
  245. method: NetMsgMethod,
  246. msg_id: Option<u64>,
  247. ) -> Result<()> {
  248. let random_id = if msg_id.is_some() { msg_id.unwrap() } else { OsRng.next_u64() };
  249. debug!(target: "raft","Role: {:?} Id: {:?}, send a msg with id: {} recipient_id: {:?} method: {:?} ",
  250. self.role, self.id, random_id, &recipient_id, &method);
  251. let net_msg = NetMsg { id: random_id, recipient_id, payload: payload.to_vec(), method };
  252. self.seen_msgs.lock().await.insert(random_id.to_string(), Utc::now().timestamp());
  253. self.p2p_sender.0.send(net_msg).await?;
  254. Ok(())
  255. }
  256. pub(super) fn reset_last_term(&mut self) -> Result<()> {
  257. self.last_term = 0;
  258. if let Some(log) = self.last_log()? {
  259. self.last_term = log.term;
  260. }
  261. Ok(())
  262. }
  263. pub(super) fn set_current_term(&mut self, i: &u64) -> Result<()> {
  264. self.datastore.current_term.insert(i)
  265. }
  266. pub(super) fn set_voted_for(&mut self, i: &Option<NodeId>) -> Result<()> {
  267. self.datastore.voted_for.insert(i)
  268. }
  269. pub(super) async fn push_commit(&mut self, commit: &[u8]) -> Result<()> {
  270. let commit: T = deserialize(commit)?;
  271. self.commits_channel.0.send(commit.clone()).await?;
  272. self.datastore.commits.insert(&commit)
  273. }
  274. pub(super) fn push_log(&mut self, log: &Log) -> Result<()> {
  275. self.datastore.logs.insert(log)
  276. }
  277. pub(super) fn push_logs(&mut self, logs: &Logs) -> Result<()> {
  278. self.datastore.logs.wipe_insert_all(&logs.to_vec())
  279. }
  280. pub(super) fn current_term(&self) -> Result<u64> {
  281. Ok(self.datastore.current_term.get_last()?.unwrap_or(0))
  282. }
  283. pub(super) fn voted_for(&self) -> Result<Option<NodeId>> {
  284. Ok(self.datastore.voted_for.get_last()?.flatten())
  285. }
  286. pub(super) fn commits_len(&self) -> u64 {
  287. self.datastore.commits.len()
  288. }
  289. fn logs(&self) -> Result<Logs> {
  290. Ok(Logs(self.datastore.logs.get_all()?))
  291. }
  292. pub(super) fn logs_len(&self) -> u64 {
  293. self.datastore.logs.len()
  294. }
  295. fn last_log(&self) -> Result<Option<Log>> {
  296. self.datastore.logs.get_last()
  297. }
  298. pub(super) fn get_log(&self, index: u64) -> Result<Log> {
  299. self.datastore.logs.get(index)
  300. }
  301. pub(super) fn slice_logs_from(&self, index: u64) -> Result<Option<Logs>> {
  302. let logs = self.logs()?;
  303. Ok(logs.slice_from(index))
  304. }
  305. pub(super) fn slice_logs_to(&self, index: u64) -> Result<Logs> {
  306. let logs = self.logs()?;
  307. Ok(logs.slice_to(index))
  308. }
  309. }