klieo-memory-sqlite 0.4.0

SQLite-backed implementations of klieo-core's memory traits.
Documentation
//! `SqliteLongTerm` — `LongTermMemory` over a SQLite table with
//! linear-scan cosine recall (default) or sqlite-vec k-NN MATCH
//! (under feature `sqlite-vec`).
//!
//! Embeddings are stored as `BLOB` (little-endian f32 array). On
//! `recall`, all rows matching the scope are loaded, the query is
//! embedded, cosine similarity is computed in Rust, results are sorted
//! and the top `k` returned. Acceptable for under ~10k facts; when the
//! `sqlite-vec` feature is enabled, a virtual shadow table is used for
//! k-NN via the MATCH operator instead.

use crate::connection::DbHandle;
use crate::embedder::Embedder;
use async_trait::async_trait;
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::{Fact, LongTermMemory, Scope};
use std::sync::Arc;

/// SQLite-backed long-term semantic memory.
pub struct SqliteLongTerm {
    db: DbHandle,
    embedder: Arc<dyn Embedder>,
}

impl SqliteLongTerm {
    pub(crate) fn new(db: DbHandle, embedder: Arc<dyn Embedder>) -> Self {
        Self { db, embedder }
    }
}

fn scope_to_kv(scope: &Scope) -> (&'static str, String) {
    match scope {
        Scope::Workspace(s) => ("workspace", s.clone()),
        Scope::Agent(s) => ("agent", s.clone()),
        Scope::Global => ("global", String::new()),
    }
}

fn embedding_to_blob(v: &[f32]) -> Vec<u8> {
    let mut out = Vec::with_capacity(v.len() * 4);
    for x in v {
        out.extend_from_slice(&x.to_le_bytes());
    }
    out
}

#[cfg(not(feature = "sqlite-vec"))]
fn blob_to_embedding(b: &[u8]) -> Result<Vec<f32>, MemoryError> {
    if !b.len().is_multiple_of(4) {
        return Err(MemoryError::Serialization(format!(
            "embedding blob length not multiple of 4: {}",
            b.len()
        )));
    }
    Ok(b.chunks_exact(4)
        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
        .collect())
}

#[cfg(not(feature = "sqlite-vec"))]
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
    if a.len() != b.len() {
        return 0.0;
    }
    let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
    let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
    let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
    if na == 0.0 || nb == 0.0 {
        // Zero-vector → undefined direction; treat as max similarity so
        // DummyEmbedder behaves predictably (FIFO-ish recall).
        return 1.0;
    }
    dot / (na * nb)
}

