smb_server_handle_store/
redb_store.rs1use std::path::Path;
8use std::sync::Arc;
9
10use async_trait::async_trait;
11use redb::{Database, ReadableTable, TableDefinition};
12
13use crate::{Guid, HandleRecord, HandleStore, StoreError};
14
15const HANDLES: TableDefinition<&[u8], &[u8]> = TableDefinition::new("handles");
16
17pub struct RedbStore {
19 db: Arc<Database>,
20}
21
22impl RedbStore {
23 pub fn open(path: impl AsRef<Path>) -> Result<RedbStore, StoreError> {
25 let db = Database::create(path).map_err(be)?;
26 let wtx = db.begin_write().map_err(be)?;
27 wtx.open_table(HANDLES).map_err(be)?;
28 wtx.commit().map_err(be)?;
29 Ok(RedbStore { db: Arc::new(db) })
30 }
31}
32
33fn be<E: std::fmt::Display>(e: E) -> StoreError {
34 StoreError::Backend(e.to_string())
35}
36
37fn se<E: std::fmt::Display>(e: E) -> StoreError {
38 StoreError::Serde(e.to_string())
39}
40
41async fn blocking<T, F>(f: F) -> Result<T, StoreError>
43where
44 F: FnOnce() -> Result<T, StoreError> + Send + 'static,
45 T: Send + 'static,
46{
47 tokio::task::spawn_blocking(f).await.map_err(be)?
48}
49
50#[async_trait]
51impl HandleStore for RedbStore {
52 async fn put(&self, record: HandleRecord) -> Result<(), StoreError> {
53 let db = self.db.clone();
54 blocking(move || {
55 let value = serde_json::to_vec(&record).map_err(se)?;
56 let wtx = db.begin_write().map_err(be)?;
57 {
58 let mut table = wtx.open_table(HANDLES).map_err(be)?;
59 table.insert(record.create_guid.as_slice(), value.as_slice()).map_err(be)?;
60 }
61 wtx.commit().map_err(be)
62 })
63 .await
64 }
65
66 async fn get(&self, create_guid: &Guid) -> Result<Option<HandleRecord>, StoreError> {
67 let db = self.db.clone();
68 let guid = *create_guid;
69 blocking(move || {
70 let rtx = db.begin_read().map_err(be)?;
71 let table = rtx.open_table(HANDLES).map_err(be)?;
72 match table.get(guid.as_slice()).map_err(be)? {
73 Some(v) => Ok(Some(serde_json::from_slice(v.value()).map_err(se)?)),
74 None => Ok(None),
75 }
76 })
77 .await
78 }
79
80 async fn take(
81 &self,
82 create_guid: &Guid,
83 match_guid: Option<Guid>,
84 now_ms: u64,
85 ) -> Result<Option<HandleRecord>, StoreError> {
86 let db = self.db.clone();
87 let guid = *create_guid;
88 blocking(move || {
89 let wtx = db.begin_write().map_err(be)?;
90 let taken = {
91 let mut table = wtx.open_table(HANDLES).map_err(be)?;
92 let current = table.get(guid.as_slice()).map_err(be)?.map(|v| v.value().to_vec());
93 match current {
94 None => None,
95 Some(bytes) => {
96 let record: HandleRecord = serde_json::from_slice(&bytes).map_err(se)?;
97 let expired = !record.is_persistent() && record.deadline_ms <= now_ms;
98 let guid_ok = match match_guid {
99 Some(g) => record.match_guid == Some(g),
100 None => true,
101 };
102 if expired || !guid_ok {
103 if expired {
104 table.remove(guid.as_slice()).map_err(be)?;
105 }
106 None
107 } else {
108 table.remove(guid.as_slice()).map_err(be)?;
109 Some(record)
110 }
111 }
112 }
113 };
114 wtx.commit().map_err(be)?;
115 Ok(taken)
116 })
117 .await
118 }
119
120 async fn reclaim(
121 &self,
122 create_guid: &Guid,
123 owner: &str,
124 now_ms: u64,
125 ) -> Result<Option<HandleRecord>, StoreError> {
126 let db = self.db.clone();
127 let guid = *create_guid;
128 let owner = owner.to_string();
129 blocking(move || {
130 let wtx = db.begin_write().map_err(be)?;
131 let claimed = {
132 let mut table = wtx.open_table(HANDLES).map_err(be)?;
133 let current = table.get(guid.as_slice()).map_err(be)?.map(|v| v.value().to_vec());
134 match current {
135 None => None,
136 Some(bytes) => {
137 let mut record: HandleRecord = serde_json::from_slice(&bytes).map_err(se)?;
138 let expired = !record.is_persistent() && record.deadline_ms <= now_ms;
139 if !record.owner_node.is_empty() && record.owner_node != owner && !expired {
140 None
141 } else {
142 record.owner_node = owner.clone();
143 if record.timeout_ms > 0 {
144 record.deadline_ms = now_ms + record.timeout_ms;
145 }
146 let value = serde_json::to_vec(&record).map_err(se)?;
147 table.insert(guid.as_slice(), value.as_slice()).map_err(be)?;
148 Some(record)
149 }
150 }
151 }
152 };
153 wtx.commit().map_err(be)?;
154 Ok(claimed)
155 })
156 .await
157 }
158
159 async fn remove(&self, create_guid: &Guid) -> Result<(), StoreError> {
160 let db = self.db.clone();
161 let guid = *create_guid;
162 blocking(move || {
163 let wtx = db.begin_write().map_err(be)?;
164 {
165 let mut table = wtx.open_table(HANDLES).map_err(be)?;
166 table.remove(guid.as_slice()).map_err(be)?;
167 }
168 wtx.commit().map_err(be)
169 })
170 .await
171 }
172
173 async fn sweep_expired(&self, now_ms: u64) -> Result<Vec<Guid>, StoreError> {
174 let db = self.db.clone();
175 blocking(move || {
176 let wtx = db.begin_write().map_err(be)?;
177 let mut expired: Vec<Guid> = Vec::new();
178 {
179 let mut table = wtx.open_table(HANDLES).map_err(be)?;
180 {
181 let iter = table.iter().map_err(be)?;
182 for entry in iter {
183 let (_k, v) = entry.map_err(be)?;
184 let record: HandleRecord = serde_json::from_slice(v.value()).map_err(se)?;
185 if !record.is_persistent() && record.deadline_ms <= now_ms {
186 expired.push(record.create_guid);
187 }
188 }
189 }
190 for guid in &expired {
191 table.remove(guid.as_slice()).map_err(be)?;
192 }
193 }
194 wtx.commit().map_err(be)?;
195 Ok(expired)
196 })
197 .await
198 }
199
200 async fn list(&self) -> Result<Vec<HandleRecord>, StoreError> {
201 let db = self.db.clone();
202 blocking(move || {
203 let rtx = db.begin_read().map_err(be)?;
204 let table = rtx.open_table(HANDLES).map_err(be)?;
205 let mut out = Vec::new();
206 for entry in table.iter().map_err(be)? {
207 let (_k, v) = entry.map_err(be)?;
208 out.push(serde_json::from_slice(v.value()).map_err(se)?);
209 }
210 Ok(out)
211 })
212 .await
213 }
214}
215
216#[cfg(test)]
217mod redb_tests {
218 use super::*;
219 use crate::sample;
220
221 fn temp_db() -> std::path::PathBuf {
222 let nanos = std::time::SystemTime::now()
223 .duration_since(std::time::UNIX_EPOCH)
224 .unwrap()
225 .as_nanos();
226 std::env::temp_dir().join(format!("smb-hs-{}-{}.redb", std::process::id(), nanos))
227 }
228
229 #[tokio::test]
230 async fn survives_reopen() {
231 let path = temp_db();
232 let rec = sample(9, 5000);
233 {
234 let store = RedbStore::open(&path).unwrap();
235 store.put(rec.clone()).await.unwrap();
236 }
237 let store = RedbStore::open(&path).unwrap();
239 assert_eq!(store.get(&rec.create_guid).await.unwrap().as_ref(), Some(&rec));
240 let _ = std::fs::remove_file(&path);
241 }
242
243 #[tokio::test]
244 async fn reclaim_and_sweep() {
245 let path = temp_db();
246 let store = RedbStore::open(&path).unwrap();
247 store.put(sample(10, 1000)).await.unwrap();
248 let guid = [10u8; 16];
249 assert!(store.reclaim(&guid, "nodeA", 0).await.unwrap().is_some());
250 assert!(store.reclaim(&guid, "nodeB", 500).await.unwrap().is_none());
251 let dropped = store.sweep_expired(2000).await.unwrap();
252 assert_eq!(dropped, vec![guid]);
253 assert!(store.get(&guid).await.unwrap().is_none());
254 let _ = std::fs::remove_file(&path);
255 }
256}