171 lines
5.2 KiB
Rust
171 lines
5.2 KiB
Rust
use redb::{Database, ReadableDatabase, TableDefinition};
|
|
use serde::{Serialize, de::DeserializeOwned};
|
|
use std::sync::{Arc, RwLock};
|
|
|
|
pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("store");
|
|
|
|
pub struct Store<T> {
|
|
pub cache: Arc<RwLock<T>>,
|
|
tx: tokio::sync::mpsc::Sender<()>,
|
|
}
|
|
|
|
impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static> Store<T> {
|
|
pub fn new(key: &str, db: Arc<Database>) -> Self {
|
|
let initial_data = Self::load_from_db(key, &db);
|
|
let cache = Arc::new(RwLock::new(initial_data));
|
|
let (tx, mut rx) = tokio::sync::mpsc::channel::<()>(1);
|
|
|
|
let db_clone = db.clone();
|
|
let key_clone = key.to_string();
|
|
let cache_clone = cache.clone();
|
|
|
|
tokio::spawn(async move {
|
|
while rx.recv().await.is_some() {
|
|
// Drain any other pending notifications so we batch writes
|
|
while rx.try_recv().is_ok() {}
|
|
|
|
let db_inner = db_clone.clone();
|
|
let key_inner = key_clone.clone();
|
|
let json_data = {
|
|
let lock = cache_clone.read().unwrap_or_else(|e| e.into_inner());
|
|
serde_json::to_vec(&*lock)
|
|
.map_err(|e| tracing::error!("Failed to serialize memory store: {}", e))
|
|
.ok()
|
|
};
|
|
|
|
if let Some(json_data) = json_data {
|
|
let _ = tokio::task::spawn_blocking(move || {
|
|
if let Ok(write_txn) = db_inner.begin_write() {
|
|
if let Ok(mut table) = write_txn.open_table(STORE_TABLE) {
|
|
let _ = table.insert(key_inner.as_str(), json_data.as_slice());
|
|
}
|
|
let _ = write_txn.commit();
|
|
}
|
|
})
|
|
.await;
|
|
}
|
|
}
|
|
});
|
|
|
|
Self { cache, tx }
|
|
}
|
|
|
|
fn load_from_db(key: &str, db: &Database) -> T {
|
|
let Ok(read_txn) = db.begin_read() else {
|
|
return T::default();
|
|
};
|
|
if let Ok(table) = read_txn.open_table(STORE_TABLE)
|
|
&& let Ok(Some(value)) = table.get(key)
|
|
&& let Ok(parsed) = serde_json::from_slice::<T>(value.value())
|
|
{
|
|
return parsed;
|
|
}
|
|
T::default()
|
|
}
|
|
|
|
pub fn read_with<F, R>(&self, f: F) -> R
|
|
where
|
|
F: FnOnce(&T) -> R,
|
|
{
|
|
let lock = self.cache.read().unwrap_or_else(|e| e.into_inner());
|
|
f(&lock)
|
|
}
|
|
|
|
pub fn modify<F: FnOnce(&mut T)>(&self, f: F) {
|
|
{
|
|
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
|
|
f(&mut lock);
|
|
}
|
|
let _ = self.tx.try_send(());
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use tempfile::NamedTempFile;
|
|
|
|
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
|
|
struct TestData {
|
|
name: String,
|
|
value: i32,
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_store_read_write() {
|
|
let temp_file = NamedTempFile::new().unwrap();
|
|
let db = Database::create(temp_file.path()).unwrap();
|
|
|
|
let write_txn = db.begin_write().unwrap();
|
|
{
|
|
write_txn.open_table(STORE_TABLE).unwrap();
|
|
}
|
|
write_txn.commit().unwrap();
|
|
|
|
let db = Arc::new(db);
|
|
let store = Store::<TestData>::new("test_key", db.clone());
|
|
|
|
assert_eq!(store.read_with(|s| s.clone()), TestData::default());
|
|
|
|
store.modify(|data| {
|
|
data.name = "Hello".to_string();
|
|
data.value = 42;
|
|
});
|
|
|
|
// Need to wait for spawn_blocking to finish
|
|
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
|
|
assert_eq!(
|
|
store.read_with(|s| s.clone()),
|
|
TestData {
|
|
name: "Hello".to_string(),
|
|
value: 42
|
|
}
|
|
);
|
|
|
|
// Load again to verify persistence
|
|
let store2 = Store::<TestData>::new("test_key", db.clone());
|
|
assert_eq!(
|
|
store2.read_with(|s| s.clone()),
|
|
TestData {
|
|
name: "Hello".to_string(),
|
|
value: 42
|
|
}
|
|
);
|
|
}
|
|
|
|
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
|
async fn test_store_concurrency() {
|
|
let temp_file = NamedTempFile::new().unwrap();
|
|
let db = Database::create(temp_file.path()).unwrap();
|
|
|
|
let write_txn = db.begin_write().unwrap();
|
|
{
|
|
write_txn.open_table(STORE_TABLE).unwrap();
|
|
}
|
|
write_txn.commit().unwrap();
|
|
|
|
let db = Arc::new(db);
|
|
let store = Arc::new(Store::<TestData>::new("concurrent_key", db.clone()));
|
|
|
|
let mut handles = vec![];
|
|
for _ in 0..50 {
|
|
let s = store.clone();
|
|
handles.push(tokio::spawn(async move {
|
|
s.modify(|data| {
|
|
data.value += 1;
|
|
});
|
|
}));
|
|
}
|
|
|
|
for h in handles {
|
|
h.await.unwrap();
|
|
}
|
|
|
|
// Wait for all blocking writes to flush
|
|
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
|
|
|
|
assert_eq!(store.read_with(|s| s.value), 50);
|
|
}
|
|
}
|