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, 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        // DummyEmbedder 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    async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
234        let id_inner = id.0;
235        self.db
236            .execute(move |conn| {
237                let tx = conn.transaction()?;
238                tx.execute(
239                    "DELETE FROM long_term_facts WHERE id = ?1",
240                    rusqlite::params![&id_inner],
241                )?;
242                #[cfg(feature = "sqlite-vec")]
243                {
244                    tx.execute(
245                        "DELETE FROM long_term_facts_vec WHERE fact_id = ?1",
246                        rusqlite::params![&id_inner],
247                    )?;
248                }
249                tx.commit()?;
250                Ok(())
251            })
252            .await
253    }
254}
255
256#[cfg(test)]
257mod tests {
258    use super::*;
259    use crate::embedder::FakeEmbedder;
260    use std::sync::Arc;
261
262    async fn fresh() -> SqliteLongTerm {
263        let db = DbHandle::open(":memory:").await.unwrap();
264        let e: Arc<dyn Embedder> = Arc::new(FakeEmbedder::new(8));
265        #[cfg(feature = "sqlite-vec")]
266        {
267            db.create_vec_table(e.dimension()).await.unwrap();
268        }
269        SqliteLongTerm::new(db, e)
270    }
271
272    fn fact(text: &str) -> Fact {
273        Fact::new(text)
274    }
275
276    /// Holds this backend to the same contract as every other
277    /// `LongTermMemory`.
278    ///
279    /// Declared `Relevance` because recall embeds the query and ranks by
280    /// cosine distance (or sqlite-vec k-NN under that feature). The
281    /// near/far pair is caller-supplied, as the suite requires, because only
282    /// this crate knows its embedder: `FakeEmbedder` is FNV-hash based, so
283    /// there is no semantic nearness to exploit — but an exact text match
284    /// embeds to an identical vector (cosine 1.0) while any other string
285    /// hashes elsewhere, so "the exact match outranks an unrelated fact" is
286    /// a real assertion about similarity ordering rather than a fixture
287    /// rigged to pass.
288    #[tokio::test]
289    async fn satisfies_long_term_conformance() {
290        use klieo_core::conformance::{self, ExpectedOrdering, Scopes};
291        let store = fresh().await;
292        let scopes = Scopes::agent("sqlite-conformance-primary", "sqlite-conformance-other");
293        let query = "sqlite conformance relevance probe";
294        conformance::long_term_memory(
295            &store,
296            &scopes,
297            ExpectedOrdering::Relevance {
298                nearer: query.to_string(),
299                farther: "an entirely unrelated fact about diesel maintenance".to_string(),
300                query: query.to_string(),
301            },
302        )
303        .await;
304    }
305
306    #[tokio::test]
307    async fn remember_then_recall_finds_exact_match() {
308        let m = fresh().await;
309        m.remember(
310            Scope::Workspace("w1".into()),
311            fact("the cat sat on the mat"),
312        )
313        .await
314        .unwrap();
315        m.remember(Scope::Workspace("w1".into()), fact("rust async runtimes"))
316            .await
317            .unwrap();
318        let hits = m
319            .recall(Scope::Workspace("w1".into()), "the cat sat on the mat", 1)
320            .await
321            .unwrap();
322        assert_eq!(hits.len(), 1);
323        assert_eq!(hits[0].text, "the cat sat on the mat");
324    }
325
326    #[tokio::test]
327    async fn recall_isolates_by_scope() {
328        let m = fresh().await;
329        m.remember(Scope::Workspace("w1".into()), fact("workspace one fact"))
330            .await
331            .unwrap();
332        m.remember(Scope::Workspace("w2".into()), fact("workspace two fact"))
333            .await
334            .unwrap();
335        let hits = m
336            .recall(Scope::Workspace("w1".into()), "any query", 10)
337            .await
338            .unwrap();
339        assert_eq!(hits.len(), 1);
340        assert_eq!(hits[0].text, "workspace one fact");
341    }
342
343    #[tokio::test]
344    async fn recall_respects_k() {
345        let m = fresh().await;
346        for i in 0..5 {
347            m.remember(Scope::Global, fact(&format!("fact-{i}")))
348                .await
349                .unwrap();
350        }
351        let hits = m.recall(Scope::Global, "any", 2).await.unwrap();
352        assert_eq!(hits.len(), 2);
353    }
354
355    #[tokio::test]
356    async fn forget_removes_fact() {
357        let m = fresh().await;
358        let id = m.remember(Scope::Global, fact("fleeting")).await.unwrap();
359        m.forget(id).await.unwrap();
360        let hits = m.recall(Scope::Global, "any", 10).await.unwrap();
361        assert!(hits.is_empty());
362    }
363
364    #[cfg(feature = "sqlite-vec")]
365    #[tokio::test]
366    async fn forget_removes_fact_from_vec_index_too() {
367        // After forget, MATCH-based recall must not return the deleted fact.
368        let m = fresh().await;
369        let id_keep = m
370            .remember(Scope::Global, Fact::new("keeper"))
371            .await
372            .unwrap();
373        let id_drop = m
374            .remember(Scope::Global, Fact::new("to be forgotten"))
375            .await
376            .unwrap();
377        m.forget(id_drop).await.unwrap();
378        let _ = id_keep; // suppress unused-binding warning
379        let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
380        // Recall MATCH-fast-path must not return the forgotten fact.
381        assert!(
382            !hits.iter().any(|f| f.text == "to be forgotten"),
383            "forget did not remove fact from vec0 index; got: {hits:?}"
384        );
385    }
386
387    #[tokio::test]
388    async fn metadata_round_trips() {
389        let m = fresh().await;
390        m.remember(
391            Scope::Global,
392            Fact::new("with metadata").with_metadata(serde_json::json!({"k": "v", "n": 7})),
393        )
394        .await
395        .unwrap();
396        let hits = m.recall(Scope::Global, "with metadata", 1).await.unwrap();
397        assert_eq!(hits.len(), 1);
398        assert_eq!(hits[0].metadata, serde_json::json!({"k": "v", "n": 7}));
399    }
400
401    #[cfg(not(feature = "sqlite-vec"))]
402    #[test]
403    fn cosine_similarity_identical_vectors_returns_one() {
404        let v = vec![1.0_f32, 0.0, 0.0];
405        let score = cosine_similarity(&v, &v);
406        assert!(
407            (score - 1.0).abs() < 1e-6,
408            "identical vectors must have similarity 1.0, got {score}"
409        );
410    }
411
412    #[cfg(not(feature = "sqlite-vec"))]
413    #[test]
414    fn cosine_similarity_orthogonal_vectors_returns_zero() {
415        let a = vec![1.0_f32, 0.0, 0.0];
416        let b = vec![0.0_f32, 1.0, 0.0];
417        let score = cosine_similarity(&a, &b);
418        assert!(
419            score.abs() < 1e-6,
420            "orthogonal vectors must have similarity 0.0, got {score}"
421        );
422    }
423}