use crate::error::LoopError;
use crate::memory::{ConsolidationStats, LoopMemory, MemoryEntry};
use std::future::Future;
use std::sync::{PoisonError, RwLock};
pub struct InMemoryStore {
entries: RwLock<Vec<MemoryEntry>>,
}
impl InMemoryStore {
#[must_use]
pub fn new() -> Self {
Self {
entries: RwLock::new(Vec::new()),
}
}
#[must_use]
pub fn with_entries(self, entries: Vec<MemoryEntry>) -> Self {
*self.entries.write().unwrap_or_else(PoisonError::into_inner) = entries;
self
}
}
impl Default for InMemoryStore {
fn default() -> Self {
Self::new()
}
}
#[allow(clippy::manual_async_fn)]
impl LoopMemory for InMemoryStore {
fn store(&self, entry: MemoryEntry) -> impl Future<Output = Result<(), LoopError>> + Send {
async move {
self.entries
.write()
.unwrap_or_else(PoisonError::into_inner)
.push(entry);
Ok(())
}
}
fn retrieve(
&self,
query: &str,
limit: usize,
) -> impl Future<Output = Result<Vec<MemoryEntry>, LoopError>> + Send {
let query = query.to_string();
async move {
let query_lower = query.to_lowercase();
let query_words: Vec<&str> = query_lower.split_whitespace().collect();
let entries = self.entries.read().unwrap_or_else(PoisonError::into_inner);
let snapshot: Vec<MemoryEntry> = entries.iter().cloned().collect();
drop(entries);
let mut scored: Vec<(f32, MemoryEntry)> = snapshot
.into_iter()
.map(|entry| {
let memory_lower = entry.memory.to_lowercase();
let tag_match = entry
.tags
.iter()
.any(|t| t.to_lowercase().contains(&query_lower));
let word_matches = query_words
.iter()
.filter(|w| memory_lower.contains(*w))
.count();
let base_score = entry.relevance;
#[allow(clippy::cast_precision_loss)]
let query_bonus = if word_matches > 0 {
word_matches as f32 / query_words.len().max(1) as f32
} else {
0.0
};
let tag_bonus = if tag_match { 0.3 } else { 0.0 };
(
base_score * 0.5 + query_bonus * 0.4 + tag_bonus + 0.1,
entry,
)
})
.collect();
scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
Ok(scored.into_iter().take(limit).map(|(_, e)| e).collect())
}
}
fn consolidate(&self) -> impl Future<Output = Result<ConsolidationStats, LoopError>> + Send {
async move {
let mut entries = self.entries.write().unwrap_or_else(PoisonError::into_inner);
let entries_before = entries.len();
entries.retain(|e| e.relevance >= 0.05);
let pruned = entries_before.saturating_sub(entries.len());
Ok(ConsolidationStats {
entries_before,
entries_after: entries.len(),
pruned,
merged: 0,
bytes_saved: 0,
})
}
}
fn len(&self) -> usize {
self.entries
.read()
.unwrap_or_else(PoisonError::into_inner)
.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::MemoryCategory;
#[tokio::test]
async fn test_store_and_retrieve() {
let store = InMemoryStore::new();
store
.store(MemoryEntry::new(
MemoryCategory::Insight,
"Prefer Glob over manual file search",
))
.await
.unwrap();
store
.store(MemoryEntry::new(
MemoryCategory::ErrorPattern,
"Edit failures often caused by stale file content",
))
.await
.unwrap();
let results = store.retrieve("Glob manual file search", 5).await.unwrap();
assert!(!results.is_empty());
assert_eq!(results[0].category, MemoryCategory::Insight);
}
#[tokio::test]
async fn test_retrieve_respects_limit() {
let store = InMemoryStore::new();
for i in 0..10 {
store
.store(MemoryEntry::new(
MemoryCategory::Fact,
format!("Fact number {i} about testing"),
))
.await
.unwrap();
}
let results = store.retrieve("testing", 3).await.unwrap();
assert_eq!(results.len(), 3);
}
#[tokio::test]
async fn test_retrieve_empty_store() {
let store = InMemoryStore::new();
let results = store.retrieve("anything", 5).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn test_len_and_is_empty() {
let store = InMemoryStore::new();
assert!(store.is_empty());
assert_eq!(store.len(), 0);
}
#[tokio::test]
async fn test_consolidate_prunes_low_relevance() {
let store = InMemoryStore::new();
let mut good_entry = MemoryEntry::new(MemoryCategory::Insight, "useful insight");
good_entry.relevance = 0.9;
store.store(good_entry).await.unwrap();
let mut bad_entry = MemoryEntry::new(MemoryCategory::Working, "temporary data");
bad_entry.relevance = 0.01;
store.store(bad_entry).await.unwrap();
assert_eq!(store.len(), 2);
let stats = store.consolidate().await.unwrap();
assert_eq!(stats.entries_before, 2);
assert_eq!(stats.pruned, 1);
assert_eq!(store.len(), 1);
}
#[tokio::test]
async fn test_with_entries() {
let entries = vec![
MemoryEntry::new(MemoryCategory::Fact, "fact 1"),
MemoryEntry::new(MemoryCategory::Fact, "fact 2"),
];
let store = InMemoryStore::new().with_entries(entries);
assert_eq!(store.len(), 2);
}
#[tokio::test]
async fn test_tag_matching_boosts_relevance() {
let store = InMemoryStore::new();
let tagged =
MemoryEntry::new(MemoryCategory::Strategy, "use iterators for loops").with_tag("rust");
store.store(tagged).await.unwrap();
store
.store(MemoryEntry::new(
MemoryCategory::Strategy,
"use caching for performance",
))
.await
.unwrap();
let results = store.retrieve("rust iterators", 2).await.unwrap();
assert!(!results.is_empty());
assert!(results[0].memory.contains("iterators"));
}
#[tokio::test]
async fn test_default_is_empty() {
let store = InMemoryStore::default();
assert!(store.is_empty());
}
#[tokio::test]
async fn test_retrieve_does_not_block_writers() {
let store = InMemoryStore::new();
for i in 0..200 {
store
.store(MemoryEntry::new(
MemoryCategory::Fact,
format!("Fact number {i} about concurrency"),
))
.await
.unwrap();
}
let retrieve_fut = store.retrieve("concurrency", 5);
let store_fut = store.store(MemoryEntry::new(
MemoryCategory::Insight,
"writer proceeds concurrently",
));
let (retrieved, store_res) = tokio::join!(retrieve_fut, store_fut);
let retrieved = retrieved.unwrap();
store_res.unwrap();
assert!(retrieved.len() <= 5);
assert_eq!(store.len(), 201); }
#[tokio::test]
async fn test_retrieve_ranking_preserved() {
let store = InMemoryStore::new();
let mut high = MemoryEntry::new(MemoryCategory::Insight, "rust rust rust rust");
high.relevance = 0.95;
let mut mid = MemoryEntry::new(MemoryCategory::Fact, "rust rust rust");
mid.relevance = 0.5;
let mut low = MemoryEntry::new(MemoryCategory::Working, "rust rust");
low.relevance = 0.1;
store.store(low.clone()).await.unwrap();
store.store(high.clone()).await.unwrap();
store.store(mid.clone()).await.unwrap();
let results = store.retrieve("rust", 3).await.unwrap();
assert_eq!(results.len(), 3);
assert!((results[0].relevance - 0.95).abs() < 1e-6);
assert!((results[1].relevance - 0.5).abs() < 1e-6);
assert!((results[2].relevance - 0.1).abs() < 1e-6);
}
}