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    #[tokio::test]
277    async fn remember_then_recall_finds_exact_match() {
278        let m = fresh().await;
279        m.remember(
280            Scope::Workspace("w1".into()),
281            fact("the cat sat on the mat"),
282        )
283        .await
284        .unwrap();
285        m.remember(Scope::Workspace("w1".into()), fact("rust async runtimes"))
286            .await
287            .unwrap();
288        let hits = m
289            .recall(Scope::Workspace("w1".into()), "the cat sat on the mat", 1)
290            .await
291            .unwrap();
292        assert_eq!(hits.len(), 1);
293        assert_eq!(hits[0].text, "the cat sat on the mat");
294    }
295
296    #[tokio::test]
297    async fn recall_isolates_by_scope() {
298        let m = fresh().await;
299        m.remember(Scope::Workspace("w1".into()), fact("workspace one fact"))
300            .await
301            .unwrap();
302        m.remember(Scope::Workspace("w2".into()), fact("workspace two fact"))
303            .await
304            .unwrap();
305        let hits = m
306            .recall(Scope::Workspace("w1".into()), "any query", 10)
307            .await
308            .unwrap();
309        assert_eq!(hits.len(), 1);
310        assert_eq!(hits[0].text, "workspace one fact");
311    }
312
313    #[tokio::test]
314    async fn recall_respects_k() {
315        let m = fresh().await;
316        for i in 0..5 {
317            m.remember(Scope::Global, fact(&format!("fact-{i}")))
318                .await
319                .unwrap();
320        }
321        let hits = m.recall(Scope::Global, "any", 2).await.unwrap();
322        assert_eq!(hits.len(), 2);
323    }
324
325    #[tokio::test]
326    async fn forget_removes_fact() {
327        let m = fresh().await;
328        let id = m.remember(Scope::Global, fact("fleeting")).await.unwrap();
329        m.forget(id).await.unwrap();
330        let hits = m.recall(Scope::Global, "any", 10).await.unwrap();
331        assert!(hits.is_empty());
332    }
333
334    #[cfg(feature = "sqlite-vec")]
335    #[tokio::test]
336    async fn forget_removes_fact_from_vec_index_too() {
337        // After forget, MATCH-based recall must not return the deleted fact.
338        let m = fresh().await;
339        let id_keep = m
340            .remember(Scope::Global, Fact::new("keeper"))
341            .await
342            .unwrap();
343        let id_drop = m
344            .remember(Scope::Global, Fact::new("to be forgotten"))
345            .await
346            .unwrap();
347        m.forget(id_drop).await.unwrap();
348        let _ = id_keep; // suppress unused-binding warning
349        let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
350        // Recall MATCH-fast-path must not return the forgotten fact.
351        assert!(
352            !hits.iter().any(|f| f.text == "to be forgotten"),
353            "forget did not remove fact from vec0 index; got: {hits:?}"
354        );
355    }
356
357    #[tokio::test]
358    async fn metadata_round_trips() {
359        let m = fresh().await;
360        m.remember(
361            Scope::Global,
362            Fact::new("with metadata").with_metadata(serde_json::json!({"k": "v", "n": 7})),
363        )
364        .await
365        .unwrap();
366        let hits = m.recall(Scope::Global, "with metadata", 1).await.unwrap();
367        assert_eq!(hits.len(), 1);
368        assert_eq!(hits[0].metadata, serde_json::json!({"k": "v", "n": 7}));
369    }
370
371    #[cfg(not(feature = "sqlite-vec"))]
372    #[test]
373    fn cosine_similarity_identical_vectors_returns_one() {
374        let v = vec![1.0_f32, 0.0, 0.0];
375        let score = cosine_similarity(&v, &v);
376        assert!(
377            (score - 1.0).abs() < 1e-6,
378            "identical vectors must have similarity 1.0, got {score}"
379        );
380    }
381
382    #[cfg(not(feature = "sqlite-vec"))]
383    #[test]
384    fn cosine_similarity_orthogonal_vectors_returns_zero() {
385        let a = vec![1.0_f32, 0.0, 0.0];
386        let b = vec![0.0_f32, 1.0, 0.0];
387        let score = cosine_similarity(&a, &b);
388        assert!(
389            score.abs() < 1e-6,
390            "orthogonal vectors must have similarity 0.0, got {score}"
391        );
392    }
393}