1use 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
20pub 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 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 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 #[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 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; let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
365 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}