klieo-memory-graph-rag 3.4.0

Graph-first RAG composer over KnowledgeGraph + LongTermMemory. Stable at 1.x per ADR-039.
Documentation
//! In-memory [`klieo_memory_graph::FilterableLongTermMemory`] with real cosine similarity.
//!
//! Non-production-scale (linear scan) — used by `App::graph_rag()` zero-infra
//! tier so demo recall matches the Qdrant tier semantically.

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>,
}

/// In-memory implementation of [`FilterableLongTermMemory`] backed by real
/// cosine similarity over an injected [`Embedder`].
///
/// Linear scan — suitable for the zero-infra demo tier or unit tests, not
/// for production workloads with large fact sets.
#[non_exhaustive]
pub struct InMemoryFilterableLongTerm {
    embedder: Arc<dyn Embedder>,
    embedder_id: String,
    facts: Mutex<HashMap<Scope, Vec<StoredFact>>>,
    seq: AtomicU64,
}

impl InMemoryFilterableLongTerm {
    /// Build a store backed by `embedder`, identified under `embedder_id`.
    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());
    }
}