243 lines
8.5 KiB
Rust
243 lines
8.5 KiB
Rust
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<String>,
|
|
pub payload: serde_json::Value,
|
|
}
|
|
|
|
pub struct MemoryState {
|
|
pub base_dir: PathBuf,
|
|
pub clipboard_watch_mode: tokio::sync::RwLock<bool>,
|
|
pub graph: Store<KnowledgeGraph>,
|
|
pub search_index: RwLock<MemoryIndex>,
|
|
pub ledger: Store<Vec<CodeChange>>,
|
|
pub sticky: Store<Vec<StickyNote>>,
|
|
pub tasks: Store<Vec<Task>>,
|
|
pub snippets: Store<Vec<Snippet>>,
|
|
pub adrs: Store<Vec<Adr>>,
|
|
pub prefs: Store<HashMap<String, Preference>>,
|
|
pub error_fixes: Store<Vec<ErrorFix>>,
|
|
pub pinned_files: Store<Vec<PinnedFile>>,
|
|
pub session_summaries: Store<Vec<SessionSummary>>,
|
|
pub handoff_memos: Store<Vec<HandoffMemo>>,
|
|
pub env_fingerprints: Store<HashMap<String, EnvFingerprint>>,
|
|
pub env_requirements: Store<Vec<EnvRequirement>>,
|
|
pub milestones: Store<Vec<Milestone>>,
|
|
pub environments: Store<Vec<EnvironmentDetail>>,
|
|
pub pr_checklists: Store<Vec<PrChecklistItem>>,
|
|
pub tech_debts: Store<Vec<TechDebt>>,
|
|
pub gates: Store<Vec<GateRecord>>,
|
|
pub context_workspaces: Store<Vec<ContextWorkspace>>,
|
|
pub recent_activities: Store<std::collections::VecDeque<serde_json::Value>>,
|
|
pub terminal_history: Store<std::collections::VecDeque<TerminalHistory>>,
|
|
pub activity_tx: tokio::sync::broadcast::Sender<String>,
|
|
pub event_bus_tx: tokio::sync::broadcast::Sender<GenericEvent>,
|
|
}
|
|
|
|
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<F, R>(&self, f: F) -> R
|
|
where
|
|
F: FnOnce(&KnowledgeGraph) -> R,
|
|
{
|
|
self.graph.read_with(f)
|
|
}
|
|
|
|
pub fn modify_graph<F: FnOnce(&mut KnowledgeGraph)>(&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<Self>) {
|
|
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");
|
|
}
|
|
}
|