Skip to main content

tatara_engine/cluster/
raft_log.rs

1use std::fmt::Debug;
2use std::ops::RangeBounds;
3use std::path::Path;
4use std::sync::Arc;
5
6use openraft::anyerror::AnyError;
7use openraft::storage::RaftLogStorage;
8use openraft::{
9    Entry, LogId, LogState, OptionalSend, RaftLogReader, StorageError, StorageIOError, Vote,
10};
11use redb::{Database, ReadableDatabase, ReadableTable, TableDefinition};
12use tokio::sync::Mutex;
13
14use super::raft_sm::TypeConfig;
15use tatara_core::cluster::types::NodeId;
16
17const LOG_TABLE: TableDefinition<u64, &[u8]> = TableDefinition::new("raft_log");
18const META_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("raft_meta");
19
20const VOTE_KEY: &str = "vote";
21const PURGE_KEY: &str = "last_purged";
22
23fn io_read_logs<E: std::error::Error + 'static>(e: &E) -> StorageError<NodeId> {
24    StorageIOError::<NodeId>::read_logs(AnyError::new(e)).into()
25}
26
27fn io_write_logs<E: std::error::Error + 'static>(e: &E) -> StorageError<NodeId> {
28    StorageIOError::<NodeId>::write_logs(AnyError::new(e)).into()
29}
30
31fn io_read_vote<E: std::error::Error + 'static>(e: &E) -> StorageError<NodeId> {
32    StorageIOError::<NodeId>::read_vote(AnyError::new(e)).into()
33}
34
35fn io_write_vote<E: std::error::Error + 'static>(e: &E) -> StorageError<NodeId> {
36    StorageIOError::<NodeId>::write_vote(AnyError::new(e)).into()
37}
38
39/// Raft log storage backed by redb (pure Rust embedded KV store).
40pub struct LogStore {
41    db: Arc<Database>,
42    /// Serialization lock — redb supports only one writer at a time.
43    write_lock: Arc<Mutex<()>>,
44}
45
46impl Clone for LogStore {
47    fn clone(&self) -> Self {
48        Self {
49            db: self.db.clone(),
50            write_lock: self.write_lock.clone(),
51        }
52    }
53}
54
55impl LogStore {
56    pub fn new(path: &Path) -> anyhow::Result<Self> {
57        if let Some(parent) = path.parent() {
58            std::fs::create_dir_all(parent)?;
59        }
60        let db = Database::create(path)?;
61
62        // Ensure tables exist
63        let write_txn = db.begin_write()?;
64        write_txn.open_table(LOG_TABLE)?;
65        write_txn.open_table(META_TABLE)?;
66        write_txn.commit()?;
67
68        Ok(Self {
69            db: Arc::new(db),
70            write_lock: Arc::new(Mutex::new(())),
71        })
72    }
73}
74
75impl RaftLogReader<TypeConfig> for LogStore {
76    async fn try_get_log_entries<RB: RangeBounds<u64> + Clone + Debug + OptionalSend>(
77        &mut self,
78        range: RB,
79    ) -> Result<Vec<Entry<TypeConfig>>, StorageError<NodeId>> {
80        let read_txn = self.db.begin_read().map_err(|e| io_read_logs(&e))?;
81        let table = read_txn
82            .open_table(LOG_TABLE)
83            .map_err(|e| io_read_logs(&e))?;
84
85        let mut entries = Vec::new();
86        let iter = table.range(range).map_err(|e| io_read_logs(&e))?;
87
88        for item in iter {
89            let (_, value) = item.map_err(|e| io_read_logs(&e))?;
90            let entry: Entry<TypeConfig> =
91                serde_json::from_slice(value.value()).map_err(|e| io_read_logs(&e))?;
92            entries.push(entry);
93        }
94
95        Ok(entries)
96    }
97}
98
99impl RaftLogStorage<TypeConfig> for LogStore {
100    type LogReader = LogStore;
101
102    async fn get_log_state(&mut self) -> Result<LogState<TypeConfig>, StorageError<NodeId>> {
103        let read_txn = self.db.begin_read().map_err(|e| io_read_logs(&e))?;
104
105        // Check for persisted purge state
106        let meta_table = read_txn
107            .open_table(META_TABLE)
108            .map_err(|e| io_read_logs(&e))?;
109        let last_purged = match meta_table.get(PURGE_KEY).map_err(|e| io_read_logs(&e))? {
110            Some(guard) => {
111                let log_id: LogId<NodeId> =
112                    serde_json::from_slice(guard.value()).map_err(|e| io_read_logs(&e))?;
113                Some(log_id)
114            }
115            None => None,
116        };
117
118        let table = read_txn
119            .open_table(LOG_TABLE)
120            .map_err(|e| io_read_logs(&e))?;
121
122        let last = table.last().map_err(|e| io_read_logs(&e))?;
123
124        let last_log_id = match last {
125            Some(entry) => {
126                let bytes = entry.1.value();
127                let log_entry: Entry<TypeConfig> =
128                    serde_json::from_slice(bytes).map_err(|e| io_read_logs(&e))?;
129                Some(log_entry.log_id)
130            }
131            None => last_purged,
132        };
133
134        Ok(LogState {
135            last_purged_log_id: last_purged,
136            last_log_id,
137        })
138    }
139
140    async fn get_log_reader(&mut self) -> Self::LogReader {
141        self.clone()
142    }
143
144    async fn save_vote(&mut self, vote: &Vote<NodeId>) -> Result<(), StorageError<NodeId>> {
145        let _lock = self.write_lock.lock().await;
146        let bytes = serde_json::to_vec(vote).map_err(|e| io_write_vote(&e))?;
147
148        let write_txn = self.db.begin_write().map_err(|e| io_write_vote(&e))?;
149        {
150            let mut table = write_txn
151                .open_table(META_TABLE)
152                .map_err(|e| io_write_vote(&e))?;
153            table
154                .insert(VOTE_KEY, bytes.as_slice())
155                .map_err(|e| io_write_vote(&e))?;
156        }
157        write_txn.commit().map_err(|e| io_write_vote(&e))?;
158
159        Ok(())
160    }
161
162    async fn read_vote(&mut self) -> Result<Option<Vote<NodeId>>, StorageError<NodeId>> {
163        let read_txn = self.db.begin_read().map_err(|e| io_read_vote(&e))?;
164        let table = read_txn
165            .open_table(META_TABLE)
166            .map_err(|e| io_read_vote(&e))?;
167
168        match table.get(VOTE_KEY).map_err(|e| io_read_vote(&e))? {
169            Some(guard) => {
170                let vote: Vote<NodeId> =
171                    serde_json::from_slice(guard.value()).map_err(|e| io_read_vote(&e))?;
172                Ok(Some(vote))
173            }
174            None => Ok(None),
175        }
176    }
177
178    async fn append<I>(
179        &mut self,
180        entries: I,
181        callback: openraft::storage::LogFlushed<TypeConfig>,
182    ) -> Result<(), StorageError<NodeId>>
183    where
184        I: IntoIterator<Item = Entry<TypeConfig>> + OptionalSend,
185    {
186        let _lock = self.write_lock.lock().await;
187        let write_txn = self.db.begin_write().map_err(|e| io_write_logs(&e))?;
188        {
189            let mut table = write_txn
190                .open_table(LOG_TABLE)
191                .map_err(|e| io_write_logs(&e))?;
192
193            for entry in entries {
194                let index = entry.log_id.index;
195                let bytes = serde_json::to_vec(&entry).map_err(|e| io_write_logs(&e))?;
196                table
197                    .insert(index, bytes.as_slice())
198                    .map_err(|e| io_write_logs(&e))?;
199            }
200        }
201        write_txn.commit().map_err(|e| io_write_logs(&e))?;
202
203        callback.log_io_completed(Ok(()));
204        Ok(())
205    }
206
207    async fn truncate(&mut self, log_id: LogId<NodeId>) -> Result<(), StorageError<NodeId>> {
208        let _lock = self.write_lock.lock().await;
209        let write_txn = self.db.begin_write().map_err(|e| io_write_logs(&e))?;
210        {
211            let mut table = write_txn
212                .open_table(LOG_TABLE)
213                .map_err(|e| io_write_logs(&e))?;
214
215            // Remove all entries from log_id.index onwards
216            let to_remove: Vec<u64> = table
217                .range(log_id.index..)
218                .map_err(|e| io_write_logs(&e))?
219                .map(|entry| entry.unwrap().0.value())
220                .collect();
221
222            for idx in to_remove {
223                table.remove(idx).map_err(|e| io_write_logs(&e))?;
224            }
225        }
226        write_txn.commit().map_err(|e| io_write_logs(&e))?;
227
228        Ok(())
229    }
230
231    async fn purge(&mut self, log_id: LogId<NodeId>) -> Result<(), StorageError<NodeId>> {
232        let _lock = self.write_lock.lock().await;
233        let write_txn = self.db.begin_write().map_err(|e| io_write_logs(&e))?;
234        {
235            let mut table = write_txn
236                .open_table(LOG_TABLE)
237                .map_err(|e| io_write_logs(&e))?;
238
239            // Remove all entries up to and including log_id.index
240            let to_remove: Vec<u64> = table
241                .range(..=log_id.index)
242                .map_err(|e| io_write_logs(&e))?
243                .map(|entry| entry.unwrap().0.value())
244                .collect();
245
246            for idx in to_remove {
247                table.remove(idx).map_err(|e| io_write_logs(&e))?;
248            }
249
250            // Persist the purge point
251            let mut meta = write_txn
252                .open_table(META_TABLE)
253                .map_err(|e| io_write_logs(&e))?;
254            let purge_bytes = serde_json::to_vec(&log_id).map_err(|e| io_write_logs(&e))?;
255            meta.insert(PURGE_KEY, purge_bytes.as_slice())
256                .map_err(|e| io_write_logs(&e))?;
257        }
258        write_txn.commit().map_err(|e| io_write_logs(&e))?;
259
260        Ok(())
261    }
262}