Skip to main content

klieo_memory_sqlite/
long_term.rs

1//! `SqliteLongTerm` — `LongTermMemory` over a SQLite table with
2//! linear-scan cosine recall (default) or sqlite-vec k-NN MATCH
3//! (under feature `sqlite-vec`).
4//!
5//! Embeddings are stored as `BLOB` (little-endian f32 array). On
6//! `recall`, all rows matching the scope are loaded, the query is
7//! embedded, cosine similarity is computed in Rust, results are sorted
8//! and the top `k` returned. Acceptable for under ~10k facts; when the
9//! `sqlite-vec` feature is enabled, a virtual shadow table is used for
10//! k-NN via the MATCH operator instead.
11
12use crate::connection::DbHandle;
13use crate::embedder::Embedder;
14use async_trait::async_trait;
15use klieo_core::error::MemoryError;
16use klieo_core::ids::FactId;
17use klieo_core::memory::{Fact, LongTermMemory, RecallSemantics, Scope};
18use std::sync::Arc;
19
20/// SQLite-backed long-term semantic memory.
21pub struct SqliteLongTerm {
22    db: DbHandle,
23    embedder: Arc<dyn Embedder>,
24}
25
26impl SqliteLongTerm {
27    pub(crate) fn new(db: DbHandle, embedder: Arc<dyn Embedder>) -> Self {
28        Self { db, embedder }
29    }
30}
31
32fn scope_to_kv(scope: &Scope) -> (&'static str, String) {
33    match scope {
34        Scope::Workspace(s) => ("workspace", s.clone()),
35        Scope::Agent(s) => ("agent", s.clone()),
36        Scope::Global => ("global", String::new()),
37    }
38}
39
40fn embedding_to_blob(v: &[f32]) -> Vec<u8> {
41    let mut out = Vec::with_capacity(v.len() * 4);
42    for x in v {
43        out.extend_from_slice(&x.to_le_bytes());
44    }
45    out
46}
47
48#[cfg(not(feature = "sqlite-vec"))]
49fn blob_to_embedding(b: &[u8]) -> Result<Vec<f32>, MemoryError> {
50    if !b.len().is_multiple_of(4) {
51        return Err(MemoryError::Serialization(format!(
52            "embedding blob length not multiple of 4: {}",
53            b.len()
54        )));
55    }
56    Ok(b.chunks_exact(4)
57        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
58        .collect())
59}
60
61#[cfg(not(feature = "sqlite-vec"))]
62fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
63    if a.len() != b.len() {
64        return 0.0;
65    }
66    let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
67    let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
68    let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
69    if na == 0.0 || nb == 0.0 {
70        // Zero-vector → undefined direction; treat as max similarity so
71        // NonRankingEmbedder behaves predictably (FIFO-ish recall).
72        return 1.0;
73    }
74    dot / (na * nb)
75}
76
77#[async_trait]
78impl LongTermMemory for SqliteLongTerm {
79    async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
80        let id = FactId(format!(
81            "fact-{}-{}",
82            ulid::Ulid::new(),
83            chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)
84        ));
85        let id_inner = id.0.clone();
86        let (kind, value) = scope_to_kv(&scope);
87        let metadata_json = serde_json::to_string(&fact.metadata)
88            .map_err(|e| MemoryError::Serialization(e.to_string()))?;
89        let embeds = self
90            .embedder
91            .embed(std::slice::from_ref(&fact.text))
92            .await?;
93        let embedding = embeds
94            .into_iter()
95            .next()
96            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
97        let dim = self.embedder.dimension();
98        if embedding.len() != dim {
99            return Err(MemoryError::Embedding(format!(
100                "embedder produced {}-dim vector, expected {dim}",
101                embedding.len()
102            )));
103        }
104        let blob = embedding_to_blob(&embedding);
105        let text = fact.text;
106        self.db
107            .execute(move |conn| {
108                let tx = conn.transaction()?;
109                tx.execute(
110                    "INSERT INTO long_term_facts (id, scope_kind, scope_value, text, metadata, embedding) VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
111                    rusqlite::params![&id_inner, kind, &value, &text, &metadata_json, &blob],
112                )?;
113                #[cfg(feature = "sqlite-vec")]
114                {
115                    tx.execute(
116                        "INSERT INTO long_term_facts_vec (fact_id, embedding) VALUES (?1, ?2)",
117                        rusqlite::params![&id_inner, &blob],
118                    )?;
119                }
120                tx.commit()?;
121                Ok(())
122            })
123            .await?;
124        Ok(id)
125    }
126
127    #[allow(unreachable_code)]
128    async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
129        if k == 0 {
130            return Ok(Vec::new());
131        }
132        let (kind, value) = scope_to_kv(&scope);
133        let query_embeds = self.embedder.embed(&[query.to_string()]).await?;
134        let query_vec = query_embeds
135            .into_iter()
136            .next()
137            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
138        let kind_owned = kind.to_string();
139        let value_owned = value;
140
141        #[cfg(feature = "sqlite-vec")]
142        {
143            // Fast path: k-NN via sqlite-vec MATCH operator. Filters
144            // post-MATCH by scope.
145            let query_blob = embedding_to_blob(&query_vec);
146            let k_i64 = k as i64;
147            let rows: Vec<(String, String)> = self
148                .db
149                .execute(move |conn| {
150                    let mut stmt = conn.prepare(
151                        r#"
152                        SELECT f.text, f.metadata
153                        FROM long_term_facts_vec v
154                        JOIN long_term_facts f ON f.id = v.fact_id
155                        WHERE v.embedding MATCH ?1 AND k = ?2
156                          AND f.scope_kind = ?3 AND f.scope_value = ?4
157                        ORDER BY v.distance ASC
158                        "#,
159                    )?;
160                    let iter = stmt.query_map(
161                        rusqlite::params![&query_blob, k_i64, &kind_owned, &value_owned],
162                        |row| Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)),
163                    )?;
164                    iter.collect::<Result<Vec<_>, _>>()
165                })
166                .await?;
167            return rows
168                .into_iter()
169                .map(|(text, metadata_json)| {
170                    let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
171                        .map_err(|e| MemoryError::Serialization(e.to_string()))?;
172                    Ok(Fact::new(text).with_metadata(metadata))
173                })
174                .collect();
175        }
176
177        // Slow path (no sqlite-vec feature): linear scan + cosine.
178        //
179        // W5.A26 / round-1 perf MED: this path is O(N) in facts per
180        // scope. For scopes above ~10 000 facts the per-call cost
181        // dominates agent recall latency; the `sqlite-vec` feature
182        // routes through a virtual-table MATCH that uses a real k-NN
183        // index. We cap the row stream at MAX_SLOW_PATH_FACTS to
184        // protect the agent from a runaway scope — callers hitting
185        // the cap get a hint to enable the fast feature.
186        #[cfg(not(feature = "sqlite-vec"))]
187        {
188            const MAX_SLOW_PATH_FACTS: usize = 10_000;
189            let row_cap = MAX_SLOW_PATH_FACTS as i64;
190            let rows: Vec<(String, String, Vec<u8>)> = self
191                .db
192                .execute(move |conn| {
193                    let mut stmt = conn.prepare(
194                        "SELECT text, metadata, embedding FROM long_term_facts \
195                         WHERE scope_kind = ?1 AND scope_value = ?2 \
196                         LIMIT ?3",
197                    )?;
198                    let iter = stmt.query_map(
199                        rusqlite::params![&kind_owned, &value_owned, row_cap],
200                        |row| {
201                            Ok((
202                                row.get::<_, String>(0)?,
203                                row.get::<_, String>(1)?,
204                                row.get::<_, Vec<u8>>(2)?,
205                            ))
206                        },
207                    )?;
208                    iter.collect::<Result<Vec<_>, _>>()
209                })
210                .await?;
211            if rows.len() == MAX_SLOW_PATH_FACTS {
212                tracing::warn!(
213                    target: "klieo.memory.sqlite",
214                    scope_kind = ?kind,
215                    cap = MAX_SLOW_PATH_FACTS,
216                    "long-term recall hit slow-path row cap — enable the \
217                     `sqlite-vec` feature for O(log N) k-NN recall"
218                );
219            }
220            let mut scored: Vec<(f32, Fact)> = Vec::with_capacity(rows.len());
221            for (text, metadata_json, blob) in rows {
222                let emb = blob_to_embedding(&blob)?;
223                let score = cosine_similarity(&query_vec, &emb);
224                let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
225                    .map_err(|e| MemoryError::Serialization(e.to_string()))?;
226                scored.push((score, Fact::new(text).with_metadata(metadata)));
227            }
228            scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
229            return Ok(scored.into_iter().take(k).map(|(_, f)| f).collect());
230        }
231    }
232
233    /// Vector recall when the injected embedder actually separates inputs;
234    /// `NonRanking` when it does not (see `NonRankingEmbedder`).
235    fn recall_semantics(&self) -> RecallSemantics {
236        klieo_embed_common::vector_recall_semantics(self.embedder.as_ref())
237    }
238
239    async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
240        let id_inner = id.0;
241        self.db
242            .execute(move |conn| {
243                let tx = conn.transaction()?;
244                tx.execute(
245                    "DELETE FROM long_term_facts WHERE id = ?1",
246                    rusqlite::params![&id_inner],
247                )?;
248                #[cfg(feature = "sqlite-vec")]
249                {
250                    tx.execute(
251                        "DELETE FROM long_term_facts_vec WHERE fact_id = ?1",
252                        rusqlite::params![&id_inner],
253                    )?;
254                }
255                tx.commit()?;
256                Ok(())
257            })
258            .await
259    }
260}
261
262#[cfg(test)]
263mod tests {
264    use super::*;
265    use crate::embedder::FakeEmbedder;
266    use std::sync::Arc;
267
268    async fn fresh() -> SqliteLongTerm {
269        let db = DbHandle::open(":memory:").await.unwrap();
270        let e: Arc<dyn Embedder> = Arc::new(FakeEmbedder::new(8));
271        #[cfg(feature = "sqlite-vec")]
272        {
273            db.create_vec_table(e.dimension()).await.unwrap();
274        }
275        SqliteLongTerm::new(db, e)
276    }
277
278    fn fact(text: &str) -> Fact {
279        Fact::new(text)
280    }
281
282    /// Holds this backend to the same contract as every other
283    /// `LongTermMemory`.
284    ///
285    /// Declared `Relevance` because recall embeds the query and ranks by
286    /// cosine distance (or sqlite-vec k-NN under that feature). The
287    /// near/far pair is caller-supplied, as the suite requires, because only
288    /// this crate knows its embedder: `FakeEmbedder` is FNV-hash based, so
289    /// there is no semantic nearness to exploit — but an exact text match
290    /// embeds to an identical vector (cosine 1.0) while any other string
291    /// hashes elsewhere, so "the exact match outranks an unrelated fact" is
292    /// a real assertion about similarity ordering rather than a fixture
293    /// rigged to pass.
294    #[tokio::test]
295    async fn satisfies_long_term_conformance() {
296        use klieo_core::conformance::{self, ExpectedOrdering, Scopes};
297        let store = fresh().await;
298        let scopes = Scopes::agent("sqlite-conformance-primary", "sqlite-conformance-other");
299        let query = "sqlite conformance relevance probe";
300        conformance::long_term_memory(
301            &store,
302            &scopes,
303            ExpectedOrdering::Relevance {
304                nearer: query.to_string(),
305                farther: "an entirely unrelated fact about diesel maintenance".to_string(),
306                query: query.to_string(),
307            },
308        )
309        .await;
310    }
311
312    #[tokio::test]
313    async fn remember_then_recall_finds_exact_match() {
314        let m = fresh().await;
315        m.remember(
316            Scope::Workspace("w1".into()),
317            fact("the cat sat on the mat"),
318        )
319        .await
320        .unwrap();
321        m.remember(Scope::Workspace("w1".into()), fact("rust async runtimes"))
322            .await
323            .unwrap();
324        let hits = m
325            .recall(Scope::Workspace("w1".into()), "the cat sat on the mat", 1)
326            .await
327            .unwrap();
328        assert_eq!(hits.len(), 1);
329        assert_eq!(hits[0].text, "the cat sat on the mat");
330    }
331
332    #[tokio::test]
333    async fn recall_isolates_by_scope() {
334        let m = fresh().await;
335        m.remember(Scope::Workspace("w1".into()), fact("workspace one fact"))
336            .await
337            .unwrap();
338        m.remember(Scope::Workspace("w2".into()), fact("workspace two fact"))
339            .await
340            .unwrap();
341        let hits = m
342            .recall(Scope::Workspace("w1".into()), "any query", 10)
343            .await
344            .unwrap();
345        assert_eq!(hits.len(), 1);
346        assert_eq!(hits[0].text, "workspace one fact");
347    }
348
349    #[tokio::test]
350    async fn recall_respects_k() {
351        let m = fresh().await;
352        for i in 0..5 {
353            m.remember(Scope::Global, fact(&format!("fact-{i}")))
354                .await
355                .unwrap();
356        }
357        let hits = m.recall(Scope::Global, "any", 2).await.unwrap();
358        assert_eq!(hits.len(), 2);
359    }
360
361    #[tokio::test]
362    async fn forget_removes_fact() {
363        let m = fresh().await;
364        let id = m.remember(Scope::Global, fact("fleeting")).await.unwrap();
365        m.forget(id).await.unwrap();
366        let hits = m.recall(Scope::Global, "any", 10).await.unwrap();
367        assert!(hits.is_empty());
368    }
369
370    #[cfg(feature = "sqlite-vec")]
371    #[tokio::test]
372    async fn forget_removes_fact_from_vec_index_too() {
373        // After forget, MATCH-based recall must not return the deleted fact.
374        let m = fresh().await;
375        let id_keep = m
376            .remember(Scope::Global, Fact::new("keeper"))
377            .await
378            .unwrap();
379        let id_drop = m
380            .remember(Scope::Global, Fact::new("to be forgotten"))
381            .await
382            .unwrap();
383        m.forget(id_drop).await.unwrap();
384        let _ = id_keep; // suppress unused-binding warning
385        let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
386        // Recall MATCH-fast-path must not return the forgotten fact.
387        assert!(
388            !hits.iter().any(|f| f.text == "to be forgotten"),
389            "forget did not remove fact from vec0 index; got: {hits:?}"
390        );
391    }
392
393    #[tokio::test]
394    async fn metadata_round_trips() {
395        let m = fresh().await;
396        m.remember(
397            Scope::Global,
398            Fact::new("with metadata").with_metadata(serde_json::json!({"k": "v", "n": 7})),
399        )
400        .await
401        .unwrap();
402        let hits = m.recall(Scope::Global, "with metadata", 1).await.unwrap();
403        assert_eq!(hits.len(), 1);
404        assert_eq!(hits[0].metadata, serde_json::json!({"k": "v", "n": 7}));
405    }
406
407    #[cfg(not(feature = "sqlite-vec"))]
408    #[test]
409    fn cosine_similarity_identical_vectors_returns_one() {
410        let v = vec![1.0_f32, 0.0, 0.0];
411        let score = cosine_similarity(&v, &v);
412        assert!(
413            (score - 1.0).abs() < 1e-6,
414            "identical vectors must have similarity 1.0, got {score}"
415        );
416    }
417
418    #[cfg(not(feature = "sqlite-vec"))]
419    #[test]
420    fn cosine_similarity_orthogonal_vectors_returns_zero() {
421        let a = vec![1.0_f32, 0.0, 0.0];
422        let b = vec![0.0_f32, 1.0, 0.0];
423        let score = cosine_similarity(&a, &b);
424        assert!(
425            score.abs() < 1e-6,
426            "orthogonal vectors must have similarity 0.0, got {score}"
427        );
428    }
429}