use crate::models::*; use crate::search::MemoryIndex; pub use crate::search::{SearchResult as UnifiedSearchResult, SearchService}; use crate::store::Store; use std::collections::HashMap; use std::path::PathBuf; use std::sync::Arc; #[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 snapshots: Store>, } pub struct CodeStores { pub ledger: Store>, pub snippets: Store>, pub adrs: Store>, pub error_fixes: Store>, pub tech_debts: Store>, pub hypotheses: Store>, } pub struct EnvironmentStores { pub env_fingerprints: Store>, pub env_requirements: Store>, pub environments: Store>, pub gates: Store>, } pub struct TelemetryStores { pub session_summaries: Store>, pub handoff_memos: Store>, pub recent_activities: Store>, pub terminal_history: Store>, pub agent_signals: Store>, } #[derive(Clone, Debug, serde::Serialize, serde::Deserialize)] pub struct CachedClipboardImage { pub file_path: String, pub file_path_wsl: String, pub captured_at_epoch_ms: u64, pub age: String, pub width: u32, pub height: u32, pub size_bytes: usize, pub ocr_text: Option, } #[derive(Clone, Debug, serde::Serialize, serde::Deserialize)] pub struct CachedClipboardText { pub text: String, pub captured_at_epoch_ms: u64, pub age: String, pub char_count: usize, } #[derive(Clone, Debug, serde::Serialize, serde::Deserialize)] #[serde(tag = "kind", rename_all = "snake_case")] pub enum ClipboardHistoryItem { Image(CachedClipboardImage), Text(CachedClipboardText), } #[derive(Default)] pub struct ClipboardCacheState { pub last_image: Option, pub last_text: Option, pub history: std::collections::VecDeque, } impl ClipboardCacheState { pub fn push_image(&mut self, image: CachedClipboardImage) { self.last_image = Some(image.clone()); self.history.push_front(ClipboardHistoryItem::Image(image)); if self.history.len() > 20 { self.history.pop_back(); } } pub fn push_text(&mut self, text: CachedClipboardText) { if let Some(ref prev) = self.last_text && prev.text == text.text { return; } self.last_text = Some(text.clone()); self.history.push_front(ClipboardHistoryItem::Text(text)); if self.history.len() > 20 { self.history.pop_back(); } } } pub fn format_age(epoch_ms: u64) -> String { let now_ms = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) .unwrap_or_default() .as_millis() as u64; let diff_secs = now_ms.saturating_sub(epoch_ms) / 1000; if diff_secs < 60 { format!("{}s ago", diff_secs) } else if diff_secs < 3600 { format!("{}m ago", diff_secs / 60) } else { format!("{}h ago", diff_secs / 3600) } } pub struct MemoryState { pub base_dir: PathBuf, pub index_commit_notify: Arc, pub ttl_notify: Arc, pub condense_notify: Arc, pub shutdown_notify: Arc, pub graph: Store, pub search_index: 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, pub clipboard_cache: Arc>, } impl MemoryState { pub fn new_in_memory() -> Self { Self::new(":memory:") } pub fn new(base_dir_str: &str) -> Self { let is_in_memory = base_dir_str == ":memory:"; let base = std::path::PathBuf::from(base_dir_str); if !is_in_memory && let Err(e) = std::fs::create_dir_all(&base) { tracing::error!("Failed to create store directory at {:?}: {}", base, e); } let db = crate::db::init_redb(&base); let search_index = if is_in_memory { crate::search::MemoryIndex::new_in_ram().expect("Failed to create RAM MemoryIndex") } else { 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)); crate::search::MemoryIndex::new_in_ram() .expect("Failed to create RAM MemoryIndex") } } }; let state = Self { ollama: Arc::new(crate::ollama::OllamaClient::new_from_env()), 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), project: ProjectStores { tasks: Store::new("tasks", db.clone()), milestones: Store::new("milestones", 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()), 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()), }, 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, clipboard_cache: Arc::new(tokio::sync::RwLock::new(ClipboardCacheState::default())), }; // 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); let threshold: usize = std::env::var("MCP_MEMORY_CONDENSE_THRESHOLD") .unwrap_or_else(|_| "100".to_string()) .parse() .unwrap_or(100); let count = self.graph.read_with(|g| g.entities.len()); if count >= threshold { self.condense_notify.notify_one(); } let _ = self.event_bus_tx.send(GenericEvent { topic: "resource:updated".to_string(), session_id: None, payload: serde_json::json!({ "uri": "memory://graph" }), }); } pub async fn get_search_index(&self) -> MemoryIndex { self.search_index.read().await.clone() } pub fn search(self: &Arc) -> SearchService { SearchService::new(self.clone()) } pub async fn rebuild_index(self: &Arc) { let is_in_memory = self.base_dir.to_str() == Some(":memory:"); 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<_> = 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: indexing {} entities, {} tasks synchronously in blocking thread", entities.len(), tasks.len() ); for e in entities { new_idx.add_entity_sync(&e); } for task in tasks { new_idx.add_task_sync(&task); } for snippet in snippets { new_idx.add_snippet_sync(&snippet); } for adr in adrs { new_idx.add_adr_sync(&adr); } new_idx }) .await { 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; self.index_commit_notify.notify_waiters(); } 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 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: truncated_details, ..Default::default() }; 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!({ "jsonrpc": "2.0", "method": "notifications/activity", "params": activity }) .to_string(); let _ = self.activity_tx.send(payload); } } pub fn broadcast_task_event(&self, event: TaskEvent) { let mut payload_val = serde_json::to_value(&event).unwrap_or_default(); if let Some(obj) = payload_val.as_object_mut() { let active_task = self.project.tasks.read_with(|tasks| { tasks .iter() .find(|t| t.is_active()) .or_else(|| { tasks .iter() .find(|t| t.status == "pending" && t.parent_id.is_none()) }) .map(|t| t.title.clone()) }); obj.insert( "active_task".to_string(), serde_json::to_value(active_task).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(truncated_details), ..Default::default() }; if let Ok(act_val) = serde_json::to_value(&activity) { activities.push_front(act_val); if activities.len() > 100 { activities.pop_back(); } } }); let generic_ev = GenericEvent { topic: "task:event".to_string(), session_id: event.session_id.clone(), payload: payload_val, }; let _ = self.event_bus_tx.send(generic_ev); let _ = self.event_bus_tx.send(GenericEvent { topic: "resource:updated".to_string(), session_id: None, payload: serde_json::json!({ "uri": "memory://tasks/active" }), }); let ws_notification = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/task/completed", "params": event }) .to_string(); let _ = self.activity_tx.send(ws_notification); let ws_resource_notification = serde_json::json!({ "jsonrpc": "2.0", "method": "notifications/resources/updated", "params": { "uri": "memory://tasks/active" } }) .to_string(); let _ = self.activity_tx.send(ws_resource_notification); } 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 { 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, ..Default::default() }); }); // 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().await; // 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); } }