raft.rs 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607
  1. use async_std::{
  2. sync::{Arc, Mutex},
  3. task,
  4. };
  5. use std::{cmp::min, collections::HashMap, net::SocketAddr, path::PathBuf, time::Duration};
  6. use async_executor::Executor;
  7. use futures::{select, FutureExt};
  8. use log::{debug, error, info, warn};
  9. use rand::{rngs::OsRng, Rng, RngCore};
  10. use crate::{
  11. net,
  12. util::serial::{deserialize, serialize, Decodable, Encodable},
  13. Error, Result,
  14. };
  15. use super::{
  16. BroadcastMsgRequest, DataStore, Log, LogRequest, LogResponse, Logs, MapLength, NetMsg,
  17. NetMsgMethod, NodeId, ProtocolRaft, Role, VoteRequest, VoteResponse,
  18. };
  19. const HEARTBEATTIMEOUT: u64 = 100;
  20. const TIMEOUT: u64 = 300;
  21. const TIMEOUT_NODES: u64 = 300;
  22. pub type Broadcast<T> = (async_channel::Sender<T>, async_channel::Receiver<T>);
  23. type Sender = (async_channel::Sender<NetMsg>, async_channel::Receiver<NetMsg>);
  24. pub struct Raft<T> {
  25. // this will be derived from the ip
  26. // if the node doesn't have an id then will become a listener and doesn't have the right
  27. // to request/response votes or response a confirmation for log
  28. id: Option<NodeId>,
  29. // these five vars should be on local storage
  30. current_term: u64,
  31. voted_for: Option<NodeId>,
  32. logs: Logs,
  33. commit_length: u64,
  34. role: Role,
  35. current_leader: Option<NodeId>,
  36. votes_received: Vec<NodeId>,
  37. sent_length: MapLength,
  38. acked_length: MapLength,
  39. nodes: Arc<Mutex<HashMap<NodeId, SocketAddr>>>,
  40. last_term: u64,
  41. sender: Sender,
  42. broadcast_msg: Broadcast<T>,
  43. broadcast_commits: Broadcast<T>,
  44. datastore: DataStore<T>,
  45. }
  46. impl<T: Decodable + Encodable + Clone> Raft<T> {
  47. pub fn new(addr: Option<SocketAddr>, db_path: PathBuf) -> Result<Self> {
  48. if db_path.to_str().is_none() {
  49. error!(target: "raft", "datastore path is incorrect");
  50. return Err(Error::ParseFailed("unable to parse pathbuf to str"))
  51. };
  52. let db_path_str = db_path.to_str().unwrap();
  53. let mut current_term = 0;
  54. let mut voted_for = None;
  55. let mut logs = Logs(vec![]);
  56. let mut commit_length = 0;
  57. let datastore = if db_path.exists() {
  58. let datastore = DataStore::new(db_path_str)?;
  59. current_term = datastore.current_term.get_last()?.unwrap_or(0);
  60. voted_for = datastore.voted_for.get_last()?.flatten();
  61. logs = Logs(datastore.logs.get_all()?);
  62. commit_length = datastore.commits_length.get_last()?.unwrap_or(0);
  63. datastore
  64. } else {
  65. DataStore::new(db_path_str)?
  66. };
  67. // broadcasting channels
  68. let broadcast_msg = async_channel::unbounded::<T>();
  69. let broadcast_commits = async_channel::unbounded::<T>();
  70. let sender = async_channel::unbounded::<NetMsg>();
  71. Ok(Self {
  72. id: addr.map(NodeId::from),
  73. current_term,
  74. voted_for,
  75. logs,
  76. commit_length,
  77. role: Role::Follower,
  78. current_leader: None,
  79. votes_received: vec![],
  80. sent_length: MapLength(HashMap::new()),
  81. acked_length: MapLength(HashMap::new()),
  82. nodes: Arc::new(Mutex::new(HashMap::new())),
  83. last_term: 0,
  84. sender,
  85. broadcast_msg,
  86. broadcast_commits,
  87. datastore,
  88. })
  89. }
  90. pub async fn start(
  91. &mut self,
  92. net_settings: net::Settings,
  93. executor: Arc<Executor<'_>>,
  94. stop_signal: async_channel::Receiver<()>,
  95. ) -> Result<()> {
  96. let (p2p_snd, receive_queues) = async_channel::unbounded::<NetMsg>();
  97. let p2p = net::P2p::new(net_settings).await;
  98. let p2p = p2p.clone();
  99. let registry = p2p.protocol_registry();
  100. let self_id = self.id.clone();
  101. registry
  102. .register(net::SESSION_ALL, move |channel, p2p| {
  103. let self_id = self_id.clone();
  104. let sender = p2p_snd.clone();
  105. async move { ProtocolRaft::init(self_id, channel, sender, p2p).await }
  106. })
  107. .await;
  108. // P2p performs seed session
  109. p2p.clone().start(executor.clone()).await?;
  110. let executor_cloned = executor.clone();
  111. let p2p_task = executor_cloned.spawn(p2p.clone().run(executor.clone()));
  112. let p2p_cloned = p2p.clone();
  113. let p2p_recv = self.sender.1.clone();
  114. let p2p_recv_task = executor.spawn(async move {
  115. loop {
  116. let msg: NetMsg = match p2p_recv.recv().await {
  117. Ok(m) => m,
  118. Err(e) => {
  119. error!(target: "raft", "error occurred while receiving a msg: {}", e);
  120. continue
  121. }
  122. };
  123. match p2p_cloned.broadcast(msg).await {
  124. Ok(_) => {}
  125. Err(e) => {
  126. error!(target: "raft", "error occurred during broadcasting a msg: {}", e);
  127. continue
  128. }
  129. }
  130. }
  131. });
  132. let self_nodes = self.nodes.clone();
  133. let p2p_cloned = p2p.clone();
  134. let self_id = self.id.clone();
  135. let load_ips_task = executor.spawn(async move {
  136. if self_id.is_none() {
  137. return
  138. }
  139. loop {
  140. debug!(target: "raft", "load node ids from p2p hosts ips");
  141. task::sleep(Duration::from_millis(TIMEOUT_NODES * 10)).await;
  142. let hosts = p2p_cloned.hosts().clone();
  143. let nodes_ip = hosts.load_all().await.clone();
  144. let mut nodes = self_nodes.lock().await;
  145. for ip in nodes_ip.iter() {
  146. nodes.insert(NodeId::from(*ip), *ip);
  147. }
  148. }
  149. });
  150. let mut rng = rand::thread_rng();
  151. let broadcast_msg_rv = self.broadcast_msg.1.clone();
  152. // send data form datastore through broadcast_commits channel
  153. let commits = self.datastore.commits.get_all()?;
  154. for commit in commits {
  155. self.broadcast_commits.0.send(commit).await?;
  156. }
  157. loop {
  158. let timeout: Duration;
  159. if self.role == Role::Leader {
  160. timeout = Duration::from_millis(HEARTBEATTIMEOUT);
  161. } else {
  162. timeout = Duration::from_millis(rng.gen_range(0..200) + TIMEOUT);
  163. }
  164. let result: Result<()>;
  165. select! {
  166. m = receive_queues.recv().fuse() => result = self.handle_method(m?).await,
  167. m = broadcast_msg_rv.recv().fuse() => result = self.broadcast_msg(&m?).await,
  168. _ = task::sleep(timeout).fuse() => {
  169. result = if self.role == Role::Leader {
  170. self.send_heartbeat().await
  171. }else {
  172. self.send_vote_request().await
  173. };
  174. },
  175. _ = stop_signal.recv().fuse() => break,
  176. }
  177. match result {
  178. Ok(_) => {}
  179. Err(e) => warn!(target: "raft", "warn: {}", e),
  180. }
  181. }
  182. warn!(target: "raft", "Raft start() Exit Signal");
  183. load_ips_task.cancel().await;
  184. p2p_recv_task.cancel().await;
  185. p2p_task.cancel().await;
  186. self.datastore.cancel().await?;
  187. Ok(())
  188. }
  189. pub fn get_commits(&self) -> async_channel::Receiver<T> {
  190. self.broadcast_commits.1.clone()
  191. }
  192. pub fn get_broadcast(&self) -> async_channel::Sender<T> {
  193. self.broadcast_msg.0.clone()
  194. }
  195. async fn broadcast_msg(&mut self, msg: &T) -> Result<()> {
  196. if self.role == Role::Leader {
  197. let msg = serialize(msg);
  198. let log = Log { msg, term: self.current_term };
  199. self.push_log(&log)?;
  200. self.acked_length.insert(&self.id.clone().unwrap(), self.logs.len());
  201. let nodes = self.nodes.lock().await.clone();
  202. for node in nodes.iter() {
  203. self.update_logs(node.0).await?;
  204. }
  205. } else {
  206. let b_msg = BroadcastMsgRequest(serialize(msg));
  207. self.send(
  208. self.current_leader.clone(),
  209. &serialize(&b_msg),
  210. NetMsgMethod::BroadcastRequest,
  211. )
  212. .await?;
  213. }
  214. info!(target: "raft", "has id: {} {:?} broadcast a msg", self.id.is_some(), self.role);
  215. Ok(())
  216. }
  217. async fn handle_method(&mut self, msg: NetMsg) -> Result<()> {
  218. match msg.method {
  219. NetMsgMethod::LogResponse => {
  220. let lr: LogResponse = deserialize(&msg.payload)?;
  221. self.receive_log_response(lr).await?;
  222. }
  223. NetMsgMethod::LogRequest => {
  224. let lr: LogRequest = deserialize(&msg.payload)?;
  225. self.receive_log_request(lr).await?;
  226. }
  227. NetMsgMethod::VoteResponse => {
  228. let vr: VoteResponse = deserialize(&msg.payload)?;
  229. self.receive_vote_response(vr).await?;
  230. }
  231. NetMsgMethod::VoteRequest => {
  232. let vr: VoteRequest = deserialize(&msg.payload)?;
  233. self.receive_vote_request(vr).await?;
  234. }
  235. NetMsgMethod::BroadcastRequest => {
  236. let vr: BroadcastMsgRequest = deserialize(&msg.payload)?;
  237. let d: T = deserialize(&vr.0)?;
  238. self.broadcast_msg(&d).await?;
  239. }
  240. }
  241. debug!(
  242. target: "raft",
  243. "{} {:?} receive msg id: {} recipient_id: {:?} method: {:?} ",
  244. self.id.is_some(), self.role, msg.id, &msg.recipient_id.is_some(), &msg.method
  245. );
  246. Ok(())
  247. }
  248. async fn send(
  249. &self,
  250. recipient_id: Option<NodeId>,
  251. payload: &[u8],
  252. method: NetMsgMethod,
  253. ) -> Result<()> {
  254. let random_id = OsRng.next_u32();
  255. debug!(
  256. target: "raft",
  257. "{} {:?} send a msg id: {} recipient_id: {:?} method: {:?} ",
  258. self.id.is_some(), self.role, random_id, &recipient_id.is_some(), &method
  259. );
  260. let net_msg = NetMsg { id: random_id, recipient_id, payload: payload.to_vec(), method };
  261. self.sender.0.send(net_msg).await?;
  262. Ok(())
  263. }
  264. async fn send_heartbeat(&self) -> Result<()> {
  265. if self.role == Role::Leader {
  266. let nodes = self.nodes.lock().await.clone();
  267. for node in nodes.iter() {
  268. self.update_logs(node.0).await?;
  269. }
  270. }
  271. Ok(())
  272. }
  273. async fn send_vote_request(&mut self) -> Result<()> {
  274. // this will prevent the listener node to become a candidate
  275. if self.id.is_none() {
  276. return Ok(())
  277. }
  278. let self_id = self.id.clone().unwrap();
  279. self.set_current_term(&(self.current_term + 1))?;
  280. self.role = Role::Candidate;
  281. self.set_voted_for(&Some(self_id.clone()))?;
  282. self.votes_received.push(self_id.clone());
  283. self.reset_last_term();
  284. let request = VoteRequest {
  285. node_id: self_id,
  286. current_term: self.current_term,
  287. log_length: self.logs.len(),
  288. last_term: self.last_term,
  289. };
  290. let payload = serialize(&request);
  291. self.send(None, &payload, NetMsgMethod::VoteRequest).await
  292. }
  293. async fn receive_vote_request(&mut self, vr: VoteRequest) -> Result<()> {
  294. if self.id.is_none() {
  295. return Ok(())
  296. }
  297. if vr.current_term > self.current_term {
  298. self.set_current_term(&vr.current_term)?;
  299. self.set_voted_for(&None)?;
  300. self.role = Role::Follower;
  301. }
  302. self.reset_last_term();
  303. // check the logs of the candidate
  304. let vote_ok = (vr.last_term > self.last_term) ||
  305. (vr.last_term == self.last_term && vr.log_length >= self.logs.len());
  306. // slef.voted_for equal to vr.node_id or is None or voted to someone else
  307. let vote = if let Some(voted_for) = self.voted_for.as_ref() {
  308. *voted_for == vr.node_id
  309. } else {
  310. true
  311. };
  312. let mut response = VoteResponse {
  313. node_id: self.id.clone().unwrap(),
  314. current_term: self.current_term,
  315. ok: false,
  316. };
  317. if vr.current_term == self.current_term && vote_ok && vote {
  318. self.set_voted_for(&Some(vr.node_id.clone()))?;
  319. response.set_ok(true);
  320. }
  321. let payload = serialize(&response);
  322. self.send(Some(vr.node_id), &payload, NetMsgMethod::VoteResponse).await
  323. }
  324. async fn receive_vote_response(&mut self, vr: VoteResponse) -> Result<()> {
  325. if self.role == Role::Candidate && vr.current_term == self.current_term && vr.ok {
  326. self.votes_received.push(vr.node_id);
  327. let nodes = self.nodes.lock().await;
  328. if self.votes_received.len() >= ((nodes.len() + 1) / 2) {
  329. self.role = Role::Leader;
  330. self.current_leader = Some(self.id.clone().unwrap());
  331. for node in nodes.iter() {
  332. self.sent_length.insert(node.0, self.logs.len());
  333. self.acked_length.insert(node.0, 0);
  334. self.update_logs(node.0).await?;
  335. }
  336. }
  337. drop(nodes);
  338. } else if vr.current_term > self.current_term {
  339. self.set_current_term(&vr.current_term)?;
  340. self.role = Role::Follower;
  341. self.set_voted_for(&None)?;
  342. }
  343. Ok(())
  344. }
  345. async fn update_logs(&self, node_id: &NodeId) -> Result<()> {
  346. let prefix_len = match self.sent_length.get(node_id) {
  347. Ok(len) => len,
  348. Err(_) => {
  349. // return if failed to index
  350. return Ok(())
  351. }
  352. };
  353. let suffix: Logs = if self.logs.slice_from(prefix_len).is_some() {
  354. self.logs.slice_from(prefix_len).unwrap()
  355. } else {
  356. return Ok(())
  357. };
  358. let mut prefix_term = 0;
  359. if prefix_len > 0 {
  360. prefix_term = self.logs.get(prefix_len - 1)?.term;
  361. }
  362. let request = LogRequest {
  363. leader_id: self.id.clone().unwrap(),
  364. current_term: self.current_term,
  365. prefix_len,
  366. prefix_term,
  367. commit_length: self.commit_length,
  368. suffix,
  369. };
  370. let payload = serialize(&request);
  371. self.send(Some(node_id.clone()), &payload, NetMsgMethod::LogRequest).await
  372. }
  373. async fn receive_log_request(&mut self, lr: LogRequest) -> Result<()> {
  374. if lr.current_term > self.current_term {
  375. self.set_current_term(&lr.current_term)?;
  376. self.set_voted_for(&None)?;
  377. }
  378. if lr.current_term == self.current_term {
  379. self.role = Role::Follower;
  380. self.current_leader = Some(lr.leader_id.clone());
  381. }
  382. let ok = (self.logs.len() >= lr.prefix_len) &&
  383. (lr.prefix_len == 0 || self.logs.get(lr.prefix_len - 1)?.term == lr.prefix_term);
  384. let mut ack = 0;
  385. if lr.current_term == self.current_term && ok {
  386. self.append_log(lr.prefix_len, lr.commit_length, &lr.suffix).await?;
  387. ack = lr.prefix_len + lr.suffix.len();
  388. }
  389. if self.id.is_none() {
  390. return Ok(())
  391. }
  392. let response = LogResponse {
  393. node_id: self.id.clone().unwrap(),
  394. current_term: self.current_term,
  395. ack,
  396. ok,
  397. };
  398. let payload = serialize(&response);
  399. self.send(Some(lr.leader_id.clone()), &payload, NetMsgMethod::LogResponse).await
  400. }
  401. async fn receive_log_response(&mut self, lr: LogResponse) -> Result<()> {
  402. if lr.current_term == self.current_term && self.role == Role::Leader {
  403. if lr.ok && lr.ack >= self.acked_length.get(&lr.node_id)? {
  404. self.sent_length.insert(&lr.node_id, lr.ack);
  405. self.acked_length.insert(&lr.node_id, lr.ack);
  406. self.commit_log().await?;
  407. } else if self.sent_length.get(&lr.node_id)? > 0 {
  408. self.sent_length.insert(&lr.node_id, self.sent_length.get(&lr.node_id)? - 1);
  409. self.update_logs(&lr.node_id).await?;
  410. }
  411. } else if lr.current_term > self.current_term {
  412. self.set_current_term(&lr.current_term)?;
  413. self.role = Role::Follower;
  414. self.set_voted_for(&None)?;
  415. }
  416. Ok(())
  417. }
  418. fn reset_last_term(&mut self) {
  419. self.last_term = 0;
  420. if let Some(log) = self.logs.0.last() {
  421. self.last_term = log.term;
  422. }
  423. }
  424. fn acks(&self, nodes: HashMap<NodeId, SocketAddr>, length: u64) -> HashMap<NodeId, SocketAddr> {
  425. nodes
  426. .into_iter()
  427. .filter(|n| {
  428. let len = self.acked_length.get(&n.0);
  429. len.is_ok() && len.unwrap() >= length
  430. })
  431. .collect()
  432. }
  433. async fn commit_log(&mut self) -> Result<()> {
  434. let nodes_ptr = self.nodes.lock().await;
  435. let min_acks = ((nodes_ptr.len() + 1) / 2) as usize;
  436. let nodes = nodes_ptr.clone();
  437. drop(nodes_ptr);
  438. let ready: Vec<u64> = self
  439. .logs
  440. .0
  441. .iter()
  442. .enumerate()
  443. .filter(|(i, _)| self.acks(nodes.clone(), *i as u64).len() >= min_acks)
  444. .map(|(i, _)| i as u64)
  445. .collect();
  446. if ready.is_empty() {
  447. return Ok(())
  448. }
  449. let max_ready = *ready.iter().max().unwrap();
  450. if max_ready > self.commit_length && self.logs.get(max_ready - 1)?.term == self.current_term
  451. {
  452. for i in self.commit_length..max_ready {
  453. self.push_commit(&self.logs.get(i)?.msg).await?;
  454. }
  455. self.set_commit_length(&max_ready)?;
  456. }
  457. Ok(())
  458. }
  459. async fn append_log(
  460. &mut self,
  461. prefix_len: u64,
  462. leader_commit: u64,
  463. suffix: &Logs,
  464. ) -> Result<()> {
  465. if !suffix.is_empty() && self.logs.len() > prefix_len {
  466. let index = min(self.logs.len(), prefix_len + suffix.len()) - 1;
  467. if self.logs.get(index)?.term != suffix.get(index - prefix_len)?.term {
  468. self.push_logs(&self.logs.slice_to(prefix_len))?;
  469. }
  470. }
  471. if prefix_len + suffix.len() > self.logs.len() {
  472. for i in (self.logs.len() - prefix_len)..(suffix.len() - 1) {
  473. self.push_log(&suffix.get(i)?)?;
  474. }
  475. }
  476. if leader_commit > self.commit_length {
  477. for i in self.commit_length..leader_commit {
  478. self.push_commit(&self.logs.get(i)?.msg).await?;
  479. }
  480. self.set_commit_length(&leader_commit)?;
  481. }
  482. Ok(())
  483. }
  484. fn set_commit_length(&mut self, i: &u64) -> Result<()> {
  485. self.commit_length = *i;
  486. self.datastore.commits_length.insert(i)
  487. }
  488. fn set_current_term(&mut self, i: &u64) -> Result<()> {
  489. self.current_term = *i;
  490. self.datastore.current_term.insert(i)
  491. }
  492. fn set_voted_for(&mut self, i: &Option<NodeId>) -> Result<()> {
  493. self.voted_for = i.clone();
  494. self.datastore.voted_for.insert(i)
  495. }
  496. async fn push_commit(&mut self, commit: &[u8]) -> Result<()> {
  497. let commit: T = deserialize(commit)?;
  498. self.broadcast_commits.0.send(commit.clone()).await?;
  499. self.datastore.commits.insert(&commit)
  500. }
  501. fn push_log(&mut self, i: &Log) -> Result<()> {
  502. self.logs.push(i);
  503. self.datastore.logs.insert(i)
  504. }
  505. fn push_logs(&mut self, i: &Logs) -> Result<()> {
  506. self.logs = i.clone();
  507. self.datastore.logs.wipe_insert_all(&i.to_vec())
  508. }
  509. }