1use 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#[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#[derive(Clone)]
44pub struct StateMachine {
45 db: Arc<Database>,
46}
47
48impl StateMachine {
49 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 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 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 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 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#[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 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 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 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 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 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 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 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 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}