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

652 lines
23 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");
/// Internal write request dispatched to the single database writer actor.
enum DbOp {
Insert(Vec<u8>),
Delete,
Batch {
inserts: Vec<(String, Vec<u8>)>,
deletes: Vec<String>,
},
}
/// Internal write request dispatched to the single database writer actor.
struct DbWriteTask {
key: String,
op: DbOp,
flushed_notifier: Arc<tokio::sync::Notify>,
oneshot_tx: Option<tokio::sync::oneshot::Sender<()>>,
}
/// Shared centralized write queue actor that handles all database writes serially with micro-batching.
#[derive(Clone)]
pub struct DbWriteQueue {
tx: tokio::sync::mpsc::Sender<DbWriteTask>,
}
static QUEUE_REGISTRY: std::sync::Mutex<Option<(Arc<Database>, DbWriteQueue)>> =
std::sync::Mutex::new(None);
fn get_or_create_queue(db: Arc<Database>) -> DbWriteQueue {
let mut reg = QUEUE_REGISTRY.lock().unwrap_or_else(|e| e.into_inner());
if let Some((ref existing_db, ref queue)) = *reg
&& Arc::ptr_eq(existing_db, &db)
&& !queue.tx.is_closed()
{
return queue.clone();
}
let new_queue = DbWriteQueue::new(db.clone());
*reg = Some((db, new_queue.clone()));
new_queue
}
impl DbWriteQueue {
pub fn new(db: Arc<Database>) -> Self {
let (tx, mut rx) = tokio::sync::mpsc::channel::<DbWriteTask>(1024);
tokio::spawn(async move {
while let Some(first_task) = rx.recv().await {
let mut batch = Vec::with_capacity(100);
batch.push(first_task);
// Gold Standard Micro-batching: Drain up to 100 accumulated tasks from queue without blocking
while batch.len() < 100 {
match rx.try_recv() {
Ok(task) => batch.push(task),
Err(_) => break,
}
}
let db_inner = db.clone();
let _ = tokio::task::spawn_blocking(move || {
match db_inner.begin_write() {
Ok(write_txn) => {
if let Ok(mut table) = write_txn.open_table(STORE_TABLE) {
for task in &batch {
match &task.op {
DbOp::Insert(data) => {
if let Err(e) =
table.insert(task.key.as_str(), data.as_slice())
{
tracing::error!(
"Failed to insert key '{}' into redb: {}",
task.key,
e
);
}
}
DbOp::Delete => {
if let Err(e) = table.remove(task.key.as_str()) {
tracing::error!(
"Failed to delete key '{}' from redb: {}",
task.key,
e
);
}
}
DbOp::Batch { inserts, deletes } => {
for del_k in deletes {
if let Err(e) = table.remove(del_k.as_str()) {
tracing::error!(
"Failed to delete batch key '{}' from redb: {}",
del_k,
e
);
}
}
for (ins_k, ins_bytes) in inserts {
if let Err(e) =
table.insert(ins_k.as_str(), ins_bytes.as_slice())
{
tracing::error!(
"Failed to insert batch key '{}' into redb: {}",
ins_k,
e
);
}
}
}
}
}
}
if let Err(e) = write_txn.commit() {
tracing::error!("Failed to commit batch to redb: {}", e);
}
}
Err(e) => {
tracing::error!(
"Failed to begin write transaction on redb writer actor: {}",
e
);
}
}
// Event-driven notification to all waiting listeners for this micro-batch
for task in batch {
if let Some(oneshot) = task.oneshot_tx {
let _ = oneshot.send(());
}
task.flushed_notifier.notify_waiters();
}
})
.await;
}
});
Self { tx }
}
pub fn push(
&self,
key: String,
data: Vec<u8>,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op(key, DbOp::Insert(data), flushed_notifier)
}
pub fn push_delete(
&self,
key: String,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op(key, DbOp::Delete, flushed_notifier)
}
pub fn push_batch(
&self,
key: String,
inserts: Vec<(String, Vec<u8>)>,
deletes: Vec<String>,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op(key, DbOp::Batch { inserts, deletes }, flushed_notifier)
}
fn push_op(
&self,
key: String,
op: DbOp,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
let task = DbWriteTask {
key,
op,
flushed_notifier,
oneshot_tx: Some(oneshot_tx),
};
if let Err(e) = self.tx.try_send(task) {
match e {
tokio::sync::mpsc::error::TrySendError::Full(task) => {
let tx = self.tx.clone();
let key = task.key.clone();
tokio::spawn(async move {
if let Err(err) = tx.send(task).await {
tracing::error!(
"DbWriteQueue fallback send failed for key '{}': {}",
key,
err
);
}
});
None
}
tokio::sync::mpsc::error::TrySendError::Closed(task) => {
tracing::error!(
"DbWriteQueue channel closed; unable to persist key '{}'",
task.key
);
None
}
}
} else {
Some(oneshot_rx)
}
}
pub async fn push_async(
&self,
key: String,
data: Vec<u8>,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op_async(key, DbOp::Insert(data), flushed_notifier)
.await
}
pub async fn push_delete_async(
&self,
key: String,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op_async(key, DbOp::Delete, flushed_notifier)
.await
}
pub async fn push_batch_async(
&self,
key: String,
inserts: Vec<(String, Vec<u8>)>,
deletes: Vec<String>,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
self.push_op_async(key, DbOp::Batch { inserts, deletes }, flushed_notifier)
.await
}
async fn push_op_async(
&self,
key: String,
op: DbOp,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
let task = DbWriteTask {
key,
op,
flushed_notifier,
oneshot_tx: Some(oneshot_tx),
};
if let Err(e) = self.tx.send(task).await {
tracing::error!(
"DbWriteQueue channel closed; unable to persist key '{}'",
e.0.key
);
None
} else {
Some(oneshot_rx)
}
}
}
pub struct Store<T> {
pub cache: Arc<RwLock<T>>,
pub flushed: Arc<tokio::sync::Notify>,
key: String,
queue: DbWriteQueue,
is_corrupted: bool,
known_granular_keys: Arc<RwLock<std::collections::HashSet<String>>>,
}
impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self {
let (initial_data, is_corrupted, known_keys) = Self::load_from_db(key, &db);
let cache = Arc::new(RwLock::new(initial_data));
let flushed = Arc::new(tokio::sync::Notify::new());
let queue = get_or_create_queue(db);
let known_granular_keys = Arc::new(RwLock::new(known_keys));
Self {
cache,
flushed,
key: key.to_string(),
queue,
is_corrupted,
known_granular_keys,
}
}
fn load_from_db(key: &str, db: &Database) -> (T, bool, std::collections::HashSet<String>) {
let Ok(read_txn) = db.begin_read() else {
tracing::error!("Failed to begin read transaction for key '{}'", key);
return (T::default(), false, std::collections::HashSet::new());
};
let Ok(table) = read_txn.open_table(STORE_TABLE) else {
return (T::default(), false, std::collections::HashSet::new());
};
// 1. Check monolithic key first as the authoritative snapshot
match table.get(key) {
Ok(Some(value)) => match serde_json::from_slice::<T>(value.value()) {
Ok(parsed) => {
let mut known = std::collections::HashSet::new();
if let Ok(val) = serde_json::to_value(&parsed) {
for k in Self::extract_granular_keys(key, &val) {
known.insert(k);
}
}
return (parsed, false, known);
}
Err(e) => {
tracing::error!(
"CRITICAL: Corrupted data for key '{}' in database: {}. Quarantine mode active: state initialized to empty default without overwriting DB key.",
key,
e
);
return (T::default(), true, std::collections::HashSet::new());
}
},
Ok(None) => {}
Err(e) => {
tracing::error!("Failed to get key '{}' from store table: {}", key, e);
}
}
// 2. Granular prefix keys fallback: format!("{}:", key)
let prefix = format!("{}:", key);
let mut items_array = Vec::new();
let mut items_map = serde_json::Map::new();
let mut found_granular = false;
let mut known = std::collections::HashSet::new();
if let Ok(range) = table.range(prefix.as_str()..) {
for (k, v) in range.flatten() {
let k_str = k.value();
if !k_str.starts_with(&prefix) {
break;
}
found_granular = true;
known.insert(k_str.to_string());
if let Ok(val) = serde_json::from_slice::<serde_json::Value>(v.value()) {
let sub_key = &k_str[prefix.len()..];
items_array.push(val.clone());
items_map.insert(sub_key.to_string(), val);
}
}
}
if found_granular {
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Array(items_array)) {
return (parsed, false, known);
}
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Object(items_map)) {
return (parsed, false, known);
}
}
(T::default(), false, std::collections::HashSet::new())
}
fn extract_granular_keys(base_key: &str, val: &serde_json::Value) -> Vec<String> {
let mut keys = Vec::new();
match val {
serde_json::Value::Array(arr) => {
for (i, item) in arr.iter().enumerate() {
let sub_key = item
.get("id")
.or_else(|| item.get("name"))
.or_else(|| item.get("title"))
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_else(|| i.to_string());
keys.push(format!("{}:{}", base_key, sub_key));
}
}
serde_json::Value::Object(map) => {
for sub_key in map.keys() {
keys.push(format!("{}:{}", base_key, sub_key));
}
}
_ => {}
}
keys
}
fn extract_granular_entries(base_key: &str, val: &serde_json::Value) -> Vec<(String, Vec<u8>)> {
let mut granular = Vec::new();
match val {
serde_json::Value::Array(arr) => {
for (i, item) in arr.iter().enumerate() {
let sub_key = item
.get("id")
.or_else(|| item.get("name"))
.or_else(|| item.get("title"))
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_else(|| i.to_string());
if let Ok(item_bytes) = serde_json::to_vec(item) {
granular.push((format!("{}:{}", base_key, sub_key), item_bytes));
}
}
}
serde_json::Value::Object(map) => {
for (sub_key, item) in map {
if let Ok(item_bytes) = serde_json::to_vec(item) {
granular.push((format!("{}:{}", base_key, sub_key), item_bytes));
}
}
}
_ => {}
}
granular
}
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)
where
T: Serialize + Clone,
{
if self.is_corrupted {
tracing::error!(
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
self.key
);
return;
}
// Fast mutation under critical lock section, only ONE clone taken, then immediately release the RwLock guard
let new_snapshot = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
f(&mut lock);
(*lock).clone()
};
match self.prepare_batch(&new_snapshot) {
Ok((batch_inserts, removed_keys)) => {
self.queue.push_batch(
self.key.clone(),
batch_inserts,
removed_keys,
self.flushed.clone(),
);
}
Err(e) => tracing::error!(
"Failed to serialize memory store for key '{}': {}",
self.key,
e
),
}
}
#[allow(clippy::type_complexity)]
fn prepare_batch(
&self,
new_snapshot: &T,
) -> Result<(Vec<(String, Vec<u8>)>, Vec<String>), serde_json::Error>
where
T: Serialize,
{
let full_bytes = serde_json::to_vec(new_snapshot)?;
let granular_entries = match serde_json::from_slice::<serde_json::Value>(&full_bytes) {
Ok(val) => Self::extract_granular_entries(&self.key, &val),
Err(_) => Vec::new(),
};
let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.clone()).collect();
let mut removed_keys = Vec::new();
{
let mut known = self
.known_granular_keys
.write()
.unwrap_or_else(|e| e.into_inner());
for old_k in known.iter() {
if !new_keys.contains(old_k) {
removed_keys.push(old_k.clone());
}
}
*known = new_keys;
}
let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
batch_inserts.extend(granular_entries);
batch_inserts.push((self.key.clone(), full_bytes));
Ok((batch_inserts, removed_keys))
}
pub async fn modify_async<F: FnOnce(&mut T)>(&self, f: F)
where
T: Serialize + Clone,
{
if self.is_corrupted {
tracing::error!(
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
self.key
);
return;
}
let new_snapshot = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
f(&mut lock);
(*lock).clone()
};
match self.prepare_batch(&new_snapshot) {
Ok((batch_inserts, removed_keys)) => {
if let Some(rx) = self
.queue
.push_batch_async(
self.key.clone(),
batch_inserts,
removed_keys,
self.flushed.clone(),
)
.await
{
let _ = rx.await;
}
}
Err(e) => tracing::error!(
"Failed to serialize memory store for key '{}': {}",
self.key,
e
),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
struct TestData {
name: String,
value: i32,
}
fn create_in_memory_test_db() -> Arc<Database> {
let db = Database::builder()
.create_with_backend(redb::backends::InMemoryBackend::new())
.unwrap();
let write_txn = db.begin_write().unwrap();
{
write_txn.open_table(STORE_TABLE).unwrap();
}
write_txn.commit().unwrap();
Arc::new(db)
}
#[tokio::test]
async fn test_store_read_write() {
let db = create_in_memory_test_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;
});
// Event-driven wait for persistence completion
store.flushed.notified().await;
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 db = create_in_memory_test_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();
}
// Event-driven wait for blocking writes to flush
store.flushed.notified().await;
assert_eq!(store.read_with(|s| s.value), 50);
}
#[tokio::test]
async fn test_store_deletion_does_not_resurrect() {
let db = create_in_memory_test_db();
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
struct Item {
id: String,
name: String,
}
let store = Store::<Vec<Item>>::new("items", db.clone());
store.modify(|items| {
items.push(Item {
id: "item1".into(),
name: "First".into(),
});
items.push(Item {
id: "item2".into(),
name: "Second".into(),
});
});
store.flushed.notified().await;
// Verify both items loaded
let store_check = Store::<Vec<Item>>::new("items", db.clone());
assert_eq!(store_check.read_with(|items| items.len()), 2);
// Delete item1
store.modify(|items| {
items.retain(|i| i.id != "item1");
});
store.flushed.notified().await;
// Reload from DB into a brand new Store instance - item1 must NOT resurrect!
let store_reloaded = Store::<Vec<Item>>::new("items", db.clone());
let remaining = store_reloaded.read_with(|items| items.clone());
assert_eq!(remaining.len(), 1);
assert_eq!(remaining[0].id, "item2");
}
}