use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::{Fact, LongTermMemory, Scope};
use klieo_embed_common::Embedder;
use klieo_memory_graph::FilterableLongTermMemory;
struct StoredFact {
id: FactId,
fact: Fact,
vector: Vec<f32>,
}
#[non_exhaustive]
pub struct InMemoryFilterableLongTerm {
embedder: Arc<dyn Embedder>,
embedder_id: String,
facts: Mutex<HashMap<Scope, Vec<StoredFact>>>,
seq: AtomicU64,
}
impl InMemoryFilterableLongTerm {
pub fn new(embedder: Arc<dyn Embedder>, embedder_id: impl Into<String>) -> Self {
Self {
embedder,
embedder_id: embedder_id.into(),
facts: Mutex::new(HashMap::new()),
seq: AtomicU64::new(0),
}
}
async fn embed_one(&self, text: &str) -> Result<Vec<f32>, MemoryError> {
let mut vecs = self.embedder.embed(&[text.to_string()]).await?;
vecs.pop()
.ok_or_else(|| MemoryError::Embedding("embedder returned no vector".into()))
}
fn ranked(
&self,
scope: &Scope,
query_vec: &[f32],
k: usize,
candidate_ids: Option<&[FactId]>,
) -> Result<Vec<Fact>, MemoryError> {
let guard = self
.facts
.lock()
.map_err(|_| MemoryError::Store("facts mutex poisoned".into()))?;
let Some(bucket) = guard.get(scope) else {
return Ok(Vec::new());
};
let mut scored: Vec<(f32, &Fact)> = bucket
.iter()
.filter(|sf| candidate_ids.is_none_or(|ids| ids.contains(&sf.id)))
.map(|sf| (cosine(query_vec, &sf.vector), &sf.fact))
.collect();
scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
Ok(scored.into_iter().take(k).map(|(_, f)| f.clone()).collect())
}
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(
a.len(),
b.len(),
"cosine: vector length mismatch ({} vs {})",
a.len(),
b.len()
);
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
0.0
} else {
dot / (norm_a * norm_b)
}
}
#[async_trait]
impl LongTermMemory for InMemoryFilterableLongTerm {
async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
let vector = self.embed_one(&fact.text).await?;
let n = self.seq.fetch_add(1, Ordering::Relaxed);
let id = FactId::new(format!("inmem-{n}"));
self.facts
.lock()
.map_err(|_| MemoryError::Store("facts mutex poisoned".into()))?
.entry(scope)
.or_default()
.push(StoredFact {
id: id.clone(),
fact,
vector,
});
Ok(id)
}
async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
let query_vec = self.embed_one(query).await?;
self.ranked(&scope, &query_vec, k, None)
}
async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
let mut guard = self
.facts
.lock()
.map_err(|_| MemoryError::Store("facts mutex poisoned".into()))?;
for bucket in guard.values_mut() {
bucket.retain(|sf| sf.id != id);
}
Ok(())
}
}
#[async_trait]
impl FilterableLongTermMemory for InMemoryFilterableLongTerm {
async fn recall_filtered(
&self,
scope: Scope,
query: &str,
k: usize,
candidate_ids: &[FactId],
) -> Result<Vec<Fact>, MemoryError> {
let query_vec = self.embed_one(query).await?;
self.ranked(&scope, &query_vec, k, Some(candidate_ids))
}
fn embedder_id(&self) -> &str {
&self.embedder_id
}
}
#[cfg(test)]
mod tests {
use super::*;
use klieo_core::memory::{Fact, Scope};
use std::sync::Arc;
struct KeywordEmbedder;
#[async_trait]
impl Embedder for KeywordEmbedder {
fn dimension(&self) -> usize {
3
}
async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError> {
Ok(texts
.iter()
.map(|t| {
let l = t.to_lowercase();
vec![
l.matches("alpha").count() as f32,
l.matches("beta").count() as f32,
l.matches("gamma").count() as f32,
]
})
.collect())
}
}
#[tokio::test]
async fn remember_then_recall_ranks_by_cosine() {
let store = InMemoryFilterableLongTerm::new(Arc::new(KeywordEmbedder), "kw-test");
let s = Scope::Workspace("w".into());
store
.remember(s.clone(), Fact::new("alpha alpha"))
.await
.unwrap();
store.remember(s.clone(), Fact::new("beta")).await.unwrap();
let hits = store.recall(s, "alpha", 1).await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].text, "alpha alpha");
}
#[tokio::test]
async fn recall_filtered_restricts_to_candidates() {
let store = InMemoryFilterableLongTerm::new(Arc::new(KeywordEmbedder), "kw-test");
let s = Scope::Workspace("w".into());
let id_a = store.remember(s.clone(), Fact::new("alpha")).await.unwrap();
let _id_b = store
.remember(s.clone(), Fact::new("alpha alpha"))
.await
.unwrap();
let hits = store
.recall_filtered(s, "alpha", 5, std::slice::from_ref(&id_a))
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].text, "alpha");
}
#[test]
fn embedder_id_is_stable() {
let store = InMemoryFilterableLongTerm::new(Arc::new(KeywordEmbedder), "kw-test");
assert_eq!(store.embedder_id(), "kw-test");
}
#[tokio::test]
async fn recall_filtered_with_empty_candidates_returns_empty() {
let store = InMemoryFilterableLongTerm::new(Arc::new(KeywordEmbedder), "kw-test");
let s = Scope::Workspace("w".into());
store.remember(s.clone(), Fact::new("alpha")).await.unwrap();
let hits = store.recall_filtered(s, "alpha", 5, &[]).await.unwrap();
assert!(
hits.is_empty(),
"empty candidate_ids must short-circuit to empty"
);
}
#[tokio::test]
async fn forget_removes_fact_from_recall() {
let store = InMemoryFilterableLongTerm::new(Arc::new(KeywordEmbedder), "kw-test");
let s = Scope::Workspace("w".into());
let id = store.remember(s.clone(), Fact::new("alpha")).await.unwrap();
store.forget(id).await.unwrap();
let hits = store.recall(s, "alpha", 5).await.unwrap();
assert!(hits.is_empty());
}
}