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::new(text).with_metadata(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::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 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; let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
350 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}