use crate::models::*; use crate::search::MemoryIndex; use crate::store::Store; 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 MemoryState { pub base_dir: PathBuf, pub clipboard_watch_mode: tokio::sync::RwLock, pub graph: Store, pub search_index: RwLock, pub ledger: Store>, pub sticky: Store>, pub tasks: Store>, pub snippets: Store>, pub adrs: Store>, pub prefs: Store>, pub error_fixes: Store>, pub pinned_files: Store>, pub session_summaries: Store>, pub handoff_memos: Store>, pub env_fingerprints: Store>, pub env_requirements: Store>, pub milestones: Store>, pub environments: Store>, pub pr_checklists: Store>, pub tech_debts: Store>, pub gates: Store>, pub context_workspaces: Store>, pub recent_activities: Store>, pub terminal_history: Store>, pub activity_tx: tokio::sync::broadcast::Sender, pub event_bus_tx: tokio::sync::broadcast::Sender, } 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); Self { 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); } }), ledger: Store::new("audit_ledger", db.clone()), sticky: Store::new("sticky_notes", db.clone()), tasks: Store::new("tasks", db.clone()), snippets: Store::new("snippets", db.clone()), adrs: Store::new("adrs", db.clone()), prefs: Store::new("preferences", db.clone()), error_fixes: Store::new("error_fixes", db.clone()), pinned_files: Store::new("pinned_files", db.clone()), session_summaries: Store::new("session_summaries", db.clone()), handoff_memos: Store::new("handoff_memos", db.clone()), env_fingerprints: Store::new("env_fingerprints", db.clone()), env_requirements: Store::new("env_requirements", db.clone()), milestones: Store::new("milestones", db.clone()), environments: Store::new("environments", db.clone()), pr_checklists: Store::new("pr_checklists", db.clone()), tech_debts: Store::new("tech_debts", db.clone()), gates: Store::new("gates", db.clone()), context_workspaces: Store::new("context_workspaces", db.clone()), recent_activities: Store::new("recent_activities", db.clone()), terminal_history: Store::new("terminal_history", db.clone()), activity_tx: tokio::sync::broadcast::channel(100).0, event_bus_tx: tokio::sync::broadcast::channel(1000).0, } } pub fn broadcast_activity(&self, message: &str) { let time = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) .unwrap_or_default() .as_millis() as u64; let item = serde_json::json!({ "time": time, "message": message }); self.recent_activities.modify(|activities| { activities.push_back(item.clone()); if activities.len() > 100 { activities.pop_front(); } }); let payload = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/activity", "params": item }) .to_string(); let _ = self.activity_tx.send(payload); } 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 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.tasks.read_with(|t| t.clone()); let snippets = self.snippets.read_with(|s| s.clone()); let adrs = self.adrs.read_with(|a| a.clone()); println!( "rebuild_index: found {} entities, {} tasks", entities.len(), tasks.len() ); let idx_clone = idx.clone(); tokio::task::spawn_blocking(move || { println!("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; } } } #[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.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.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(); println!( "Index reader doc count: {}", idx.reader.searcher().num_docs() ); let all_docs = idx.search("Test", None).expect("Search failed"); println!("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"); } }