use crate::embedder::f32_to_bytes;
use crate::errors::AppError;
use crate::storage::utils::with_busy_retry;
use rusqlite::{params, Connection};
pub fn upsert_vec(
conn: &Connection,
memory_id: i64,
namespace: &str,
_memory_type: &str,
embedding: &[f32],
_name: &str,
_snippet: &str,
) -> Result<(), AppError> {
if embedding.is_empty() {
tracing::debug!(
memory_id,
"empty memory embedding: skipping memory_embeddings row (backfill via enrich re-embed)"
);
return Ok(());
}
let embedding_bytes = f32_to_bytes(embedding);
with_busy_retry(|| {
conn.execute(
"DELETE FROM memory_embeddings WHERE memory_id = ?1",
params![memory_id],
)?;
conn.execute(
"INSERT INTO memory_embeddings(memory_id, namespace, embedding, source, model, dim)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
params![
memory_id,
namespace,
&embedding_bytes,
"llm-headless",
crate::constants::SQLITE_GRAPHRAG_VERSION,
crate::constants::embedding_dim() as i64,
],
)?;
Ok(())
})
}
pub fn delete_vec(conn: &Connection, memory_id: i64) -> Result<(), AppError> {
conn.execute(
"DELETE FROM memory_embeddings WHERE memory_id = ?1",
params![memory_id],
)?;
Ok(())
}
pub fn knn_search(
conn: &Connection,
embedding: &[f32],
namespaces: &[String],
memory_type: Option<&str>,
k: usize,
) -> Result<Vec<(i64, f32)>, AppError> {
if embedding.len() != crate::constants::embedding_dim() {
return Err(AppError::Embedding(
crate::i18n::validation::embedding_knn_search_dim_mismatch(
embedding.len(),
crate::constants::embedding_dim(),
),
));
}
let placeholders = (0..namespaces.len())
.map(|_| "?")
.collect::<Vec<_>>()
.join(",");
let ns_clause = if namespaces.is_empty() {
String::new()
} else {
format!(" WHERE e.namespace IN ({placeholders})")
};
let sql = if memory_type.is_some() {
let type_clause = if namespaces.is_empty() {
" WHERE m.type = ?"
} else {
" AND m.type = ?"
};
format!(
"SELECT e.memory_id, e.embedding, e.namespace FROM memory_embeddings e \
LEFT JOIN memories m ON m.id = e.memory_id{ns_clause}{type_clause}"
)
} else {
format!("SELECT e.memory_id, e.embedding, e.namespace FROM memory_embeddings e{ns_clause}")
};
let mut stmt = conn.prepare(&sql)?;
let mut raw_params: Vec<Box<dyn rusqlite::ToSql>> = Vec::new();
for ns in namespaces {
raw_params.push(Box::new(ns.clone()));
}
if let Some(mt) = memory_type {
raw_params.push(Box::new(mt.to_string()));
}
let param_refs: Vec<&dyn rusqlite::ToSql> = raw_params.iter().map(|b| b.as_ref()).collect();
let rows = stmt.query_map(param_refs.as_slice(), |r| {
let id: i64 = r.get(0)?;
let bytes: Vec<u8> = r.get(1)?;
let ns: String = r.get(2)?;
Ok((id, bytes, ns))
})?;
let mut candidates: Vec<(i64, f32)> = Vec::new();
for row in rows {
let (id, bytes, ns) = row?;
let stored = crate::embedder::bytes_to_f32(&bytes);
if stored.len() != embedding.len() {
continue;
}
let sim = crate::similarity::cosine_similarity(embedding, &stored);
let dist = crate::similarity::similarity_to_distance(sim);
let _ = ns; candidates.push((id, dist));
}
candidates.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
candidates.truncate(k);
Ok(candidates)
}