use crate::connection::DbHandle;
use crate::embedder::Embedder;
use async_trait::async_trait;
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::{Fact, LongTermMemory, Scope};
use std::sync::Arc;
pub struct SqliteLongTerm {
db: DbHandle,
embedder: Arc<dyn Embedder>,
}
impl SqliteLongTerm {
pub(crate) fn new(db: DbHandle, embedder: Arc<dyn Embedder>) -> Self {
Self { db, embedder }
}
}
fn scope_to_kv(scope: &Scope) -> (&'static str, String) {
match scope {
Scope::Workspace(s) => ("workspace", s.clone()),
Scope::Agent(s) => ("agent", s.clone()),
Scope::Global => ("global", String::new()),
}
}
fn embedding_to_blob(v: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(v.len() * 4);
for x in v {
out.extend_from_slice(&x.to_le_bytes());
}
out
}
#[cfg(not(feature = "sqlite-vec"))]
fn blob_to_embedding(b: &[u8]) -> Result<Vec<f32>, MemoryError> {
if b.len() % 4 != 0 {
return Err(MemoryError::Serialization(format!(
"embedding blob length not multiple of 4: {}",
b.len()
)));
}
Ok(b.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect())
}
#[cfg(not(feature = "sqlite-vec"))]
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na == 0.0 || nb == 0.0 {
return 1.0;
}
dot / (na * nb)
}
#[async_trait]
impl LongTermMemory for SqliteLongTerm {
async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
let id = FactId(format!(
"fact-{}-{}",
ulid::Ulid::new(),
chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)
));
let id_inner = id.0.clone();
let (kind, value) = scope_to_kv(&scope);
let metadata_json = serde_json::to_string(&fact.metadata)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let embeds = self
.embedder
.embed(std::slice::from_ref(&fact.text))
.await?;
let embedding = embeds
.into_iter()
.next()
.ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
let dim = self.embedder.dimension();
if embedding.len() != dim {
return Err(MemoryError::Embedding(format!(
"embedder produced {}-dim vector, expected {dim}",
embedding.len()
)));
}
let blob = embedding_to_blob(&embedding);
let text = fact.text;
self.db
.execute(move |conn| {
let tx = conn.transaction()?;
tx.execute(
"INSERT INTO long_term_facts (id, scope_kind, scope_value, text, metadata, embedding) VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
rusqlite::params![&id_inner, kind, &value, &text, &metadata_json, &blob],
)?;
#[cfg(feature = "sqlite-vec")]
{
tx.execute(
"INSERT INTO long_term_facts_vec (fact_id, embedding) VALUES (?1, ?2)",
rusqlite::params![&id_inner, &blob],
)?;
}
tx.commit()?;
Ok(())
})
.await?;
Ok(id)
}
#[allow(unreachable_code)]
async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
if k == 0 {
return Ok(Vec::new());
}
let (kind, value) = scope_to_kv(&scope);
let query_embeds = self.embedder.embed(&[query.to_string()]).await?;
let query_vec = query_embeds
.into_iter()
.next()
.ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
let kind_owned = kind.to_string();
let value_owned = value;
#[cfg(feature = "sqlite-vec")]
{
let query_blob = embedding_to_blob(&query_vec);
let k_i64 = k as i64;
let rows: Vec<(String, String)> = self
.db
.execute(move |conn| {
let mut stmt = conn.prepare(
r#"
SELECT f.text, f.metadata
FROM long_term_facts_vec v
JOIN long_term_facts f ON f.id = v.fact_id
WHERE v.embedding MATCH ?1 AND k = ?2
AND f.scope_kind = ?3 AND f.scope_value = ?4
ORDER BY v.distance ASC
"#,
)?;
let iter = stmt.query_map(
rusqlite::params![&query_blob, k_i64, &kind_owned, &value_owned],
|row| Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)),
)?;
iter.collect::<Result<Vec<_>, _>>()
})
.await?;
return rows
.into_iter()
.map(|(text, metadata_json)| {
let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
Ok(Fact { text, metadata })
})
.collect();
}
#[cfg(not(feature = "sqlite-vec"))]
{
const MAX_SLOW_PATH_FACTS: usize = 10_000;
let row_cap = MAX_SLOW_PATH_FACTS as i64;
let rows: Vec<(String, String, Vec<u8>)> = self
.db
.execute(move |conn| {
let mut stmt = conn.prepare(
"SELECT text, metadata, embedding FROM long_term_facts \
WHERE scope_kind = ?1 AND scope_value = ?2 \
LIMIT ?3",
)?;
let iter = stmt.query_map(
rusqlite::params![&kind_owned, &value_owned, row_cap],
|row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, String>(1)?,
row.get::<_, Vec<u8>>(2)?,
))
},
)?;
iter.collect::<Result<Vec<_>, _>>()
})
.await?;
if rows.len() == MAX_SLOW_PATH_FACTS {
tracing::warn!(
target: "klieo.memory.sqlite",
scope_kind = ?kind,
cap = MAX_SLOW_PATH_FACTS,
"long-term recall hit slow-path row cap — enable the \
`sqlite-vec` feature for O(log N) k-NN recall"
);
}
let mut scored: Vec<(f32, Fact)> = Vec::with_capacity(rows.len());
for (text, metadata_json, blob) in rows {
let emb = blob_to_embedding(&blob)?;
let score = cosine_similarity(&query_vec, &emb);
let metadata: serde_json::Value = serde_json::from_str(&metadata_json)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
scored.push((score, Fact { text, metadata }));
}
scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
return Ok(scored.into_iter().take(k).map(|(_, f)| f).collect());
}
}
async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
let id_inner = id.0;
self.db
.execute(move |conn| {
let tx = conn.transaction()?;
tx.execute(
"DELETE FROM long_term_facts WHERE id = ?1",
rusqlite::params![&id_inner],
)?;
#[cfg(feature = "sqlite-vec")]
{
tx.execute(
"DELETE FROM long_term_facts_vec WHERE fact_id = ?1",
rusqlite::params![&id_inner],
)?;
}
tx.commit()?;
Ok(())
})
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embedder::FakeEmbedder;
use std::sync::Arc;
async fn fresh() -> SqliteLongTerm {
let db = DbHandle::open(":memory:").await.unwrap();
let e: Arc<dyn Embedder> = Arc::new(FakeEmbedder::new(8));
#[cfg(feature = "sqlite-vec")]
{
db.create_vec_table(e.dimension()).await.unwrap();
}
SqliteLongTerm::new(db, e)
}
fn fact(text: &str) -> Fact {
Fact {
text: text.into(),
metadata: serde_json::Value::Null,
}
}
#[tokio::test]
async fn remember_then_recall_finds_exact_match() {
let m = fresh().await;
m.remember(
Scope::Workspace("w1".into()),
fact("the cat sat on the mat"),
)
.await
.unwrap();
m.remember(Scope::Workspace("w1".into()), fact("rust async runtimes"))
.await
.unwrap();
let hits = m
.recall(Scope::Workspace("w1".into()), "the cat sat on the mat", 1)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].text, "the cat sat on the mat");
}
#[tokio::test]
async fn recall_isolates_by_scope() {
let m = fresh().await;
m.remember(Scope::Workspace("w1".into()), fact("workspace one fact"))
.await
.unwrap();
m.remember(Scope::Workspace("w2".into()), fact("workspace two fact"))
.await
.unwrap();
let hits = m
.recall(Scope::Workspace("w1".into()), "any query", 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].text, "workspace one fact");
}
#[tokio::test]
async fn recall_respects_k() {
let m = fresh().await;
for i in 0..5 {
m.remember(Scope::Global, fact(&format!("fact-{i}")))
.await
.unwrap();
}
let hits = m.recall(Scope::Global, "any", 2).await.unwrap();
assert_eq!(hits.len(), 2);
}
#[tokio::test]
async fn forget_removes_fact() {
let m = fresh().await;
let id = m.remember(Scope::Global, fact("fleeting")).await.unwrap();
m.forget(id).await.unwrap();
let hits = m.recall(Scope::Global, "any", 10).await.unwrap();
assert!(hits.is_empty());
}
#[cfg(feature = "sqlite-vec")]
#[tokio::test]
async fn forget_removes_fact_from_vec_index_too() {
let m = fresh().await;
let id_keep = m
.remember(
Scope::Global,
Fact {
text: "keeper".into(),
metadata: serde_json::Value::Null,
},
)
.await
.unwrap();
let id_drop = m
.remember(
Scope::Global,
Fact {
text: "to be forgotten".into(),
metadata: serde_json::Value::Null,
},
)
.await
.unwrap();
m.forget(id_drop).await.unwrap();
let _ = id_keep; let hits = m.recall(Scope::Global, "to be forgotten", 5).await.unwrap();
assert!(
!hits.iter().any(|f| f.text == "to be forgotten"),
"forget did not remove fact from vec0 index; got: {hits:?}"
);
}
#[tokio::test]
async fn metadata_round_trips() {
let m = fresh().await;
m.remember(
Scope::Global,
Fact {
text: "with metadata".into(),
metadata: serde_json::json!({"k": "v", "n": 7}),
},
)
.await
.unwrap();
let hits = m.recall(Scope::Global, "with metadata", 1).await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].metadata, serde_json::json!({"k": "v", "n": 7}));
}
}