Files
mcp-memory/server/src/state.rs
T

479 lines
17 KiB
Rust

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<String>,
pub payload: serde_json::Value,
}
pub struct ProjectStores {
pub tasks: Store<Vec<Task>>,
pub milestones: Store<Vec<Milestone>>,
pub pr_checklists: Store<Vec<PrChecklistItem>>,
pub context_workspaces: Store<Vec<ContextWorkspace>>,
pub pinned_files: Store<Vec<PinnedFile>>,
pub snapshots: Store<Vec<StateSnapshot>>,
}
pub struct CodeStores {
pub ledger: Store<Vec<CodeChange>>,
pub snippets: Store<Vec<Snippet>>,
pub adrs: Store<Vec<Adr>>,
pub error_fixes: Store<Vec<ErrorFix>>,
pub tech_debts: Store<Vec<TechDebt>>,
pub sticky: Store<Vec<StickyNote>>,
pub hypotheses: Store<Vec<Hypothesis>>,
}
pub struct EnvironmentStores {
pub env_fingerprints: Store<HashMap<String, EnvFingerprint>>,
pub env_requirements: Store<Vec<EnvRequirement>>,
pub environments: Store<Vec<EnvironmentDetail>>,
pub gates: Store<Vec<GateRecord>>,
pub prefs: Store<HashMap<String, Preference>>,
}
pub struct TelemetryStores {
pub session_summaries: Store<Vec<SessionSummary>>,
pub handoff_memos: Store<Vec<HandoffMemo>>,
pub recent_activities: Store<std::collections::VecDeque<serde_json::Value>>,
pub terminal_history: Store<std::collections::VecDeque<TerminalHistory>>,
pub agent_signals: Store<Vec<AgentSignal>>,
}
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 vector_db: tokio::sync::RwLock<Option<VectorDB>>,
pub project: ProjectStores,
pub code: CodeStores,
pub env: EnvironmentStores,
pub telemetry: TelemetryStores,
pub activity_tx: tokio::sync::broadcast::Sender<String>,
pub event_bus_tx: tokio::sync::broadcast::Sender<GenericEvent>,
pub ollama: Arc<crate::ollama::OllamaClient>,
}
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<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 fn search(self: &Arc<Self>) -> SearchService {
SearchService::new(self.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.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<ActivityRecord> = 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<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");
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<MemoryState>,
}
impl SearchService {
pub fn new(state: Arc<MemoryState>) -> Self {
Self { state }
}
pub async fn semantic_search(
&self,
query: &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 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::<String>(),
));
}
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<Vec<UnifiedSearchResult>> {
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)
}
}