use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ReembedWritePolicy {
Pause,
Queue,
}
impl Default for ReembedWritePolicy {
fn default() -> Self {
ReembedWritePolicy::Queue
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ReembedPhase {
Probing,
Encoding,
Rebuilding,
Swapping,
Verifying,
Completed,
Aborted,
}
impl ReembedPhase {
pub fn as_str(&self) -> &'static str {
match self {
ReembedPhase::Probing => "Probing",
ReembedPhase::Encoding => "Encoding",
ReembedPhase::Rebuilding => "Rebuilding",
ReembedPhase::Swapping => "Swapping",
ReembedPhase::Verifying => "Verifying",
ReembedPhase::Completed => "Completed",
ReembedPhase::Aborted => "Aborted",
}
}
pub fn parse(s: &str) -> Option<Self> {
match s {
"Probing" => Some(ReembedPhase::Probing),
"Encoding" => Some(ReembedPhase::Encoding),
"Rebuilding" => Some(ReembedPhase::Rebuilding),
"Swapping" => Some(ReembedPhase::Swapping),
"Verifying" => Some(ReembedPhase::Verifying),
"Completed" => Some(ReembedPhase::Completed),
"Aborted" => Some(ReembedPhase::Aborted),
_ => None,
}
}
pub fn is_terminal(&self) -> bool {
matches!(self, ReembedPhase::Completed | ReembedPhase::Aborted)
}
}
impl fmt::Display for ReembedPhase {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReembedProgress {
pub phase: ReembedPhase,
pub processed: u64,
pub total: Option<u64>,
pub elapsed_ms: u64,
pub namespace: Option<String>,
}
pub struct ReembedOptions {
pub namespace: Option<String>,
pub progress_cb: Option<Box<dyn Fn(ReembedProgress) + Send + Sync>>,
pub on_phase_complete: Option<Box<dyn Fn(ReembedPhase, &ReembedStatus) + Send + Sync>>,
pub batch_size: usize,
pub write_policy: ReembedWritePolicy,
pub hnsw_m: Option<u32>,
pub hnsw_ef_construction: Option<u32>,
pub hnsw_ef_search: Option<u32>,
pub resume_from_checkpoint: bool,
pub dry_run: bool,
}
impl Default for ReembedOptions {
fn default() -> Self {
ReembedOptions {
namespace: None,
progress_cb: None,
on_phase_complete: None,
batch_size: 256,
write_policy: ReembedWritePolicy::default(),
hnsw_m: None,
hnsw_ef_construction: None,
hnsw_ef_search: None,
resume_from_checkpoint: true,
dry_run: false,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct NamespaceReembedStats {
pub encoded_count: u64,
pub skipped_count: u64,
pub duration_ms: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReembedReport {
pub generation: u64,
pub encoded_count: u64,
pub skipped_count: u64,
pub duration: Duration,
pub old_embedder: String,
pub old_embedder_digest: String,
pub new_embedder: String,
pub new_embedder_digest: String,
pub old_dim: usize,
pub new_dim: usize,
pub build_hwm: u64,
pub per_namespace: HashMap<String, NamespaceReembedStats>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReembedStatus {
pub generation: u64,
pub phase: ReembedPhase,
pub old_embedder: String,
pub old_embedder_digest: String,
pub new_embedder: String,
pub new_embedder_digest: String,
pub old_dim: usize,
pub new_dim: usize,
pub memories_total: u64,
pub memories_encoded: u64,
pub queued_writes: u64,
pub checkpoint_rid: Option<String>,
pub started_at: SystemTime,
pub last_event_at: SystemTime,
pub last_error: Option<String>,
pub write_policy: ReembedWritePolicy,
}
#[derive(Debug, Clone)]
pub enum EmbeddingProvenance {
Known {
name: Option<String>,
digest: String,
dim: usize,
},
ExternalOrUnknown { dim: usize },
}
impl EmbeddingProvenance {
pub fn dim(&self) -> usize {
match self {
EmbeddingProvenance::Known { dim, .. } => *dim,
EmbeddingProvenance::ExternalOrUnknown { dim } => *dim,
}
}
pub fn digest(&self) -> Option<&str> {
match self {
EmbeddingProvenance::Known { digest, .. } => Some(digest),
EmbeddingProvenance::ExternalOrUnknown { .. } => None,
}
}
pub fn name(&self) -> Option<&str> {
match self {
EmbeddingProvenance::Known { name, .. } => name.as_deref(),
EmbeddingProvenance::ExternalOrUnknown { .. } => None,
}
}
}
pub struct SearchState {
pub index_embedding: EmbeddingProvenance,
pub embedder: Option<Arc<dyn crate::types::Embedder + Send + Sync>>,
pub runtime_embedder_name: Option<String>,
pub runtime_embedder_digest: Option<String>,
pub generation: u64,
pub covers_through_seq: u64,
pub hnsw_m: u32,
pub hnsw_ef_construction: u32,
pub hnsw_ef_search: u32,
pub vec_index: Arc<crate::vector::delta_index::DeltaIndex>,
}
impl SearchState {
pub fn initial(
dim: usize,
hnsw_m: u32,
hnsw_ef_construction: u32,
hnsw_ef_search: u32,
vec_index: Arc<crate::vector::delta_index::DeltaIndex>,
) -> Self {
SearchState {
index_embedding: EmbeddingProvenance::ExternalOrUnknown { dim },
embedder: None,
runtime_embedder_name: None,
runtime_embedder_digest: None,
generation: 0,
covers_through_seq: 0,
hnsw_m,
hnsw_ef_construction,
hnsw_ef_search,
vec_index,
}
}
pub fn dim(&self) -> usize {
self.index_embedding.dim()
}
pub fn has_runtime_embedder(&self) -> bool {
self.embedder.is_some()
}
}
impl fmt::Debug for SearchState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SearchState")
.field("index_embedding", &self.index_embedding)
.field("runtime_embedder_name", &self.runtime_embedder_name)
.field("runtime_embedder_digest", &self.runtime_embedder_digest)
.field("has_embedder", &self.embedder.is_some())
.field("generation", &self.generation)
.field("covers_through_seq", &self.covers_through_seq)
.field("hnsw_m", &self.hnsw_m)
.field("hnsw_ef_construction", &self.hnsw_ef_construction)
.field("hnsw_ef_search", &self.hnsw_ef_search)
.finish()
}
}
use rusqlite::params;
use crate::error::{Result, YantrikDbError};
use crate::serde_helpers::{deserialize_f32, serialize_f32};
use crate::vector::hnsw::HnswIndex;
use crate::YantrikDB;
impl YantrikDB {
pub fn reembed(
&self,
new_embedder_name: &str,
options: ReembedOptions,
) -> Result<ReembedReport> {
let probing_state_for_name_check = self.search_state.load_full();
let active_name_for_check = probing_state_for_name_check
.runtime_embedder_name
.clone()
.unwrap_or_default();
if active_name_for_check == new_embedder_name {
return self.reembed_with_embedder(new_embedder_name, None, options);
}
drop(probing_state_for_name_check);
#[cfg(feature = "embedder-download")]
let resolved: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> = {
use crate::embedder::DownloadedEmbedder;
let downloaded = DownloadedEmbedder::fetch(new_embedder_name).map_err(|e| {
YantrikDbError::Inference(format!(
"reembed: failed to resolve embedder {new_embedder_name:?}: {e}"
))
})?;
std::sync::Arc::new(downloaded)
};
#[cfg(not(feature = "embedder-download"))]
let resolved: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> = {
return Err(YantrikDbError::Inference(format!(
"reembed by name {new_embedder_name:?} requires the `embedder-download` \
cargo feature (enabled by default). Slim builds must call \
reembed_with_embedder() directly with a Box<dyn Embedder>."
)));
};
self.reembed_with_embedder(new_embedder_name, Some(resolved), options)
}
pub(crate) fn reembed_with_embedder(
&self,
new_embedder_name: &str,
pre_resolved: Option<std::sync::Arc<dyn crate::types::Embedder + Send + Sync>>,
options: ReembedOptions,
) -> Result<ReembedReport> {
let started_at = SystemTime::now();
let progress_cb = options.progress_cb.as_ref();
let on_phase_complete = options.on_phase_complete.as_ref();
let _index_guard = self.index_write_lock.lock();
let probing_state = self.search_state.load_full();
let active_dim = probing_state.dim();
let active_digest = probing_state.runtime_embedder_digest.clone();
let active_name = probing_state
.runtime_embedder_name
.clone()
.unwrap_or_default();
let same_name = active_name == new_embedder_name;
if same_name {
let duration = started_at.elapsed().unwrap_or_default();
let _ = progress_cb;
let _ = on_phase_complete;
return Ok(ReembedReport {
generation: probing_state.generation,
encoded_count: 0,
skipped_count: 0,
duration,
old_embedder: active_name.clone(),
old_embedder_digest: active_digest.clone().unwrap_or_default(),
new_embedder: active_name,
new_embedder_digest: active_digest.unwrap_or_default(),
old_dim: active_dim,
new_dim: active_dim,
build_hwm: probing_state.covers_through_seq,
per_namespace: HashMap::new(),
});
}
let next_generation = probing_state.generation + 1;
let probing_event_ts = systime_to_unix_secs(started_at);
self.write_reembed_event(
next_generation,
ReembedPhase::Probing,
probing_event_ts,
&serde_json::json!({
"old_embedder": active_name,
"new_embedder_name": new_embedder_name,
"old_dim": active_dim,
"namespace": options.namespace,
}),
)?;
let new_embedder = pre_resolved.ok_or_else(|| {
YantrikDbError::Inference(
"reembed_with_embedder: pre_resolved=None reached past same-name short-circuit \
— engine invariant violation"
.to_string(),
)
})?;
let new_dim = new_embedder.dim();
let new_embedder_digest = new_embedder.fingerprint().unwrap_or_default();
let new_embedder_name_resolved = new_embedder
.name()
.unwrap_or_else(|| new_embedder_name.to_string());
if new_dim != active_dim {
self.write_reembed_event(
next_generation,
ReembedPhase::Aborted,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"reason": "cross-dim reembed not yet supported",
"active_dim": active_dim,
"new_dim": new_dim,
}),
)?;
return Err(YantrikDbError::Inference(format!(
"reembed: new embedder {new_embedder_name_resolved:?} dim={new_dim} \
differs from active dim={active_dim}. Cross-dim reembed is not yet \
supported (engine's standalone embedding_dim field still gates \
record_with_rid/replication paths). Workaround: open a new database \
with YantrikDB::new(path, {new_dim}) and copy memories via export/import."
)));
}
let total_to_encode: u64 = {
let conn = self.read_conn();
let n: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memories \
WHERE consolidation_status = 'active' \
AND embedding IS NOT NULL \
AND COALESCE(embedding_generation, 0) < ?1",
params![next_generation as i64],
|r| r.get(0),
)
.unwrap_or(0);
n.max(0) as u64
};
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Probing,
processed: 0,
total: Some(total_to_encode),
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
if options.dry_run {
let duration = started_at.elapsed().unwrap_or_default();
return Ok(ReembedReport {
generation: probing_state.generation,
encoded_count: 0,
skipped_count: total_to_encode,
duration,
old_embedder: active_name.clone(),
old_embedder_digest: active_digest.clone().unwrap_or_default(),
new_embedder: new_embedder_name_resolved.clone(),
new_embedder_digest,
old_dim: active_dim,
new_dim,
build_hwm: probing_state.covers_through_seq,
per_namespace: HashMap::new(),
});
}
self.write_reembed_state_meta(&serde_json::json!({
"generation": next_generation,
"phase": "Encoding",
"old_embedder": active_name,
"new_embedder_name": new_embedder_name_resolved,
"old_dim": active_dim,
"new_dim": new_dim,
"started_at_unix": probing_event_ts,
"total_to_encode": total_to_encode,
"namespace": options.namespace,
"write_policy": match options.write_policy {
ReembedWritePolicy::Queue => "Queue",
ReembedWritePolicy::Pause => "Pause",
},
}))?;
{
let conn = self.conn();
conn.execute(
"UPDATE memories SET embedding_new = NULL, embedding_new_model = NULL \
WHERE embedding_new IS NOT NULL",
[],
)?;
}
let encoding_start_ts = systime_to_unix_secs(SystemTime::now());
self.write_reembed_event(
next_generation,
ReembedPhase::Encoding,
encoding_start_ts,
&serde_json::json!({
"total": total_to_encode,
"batch_size": options.batch_size,
}),
)?;
let batch_size = options.batch_size.max(1);
let mut processed: u64 = 0;
let mut offset: usize = 0;
loop {
let batch: Vec<(String, String)> = {
let conn = self.read_conn();
let mut stmt = conn.prepare(
"SELECT rid, text 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 = stmt
.query_map(
params![next_generation as i64, batch_size as i64, offset as i64],
|r| Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?)),
)?
.collect::<std::result::Result<Vec<_>, _>>()?;
drop(stmt);
drop(conn);
rows
};
if batch.is_empty() {
break;
}
let mut encoded_pairs: Vec<(String, Vec<u8>, String)> = Vec::with_capacity(batch.len());
for (rid, stored_text) in &batch {
let plain = self.decrypt_text(stored_text)?;
let new_emb = new_embedder.embed(&plain).map_err(|e| {
YantrikDbError::Inference(format!(
"reembed: embedder failed on rid {rid:?}: {e}"
))
})?;
if new_emb.len() != new_dim {
return Err(YantrikDbError::Inference(format!(
"reembed: embedder returned vector of len {} but reports dim {}; \
engine cannot trust this embedder",
new_emb.len(),
new_dim
)));
}
let blob = serialize_f32(&new_emb);
let encrypted = self.encrypt_embedding(&blob)?;
encoded_pairs.push((rid.clone(), encrypted, stored_text.clone()));
}
{
let conn = self.conn();
conn.execute_batch("SAVEPOINT reembed_encoding_batch")?;
let write_result: Result<()> = (|| {
for (rid, encrypted, read_text) in &encoded_pairs {
conn.execute(
"UPDATE memories SET embedding_new = ?1, embedding_new_model = ?2 \
WHERE rid = ?3 AND text = ?4",
params![
encrypted,
new_embedder_name_resolved.as_str(),
rid,
read_text
],
)?;
}
Ok(())
})();
match write_result {
Ok(()) => {
conn.execute_batch("RELEASE reembed_encoding_batch")?;
}
Err(e) => {
let _ = conn.execute_batch("ROLLBACK TO reembed_encoding_batch");
let _ = conn.execute_batch("RELEASE reembed_encoding_batch");
return Err(e);
}
}
}
processed += encoded_pairs.len() as u64;
offset += batch.len();
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Encoding,
processed,
total: Some(total_to_encode),
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
}
self.write_reembed_event(
next_generation,
ReembedPhase::Encoding,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"encoded_count": processed,
"completed": true,
}),
)?;
let new_hnsw_m = options.hnsw_m.unwrap_or(probing_state.hnsw_m);
let new_hnsw_efc = options
.hnsw_ef_construction
.unwrap_or(probing_state.hnsw_ef_construction);
let new_hnsw_efs = options
.hnsw_ef_search
.unwrap_or(probing_state.hnsw_ef_search);
let rebuilding_start_ts = systime_to_unix_secs(SystemTime::now());
self.write_reembed_event(
next_generation,
ReembedPhase::Rebuilding,
rebuilding_start_ts,
&serde_json::json!({
"expected_count": processed,
"hnsw_m": new_hnsw_m,
"hnsw_ef_construction": new_hnsw_efc,
"hnsw_ef_search": new_hnsw_efs,
}),
)?;
self.update_reembed_state_phase("Rebuilding")?;
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Rebuilding,
processed: 0,
total: Some(processed),
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
let mut new_hnsw = HnswIndex::with_params(
new_dim,
new_hnsw_m as usize,
new_hnsw_efc as usize,
new_hnsw_efs as usize,
);
let mut rebuilt: u64 = 0;
let mut rebuild_offset: usize = 0;
loop {
let batch: Vec<(String, Vec<u8>)> = {
let conn = self.read_conn();
let mut stmt = conn.prepare(
"SELECT rid, embedding_new FROM memories \
WHERE consolidation_status = 'active' \
AND embedding_new IS NOT NULL \
ORDER BY rid \
LIMIT ?1 OFFSET ?2",
)?;
let rows = stmt
.query_map(params![batch_size as i64, rebuild_offset as i64], |r| {
Ok((r.get::<_, String>(0)?, r.get::<_, Vec<u8>>(1)?))
})?
.collect::<std::result::Result<Vec<_>, _>>()?;
drop(stmt);
drop(conn);
rows
};
if batch.is_empty() {
break;
}
for (rid, stored_blob) in &batch {
let decrypted = self.decrypt_embedding(stored_blob)?;
let vec = deserialize_f32(&decrypted);
if vec.len() != new_dim {
return Err(YantrikDbError::Inference(format!(
"reembed Rebuilding: staged embedding_new for rid {rid:?} has len {} \
but expected new_dim {new_dim}",
vec.len()
)));
}
new_hnsw.insert(rid, &vec)?;
rebuilt += 1;
}
rebuild_offset += batch.len();
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Rebuilding,
processed: rebuilt,
total: Some(processed),
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
}
self.write_reembed_event(
next_generation,
ReembedPhase::Rebuilding,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"rebuilt_count": rebuilt,
"completed": true,
}),
)?;
let swapping_start_ts = systime_to_unix_secs(SystemTime::now());
self.write_reembed_event(
next_generation,
ReembedPhase::Swapping,
swapping_start_ts,
&serde_json::json!({}),
)?;
self.update_reembed_state_phase("Swapping")?;
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Swapping,
processed: 0,
total: None,
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
self.write_router.switch_to_queueing();
self.write_router.wait_for_no_sync_writers();
let covers_through_seq = self.vec_seq.load(std::sync::atomic::Ordering::Acquire);
let tail_rows: Vec<(String, String)> = {
let conn = self.read_conn();
let mut stmt = conn.prepare(
"SELECT rid, text FROM memories \
WHERE consolidation_status = 'active' \
AND embedding IS NOT NULL \
AND embedding_new IS NULL \
AND COALESCE(embedding_generation, 0) < ?1",
)?;
let rows = stmt
.query_map(params![next_generation as i64], |r| {
Ok((r.get::<_, String>(0)?, r.get::<_, String>(1)?))
})?
.collect::<std::result::Result<Vec<_>, _>>()?;
drop(stmt);
drop(conn);
rows
};
for (rid, stored_text) in &tail_rows {
let plain = self.decrypt_text(stored_text)?;
let new_emb = new_embedder.embed(&plain).map_err(|e| {
YantrikDbError::Inference(format!(
"reembed tail-catchup: embedder failed on rid {rid:?}: {e}"
))
})?;
if new_emb.len() != new_dim {
self.write_router.switch_to_normal();
return Err(YantrikDbError::Inference(format!(
"reembed tail-catchup: embedder returned vector of len {} but reports dim {}",
new_emb.len(),
new_dim
)));
}
let blob = serialize_f32(&new_emb);
let encrypted = self.encrypt_embedding(&blob)?;
{
let conn = self.conn();
conn.execute(
"UPDATE memories SET embedding_new = ?1, embedding_new_model = ?2 \
WHERE rid = ?3",
params![encrypted, new_embedder_name_resolved.as_str(), rid],
)?;
}
new_hnsw.insert(rid, &new_emb)?;
}
let tail_caught = tail_rows.len() as u64;
let total_swapped: i64 = {
let conn = self.conn();
conn.execute_batch("SAVEPOINT reembed_swap")?;
let result: Result<i64> = (|| {
conn.execute(
"INSERT OR REPLACE INTO meta (key, value) VALUES ('active_generation', ?1)",
params![next_generation.to_string()],
)?;
let n = conn.execute(
"UPDATE memories \
SET embedding = embedding_new, \
embedding_generation = ?1, \
embedding_new = NULL, \
embedding_new_model = NULL \
WHERE embedding_new IS NOT NULL",
params![next_generation as i64],
)?;
Ok(n as i64)
})();
match result {
Ok(n) => {
conn.execute_batch("RELEASE reembed_swap")?;
n
}
Err(e) => {
let _ = conn.execute_batch("ROLLBACK TO reembed_swap");
let _ = conn.execute_batch("RELEASE reembed_swap");
self.write_router.switch_to_normal();
return Err(e);
}
}
};
let new_delta_index = {
let delta_max = std::env::var("YANTRIKDB_DELTA_MAX")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(crate::vector::delta_index::DEFAULT_DELTA_MAX);
let max_dirty_age = std::env::var("YANTRIKDB_MAX_DIRTY_AGE_SECS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.map(std::time::Duration::from_secs)
.unwrap_or(crate::vector::delta_index::DEFAULT_MAX_DIRTY_AGE);
std::sync::Arc::new(crate::vector::delta_index::DeltaIndex::from_cold_with_age(
new_hnsw,
delta_max,
max_dirty_age,
))
};
let new_search_state = SearchState {
index_embedding: EmbeddingProvenance::Known {
name: Some(new_embedder_name_resolved.clone()),
digest: new_embedder_digest.clone(),
dim: new_dim,
},
embedder: Some(std::sync::Arc::clone(&new_embedder)),
runtime_embedder_name: Some(new_embedder_name_resolved.clone()),
runtime_embedder_digest: Some(new_embedder_digest.clone()),
generation: next_generation,
covers_through_seq,
hnsw_m: new_hnsw_m,
hnsw_ef_construction: new_hnsw_efc,
hnsw_ef_search: new_hnsw_efs,
vec_index: new_delta_index,
};
self.try_publish_search_state(new_search_state)?;
self.write_router.switch_to_normal();
self.write_reembed_event(
next_generation,
ReembedPhase::Swapping,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"covers_through_seq": covers_through_seq,
"tail_caught": tail_caught,
"rows_swapped": total_swapped,
"completed": true,
}),
)?;
let verifying_start_ts = systime_to_unix_secs(SystemTime::now());
self.write_reembed_event(
next_generation,
ReembedPhase::Verifying,
verifying_start_ts,
&serde_json::json!({}),
)?;
self.update_reembed_state_phase("Verifying")?;
let stragglers: i64 = {
let conn = self.read_conn();
conn.query_row(
"SELECT COUNT(*) FROM memories \
WHERE consolidation_status = 'active' \
AND embedding IS NOT NULL \
AND COALESCE(embedding_generation, 0) < ?1",
params![next_generation as i64],
|r| r.get(0),
)
.unwrap_or(0)
};
if let Some(cb) = progress_cb {
cb(ReembedProgress {
phase: ReembedPhase::Verifying,
processed: total_swapped as u64,
total: Some(total_swapped as u64),
elapsed_ms: started_at.elapsed().unwrap_or_default().as_millis() as u64,
namespace: options.namespace.clone(),
});
}
self.write_reembed_event(
next_generation,
ReembedPhase::Verifying,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"stragglers_under_old_gen": stragglers,
"completed": true,
}),
)?;
self.clear_reembed_state_meta()?;
self.write_reembed_event(
next_generation,
ReembedPhase::Completed,
systime_to_unix_secs(SystemTime::now()),
&serde_json::json!({
"encoded_count": processed,
"tail_caught": tail_caught,
"rows_swapped": total_swapped,
"covers_through_seq": covers_through_seq,
}),
)?;
let _ = on_phase_complete;
let duration = started_at.elapsed().unwrap_or_default();
Ok(ReembedReport {
generation: next_generation,
encoded_count: processed + tail_caught,
skipped_count: 0,
duration,
old_embedder: active_name,
old_embedder_digest: active_digest.unwrap_or_default(),
new_embedder: new_embedder_name_resolved,
new_embedder_digest,
old_dim: active_dim,
new_dim,
build_hwm: covers_through_seq,
per_namespace: HashMap::new(),
})
}
pub fn reembed_status(&self) -> Option<ReembedStatus> {
let conn = self.read_conn();
let payload_json: Option<String> = conn
.query_row(
"SELECT value FROM meta WHERE key = 'reembed_state'",
[],
|row| row.get(0),
)
.ok();
let payload: serde_json::Value = serde_json::from_str(&payload_json?).ok()?;
let generation = payload.get("generation")?.as_u64()?;
let phase = ReembedPhase::parse(payload.get("phase")?.as_str()?)?;
let old_name = payload
.get("old_embedder")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let new_name = payload
.get("new_embedder_name")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let old_dim = payload.get("old_dim").and_then(|v| v.as_u64()).unwrap_or(0) as usize;
let started_at_unix = payload
.get("started_at_unix")
.and_then(|v| v.as_f64())
.unwrap_or(0.0);
let started_at =
std::time::UNIX_EPOCH + std::time::Duration::from_secs_f64(started_at_unix.max(0.0));
let write_policy = match payload.get("write_policy").and_then(|v| v.as_str()) {
Some("Pause") => ReembedWritePolicy::Pause,
_ => ReembedWritePolicy::Queue,
};
Some(ReembedStatus {
generation,
phase,
old_embedder: old_name,
old_embedder_digest: String::new(),
new_embedder: new_name,
new_embedder_digest: String::new(),
old_dim,
new_dim: old_dim,
memories_total: 0,
memories_encoded: 0,
queued_writes: 0,
checkpoint_rid: None,
started_at,
last_event_at: started_at,
last_error: None,
write_policy,
})
}
pub(crate) fn write_reembed_event(
&self,
generation: u64,
phase: ReembedPhase,
timestamp_unix_secs: f64,
payload: &serde_json::Value,
) -> Result<()> {
use rusqlite::params;
let payload_str = serde_json::to_string(payload)?;
let conn = self.conn();
conn.execute(
"INSERT INTO reembed_events (generation, phase, timestamp, payload_json) \
VALUES (?1, ?2, ?3, ?4)",
params![
generation as i64,
phase.as_str(),
timestamp_unix_secs,
payload_str
],
)?;
Ok(())
}
pub(crate) fn write_reembed_state_meta(&self, state_json: &serde_json::Value) -> Result<()> {
use rusqlite::params;
let s = serde_json::to_string(state_json)?;
let conn = self.conn();
conn.execute(
"INSERT OR REPLACE INTO meta (key, value) VALUES ('reembed_state', ?1)",
params![s],
)?;
Ok(())
}
pub(crate) fn clear_reembed_state_meta(&self) -> Result<()> {
let conn = self.conn();
conn.execute("DELETE FROM meta WHERE key = 'reembed_state'", [])?;
Ok(())
}
pub(crate) fn update_reembed_state_phase(&self, new_phase: &str) -> Result<()> {
use rusqlite::params;
let conn = self.conn();
let raw: Option<String> = conn
.query_row(
"SELECT value FROM meta WHERE key = 'reembed_state'",
[],
|r| r.get::<_, String>(0),
)
.ok();
let Some(s) = raw else {
return Ok(());
};
let mut state: serde_json::Value =
serde_json::from_str(&s).unwrap_or_else(|_| serde_json::json!({}));
if let Some(map) = state.as_object_mut() {
map.insert(
"phase".to_string(),
serde_json::Value::String(new_phase.to_string()),
);
}
let s_new = serde_json::to_string(&state)?;
conn.execute(
"INSERT OR REPLACE INTO meta (key, value) VALUES ('reembed_state', ?1)",
params![s_new],
)?;
Ok(())
}
}
fn systime_to_unix_secs(t: SystemTime) -> f64 {
t.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn phase_string_round_trip() {
for &phase in &[
ReembedPhase::Probing,
ReembedPhase::Encoding,
ReembedPhase::Rebuilding,
ReembedPhase::Swapping,
ReembedPhase::Verifying,
ReembedPhase::Completed,
ReembedPhase::Aborted,
] {
let s = phase.as_str();
let parsed = ReembedPhase::parse(s).unwrap_or_else(|| {
panic!("ReembedPhase::parse({s:?}) returned None; persisted state would be unparseable on restart")
});
assert_eq!(parsed, phase, "round-trip mismatch for {phase:?}");
}
}
#[test]
fn phase_parse_rejects_unknown() {
assert!(ReembedPhase::parse("Nonsense").is_none());
assert!(ReembedPhase::parse("").is_none());
assert!(ReembedPhase::parse("probing").is_none());
}
#[test]
fn phase_terminal_classification() {
assert!(ReembedPhase::Completed.is_terminal());
assert!(ReembedPhase::Aborted.is_terminal());
assert!(!ReembedPhase::Probing.is_terminal());
assert!(!ReembedPhase::Encoding.is_terminal());
assert!(!ReembedPhase::Rebuilding.is_terminal());
assert!(!ReembedPhase::Swapping.is_terminal());
assert!(!ReembedPhase::Verifying.is_terminal());
}
#[test]
fn write_policy_default_is_queue() {
assert_eq!(ReembedWritePolicy::default(), ReembedWritePolicy::Queue);
}
#[test]
fn options_default_shape_matches_locked_design() {
let opts = ReembedOptions::default();
assert!(opts.namespace.is_none());
assert!(opts.progress_cb.is_none());
assert!(opts.on_phase_complete.is_none());
assert_eq!(opts.batch_size, 256);
assert_eq!(opts.write_policy, ReembedWritePolicy::Queue);
assert!(opts.hnsw_m.is_none());
assert!(opts.hnsw_ef_construction.is_none());
assert!(opts.hnsw_ef_search.is_none());
assert!(opts.resume_from_checkpoint);
assert!(!opts.dry_run);
}
#[test]
fn provenance_known_dim_digest_name_accessors() {
let p = EmbeddingProvenance::Known {
name: Some("potion-base-2M".to_string()),
digest: "sha256:abc123".to_string(),
dim: 64,
};
assert_eq!(p.dim(), 64);
assert_eq!(p.digest(), Some("sha256:abc123"));
assert_eq!(p.name(), Some("potion-base-2M"));
}
#[test]
fn provenance_external_or_unknown_has_dim_no_digest_no_name() {
let p = EmbeddingProvenance::ExternalOrUnknown { dim: 384 };
assert_eq!(p.dim(), 384);
assert_eq!(p.digest(), None);
assert_eq!(p.name(), None);
}
#[test]
fn search_state_initial_is_external_or_unknown_with_no_embedder() {
let vec_index = Arc::new(crate::vector::delta_index::DeltaIndex::new(384));
let s = SearchState::initial(384, 16, 200, 50, vec_index);
assert_eq!(s.dim(), 384);
assert!(matches!(
s.index_embedding,
EmbeddingProvenance::ExternalOrUnknown { dim: 384 }
));
assert!(!s.has_runtime_embedder());
assert!(s.embedder.is_none());
assert!(s.runtime_embedder_name.is_none());
assert!(s.runtime_embedder_digest.is_none());
assert_eq!(s.generation, 0);
assert_eq!(s.covers_through_seq, 0);
assert_eq!(s.hnsw_m, 16);
assert_eq!(s.hnsw_ef_construction, 200);
assert_eq!(s.hnsw_ef_search, 50);
assert_eq!(s.vec_index.len(), 0);
}
struct PhaseTestEmbedder {
pub dim: usize,
pub fp: String,
pub name: String,
}
impl crate::types::Embedder for PhaseTestEmbedder {
fn embed(
&self,
_text: &str,
) -> std::result::Result<Vec<f32>, Box<dyn std::error::Error + Send + Sync>> {
Ok(vec![0.1_f32; self.dim])
}
fn dim(&self) -> usize {
self.dim
}
fn fingerprint(&self) -> Option<String> {
Some(self.fp.clone())
}
fn name(&self) -> Option<String> {
Some(self.name.clone())
}
}
struct PhaseTestEmbedderSentinel {
pub dim: usize,
pub fp: String,
pub name: String,
pub sentinel: f32,
}
impl crate::types::Embedder for PhaseTestEmbedderSentinel {
fn embed(
&self,
_text: &str,
) -> std::result::Result<Vec<f32>, Box<dyn std::error::Error + Send + Sync>> {
let mut v = vec![0.0_f32; self.dim];
if !v.is_empty() {
v[0] = self.sentinel;
}
Ok(v)
}
fn dim(&self) -> usize {
self.dim
}
fn fingerprint(&self) -> Option<String> {
Some(self.fp.clone())
}
fn name(&self) -> Option<String> {
Some(self.name.clone())
}
}
#[test]
fn reembed_same_name_is_no_op() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:same".to_string(),
name: "model-x".to_string(),
}))
.unwrap();
let report = db.reembed("model-x", ReembedOptions::default()).unwrap();
assert_eq!(report.encoded_count, 0, "same-name no-op encodes 0");
assert_eq!(report.skipped_count, 0);
assert_eq!(report.old_embedder, "model-x");
assert_eq!(report.new_embedder, "model-x");
assert!(
db.reembed_status().is_none(),
"no-op must not leave reembed_status"
);
}
#[test]
fn reembed_dry_run_returns_predicted_report_and_clears_state() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:original".to_string(),
name: "original-model".to_string(),
}))
.unwrap();
let synthetic: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> =
std::sync::Arc::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:target".to_string(),
name: "target-model".to_string(),
});
let opts = ReembedOptions {
dry_run: true,
..Default::default()
};
let report = db
.reembed_with_embedder("target-model", Some(synthetic), opts)
.unwrap();
assert_eq!(report.old_embedder, "original-model");
assert_eq!(report.new_embedder, "target-model");
assert!(
db.reembed_status().is_none(),
"dry-run must NOT write meta.reembed_state"
);
let event_count: i64 = db
.conn()
.query_row(
"SELECT COUNT(*) FROM reembed_events WHERE phase = 'Probing'",
[],
|row| row.get(0),
)
.unwrap();
assert!(
event_count >= 1,
"Probing event must be in audit log even for dry-run"
);
}
#[test]
fn reembed_phase_2_full_run_promotes_to_new_generation() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:original".to_string(),
name: "original-model".to_string(),
}))
.unwrap();
for i in 0..3 {
db.record(
&format!("row {i}"),
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec![0.1_f32; 8],
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
}
let synthetic: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> =
std::sync::Arc::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:target".to_string(),
name: "target-model".to_string(),
});
let report = db
.reembed_with_embedder("target-model", Some(synthetic), ReembedOptions::default())
.unwrap();
assert_eq!(report.generation, 1, "new generation = 1 (was 0)");
assert_eq!(report.encoded_count, 3, "all 3 rows re-encoded");
assert_eq!(report.old_embedder, "original-model");
assert_eq!(report.new_embedder, "target-model");
assert_eq!(report.old_dim, 8);
assert_eq!(report.new_dim, 8);
assert!(
report.build_hwm >= 3,
"covers_through_seq covers all writes"
);
let state = db.search_state.load_full();
assert_eq!(state.generation, 1, "in-memory generation advanced");
assert_eq!(state.runtime_embedder_name.as_deref(), Some("target-model"));
assert!(
matches!(
&state.index_embedding,
EmbeddingProvenance::Known { digest, .. } if digest == "sha256:target"
),
"provenance is Known under the new embedder digest"
);
let conn = db.conn();
let active_gen: String = conn
.query_row(
"SELECT value FROM meta WHERE key = 'active_generation'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(active_gen, "1", "meta.active_generation bumped in SQL");
let row_gens: Vec<i64> = conn
.prepare(
"SELECT embedding_generation FROM memories WHERE consolidation_status = 'active'",
)
.unwrap()
.query_map([], |r| r.get(0))
.unwrap()
.collect::<std::result::Result<Vec<_>, _>>()
.unwrap();
assert_eq!(row_gens, vec![1, 1, 1], "every row stamped at new gen");
let staging_count: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memories \
WHERE embedding_new IS NOT NULL OR embedding_new_model IS NOT NULL",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(
staging_count, 0,
"staging columns must be cleared post-swap"
);
drop(conn);
assert!(
db.reembed_status().is_none(),
"meta.reembed_state must be cleared on Completed"
);
let conn = db.conn();
let completed_count: i64 = conn
.query_row(
"SELECT COUNT(*) FROM reembed_events WHERE phase = 'Completed'",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(completed_count, 1, "Completed event must be logged");
assert!(
db.write_router.try_enter_sync_writer().is_some(),
"WriteRouter must be Normal after reembed completion"
);
}
#[test]
fn layer_5_materializer_drains_queued_record_under_new_embedder() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0-name".to_string(),
sentinel: 0.42,
}))
.unwrap();
db.write_router.switch_to_queueing();
let queued_rid = db
.record_text(
"queue-drain-test-text",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({"k": "v"}),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let mem_count: i64 = db
.conn()
.query_row(
"SELECT COUNT(*) FROM memories WHERE rid = ?1",
[&queued_rid],
|r| r.get(0),
)
.unwrap();
assert_eq!(mem_count, 0, "queue path does not write memories yet");
let old_state = db.search_state.load_full();
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1-name".to_string(),
sentinel: 0.99,
});
let new_state = SearchState {
index_embedding: EmbeddingProvenance::Known {
name: Some("E1-name".to_string()),
digest: "sha256:E1".to_string(),
dim: 8,
},
embedder: Some(Arc::clone(&e1)),
runtime_embedder_name: Some("E1-name".to_string()),
runtime_embedder_digest: Some("sha256:E1".to_string()),
generation: 1,
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: Arc::clone(&old_state.vec_index),
};
db.try_publish_search_state(new_state).unwrap();
db.write_router.switch_to_normal();
let n_applied = db.apply_pending_ops_once(100).unwrap();
assert!(
n_applied >= 1,
"Layer 5 materializer must drain the queued record (applied >= 1, got {n_applied})"
);
let conn = db.conn();
let (row_gen, stored_emb_blob): (i64, Vec<u8>) = conn
.query_row(
"SELECT embedding_generation, embedding FROM memories WHERE rid = ?1",
[&queued_rid],
|r| Ok((r.get(0)?, r.get(1)?)),
)
.unwrap();
assert_eq!(row_gen, 1, "Layer 5 stamped row at new generation 1");
let plain = db.decrypt_embedding(&stored_emb_blob).unwrap();
let vec = crate::serde_helpers::deserialize_f32(&plain);
assert!(
(vec[0] - 0.99).abs() < 1e-6,
"row encoded under E1 (sentinel 0.99), got vec[0]={}",
vec[0]
);
let applied: i64 = conn
.query_row(
"SELECT applied FROM oplog WHERE target_rid = ?1 AND op_type = 'record'",
[&queued_rid],
|r| r.get(0),
)
.unwrap();
assert_eq!(applied, 1, "queued oplog row marked applied=1 after drain");
}
#[test]
fn layer_5_defers_drain_while_reembed_in_flight() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0-name".to_string(),
sentinel: 0.42,
}))
.unwrap();
db.write_router.switch_to_queueing();
let queued_rid = db
.record_text(
"deferred-drain-test",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
db.write_reembed_state_meta(&serde_json::json!({
"generation": 1,
"phase": "Encoding",
}))
.unwrap();
let _ = db.apply_pending_ops_once(100).unwrap();
let applied: i64 = db
.conn()
.query_row(
"SELECT applied FROM oplog WHERE target_rid = ?1 AND op_type = 'record'",
[&queued_rid],
|r| r.get(0),
)
.unwrap();
assert_eq!(
applied, 0,
"Layer 5 must defer drain while reembed in flight; got applied={applied}"
);
}
#[test]
fn reembed_phase_2_rejects_cross_dim_with_clear_error() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:active".to_string(),
name: "active".to_string(),
}))
.unwrap();
let cross_dim: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> =
std::sync::Arc::new(PhaseTestEmbedder {
dim: 16, fp: "sha256:cross".to_string(),
name: "cross-dim-target".to_string(),
});
let err = db
.reembed_with_embedder(
"cross-dim-target",
Some(cross_dim),
ReembedOptions::default(),
)
.unwrap_err();
match err {
crate::error::YantrikDbError::Inference(msg) => {
assert!(
msg.contains("Cross-dim reembed is not yet supported"),
"expected cross-dim guard message, got: {msg}"
);
assert!(
msg.contains("dim=8") && msg.contains("dim=16"),
"msg: {msg}"
);
}
other => panic!("expected Inference, got {other:?}"),
}
assert!(db.reembed_status().is_none());
}
#[test]
fn reembed_phase_2_part_a_with_progress_callback_emits_total() {
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:active".to_string(),
name: "active".to_string(),
}))
.unwrap();
for i in 0..3 {
db.record(
&format!("row {i}"),
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&vec![0.1_f32; 8],
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
}
use std::sync::{Arc, Mutex};
let events: Arc<Mutex<Vec<ReembedProgress>>> = Arc::new(Mutex::new(Vec::new()));
let events_clone = Arc::clone(&events);
let progress_cb: Box<dyn Fn(ReembedProgress) + Send + Sync> =
Box::new(move |p| events_clone.lock().unwrap().push(p));
let synthetic: std::sync::Arc<dyn crate::types::Embedder + Send + Sync> =
std::sync::Arc::new(PhaseTestEmbedder {
dim: 8,
fp: "sha256:target".to_string(),
name: "target".to_string(),
});
let opts = ReembedOptions {
progress_cb: Some(progress_cb),
..Default::default()
};
let _ = db.reembed_with_embedder("target", Some(synthetic), opts);
let captured = events.lock().unwrap().clone();
assert!(!captured.is_empty(), "at least one progress event fired");
let first = &captured[0];
assert!(
matches!(first.phase, ReembedPhase::Probing),
"first event must be Probing, got {:?}",
first.phase
);
assert_eq!(
first.total,
Some(3),
"Probing must populate total so CLIs show 0/N immediately (hermes ask 1)"
);
assert_eq!(first.processed, 0, "Probing fires with processed=0");
let encoding_events: Vec<&ReembedProgress> = captured
.iter()
.filter(|e| matches!(e.phase, ReembedPhase::Encoding))
.collect();
for e in &encoding_events {
assert_eq!(e.total, Some(3));
assert!(e.processed <= 3);
}
if let Some(last) = encoding_events.last() {
assert_eq!(last.processed, 3, "final Encoding event reflects all rows");
}
}
#[test]
fn search_state_dim_derives_from_provenance_not_a_separate_field() {
let vec_index_known = Arc::new(crate::vector::delta_index::DeltaIndex::new(768));
let s_known = SearchState {
index_embedding: EmbeddingProvenance::Known {
name: None,
digest: "x".into(),
dim: 768,
},
embedder: None,
runtime_embedder_name: None,
runtime_embedder_digest: None,
generation: 5,
covers_through_seq: 1000,
hnsw_m: 16,
hnsw_ef_construction: 200,
hnsw_ef_search: 50,
vec_index: vec_index_known,
};
assert_eq!(s_known.dim(), 768);
let vec_index_unknown = Arc::new(crate::vector::delta_index::DeltaIndex::new(128));
let s_unknown = SearchState {
index_embedding: EmbeddingProvenance::ExternalOrUnknown { dim: 128 },
embedder: None,
runtime_embedder_name: None,
runtime_embedder_digest: None,
generation: 0,
covers_through_seq: 0,
hnsw_m: 16,
hnsw_ef_construction: 200,
hnsw_ef_search: 50,
vec_index: vec_index_unknown,
};
assert_eq!(s_unknown.dim(), 128);
}
#[test]
fn reembed_post_swap_recall_returns_rows_under_new_embedder() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0-name".to_string(),
sentinel: 0.42,
}))
.unwrap();
let rids: Vec<String> = (0..5)
.map(|i| {
db.record_text(
&format!("layer-8 row {i}"),
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({"i": i}),
"default",
0.8,
"general",
"user",
None,
)
.unwrap()
})
.collect();
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1-name".to_string(),
sentinel: 0.99,
});
let report = db
.reembed_with_embedder("E1-name", Some(e1), ReembedOptions::default())
.unwrap();
assert_eq!(report.generation, 1);
assert_eq!(report.encoded_count, 5);
let query = vec![0.99_f32, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let results = db
.recall(
&query, 5, None, None, false, true, None, true, None, None, None, None, None, false,
)
.unwrap();
assert_eq!(
results.len(),
5,
"recall must return all 5 planted rows after reembed; got {}",
results.len()
);
for r in &results {
assert!(
rids.contains(&r.rid),
"recall returned unexpected rid {:?}",
r.rid
);
}
}
#[test]
fn reembed_on_empty_engine_returns_gracefully() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0".to_string(),
sentinel: 0.42,
}))
.unwrap();
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1".to_string(),
sentinel: 0.99,
});
let report = db
.reembed_with_embedder("E1", Some(e1), ReembedOptions::default())
.unwrap();
assert_eq!(report.encoded_count, 0);
assert_eq!(report.generation, 1);
let conn = db.conn();
for phase in [
"Probing",
"Encoding",
"Rebuilding",
"Swapping",
"Verifying",
"Completed",
] {
let count: i64 = conn
.query_row(
"SELECT COUNT(*) FROM reembed_events WHERE phase = ?1 AND generation = 1",
params![phase],
|r| r.get(0),
)
.unwrap();
assert!(
count >= 1,
"expected {phase} event for empty-engine reembed (got {count})"
);
}
}
#[test]
fn reembed_sequential_runs_advance_generation_monotonically() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0".to_string(),
sentinel: 0.10,
}))
.unwrap();
let rid = db
.record_text(
"seq",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1".to_string(),
sentinel: 0.50,
});
let r1 = db
.reembed_with_embedder("E1", Some(e1), ReembedOptions::default())
.unwrap();
assert_eq!(r1.generation, 1);
let e2: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E2".to_string(),
name: "E2".to_string(),
sentinel: 0.90,
});
let r2 = db
.reembed_with_embedder("E2", Some(e2), ReembedOptions::default())
.unwrap();
assert_eq!(r2.generation, 2);
assert_eq!(
r2.encoded_count, 1,
"second reembed re-encodes the planted row"
);
let row_gen: i64 = db
.conn()
.query_row(
"SELECT embedding_generation FROM memories WHERE rid = ?1",
[&rid],
|r| r.get(0),
)
.unwrap();
assert_eq!(row_gen, 2);
assert_eq!(db.search_state.load().generation, 2);
}
#[test]
fn reembed_audit_event_sequence_is_complete_and_ordered() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0".to_string(),
sentinel: 0.10,
}))
.unwrap();
let _ = db.record_text(
"audit",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
"default",
0.8,
"general",
"user",
None,
);
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1".to_string(),
sentinel: 0.50,
});
let _ = db
.reembed_with_embedder("E1", Some(e1), ReembedOptions::default())
.unwrap();
let phases: Vec<String> = {
let conn = db.conn();
let mut stmt = conn
.prepare("SELECT phase FROM reembed_events WHERE generation = 1 ORDER BY timestamp")
.unwrap();
stmt.query_map([], |r| r.get::<_, String>(0))
.unwrap()
.collect::<std::result::Result<Vec<_>, _>>()
.unwrap()
};
assert_eq!(phases.first().map(String::as_str), Some("Probing"));
for required in [
"Encoding",
"Rebuilding",
"Swapping",
"Verifying",
"Completed",
] {
assert!(
phases.iter().any(|p| p == required),
"audit log missing {required} event; phases recorded: {phases:?}"
);
}
let last_meaningful = phases.last().map(String::as_str);
assert_eq!(
last_meaningful,
Some("Completed"),
"Completed must be the terminal event for the reembed run; got {last_meaningful:?}"
);
}
#[test]
fn reembed_re_encodes_rows_from_multiple_namespaces() {
use std::sync::Arc;
let mut db = crate::YantrikDB::new(":memory:", 8).unwrap();
db.set_embedder(Box::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E0".to_string(),
name: "E0".to_string(),
sentinel: 0.10,
}))
.unwrap();
for ns in ["alpha", "beta", "gamma"] {
for i in 0..2 {
db.record_text(
&format!("{ns} row {i}"),
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
ns,
0.8,
"general",
"user",
None,
)
.unwrap();
}
}
let e1: Arc<dyn crate::types::Embedder + Send + Sync> =
Arc::new(PhaseTestEmbedderSentinel {
dim: 8,
fp: "sha256:E1".to_string(),
name: "E1".to_string(),
sentinel: 0.90,
});
let report = db
.reembed_with_embedder("E1", Some(e1), ReembedOptions::default())
.unwrap();
assert_eq!(
report.encoded_count, 6,
"6 rows across 3 namespaces must all be re-encoded"
);
let conn = db.conn();
for ns in ["alpha", "beta", "gamma"] {
let count: i64 = conn
.query_row(
"SELECT COUNT(*) FROM memories WHERE namespace = ?1 \
AND embedding_generation = 1 AND consolidation_status = 'active'",
params![ns],
|r| r.get(0),
)
.unwrap();
assert_eq!(
count, 2,
"namespace {ns} must have 2 rows at gen 1, got {count}"
);
}
}
}