579 lines
20 KiB
Rust
579 lines
20 KiB
Rust
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<String>,
|
|
pub payload: serde_json::Value,
|
|
}
|
|
|
|
pub struct ProjectStores {
|
|
pub tasks: Store<Vec<Task>>,
|
|
pub milestones: Store<Vec<Milestone>>,
|
|
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 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 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>>,
|
|
}
|
|
|
|
#[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<String>,
|
|
}
|
|
|
|
#[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<CachedClipboardImage>,
|
|
pub last_text: Option<CachedClipboardText>,
|
|
pub history: std::collections::VecDeque<ClipboardHistoryItem>,
|
|
}
|
|
|
|
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<tokio::sync::Notify>,
|
|
pub ttl_notify: Arc<tokio::sync::Notify>,
|
|
pub condense_notify: Arc<tokio::sync::Notify>,
|
|
pub shutdown_notify: Arc<tokio::sync::Notify>,
|
|
pub graph: Store<KnowledgeGraph>,
|
|
pub search_index: tokio::sync::RwLock<MemoryIndex>,
|
|
|
|
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>,
|
|
pub clipboard_cache: Arc<tokio::sync::RwLock<ClipboardCacheState>>,
|
|
}
|
|
|
|
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<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);
|
|
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<Self>) -> SearchService {
|
|
SearchService::new(self.clone())
|
|
}
|
|
|
|
pub async fn rebuild_index(self: &Arc<Self>) {
|
|
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<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);
|
|
}
|
|
}
|