tatara_engine/cluster/
raft_node.rs1use anyhow::{Context, Result};
2use openraft::{BasicNode, Config, Raft};
3use std::collections::BTreeMap;
4use std::path::PathBuf;
5use std::sync::Arc;
6use tokio::sync::RwLock;
7use tracing::info;
8
9use super::network::RaftHttpNetwork;
10use super::raft_log::LogStore;
11use super::raft_sm::{StateMachine, StateMachineData, TypeConfig};
12use tatara_core::cluster::types::{ClusterCommand, ClusterResponse, NodeId};
13
14pub type TataraRaft = Raft<TypeConfig>;
15
16pub struct RaftCluster {
18 pub raft: TataraRaft,
19 pub state_machine: Arc<RwLock<StateMachineData>>,
20 pub network: RaftHttpNetwork,
21 node_id: NodeId,
22}
23
24impl RaftCluster {
25 pub async fn start(node_id: NodeId, _raft_addr: &str, data_dir: &PathBuf) -> Result<Self> {
27 let config = Config {
28 heartbeat_interval: 500,
29 election_timeout_min: 1500,
30 election_timeout_max: 3000,
31 max_in_snapshot_log_to_keep: 500,
32 ..Default::default()
33 };
34 let config = Arc::new(config.validate().context("Invalid Raft config")?);
35
36 let log_path = data_dir.join("raft.redb");
37 let log_store = LogStore::new(&log_path).context("Failed to open Raft log store")?;
38
39 let state_machine = StateMachine::new();
40 let sm_data = state_machine.state();
41
42 let network = RaftHttpNetwork::new();
43
44 let raft = Raft::new(node_id, config, network.clone(), log_store, state_machine)
45 .await
46 .context("Failed to initialize Raft")?;
47
48 info!(node_id = node_id, "Raft node initialized");
49
50 Ok(Self {
51 raft,
52 state_machine: sm_data,
53 network,
54 node_id,
55 })
56 }
57
58 pub async fn bootstrap_single(&self, raft_addr: &str) -> Result<()> {
60 let mut members = BTreeMap::new();
61 members.insert(
62 self.node_id,
63 BasicNode {
64 addr: raft_addr.to_string(),
65 },
66 );
67
68 self.raft
69 .initialize(members)
70 .await
71 .map_err(|e| anyhow::anyhow!("Raft bootstrap failed: {}", e))?;
72
73 info!(
74 node_id = self.node_id,
75 "Bootstrapped single-node Raft cluster"
76 );
77 Ok(())
78 }
79
80 pub async fn add_voter(&self, node_id: NodeId, addr: &str) -> Result<()> {
82 let node = BasicNode {
83 addr: addr.to_string(),
84 };
85
86 self.raft
88 .add_learner(node_id, node, true)
89 .await
90 .map_err(|e| anyhow::anyhow!("Failed to add learner: {}", e))?;
91
92 let members = self.current_members().await?;
94 let mut new_members = members;
95 new_members.insert(node_id);
96
97 self.raft
98 .change_membership(new_members, false)
99 .await
100 .map_err(|e| anyhow::anyhow!("Failed to change membership: {}", e))?;
101
102 info!(
103 node_id = node_id,
104 addr = addr,
105 "Added voter to Raft cluster"
106 );
107 Ok(())
108 }
109
110 pub async fn add_learner(&self, node_id: NodeId, addr: &str) -> Result<()> {
112 let node = BasicNode {
113 addr: addr.to_string(),
114 };
115
116 self.raft
117 .add_learner(node_id, node, true)
118 .await
119 .map_err(|e| anyhow::anyhow!("Failed to add learner: {}", e))?;
120
121 info!(node_id = node_id, "Added learner to Raft cluster");
122 Ok(())
123 }
124
125 pub async fn write(&self, cmd: ClusterCommand) -> Result<ClusterResponse> {
127 let resp = self
128 .raft
129 .client_write(cmd)
130 .await
131 .map_err(|e| anyhow::anyhow!("Raft write failed: {}", e))?;
132
133 Ok(resp.data)
134 }
135
136 pub async fn read_state(&self) -> Result<Arc<RwLock<StateMachineData>>> {
138 self.raft
140 .ensure_linearizable()
141 .await
142 .map_err(|e| anyhow::anyhow!("Linearizable read failed: {}", e))?;
143
144 Ok(self.state_machine.clone())
145 }
146
147 pub async fn read_local(&self) -> Arc<RwLock<StateMachineData>> {
149 self.state_machine.clone()
150 }
151
152 pub fn read_local_sync(&self) -> Arc<RwLock<StateMachineData>> {
155 self.state_machine.clone()
156 }
157
158 pub async fn is_leader(&self) -> bool {
160 self.raft.ensure_linearizable().await.is_ok()
161 }
162
163 pub async fn current_leader(&self) -> Option<NodeId> {
165 self.raft.current_leader().await
166 }
167
168 async fn current_members(&self) -> Result<std::collections::BTreeSet<NodeId>> {
170 let metrics = self.raft.metrics().borrow().clone();
171 Ok(metrics.membership_config.membership().voter_ids().collect())
172 }
173
174 pub fn node_id(&self) -> NodeId {
175 self.node_id
176 }
177}