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 relation_count = graph.relations.len();
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 adr_count = adrs.len();
let tech_debts = state_clone.code.tech_debts.cache.read().unwrap();
@@ -464,10 +464,7 @@ mod tests {
async fn test_ping_endpoint() {
let (app, _, _dir) = setup_app().await;
let request = Request::builder()
.uri("/ping")
.body(Body::empty())
.unwrap();
let request = Request::builder().uri("/ping").body(Body::empty()).unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
@@ -611,12 +608,15 @@ mod tests {
.method("POST")
.uri("/gate/set")
.header("content-type", "application/json")
.body(Body::from(serde_json::json!({
.body(Body::from(
serde_json::json!({
"action": "deploy",
"target": "prod",
"authorize": true,
"reason": "Tests passed"
}).to_string()))
})
.to_string(),
))
.unwrap();
let response = app.oneshot(set_req).await.unwrap();
@@ -655,10 +655,7 @@ mod tests {
for ep in endpoints {
let app_inst = create_router(app_state.clone());
let req = Request::builder()
.uri(ep)
.body(Body::empty())
.unwrap();
let req = Request::builder().uri(ep).body(Body::empty()).unwrap();
let resp = app_inst.oneshot(req).await.unwrap();
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 resp_html = app_html.oneshot(req_html).await.unwrap();
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"));
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();
assert!(html_str.contains("<script src=\"/dashboard.js\"></script>"));
// Test GET /dashboard.js
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();
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"));
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();
assert!(js_str.contains("escapeHtml"));
assert!(js_str.contains("setupWS"));
@@ -722,15 +736,24 @@ mod tests {
// 1. Initial GET /api/activity should return empty array []
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();
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();
assert_eq!(json_str.trim(), "[]");
// 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
let app_term = create_router(app_state.clone());
@@ -738,13 +761,16 @@ mod tests {
.method("POST")
.uri("/terminal/telemetry")
.header("content-type", "application/json")
.body(Body::from(serde_json::json!({
.body(Body::from(
serde_json::json!({
"command": "cargo test --workspace",
"exit_code": 0,
"cwd": "C:\\workspace\\mcp-memory",
"os": "windows",
"timestamp": 1728130000000u64
}).to_string()))
})
.to_string(),
))
.unwrap();
let term_resp = app_term.oneshot(term_req).await.unwrap();
assert_eq!(term_resp.status(), StatusCode::OK);
@@ -755,21 +781,29 @@ mod tests {
.method("POST")
.uri("/nvim/telemetry")
.header("content-type", "application/json")
.body(Body::from(serde_json::json!({
.body(Body::from(
serde_json::json!({
"session_id": "test_session_1",
"event": "BufWritePost",
"file": "server/src/dashboard.ts"
}).to_string()))
})
.to_string(),
))
.unwrap();
let nvim_resp = app_nvim.oneshot(nvim_req).await.unwrap();
assert_eq!(nvim_resp.status(), StatusCode::OK);
// 5. GET /api/activity should return all recorded items
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();
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();
assert!(json_str2.contains("CODE_CHANGE"));
@@ -781,19 +815,18 @@ mod tests {
// 6. Verify /dashboard.js contains empty state text and stream logic
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();
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();
assert!(js_str.contains("No recent activity recorded yet."));
assert!(js_str.contains("/api/activity/stream"));
assert!(js_str.contains("parseActivityPayload"));
}
}
+27 -8
View File
@@ -14,22 +14,42 @@ pub fn init_redb(base: &Path) -> Arc<Database> {
} else {
let redb_path = base.join("mcp_store.redb");
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) {
Ok(db) => Arc::new(db),
Ok(db) => {
db_opt = Some(Arc::new(db));
break;
}
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!(
"Failed to open existing redb database at {:?}: {}. Attempting to recreate database.",
redb_path, open_err
redb_path, last_open_err
);
tracing::warn!("{}", err_msg);
match redb::Database::create(&redb_path) {
Ok(db) => Arc::new(db),
Err(create_err) => {
if std::env::var("MCP_ALLOW_TMP_FALLBACK").unwrap_or_default() == "1" {
let temp_path = std::env::temp_dir().join(format!(
"mcp_store_fallback_{}.redb",
uuid::Uuid::new_v4()
));
let temp_path = std::env::temp_dir()
.join(format!("mcp_store_fallback_{}.redb", uuid::Uuid::new_v4()));
tracing::warn!(
"CRITICAL PERSISTENCE ALERT: MCP_ALLOW_TMP_FALLBACK=1 set. Using temporary redb database {:?}. Changes will be discarded upon application exit.",
temp_path
@@ -41,13 +61,12 @@ pub fn init_redb(base: &Path) -> Arc<Database> {
} else {
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.",
redb_path, open_err, create_err
redb_path, last_open_err, create_err
);
}
}
}
}
}
} else {
match redb::Database::create(&redb_path) {
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() {
Some("all") => true,
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 {
Some(branch) => t.git_branch.is_none() || t.git_branch.as_deref() == Some(branch.as_str()),
+44 -20
View File
@@ -139,70 +139,94 @@ pub async fn condense_graph_worker(state: Arc<MemoryState>) {
.unwrap_or_default()
.as_secs();
let mut condensed_sticky_content = String::new();
state.code.sticky.modify(|notes| {
let sticky_condensation = state.code.sticky.read_with(|notes| {
if notes.len() > threshold {
notes.sort_by_key(|n| n.timestamp);
let to_remove = notes.len() - (threshold / 2);
let removed: Vec<_> = notes.drain(0..to_remove).collect();
for r in removed {
condensed_sticky_content.push_str(&format!("{}\n", r.content));
let mut sorted = notes.clone();
sorted.sort_by_key(|n| n.timestamp);
let to_remove = sorted.len() - (threshold / 2);
let removed: Vec<_> = sorted.into_iter().take(to_remove).collect();
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() {
state.modify_graph(|graph| {
if let Some((content, ids)) = sticky_condensation {
if !content.is_empty() {
let name = format!("StickyNote History {}", now);
state.modify_graph(|graph| {
graph.entities.insert(
name.clone(),
crate::models::Entity {
name: name.clone(),
entity_type: "Historical Summary".to_string(),
observations: vec![condensed_sticky_content],
observations: vec![content],
namespace: crate::models::default_namespace(),
git_branch: None,
..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.");
}
}
let mut condensed_snippet_content = String::new();
state.code.snippets.modify(|snippets| {
let snippet_condensation = state.code.snippets.read_with(|snippets| {
if snippets.len() > threshold {
snippets.sort_by_key(|s| s.updated_at);
let to_remove = snippets.len() - (threshold / 2);
let removed: Vec<_> = snippets.drain(0..to_remove).collect();
for r in removed {
condensed_snippet_content.push_str(&format!(
let mut sorted = snippets.clone();
sorted.sort_by_key(|s| s.updated_at);
let to_remove = sorted.len() - (threshold / 2);
let removed: Vec<_> = sorted.into_iter().take(to_remove).collect();
let mut content = String::new();
let mut names = Vec::new();
for r in &removed {
content.push_str(&format!(
"Name: {}\nDesc: {}\nCode: {}\n",
r.name, r.description, r.code
));
names.push(r.name.clone());
}
Some((content, names))
} else {
None
}
});
if !condensed_snippet_content.is_empty() {
state.modify_graph(|graph| {
if let Some((content, names)) = snippet_condensation {
if !content.is_empty() {
let name = format!("Snippet History {}", now);
state.modify_graph(|graph| {
graph.entities.insert(
name.clone(),
crate::models::Entity {
name: name.clone(),
entity_type: "Historical Summary".to_string(),
observations: vec![condensed_snippet_content],
observations: vec![content],
namespace: crate::models::default_namespace(),
git_branch: None,
..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.");
}
}
}
}
pub async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
let state_for_index = Arc::clone(&state);
+10 -7
View File
@@ -513,12 +513,12 @@ impl SearchService {
let mut uncached_meta = Vec::new();
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 desc = snippet.description.clone();
if let Some(ref emb) = snippet.embedding {
cached_items.push((title, "snippet".to_string(), desc, emb.clone()));
} else {
} else if uncached_texts.len() < 50 {
uncached_texts.push(format!(
"{} {} {}",
snippet.name, snippet.description, snippet.code
@@ -529,7 +529,10 @@ impl SearchService {
});
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>();
uncached_texts.push(note.content.clone());
uncached_meta.push((
@@ -540,7 +543,7 @@ impl SearchService {
}
});
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 entity.namespace != ns {
continue;
@@ -551,7 +554,7 @@ impl SearchService {
let desc = format!("{}: {}", entity.entity_type, obs);
if let Some(ref emb) = entity.embedding {
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_meta.push((title, "entity".to_string(), desc));
}
@@ -559,12 +562,12 @@ impl SearchService {
});
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 desc = fix.solution.clone();
if let Some(ref emb) = fix.embedding {
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_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 {
Insert(Vec<u8>),
Delete,
Batch {
inserts: Vec<(String, Vec<u8>)>,
deletes: Vec<String>,
},
}
/// 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)
}
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,
@@ -191,6 +227,17 @@ impl DbWriteQueue {
.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,
@@ -222,14 +269,16 @@ pub struct Store<T> {
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) = 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 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,
@@ -237,29 +286,38 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
key: key.to_string(),
queue,
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 {
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 {
return (T::default(), false);
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) => 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) => {
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);
return (T::default(), true, std::collections::HashSet::new());
}
},
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_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 entry in range {
@@ -282,6 +341,7 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
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());
@@ -293,14 +353,14 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
if found_granular {
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)) {
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> {
@@ -377,49 +437,43 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
return;
}
// Fast mutation under critical lock section, then immediately release the RwLock guard
let (old_snapshot, new_snapshot) = {
// 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());
let old = (*lock).clone();
f(&mut lock);
let new = (*lock).clone();
(old, new)
(*lock).clone()
};
// 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 granular_entries = serde_json::to_value(&new_snapshot)
.map(|val| Self::extract_granular_entries(&self.key, &val))
.unwrap_or_default();
let new_keys: std::collections::HashSet<&str> =
granular_entries.iter().map(|(k, _)| k.as_str()).collect();
let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.clone()).collect();
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());
}
}
*known = new_keys;
}
match full_bytes_res {
Ok(data) => {
// Delete removed granular entries so they don't resurrect on restart
for del_key in removed_keys {
self.queue.push_delete(del_key, self.flushed.clone());
}
// Queue granular entries
let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
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
.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()
{
tracing::warn!(
@@ -433,7 +487,7 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
handle.spawn(async move {
let _ = tokio::time::timeout(
std::time::Duration::from_secs(10),
queue.push_async(key, data, flushed),
queue.push_batch_async(key, batch_inserts, removed_keys, flushed),
)
.await;
});
@@ -460,47 +514,41 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
return;
}
let (old_snapshot, new_snapshot) = {
let new_snapshot = {
let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
let old = (*lock).clone();
f(&mut lock);
let new = (*lock).clone();
(old, new)
(*lock).clone()
};
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 granular_entries = serde_json::to_value(&new_snapshot)
.map(|val| Self::extract_granular_entries(&self.key, &val))
.unwrap_or_default();
let new_keys: std::collections::HashSet<&str> =
granular_entries.iter().map(|(k, _)| k.as_str()).collect();
let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.clone()).collect();
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());
}
}
*known = new_keys;
}
match full_bytes_res {
Ok(data) => {
for del_key in removed_keys {
self.queue
.push_delete_async(del_key, self.flushed.clone())
.await;
}
let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1);
for (g_key, g_bytes) in granular_entries {
self.queue
.push_async(g_key, g_bytes, self.flushed.clone())
.await;
batch_inserts.push((g_key, g_bytes));
}
batch_inserts.push((self.key.clone(), data));
if let Some(rx) = self
.queue
.push_async(self.key.clone(), data, self.flushed.clone())
.push_batch_async(self.key.clone(), batch_inserts, removed_keys, self.flushed.clone())
.await
{
let _ = rx.await;