#[async_trait]
impl LongTermMemory for SqliteLongTerm {
    async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
        let id = FactId(format!(
            "fact-{}-{}",
            ulid::Ulid::new(),
            chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)
        ));
        let id_inner = id.0.clone();
        let (kind, value) = scope_to_kv(&scope);
        let metadata_json = serde_json::to_string(&fact.metadata)
            .map_err(|e| MemoryError::Serialization(e.to_string()))?;
        let embeds = self
            .embedder
            .embed(std::slice::from_ref(&fact.text))
            .await?;
        let embedding = embeds
            .into_iter()
            .next()
            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
        let dim = self.embedder.dimension();
        if embedding.len() != dim {
            return Err(MemoryError::Embedding(format!(
                "embedder produced {}-dim vector, expected {dim}",
                embedding.len()
            )));
        }
        let blob = embedding_to_blob(&embedding);
        let text = fact.text;
        self.db
            .execute(move |conn| {
                let tx = conn.transaction()?;
                tx.execute(
                    "INSERT INTO long_term_facts (id, scope_kind, scope_value, text, metadata, embedding) VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
                    rusqlite::params![&id_inner, kind, &value, &text, &metadata_json, &blob],
                )?;
                #[cfg(feature = "sqlite-vec")]
                {
                    tx.execute(
                        "INSERT INTO long_term_facts_vec (fact_id, embedding) VALUES (?1, ?2)",
                        rusqlite::params![&id_inner, &blob],
                    )?;
                }
                tx.commit()?;
                Ok(())
            })
            .await?;
        Ok(id)
    }

    #[allow(unreachable_code)]
    async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
        if k == 0 {
            return Ok(Vec::new());
        }
        let (kind, value) = scope_to_kv(&scope);
        let query_embeds = self.embedder.embed(&[query.to_string()]).await?;
        let query_vec = query_embeds
            .into_iter()
            .next()
            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
        let kind_owned = kind.to_string();
        let value_owned = value;

        #[cfg(feature = "sqlite-vec")]
        {
            // Fast path: k-NN via sqlite-vec MATCH operator. Filters
            // post-MATCH by scope.
            let query_blob = embedding_to_blob(&query_vec);
            let k_i64 = k as i64;
            let rows: Vec<(String, String)> = self
                .db
                .execute(move |conn| {
                    let mut stmt = conn.prepare(
                        r#"
                        SELECT f.text, f.metadata
                        FROM long_term_facts_vec v
                        JOIN long_term_facts f ON f.id = v.fact_id
                        WHERE v.embedding MATCH ?1 AND k = ?2
                          AND f.scope_kind = ?3 AND f.scope_value = ?4
                        ORDER BY v.distance ASC
                        "#,
                    )?;
                    let iter = stmt.query_map(
                        rusqlite::params![&query_blob, k_i64, &kind_owned, &value_owned],
                        |row| Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)),
                    )?;
                    iter.collect::<Result<Vec<_>, _>>()
                })
                .await?;
            return rows
                .into_iter()
                .map(|(text, metadata_json)| {
                    let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
                        .map_err(|e| MemoryError::Serialization(e.to_string()))?;
                    Ok(Fact { text, metadata })
                })
                .collect();
        }

        // Slow path (no sqlite-vec feature): linear scan + cosine.
        //
        // W5.A26 / round-1 perf MED: this path is O(N) in facts per
        // scope. For scopes above ~10 000 facts the per-call cost
        // dominates agent recall latency; the `sqlite-vec` feature
        // routes through a virtual-table MATCH that uses a real k-NN
        // index. We cap the row stream at MAX_SLOW_PATH_FACTS to
        // protect the agent from a runaway scope — callers hitting
        // the cap get a hint to enable the fast feature.
        #[cfg(not(feature = "sqlite-vec"))]
        {
            const MAX_SLOW_PATH_FACTS: usize = 10_000;
            let row_cap = MAX_SLOW_PATH_FACTS as i64;
            let rows: Vec<(String, String, Vec<u8>)> = self
                .db
                .execute(move |conn| {
                    let mut stmt = conn.prepare(
                        "SELECT text, metadata, embedding FROM long_term_facts \
                         WHERE scope_kind = ?1 AND scope_value = ?2 \
                         LIMIT ?3",
                    )?;
                    let iter = stmt.query_map(
                        rusqlite::params![&kind_owned, &value_owned, row_cap],
                        |row| {
                            Ok((
                                row.get::<_, String>(0)?,
                                row.get::<_, String>(1)?,
                                row.get::<_, Vec<u8>>(2)?,
                            ))
                        },
                    )?;
                    iter.collect::<Result<Vec<_>, _>>()
                })
                .await?;
            if rows.len() == MAX_SLOW_PATH_FACTS {
                tracing::warn!(
                    target: "klieo.memory.sqlite",
                    scope_kind = ?kind,
                    cap = MAX_SLOW_PATH_FACTS,
                    "long-term recall hit slow-path row cap — enable the \
                     `sqlite-vec` feature for O(log N) k-NN recall"
                );
            }
            let mut scored: Vec<(f32, Fact)> = Vec::with_capacity(rows.len());
            for (text, metadata_json, blob) in rows {
                let emb = blob_to_embedding(&blob)?;
                let score = cosine_similarity(&query_vec, &emb);
                let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
                    .map_err(|e| MemoryError::Serialization(e.to_string()))?;
                scored.push((score, Fact { text, metadata }));
            }
            scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
            return Ok(scored.into_iter().take(k).map(|(_, f)| f).collect());
        }
    }

    async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
        let id_inner = id.0;
        self.db
            .execute(move |conn| {
                let tx = conn.transaction()?;
                tx.execute(
                    "DELETE FROM long_term_facts WHERE id = ?1",
                    rusqlite::params![&id_inner],
                )?;
                #[cfg(feature = "sqlite-vec")]
                {
                    tx.execute(
                        "DELETE FROM long_term_facts_vec WHERE fact_id = ?1",
                        rusqlite::params![&id_inner],
                    )?;
                }
                tx.commit()?;
                Ok(())
            })
            .await
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::embedder::FakeEmbedder;
    use std::sync::Arc;

    async fn fresh() -> SqliteLongTerm {
        let db = DbHandle::open(":memory:").await.unwrap();
        let e: Arc<dyn Embedder> = Arc::new(FakeEmbedder::new(8));
        #[cfg(feature = "sqlite-vec")]
        {
            db.create_vec_table(e.dimension()).await.unwrap();
        }
        SqliteLongTerm::new(db, e)
    }

    fn fact(text: &str) -> Fact {
        Fact {
            text: text.into(),
            metadata: serde_json::Value::Null,
        }
    }

    #[tokio::test]
    async fn remember_then_recall_finds_exact_match() {
        let m = fresh().await;
        m.remember(
            Scope::Workspace("w1".into()),
            fact("the cat sat on the mat"),
        )
        .await
        .unwrap();
        m.remember(Scope::Workspace("w1".into()), fact("rust async runtimes"))
            .await
            .unwrap();
        let hits = m
            .recall(Scope::Workspace("w1".into()), "the cat sat on the mat", 1)
            .await
            .unwrap();
        assert_eq!(hits.len(), 1);
        assert_eq!(hits[0].text, "the cat sat on the mat");
    }

    #[tokio::test]
    async fn recall_isolates_by_scope() {
        let m = fresh().await;
        m.remember(Scope::Workspace("w1".into()), fact("workspace one fact"))
            .await
            .unwrap();
        m.remember(Scope::Workspace("w2".into()), fact("workspace two fact"))
            .await
            .unwrap();
        let hits = m
            .recall(Scope::Workspace("w1".into()), "any query", 10)
            .await
            .unwrap();
        assert_eq!(hits.len(), 1);
        assert_eq!(hits[0].text, "workspace one fact");
    }

    #[tokio::test]
    async fn recall_respects_k() {
        let m = fresh().await;
        for i in 0..5 {
            m.remember(Scope::Global, fact(&format!("fact-{i}")))
                .await
                .unwrap();
        }
        let hits = m.recall(Scope::Global, "any", 2).await.unwrap();
        assert_eq!(hits.len(), 2);
    }

    #[tokio::test]
    async fn forget_removes_fact() {
        let m = fresh().await;
        let id = m.remember(Scope::Global, fact("fleeting")).await.unwrap();
        m.forget(id).await.unwrap();
        let hits = m.recall(Scope::Global, "any", 10).await.unwrap();
        assert!(hits.is_empty());
    }

    #[cfg(feature = "sqlite-vec")]
    #[tokio::test]
    async fn forget_removes_fact_from_vec_index_too() {
        // After forget, MATCH-based recall must not return the deleted fact.
        let m = fresh().await;
        let id_keep = m
            .remember(
                Scope::Global,
                Fact {
                    text: "keeper".into(),
                    metadata: serde_json::Value::Null,
                },
            )
            .await
            .unwrap();
        let id_drop = m
            .remember(
                Scope::Global,
                Fact {
                    text: "to be forgotten".into(),
                    metadata: serde_json::Value::Null,
                },
            )
            .await
            .unwrap();
        m.forget(id_drop).await.unwrap();
        let _ = id_keep; // suppress unused-binding warning
        let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
        // Recall MATCH-fast-path must not return the forgotten fact.
        assert!(
            !hits.iter().any(|f| f.text == "to be forgotten"),
            "forget did not remove fact from vec0 index; got: {hits:?}"
        );
    }

    #[tokio::test]
    async fn metadata_round_trips() {
        let m = fresh().await;
        m.remember(
            Scope::Global,
            Fact {
                text: "with metadata".into(),
                metadata: serde_json::json!({"k": "v", "n": 7}),
            },
        )
        .await
        .unwrap();
        let hits = m.recall(Scope::Global, "with metadata", 1).await.unwrap();
        assert_eq!(hits.len(), 1);
        assert_eq!(hits[0].metadata, serde_json::json!({"k": "v", "n": 7}));
    }
}