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, RwLock}; #[derive(Clone, Debug, serde::Serialize, serde::Deserialize)] pub struct GenericEvent { pub topic: String, pub session_id: Option, pub payload: serde_json::Value, } pub struct ProjectStores { pub tasks: Store>, pub milestones: Store>, pub pr_checklists: Store>, pub context_workspaces: Store>, pub pinned_files: Store>, pub snapshots: Store>, } pub struct CodeStores { pub ledger: Store>, pub snippets: Store>, pub adrs: Store>, pub error_fixes: Store>, pub tech_debts: Store>, pub sticky: Store>, pub hypotheses: Store>, } pub struct EnvironmentStores { pub env_fingerprints: Store>, pub env_requirements: Store>, pub environments: Store>, pub gates: Store>, pub prefs: Store>, } pub struct TelemetryStores { pub session_summaries: Store>, pub handoff_memos: Store>, pub recent_activities: Store>, pub terminal_history: Store>, pub agent_signals: Store>, } pub struct MemoryState { pub base_dir: PathBuf, pub clipboard_watch_mode: tokio::sync::RwLock, pub graph: Store, pub search_index: RwLock, pub vector_db: tokio::sync::RwLock>, pub project: ProjectStores, pub code: CodeStores, pub env: EnvironmentStores, pub telemetry: TelemetryStores, pub activity_tx: tokio::sync::broadcast::Sender, pub event_bus_tx: tokio::sync::broadcast::Sender, pub ollama: Arc, } impl MemoryState { pub fn new(base_dir_str: &str) -> Self { let base = std::path::PathBuf::from(base_dir_str); std::fs::create_dir_all(&base).expect("Failed to create store dir"); let db = crate::db::init_redb(&base); let state = Self { ollama: Arc::new(crate::ollama::OllamaClient::new_from_env()), clipboard_watch_mode: tokio::sync::RwLock::new(false), graph: Store::new("knowledge_graph_master", db.clone()), base_dir: base.clone(), search_index: RwLock::new(match crate::search::MemoryIndex::new(&base) { Ok(idx) => idx, Err(e) => { let log_path = dirs::home_dir() .unwrap_or_default() .join(".gemini/mcp_memory/daemon_error.log"); let _ = std::fs::write(&log_path, format!("Failed to create MemoryIndex: {}\n", e)); std::process::exit(1); } }), vector_db: tokio::sync::RwLock::new(None), project: ProjectStores { tasks: Store::new("tasks", db.clone()), milestones: Store::new("milestones", db.clone()), pr_checklists: Store::new("pr_checklists", db.clone()), context_workspaces: Store::new("context_workspaces", db.clone()), pinned_files: Store::new("pinned_files", db.clone()), snapshots: Store::new("state_snapshots", db.clone()), }, code: CodeStores { ledger: Store::new("audit_ledger", db.clone()), snippets: Store::new("snippets", db.clone()), adrs: Store::new("adrs", db.clone()), error_fixes: Store::new("error_fixes", db.clone()), tech_debts: Store::new("tech_debts", db.clone()), sticky: Store::new("sticky_notes", db.clone()), hypotheses: Store::new("hypotheses", db.clone()), }, env: EnvironmentStores { env_fingerprints: Store::new("env_fingerprints", db.clone()), env_requirements: Store::new("env_requirements", db.clone()), environments: Store::new("environments", db.clone()), gates: Store::new("gates", db.clone()), prefs: Store::new("preferences", db.clone()), }, telemetry: TelemetryStores { session_summaries: Store::new("session_summaries", db.clone()), handoff_memos: Store::new("handoff_memos", db.clone()), recent_activities: Store::new("recent_activities", db.clone()), terminal_history: Store::new("terminal_history", db.clone()), agent_signals: Store::new("agent_signals", db.clone()), }, activity_tx: tokio::sync::broadcast::channel(100).0, event_bus_tx: tokio::sync::broadcast::channel(1000).0, }; // Normalize pre-existing graph entity and relation types state.graph.modify(|g| { for entity in g.entities.values_mut() { 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); } }); state } pub fn broadcast_activity(&self, category: &str, message: &str) { self.record_activity(category, message, None); } pub fn read_graph(&self, f: F) -> R where F: FnOnce(&KnowledgeGraph) -> R, { self.graph.read_with(f) } pub fn modify_graph(&self, update_fn: F) { self.graph.modify(update_fn); } pub fn get_search_index(&self) -> MemoryIndex { self.search_index .read() .unwrap_or_else(|e| e.into_inner()) .clone() } pub fn search(self: &Arc) -> SearchService { SearchService::new(self.clone()) } pub async fn rebuild_index(self: &Arc) { let idx = self.search_index.read().unwrap().clone(); idx.delete_all(); 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()); tracing::info!( "rebuild_index: found {} entities, {} tasks", entities.len(), tasks.len() ); let idx_clone = idx.clone(); tokio::task::spawn_blocking(move || { // tracing::info!("spawn_blocking started in rebuild_index"); for e in entities { idx_clone.add_entity_sync(&e); } for task in tasks { idx_clone.add_task_sync(&task); } for snippet in snippets { idx_clone.add_snippet_sync(&snippet); } for adr in adrs { idx_clone.add_adr_sync(&adr); } }) .await .unwrap_or_else(|e| { tracing::error!("Failed to join tantivy index rebuild thread: {}", e); }); let _ = idx.commit().await; if let Ok(mut w) = self.search_index.write() { *w = idx; } } pub fn record_activity(&self, category: &str, summary: &str, details: Option<&str>) { let ts = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) .unwrap_or_default() .as_millis() as u64; let category_upper = category.to_uppercase(); let activity = ActivityRecord { timestamp: ts, category: category_upper, summary: summary.to_string(), details: details.map(|s| s.to_string()), }; 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(); } }); let payload = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/activity", "params": activity }) .to_string(); let _ = self.activity_tx.send(payload); } pub fn record_terminal_history(&self, payload: TerminalHistory) { self.telemetry.terminal_history.modify(|history| { history.push_front(payload); if history.len() > 100 { history.pop_back(); } }); } } #[cfg(test)] mod tests { use super::*; use tempfile::tempdir; #[tokio::test] async fn test_memory_state_initialization() { let dir = tempdir().unwrap(); let state = MemoryState::new(dir.path().to_str().unwrap()); // Ensure state fields are properly initialized assert_eq!(state.base_dir, dir.path()); // Write a test value state.project.tasks.modify(|tasks| { tasks.push(Task { id: "123".to_string(), title: "Test Task".to_string(), description: "".to_string(), status: "active".to_string(), created_at: 0, updated_at: 0, dependencies: vec![], acceptance_criteria: vec![], git_branch: None, parent_id: None, expires_at: None, }); }); // Ensure it is saved state.project.tasks.read_with(|tasks| { assert_eq!(tasks.len(), 1); assert_eq!(tasks[0].id, "123"); }); // Test rebuild index let arc_state = Arc::new(state); arc_state.rebuild_index().await; // Check search index initialization let idx = arc_state.search_index.read().unwrap(); // Force reload reader to ensure it sees the commit made by rebuild_index idx.reader.reload().unwrap(); // tracing::info!( // "Index reader doc count: {}", // idx.reader.searcher().num_docs() // ); let _all_docs = idx.search("Test", None).expect("Search failed"); // tracing::info!("All docs for 'Test': {:?}", all_docs); // Verify the task added synchronously is actually searchable let results = idx.search("Test", None).expect("Search failed"); assert_eq!(results.len(), 1, "Expected exactly 1 search result"); assert_eq!( results[0].0, "123", "Expected the result to be the task we just added" ); assert_eq!(results[0].1, "task", "Expected document type to be task"); } #[tokio::test] async fn test_record_and_broadcast_activity() { let dir = tempdir().unwrap(); let state = MemoryState::new(dir.path().to_str().unwrap()); 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")); // Verify recent_activities store let activities: Vec = state.telemetry.recent_activities.read_with(|act| { act.iter() .filter_map(|v| serde_json::from_value(v.clone()).ok()) .collect() }); 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"); // 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"); assert_eq!(broadcast_val["jsonrpc"], "2.0"); assert_eq!(broadcast_val["method"], "notifications/activity"); assert_eq!(broadcast_val["params"]["category"], "CODE_CHANGE"); // 2. Broadcast an activity without details state.broadcast_activity("task", "Completed live activity fix"); let activities_updated: Vec = 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"); assert_eq!(activities_updated[0].summary, "Completed live activity fix"); assert_eq!(activities_updated[0].details, None); assert!(activities_updated[0].timestamp >= activities_updated[1].timestamp); } } 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 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, }); } } } if !vdb_search { let mut texts_to_embed = Vec::new(); let mut metadata = Vec::new(); let snippets = self.state.code.snippets.read_with(|snips| snips.clone()); for snippet in snippets { let combined = format!("{} {} {}", snippet.name, snippet.description, snippet.code); texts_to_embed.push(combined); metadata.push((snippet.name, "snippet".to_string(), snippet.description)); } let sticky = self.state.code.sticky.read_with(|s| s.clone()); for note in sticky { texts_to_embed.push(note.content.clone()); metadata.push(( "StickyNote".to_string(), "sticky".to_string(), note.content.chars().take(200).collect::(), )); } if let Ok(embeddings) = generate_embeddings_async(texts_to_embed).await { for (emb, meta) in embeddings.into_iter().zip(metadata) { 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 fn keyword_search( &self, query: &str, filter_namespace: Option<&str>, limit: usize, ) -> crate::error::Result> { let idx = self.state.get_search_index(); 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) } }