refactor: apply 5-pass audit optimizations across mcp-memory codebase

This commit is contained in:
Riz Ashraf committed 2026-10-06 06:05:38 +01:00
1 parent 924b6d09fa
commit 5bd8b1587a
43 files changed
+1866 -1658

No files matched your search

+138 -42
View File
@@ -8,7 +8,8 @@ pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("stor
struct DbWriteTask {
key: String,
data: Vec<u8>,
flushed: Arc<tokio::sync::Notify>,
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.
@@ -33,7 +34,7 @@ fn get_or_create_queue(db: Arc<Database>) -> DbWriteQueue {
impl DbWriteQueue {
pub fn new(db: Arc<Database>) -> Self {
let (tx, mut rx) = tokio::sync::mpsc::channel::<DbWriteTask>(2048);
let (tx, mut rx) = tokio::sync::mpsc::channel::<DbWriteTask>(1024);
tokio::spawn(async move {
while let Some(first_task) = rx.recv().await {
@@ -69,7 +70,10 @@ impl DbWriteQueue {
// Event-driven notification to all waiting listeners for this micro-batch
for task in batch {
task.flushed.notify_waiters();
if let Some(oneshot) = task.oneshot_tx {
let _ = oneshot.send(());
}
task.flushed_notifier.notify_waiters();
}
})
.await;
@@ -79,14 +83,46 @@ impl DbWriteQueue {
Self { tx }
}
pub fn push(&self, key: String, data: Vec<u8>, flushed: Arc<tokio::sync::Notify>) {
let task = DbWriteTask { key, data, flushed };
pub fn push(
&self,
key: String,
data: Vec<u8>,
flushed_notifier: Arc<tokio::sync::Notify>,
) -> Option<tokio::sync::oneshot::Receiver<()>> {
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
let task = DbWriteTask {
key,
data,
flushed_notifier,
oneshot_tx: Some(oneshot_tx),
};
if let Err(e) = self.tx.try_send(task) {
let task = e.into_inner();
let tx = self.tx.clone();
tokio::spawn(async move {
let _ = tx.send(task).await;
});
let key = e.into_inner().key;
tracing::error!("DbWriteQueue channel full or closed; unable to persist key '{}'", 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<()>> {
let (oneshot_tx, oneshot_rx) = tokio::sync::oneshot::channel();
let task = DbWriteTask {
key,
data,
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)
}
}
}
@@ -96,11 +132,12 @@ pub struct Store<T> {
pub flushed: Arc<tokio::sync::Notify>,
key: String,
queue: DbWriteQueue,
is_corrupted: bool,
}
impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static> Store<T> {
impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self {
let initial_data = Self::load_from_db(key, &db);
let (initial_data, is_corrupted) = 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);
@@ -110,20 +147,38 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static>
flushed,
key: key.to_string(),
queue,
is_corrupted,
}
}
fn load_from_db(key: &str, db: &Database) -> T {
fn load_from_db(key: &str, db: &Database) -> (T, bool) {
let Ok(read_txn) = db.begin_read() else {
return T::default();
tracing::error!("Failed to begin read transaction for key '{}'", key);
return (T::default(), false);
};
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;
match read_txn.open_table(STORE_TABLE) {
Ok(table) => match table.get(key) {
Ok(Some(value)) => match serde_json::from_slice::<T>(value.value()) {
Ok(parsed) => (parsed, false),
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
);
(T::default(), true)
}
},
Ok(None) => (T::default(), false),
Err(e) => {
tracing::error!("Failed to get key '{}' from store table: {}", key, e);
(T::default(), false)
}
},
Err(e) => {
tracing::error!("Failed to open STORE_TABLE for key '{}': {}", key, e);
(T::default(), false)
}
}
T::default()
}
pub fn read_with<F, R>(&self, f: F) -> R
@@ -134,15 +189,63 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static>
f(&lock)
}
pub fn modify<F: FnOnce(&mut T)>(&self, f: F) {
let cloned_data = {
pub fn modify<F: FnOnce(&mut T)>(&self, f: F)
where
T: Serialize,
{
if self.is_corrupted {
tracing::error!(
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
self.key
);
return;
}
let serialized_res = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
f(&mut lock);
lock.clone()
serde_json::to_vec(&*lock)
};
match serde_json::to_vec(&cloned_data) {
Ok(data) => self.queue.push(self.key.clone(), data, self.flushed.clone()),
match serialized_res {
Ok(data) => {
if self.queue.push(self.key.clone(), data.clone(), self.flushed.clone()).is_none() {
let queue = self.queue.clone();
let key = self.key.clone();
let flushed = self.flushed.clone();
tokio::spawn(async move {
let _ = queue.push_async(key, data, flushed).await;
});
}
}
Err(e) => tracing::error!("Failed to serialize memory store for key '{}': {}", self.key, e),
}
}
pub async fn modify_async<F: FnOnce(&mut T)>(&self, f: F)
where
T: Serialize,
{
if self.is_corrupted {
tracing::error!(
"CRITICAL: Refusing to persist changes for corrupted store key '{}' to prevent data loss.",
self.key
);
return;
}
let serialized_res = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
f(&mut lock);
serde_json::to_vec(&*lock)
};
match serialized_res {
Ok(data) => {
if let Some(rx) = self.queue.push_async(self.key.clone(), data, self.flushed.clone()).await {
let _ = rx.await;
}
}
Err(e) => tracing::error!("Failed to serialize memory store for key '{}': {}", self.key, e),
}
}
@@ -151,7 +254,6 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static>
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[derive(Serialize, serde::Deserialize, Clone, Default, PartialEq, Debug)]
struct TestData {
@@ -159,18 +261,21 @@ mod tests {
value: i32,
}
#[tokio::test]
async fn test_store_read_write() {
let temp_file = NamedTempFile::new().unwrap();
let db = Database::create(temp_file.path()).unwrap();
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)
}
let db = 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());
@@ -195,16 +300,7 @@ mod tests {
#[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 db = create_in_memory_test_db();
let store = Arc::new(Store::<TestData>::new("concurrent_key", db.clone()));
let mut handles = vec![];