perf: optimize store batch serialization, zero-clone semantic search, and lock contention

This commit is contained in:
Riz Ashraf committed 2026-10-07 22:14:53 +01:00
1 parent 3b08f45618
commit 35c802c1b8
6 files changed
+205 -309

No files matched your search

+2 -15
View File
@@ -105,8 +105,8 @@ pub async fn handle_socket(socket: WebSocket, state: Arc<AppState>, _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,26 +121,13 @@ pub async fn handle_socket(socket: WebSocket, state: Arc<AppState>, _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 {
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
);
}
} else {
tracing::warn!(
"Could not find client_tx for session_id {} when trying to send response",
session_id_clone
);
}
}
} else {
tracing::warn!(
+10 -61
View File
@@ -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<String, (f64, MatchItem)> =
let mut rrf_scores: std::collections::HashMap<String, (f64, crate::search::SearchResult)> =
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<MatchItem> = ranked_items.into_iter().map(|(_, item)| item).collect();
let matches: Vec<crate::search::SearchResult> = 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!({
+8 -13
View File
@@ -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<serde_json::Value> {
let (mut out, b64_opt) =
tokio::task::spawn_blocking(move || -> crate::error::Result<(serde_json::Map<String, Value>, Option<String>)> {
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::<String, crate::error::AppError>(serde_json::to_string_pretty(&final_obj)?)
Ok::<String, crate::error::AppError>(serde_json::to_string_pretty(&Value::Object(out))?)
}
ClipboardAction::Write => {
let text_opt = req.text;
+150
View File
@@ -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<MemoryState>,
}
impl SearchService {
pub fn new(state: Arc<MemoryState>) -> Self {
Self { state }
}
pub async fn semantic_search(
&self,
query: &str,
filter_namespace: Option<&str>,
limit: usize,
) -> crate::error::Result<Vec<SearchResult>> {
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<Vec<SearchResult>> {
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::*;
+1 -140
View File
@@ -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<MemoryState>,
}
impl SearchService {
pub fn new(state: Arc<MemoryState>) -> Self {
Self { state }
}
pub async fn semantic_search(
&self,
query: &str,
filter_namespace: Option<&str>,
limit: usize,
) -> crate::error::Result<Vec<UnifiedSearchResult>> {
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<Vec<UnifiedSearchResult>> {
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};
+34 -80
View File
@@ -444,14 +444,35 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
(*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::<serde_json::Value>(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<u8>)>, Vec<String>), serde_json::Error>
where
T: Serialize,
{
let full_bytes = serde_json::to_vec(new_snapshot)?;
let granular_entries = match serde_json::from_slice::<serde_json::Value>(&full_bytes) {
Ok(val) => Self::extract_granular_entries(&self.key, &val),
Err(_) => Vec::new(),
};
let new_keys: std::collections::HashSet<String> =
granular_entries.iter().map(|(k, _)| k.clone()).collect();
@@ -469,48 +490,11 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
*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()));
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<F: FnOnce(&mut T)>(&self, f: F)
@@ -531,38 +515,8 @@ impl<T: DeserializeOwned + Default + Serialize + Send + Sync + 'static> Store<T>
(*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::<serde_json::Value>(bytes).ok())
.map(|val| Self::extract_granular_entries(&self.key, &val))
.unwrap_or_default();
let new_keys: std::collections::HashSet<String> =
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(