refactor: resolve critical concurrency, data loss, schema, and search audit passes

This commit is contained in:
Riz Ashraf committed 2026-10-07 07:09:28 +01:00
1 parent e4a0fe72df
commit d80915635f
7 files changed
+679 -283

No files matched your search

+66 -33
View File
@@ -34,7 +34,7 @@ pub fn create_router(app_state: Arc<AppState>) -> Router {
let entity_count = graph.entities.len(); let entity_count = graph.entities.len();
let relation_count = graph.relations.len(); let relation_count = graph.relations.len();
let tasks = state_clone.project.tasks.cache.read().unwrap(); let tasks = state_clone.project.tasks.cache.read().unwrap();
let active_tasks = tasks.iter().filter(|t| t.status != "completed" && t.status != "done").count(); let active_tasks = tasks.iter().filter(|t| t.is_active()).count();
let adrs = state_clone.code.adrs.cache.read().unwrap(); let adrs = state_clone.code.adrs.cache.read().unwrap();
let adr_count = adrs.len(); let adr_count = adrs.len();
let tech_debts = state_clone.code.tech_debts.cache.read().unwrap(); let tech_debts = state_clone.code.tech_debts.cache.read().unwrap();
@@ -464,10 +464,7 @@ mod tests {
async fn test_ping_endpoint() { async fn test_ping_endpoint() {
let (app, _, _dir) = setup_app().await; let (app, _, _dir) = setup_app().await;
let request = Request::builder() let request = Request::builder().uri("/ping").body(Body::empty()).unwrap();
.uri("/ping")
.body(Body::empty())
.unwrap();
let response = app.oneshot(request).await.unwrap(); let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
@@ -611,12 +608,15 @@ mod tests {
.method("POST") .method("POST")
.uri("/gate/set") .uri("/gate/set")
.header("content-type", "application/json") .header("content-type", "application/json")
.body(Body::from(serde_json::json!({ .body(Body::from(
serde_json::json!({
"action": "deploy", "action": "deploy",
"target": "prod", "target": "prod",
"authorize": true, "authorize": true,
"reason": "Tests passed" "reason": "Tests passed"
}).to_string())) })
.to_string(),
))
.unwrap(); .unwrap();
let response = app.oneshot(set_req).await.unwrap(); let response = app.oneshot(set_req).await.unwrap();
@@ -655,10 +655,7 @@ mod tests {
for ep in endpoints { for ep in endpoints {
let app_inst = create_router(app_state.clone()); let app_inst = create_router(app_state.clone());
let req = Request::builder() let req = Request::builder().uri(ep).body(Body::empty()).unwrap();
.uri(ep)
.body(Body::empty())
.unwrap();
let resp = app_inst.oneshot(req).await.unwrap(); let resp = app_inst.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK, "Failed endpoint: {}", ep); assert_eq!(resp.status(), StatusCode::OK, "Failed endpoint: {}", ep);
} }
@@ -693,22 +690,39 @@ mod tests {
let req_html = Request::builder().uri("/").body(Body::empty()).unwrap(); let req_html = Request::builder().uri("/").body(Body::empty()).unwrap();
let resp_html = app_html.oneshot(req_html).await.unwrap(); let resp_html = app_html.oneshot(req_html).await.unwrap();
assert_eq!(resp_html.status(), StatusCode::OK); assert_eq!(resp_html.status(), StatusCode::OK);
let content_type = resp_html.headers().get(axum::http::header::CONTENT_TYPE).unwrap().to_str().unwrap(); let content_type = resp_html
.headers()
.get(axum::http::header::CONTENT_TYPE)
.unwrap()
.to_str()
.unwrap();
assert!(content_type.contains("text/html")); assert!(content_type.contains("text/html"));
let body_bytes = axum::body::to_bytes(resp_html.into_body(), usize::MAX).await.unwrap(); let body_bytes = axum::body::to_bytes(resp_html.into_body(), usize::MAX)
.await
.unwrap();
let html_str = String::from_utf8(body_bytes.to_vec()).unwrap(); let html_str = String::from_utf8(body_bytes.to_vec()).unwrap();
assert!(html_str.contains("<script src=\"/dashboard.js\"></script>")); assert!(html_str.contains("<script src=\"/dashboard.js\"></script>"));
// Test GET /dashboard.js // Test GET /dashboard.js
let app_js = create_router(app_state.clone()); let app_js = create_router(app_state.clone());
let req_js = Request::builder().uri("/dashboard.js").body(Body::empty()).unwrap(); let req_js = Request::builder()
.uri("/dashboard.js")
.body(Body::empty())
.unwrap();
let resp_js = app_js.oneshot(req_js).await.unwrap(); let resp_js = app_js.oneshot(req_js).await.unwrap();
assert_eq!(resp_js.status(), StatusCode::OK); assert_eq!(resp_js.status(), StatusCode::OK);
let content_type_js = resp_js.headers().get(axum::http::header::CONTENT_TYPE).unwrap().to_str().unwrap(); let content_type_js = resp_js
.headers()
.get(axum::http::header::CONTENT_TYPE)
.unwrap()
.to_str()
.unwrap();
assert!(content_type_js.contains("application/javascript")); assert!(content_type_js.contains("application/javascript"));
let js_bytes = axum::body::to_bytes(resp_js.into_body(), usize::MAX).await.unwrap(); let js_bytes = axum::body::to_bytes(resp_js.into_body(), usize::MAX)
.await
.unwrap();
let js_str = String::from_utf8(js_bytes.to_vec()).unwrap(); let js_str = String::from_utf8(js_bytes.to_vec()).unwrap();
assert!(js_str.contains("escapeHtml")); assert!(js_str.contains("escapeHtml"));
assert!(js_str.contains("setupWS")); assert!(js_str.contains("setupWS"));
@@ -722,15 +736,24 @@ mod tests {
// 1. Initial GET /api/activity should return empty array [] // 1. Initial GET /api/activity should return empty array []
let app_act = create_router(app_state.clone()); let app_act = create_router(app_state.clone());
let req_act = Request::builder().uri("/api/activity").body(Body::empty()).unwrap(); let req_act = Request::builder()
.uri("/api/activity")
.body(Body::empty())
.unwrap();
let resp_act = app_act.oneshot(req_act).await.unwrap(); let resp_act = app_act.oneshot(req_act).await.unwrap();
assert_eq!(resp_act.status(), StatusCode::OK); assert_eq!(resp_act.status(), StatusCode::OK);
let body_bytes = axum::body::to_bytes(resp_act.into_body(), usize::MAX).await.unwrap(); let body_bytes = axum::body::to_bytes(resp_act.into_body(), usize::MAX)
.await
.unwrap();
let json_str = String::from_utf8(body_bytes.to_vec()).unwrap(); let json_str = String::from_utf8(body_bytes.to_vec()).unwrap();
assert_eq!(json_str.trim(), "[]"); assert_eq!(json_str.trim(), "[]");
// 2. Record a direct activity // 2. Record a direct activity
app_state.handler.state.record_activity("code_change", "Refactored live activity UI", Some("Added empty state placeholder")); app_state.handler.state.record_activity(
"code_change",
"Refactored live activity UI",
Some("Added empty state placeholder"),
);
// 3. Post terminal telemetry activity // 3. Post terminal telemetry activity
let app_term = create_router(app_state.clone()); let app_term = create_router(app_state.clone());
@@ -738,13 +761,16 @@ mod tests {
.method("POST") .method("POST")
.uri("/terminal/telemetry") .uri("/terminal/telemetry")
.header("content-type", "application/json") .header("content-type", "application/json")
.body(Body::from(serde_json::json!({ .body(Body::from(
serde_json::json!({
"command": "cargo test --workspace", "command": "cargo test --workspace",
"exit_code": 0, "exit_code": 0,
"cwd": "C:\\workspace\\mcp-memory", "cwd": "C:\\workspace\\mcp-memory",
"os": "windows", "os": "windows",
"timestamp": 1728130000000u64 "timestamp": 1728130000000u64
}).to_string())) })
.to_string(),
))
.unwrap(); .unwrap();
let term_resp = app_term.oneshot(term_req).await.unwrap(); let term_resp = app_term.oneshot(term_req).await.unwrap();
assert_eq!(term_resp.status(), StatusCode::OK); assert_eq!(term_resp.status(), StatusCode::OK);
@@ -755,21 +781,29 @@ mod tests {
.method("POST") .method("POST")
.uri("/nvim/telemetry") .uri("/nvim/telemetry")
.header("content-type", "application/json") .header("content-type", "application/json")
.body(Body::from(serde_json::json!({ .body(Body::from(
serde_json::json!({
"session_id": "test_session_1", "session_id": "test_session_1",
"event": "BufWritePost", "event": "BufWritePost",
"file": "server/src/dashboard.ts" "file": "server/src/dashboard.ts"
}).to_string())) })
.to_string(),
))
.unwrap(); .unwrap();
let nvim_resp = app_nvim.oneshot(nvim_req).await.unwrap(); let nvim_resp = app_nvim.oneshot(nvim_req).await.unwrap();
assert_eq!(nvim_resp.status(), StatusCode::OK); assert_eq!(nvim_resp.status(), StatusCode::OK);
// 5. GET /api/activity should return all recorded items // 5. GET /api/activity should return all recorded items
let app_act2 = create_router(app_state.clone()); let app_act2 = create_router(app_state.clone());
let req_act2 = Request::builder().uri("/api/activity").body(Body::empty()).unwrap(); let req_act2 = Request::builder()
.uri("/api/activity")
.body(Body::empty())
.unwrap();
let resp_act2 = app_act2.oneshot(req_act2).await.unwrap(); let resp_act2 = app_act2.oneshot(req_act2).await.unwrap();
assert_eq!(resp_act2.status(), StatusCode::OK); assert_eq!(resp_act2.status(), StatusCode::OK);
let body_bytes2 = axum::body::to_bytes(resp_act2.into_body(), usize::MAX).await.unwrap(); let body_bytes2 = axum::body::to_bytes(resp_act2.into_body(), usize::MAX)
.await
.unwrap();
let json_str2 = String::from_utf8(body_bytes2.to_vec()).unwrap(); let json_str2 = String::from_utf8(body_bytes2.to_vec()).unwrap();
assert!(json_str2.contains("CODE_CHANGE")); assert!(json_str2.contains("CODE_CHANGE"));
@@ -781,19 +815,18 @@ mod tests {
// 6. Verify /dashboard.js contains empty state text and stream logic // 6. Verify /dashboard.js contains empty state text and stream logic
let app_js = create_router(app_state.clone()); let app_js = create_router(app_state.clone());
let req_js = Request::builder().uri("/dashboard.js").body(Body::empty()).unwrap(); let req_js = Request::builder()
.uri("/dashboard.js")
.body(Body::empty())
.unwrap();
let resp_js = app_js.oneshot(req_js).await.unwrap(); let resp_js = app_js.oneshot(req_js).await.unwrap();
assert_eq!(resp_js.status(), StatusCode::OK); assert_eq!(resp_js.status(), StatusCode::OK);
let js_bytes = axum::body::to_bytes(resp_js.into_body(), usize::MAX).await.unwrap(); let js_bytes = axum::body::to_bytes(resp_js.into_body(), usize::MAX)
.await
.unwrap();
let js_str = String::from_utf8(js_bytes.to_vec()).unwrap(); let js_str = String::from_utf8(js_bytes.to_vec()).unwrap();
assert!(js_str.contains("No recent activity recorded yet.")); assert!(js_str.contains("No recent activity recorded yet."));
assert!(js_str.contains("/api/activity/stream")); assert!(js_str.contains("/api/activity/stream"));
assert!(js_str.contains("parseActivityPayload")); assert!(js_str.contains("parseActivityPayload"));
} }
} }
+27 -8
View File
@@ -14,22 +14,42 @@ pub fn init_redb(base: &Path) -> Arc<Database> {
} else { } else {
let redb_path = base.join("mcp_store.redb"); let redb_path = base.join("mcp_store.redb");
if redb_path.exists() { if redb_path.exists() {
let mut db_opt = None;
let mut last_open_err = String::new();
for attempt in 1..=3 {
match redb::Database::open(&redb_path) { match redb::Database::open(&redb_path) {
Ok(db) => Arc::new(db), Ok(db) => {
db_opt = Some(Arc::new(db));
break;
}
Err(open_err) => { Err(open_err) => {
last_open_err = open_err.to_string();
if attempt < 3 {
tracing::warn!(
"Transient lock contention opening redb at {:?} (attempt {}/3: {}). Retrying...",
redb_path,
attempt,
open_err
);
std::thread::sleep(std::time::Duration::from_millis(150));
}
}
}
}
if let Some(db) = db_opt {
db
} else {
let err_msg = format!( let err_msg = format!(
"Failed to open existing redb database at {:?}: {}. Attempting to recreate database.", "Failed to open existing redb database at {:?}: {}. Attempting to recreate database.",
redb_path, open_err redb_path, last_open_err
); );
tracing::warn!("{}", err_msg); tracing::warn!("{}", err_msg);
match redb::Database::create(&redb_path) { match redb::Database::create(&redb_path) {
Ok(db) => Arc::new(db), Ok(db) => Arc::new(db),
Err(create_err) => { Err(create_err) => {
if std::env::var("MCP_ALLOW_TMP_FALLBACK").unwrap_or_default() == "1" { if std::env::var("MCP_ALLOW_TMP_FALLBACK").unwrap_or_default() == "1" {
let temp_path = std::env::temp_dir().join(format!( let temp_path = std::env::temp_dir()
"mcp_store_fallback_{}.redb", .join(format!("mcp_store_fallback_{}.redb", uuid::Uuid::new_v4()));
uuid::Uuid::new_v4()
));
tracing::warn!( tracing::warn!(
"CRITICAL PERSISTENCE ALERT: MCP_ALLOW_TMP_FALLBACK=1 set. Using temporary redb database {:?}. Changes will be discarded upon application exit.", "CRITICAL PERSISTENCE ALERT: MCP_ALLOW_TMP_FALLBACK=1 set. Using temporary redb database {:?}. Changes will be discarded upon application exit.",
temp_path temp_path
@@ -41,13 +61,12 @@ pub fn init_redb(base: &Path) -> Arc<Database> {
} else { } else {
panic!( panic!(
"CRITICAL PERSISTENCE FAILURE: Unable to open or create primary database at {:?}: (open: {}, create: {}). To prevent silent data loss on restart, the application cannot start without accessible persistence.", "CRITICAL PERSISTENCE FAILURE: Unable to open or create primary database at {:?}: (open: {}, create: {}). To prevent silent data loss on restart, the application cannot start without accessible persistence.",
redb_path, open_err, create_err redb_path, last_open_err, create_err
); );
} }
} }
} }
} }
}
} else { } else {
match redb::Database::create(&redb_path) { match redb::Database::create(&redb_path) {
Ok(db) => Arc::new(db), Ok(db) => Arc::new(db),
File diff suppressed because it is too large. Load diff
+1 -1
View File
@@ -269,7 +269,7 @@ impl McpTool for TasksHandler {
let status_match = match req.status.as_deref() { let status_match = match req.status.as_deref() {
Some("all") => true, Some("all") => true,
Some(s) => t.status.eq_ignore_ascii_case(s), Some(s) => t.status.eq_ignore_ascii_case(s),
None => t.status != "done" && t.status != "completed", None => t.is_active(),
}; };
let branch_match = match &req.git_branch { let branch_match = match &req.git_branch {
Some(branch) => t.git_branch.is_none() || t.git_branch.as_deref() == Some(branch.as_str()), Some(branch) => t.git_branch.is_none() || t.git_branch.as_deref() == Some(branch.as_str()),
+44 -20
View File
@@ -139,69 +139,93 @@ pub async fn condense_graph_worker(state: Arc<MemoryState>) {
.unwrap_or_default() .unwrap_or_default()
.as_secs(); .as_secs();
let mut condensed_sticky_content = String::new(); let sticky_condensation = state.code.sticky.read_with(|notes| {
state.code.sticky.modify(|notes| {
if notes.len() > threshold { if notes.len() > threshold {
notes.sort_by_key(|n| n.timestamp); let mut sorted = notes.clone();
let to_remove = notes.len() - (threshold / 2); sorted.sort_by_key(|n| n.timestamp);
let removed: Vec<_> = notes.drain(0..to_remove).collect(); let to_remove = sorted.len() - (threshold / 2);
for r in removed { let removed: Vec<_> = sorted.into_iter().take(to_remove).collect();
condensed_sticky_content.push_str(&format!("{}\n", r.content)); let mut content = String::new();
let mut ids = Vec::new();
for r in &removed {
content.push_str(&format!("{}\n", r.content));
ids.push(r.id.clone());
} }
Some((content, ids))
} else {
None
} }
}); });
if !condensed_sticky_content.is_empty() { if let Some((content, ids)) = sticky_condensation {
state.modify_graph(|graph| { if !content.is_empty() {
let name = format!("StickyNote History {}", now); let name = format!("StickyNote History {}", now);
state.modify_graph(|graph| {
graph.entities.insert( graph.entities.insert(
name.clone(), name.clone(),
crate::models::Entity { crate::models::Entity {
name: name.clone(), name: name.clone(),
entity_type: "Historical Summary".to_string(), entity_type: "Historical Summary".to_string(),
observations: vec![condensed_sticky_content], observations: vec![content],
namespace: crate::models::default_namespace(), namespace: crate::models::default_namespace(),
git_branch: None, git_branch: None,
..Default::default() ..Default::default()
}, },
); );
}); });
let id_set: std::collections::HashSet<String> = ids.into_iter().collect();
state.code.sticky.modify(|notes| {
notes.retain(|n| !id_set.contains(&n.id));
});
tracing::info!("Condensed sticky notes into Historical Summary."); tracing::info!("Condensed sticky notes into Historical Summary.");
} }
}
let mut condensed_snippet_content = String::new(); let snippet_condensation = state.code.snippets.read_with(|snippets| {
state.code.snippets.modify(|snippets| {
if snippets.len() > threshold { if snippets.len() > threshold {
snippets.sort_by_key(|s| s.updated_at); let mut sorted = snippets.clone();
let to_remove = snippets.len() - (threshold / 2); sorted.sort_by_key(|s| s.updated_at);
let removed: Vec<_> = snippets.drain(0..to_remove).collect(); let to_remove = sorted.len() - (threshold / 2);
for r in removed { let removed: Vec<_> = sorted.into_iter().take(to_remove).collect();
condensed_snippet_content.push_str(&format!( let mut content = String::new();
let mut names = Vec::new();
for r in &removed {
content.push_str(&format!(
"Name: {}\nDesc: {}\nCode: {}\n", "Name: {}\nDesc: {}\nCode: {}\n",
r.name, r.description, r.code r.name, r.description, r.code
)); ));
names.push(r.name.clone());
} }
Some((content, names))
} else {
None
} }
}); });
if !condensed_snippet_content.is_empty() { if let Some((content, names)) = snippet_condensation {
state.modify_graph(|graph| { if !content.is_empty() {
let name = format!("Snippet History {}", now); let name = format!("Snippet History {}", now);
state.modify_graph(|graph| {
graph.entities.insert( graph.entities.insert(
name.clone(), name.clone(),
crate::models::Entity { crate::models::Entity {
name: name.clone(), name: name.clone(),
entity_type: "Historical Summary".to_string(), entity_type: "Historical Summary".to_string(),
observations: vec![condensed_snippet_content], observations: vec![content],
namespace: crate::models::default_namespace(), namespace: crate::models::default_namespace(),
git_branch: None, git_branch: None,
..Default::default() ..Default::default()
}, },
); );
}); });
let name_set: std::collections::HashSet<String> = names.into_iter().collect();
state.code.snippets.modify(|snippets| {
snippets.retain(|s| !name_set.contains(&s.name));
});
tracing::info!("Condensed snippets into Historical Summary."); tracing::info!("Condensed snippets into Historical Summary.");
} }
} }
}
} }
pub async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> { pub async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
+10 -7
View File
@@ -513,12 +513,12 @@ impl SearchService {
let mut uncached_meta = Vec::new(); let mut uncached_meta = Vec::new();
self.state.code.snippets.read_with(|snips| { self.state.code.snippets.read_with(|snips| {
for snippet in snips.iter().take(50) { for snippet in snips.iter() {
let title = snippet.name.clone(); let title = snippet.name.clone();
let desc = snippet.description.clone(); let desc = snippet.description.clone();
if let Some(ref emb) = snippet.embedding { if let Some(ref emb) = snippet.embedding {
cached_items.push((title, "snippet".to_string(), desc, emb.clone())); cached_items.push((title, "snippet".to_string(), desc, emb.clone()));
} else { } else if uncached_texts.len() < 50 {
uncached_texts.push(format!( uncached_texts.push(format!(
"{} {} {}", "{} {} {}",
snippet.name, snippet.description, snippet.code snippet.name, snippet.description, snippet.code
@@ -529,7 +529,10 @@ impl SearchService {
}); });
self.state.code.sticky.read_with(|sticky| { self.state.code.sticky.read_with(|sticky| {
for note in sticky.iter().take(50) { for note in sticky.iter() {
if uncached_texts.len() >= 50 {
break;
}
let content_preview = note.content.chars().take(200).collect::<String>(); let content_preview = note.content.chars().take(200).collect::<String>();
uncached_texts.push(note.content.clone()); uncached_texts.push(note.content.clone());
uncached_meta.push(( uncached_meta.push((
@@ -540,7 +543,7 @@ impl SearchService {
} }
}); });
self.state.read_graph(|graph| { self.state.read_graph(|graph| {
for entity in graph.entities.values().take(50) { for entity in graph.entities.values() {
if let Some(ns) = filter_namespace { if let Some(ns) = filter_namespace {
if entity.namespace != ns { if entity.namespace != ns {
continue; continue;
@@ -551,7 +554,7 @@ impl SearchService {
let desc = format!("{}: {}", entity.entity_type, obs); let desc = format!("{}: {}", entity.entity_type, obs);
if let Some(ref emb) = entity.embedding { if let Some(ref emb) = entity.embedding {
cached_items.push((title, "entity".to_string(), desc, emb.clone())); cached_items.push((title, "entity".to_string(), desc, emb.clone()));
} else { } else if uncached_texts.len() < 50 {
uncached_texts.push(format!("{} {} {}", entity.name, entity.entity_type, obs)); uncached_texts.push(format!("{} {} {}", entity.name, entity.entity_type, obs));
uncached_meta.push((title, "entity".to_string(), desc)); uncached_meta.push((title, "entity".to_string(), desc));
} }
@@ -559,12 +562,12 @@ impl SearchService {
}); });
self.state.code.error_fixes.read_with(|fixes| { self.state.code.error_fixes.read_with(|fixes| {
for fix in fixes.iter().take(50) { for fix in fixes.iter() {
let title = fix.signature.clone(); let title = fix.signature.clone();
let desc = fix.solution.clone(); let desc = fix.solution.clone();
if let Some(ref emb) = fix.embedding { if let Some(ref emb) = fix.embedding {
cached_items.push((title, "error_fix".to_string(), desc, emb.clone())); cached_items.push((title, "error_fix".to_string(), desc, emb.clone()));
} else { } else if uncached_texts.len() < 50 {
uncached_texts.push(format!("{} {}", fix.signature, fix.solution)); uncached_texts.push(format!("{} {}", fix.signature, fix.solution));
uncached_meta.push((title, "error_fix".to_string(), desc)); uncached_meta.push((title, "error_fix".to_string(), desc));
} }
+100 -52
View File
@@ -8,6 +8,10 @@ pub const STORE_TABLE: TableDefinition<&str, &[u8]> = TableDefinition::new("stor
enum DbOp { enum DbOp {
Insert(Vec<u8>), Insert(Vec<u8>),
Delete, Delete,
Batch {
inserts: Vec<(String, Vec<u8>)>,
deletes: Vec<String>,
},
} }
/// Internal write request dispatched to the single database writer actor. /// Internal write request dispatched to the single database writer actor.
@@ -83,6 +87,28 @@ impl DbWriteQueue {
); );
} }
} }
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
);
}
}
}
} }
} }
} }
@@ -130,6 +156,16 @@ impl DbWriteQueue {
self.push_op(key, DbOp::Delete, flushed_notifier) 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( fn push_op(
&self, &self,
key: String, key: String,
@@ -191,6 +227,17 @@ impl DbWriteQueue {
.await .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( async fn push_op_async(
&self, &self,
key: String, key: String,
@@ -222,14 +269,16 @@ pub struct Store<T> {
key: String, key: String,
queue: DbWriteQueue, queue: DbWriteQueue,
is_corrupted: bool, is_corrupted: bool,
known_granular_keys: Arc<RwLock<std::collections::HashSet<String>>>,
} }
impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T> { impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T> {
pub fn new(key: &str, db: Arc<Database>) -> Self { pub fn new(key: &str, db: Arc<Database>) -> Self {
let (initial_data, is_corrupted) = Self::load_from_db(key, &db); let (initial_data, is_corrupted, known_keys) = Self::load_from_db(key, &db);
let cache = Arc::new(RwLock::new(initial_data)); let cache = Arc::new(RwLock::new(initial_data));
let flushed = Arc::new(tokio::sync::Notify::new()); let flushed = Arc::new(tokio::sync::Notify::new());
let queue = get_or_create_queue(db); let queue = get_or_create_queue(db);
let known_granular_keys = Arc::new(RwLock::new(known_keys));
Self { Self {
cache, cache,
@@ -237,29 +286,38 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
key: key.to_string(), key: key.to_string(),
queue, queue,
is_corrupted, is_corrupted,
known_granular_keys,
} }
} }
fn load_from_db(key: &str, db: &Database) -> (T, bool) { fn load_from_db(key: &str, db: &Database) -> (T, bool, std::collections::HashSet<String>) {
let Ok(read_txn) = db.begin_read() else { let Ok(read_txn) = db.begin_read() else {
tracing::error!("Failed to begin read transaction for key '{}'", key); tracing::error!("Failed to begin read transaction for key '{}'", key);
return (T::default(), false); return (T::default(), false, std::collections::HashSet::new());
}; };
let Ok(table) = read_txn.open_table(STORE_TABLE) else { let Ok(table) = read_txn.open_table(STORE_TABLE) else {
return (T::default(), false); return (T::default(), false, std::collections::HashSet::new());
}; };
// 1. Check monolithic key first as the authoritative snapshot // 1. Check monolithic key first as the authoritative snapshot
match table.get(key) { match table.get(key) {
Ok(Some(value)) => match serde_json::from_slice::<T>(value.value()) { Ok(Some(value)) => match serde_json::from_slice::<T>(value.value()) {
Ok(parsed) => return (parsed, false), 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) => { Err(e) => {
tracing::error!( tracing::error!(
"CRITICAL: Corrupted data for key '{}' in database: {}. Quarantine mode active: state initialized to empty default without overwriting DB key.", "CRITICAL: Corrupted data for key '{}' in database: {}. Quarantine mode active: state initialized to empty default without overwriting DB key.",
key, key,
e e
); );
return (T::default(), true); return (T::default(), true, std::collections::HashSet::new());
} }
}, },
Ok(None) => {} Ok(None) => {}
@@ -273,6 +331,7 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
let mut items_array = Vec::new(); let mut items_array = Vec::new();
let mut items_map = serde_json::Map::new(); let mut items_map = serde_json::Map::new();
let mut found_granular = false; let mut found_granular = false;
let mut known = std::collections::HashSet::new();
if let Ok(range) = table.range(prefix.as_str()..) { if let Ok(range) = table.range(prefix.as_str()..) {
for entry in range { for entry in range {
@@ -282,6 +341,7 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
break; break;
} }
found_granular = true; found_granular = true;
known.insert(k_str.to_string());
if let Ok(val) = serde_json::from_slice::<serde_json::Value>(v.value()) { if let Ok(val) = serde_json::from_slice::<serde_json::Value>(v.value()) {
let sub_key = &k_str[prefix.len()..]; let sub_key = &k_str[prefix.len()..];
items_array.push(val.clone()); items_array.push(val.clone());
@@ -293,14 +353,14 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
if found_granular { if found_granular {
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Array(items_array)) { if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Array(items_array)) {
return (parsed, false); return (parsed, false, known);
} }
if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Object(items_map)) { if let Ok(parsed) = serde_json::from_value::<T>(serde_json::Value::Object(items_map)) {
return (parsed, false); return (parsed, false, known);
} }
} }
(T::default(), false) (T::default(), false, std::collections::HashSet::new())
} }
fn extract_granular_keys(base_key: &str, val: &serde_json::Value) -> Vec<String> { fn extract_granular_keys(base_key: &str, val: &serde_json::Value) -> Vec<String> {
@@ -377,49 +437,43 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
return; return;
} }
// Fast mutation under critical lock section, then immediately release the RwLock guard // Fast mutation under critical lock section, only ONE clone taken, then immediately release the RwLock guard
let (old_snapshot, new_snapshot) = { let new_snapshot = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner()); let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
let old = (*lock).clone();
f(&mut lock); f(&mut lock);
let new = (*lock).clone(); (*lock).clone()
(old, new)
}; };
// Expensive serialization and granular extraction run completely unblocked outside the lock // Expensive serialization and granular extraction run completely unblocked outside the lock
let old_keys = serde_json::to_value(&old_snapshot)
.map(|val| Self::extract_granular_keys(&self.key, &val))
.unwrap_or_default();
let full_bytes_res = serde_json::to_vec(&new_snapshot); let full_bytes_res = serde_json::to_vec(&new_snapshot);
let granular_entries = serde_json::to_value(&new_snapshot) let granular_entries = serde_json::to_value(&new_snapshot)
.map(|val| Self::extract_granular_entries(&self.key, &val)) .map(|val| Self::extract_granular_entries(&self.key, &val))
.unwrap_or_default(); .unwrap_or_default();
let new_keys: std::collections::HashSet<&str> = let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.as_str()).collect(); granular_entries.iter().map(|(k, _)| k.clone()).collect();
let mut removed_keys = Vec::new(); let mut removed_keys = Vec::new();
for old_k in &old_keys { {
if !new_keys.contains(old_k.as_str()) { 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()); removed_keys.push(old_k.clone());
} }
} }
*known = new_keys;
}
match full_bytes_res { match full_bytes_res {
Ok(data) => { Ok(data) => {
// Delete removed granular entries so they don't resurrect on restart let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
for del_key in removed_keys {
self.queue.push_delete(del_key, self.flushed.clone());
}
// Queue granular entries
for (g_key, g_bytes) in granular_entries { for (g_key, g_bytes) in granular_entries {
self.queue.push(g_key, g_bytes, self.flushed.clone()); batch_inserts.push((g_key, g_bytes));
} }
batch_inserts.push((self.key.clone(), data.clone()));
if self if self
.queue .queue
.push(self.key.clone(), data.clone(), self.flushed.clone()) .push_batch(self.key.clone(), batch_inserts.clone(), removed_keys.clone(), self.flushed.clone())
.is_none() .is_none()
{ {
tracing::warn!( tracing::warn!(
@@ -433,7 +487,7 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
handle.spawn(async move { handle.spawn(async move {
let _ = tokio::time::timeout( let _ = tokio::time::timeout(
std::time::Duration::from_secs(10), std::time::Duration::from_secs(10),
queue.push_async(key, data, flushed), queue.push_batch_async(key, batch_inserts, removed_keys, flushed),
) )
.await; .await;
}); });
@@ -460,47 +514,41 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
return; return;
} }
let (old_snapshot, new_snapshot) = { let new_snapshot = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner()); let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
let old = (*lock).clone();
f(&mut lock); f(&mut lock);
let new = (*lock).clone(); (*lock).clone()
(old, new)
}; };
let old_keys = serde_json::to_value(&old_snapshot)
.map(|val| Self::extract_granular_keys(&self.key, &val))
.unwrap_or_default();
let full_bytes_res = serde_json::to_vec(&new_snapshot); let full_bytes_res = serde_json::to_vec(&new_snapshot);
let granular_entries = serde_json::to_value(&new_snapshot) let granular_entries = serde_json::to_value(&new_snapshot)
.map(|val| Self::extract_granular_entries(&self.key, &val)) .map(|val| Self::extract_granular_entries(&self.key, &val))
.unwrap_or_default(); .unwrap_or_default();
let new_keys: std::collections::HashSet<&str> = let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.as_str()).collect(); granular_entries.iter().map(|(k, _)| k.clone()).collect();
let mut removed_keys = Vec::new(); let mut removed_keys = Vec::new();
for old_k in &old_keys { {
if !new_keys.contains(old_k.as_str()) { 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()); removed_keys.push(old_k.clone());
} }
} }
*known = new_keys;
}
match full_bytes_res { match full_bytes_res {
Ok(data) => { Ok(data) => {
for del_key in removed_keys { let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
self.queue
.push_delete_async(del_key, self.flushed.clone())
.await;
}
for (g_key, g_bytes) in granular_entries { for (g_key, g_bytes) in granular_entries {
self.queue batch_inserts.push((g_key, g_bytes));
.push_async(g_key, g_bytes, self.flushed.clone())
.await;
} }
batch_inserts.push((self.key.clone(), data));
if let Some(rx) = self if let Some(rx) = self
.queue .queue
.push_async(self.key.clone(), data, self.flushed.clone()) .push_batch_async(self.key.clone(), batch_inserts, removed_keys, self.flushed.clone())
.await .await
{ {
let _ = rx.await; let _ = rx.await;