use anyhow::Result;
use std::collections::{HashMap, HashSet};
use std::sync::{Arc, RwLock};
use tracing::{debug, info};
use super::temporal::TemporalRecord;
use crate::analysis::vector_store::{EmbeddingProvider, TfIdfEmbeddingProvider, VectorIndex};
fn record_text(record: &TemporalRecord) -> String {
let mut text = record.content.summary.clone();
for fact in &record.content.key_facts {
text.push(' ');
text.push_str(fact);
}
for entity in &record.content.entities {
text.push(' ');
text.push_str(entity);
}
for action in &record.content.actions {
text.push(' ');
text.push_str(action);
}
for outcome in &record.content.outcomes {
text.push(' ');
text.push_str(outcome);
}
for insight in &record.content.insights {
text.push(' ');
text.push_str(insight);
}
text
}
pub struct LongTermStore {
collection_name: String,
storage_dir: std::path::PathBuf,
embedder: Arc<dyn EmbeddingProvider>,
embeddings: RwLock<HashMap<String, Vec<f32>>>,
records: RwLock<HashMap<String, TemporalRecord>>,
}
impl LongTermStore {
pub fn new(storage_dir: std::path::PathBuf) -> Self {
Self::with_embedder(storage_dir, Arc::new(TfIdfEmbeddingProvider::default()))
}
pub fn with_embedder(
storage_dir: std::path::PathBuf,
embedder: Arc<dyn EmbeddingProvider>,
) -> Self {
Self {
collection_name: "consolidated".into(),
storage_dir,
embedder,
embeddings: RwLock::new(HashMap::new()),
records: RwLock::new(HashMap::new()),
}
}
pub async fn store(&self, records: &[TemporalRecord]) -> Result<StoreResult> {
std::fs::create_dir_all(&self.storage_dir)?;
let mut stored = 0;
let mut errors = Vec::new();
for record in records {
let path = self.storage_dir.join(format!("{}.json", record.id));
match serde_json::to_string_pretty(record) {
Ok(json) => {
if let Err(e) = std::fs::write(&path, &json) {
errors.push(format!("Failed to write {}: {e}", record.id));
} else {
stored += 1;
debug!(
record_id = %record.id,
sources = record.source_ids.len(),
"Stored consolidated record"
);
}
}
Err(e) => {
errors.push(format!("Failed to serialize {}: {e}", record.id));
}
}
}
for record in records {
let text = record_text(record);
match self.embedder.embed(&text).await {
Ok(embedding) => {
let mut emb_map = self.embeddings.write().unwrap_or_else(|e| e.into_inner());
emb_map.insert(record.id.clone(), embedding);
let mut rec_map = self.records.write().unwrap_or_else(|e| e.into_inner());
rec_map.insert(record.id.clone(), record.clone());
}
Err(e) => {
errors.push(format!("Failed to embed {}: {e}", record.id));
}
}
}
info!(
"Stored {stored}/{} records to {}",
records.len(),
self.storage_dir.display(),
);
Ok(StoreResult {
stored,
errors,
collection: self.collection_name.clone(),
})
}
pub fn load_all(&self) -> Result<Vec<TemporalRecord>> {
let mut records = Vec::new();
if !self.storage_dir.exists() {
return Ok(records);
}
for entry in std::fs::read_dir(&self.storage_dir)? {
let entry = entry?;
let path = entry.path();
if path.extension().and_then(|e| e.to_str()) == Some("json") {
let content = std::fs::read_to_string(&path)?;
match serde_json::from_str::<TemporalRecord>(&content) {
Ok(record) => records.push(record),
Err(e) => {
tracing::warn!("Failed to parse {}: {e}", path.display());
}
}
}
}
records.sort_by_key(|r| r.sequence_order);
info!("Loaded {} consolidated records", records.len());
Ok(records)
}
pub fn query_by_tag(&self, tag: &str) -> Result<Vec<TemporalRecord>> {
let all = self.load_all()?;
Ok(all
.into_iter()
.filter(|r| r.tags.iter().any(|t| t == tag))
.collect())
}
pub async fn index_existing(&self) -> Result<usize> {
let records = self.load_all()?;
let count = records.len();
for record in &records {
let text = record_text(record);
match self.embedder.embed(&text).await {
Ok(embedding) => {
let mut emb_map = self.embeddings.write().unwrap_or_else(|e| e.into_inner());
emb_map.insert(record.id.clone(), embedding);
let mut rec_map = self.records.write().unwrap_or_else(|e| e.into_inner());
rec_map.insert(record.id.clone(), record.clone());
}
Err(e) => {
tracing::warn!("Failed to embed {}: {e}", record.id);
}
}
}
info!("Indexed {} existing records for retrieval", count);
Ok(count)
}
pub async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<TemporalRecord>> {
if k == 0 {
return Ok(Vec::new());
}
let needs_index = {
let emb_map = self.embeddings.read().unwrap_or_else(|e| e.into_inner());
emb_map.is_empty()
};
if needs_index {
self.index_existing().await?;
}
let query_embedding = self.embedder.embed(query).await?;
let emb_map = self.embeddings.read().unwrap_or_else(|e| e.into_inner());
let mut scored: Vec<(String, f32)> = emb_map
.iter()
.map(|(id, emb)| {
let sim = VectorIndex::cosine_similarity(&query_embedding, emb);
(id.clone(), sim)
})
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let semantic_hits: Vec<(String, f32)> = scored.into_iter().take(k).collect();
drop(emb_map);
let mut result_ids: Vec<String> = semantic_hits.iter().map(|(id, _)| id.clone()).collect();
let mut seen: HashSet<String> = result_ids.iter().cloned().collect();
let rec_map = self.records.read().unwrap_or_else(|e| e.into_inner());
for (id, _) in &semantic_hits {
if let Some(record) = rec_map.get(id) {
for parent in &record.causal_parents {
if !seen.contains(parent) && rec_map.contains_key(parent) {
seen.insert(parent.clone());
result_ids.push(parent.clone());
}
}
for child in &record.causal_children {
if !seen.contains(child) && rec_map.contains_key(child) {
seen.insert(child.clone());
result_ids.push(child.clone());
}
}
}
}
let cap = k * 2;
if result_ids.len() > cap {
result_ids.truncate(cap);
}
let results: Vec<TemporalRecord> = result_ids
.iter()
.filter_map(|id| rec_map.get(id).cloned())
.collect();
drop(rec_map);
Ok(results)
}
}
#[derive(Debug, Clone)]
pub struct StoreResult {
pub stored: usize,
pub errors: Vec<String>,
pub collection: String,
}
#[cfg(test)]
#[path = "../../tests/unit/consolidation/store/store_test.rs"]
mod tests;