Files
mcp-memory/server/src/store.rs
T

162 lines
4.7 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: RwLock<T>,
tx: tokio::sync::mpsc::UnboundedSender<Vec<u8>>,
}
impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self {
let initial_data = Self::load_from_db(key, &db);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<Vec<u8>>();
let db_clone = db.clone();
let key_clone = key.to_string();
tokio::spawn(async move {
while let Some(json_data) = rx.recv().await {
let db_inner = db_clone.clone();
let key_inner = key_clone.clone();
let _ = tokio::task::spawn_blocking(move || {
let write_txn = db_inner.begin_write().unwrap();
{
let mut table = write_txn.open_table(STORE_TABLE).unwrap();
table.insert(key_inner.as_str(), json_data.as_slice()).unwrap();
}
write_txn.commit().unwrap();
}).await;
}
});
Self {
cache: RwLock::new(initial_data),
tx,
}
}
fn load_from_db(key: &str, db: &Database) -> T {
let read_txn = db.begin_read().unwrap();
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(&self) -> T {
let lock = self.cache.read().unwrap();
lock.clone()
}
pub fn read_with<F, R>(&self, f: F) -> R
where
F: FnOnce(&T) -> R,
{
let lock = self.cache.read().unwrap();
f(&lock)
}
pub fn modify<F: FnOnce(&mut T)>(&self, f: F) {
let json_data = {
let mut lock = self.cache.write().unwrap();
f(&mut lock);
// Serialize while holding lock to avoid expensive deep clone of T
serde_json::to_vec(&*lock).unwrap()
};
let _ = self.tx.send(json_data);
}
}
#[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(), 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(),
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(),
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().value, 50);
}
}