From 35c802c1b81844d44c47c3d20eb85dca794ad114 Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Wed, 7 Oct 2026 22:14:53 +0100 Subject: [PATCH] perf: optimize store batch serialization, zero-clone semantic search, and lock contention --- server/src/api/ws.rs | 25 ++---- server/src/handlers/meta.rs | 71 +++------------- server/src/handlers/vision.rs | 21 ++--- server/src/search.rs | 150 ++++++++++++++++++++++++++++++++++ server/src/state.rs | 141 +------------------------------- server/src/store.rs | 116 ++++++++------------------ 6 files changed, 210 insertions(+), 314 deletions(-) diff --git a/server/src/api/ws.rs b/server/src/api/ws.rs index 935676a..2900eb9 100644 --- a/server/src/api/ws.rs +++ b/server/src/api/ws.rs @@ -105,8 +105,8 @@ pub async fn handle_socket(socket: WebSocket, state: Arc, _client_type }); let handler = Arc::clone(&state.handler); - let state_clone = Arc::clone(&state); let session_id_clone = session_id.clone(); + let response_tx = tx.clone(); let recv_task = tokio::spawn(async move { while let Some(msg_result) = receiver.next().await { @@ -121,24 +121,11 @@ pub async fn handle_socket(socket: WebSocket, state: Arc, _client_type // Process MCP request if let Some(response) = handler.handle_request(payload).await { let res_str = serde_json::to_string(&response).unwrap_or_else(|e| format!(r#"{{\"jsonrpc\":\"2.0\",\"id\":null,\"error\":{{\"code\":-32603,\"message\":\"{}\"}}}}"#, e)); - let tx_opt = state_clone - .clients - .read() - .unwrap_or_else(|e| e.into_inner()) - .get(&session_id_clone) - .cloned(); - if let Some(client_tx) = tx_opt { - if let Err(e) = client_tx.send(res_str).await { - tracing::error!( - "Failed to send response to client channel for session {}: {}", - session_id_clone, - e - ); - } - } else { - tracing::warn!( - "Could not find client_tx for session_id {} when trying to send response", - session_id_clone + if let Err(e) = response_tx.send(res_str).await { + tracing::error!( + "Failed to send response to client channel for session {}: {}", + session_id_clone, + e ); } } diff --git a/server/src/handlers/meta.rs b/server/src/handlers/meta.rs index 0cd979a..443e850 100644 --- a/server/src/handlers/meta.rs +++ b/server/src/handlers/meta.rs @@ -287,8 +287,8 @@ impl McpTool for LogCodeChangeHandler { author: req.author, session_id: req.session_id, vcs_type: detected_vcs, - revision: effective_rev.clone(), - branch: effective_branch.clone(), + revision: effective_rev, + branch: effective_branch, repository_root: req.repository_root, }); if ledger.len() > 500 { @@ -300,42 +300,9 @@ impl McpTool for LogCodeChangeHandler { &format!("Modified {}", req.file_path), Some(&description), ); - - let recon = crate::handlers::reconciliation::reconcile_commit_or_code_change( - &state, - &description, - Some(&req.file_path), - effective_rev.as_deref(), - effective_branch.as_deref(), - ) - .await; - - let mut recon_notes = Vec::new(); - if !recon.implemented_adrs.is_empty() { - recon_notes.push(format!("Implemented ADRs: {}", recon.implemented_adrs.join(", "))); - } - if !recon.resolved_tech_debts.is_empty() { - recon_notes.push(format!("Resolved TechDebt: {}", recon.resolved_tech_debts.join(", "))); - } - if !recon.completed_tasks.is_empty() { - recon_notes.push(format!("Completed Tasks: {}", recon.completed_tasks.join(", "))); - } - if !recon.unblocked_tasks.is_empty() { - recon_notes.push(format!("Unblocked Tasks: {}", recon.unblocked_tasks.join(", "))); - } - if !recon.updated_milestones.is_empty() { - recon_notes.push(format!("Updated Milestones: {}", recon.updated_milestones.join(", "))); - } - - let recon_suffix = if recon_notes.is_empty() { - String::new() - } else { - format!(" [{}]", recon_notes.join(" | ")) - }; - Ok(format!( - "Logged code change for {}: {}{}", - req.file_path, description, recon_suffix + "Logged code change for {}: {}", + req.file_path, description )) } } @@ -811,16 +778,7 @@ impl McpTool for OmniSearchHandler { .unwrap_or_default(); // Reciprocal Rank Fusion (RRF) algorithm - #[allow(dead_code)] - #[derive(Clone)] - struct MatchItem { - id: String, - doc_type: String, - title: String, - body: String, - } - - let mut rrf_scores: std::collections::HashMap = + let mut rrf_scores: std::collections::HashMap = std::collections::HashMap::new(); for (rank, (id, doc_type, title, body, _score)) in keyword_matches.into_iter().enumerate() { @@ -829,11 +787,12 @@ impl McpTool for OmniSearchHandler { id.clone(), ( score, - MatchItem { + crate::search::SearchResult { id, doc_type, title, body, + score: 0.0, }, ), ); @@ -845,20 +804,14 @@ impl McpTool for OmniSearchHandler { if let Some(existing) = rrf_scores.get_mut(&item_id) { existing.0 += score; } else { - let item = MatchItem { - id: v_match.id.clone(), - doc_type: v_match.doc_type, - title: v_match.title, - body: v_match.body, - }; - rrf_scores.insert(item_id, (score, item)); + rrf_scores.insert(item_id, (score, v_match)); } } let mut ranked_items: Vec<_> = rrf_scores.into_values().collect(); ranked_items.sort_by(|a, b| b.0.total_cmp(&a.0)); - let matches: Vec = ranked_items.into_iter().map(|(_, item)| item).collect(); + let matches: Vec = ranked_items.into_iter().map(|(_, item)| item).collect(); let kg_json = state.read_graph(|full| { let mut kg_results = serde_json::Map::new(); @@ -1203,11 +1156,7 @@ impl McpTool for GetProjectHealthHandler { let active_milestones = state.project.milestones.read_with(|milestones| { milestones .iter() - .filter(|m| { - m.namespace == req.namespace - && !m.status.eq_ignore_ascii_case("done") - && !m.status.eq_ignore_ascii_case("completed") - }) + .filter(|m| m.namespace == req.namespace && m.status != "done") .count() }); let report = serde_json::json!({ diff --git a/server/src/handlers/vision.rs b/server/src/handlers/vision.rs index aac2582..9cacbc2 100644 --- a/server/src/handlers/vision.rs +++ b/server/src/handlers/vision.rs @@ -104,9 +104,10 @@ impl McpTool for ClipboardHandler { let req: ClipboardTool = serde_json::from_value(args).map_err(|e| e.to_string())?; match req.action { ClipboardAction::Read => { - let result = - tokio::task::spawn_blocking(move || -> crate::error::Result { + let (mut out, b64_opt) = + tokio::task::spawn_blocking(move || -> crate::error::Result<(serde_json::Map, Option)> { let mut out = serde_json::Map::new(); + let mut b64_opt = None; if let Some(text) = get_native_clipboard_text() { out.insert("text".into(), json!(text)); @@ -125,7 +126,7 @@ impl McpTool for ClipboardHandler { let bytes = jpeg_bytes.into_inner(); use base64::Engine; let b64 = base64::engine::general_purpose::STANDARD.encode(&bytes); - out.insert("image_base64".into(), json!(b64)); + b64_opt = Some(b64); let cache_dir = dirs::home_dir() .unwrap_or_default() @@ -146,17 +147,12 @@ impl McpTool for ClipboardHandler { } } } - Ok(Value::Object(out)) + Ok((out, b64_opt)) }) .await .map_err(|e| crate::error::AppError::Internal(format!("Task panic: {}", e)))??; - let mut final_obj = result; - if let Some(b64) = final_obj.get("image_base64").and_then(|v| v.as_str()) { - let b64_str = b64.to_string(); - if let Some(obj) = final_obj.as_object_mut() { - obj.remove("image_base64"); - } + if let Some(b64_str) = b64_opt { if state.ollama.is_available().await && let Ok(analysis) = state .ollama @@ -165,14 +161,13 @@ impl McpTool for ClipboardHandler { &b64_str, ) .await - && let Some(obj) = final_obj.as_object_mut() { - obj.insert("image_analysis".to_string(), json!(analysis.trim())); + out.insert("image_analysis".to_string(), json!(analysis.trim())); } } state.record_activity("clipboard", "Read contents from OS clipboard", None); - Ok::(serde_json::to_string_pretty(&final_obj)?) + Ok::(serde_json::to_string_pretty(&Value::Object(out))?) } ClipboardAction::Write => { let text_opt = req.text; diff --git a/server/src/search.rs b/server/src/search.rs index 42fa3f3..fe879f5 100644 --- a/server/src/search.rs +++ b/server/src/search.rs @@ -1,4 +1,6 @@ +use crate::embedding::{cosine_similarity, generate_embedding_async, generate_embeddings_async}; use crate::models::{Adr, Entity, Snippet, Task}; +use crate::state::MemoryState; use std::sync::{Arc, Mutex}; use tantivy::schema::*; use tantivy::{Index, IndexReader, IndexWriter, ReloadPolicy, doc}; @@ -432,6 +434,154 @@ impl MemoryIndex { } } +pub struct SearchService { + state: Arc, +} + +impl SearchService { + pub fn new(state: Arc) -> Self { + Self { state } + } + + pub async fn semantic_search( + &self, + query: &str, + filter_namespace: Option<&str>, + limit: usize, + ) -> crate::error::Result> { + let query_emb = generate_embedding_async(query.to_string()) + .await + .unwrap_or_default(); + if query_emb.is_empty() { + return Ok(Vec::new()); + } + + let mut results = Vec::new(); + let mut uncached_texts = Vec::new(); + let mut uncached_meta = Vec::new(); + + self.state.code.snippets.read_with(|snips| { + for snippet in snips.iter() { + if let Some(ref emb) = snippet.embedding { + let sim = cosine_similarity(&query_emb, emb); + results.push(SearchResult { + id: snippet.name.clone(), + doc_type: "snippet".to_string(), + title: snippet.name.clone(), + body: snippet.description.clone(), + score: sim, + }); + } else if uncached_texts.len() < 50 { + uncached_texts.push(format!( + "{} {} {}", + snippet.name, snippet.description, snippet.code + )); + uncached_meta.push(( + snippet.name.clone(), + "snippet".to_string(), + snippet.description.clone(), + )); + } + } + }); + + self.state.read_graph(|graph| { + for entity in graph.entities.values() { + if let Some(ns) = filter_namespace { + if entity.namespace != ns { + continue; + } + } + let obs = entity.observations.join("; "); + let desc = format!("{}: {}", entity.entity_type, obs); + if let Some(ref emb) = entity.embedding { + let sim = cosine_similarity(&query_emb, emb); + results.push(SearchResult { + id: entity.name.clone(), + doc_type: "entity".to_string(), + title: entity.name.clone(), + body: desc, + score: sim, + }); + } else if uncached_texts.len() < 50 { + uncached_texts.push(format!("{} {} {}", entity.name, entity.entity_type, obs)); + uncached_meta.push((entity.name.clone(), "entity".to_string(), desc)); + } + } + }); + + self.state.code.error_fixes.read_with(|fixes| { + for fix in fixes.iter() { + if let Some(ref emb) = fix.embedding { + let sim = cosine_similarity(&query_emb, emb); + results.push(SearchResult { + id: fix.signature.clone(), + doc_type: "error_fix".to_string(), + title: fix.signature.clone(), + body: fix.solution.clone(), + score: sim, + }); + } else if uncached_texts.len() < 50 { + uncached_texts.push(format!("{} {}", fix.signature, fix.solution)); + uncached_meta.push(( + fix.signature.clone(), + "error_fix".to_string(), + fix.solution.clone(), + )); + } + } + }); + + if !uncached_texts.is_empty() + && let Ok(embeddings) = generate_embeddings_async(uncached_texts).await + { + for (emb, (title, doc_type, body)) in embeddings.into_iter().zip(uncached_meta) { + let sim = cosine_similarity(&query_emb, &emb); + results.push(SearchResult { + id: title.clone(), + doc_type, + title, + body, + score: sim, + }); + } + } + + results.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + results.truncate(limit); + + Ok(results) + } + + pub async fn keyword_search( + &self, + query: &str, + filter_namespace: Option<&str>, + limit: usize, + ) -> crate::error::Result> { + let idx = self.state.get_search_index().await; + let matches = idx + .search(query, filter_namespace) + .map_err(|e| crate::error::AppError::Internal(e.to_string()))?; + + let mut results = Vec::new(); + for (id, doc_type, title, body, score) in matches.into_iter().take(limit) { + results.push(SearchResult { + id, + doc_type, + title, + body, + score, + }); + } + Ok(results) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/server/src/state.rs b/server/src/state.rs index dd82061..cabfdae 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -469,143 +469,4 @@ mod tests { } } -use crate::embedding::{cosine_similarity, generate_embedding_async, generate_embeddings_async}; - -pub struct UnifiedSearchResult { - pub id: String, - pub doc_type: String, - pub title: String, - pub body: String, - pub score: f32, -} - -pub struct SearchService { - state: Arc, -} - -impl SearchService { - pub fn new(state: Arc) -> Self { - Self { state } - } - - pub async fn semantic_search( - &self, - query: &str, - filter_namespace: Option<&str>, - limit: usize, - ) -> crate::error::Result> { - let query_emb = generate_embedding_async(query.to_string()) - .await - .unwrap_or_default(); - let mut results = Vec::new(); - let mut cached_items = Vec::new(); - let mut uncached_texts = Vec::new(); - let mut uncached_meta = Vec::new(); - - self.state.code.snippets.read_with(|snips| { - 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 if uncached_texts.len() < 50 { - uncached_texts.push(format!( - "{} {} {}", - snippet.name, snippet.description, snippet.code - )); - uncached_meta.push((title, "snippet".to_string(), desc)); - } - } - }); - self.state.read_graph(|graph| { - for entity in graph.entities.values() { - if let Some(ns) = filter_namespace { - if entity.namespace != ns { - continue; - } - } - let title = entity.name.clone(); - let obs = entity.observations.join("; "); - 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 if uncached_texts.len() < 50 { - uncached_texts.push(format!("{} {} {}", entity.name, entity.entity_type, obs)); - uncached_meta.push((title, "entity".to_string(), desc)); - } - } - }); - - self.state.code.error_fixes.read_with(|fixes| { - 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 if uncached_texts.len() < 50 { - uncached_texts.push(format!("{} {}", fix.signature, fix.solution)); - uncached_meta.push((title, "error_fix".to_string(), desc)); - } - } - }); - - for (title, doc_type, body, emb) in cached_items { - let sim = cosine_similarity(&query_emb, &emb); - results.push(UnifiedSearchResult { - id: title.clone(), - doc_type, - title, - body, - score: sim, - }); - } - - if !uncached_texts.is_empty() - && let Ok(embeddings) = generate_embeddings_async(uncached_texts).await - { - for (emb, meta) in embeddings.into_iter().zip(uncached_meta) { - let sim = cosine_similarity(&query_emb, &emb); - results.push(UnifiedSearchResult { - id: meta.0.clone(), - doc_type: meta.1.clone(), - title: meta.0, - body: meta.2, - score: sim, - }); - } - } - - results.sort_by(|a, b| { - b.score - .partial_cmp(&a.score) - .unwrap_or(std::cmp::Ordering::Equal) - }); - results.truncate(limit); - - Ok(results) - } - - pub async fn keyword_search( - &self, - query: &str, - filter_namespace: Option<&str>, - limit: usize, - ) -> crate::error::Result> { - let idx = self.state.get_search_index().await; - let matches = idx - .search(query, filter_namespace) - .map_err(|e| crate::error::AppError::Internal(e.to_string()))?; - - let mut results = Vec::new(); - for (id, doc_type, title, body, score) in matches.into_iter().take(limit) { - results.push(UnifiedSearchResult { - id, - doc_type, - title, - body, - score, - }); - } - Ok(results) - } -} +pub use crate::search::{SearchResult as UnifiedSearchResult, SearchService}; diff --git a/server/src/store.rs b/server/src/store.rs index aeae672..41913b8 100644 --- a/server/src/store.rs +++ b/server/src/store.rs @@ -444,14 +444,35 @@ impl Store (*lock).clone() }; - // Expensive serialization and granular extraction run completely unblocked outside the lock - let full_bytes_res = serde_json::to_vec(&new_snapshot); - let granular_entries = full_bytes_res - .as_ref() - .ok() - .and_then(|bytes| serde_json::from_slice::(bytes).ok()) - .map(|val| Self::extract_granular_entries(&self.key, &val)) - .unwrap_or_default(); + 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 + ), + } + } + + fn prepare_batch( + &self, + new_snapshot: &T, + ) -> Result<(Vec<(String, Vec)>, Vec), serde_json::Error> + where + T: Serialize, + { + let full_bytes = serde_json::to_vec(new_snapshot)?; + let granular_entries = match serde_json::from_slice::(&full_bytes) { + Ok(val) => Self::extract_granular_entries(&self.key, &val), + Err(_) => Vec::new(), + }; let new_keys: std::collections::HashSet = granular_entries.iter().map(|(k, _)| k.clone()).collect(); @@ -469,48 +490,11 @@ impl Store *known = new_keys; } - match full_bytes_res { - Ok(data) => { - let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1); - for (g_key, g_bytes) in granular_entries { - batch_inserts.push((g_key, g_bytes)); - } - batch_inserts.push((self.key.clone(), data.clone())); + let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1); + batch_inserts.extend(granular_entries); + batch_inserts.push((self.key.clone(), full_bytes)); - if self - .queue - .push_batch( - self.key.clone(), - batch_inserts.clone(), - removed_keys.clone(), - self.flushed.clone(), - ) - .is_none() - { - tracing::warn!( - "DbWriteQueue channel full for key '{}'. Applying backpressure fallback.", - self.key - ); - let queue = self.queue.clone(); - let key = self.key.clone(); - let flushed = self.flushed.clone(); - if let Ok(handle) = tokio::runtime::Handle::try_current() { - handle.spawn(async move { - let _ = tokio::time::timeout( - std::time::Duration::from_secs(10), - queue.push_batch_async(key, batch_inserts, removed_keys, flushed), - ) - .await; - }); - } - } - } - Err(e) => tracing::error!( - "Failed to serialize memory store for key '{}': {}", - self.key, - e - ), - } + Ok((batch_inserts, removed_keys)) } pub async fn modify_async(&self, f: F) @@ -531,38 +515,8 @@ impl Store (*lock).clone() }; - let full_bytes_res = serde_json::to_vec(&new_snapshot); - let granular_entries = full_bytes_res - .as_ref() - .ok() - .and_then(|bytes| serde_json::from_slice::(bytes).ok()) - .map(|val| Self::extract_granular_entries(&self.key, &val)) - .unwrap_or_default(); - - let new_keys: std::collections::HashSet = - 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; - } - - match full_bytes_res { - Ok(data) => { - let mut batch_inserts = Vec::with_capacity(granular_entries.len() + 1); - for (g_key, g_bytes) in granular_entries { - batch_inserts.push((g_key, g_bytes)); - } - batch_inserts.push((self.key.clone(), data)); - + match self.prepare_batch(&new_snapshot) { + Ok((batch_inserts, removed_keys)) => { if let Some(rx) = self .queue .push_batch_async(