use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, OnceLock, RwLock};
use rusqlite::Connection;
use serde::{Deserialize, Serialize};
use usearch::{Index, IndexOptions, MetricKind, ScalarKind};
use kimetsu_core::KimetsuResult;
const SCHEMA_VERSION: u32 = 2;
const CONNECTIVITY: usize = 16;
const EXPANSION_ADD: usize = 128;
const EXPANSION_SEARCH: usize = 64;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
struct Manifest {
schema_version: u32,
dim: usize,
model_id: String,
max_rowid_indexed: i64,
count: usize,
quant: String,
}
fn ann_scalar_kind() -> ScalarKind {
match std::env::var("KIMETSU_ANN_QUANTIZATION").ok().as_deref() {
Some("f32") => ScalarKind::F32,
Some("i8") => ScalarKind::I8,
Some("f16") | None => ScalarKind::F16,
Some(other) => {
eprintln!("kimetsu-brain: unknown KIMETSU_ANN_QUANTIZATION '{other}', using f16");
ScalarKind::F16
}
}
}
fn scalar_kind_id(k: ScalarKind) -> &'static str {
match k {
ScalarKind::F32 => "f32",
ScalarKind::F16 => "f16",
ScalarKind::I8 => "i8",
_ => "other",
}
}
fn index_options(dim: usize) -> IndexOptions {
IndexOptions {
dimensions: dim,
metric: MetricKind::Cos,
quantization: ann_scalar_kind(),
connectivity: CONNECTIVITY,
expansion_add: EXPANSION_ADD,
expansion_search: EXPANSION_SEARCH,
multi: false,
}
}
fn build_threads() -> usize {
std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
.clamp(1, 16)
}
fn parallel_add(index: &Index, rows: &[(i64, Vec<f32>)]) -> KimetsuResult<()> {
if rows.is_empty() {
return Ok(());
}
let nthreads = build_threads().min(rows.len());
let chunk = rows.len().div_ceil(nthreads);
let err: std::sync::Mutex<Option<String>> = std::sync::Mutex::new(None);
std::thread::scope(|s| {
for part in rows.chunks(chunk) {
let err = &err;
s.spawn(move || {
for (rowid, vec) in part {
if let Err(e) = index.add(*rowid as u64, vec) {
let mut g = err.lock().unwrap_or_else(|p| p.into_inner());
if g.is_none() {
*g = Some(format!("usearch add: {e}"));
}
return;
}
}
});
}
});
match err.into_inner().unwrap_or_else(|p| p.into_inner()) {
Some(e) => Err(e.into()),
None => Ok(()),
}
}
type Handle = Arc<RwLock<AnnIndex>>;
fn registry() -> &'static Mutex<HashMap<PathBuf, Handle>> {
static REG: OnceLock<Mutex<HashMap<PathBuf, Handle>>> = OnceLock::new();
REG.get_or_init(|| Mutex::new(HashMap::new()))
}
fn build_lock_for(key: &Path) -> Arc<Mutex<()>> {
static LOCKS: OnceLock<Mutex<HashMap<PathBuf, Arc<Mutex<()>>>>> = OnceLock::new();
let m = LOCKS.get_or_init(|| Mutex::new(HashMap::new()));
let mut g = m.lock().unwrap_or_else(|p| p.into_inner());
g.entry(key.to_path_buf())
.or_insert_with(|| Arc::new(Mutex::new(())))
.clone()
}
fn spawn_save(handle: Handle) {
std::thread::spawn(move || {
let guard = handle.read().unwrap_or_else(|p| p.into_inner());
if let Err(e) = guard.save() {
eprintln!("kimetsu-brain: ann background save failed: {e}");
}
});
}
pub fn handle_for_query(conn: &Connection, dim: usize, model_id: &str) -> KimetsuResult<Handle> {
let Some(key) = AnnIndex::sidecar_for(conn) else {
return Ok(Arc::new(RwLock::new(AnnIndex::build_from_conn(
conn, dim, model_id,
)?)));
};
let handle = get_or_build_handle(&key, conn, dim, model_id)?;
reconcile_if_stale(&handle, conn)?;
Ok(handle)
}
fn get_or_build_handle(
key: &Path,
conn: &Connection,
dim: usize,
model_id: &str,
) -> KimetsuResult<Handle> {
{
let reg = registry().lock().unwrap_or_else(|p| p.into_inner());
if let Some(h) = reg.get(key) {
return Ok(h.clone());
}
}
let bl = build_lock_for(key);
let _g = bl.lock().unwrap_or_else(|p| p.into_inner());
{
let reg = registry().lock().unwrap_or_else(|p| p.into_inner());
if let Some(h) = reg.get(key) {
return Ok(h.clone());
}
}
let idx = AnnIndex::open_or_build(conn, dim, model_id)?;
let handle: Handle = Arc::new(RwLock::new(idx));
registry()
.lock()
.unwrap_or_else(|p| p.into_inner())
.insert(key.to_path_buf(), handle.clone());
spawn_save(handle.clone());
Ok(handle)
}
fn reconcile_if_stale(handle: &Handle, conn: &Connection) -> KimetsuResult<()> {
let stale = {
let idx = handle.read().unwrap_or_else(|p| p.into_inner());
idx.is_stale(conn)?
};
if stale {
let mut idx = handle.write().unwrap_or_else(|p| p.into_inner());
if idx.is_stale(conn)? {
idx.reconcile(conn)?;
}
}
Ok(())
}
pub fn warm(conn: &Connection, dim: usize, model_id: &str) -> KimetsuResult<()> {
handle_for_query(conn, dim, model_id).map(|_| ())
}
pub fn cached_handle(conn: &Connection) -> Option<Handle> {
let sidecar = AnnIndex::sidecar_for(conn)?;
let reg = registry().lock().unwrap_or_else(|p| p.into_inner());
reg.get(&sidecar).cloned()
}
pub fn on_supersede(conn: &Connection, memory_id: &str) {
on_invalidate(conn, memory_id);
}
pub fn on_invalidate(conn: &Connection, memory_id: &str) {
let Some(handle) = cached_handle(conn) else {
return;
};
let rowid: Option<i64> = conn
.query_row(
"SELECT rowid FROM memories WHERE memory_id = ?1",
rusqlite::params![memory_id],
|r| r.get(0),
)
.ok();
if let Some(rowid) = rowid {
let mut guard = handle.write().unwrap_or_else(|p| p.into_inner());
let _ = guard.remove(rowid);
}
}
pub fn invalidate_sidecar(conn: &Connection) {
if let Some(sidecar) = AnnIndex::sidecar_for(conn) {
registry()
.lock()
.unwrap_or_else(|p| p.into_inner())
.remove(&sidecar);
let _ = std::fs::remove_file(&sidecar);
let _ = std::fs::remove_file(AnnIndex::manifest_path(&sidecar));
}
}
pub fn save_all() {
let reg = registry().lock().unwrap_or_else(|p| p.into_inner());
for handle in reg.values() {
let guard = handle.read().unwrap_or_else(|p| p.into_inner());
if let Err(e) = guard.save() {
eprintln!("kimetsu-brain: ann save_all failed: {e}");
}
}
}
pub struct AnnIndex {
index: Index,
dim: usize,
model_id: String,
sidecar: Option<PathBuf>,
max_rowid_indexed: i64,
}
impl AnnIndex {
pub fn len(&self) -> usize {
self.index.size()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn is_stale(&self, conn: &Connection) -> KimetsuResult<bool> {
let max_rowid: i64 =
conn.query_row("SELECT COALESCE(MAX(rowid), 0) FROM memories", [], |r| {
r.get(0)
})?;
Ok(max_rowid > self.max_rowid_indexed)
}
pub fn build_from_conn(conn: &Connection, dim: usize, model_id: &str) -> KimetsuResult<Self> {
let index = Index::new(&index_options(dim)).map_err(|e| format!("usearch new: {e}"))?;
let mut me = Self {
index,
dim,
model_id: model_id.to_string(),
sidecar: None,
max_rowid_indexed: 0,
};
me.reserve_and_load_active(conn)?;
Ok(me)
}
fn reserve_and_load_active(&mut self, conn: &Connection) -> KimetsuResult<()> {
let count: i64 = conn.query_row(
"SELECT COUNT(*) FROM memories
WHERE invalidated_at IS NULL AND superseded_by IS NULL
AND embedding IS NOT NULL AND embedding_model = ?1",
rusqlite::params![self.model_id],
|r| r.get(0),
)?;
if count > 0 {
self.index
.reserve(count as usize)
.map_err(|e| format!("usearch reserve: {e}"))?;
}
const BUILD_CHUNK: usize = 16384;
let mut stmt = conn.prepare(
"SELECT rowid, embedding FROM memories
WHERE invalidated_at IS NULL AND superseded_by IS NULL
AND embedding IS NOT NULL AND embedding_model = ?1
ORDER BY rowid",
)?;
let mut rows_iter = stmt.query(rusqlite::params![self.model_id])?;
let mut batch: Vec<(i64, Vec<f32>)> = Vec::with_capacity(BUILD_CHUNK);
let mut max_rowid = self.max_rowid_indexed;
loop {
let row = rows_iter.next()?;
let done = row.is_none();
if let Some(row) = row {
let rowid: i64 = row.get(0)?;
let blob: Vec<u8> = row.get(1)?;
if rowid > max_rowid {
max_rowid = rowid;
}
if blob.len() == self.dim * 4
&& let Ok(vec) = crate::embeddings::decode_embedding(&blob, Some(self.dim))
{
batch.push((rowid, vec));
}
}
if batch.len() >= BUILD_CHUNK || (done && !batch.is_empty()) {
parallel_add(&self.index, &batch)?;
batch.clear();
}
if done {
break;
}
}
self.max_rowid_indexed = max_rowid;
Ok(())
}
pub fn search(&self, query: &[f32], k: usize) -> KimetsuResult<Vec<(i64, f32)>> {
if k == 0 || self.is_empty() {
return Ok(Vec::new());
}
let matches = self
.index
.search(query, k)
.map_err(|e| format!("usearch search: {e}"))?;
Ok(matches
.keys
.into_iter()
.zip(matches.distances)
.map(|(key, dist)| (key as i64, dist))
.collect())
}
pub fn add(&mut self, rowid: i64, vector: &[f32]) -> KimetsuResult<()> {
if vector.len() != self.dim {
return Err(format!("ann add: dim {} != index dim {}", vector.len(), self.dim).into());
}
if self.index.contains(rowid as u64) {
self.index
.remove(rowid as u64)
.map_err(|e| format!("usearch remove (upsert): {e}"))?;
}
if self.index.size() + 1 > self.index.capacity() {
self.index
.reserve((self.index.capacity() + 1).max(64) * 2)
.map_err(|e| format!("usearch reserve (grow): {e}"))?;
}
self.index
.add(rowid as u64, vector)
.map_err(|e| format!("usearch add: {e}"))?;
if rowid > self.max_rowid_indexed {
self.max_rowid_indexed = rowid;
}
Ok(())
}
pub fn remove(&mut self, rowid: i64) -> KimetsuResult<()> {
if self.index.contains(rowid as u64) {
self.index
.remove(rowid as u64)
.map_err(|e| format!("usearch remove: {e}"))?;
}
Ok(())
}
fn sidecar_for(conn: &Connection) -> Option<PathBuf> {
match conn.path() {
Some(p) if !p.is_empty() && p != ":memory:" => {
let db = std::fs::canonicalize(p).unwrap_or_else(|_| PathBuf::from(p));
Some(db.with_extension("usearch"))
}
_ => None,
}
}
fn manifest_path(sidecar: &Path) -> PathBuf {
let mut s = sidecar.as_os_str().to_owned();
s.push(".json");
PathBuf::from(s)
}
fn tmp_sibling(path: &Path) -> PathBuf {
static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let seq = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut name = path.file_name().map(|n| n.to_owned()).unwrap_or_default();
name.push(format!(".{}-{}.tmp", std::process::id(), seq));
path.with_file_name(name)
}
fn manifest(&self) -> Manifest {
Manifest {
schema_version: SCHEMA_VERSION,
dim: self.dim,
model_id: self.model_id.clone(),
max_rowid_indexed: self.max_rowid_indexed,
count: self.len(),
quant: scalar_kind_id(ann_scalar_kind()).to_string(),
}
}
pub fn save(&self) -> KimetsuResult<()> {
let Some(sidecar) = &self.sidecar else {
return Ok(());
};
let index_tmp = Self::tmp_sibling(sidecar);
self.index
.save(index_tmp.to_string_lossy().as_ref())
.map_err(|e| format!("usearch save: {e}"))?;
std::fs::rename(&index_tmp, sidecar).map_err(|e| {
let _ = std::fs::remove_file(&index_tmp);
format!("usearch rename: {e}")
})?;
let manifest_path = Self::manifest_path(sidecar);
let manifest_tmp = Self::tmp_sibling(&manifest_path);
let manifest =
serde_json::to_vec(&self.manifest()).map_err(|e| format!("manifest serialize: {e}"))?;
std::fs::write(&manifest_tmp, manifest).map_err(|e| format!("manifest write: {e}"))?;
std::fs::rename(&manifest_tmp, &manifest_path).map_err(|e| {
let _ = std::fs::remove_file(&manifest_tmp);
format!("manifest rename: {e}")
})?;
Ok(())
}
pub fn open_or_build(conn: &Connection, dim: usize, model_id: &str) -> KimetsuResult<Self> {
let sidecar = Self::sidecar_for(conn);
if let Some(path) = &sidecar
&& path.exists()
&& let Some(loaded) = Self::try_load(path, dim, model_id)?
{
let mut idx = loaded;
idx.reconcile(conn)?;
return Ok(idx);
}
let mut idx = Self::build_from_conn(conn, dim, model_id)?;
idx.sidecar = sidecar;
Ok(idx)
}
fn try_load(sidecar: &Path, dim: usize, model_id: &str) -> KimetsuResult<Option<Self>> {
let manifest_bytes = match std::fs::read(Self::manifest_path(sidecar)) {
Ok(b) => b,
Err(_) => return Ok(None),
};
let manifest: Manifest = match serde_json::from_slice(&manifest_bytes) {
Ok(m) => m,
Err(_) => return Ok(None),
};
if manifest.schema_version != SCHEMA_VERSION
|| manifest.dim != dim
|| manifest.model_id != model_id
|| manifest.quant != scalar_kind_id(ann_scalar_kind())
{
return Ok(None);
}
let index = Index::new(&index_options(dim)).map_err(|e| format!("usearch new: {e}"))?;
if index.load(sidecar.to_string_lossy().as_ref()).is_err() {
return Ok(None); }
if index.size() != manifest.count {
return Ok(None);
}
Ok(Some(Self {
index,
dim,
model_id: model_id.to_string(),
sidecar: Some(sidecar.to_path_buf()),
max_rowid_indexed: manifest.max_rowid_indexed,
}))
}
pub fn reconcile(&mut self, conn: &Connection) -> KimetsuResult<()> {
const RECONCILE_CHUNK: usize = 16384;
let delta_count: i64 = conn.query_row(
"SELECT COUNT(*) FROM memories
WHERE invalidated_at IS NULL AND superseded_by IS NULL
AND embedding IS NOT NULL
AND embedding_model = ?1 AND rowid > ?2",
rusqlite::params![self.model_id, self.max_rowid_indexed],
|r| r.get(0),
)?;
if delta_count > 0 {
self.index
.reserve(self.index.size() + delta_count as usize)
.map_err(|e| format!("usearch reserve: {e}"))?;
}
let mut stmt = conn.prepare(
"SELECT rowid, embedding FROM memories
WHERE invalidated_at IS NULL AND superseded_by IS NULL
AND embedding IS NOT NULL
AND embedding_model = ?1 AND rowid > ?2 ORDER BY rowid",
)?;
let mut rows_iter = stmt.query(rusqlite::params![self.model_id, self.max_rowid_indexed])?;
let mut batch: Vec<(i64, Vec<f32>)> = Vec::with_capacity(RECONCILE_CHUNK);
let mut max_rowid = self.max_rowid_indexed;
loop {
let row = rows_iter.next()?;
let done = row.is_none();
if let Some(row) = row {
let rowid: i64 = row.get(0)?;
let blob: Vec<u8> = row.get(1)?;
if rowid > max_rowid {
max_rowid = rowid;
}
if blob.len() == self.dim * 4
&& let Ok(vec) = crate::embeddings::decode_embedding(&blob, Some(self.dim))
{
batch.push((rowid, vec));
}
}
if batch.len() >= RECONCILE_CHUNK || (done && !batch.is_empty()) {
parallel_add(&self.index, &batch)?;
batch.clear();
}
if done {
break;
}
}
drop(rows_iter);
drop(stmt);
self.max_rowid_indexed = max_rowid;
let gone: Vec<i64> = {
let mut stmt = conn.prepare(
"SELECT rowid FROM memories
WHERE (invalidated_at IS NOT NULL OR superseded_by IS NOT NULL)
AND rowid <= ?1",
)?;
stmt.query_map(rusqlite::params![self.max_rowid_indexed], |r| {
r.get::<_, i64>(0)
})?
.filter_map(|r| r.ok())
.collect()
};
for rowid in gone {
self.remove(rowid)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embeddings::encode_embedding;
fn with_quant<R>(val: Option<&str>, f: impl FnOnce() -> R) -> R {
let _guard = crate::user_brain::test_env_lock()
.lock()
.unwrap_or_else(|p| p.into_inner());
let prev = std::env::var("KIMETSU_ANN_QUANTIZATION").ok();
unsafe {
match val {
Some(v) => std::env::set_var("KIMETSU_ANN_QUANTIZATION", v),
None => std::env::remove_var("KIMETSU_ANN_QUANTIZATION"),
}
}
let out = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
unsafe {
match prev {
Some(v) => std::env::set_var("KIMETSU_ANN_QUANTIZATION", v),
None => std::env::remove_var("KIMETSU_ANN_QUANTIZATION"),
}
}
match out {
Ok(r) => r,
Err(e) => std::panic::resume_unwind(e),
}
}
fn seed_random(n: usize, dim: usize, model: &str) -> (Connection, Vec<(i64, Vec<f32>)>) {
use crate::embeddings::decode_embedding;
let conn = Connection::open_in_memory().expect("open");
crate::schema::initialize(&conn).expect("init");
let mut state: u64 = 0x9E3779B97F4A7C15;
let mut next = || {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as f32 / (1u64 << 31) as f32) - 1.0
};
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| next()).collect();
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,?4)",
rusqlite::params![format!("m-{i:06}"), "t", encode_embedding(&v), model],
)
.expect("insert");
}
let mut stmt = conn
.prepare("SELECT rowid, embedding FROM memories")
.unwrap();
let rows = stmt
.query_map([], |r| Ok((r.get::<_, i64>(0)?, r.get::<_, Vec<u8>>(1)?)))
.unwrap();
let mut vectors: Vec<(i64, Vec<f32>)> = Vec::new();
for row in rows {
let (rowid, blob) = row.unwrap();
vectors.push((rowid, decode_embedding(&blob, Some(dim)).unwrap()));
}
drop(stmt);
(conn, vectors)
}
fn measure_recall(idx: &AnnIndex, vectors: &[(i64, Vec<f32>)], k: usize, trials: usize) -> f32 {
use crate::embeddings::cosine_similarity;
let mut hit = 0usize;
let mut total = 0usize;
for t in 0..trials {
let q = &vectors[t * 7 % vectors.len()].1;
let mut scored: Vec<(i64, f32)> = vectors
.iter()
.map(|(id, v)| (*id, cosine_similarity(q, v)))
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let exact: std::collections::HashSet<i64> =
scored.iter().take(k).map(|(id, _)| *id).collect();
let ann: std::collections::HashSet<i64> = idx
.search(q, k)
.unwrap()
.into_iter()
.map(|(id, _)| id)
.collect();
hit += exact.intersection(&ann).count();
total += k;
}
hit as f32 / total as f32
}
#[test]
fn default_quant_is_f16() {
with_quant(None, || {
assert!(matches!(ann_scalar_kind(), ScalarKind::F16));
assert_eq!(scalar_kind_id(ScalarKind::F16), "f16");
assert_eq!(scalar_kind_id(ScalarKind::F32), "f32");
assert_eq!(scalar_kind_id(ScalarKind::I8), "i8");
});
}
#[test]
fn recall_guard_holds_under_f16() {
with_quant(Some("f16"), || {
let dim = 16;
let (conn, vectors) = seed_random(5000, dim, "stub");
let idx = AnnIndex::build_from_conn(&conn, dim, "stub").expect("build");
let recall = measure_recall(&idx, &vectors, 10, 50);
assert!(recall >= 0.95, "f16 recall@10 = {recall} (want >= 0.95)");
});
}
#[test]
fn i8_quant_builds_and_searches() {
with_quant(Some("i8"), || {
assert!(matches!(ann_scalar_kind(), ScalarKind::I8));
let dim = 16;
let (conn, vectors) = seed_random(5000, dim, "stub");
let idx = AnnIndex::build_from_conn(&conn, dim, "stub").expect("build");
assert_eq!(idx.len(), 5000, "i8 index covers all rows");
let hits = idx.search(&vectors[0].1, 10).expect("search");
assert_eq!(hits.len(), 10, "i8 search returns k results");
let recall = measure_recall(&idx, &vectors, 10, 50);
assert!(recall >= 0.85, "i8 recall@10 = {recall} (want >= 0.85)");
});
}
#[test]
fn manifest_quant_mismatch_forces_rebuild() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
with_quant(Some("f16"), || {
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
for i in 0..10usize {
insert_row(&conn, i, dim);
}
let idx = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("build");
idx.save().expect("save");
assert_eq!(idx.manifest().quant, "f16");
assert!(db.with_extension("usearch").exists(), "f16 sidecar written");
});
with_quant(Some("i8"), || {
let conn = Connection::open(&db).expect("open");
let sidecar = db.with_extension("usearch");
assert!(
AnnIndex::try_load(&sidecar, dim, "stub-d8")
.expect("try_load")
.is_none(),
"f16 sidecar must be rejected when i8 is active"
);
let idx = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("rebuild");
assert_eq!(idx.manifest().quant, "i8");
assert_eq!(idx.len(), 10, "rebuilt i8 index covers all rows");
});
}
fn seed_conn(n: usize, dim: usize, model: &str) -> Connection {
let conn = Connection::open_in_memory().expect("open");
crate::schema::initialize(&conn).expect("init");
for i in 0..n {
let mut v = vec![0.01f32; dim];
v[i % dim] = 1.0;
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,?4)",
rusqlite::params![
format!("m-{i:06}"),
format!("text {i}"),
encode_embedding(&v),
model
],
)
.expect("insert");
}
conn
}
#[test]
fn build_from_conn_indexes_all_active_rows() {
let dim = 8;
let conn = seed_conn(50, dim, "stub-d8");
let idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
assert_eq!(idx.len(), 50, "all 50 active rows indexed");
}
#[test]
fn search_returns_nearest_rowid_first() {
with_quant(Some("f16"), || {
let dim = 8;
let conn = seed_conn(dim, dim, "stub-d8"); let idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
let mut q = vec![0.0f32; dim];
q[3] = 1.0;
let hits = idx.search(&q, 3).expect("search");
assert!(!hits.is_empty(), "got candidates");
let (rowid, _dist) = hits[0];
let mid: String = conn
.query_row(
"SELECT memory_id FROM memories WHERE rowid = ?1",
rusqlite::params![rowid],
|r| r.get(0),
)
.expect("map rowid");
assert_eq!(mid, "m-000003");
});
}
#[test]
fn add_is_upsert_and_remove_drops() {
let dim = 8;
let conn = seed_conn(4, dim, "stub-d8");
let mut idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
assert_eq!(idx.len(), 4);
let mut v = vec![0.0f32; dim];
v[0] = 1.0;
idx.add(1, &v).expect("upsert");
assert_eq!(idx.len(), 4, "upsert must not grow the index");
idx.add(999, &v).expect("add new");
assert_eq!(idx.len(), 5);
idx.remove(999).expect("remove");
assert_eq!(idx.len(), 4);
}
#[test]
fn save_then_open_reuses_sidecar_and_search_matches() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open file db");
crate::schema::initialize(&conn).expect("init");
for i in 0..20usize {
let mut v = vec![0.01f32; dim];
v[i % dim] = 1.0;
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,'stub-d8')",
rusqlite::params![format!("m-{i:06}"), format!("t{i}"), crate::embeddings::encode_embedding(&v)],
).expect("insert");
}
let idx = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("build");
idx.save().expect("save");
assert!(db.with_extension("usearch").exists(), "sidecar written");
let idx2 = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("load");
assert_eq!(idx2.len(), 20);
let mut q = vec![0.0f32; dim];
q[2] = 1.0;
assert!(!idx2.search(&q, 5).expect("search").is_empty());
}
#[test]
fn manifest_model_mismatch_forces_rebuild() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
AnnIndex::open_or_build(&conn, dim, "model-a")
.expect("a")
.save()
.expect("save");
let idx = AnnIndex::open_or_build(&conn, dim, "model-b").expect("b");
assert_eq!(idx.len(), 0, "rebuilt for model-b which has no rows");
}
#[test]
fn reconcile_adds_new_and_removes_invalidated() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
let insert = |conn: &Connection, i: usize| {
let mut v = vec![0.01f32; dim];
v[i % dim] = 1.0;
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,'stub-d8')",
rusqlite::params![format!("m-{i:06}"), format!("t{i}"), crate::embeddings::encode_embedding(&v)],
).expect("insert");
};
for i in 0..10 {
insert(&conn, i);
}
let idx = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("build");
idx.save().expect("save");
assert_eq!(idx.len(), 10);
for i in 10..15 {
insert(&conn, i);
}
conn.execute("UPDATE memories SET invalidated_at='2026-02-01T00:00:00Z' WHERE memory_id IN ('m-000000','m-000001')", []).expect("invalidate");
let idx2 = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("reopen");
assert_eq!(idx2.len(), 13);
}
#[test]
fn recall_at_10_is_at_least_0_95_vs_brute_force() {
with_quant(Some("f16"), || {
use crate::embeddings::{cosine_similarity, decode_embedding};
let dim = 16;
let n = 5000usize;
let conn = Connection::open_in_memory().expect("open");
crate::schema::initialize(&conn).expect("init");
let mut state: u64 = 0x9E3779B97F4A7C15;
let mut next = || {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as f32 / (1u64 << 31) as f32) - 1.0
};
let mut vectors: Vec<(i64, Vec<f32>)> = Vec::new();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| next()).collect();
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,'stub')",
rusqlite::params![format!("m-{i:06}"), "t", crate::embeddings::encode_embedding(&v)],
).expect("insert");
}
let mut stmt = conn
.prepare("SELECT rowid, embedding FROM memories")
.unwrap();
let rows = stmt
.query_map([], |r| Ok((r.get::<_, i64>(0)?, r.get::<_, Vec<u8>>(1)?)))
.unwrap();
for row in rows {
let (rowid, blob) = row.unwrap();
vectors.push((rowid, decode_embedding(&blob, Some(dim)).unwrap()));
}
let idx = AnnIndex::build_from_conn(&conn, dim, "stub").expect("build");
let trials = 50;
let k = 10;
let mut hit = 0usize;
let mut total = 0usize;
for t in 0..trials {
let q = &vectors[t * 7 % vectors.len()].1;
let mut scored: Vec<(i64, f32)> = vectors
.iter()
.map(|(id, v)| (*id, cosine_similarity(q, v)))
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let exact: std::collections::HashSet<i64> =
scored.iter().take(k).map(|(id, _)| *id).collect();
let ann: std::collections::HashSet<i64> = idx
.search(q, k)
.unwrap()
.into_iter()
.map(|(id, _)| id)
.collect();
hit += exact.intersection(&ann).count();
total += k;
}
let recall = hit as f32 / total as f32;
assert!(recall >= 0.95, "recall@10 = {recall} (want >= 0.95)");
});
}
fn insert_row(conn: &Connection, i: usize, dim: usize) {
let mut v = vec![0.01f32; dim];
v[i % dim] = 1.0;
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES (?1,'project','fact',?2,?2,1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?3,'stub-d8')",
rusqlite::params![format!("m-{i:06}"), format!("t{i}"), encode_embedding(&v)],
)
.expect("insert");
}
#[test]
fn parallel_build_indexes_all_rows() {
with_quant(Some("f16"), || {
let dim = 16;
let n = 2000;
let (conn, vectors) = seed_random(n, dim, "stub-d16");
let idx = AnnIndex::build_from_conn(&conn, dim, "stub-d16").expect("build");
assert_eq!(idx.len(), n, "all {n} rows indexed by parallel build");
let recall = measure_recall(&idx, &vectors, 10, 100);
assert!(
recall >= 0.9,
"parallel-built index recall@10 = {recall} (want >= 0.9)"
);
});
}
#[test]
fn is_stale_detects_new_rows() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
for i in 0..10 {
insert_row(&conn, i, dim);
}
let idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
assert!(!idx.is_stale(&conn).expect("stale check"), "fresh index");
for i in 10..15 {
insert_row(&conn, i, dim);
}
assert!(
idx.is_stale(&conn).expect("stale check"),
"stale after out-of-band inserts"
);
}
#[test]
fn malformed_embedding_rows_do_not_keep_index_stale() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
insert_row(&conn, 0, dim);
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES ('m-bad','project','fact','bad','bad',1.0,'{}',
'2026-01-01T00:00:00Z',0,0.0,?1,'stub-d8')",
rusqlite::params![vec![1_u8, 2, 3]],
)
.expect("insert malformed embedding");
let idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
assert_eq!(idx.len(), 1, "only the valid vector is indexed");
assert!(
!idx.is_stale(&conn).expect("stale check"),
"malformed high-rowid embeddings should not force repeated reconcile"
);
}
#[test]
fn reconcile_advances_watermark_past_malformed_delta() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
insert_row(&conn, 0, dim);
let mut idx = AnnIndex::build_from_conn(&conn, dim, "stub-d8").expect("build");
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES ('m-bad-delta','project','fact','bad','bad',1.0,'{}',
'2026-01-01T00:00:00Z',0,0.0,?1,'stub-d8')",
rusqlite::params![vec![1_u8, 2, 3]],
)
.expect("insert malformed embedding");
idx.reconcile(&conn).expect("reconcile");
assert_eq!(idx.len(), 1, "malformed delta is skipped");
assert!(
!idx.is_stale(&conn).expect("stale check"),
"malformed delta should be skipped once, not retried forever"
);
}
#[test]
fn handle_for_query_reconciles_cached_index_on_new_rows() {
let dim = 8;
let dir = tempfile::tempdir().expect("tmp");
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).expect("open");
crate::schema::initialize(&conn).expect("init");
let n = 8usize;
for i in 0..n {
insert_row(&conn, i, dim);
}
let h1 = handle_for_query(&conn, dim, "stub-d8").expect("build");
assert_eq!(h1.read().unwrap().len(), n, "initial build covers all rows");
let m = 5usize;
for i in n..n + m {
insert_row(&conn, i, dim);
}
let h2 = handle_for_query(&conn, dim, "stub-d8").expect("requery");
assert!(Arc::ptr_eq(&h1, &h2), "same cached handle reused");
assert_eq!(
h2.read().unwrap().len(),
n + m,
"cached index reconciled to include bulk-added rows"
);
}
#[test]
fn registry_caches_per_ondisk_db_and_transient_for_memory() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
let h1 = handle_for_query(&conn, dim, "stub-d8").unwrap();
let h2 = handle_for_query(&conn, dim, "stub-d8").unwrap();
assert!(Arc::ptr_eq(&h1, &h2), "same db → cached handle");
let mem = Connection::open_in_memory().unwrap();
crate::schema::initialize(&mem).unwrap();
let hm = handle_for_query(&mem, dim, "stub-d8").unwrap();
assert_eq!(hm.read().unwrap().len(), 0);
}
#[test]
fn on_invalidate_removes_from_cached_index() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
let mut v = vec![0.0f32; dim];
v[0] = 1.0;
conn.execute(
"INSERT INTO memories
(memory_id, scope, kind, text, normalized_text, confidence,
provenance_snapshot_json, created_at, use_count, usefulness_score,
embedding, embedding_model)
VALUES ('m-x','project','fact','t','t',1.0,'{}','2026-01-01T00:00:00Z',0,0.0,?1,'stub-d8')",
rusqlite::params![crate::embeddings::encode_embedding(&v)],
).unwrap();
let h = handle_for_query(&conn, dim, "stub-d8").unwrap();
assert_eq!(h.read().unwrap().len(), 1);
conn.execute(
"UPDATE memories SET invalidated_at='2026-02-01T00:00:00Z' WHERE memory_id='m-x'",
[],
)
.unwrap();
on_invalidate(&conn, "m-x");
assert_eq!(cached_handle(&conn).unwrap().read().unwrap().len(), 0);
}
#[test]
fn concurrent_same_key_builds_once() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
let n = 12usize;
for i in 0..n {
insert_row(&conn, i, dim);
}
drop(conn); let db_path = db.clone();
let threads = 8usize;
let mut handles = Vec::new();
for _ in 0..threads {
let p = db_path.clone();
handles.push(std::thread::spawn(move || {
let conn = Connection::open(&p).unwrap();
handle_for_query(&conn, dim, "stub-d8").unwrap()
}));
}
let results: Vec<Handle> = handles.into_iter().map(|h| h.join().unwrap()).collect();
let first = results[0].clone();
for h in &results[1..] {
assert!(
Arc::ptr_eq(&first, h),
"all concurrent builds must share ONE handle (built once)"
);
}
assert_eq!(
first.read().unwrap().len(),
n,
"shared index covers all rows"
);
}
#[test]
fn build_persists_sidecar() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
for i in 0..10 {
insert_row(&conn, i, dim);
}
let _h = handle_for_query(&conn, dim, "stub-d8").unwrap();
let sidecar = db.with_extension("usearch");
let mut appeared = false;
for _ in 0..100 {
if sidecar.exists() {
appeared = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(20));
}
assert!(
appeared,
"background save must persist the sidecar within ~2s"
);
let loaded = AnnIndex::open_or_build(&conn, dim, "stub-d8").expect("load sidecar");
assert_eq!(loaded.len(), 10, "reloaded sidecar covers all rows");
}
#[test]
fn warm_caches_handle() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
for i in 0..6 {
insert_row(&conn, i, dim);
}
assert!(cached_handle(&conn).is_none(), "cold before warm");
warm(&conn, dim, "stub-d8").expect("warm");
assert!(
cached_handle(&conn).is_some(),
"warm must build + cache the handle"
);
}
#[test]
fn invalidate_sidecar_removes_file_and_cache() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
let _ = handle_for_query(&conn, dim, "stub-d8").unwrap();
drop(
handle_for_query(&conn, dim, "stub-d8")
.unwrap()
.read()
.unwrap(),
);
cached_handle(&conn)
.unwrap()
.read()
.unwrap()
.save()
.unwrap();
assert!(db.with_extension("usearch").exists());
invalidate_sidecar(&conn);
assert!(!db.with_extension("usearch").exists());
assert!(cached_handle(&conn).is_none());
}
#[test]
fn concurrent_saves_do_not_collide() {
let dim = 8;
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("brain.db");
let conn = Connection::open(&db).unwrap();
crate::schema::initialize(&conn).unwrap();
let handle = handle_for_query(&conn, dim, "stub-d8").unwrap();
std::thread::scope(|s| {
for _ in 0..8 {
let h = handle.clone();
s.spawn(move || {
h.read().unwrap().save().unwrap();
});
}
});
assert!(db.with_extension("usearch").exists());
}
}