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