use std::collections::HashMap;
use chrono::Utc;
use crate::embed::{cosine_similarity, Embedder};
use crate::memory::{MemoryKind, MemoryRecord};
use crate::store::MemoryStore;
const VECTOR_SIMILARITY_THRESHOLD: f32 = 0.3;
const RRF_K: f64 = 60.0;
#[derive(Debug, Clone, serde::Serialize)]
pub struct RecallResult {
pub memory: MemoryRecord,
pub score: f64,
}
const FACT_DECAY_LAMBDA_PER_DAY: f64 = 0.02;
const EPISODE_DECAY_LAMBDA_PER_DAY: f64 = 0.06;
const FACT_TOUCH_BOOST: f64 = 0.10;
const EPISODE_TOUCH_BOOST: f64 = 0.20;
const CANDIDATE_MULTIPLIER: usize = 10;
const MIN_CANDIDATES: usize = 50;
const SPREAD_FACTOR: f64 = 0.15;
const RECENCY_HALF_LIFE_HOURS: f64 = 168.0;
const RECENCY_FLOOR: f64 = 0.3;
pub fn recall(
store: &MemoryStore,
query: &str,
embedder: &dyn Embedder,
limit: usize,
) -> Result<Vec<RecallResult>, RecallError> {
recall_with_tag_filter(store, query, embedder, limit, None)
}
pub fn recall_with_tag_filter(
store: &MemoryStore,
query: &str,
embedder: &dyn Embedder,
limit: usize,
tag_filter: Option<&str>,
) -> Result<Vec<RecallResult>, RecallError> {
recall_with_tag_filter_ns(store, query, embedder, limit, tag_filter, "default")
}
pub fn recall_with_tag_filter_ns(
store: &MemoryStore,
query: &str,
embedder: &dyn Embedder,
limit: usize,
tag_filter: Option<&str>,
namespace: &str,
) -> Result<Vec<RecallResult>, RecallError> {
let mut all_memories = store.all_memories_with_text_ns(namespace).map_err(RecallError::Db)?;
if let Some(tag) = tag_filter {
let tag_lower = tag.to_lowercase();
all_memories.retain(|(mem, _)| {
mem.tags.iter().any(|t| t.to_lowercase() == tag_lower)
});
}
if all_memories.is_empty() {
return Ok(vec![]);
}
let now = Utc::now();
let max_access = all_memories
.iter()
.map(|(m, _)| m.access_count)
.max()
.unwrap_or(0);
let bm25_ranked = bm25_search(query, &all_memories);
let query_embedding = embedder
.embed_one(query)
.map_err(|e| RecallError::Embedding(e.to_string()))?;
let vector_ranked = vector_search(&query_embedding, &all_memories);
let fused = rrf(&bm25_ranked, &vector_ranked);
let candidate_count = (limit.saturating_mul(CANDIDATE_MULTIPLIER)).max(MIN_CANDIDATES);
let candidates = fused.into_iter().take(candidate_count);
let mut results: Vec<RecallResult> = candidates
.map(|(idx, rrf_score)| {
let mem = &all_memories[idx].0;
let decayed_strength = effective_strength(mem, now);
let recency = recency_boost(mem, now);
let access = access_weight(mem, max_access);
RecallResult {
memory: mem.clone(),
score: rrf_score * decayed_strength * recency * access,
}
})
.collect();
spread_activation(&mut results, SPREAD_FACTOR);
temporal_cooccurrence_boost(&mut results);
results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(limit);
for result in &results {
let mem = &result.memory;
let decayed = effective_strength(mem, now);
let boosted = (decayed + touch_boost(mem)).min(1.0);
store
.touch_memory_with_strength(mem.id, boosted, now)
.map_err(RecallError::Db)?;
}
Ok(results)
}
fn recency_boost(mem: &MemoryRecord, now: chrono::DateTime<Utc>) -> f64 {
let hours_ago = (now - mem.created_at).num_seconds().max(0) as f64 / 3600.0;
let raw = 1.0 / (1.0 + (hours_ago / RECENCY_HALF_LIFE_HOURS).powf(0.8));
raw.max(RECENCY_FLOOR)
}
fn access_weight(mem: &MemoryRecord, max_access: i64) -> f64 {
if max_access <= 0 {
return 1.0;
}
let norm = (mem.access_count as f64 + 1.0).log2() / (max_access as f64 + 1.0).log2();
1.0 + norm }
fn spread_activation(results: &mut Vec<RecallResult>, factor: f64) {
let mut entity_index: HashMap<String, Vec<usize>> = HashMap::new();
for (i, r) in results.iter().enumerate() {
if let MemoryKind::Fact(f) = &r.memory.kind {
let subj = f.subject.to_lowercase();
let obj = f.object.to_lowercase();
entity_index.entry(subj).or_default().push(i);
entity_index.entry(obj).or_default().push(i);
}
}
let mut boosts: HashMap<usize, f64> = HashMap::new();
for (i, r) in results.iter().enumerate() {
if let MemoryKind::Fact(f) = &r.memory.kind {
let entities = [f.subject.to_lowercase(), f.object.to_lowercase()];
for entity in &entities {
if let Some(neighbors) = entity_index.get(entity) {
for &ni in neighbors {
if ni != i {
*boosts.entry(ni).or_insert(0.0) += r.score * factor;
}
}
}
}
}
}
for (idx, boost) in boosts {
if idx < results.len() {
results[idx].score += boost;
}
}
}
fn temporal_cooccurrence_boost(results: &mut Vec<RecallResult>) {
if results.len() < 2 {
return;
}
let mut sorted_indices: Vec<usize> = (0..results.len()).collect();
sorted_indices.sort_by(|&a, &b| {
results[b]
.score
.partial_cmp(&results[a].score)
.unwrap_or(std::cmp::Ordering::Equal)
});
let anchor_count = sorted_indices.len().min(5);
let anchors: Vec<(usize, f64, chrono::DateTime<Utc>)> = sorted_indices[..anchor_count]
.iter()
.map(|&i| (i, results[i].score, results[i].memory.created_at))
.collect();
let mut boosts: HashMap<usize, f64> = HashMap::new();
for (ai, a_score, a_time) in &anchors {
for (j, r) in results.iter().enumerate() {
if j == *ai {
continue;
}
let gap_minutes = (*a_time - r.memory.created_at)
.num_minutes()
.unsigned_abs() as f64;
if gap_minutes < 30.0 {
let proximity = 0.1 * (1.0 - gap_minutes / 30.0);
*boosts.entry(j).or_insert(0.0) += a_score * proximity;
}
}
}
for (idx, boost) in boosts {
if idx < results.len() {
results[idx].score += boost;
}
}
}
fn kind_decay_lambda_per_day(mem: &MemoryRecord) -> f64 {
match &mem.kind {
MemoryKind::Fact(_) => FACT_DECAY_LAMBDA_PER_DAY,
MemoryKind::Episode(_) => EPISODE_DECAY_LAMBDA_PER_DAY,
}
}
fn touch_boost(mem: &MemoryRecord) -> f64 {
match &mem.kind {
MemoryKind::Fact(_) => FACT_TOUCH_BOOST,
MemoryKind::Episode(_) => EPISODE_TOUCH_BOOST,
}
}
fn effective_strength(mem: &MemoryRecord, now: chrono::DateTime<Utc>) -> f64 {
let elapsed_secs = (now - mem.last_accessed_at).num_seconds().max(0) as f64;
let elapsed_days = elapsed_secs / 86_400.0;
let lambda = kind_decay_lambda_per_day(mem);
let effective_lambda = lambda / (1.0 + mem.importance);
(mem.strength * (-effective_lambda * elapsed_days).exp()).clamp(0.0, 1.0)
}
fn bm25_search(query: &str, memories: &[(MemoryRecord, String)]) -> Vec<(usize, f32)> {
use bm25::{Document, Language, SearchEngineBuilder};
let documents: Vec<Document<usize>> = memories
.iter()
.enumerate()
.map(|(i, (_, text))| Document {
id: i,
contents: text.clone(),
})
.collect();
let engine: bm25::SearchEngine<usize> =
SearchEngineBuilder::with_documents(Language::English, documents)
.b(0.5)
.build();
engine
.search(query, memories.len())
.into_iter()
.map(|r| (r.document.id, r.score))
.collect()
}
fn vector_search(query_emb: &[f32], memories: &[(MemoryRecord, String)]) -> Vec<(usize, f32)> {
let mut scored: Vec<(usize, f32)> = memories
.iter()
.enumerate()
.filter_map(|(i, (mem, _))| {
let emb = mem.embedding.as_ref()?;
let sim = cosine_similarity(query_emb, emb);
if sim > VECTOR_SIMILARITY_THRESHOLD {
Some((i, sim))
} else {
None
}
})
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored
}
fn rrf(list_a: &[(usize, f32)], list_b: &[(usize, f32)]) -> Vec<(usize, f64)> {
let mut scores: HashMap<usize, f64> = HashMap::new();
for (rank, &(idx, _)) in list_a.iter().enumerate() {
*scores.entry(idx).or_insert(0.0) += 1.0 / (RRF_K + rank as f64 + 1.0);
}
for (rank, &(idx, _)) in list_b.iter().enumerate() {
*scores.entry(idx).or_insert(0.0) += 1.0 / (RRF_K + rank as f64 + 1.0);
}
let mut results: Vec<(usize, f64)> = scores.into_iter().collect();
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results
}
#[derive(Debug, thiserror::Error)]
pub enum RecallError {
#[error("database error: {0}")]
Db(rusqlite::Error),
#[error("embedding error: {0}")]
Embedding(String),
}
impl From<rusqlite::Error> for RecallError {
fn from(e: rusqlite::Error) -> Self {
RecallError::Db(e)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embed::{EmbedError, Embedding};
struct MockEmbedder;
impl Embedder for MockEmbedder {
fn embed(&self, texts: &[&str]) -> Result<Vec<Embedding>, EmbedError> {
Ok(texts
.iter()
.map(|t| {
if t.contains("alpha") {
vec![1.0, 0.0]
} else {
vec![0.0, 1.0]
}
})
.collect())
}
fn dimension(&self) -> usize {
2
}
}
#[test]
fn effective_strength_decays_by_kind() {
let store = MemoryStore::open_in_memory().unwrap();
let fact_id = store.remember_fact("Jared", "builds", "Gen", Some(&[1.0, 0.0])).unwrap();
let ep_id = store
.remember_episode("alpha project context", Some(&[1.0, 0.0]))
.unwrap();
let old_time = (Utc::now() - chrono::Duration::days(10)).to_rfc3339();
store
.conn()
.execute(
"UPDATE memories SET last_accessed_at = ?1 WHERE id IN (?2, ?3)",
rusqlite::params![old_time, fact_id, ep_id],
)
.unwrap();
let fact = store.get_memory(fact_id).unwrap().unwrap();
let episode = store.get_memory(ep_id).unwrap().unwrap();
let sf = effective_strength(&fact, Utc::now());
let se = effective_strength(&episode, Utc::now());
assert!(sf > se, "facts should decay slower than episodes");
}
#[test]
fn recency_boost_favors_recent_over_old() {
let store = MemoryStore::open_in_memory().unwrap();
let now = Utc::now();
let recent_id = store
.remember_episode("alpha project is great", Some(&[1.0, 0.0]))
.unwrap();
let old_id = store
.remember_episode("alpha project is great", Some(&[1.0, 0.0]))
.unwrap();
let old_time = (now - chrono::Duration::days(30)).to_rfc3339();
store
.conn()
.execute(
"UPDATE memories SET created_at = ?1, last_accessed_at = ?1 WHERE id = ?2",
rusqlite::params![old_time, old_id],
)
.unwrap();
let results = recall(&store, "alpha project", &MockEmbedder, 2).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].memory.id, recent_id);
assert!(results[0].score > results[1].score);
}
#[test]
fn recency_boost_has_floor_old_memories_still_appear() {
let store = MemoryStore::open_in_memory().unwrap();
let now = Utc::now();
let id = store
.remember_episode("alpha ancient knowledge", Some(&[1.0, 0.0]))
.unwrap();
let ancient_time = (now - chrono::Duration::days(365)).to_rfc3339();
store
.conn()
.execute(
"UPDATE memories SET created_at = ?1, last_accessed_at = ?1 WHERE id = ?2",
rusqlite::params![ancient_time, id],
)
.unwrap();
let mem = store.get_memory(id).unwrap().unwrap();
let boost = recency_boost(&mem, now);
assert!(boost >= RECENCY_FLOOR, "recency boost {} should be >= floor {}", boost, RECENCY_FLOOR);
}
#[test]
fn access_weight_boosts_frequently_recalled_memories() {
let store = MemoryStore::open_in_memory().unwrap();
let hot_id = store
.remember_episode("alpha hot memory", Some(&[1.0, 0.0]))
.unwrap();
let cold_id = store
.remember_episode("alpha cold memory", Some(&[1.0, 0.0]))
.unwrap();
store
.conn()
.execute(
"UPDATE memories SET access_count = 20 WHERE id = ?1",
rusqlite::params![hot_id],
)
.unwrap();
let hot = store.get_memory(hot_id).unwrap().unwrap();
let cold = store.get_memory(cold_id).unwrap().unwrap();
let hot_w = access_weight(&hot, 20);
let cold_w = access_weight(&cold, 20);
assert!(hot_w > cold_w, "hot ({}) should weigh more than cold ({})", hot_w, cold_w);
assert!(hot_w >= 1.0 && hot_w <= 2.0, "access weight should be in [1.0, 2.0], got {}", hot_w);
assert!(cold_w >= 1.0, "cold access weight should be >= 1.0, got {}", cold_w);
}
#[test]
fn access_weight_is_bounded() {
let store = MemoryStore::open_in_memory().unwrap();
let id = store
.remember_episode("alpha bounded test", Some(&[1.0, 0.0]))
.unwrap();
store
.conn()
.execute(
"UPDATE memories SET access_count = 1000 WHERE id = ?1",
rusqlite::params![id],
)
.unwrap();
let mem = store.get_memory(id).unwrap().unwrap();
let w = access_weight(&mem, 1000);
assert!(w <= 2.0, "access weight should never exceed 2.0, got {}", w);
}
#[test]
fn spreading_activation_boosts_related_facts() {
let mut results = vec![
RecallResult {
memory: make_fact_record(1, "Jared", "has_pet", "Tortellini"),
score: 1.0,
},
RecallResult {
memory: make_fact_record(2, "Tortellini", "is_a", "dog"),
score: 0.1, },
RecallResult {
memory: make_fact_record(3, "Abby", "likes", "cats"),
score: 0.1, },
];
let original_related = results[1].score;
let original_unrelated = results[2].score;
spread_activation(&mut results, SPREAD_FACTOR);
assert!(
results[1].score > original_related,
"related fact should be boosted: {} > {}",
results[1].score,
original_related
);
assert_eq!(
results[2].score, original_unrelated,
"unrelated fact should not be boosted"
);
}
#[test]
fn spreading_activation_is_bidirectional() {
let mut results = vec![
RecallResult {
memory: make_fact_record(1, "Jared", "works_at", "Microsoft"),
score: 0.8,
},
RecallResult {
memory: make_fact_record(2, "Microsoft", "located_in", "Seattle"),
score: 0.3,
},
];
let score_a_before = results[0].score;
let score_b_before = results[1].score;
spread_activation(&mut results, SPREAD_FACTOR);
assert!(results[1].score > score_b_before);
assert!(results[0].score > score_a_before);
}
#[test]
fn spreading_activation_does_not_self_boost() {
let mut results = vec![
RecallResult {
memory: make_fact_record(1, "Jared", "builds", "Gen"),
score: 1.0,
},
];
spread_activation(&mut results, SPREAD_FACTOR);
assert!((results[0].score - 1.0).abs() < f64::EPSILON);
}
#[test]
fn temporal_cooccurrence_boosts_same_session_memories() {
let now = Utc::now();
let mut results = vec![
RecallResult {
memory: make_timed_episode(1, "alpha anchor memory", now),
score: 1.0,
},
RecallResult {
memory: make_timed_episode(2, "alpha nearby memory", now - chrono::Duration::minutes(5)),
score: 0.2,
},
RecallResult {
memory: make_timed_episode(3, "alpha distant memory", now - chrono::Duration::hours(3)),
score: 0.2,
},
];
let nearby_before = results[1].score;
let distant_before = results[2].score;
temporal_cooccurrence_boost(&mut results);
assert!(
results[1].score > nearby_before,
"nearby memory should be boosted: {} > {}",
results[1].score,
nearby_before
);
assert_eq!(
results[2].score, distant_before,
"distant memory (>30min) should not be boosted"
);
}
#[test]
fn temporal_cooccurrence_scales_with_proximity() {
let now = Utc::now();
let mut results = vec![
RecallResult {
memory: make_timed_episode(1, "alpha anchor", now),
score: 1.0,
},
RecallResult {
memory: make_timed_episode(2, "alpha very close", now - chrono::Duration::minutes(2)),
score: 0.1,
},
RecallResult {
memory: make_timed_episode(3, "alpha further", now - chrono::Duration::minutes(25)),
score: 0.1,
},
];
temporal_cooccurrence_boost(&mut results);
assert!(
results[1].score > results[2].score,
"closer memory ({}) should score higher than further one ({})",
results[1].score,
results[2].score
);
}
#[test]
fn full_recall_pipeline_ranks_recent_accessed_related_higher() {
let store = MemoryStore::open_in_memory().unwrap();
let now = Utc::now();
store.remember_fact("Jared", "has_pet", "Tortellini", Some(&[1.0, 0.0])).unwrap();
store.remember_fact("Tortellini", "is_a", "dog", Some(&[1.0, 0.0])).unwrap();
let old_id = store.remember_fact("weather", "is", "sunny", Some(&[0.5, 0.5])).unwrap();
let old_time = (now - chrono::Duration::days(60)).to_rfc3339();
store
.conn()
.execute(
"UPDATE memories SET created_at = ?1, last_accessed_at = ?1 WHERE id = ?2",
rusqlite::params![old_time, old_id],
)
.unwrap();
let results = recall(&store, "alpha", &MockEmbedder, 10).unwrap();
if results.len() >= 3 {
let weather_pos = results.iter().position(|r| r.memory.id == old_id);
if let Some(pos) = weather_pos {
assert!(pos >= 2, "old unrelated memory should rank below related recent ones, was at position {}", pos);
}
}
}
fn make_fact_record(id: i64, subj: &str, rel: &str, obj: &str) -> MemoryRecord {
MemoryRecord {
id,
kind: MemoryKind::Fact(crate::memory::Fact {
subject: subj.to_string(),
relation: rel.to_string(),
object: obj.to_string(),
}),
strength: 1.0,
created_at: Utc::now(),
last_accessed_at: Utc::now(),
access_count: 0,
embedding: None,
tags: vec![],
source: None,
session_id: None,
channel: None,
importance: 0.5,
namespace: "default".to_string(),
checksum: None,
}
}
fn make_timed_episode(id: i64, text: &str, time: chrono::DateTime<Utc>) -> MemoryRecord {
MemoryRecord {
id,
kind: MemoryKind::Episode(crate::memory::Episode {
text: text.to_string(),
}),
strength: 1.0,
created_at: time,
last_accessed_at: time,
access_count: 0,
embedding: None,
tags: vec![],
source: None,
session_id: None,
channel: None,
importance: 0.5,
namespace: "default".to_string(),
checksum: None,
}
}
#[test]
fn recall_touch_applies_decay_then_reinforcement() {
let store = MemoryStore::open_in_memory().unwrap();
let id = store
.remember_episode("alpha memory to recall", Some(&[1.0, 0.0]))
.unwrap();
let old_time = (Utc::now() - chrono::Duration::days(30)).to_rfc3339();
store
.conn()
.execute(
"UPDATE memories SET strength = 1.0, last_accessed_at = ?1 WHERE id = ?2",
rusqlite::params![old_time, id],
)
.unwrap();
let results = recall(&store, "alpha", &MockEmbedder, 1).unwrap();
assert_eq!(results.len(), 1);
let after = store.get_memory(id).unwrap().unwrap();
assert!(after.access_count >= 1);
assert!(after.strength < 1.0);
assert!(after.strength > 0.2);
}
#[test]
fn recall_with_tag_filter_returns_only_tagged_memories() {
let store = MemoryStore::open_in_memory().unwrap();
store.remember_fact_with_tags("Jared", "likes", "alpha", Some(&[1.0, 0.0]), &["preference".to_string()]).unwrap();
store.remember_fact_with_tags("Jared", "uses", "alpha", Some(&[1.0, 0.0]), &["technical".to_string()]).unwrap();
store.remember_fact("weather", "is", "alpha", Some(&[1.0, 0.0])).unwrap();
let all_results = recall(&store, "alpha", &MockEmbedder, 10).unwrap();
assert_eq!(all_results.len(), 3);
let filtered = recall_with_tag_filter(&store, "alpha", &MockEmbedder, 10, Some("preference")).unwrap();
assert_eq!(filtered.len(), 1);
assert_eq!(filtered[0].memory.tags, vec!["preference"]);
let filtered = recall_with_tag_filter(&store, "alpha", &MockEmbedder, 10, Some("technical")).unwrap();
assert_eq!(filtered.len(), 1);
assert_eq!(filtered[0].memory.tags, vec!["technical"]);
}
#[test]
fn recall_with_tag_filter_is_case_insensitive() {
let store = MemoryStore::open_in_memory().unwrap();
store.remember_fact_with_tags("Jared", "likes", "alpha", Some(&[1.0, 0.0]), &["Preference".to_string()]).unwrap();
let results = recall_with_tag_filter(&store, "alpha", &MockEmbedder, 10, Some("preference")).unwrap();
assert_eq!(results.len(), 1);
}
#[test]
fn recall_with_no_tag_filter_returns_all() {
let store = MemoryStore::open_in_memory().unwrap();
store.remember_fact_with_tags("Jared", "likes", "alpha", Some(&[1.0, 0.0]), &["preference".to_string()]).unwrap();
store.remember_fact("weather", "is", "alpha", Some(&[1.0, 0.0])).unwrap();
let results = recall_with_tag_filter(&store, "alpha", &MockEmbedder, 10, None).unwrap();
assert_eq!(results.len(), 2);
}
}