refactor: apply 5-pass audit optimizations across mcp-memory codebase
This commit is contained in:
1 parent
924b6d09fa
commit
5bd8b1587a
43 files changed
+1866
-1658
No files matched your search
+138
-42
@@ -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![];
|
||||
|
||||
Reference in new issue
Block a user