Skip to main content

mocra_cluster/
state_machine.rs

1//! A redb-backed replicated state machine.
2//!
3//! Commands mutate the persistent state deterministically through [`apply`](StateMachine::apply).
4//! This is exactly the state machine Raft replicates: every node calls `apply` once a log entry is
5//! committed and arrives at the same result.
6//!
7//! Tables: `kv` (general-purpose KV), `locks` (distributed locks), `meta` (the fencing counter).
8
9use std::path::Path;
10use std::sync::Arc;
11
12use redb::{Database, ReadableTable, TableDefinition};
13
14use crate::cmd::{Cmd, CmdResult, Lock};
15
16const KV: TableDefinition<&[u8], &[u8]> = TableDefinition::new("kv");
17const LOCKS: TableDefinition<&str, &[u8]> = TableDefinition::new("locks");
18const META: TableDefinition<&str, u64> = TableDefinition::new("meta");
19const FENCING_COUNTER: &str = "fencing_counter";
20
21/// State machine errors.
22#[derive(Debug, thiserror::Error)]
23pub enum StateMachineError {
24    #[error("redb: {0}")]
25    Redb(String),
26    #[error("codec: {0}")]
27    Codec(String),
28}
29
30fn redb_err<E: std::fmt::Display>(e: E) -> StateMachineError {
31    StateMachineError::Redb(e.to_string())
32}
33
34fn encode_lock(l: &Lock) -> Result<Vec<u8>, StateMachineError> {
35    rmp_serde::to_vec(l).map_err(|e| StateMachineError::Codec(e.to_string()))
36}
37
38fn decode_lock(b: &[u8]) -> Result<Lock, StateMachineError> {
39    rmp_serde::from_slice(b).map_err(|e| StateMachineError::Codec(e.to_string()))
40}
41
42/// A redb-backed replicated state machine.
43#[derive(Clone)]
44pub struct StateMachine {
45    db: Arc<Database>,
46}
47
48impl StateMachine {
49    /// Open (or create) a redb-backed state machine, creating the required tables up front.
50    pub fn open(path: impl AsRef<Path>) -> Result<Self, StateMachineError> {
51        let db = Database::create(path).map_err(redb_err)?;
52        let w = db.begin_write().map_err(redb_err)?;
53        {
54            w.open_table(KV).map_err(redb_err)?;
55            w.open_table(LOCKS).map_err(redb_err)?;
56            w.open_table(META).map_err(redb_err)?;
57        }
58        w.commit().map_err(redb_err)?;
59        Ok(Self { db: Arc::new(db) })
60    }
61
62    /// Apply a single command deterministically (called by every node once Raft has committed it).
63    pub fn apply(&self, cmd: &Cmd) -> Result<CmdResult, StateMachineError> {
64        let w = self.db.begin_write().map_err(redb_err)?;
65        let result = match cmd {
66            Cmd::Set { key, value } => {
67                let mut t = w.open_table(KV).map_err(redb_err)?;
68                t.insert(key.as_slice(), value.as_slice())
69                    .map_err(redb_err)?;
70                CmdResult::Ok
71            }
72            Cmd::Delete { key } => {
73                let mut t = w.open_table(KV).map_err(redb_err)?;
74                t.remove(key.as_slice()).map_err(redb_err)?;
75                CmdResult::Ok
76            }
77            Cmd::Cas { key, expect, value } => {
78                let mut t = w.open_table(KV).map_err(redb_err)?;
79                let cur = t
80                    .get(key.as_slice())
81                    .map_err(redb_err)?
82                    .map(|g| g.value().to_vec());
83                if cur.as_deref() == expect.as_deref() {
84                    t.insert(key.as_slice(), value.as_slice())
85                        .map_err(redb_err)?;
86                    CmdResult::Bool(true)
87                } else {
88                    CmdResult::Bool(false)
89                }
90            }
91            Cmd::AcquireLock {
92                key,
93                holder,
94                now_ms,
95                ttl_ms,
96            } => {
97                let mut locks = w.open_table(LOCKS).map_err(redb_err)?;
98                let cur = match locks.get(key.as_str()).map_err(redb_err)? {
99                    Some(g) => Some(decode_lock(g.value())?),
100                    None => None,
101                };
102                let free = match &cur {
103                    None => true,
104                    Some(l) => l.expire_at_ms <= *now_ms || l.holder == *holder,
105                };
106                if free {
107                    let token = {
108                        let mut meta = w.open_table(META).map_err(redb_err)?;
109                        let n = meta
110                            .get(FENCING_COUNTER)
111                            .map_err(redb_err)?
112                            .map(|g| g.value())
113                            .unwrap_or(0)
114                            + 1;
115                        meta.insert(FENCING_COUNTER, n).map_err(redb_err)?;
116                        n
117                    };
118                    let lock = Lock {
119                        holder: holder.clone(),
120                        expire_at_ms: now_ms + ttl_ms,
121                        fencing_token: token,
122                    };
123                    locks
124                        .insert(key.as_str(), encode_lock(&lock)?.as_slice())
125                        .map_err(redb_err)?;
126                    CmdResult::Fencing(Some(token))
127                } else {
128                    CmdResult::Fencing(None)
129                }
130            }
131            Cmd::RenewLock {
132                key,
133                holder,
134                now_ms,
135                ttl_ms,
136            } => {
137                let mut locks = w.open_table(LOCKS).map_err(redb_err)?;
138                let cur = match locks.get(key.as_str()).map_err(redb_err)? {
139                    Some(g) => Some(decode_lock(g.value())?),
140                    None => None,
141                };
142                match cur {
143                    Some(mut l) if l.holder == *holder && l.expire_at_ms > *now_ms => {
144                        l.expire_at_ms = now_ms + ttl_ms;
145                        locks
146                            .insert(key.as_str(), encode_lock(&l)?.as_slice())
147                            .map_err(redb_err)?;
148                        CmdResult::Bool(true)
149                    }
150                    _ => CmdResult::Bool(false),
151                }
152            }
153            Cmd::ReleaseLock { key, holder } => {
154                let mut locks = w.open_table(LOCKS).map_err(redb_err)?;
155                let held_by_holder = match locks.get(key.as_str()).map_err(redb_err)? {
156                    Some(g) => decode_lock(g.value())?.holder == *holder,
157                    None => false,
158                };
159                if held_by_holder {
160                    locks.remove(key.as_str()).map_err(redb_err)?;
161                }
162                CmdResult::Ok
163            }
164        };
165        w.commit().map_err(redb_err)?;
166        Ok(result)
167    }
168
169    /// Read a KV pair (a local read; linearizability is guaranteed by the Raft read-index above).
170    pub fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>, StateMachineError> {
171        let r = self.db.begin_read().map_err(redb_err)?;
172        let t = r.open_table(KV).map_err(redb_err)?;
173        Ok(t.get(key).map_err(redb_err)?.map(|g| g.value().to_vec()))
174    }
175
176    /// For snapshots: serialize the entire business state (kv + locks) into bytes.
177    pub fn dump(&self) -> Result<Vec<u8>, StateMachineError> {
178        let r = self.db.begin_read().map_err(redb_err)?;
179        let mut kv = Vec::new();
180        {
181            let t = r.open_table(KV).map_err(redb_err)?;
182            for item in t.iter().map_err(redb_err)? {
183                let (k, v) = item.map_err(redb_err)?;
184                kv.push((k.value().to_vec(), v.value().to_vec()));
185            }
186        }
187        let mut locks = Vec::new();
188        {
189            let t = r.open_table(LOCKS).map_err(redb_err)?;
190            for item in t.iter().map_err(redb_err)? {
191                let (k, v) = item.map_err(redb_err)?;
192                locks.push((k.value().to_string(), v.value().to_vec()));
193            }
194        }
195        rmp_serde::to_vec(&SmDump { kv, locks })
196            .map_err(|e| StateMachineError::Codec(e.to_string()))
197    }
198
199    /// Restore from a snapshot: clear kv / locks, then load the contents.
200    pub fn restore(&self, bytes: &[u8]) -> Result<(), StateMachineError> {
201        let dump: SmDump =
202            rmp_serde::from_slice(bytes).map_err(|e| StateMachineError::Codec(e.to_string()))?;
203        let w = self.db.begin_write().map_err(redb_err)?;
204        {
205            let mut t = w.open_table(KV).map_err(redb_err)?;
206            t.retain(|_, _| false).map_err(redb_err)?;
207            for (k, v) in &dump.kv {
208                t.insert(k.as_slice(), v.as_slice()).map_err(redb_err)?;
209            }
210        }
211        {
212            let mut t = w.open_table(LOCKS).map_err(redb_err)?;
213            t.retain(|_, _| false).map_err(redb_err)?;
214            for (k, v) in &dump.locks {
215                t.insert(k.as_str(), v.as_slice()).map_err(redb_err)?;
216            }
217        }
218        w.commit().map_err(redb_err)?;
219        Ok(())
220    }
221}
222
223/// The serializable representation of a state machine snapshot (all of kv + locks).
224#[derive(serde::Serialize, serde::Deserialize)]
225struct SmDump {
226    kv: Vec<(Vec<u8>, Vec<u8>)>,
227    locks: Vec<(String, Vec<u8>)>,
228}
229
230#[cfg(test)]
231mod tests {
232    use super::*;
233
234    fn sm(dir: &tempfile::TempDir) -> StateMachine {
235        StateMachine::open(dir.path().join("sm.redb")).unwrap()
236    }
237
238    #[test]
239    fn kv_and_cas() {
240        let dir = tempfile::tempdir().unwrap();
241        let sm = sm(&dir);
242        sm.apply(&Cmd::Set {
243            key: b"x".to_vec(),
244            value: b"1".to_vec(),
245        })
246        .unwrap();
247        assert_eq!(sm.get(b"x").unwrap(), Some(b"1".to_vec()));
248
249        // CAS fails (expect does not match).
250        assert_eq!(
251            sm.apply(&Cmd::Cas {
252                key: b"x".to_vec(),
253                expect: Some(b"9".to_vec()),
254                value: b"2".to_vec()
255            })
256            .unwrap(),
257            CmdResult::Bool(false)
258        );
259        // CAS succeeds.
260        assert_eq!(
261            sm.apply(&Cmd::Cas {
262                key: b"x".to_vec(),
263                expect: Some(b"1".to_vec()),
264                value: b"2".to_vec()
265            })
266            .unwrap(),
267            CmdResult::Bool(true)
268        );
269        assert_eq!(sm.get(b"x").unwrap(), Some(b"2".to_vec()));
270    }
271
272    #[test]
273    fn lock_fencing_and_expiry() {
274        let dir = tempfile::tempdir().unwrap();
275        let sm = sm(&dir);
276
277        // a acquires the lock -> fencing token 1.
278        assert_eq!(
279            sm.apply(&Cmd::AcquireLock {
280                key: "k".into(),
281                holder: "a".into(),
282                now_ms: 1000,
283                ttl_ms: 5000
284            })
285            .unwrap(),
286            CmdResult::Fencing(Some(1))
287        );
288        // b is rejected while the lock has not expired.
289        assert_eq!(
290            sm.apply(&Cmd::AcquireLock {
291                key: "k".into(),
292                holder: "b".into(),
293                now_ms: 2000,
294                ttl_ms: 5000
295            })
296            .unwrap(),
297            CmdResult::Fencing(None)
298        );
299        // After expiry b acquires it -> the fencing token increments to 2.
300        assert_eq!(
301            sm.apply(&Cmd::AcquireLock {
302                key: "k".into(),
303                holder: "b".into(),
304                now_ms: 7000,
305                ttl_ms: 5000
306            })
307            .unwrap(),
308            CmdResult::Fencing(Some(2))
309        );
310        // a's lease renewal fails (it is no longer the holder).
311        assert_eq!(
312            sm.apply(&Cmd::RenewLock {
313                key: "k".into(),
314                holder: "a".into(),
315                now_ms: 8000,
316                ttl_ms: 5000
317            })
318            .unwrap(),
319            CmdResult::Bool(false)
320        );
321        // b's lease renewal succeeds.
322        assert_eq!(
323            sm.apply(&Cmd::RenewLock {
324                key: "k".into(),
325                holder: "b".into(),
326                now_ms: 9000,
327                ttl_ms: 5000
328            })
329            .unwrap(),
330            CmdResult::Bool(true)
331        );
332        // Once b releases it, a can acquire it.
333        sm.apply(&Cmd::ReleaseLock {
334            key: "k".into(),
335            holder: "b".into(),
336        })
337        .unwrap();
338        assert_eq!(
339            sm.apply(&Cmd::AcquireLock {
340                key: "k".into(),
341                holder: "a".into(),
342                now_ms: 10000,
343                ttl_ms: 5000
344            })
345            .unwrap(),
346            CmdResult::Fencing(Some(3))
347        );
348    }
349}