use crate::storage::{EdgeStorage, MemoryStorage};
use anyhow::Result;
use async_trait::async_trait;
use mr_common::{EdgeType, Memory, MemoryEdge};
use uuid::Uuid;
#[async_trait]
pub trait EdgeGenerator: Send + Sync {
async fn generate(
&self,
new_memory: &Memory,
candidate_memories: &[Memory],
edge_store: &dyn EdgeStorage,
) -> Result<Vec<MemoryEdge>>;
}
pub struct MockEdgeGenerator {
min_similarity: f32,
max_edges_per_memory: usize,
}
impl MockEdgeGenerator {
pub fn new(min_similarity: f32, max_edges_per_memory: usize) -> Self {
Self {
min_similarity: min_similarity.clamp(0.0, 1.0),
max_edges_per_memory,
}
}
pub fn with_defaults() -> Self {
Self::new(0.5, 10)
}
fn calculate_similarity(content_a: &str, content_b: &str) -> f32 {
let a_lower = content_a.to_lowercase();
let b_lower = content_b.to_lowercase();
let words_a: std::collections::HashSet<&str> = a_lower.split_whitespace().collect();
let words_b: std::collections::HashSet<&str> = b_lower.split_whitespace().collect();
if words_a.is_empty() || words_b.is_empty() {
return 0.0;
}
let intersection = words_a.intersection(&words_b).count();
let union = words_a.union(&words_b).count();
if union == 0 {
0.0
} else {
intersection as f32 / union as f32
}
}
fn infer_edge_type(&self, _content_a: &str, _content_b: &str) -> EdgeType {
EdgeType::Similar
}
fn generate_evidence(edge_type: EdgeType, similarity: f32) -> String {
format!(
"Auto-generated {} edge (similarity: {:.2})",
edge_type, similarity
)
}
}
#[async_trait]
impl EdgeGenerator for MockEdgeGenerator {
async fn generate(
&self,
new_memory: &Memory,
candidate_memories: &[Memory],
edge_store: &dyn EdgeStorage,
) -> Result<Vec<MemoryEdge>> {
let existing_edges = edge_store.neighbors(&new_memory.id).await?;
let existing_targets: std::collections::HashSet<Uuid> = existing_edges
.iter()
.map(|e| {
if e.source_id == new_memory.id {
e.target_id
} else {
e.source_id
}
})
.collect();
let mut edges = Vec::new();
let mut edge_count = existing_edges.len();
for candidate in candidate_memories {
if edge_count >= self.max_edges_per_memory {
break;
}
if existing_targets.contains(&candidate.id) {
continue;
}
if new_memory.id == candidate.id {
continue;
}
let similarity = Self::calculate_similarity(&new_memory.content, &candidate.content);
if similarity >= self.min_similarity {
let edge_type = self.infer_edge_type(&new_memory.content, &candidate.content);
let evidence = Self::generate_evidence(edge_type, similarity);
let edge =
MemoryEdge::new(new_memory.id, candidate.id, edge_type, similarity, evidence)
.auto_generated();
edges.push(edge);
edge_count += 1;
}
}
Ok(edges)
}
}
pub struct ProjectEdgeGenerator {
max_edges_per_project: usize,
}
impl ProjectEdgeGenerator {
pub fn new(max_edges_per_project: usize) -> Self {
Self {
max_edges_per_project,
}
}
pub fn with_defaults() -> Self {
Self::new(5)
}
pub async fn generate_project_edges(
&self,
memory: &Memory,
memory_store: &dyn MemoryStorage,
edge_store: &dyn EdgeStorage,
) -> Result<Vec<MemoryEdge>> {
let project_id = match memory.project_id {
Some(id) => id,
None => return Ok(Vec::new()),
};
let project_memories = memory_store.list_by_project(&project_id).await?;
let existing = edge_store.neighbors(&memory.id).await?;
let existing_ids: std::collections::HashSet<Uuid> = existing
.iter()
.filter(|e| e.edge_type == EdgeType::SameProject)
.map(|e| {
if e.source_id == memory.id {
e.target_id
} else {
e.source_id
}
})
.collect();
let mut edges = Vec::new();
for other in project_memories {
if other.id == memory.id {
continue;
}
if existing_ids.contains(&other.id) {
continue;
}
if edges.len() >= self.max_edges_per_project {
break;
}
edges.push(
MemoryEdge::new(
memory.id,
other.id,
EdgeType::SameProject,
1.0,
format!("Same project: {:?}", memory.project_id),
)
.auto_generated(),
);
}
Ok(edges)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::{EdgeStore, RocksDBStore};
use mr_common::MemoryType;
use tempfile::tempdir;
async fn create_test_edge_store() -> EdgeStore {
let dir = tempdir().unwrap();
let rocksdb = RocksDBStore::open(dir.path()).unwrap();
EdgeStore::new(rocksdb)
}
fn create_test_memory(content: &str) -> Memory {
Memory::new(content.to_string(), MemoryType::Knowledge)
}
#[tokio::test]
async fn test_mock_generator_similar_edges() {
let edge_store = create_test_edge_store().await;
let generator = MockEdgeGenerator::with_defaults();
let memory_a = create_test_memory("Tokio is a great async runtime for Rust");
let memory_b = create_test_memory("async-std is another async runtime for Rust");
let memory_c = create_test_memory("Database optimization techniques");
let edges = generator
.generate(&memory_a, &[memory_b.clone(), memory_c], &edge_store)
.await
.unwrap();
assert!(edges.len() >= 1);
assert!(edges[0].edge_type == EdgeType::Similar);
assert!(edges[0].auto_generated);
}
#[tokio::test]
async fn test_mock_generator_respects_max_edges() {
let edge_store = create_test_edge_store().await;
let generator = MockEdgeGenerator::new(0.3, 2);
let memory = create_test_memory("Test memory with keywords");
let candidates: Vec<Memory> = (0..10)
.map(|i| create_test_memory(&format!("Test memory with keywords {}", i)))
.collect();
let edges = generator
.generate(&memory, &candidates, &edge_store)
.await
.unwrap();
assert!(edges.len() <= 2);
}
#[tokio::test]
async fn test_mock_generator_avoids_self_edge() {
let edge_store = create_test_edge_store().await;
let generator = MockEdgeGenerator::with_defaults();
let memory = create_test_memory("Test content");
let edges = generator
.generate(&memory, &[memory.clone()], &edge_store)
.await
.unwrap();
assert!(edges.is_empty());
}
#[tokio::test]
async fn test_mock_generator_low_similarity_filtered() {
let edge_store = create_test_edge_store().await;
let generator = MockEdgeGenerator::new(0.8, 10);
let memory_a = create_test_memory("Rust programming language");
let memory_b = create_test_memory("Cooking recipes for dinner");
let edges = generator
.generate(&memory_a, &[memory_b], &edge_store)
.await
.unwrap();
assert!(edges.is_empty());
}
}