use anyhow::{Context, Result};
use rusqlite::Connection;
use super::super::embedding::EmbeddingProfile;
use super::super::vector_candidates::memory_filter_conditions;
use super::{vec_extension_loaded, VectorHit, VectorSearchFilters};
pub(crate) const VEC_INDEX_BACKFILL_BATCH_SIZE: usize = 512;
const STATE_TABLE: &str = "memory_embedding_vec_state";
fn ensure_state_table(conn: &Connection) -> Result<()> {
conn.execute_batch(&format!(
"CREATE TABLE IF NOT EXISTS {STATE_TABLE} (
dimensions INTEGER PRIMARY KEY,
last_memory_id INTEGER NOT NULL DEFAULT 0,
done INTEGER NOT NULL DEFAULT 0,
updated_at_epoch INTEGER NOT NULL
)"
))?;
Ok(())
}
fn vec_table_name(dimensions: usize) -> String {
format!("memory_embedding_vec_{dimensions}")
}
fn vec_table_exists(conn: &Connection, dimensions: usize) -> Result<bool> {
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?1",
[vec_table_name(dimensions)],
|row| row.get(0),
)?;
Ok(count > 0)
}
fn create_vec_table(conn: &Connection, dimensions: usize) -> Result<()> {
anyhow::ensure!(
(1..=65_536).contains(&dimensions),
"embedding dimensions out of indexable range: {dimensions}"
);
conn.execute_batch(&format!(
"CREATE VIRTUAL TABLE IF NOT EXISTS {} USING vec0(
memory_id INTEGER PRIMARY KEY,
embedding float[{dimensions}] distance_metric=cosine,
+model TEXT
)",
vec_table_name(dimensions)
))?;
Ok(())
}
pub(crate) fn ensure_vec_index(conn: &Connection) -> Result<()> {
if !vec_extension_loaded(conn) {
return Ok(());
}
ensure_state_table(conn)?;
let mut profiles: Vec<i64> = {
let mut stmt =
conn.prepare("SELECT DISTINCT dimensions FROM memory_embeddings ORDER BY dimensions")?;
let rows = stmt.query_map([], |row| row.get::<_, i64>(0))?;
crate::db::query::collect_rows(rows)?
};
let builtin = super::EMBEDDING_DIMENSIONS as i64;
if !profiles.contains(&builtin) {
profiles.push(builtin);
}
for dimensions in profiles {
anyhow::ensure!(
dimensions > 0,
"memory_embeddings carries non-positive dimensions {dimensions}"
);
ensure_vec_index_profile(conn, dimensions as usize)?;
}
Ok(())
}
fn ensure_vec_index_profile(conn: &Connection, dimensions: usize) -> Result<()> {
let state: Option<(i64, i64)> = conn
.query_row(
&format!("SELECT last_memory_id, done FROM {STATE_TABLE} WHERE dimensions = ?1"),
[dimensions as i64],
|row| Ok((row.get(0)?, row.get(1)?)),
)
.ok();
let (cursor, done) = state.unwrap_or((0, 0));
if done == 1 {
return Ok(());
}
create_vec_table(conn, dimensions)?;
let table = vec_table_name(dimensions);
let mut stmt = conn.prepare(
"SELECT memory_id, embedding, model
FROM memory_embeddings
WHERE dimensions = ?1 AND memory_id > ?2
ORDER BY memory_id
LIMIT ?3",
)?;
let rows = stmt.query_map(
rusqlite::params![
dimensions as i64,
cursor,
VEC_INDEX_BACKFILL_BATCH_SIZE as i64
],
|row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, Vec<u8>>(1)?,
row.get::<_, String>(2)?,
))
},
)?;
let batch = crate::db::query::collect_rows(rows)?;
let mut delete = conn.prepare(&format!("DELETE FROM {table} WHERE memory_id = ?1"))?;
let mut insert = conn.prepare(&format!(
"INSERT INTO {table} (memory_id, embedding, model) VALUES (?1, ?2, ?3)"
))?;
let mut advanced = cursor;
for (memory_id, embedding, model) in &batch {
delete.execute([memory_id])?;
insert.execute(rusqlite::params![memory_id, embedding, model])?;
advanced = *memory_id;
}
let finished = (batch.len() < VEC_INDEX_BACKFILL_BATCH_SIZE) as i64;
conn.execute(
&format!(
"INSERT INTO {STATE_TABLE} (dimensions, last_memory_id, done, updated_at_epoch)
VALUES (?1, ?2, ?3, ?4)
ON CONFLICT(dimensions) DO UPDATE SET
last_memory_id = excluded.last_memory_id,
done = excluded.done,
updated_at_epoch = excluded.updated_at_epoch"
),
rusqlite::params![
dimensions as i64,
advanced,
finished,
chrono::Utc::now().timestamp()
],
)?;
Ok(())
}
pub(crate) fn vec_index_ready(conn: &Connection, profile: EmbeddingProfile<'_>) -> Result<bool> {
if !vec_extension_loaded(conn) || !vec_table_exists(conn, profile.dimensions)? {
return Ok(false);
}
let done: Option<i64> = conn
.query_row(
&format!("SELECT done FROM {STATE_TABLE} WHERE dimensions = ?1"),
[profile.dimensions as i64],
|row| row.get(0),
)
.ok();
Ok(done == Some(1))
}
pub(crate) fn sync_vec_upsert(
conn: &Connection,
memory_id: i64,
model: &str,
dimensions: usize,
) -> Result<()> {
if !vec_extension_loaded(conn) || !vec_table_exists(conn, dimensions)? {
return Ok(());
}
let table = vec_table_name(dimensions);
conn.execute(
&format!("DELETE FROM {table} WHERE memory_id = ?1"),
[memory_id],
)
.with_context(|| format!("clear vec index row for memory id={memory_id}"))?;
conn.execute(
&format!(
"INSERT INTO {table} (memory_id, embedding, model)
SELECT memory_id, embedding, model
FROM memory_embeddings
WHERE memory_id = ?1 AND model = ?2 AND dimensions = ?3"
),
rusqlite::params![memory_id, model, dimensions as i64],
)
.with_context(|| format!("sync vec index for memory id={memory_id}"))?;
Ok(())
}
pub(crate) fn sync_vec_upsert_batch(
conn: &Connection,
dimensions: usize,
memory_ids: &[i64],
) -> Result<()> {
if memory_ids.is_empty() || !vec_extension_loaded(conn) || !vec_table_exists(conn, dimensions)?
{
return Ok(());
}
let table = vec_table_name(dimensions);
let id_values: Vec<Box<dyn rusqlite::types::ToSql>> = memory_ids
.iter()
.map(|id| Box::new(*id) as Box<dyn rusqlite::types::ToSql>)
.collect();
let delete_placeholders = (1..=memory_ids.len())
.map(|index| format!("?{index}"))
.collect::<Vec<_>>()
.join(", ");
let id_refs = crate::db::to_sql_refs(&id_values);
conn.execute(
&format!("DELETE FROM {table} WHERE memory_id IN ({delete_placeholders})"),
id_refs.as_slice(),
)?;
let insert_placeholders = (2..=memory_ids.len() + 1)
.map(|index| format!("?{index}"))
.collect::<Vec<_>>()
.join(", ");
let mut values: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(dimensions as i64)];
values.extend(
memory_ids
.iter()
.map(|id| Box::new(*id) as Box<dyn rusqlite::types::ToSql>),
);
let refs = crate::db::to_sql_refs(&values);
conn.execute(
&format!(
"INSERT INTO {table} (memory_id, embedding, model)
SELECT memory_id, embedding, model
FROM memory_embeddings
WHERE dimensions = ?1 AND memory_id IN ({insert_placeholders})"
),
refs.as_slice(),
)?;
Ok(())
}
pub(crate) fn sync_vec_keep_only_profile(
conn: &Connection,
model: &str,
dimensions: usize,
) -> Result<()> {
if !vec_extension_loaded(conn) {
return Ok(());
}
let target_table = vec_table_name(dimensions);
for name in existing_vec_tables(conn)? {
if name == target_table {
conn.execute(
&format!("DELETE FROM \"{name}\" WHERE model != ?1"),
[model],
)?;
} else {
conn.execute_batch(&format!("DROP TABLE \"{name}\""))?;
}
}
ensure_state_table(conn)?;
conn.execute(
&format!("DELETE FROM {STATE_TABLE} WHERE dimensions != ?1"),
[dimensions as i64],
)?;
Ok(())
}
fn existing_vec_tables(conn: &Connection) -> Result<Vec<String>> {
let mut stmt = conn.prepare(
"SELECT name FROM sqlite_master
WHERE type = 'table' AND name LIKE 'memory_embedding_vec_%'",
)?;
let rows = stmt.query_map([], |row| row.get::<_, String>(0))?;
Ok(crate::db::query::collect_rows(rows)?
.into_iter()
.filter(|name| {
name.strip_prefix("memory_embedding_vec_")
.is_some_and(|suffix| {
!suffix.is_empty() && suffix.bytes().all(|b| b.is_ascii_digit())
})
})
.collect())
}
pub(crate) fn knn_candidates(
conn: &Connection,
query_embedding: &[f32],
profile: EmbeddingProfile<'_>,
filters: VectorSearchFilters<'_>,
candidate_limit: usize,
) -> Result<Option<Vec<VectorHit>>> {
if !vec_index_ready(conn, profile)? {
return Ok(None);
}
anyhow::ensure!(
query_embedding.len() == profile.dimensions,
"query embedding must be {} dimensions, got {}",
profile.dimensions,
query_embedding.len()
);
let mut blob = Vec::with_capacity(std::mem::size_of_val(query_embedding));
for value in query_embedding {
blob.extend_from_slice(&value.to_le_bytes());
}
let (conditions, filter_values) = memory_filter_conditions(filters, 5);
let mut values: Vec<Box<dyn rusqlite::types::ToSql>> = vec![
Box::new(blob),
Box::new(candidate_limit as i64),
Box::new(profile.model.to_string()),
Box::new(profile.dimensions as i64),
];
values.extend(filter_values);
let sql = format!(
"SELECT v.memory_id, v.distance
FROM {} v
JOIN memories m ON m.id = v.memory_id
WHERE v.embedding MATCH ?1 AND k = ?2
AND EXISTS (
SELECT 1 FROM memory_embeddings e
WHERE e.memory_id = v.memory_id AND e.model = ?3 AND e.dimensions = ?4
)
AND {}
ORDER BY v.distance",
vec_table_name(profile.dimensions),
conditions.join(" AND ")
);
let refs = crate::db::to_sql_refs(&values);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(refs.as_slice(), |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, f32>(1)?))
})?;
let mut hits = crate::db::query::collect_rows(rows)?
.into_iter()
.map(|(memory_id, distance)| VectorHit {
memory_id,
distance,
})
.collect::<Vec<_>>();
hits.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.memory_id.cmp(&b.memory_id))
});
Ok(Some(hits))
}