use std::collections::HashMap;
use std::sync::Arc;
use tracing::info;
use crate::storage::MemoryStorage;
use super::{upsert_summary, PhaseResult};
#[allow(dead_code)]
pub struct CrossProjectExtractor {
storage: Arc<dyn MemoryStorage>,
batch_size: usize,
batch_interval_ms: u64,
}
impl CrossProjectExtractor {
pub fn new(storage: Arc<dyn MemoryStorage>, batch_size: usize, batch_interval_ms: u64) -> Self {
Self {
storage,
batch_size,
batch_interval_ms,
}
}
pub async fn execute(&self) -> PhaseResult {
info!(target: "dream", "Dream Phase 1 started: CrossProjectExtract");
let memories = match self.storage.list(10000).await {
Ok(m) => m,
Err(e) => return PhaseResult::err("CrossProjectExtract", e.to_string()),
};
if memories.is_empty() {
info!(target: "dream", "Dream Phase 1 skipped: no memories found");
return PhaseResult::ok("CrossProjectExtract", 0, 0);
}
let mut tag_project_count: HashMap<String, usize> = HashMap::new();
for memory in &memories {
let project_key = memory
.project_id
.map(|p| p.to_string())
.unwrap_or_else(|| "global".to_string());
for tag in &memory.tags {
let key = format!("{}:{}", project_key, tag);
*tag_project_count.entry(key).or_insert(0) += 1;
}
}
let mut cross_project_tags: HashMap<String, usize> = HashMap::new();
for key in tag_project_count.keys() {
let parts: Vec<&str> = key.splitn(2, ':').collect();
if parts.len() == 2 {
let tag = parts[1];
*cross_project_tags.entry(tag.to_string()).or_insert(0) += 1;
}
}
let cross_project_tags: Vec<String> = cross_project_tags
.iter()
.filter(|(_, count)| **count > 1)
.map(|(tag, _)| tag.clone())
.collect();
let processed = memories.len();
if cross_project_tags.is_empty() {
info!(target: "dream", "Dream Phase 1 completed: no cross-project tags found");
return PhaseResult::ok("CrossProjectExtract", processed, 0);
}
info!(
target: "dream",
"Dream Phase 1: found {} cross-project tags",
cross_project_tags.len()
);
let summary_content = format!(
"Cross-project common themes: {}",
cross_project_tags.join(", ")
);
let (created, deleted) =
match upsert_summary(&self.storage, "cross-project", summary_content, 0.3).await {
Ok(result) => result,
Err(e) => {
tracing::warn!("Failed to upsert cross-project summary: {}", e);
(0, 0)
}
};
info!(
target: "dream",
"Dream Phase 1 completed: processed {} memories, created {}, removed {}",
processed, created, deleted
);
PhaseResult::ok("CrossProjectExtract", processed, created)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::{MemoryStore, RocksDBStore};
use mr_common::types::{Memory, MemoryType};
use tempfile::tempdir;
#[tokio::test]
async fn test_cross_project_extractor_empty() {
let dir = tempdir().unwrap();
let rocksdb = RocksDBStore::open(dir.path()).unwrap();
let storage = Arc::new(MemoryStore::new(std::sync::Arc::new(rocksdb)));
let extractor = CrossProjectExtractor::new(storage, 100, 10);
let result = extractor.execute().await;
assert!(result.success);
assert_eq!(result.processed_count, 0);
}
#[tokio::test]
async fn test_cross_project_extractor_with_memories() {
let dir = tempdir().unwrap();
let rocksdb = RocksDBStore::open(dir.path()).unwrap();
let storage = Arc::new(MemoryStore::new(std::sync::Arc::new(rocksdb)));
let mut m1 = Memory::new("memory 1".to_string(), MemoryType::Knowledge);
m1.tags = vec!["rust".to_string()];
m1.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m1).await.unwrap();
let mut m2 = Memory::new("memory 2".to_string(), MemoryType::Knowledge);
m2.tags = vec!["rust".to_string()];
m2.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m2).await.unwrap();
let extractor = CrossProjectExtractor::new(storage, 100, 10);
let result = extractor.execute().await;
assert!(result.success);
assert_eq!(result.processed_count, 2);
assert_eq!(result.created_count, 1);
}
#[tokio::test]
async fn test_cross_project_extractor_upsert_idempotent() {
let dir = tempdir().unwrap();
let rocksdb = RocksDBStore::open(dir.path()).unwrap();
let storage: Arc<dyn MemoryStorage> =
Arc::new(MemoryStore::new(std::sync::Arc::new(rocksdb)));
let mut m1 = Memory::new("memory 1".to_string(), MemoryType::Knowledge);
m1.tags = vec!["rust".to_string()];
m1.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m1).await.unwrap();
let mut m2 = Memory::new("memory 2".to_string(), MemoryType::Knowledge);
m2.tags = vec!["rust".to_string()];
m2.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m2).await.unwrap();
let extractor = CrossProjectExtractor::new(storage.clone(), 100, 10);
let r1 = extractor.execute().await;
assert!(r1.success);
assert_eq!(r1.created_count, 1);
let summary = storage.list_by_tag("cross-project", 10).await.unwrap();
assert_eq!(summary.len(), 1);
assert_eq!(summary[0].source, mr_common::types::MemorySource::System);
assert_eq!(summary[0].importance, 0.3);
let r2 = extractor.execute().await;
assert!(r2.success);
assert_eq!(r2.created_count, 0);
let summary = storage.list_by_tag("cross-project", 10).await.unwrap();
assert_eq!(summary.len(), 1);
}
#[tokio::test]
async fn test_cross_project_extractor_upsert_rolls_back_extra() {
let dir = tempdir().unwrap();
let rocksdb = RocksDBStore::open(dir.path()).unwrap();
let storage: Arc<dyn MemoryStorage> =
Arc::new(MemoryStore::new(std::sync::Arc::new(rocksdb)));
let mut m1 = Memory::new("memory 1".to_string(), MemoryType::Knowledge);
m1.tags = vec!["rust".to_string()];
m1.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m1).await.unwrap();
let mut m2 = Memory::new("memory 2".to_string(), MemoryType::Knowledge);
m2.tags = vec!["rust".to_string()];
m2.project_id = Some(uuid::Uuid::new_v4());
storage.save(&m2).await.unwrap();
for i in 0..3 {
let mut stale = Memory::new(format!("stale summary {}", i), MemoryType::Knowledge);
stale.tags = vec!["cross-project".to_string()];
stale.importance = 0.8;
stale.created_at = chrono::Utc::now() - chrono::Duration::hours(i as i64 + 1);
storage.save(&stale).await.unwrap();
}
let extractor = CrossProjectExtractor::new(storage.clone(), 100, 10);
let result = extractor.execute().await;
assert!(result.success);
assert_eq!(result.created_count, 0);
let summary = storage.list_by_tag("cross-project", 100).await.unwrap();
assert_eq!(summary.len(), 1);
assert!(summary[0]
.content
.starts_with("Cross-project common themes: rust"));
assert_eq!(summary[0].source, mr_common::types::MemorySource::System);
assert_eq!(summary[0].importance, 0.3);
}
}