Skip to main content

smb_server_handle_store/
redb_store.rs

1//! Embedded single-node durable backend for [`HandleStore`], built on `redb`.
2//!
3//! Handle records survive a process restart, which is what the single-node
4//! persistent-handle reclaim path ([MS-SMB2] ยง3.3.5.9.7) needs. redb calls are
5//! blocking, so they run on the blocking pool to keep the async reactor free.
6
7use 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
17/// Durable handle store backed by an embedded redb database file.
18pub struct RedbStore {
19    db: Arc<Database>,
20}
21
22impl RedbStore {
23    /// Open (creating if needed) the database at `path`.
24    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
41/// Run a blocking redb closure off the async reactor.
42async 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        // Re-open the same file: the record must still be there.
238        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}