use crate::error::Result;
use rusqlite::params;
use std::collections::HashMap;
use super::YantrikDB;
fn decompress_if_cold(bytes: Vec<u8>) -> Vec<u8> {
if crate::compression::is_compressed(&bytes) {
crate::serde_helpers::serialize_f32(&crate::compression::decompress_embedding(&bytes))
} else {
bytes
}
}
#[derive(Debug, Clone)]
pub struct EmbeddingWithGeneration {
pub bytes: Vec<u8>,
pub generation: u64,
}
pub(crate) struct DurableEmbeddingStore<'a> {
db: &'a YantrikDB,
}
impl<'a> DurableEmbeddingStore<'a> {
pub fn new(db: &'a YantrikDB) -> Self {
DurableEmbeddingStore { db }
}
pub fn read_embeddings_for_rids(
&self,
rids: &[&str],
) -> Result<HashMap<String, EmbeddingWithGeneration>> {
if rids.is_empty() {
return Ok(HashMap::new());
}
let placeholders: String = (0..rids.len())
.map(|i| format!("?{}", i + 1))
.collect::<Vec<_>>()
.join(",");
let sql = format!(
"SELECT rid, embedding, COALESCE(embedding_generation, 0) \
FROM memories WHERE rid IN ({placeholders})"
);
let mut param_values: Vec<Box<dyn rusqlite::types::ToSql>> = Vec::new();
for r in rids {
param_values.push(Box::new(r.to_string()));
}
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|p| p.as_ref()).collect();
let conn = self.db.read_conn();
let mut stmt = conn.prepare(&sql)?;
let rows: Vec<(String, Option<Vec<u8>>, i64)> = stmt
.query_map(params_ref.as_slice(), |row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, Option<Vec<u8>>>(1)?,
row.get::<_, i64>(2)?,
))
})?
.collect::<std::result::Result<Vec<_>, _>>()?;
drop(stmt);
drop(conn);
let mut map = HashMap::new();
for (rid, stored_emb, generation) in rows {
let Some(stored_emb) = stored_emb else {
continue;
};
let bytes = self.db.decrypt_embedding(&stored_emb)?;
let bytes = decompress_if_cold(bytes);
map.insert(
rid,
EmbeddingWithGeneration {
bytes,
generation: generation as u64,
},
);
}
Ok(map)
}
pub fn read_embeddings_under_generation(
&self,
active_generation: u64,
limit: usize,
offset: usize,
) -> Result<Vec<(String, EmbeddingWithGeneration)>> {
let conn = self.db.read_conn();
let mut stmt = conn.prepare(
"SELECT rid, embedding, COALESCE(embedding_generation, 0) \
FROM memories \
WHERE consolidation_status = 'active' \
AND embedding IS NOT NULL \
AND COALESCE(embedding_generation, 0) < ?1 \
ORDER BY rid \
LIMIT ?2 OFFSET ?3",
)?;
let rows: Vec<(String, Vec<u8>, i64)> = stmt
.query_map(
params![active_generation as i64, limit as i64, offset as i64],
|row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, Vec<u8>>(1)?,
row.get::<_, i64>(2)?,
))
},
)?
.collect::<std::result::Result<Vec<_>, _>>()?;
drop(stmt);
drop(conn);
let mut out = Vec::with_capacity(rows.len());
for (rid, stored_emb, generation) in rows {
let bytes = self.db.decrypt_embedding(&stored_emb)?;
let bytes = decompress_if_cold(bytes);
out.push((
rid,
EmbeddingWithGeneration {
bytes,
generation: generation as u64,
},
));
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::YantrikDB;
fn vec_seed(seed: f32, dim: usize) -> Vec<f32> {
(0..dim).map(|i| seed + (i as f32) * 0.001).collect()
}
#[test]
fn read_embeddings_for_rids_returns_bytes_and_generation() {
let db = YantrikDB::new(":memory:", 8).unwrap();
let rid = db
.record(
"durable embed read",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec_seed(0.5, 8),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let store = DurableEmbeddingStore::new(&db);
let map = store.read_embeddings_for_rids(&[&rid]).unwrap();
assert_eq!(map.len(), 1, "exact match");
let entry = map.get(&rid).unwrap();
assert!(!entry.bytes.is_empty(), "bytes returned");
assert_eq!(entry.generation, 0, "fresh engine writes rows at gen 0");
}
#[test]
fn read_embeddings_for_rids_reflects_generation_advance() {
let db = YantrikDB::new(":memory:", 8).unwrap();
let old_state = db.search_state.load_full();
let advanced = crate::engine::reembed::SearchState {
index_embedding: old_state.index_embedding.clone(),
embedder: old_state.embedder.clone(),
runtime_embedder_name: old_state.runtime_embedder_name.clone(),
runtime_embedder_digest: old_state.runtime_embedder_digest.clone(),
generation: 11,
covers_through_seq: old_state.covers_through_seq,
hnsw_m: old_state.hnsw_m,
hnsw_ef_construction: old_state.hnsw_ef_construction,
hnsw_ef_search: old_state.hnsw_ef_search,
vec_index: std::sync::Arc::clone(&old_state.vec_index),
};
db.try_publish_search_state(advanced).unwrap();
let rid = db
.record(
"post-advance",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec_seed(0.5, 8),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let store = DurableEmbeddingStore::new(&db);
let map = store.read_embeddings_for_rids(&[&rid]).unwrap();
assert_eq!(map.get(&rid).unwrap().generation, 11);
}
#[test]
fn read_embeddings_under_generation_filters_by_strict_less_than() {
let db = YantrikDB::new(":memory:", 8).unwrap();
let rid_old = db
.record(
"old gen",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec_seed(0.1, 8),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let old_state = db.search_state.load_full();
let advanced = crate::engine::reembed::SearchState {
index_embedding: old_state.index_embedding.clone(),
embedder: old_state.embedder.clone(),
runtime_embedder_name: old_state.runtime_embedder_name.clone(),
runtime_embedder_digest: old_state.runtime_embedder_digest.clone(),
generation: 5,
covers_through_seq: old_state.covers_through_seq,
hnsw_m: old_state.hnsw_m,
hnsw_ef_construction: old_state.hnsw_ef_construction,
hnsw_ef_search: old_state.hnsw_ef_search,
vec_index: std::sync::Arc::clone(&old_state.vec_index),
};
db.try_publish_search_state(advanced).unwrap();
let rid_new = db
.record(
"new gen",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec_seed(0.2, 8),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let store = DurableEmbeddingStore::new(&db);
let rows = store.read_embeddings_under_generation(5, 1000, 0).unwrap();
let rids: Vec<&str> = rows.iter().map(|(r, _)| r.as_str()).collect();
assert!(
rids.contains(&rid_old.as_str()),
"rid_old (gen 0 < 5) must appear"
);
assert!(
!rids.contains(&rid_new.as_str()),
"rid_new (gen 5 not strictly less than 5) must NOT appear"
);
let rows = store.read_embeddings_under_generation(6, 1000, 0).unwrap();
let rids: Vec<&str> = rows.iter().map(|(r, _)| r.as_str()).collect();
assert!(rids.contains(&rid_old.as_str()));
assert!(rids.contains(&rid_new.as_str()));
}
#[test]
fn recall_rs_has_no_raw_sql_embedding_reads() {
let src = include_str!("recall.rs");
let mut offenders: Vec<&str> = Vec::new();
for line in src.lines() {
let trimmed = line.trim();
if !(trimmed.starts_with("//") || trimmed.starts_with("///")) {
let lower = trimmed.to_ascii_lowercase();
for cand in [
"select embedding ",
"select embedding,",
"select embedding\"",
"select embedding\\",
", embedding ",
", embedding,",
", embedding\"",
", embedding\\",
", embedding\n",
] {
if lower.contains(cand) {
offenders.push(line);
break;
}
}
}
}
assert!(
offenders.is_empty(),
"brainstorm-4 §5 audit: recall.rs grew {} raw SQL embedding read(s); \
route through engine::durable_embeddings::DurableEmbeddingStore. \
Offending line(s):\n{}",
offenders.len(),
offenders.join("\n")
);
}
}