feat(embedding,vision): candle embeddings with offline fallback, on-demand clipboard vision capture, and concurrency audit

This commit is contained in:
Riz Ashraf committed 2026-10-07 06:36:09 +01:00
1 parent 5bd8b1587a
commit e4a0fe72df
47 files changed
+6292 -3503

No files matched your search

+189 -125
View File
@@ -1,7 +1,6 @@
use crate::models::*;
use crate::search::MemoryIndex;
use crate::store::Store;
use crate::vector_db::VectorDB;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
@@ -50,14 +49,12 @@ pub struct TelemetryStores {
pub struct MemoryState {
pub base_dir: PathBuf,
pub clipboard_watch_mode: tokio::sync::RwLock<bool>,
pub clipboard_notify: Arc<tokio::sync::Notify>,
pub index_commit_notify: Arc<tokio::sync::Notify>,
pub ttl_notify: Arc<tokio::sync::Notify>,
pub condense_notify: Arc<tokio::sync::Notify>,
pub shutdown_notify: Arc<tokio::sync::Notify>,
pub graph: Store<KnowledgeGraph>,
pub search_index: tokio::sync::RwLock<MemoryIndex>,
pub vector_db: tokio::sync::RwLock<Option<VectorDB>>,
pub project: ProjectStores,
pub code: CodeStores,
@@ -96,22 +93,21 @@ impl MemoryState {
.join(".gemini/mcp_memory/daemon_error.log");
let _ =
std::fs::write(&log_path, format!("Failed to create MemoryIndex: {}\n", e));
crate::search::MemoryIndex::new_in_ram().expect("Failed to create RAM MemoryIndex")
crate::search::MemoryIndex::new_in_ram()
.expect("Failed to create RAM MemoryIndex")
}
}
};
let state = Self {
ollama: Arc::new(crate::ollama::OllamaClient::new_from_env()),
clipboard_watch_mode: tokio::sync::RwLock::new(false),
clipboard_notify: Arc::new(tokio::sync::Notify::new()),
index_commit_notify: Arc::new(tokio::sync::Notify::new()),
ttl_notify: Arc::new(tokio::sync::Notify::new()),
condense_notify: Arc::new(tokio::sync::Notify::new()),
shutdown_notify: Arc::new(tokio::sync::Notify::new()),
graph: Store::new("knowledge_graph_master", db.clone()),
base_dir: base.clone(),
search_index: tokio::sync::RwLock::new(search_index),
vector_db: tokio::sync::RwLock::new(None),
project: ProjectStores {
tasks: Store::new("tasks", db.clone()),
@@ -155,7 +151,8 @@ impl MemoryState {
entity.entity_type = crate::models::normalize_entity_type(&entity.entity_type);
}
for relation in g.relations.iter_mut() {
relation.relation_type = crate::models::normalize_relation_type(&relation.relation_type);
relation.relation_type =
crate::models::normalize_relation_type(&relation.relation_type);
}
});
@@ -187,50 +184,64 @@ impl MemoryState {
pub async fn rebuild_index(self: &Arc<Self>) {
let is_in_memory = self.base_dir.to_str() == Some(":memory:");
let new_idx = if is_in_memory {
crate::search::MemoryIndex::new_in_ram().expect("Failed to create RAM MemoryIndex for rebuild")
} else {
match crate::search::MemoryIndex::new(&self.base_dir) {
Ok(idx) => idx,
Err(e) => {
tracing::warn!("Failed to create disk MemoryIndex for rebuild ({}), falling back to RAM", e);
crate::search::MemoryIndex::new_in_ram().expect("Failed to create RAM MemoryIndex for rebuild")
let base_dir = self.base_dir.clone();
let state_clone = Arc::clone(self);
// Offload full clone and synchronous Tantivy doc indexing off the async Tokio reactor
let new_idx = match tokio::task::spawn_blocking(move || {
let new_idx = if is_in_memory {
crate::search::MemoryIndex::new_in_ram()
.expect("Failed to create RAM MemoryIndex for rebuild")
} else {
match crate::search::MemoryIndex::new(&base_dir) {
Ok(idx) => idx,
Err(e) => {
tracing::warn!(
"Failed to create disk MemoryIndex for rebuild ({}), falling back to RAM",
e
);
crate::search::MemoryIndex::new_in_ram()
.expect("Failed to create RAM MemoryIndex for rebuild")
}
}
}
};
};
let _ = new_idx.clear();
let entities: Vec<_> = self
.graph
.read_with(|g| g.entities.values().cloned().collect());
let tasks = self.project.tasks.read_with(|t| t.clone());
let snippets = self.code.snippets.read_with(|s| s.clone());
let adrs = self.code.adrs.read_with(|a| a.clone());
let entities: Vec<_> = state_clone
.graph
.read_with(|g| g.entities.values().cloned().collect());
let tasks = state_clone.project.tasks.read_with(|t| t.clone());
let snippets = state_clone.code.snippets.read_with(|s| s.clone());
let adrs = state_clone.code.adrs.read_with(|a| a.clone());
tracing::info!(
"rebuild_index: found {} entities, {} tasks",
entities.len(),
tasks.len()
);
tracing::info!(
"rebuild_index: indexing {} entities, {} tasks synchronously in blocking thread",
entities.len(),
tasks.len()
);
let idx_clone = new_idx.clone();
tokio::task::spawn_blocking(move || {
for e in entities {
idx_clone.add_entity_sync(&e);
new_idx.add_entity_sync(&e);
}
for task in tasks {
idx_clone.add_task_sync(&task);
new_idx.add_task_sync(&task);
}
for snippet in snippets {
idx_clone.add_snippet_sync(&snippet);
new_idx.add_snippet_sync(&snippet);
}
for adr in adrs {
idx_clone.add_adr_sync(&adr);
new_idx.add_adr_sync(&adr);
}
new_idx
})
.await
.unwrap_or_else(|e| {
tracing::error!("Failed to join tantivy index rebuild thread: {}", e);
});
{
Ok(idx) => idx,
Err(e) => {
tracing::error!("Failed to join tantivy index rebuild thread: {}", e);
return;
}
};
let _ = new_idx.commit().await;
*self.search_index.write().await = new_idx;
@@ -244,22 +255,30 @@ impl MemoryState {
.as_millis() as u64;
let category_upper = category.to_uppercase();
let truncated_details = details.map(|s| {
if s.len() > 4096 {
format!("{}... [truncated]", &s[..4096])
} else {
s.to_string()
}
});
let activity = ActivityRecord {
timestamp: ts,
category: category_upper,
summary: summary.to_string(),
details: details.map(|s| s.to_string()),
details: truncated_details,
..Default::default()
};
let record_val = serde_json::to_value(&activity).unwrap_or_default();
self.telemetry.recent_activities.modify(|activities| {
activities.push_front(record_val.clone());
if activities.len() > 100 {
activities.pop_back();
}
});
if let Ok(record_val) = serde_json::to_value(&activity) {
self.telemetry.recent_activities.modify(|activities| {
activities.push_front(record_val);
if activities.len() > 100 {
activities.pop_back();
}
});
}
if self.activity_tx.receiver_count() > 0 {
let payload = serde_json::json!({
@@ -277,16 +296,26 @@ impl MemoryState {
let payload_val = serde_json::to_value(&event).unwrap_or_default();
let summary_str = format!("Task {} -> {}", event.task_id, event.status);
let details_str = payload_val.to_string();
let truncated_details = if details_str.len() > 4096 {
format!("{}... [truncated]", &details_str[..4096])
} else {
details_str
};
self.telemetry.recent_activities.modify(|activities| {
let activity = ActivityRecord {
timestamp: event.timestamp,
category: "TASK_EVENT".to_string(),
summary: summary_str,
details: Some(payload_val.to_string()),
details: Some(truncated_details),
..Default::default()
};
activities.push_front(serde_json::to_value(&activity).unwrap_or_default());
if activities.len() > 100 {
activities.pop_back();
if let Ok(act_val) = serde_json::to_value(&activity) {
activities.push_front(act_val);
if activities.len() > 100 {
activities.pop_back();
}
}
});
@@ -316,7 +345,10 @@ impl MemoryState {
let _ = self.activity_tx.send(ws_resource_notification);
}
pub fn record_terminal_history(&self, payload: TerminalHistory) {
pub fn record_terminal_history(&self, mut payload: TerminalHistory) {
if payload.command.len() > 2048 {
payload.command = format!("{}... [truncated]", &payload.command[..2048]);
}
self.telemetry.terminal_history.modify(|history| {
history.push_front(payload);
if history.len() > 100 {
@@ -353,6 +385,7 @@ mod tests {
git_branch: None,
parent_id: None,
expires_at: None,
..Default::default()
});
});
@@ -395,7 +428,11 @@ mod tests {
let mut rx = state.activity_tx.subscribe();
// 1. Record an activity with details
state.record_activity("code_change", "Refactored state.rs", Some("Updated ActivityRecord schema"));
state.record_activity(
"code_change",
"Refactored state.rs",
Some("Updated ActivityRecord schema"),
);
// Verify recent_activities store
let activities: Vec<ActivityRecord> = state.telemetry.recent_activities.read_with(|act| {
@@ -407,12 +444,19 @@ mod tests {
assert_eq!(activities.len(), 1);
assert_eq!(activities[0].category, "CODE_CHANGE");
assert_eq!(activities[0].summary, "Refactored state.rs");
assert_eq!(activities[0].details, Some("Updated ActivityRecord schema".to_string()));
assert!(activities[0].timestamp > 1_700_000_000_000, "Timestamp must be in epoch milliseconds");
assert_eq!(
activities[0].details,
Some("Updated ActivityRecord schema".to_string())
);
assert!(
activities[0].timestamp > 1_700_000_000_000,
"Timestamp must be in epoch milliseconds"
);
// Verify broadcast channel message
let broadcast_msg = rx.recv().await.expect("Expected broadcast notification");
let broadcast_val: serde_json::Value = serde_json::from_str(&broadcast_msg).expect("Valid JSON");
let broadcast_val: serde_json::Value =
serde_json::from_str(&broadcast_msg).expect("Valid JSON");
assert_eq!(broadcast_val["jsonrpc"], "2.0");
assert_eq!(broadcast_val["method"], "notifications/activity");
assert_eq!(broadcast_val["params"]["category"], "CODE_CHANGE");
@@ -420,11 +464,12 @@ mod tests {
// 2. Broadcast an activity without details
state.broadcast_activity("task", "Completed live activity fix");
let activities_updated: Vec<ActivityRecord> = state.telemetry.recent_activities.read_with(|act| {
act.iter()
.filter_map(|v| serde_json::from_value(v.clone()).ok())
.collect()
});
let activities_updated: Vec<ActivityRecord> =
state.telemetry.recent_activities.read_with(|act| {
act.iter()
.filter_map(|v| serde_json::from_value(v.clone()).ok())
.collect()
});
assert_eq!(activities_updated.len(), 2);
assert_eq!(activities_updated[0].category, "TASK");
@@ -456,90 +501,109 @@ impl SearchService {
pub async fn semantic_search(
&self,
query: &str,
_filter_namespace: Option<&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();
let mut vdb_search = false;
if let Some(vdb) = &*self.state.vector_db.read().await {
vdb_search = true;
if let Ok(search_results) = vdb.search(query_emb.clone(), limit as u64).await {
for res in search_results {
results.push(UnifiedSearchResult {
id: res.id.clone(),
doc_type: res.doc_type.clone(),
title: res.id,
body: res.text,
score: res.score,
});
self.state.code.snippets.read_with(|snips| {
for snippet in snips.iter().take(50) {
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 {
uncached_texts.push(format!(
"{} {} {}",
snippet.name, snippet.description, snippet.code
));
uncached_meta.push((title, "snippet".to_string(), desc));
}
}
}
});
if !vdb_search {
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().take(50) {
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 {
uncached_texts.push(format!("{} {} {}", snippet.name, snippet.description, snippet.code));
uncached_meta.push((title, "snippet".to_string(), desc));
self.state.code.sticky.read_with(|sticky| {
for note in sticky.iter().take(50) {
let content_preview = note.content.chars().take(200).collect::<String>();
uncached_texts.push(note.content.clone());
uncached_meta.push((
"StickyNote".to_string(),
"sticky".to_string(),
content_preview,
));
}
});
self.state.read_graph(|graph| {
for entity in graph.entities.values().take(50) {
if let Some(ns) = filter_namespace {
if entity.namespace != ns {
continue;
}
}
});
self.state.code.sticky.read_with(|sticky| {
for note in sticky.iter().take(50) {
let content_preview = note.content.chars().take(200).collect::<String>();
uncached_texts.push(note.content.clone());
uncached_meta.push(("StickyNote".to_string(), "sticky".to_string(), content_preview));
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 {
uncached_texts.push(format!("{} {} {}", entity.name, entity.entity_type, obs));
uncached_meta.push((title, "entity".to_string(), desc));
}
});
}
});
for (title, doc_type, body, emb) in cached_items {
self.state.code.error_fixes.read_with(|fixes| {
for fix in fixes.iter().take(50) {
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 {
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: title.clone(),
doc_type,
title,
body,
id: meta.0.clone(),
doc_type: meta.1.clone(),
title: meta.0,
body: meta.2,
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);
}
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results.truncate(limit);
Ok(results)
}