use std::sync::Arc;
use fathomdb_embedder::EmbedderEvent;
#[cfg(feature = "operator")]
use fathomdb_embedder::MeanRecomputeTrigger;
use fathomdb_embedder_api::{Embedder, EmbedderError, EmbedderIdentity, Vector};
use fathomdb_engine::{EmbedderChoice, Engine, PreparedWrite, MEAN_VEC_PIN_THRESHOLD};
use rusqlite::Connection;
use tempfile::TempDir;
const DIM: u32 = 384;
const BGE_NAME: &str = "fathomdb-bge-small-en-v1.5";
const BGE_REV: &str = "5c38ec7c405ec4b44b94cc5a9bb96e735b38267a";
#[derive(Clone, Debug)]
struct SimulatedBgeEmbedder {
identity: EmbedderIdentity,
}
impl Default for SimulatedBgeEmbedder {
fn default() -> Self {
Self { identity: EmbedderIdentity::new(BGE_NAME, BGE_REV, DIM) }
}
}
impl Embedder for SimulatedBgeEmbedder {
fn identity(&self) -> EmbedderIdentity {
self.identity.clone()
}
fn embed(&self, input: &str) -> Result<Vector, EmbedderError> {
Ok(topic_vector(input))
}
}
fn hash64(input: &str) -> u64 {
let mut seed: u64 = 0xcbf29ce484222325;
for b in input.bytes() {
seed ^= u64::from(b);
seed = seed.wrapping_mul(0x100000001b3);
}
seed
}
fn topic_vector(input: &str) -> Vector {
let (topic, rest) = match input.split_once(':') {
Some(("A", rest)) => (Some(false), rest),
Some(("B", rest)) => (Some(true), rest),
_ => (None, input),
};
let seed = hash64(rest);
let mut v = vec![0.0f32; DIM as usize];
for (i, slot) in v.iter_mut().enumerate() {
let mixed = seed.wrapping_add(i as u64).wrapping_mul(2654435761);
*slot = (((mixed >> 8) as u32 as f32) / (u32::MAX as f32) - 0.5) * 0.7;
}
if let Some(is_b) = topic {
let half = (DIM / 2) as usize;
for (i, slot) in v.iter_mut().enumerate() {
let in_b_block = i >= half;
if in_b_block == is_b {
*slot += 0.4;
} else {
*slot -= 0.4;
}
}
}
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-9);
for slot in &mut v {
*slot /= norm;
}
v
}
fn fixture_path(name: &str) -> (TempDir, std::path::PathBuf) {
let dir = TempDir::new().unwrap();
let path = dir.path().join(format!("{name}.sqlite"));
(dir, path)
}
fn open_caller(
path: &std::path::Path,
embedder: Arc<dyn Embedder>,
) -> fathomdb_engine::OpenedEngine {
Engine::open_with_choice(path, EmbedderChoice::Caller(embedder)).expect("open")
}
fn write_docs<F: Fn(usize) -> String>(engine: &Engine, count: usize, batch: usize, body: F) {
let mut written = 0usize;
while written < count {
let take = batch.min(count - written);
let nodes: Vec<PreparedWrite> = (0..take)
.map(|i| PreparedWrite::Node {
kind: "doc".to_string(),
body: body(written + i),
source_id: fathomdb_engine::SourceId::new("test:fixture").expect("test source id"),
logical_id: None,
state: fathomdb_engine::InitialState::Active,
reason: None,
valid_from: None,
valid_until: None,
})
.collect();
engine.write(&nodes).expect("production write");
written += take;
engine.drain(60_000).expect("drain");
}
}
fn read_mean_vec(path: &std::path::Path) -> Option<Vec<u8>> {
let conn = Connection::open(path).expect("reopen");
conn.query_row(
"SELECT mean_vec FROM _fathomdb_embedder_profiles WHERE profile = 'default'",
[],
|row| row.get::<_, Option<Vec<u8>>>(0),
)
.expect("mean_vec query")
}
#[cfg(feature = "operator")]
fn decode_f32(blob: &[u8]) -> Vec<f32> {
blob.chunks_exact(4).map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]])).collect()
}
#[cfg(feature = "operator")]
fn subtract(v: &[f32], mean: &[f32]) -> Vec<f32> {
v.iter().zip(mean).map(|(a, b)| a - b).collect()
}
#[cfg(feature = "operator")]
fn quantize_binary(conn: &Connection, vec: &[f32]) -> Vec<u8> {
let json = serde_json::to_string(vec).expect("json");
conn.query_row("SELECT vec_quantize_binary(vec_f32(?1))", [json], |r| r.get::<_, Vec<u8>>(0))
.expect("vec_quantize_binary")
}
#[cfg(feature = "operator")]
fn closed_form_mean(conn: &Connection) -> Vec<f32> {
let mut stmt = conn.prepare("SELECT embedding FROM vector_default ORDER BY rowid").unwrap();
let rows: Vec<Vec<u8>> =
stmt.query_map([], |r| r.get::<_, Vec<u8>>(0)).unwrap().filter_map(Result::ok).collect();
let mut sum = vec![0.0f64; DIM as usize];
for blob in &rows {
for (s, x) in sum.iter_mut().zip(decode_f32(blob)) {
*s += f64::from(x);
}
}
let n = rows.len().max(1) as f64;
sum.iter().map(|s| (s / n) as f32).collect()
}
#[cfg(feature = "operator")]
fn cosine(a: &[f32], b: &[f32]) -> f32 {
let dot: f64 = a.iter().zip(b).map(|(x, y)| f64::from(*x) * f64::from(*y)).sum();
let na: f64 = a.iter().map(|x| f64::from(*x) * f64::from(*x)).sum::<f64>().sqrt();
let nb: f64 = b.iter().map(|x| f64::from(*x) * f64::from(*x)).sum::<f64>().sqrt();
(dot / (na * nb).max(1e-12)) as f32
}
#[test]
fn topic_pivot_does_not_auto_recompute_mid_ingest() {
let (_dir, path) = fixture_path("pr2bc_no_auto_drift");
let opened = open_caller(&path, Arc::new(SimulatedBgeEmbedder::default()));
let engine = opened.engine;
engine.configure_vector_kind_for_test("doc").expect("vector kind");
write_docs(&engine, MEAN_VEC_PIN_THRESHOLD as usize, 64, |i| format!("A:{i}"));
let _ = engine.drain_embedder_events();
let mean_after_pin = read_mean_vec(&path).expect("mean pinned after topic-A");
write_docs(&engine, 600, 64, |i| format!("B:{i}"));
let events = engine.drain_embedder_events().expect("drain events");
let any_recomputed =
events.iter().any(|e| matches!(e, EmbedderEvent::MeanVecRecomputed { .. }));
assert!(
!any_recomputed,
"no MeanVecRecomputed may be emitted during ingest (auto detector carved out); events={events:?}"
);
let mean_after_flood = read_mean_vec(&path).expect("mean still pinned after topic-B flood");
assert_eq!(
mean_after_pin, mean_after_flood,
"pinned mean must be unchanged after the topic-B flood (no auto-recompute)"
);
engine.close().expect("close");
}
#[cfg(feature = "operator")]
#[test]
fn manual_recompute_matches_closed_form_and_requantizes_all() {
let (_dir, path) = fixture_path("pr2b_mechanical");
let opened = open_caller(&path, Arc::new(SimulatedBgeEmbedder::default()));
let engine = opened.engine;
engine.configure_vector_kind_for_test("doc").expect("vector kind");
write_docs(&engine, MEAN_VEC_PIN_THRESHOLD as usize, 64, |i| format!("A:{i}"));
write_docs(&engine, 64, 64, |i| format!("B:{i}"));
let _ = engine.drain_embedder_events();
let report = engine.recompute_mean().expect("manual recompute");
assert!(report.mean_was_pinned, "recompute must observe the prior pin");
assert_eq!(report.dim, DIM);
let total = engine.vector_row_count_for_test().expect("count");
assert_eq!(
report.doc_count_requantized, total,
"recompute must re-quantize every row, got {} of {total}",
report.doc_count_requantized
);
engine.close().expect("close");
let conn = Connection::open(&path).expect("reopen");
let pinned = decode_f32(&read_mean_vec(&path).expect("mean pinned"));
let want = closed_form_mean(&conn);
let cos = cosine(&pinned, &want);
assert!(cos > 0.9999, "pinned mean must match closed-form full-corpus mean, cos={cos}");
let mut stmt =
conn.prepare("SELECT rowid, embedding, embedding_bin FROM vector_default").unwrap();
let rows: Vec<(i64, Vec<u8>, Vec<u8>)> = stmt
.query_map([], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)))
.unwrap()
.filter_map(Result::ok)
.collect();
assert!(rows.len() as u64 == total);
for (rid, emb, bin) in &rows {
let want_bits = quantize_binary(&conn, &subtract(&decode_f32(emb), &pinned));
assert_eq!(bin, &want_bits, "row {rid} sign-bits inconsistent with re-pinned mean");
}
}
#[cfg(feature = "operator")]
#[test]
fn recompute_fault_rolls_back_fully() {
let (_dir, path) = fixture_path("pr2b_atomicity");
let opened = open_caller(&path, Arc::new(SimulatedBgeEmbedder::default()));
let engine = opened.engine;
engine.configure_vector_kind_for_test("doc").expect("vector kind");
write_docs(&engine, MEAN_VEC_PIN_THRESHOLD as usize, 64, |i| format!("A:{i}"));
write_docs(&engine, 64, 64, |i| format!("B:{i}"));
let _ = engine.drain_embedder_events();
let mean_before = read_mean_vec(&path).expect("mean pinned before");
engine.force_next_recompute_failure_for_test();
let err = engine.recompute_mean();
assert!(err.is_err(), "injected fault must surface as an error, got {err:?}");
engine.close().expect("close");
let mean_after = read_mean_vec(&path).expect("mean still pinned after rollback");
assert_eq!(mean_before, mean_after, "mean_vec must be unchanged after a rolled-back recompute");
let conn = Connection::open(&path).expect("reopen");
let pinned = decode_f32(&mean_after);
let mut stmt =
conn.prepare("SELECT rowid, embedding, embedding_bin FROM vector_default").unwrap();
let rows: Vec<(i64, Vec<u8>, Vec<u8>)> = stmt
.query_map([], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)))
.unwrap()
.filter_map(Result::ok)
.collect();
for (rid, emb, bin) in &rows {
let want_bits = quantize_binary(&conn, &subtract(&decode_f32(emb), &pinned));
assert_eq!(bin, &want_bits, "row {rid} must remain centered under the ORIGINAL mean");
}
}
#[cfg(feature = "operator")]
#[test]
fn recompute_event_fields_and_post_commit_publish() {
let (_dir, path) = fixture_path("pr2b_events");
let opened = open_caller(&path, Arc::new(SimulatedBgeEmbedder::default()));
let engine = opened.engine;
engine.configure_vector_kind_for_test("doc").expect("vector kind");
write_docs(&engine, MEAN_VEC_PIN_THRESHOLD as usize, 64, |i| format!("A:{i}"));
write_docs(&engine, 64, 64, |i| format!("B:{i}"));
let _ = engine.drain_embedder_events();
let report = engine.recompute_mean().expect("manual recompute");
let events = engine.drain_embedder_events().expect("drain");
let recomputed: Vec<&EmbedderEvent> =
events.iter().filter(|e| matches!(e, EmbedderEvent::MeanVecRecomputed { .. })).collect();
assert_eq!(recomputed.len(), 1, "exactly one MeanVecRecomputed, got {events:?}");
match recomputed[0] {
EmbedderEvent::MeanVecRecomputed { dim, doc_count, trigger } => {
assert_eq!(*dim, DIM);
assert_eq!(*doc_count, report.doc_count_requantized);
assert_eq!(*trigger, MeanRecomputeTrigger::Manual);
}
other => panic!("expected MeanVecRecomputed, got {other:?}"),
}
assert!(engine.drain_embedder_events().expect("drain2").is_empty());
engine.close().expect("close");
}
#[cfg(feature = "operator")]
#[derive(Clone, Debug)]
struct NonMcEmbedder {
identity: EmbedderIdentity,
}
#[cfg(feature = "operator")]
impl Default for NonMcEmbedder {
fn default() -> Self {
Self { identity: EmbedderIdentity::new("fathomdb-noop", "0.6.0-scaffold", DIM) }
}
}
#[cfg(feature = "operator")]
impl Embedder for NonMcEmbedder {
fn identity(&self) -> EmbedderIdentity {
self.identity.clone()
}
fn embed(&self, input: &str) -> Result<Vector, EmbedderError> {
Ok(topic_vector(input))
}
}
#[cfg(feature = "operator")]
#[test]
fn recompute_rejects_non_mc_identity() {
let (_dir, path) = fixture_path("pr2b_noop");
let opened = open_caller(&path, Arc::new(NonMcEmbedder::default()));
let engine = opened.engine;
let err = engine.recompute_mean();
assert!(err.is_err(), "non-MC identity must not be recomputable, got {err:?}");
}