use std::sync::Arc;
use fathomdb_embedder::EmbedderEvent;
use fathomdb_embedder_api::{Embedder, EmbedderError, EmbedderIdentity, Vector};
use fathomdb_engine::{EmbedderChoice, Engine, MEAN_VEC_PIN_THRESHOLD};
use rusqlite::Connection;
use tempfile::TempDir;
#[allow(unused_macros)]
macro_rules! skip_if_no_network {
() => {
if std::env::var("FATHOMDB_SKIP_NETWORK_TESTS").is_ok() {
eprintln!("[skip] FATHOMDB_SKIP_NETWORK_TESTS set; skipping test");
return;
}
};
}
fn fixture_path(name: &str) -> (TempDir, std::path::PathBuf) {
let dir = TempDir::new().unwrap();
let path = dir.path().join(format!("{name}.sqlite"));
(dir, path)
}
#[test]
fn default_embedder_constants_flipped_to_bge_small() {
let (_dir, path) = fixture_path("eu5b_constants");
let opened =
Engine::open(&path).expect("open should succeed even pre-fetch (no Default choice)");
assert_eq!(opened.report.default_embedder.name, "fathomdb-bge-small-en-v1.5");
assert_eq!(opened.report.default_embedder.revision, "5c38ec7c405ec4b44b94cc5a9bb96e735b38267a");
assert_eq!(opened.report.default_embedder.dimension, 384);
}
#[cfg(feature = "default-embedder")]
#[test]
fn embedder_choice_default_succeeds_with_bge_identity() {
skip_if_no_network!();
let (_dir, path) = fixture_path("eu5b_default_succeeds");
let opened = Engine::open_with_choice(&path, EmbedderChoice::Default)
.expect("EmbedderChoice::Default must succeed after EU-5b");
assert_eq!(opened.report.default_embedder.name, "fathomdb-bge-small-en-v1.5");
}
#[cfg(feature = "default-embedder")]
#[test]
fn open_report_embedder_mean_centering_required_true_for_default() {
skip_if_no_network!();
let (_dir, path) = fixture_path("eu5b_mc_required");
let opened = Engine::open_with_choice(&path, EmbedderChoice::Default)
.expect("EmbedderChoice::Default must succeed after EU-5b");
assert!(
opened.report.embedder_mean_centering_required,
"Default identity must report MC-required after the EU-5b flip"
);
}
#[cfg(feature = "default-embedder")]
#[test]
fn open_report_embedder_events_populated_for_default() {
skip_if_no_network!();
let (_dir, path) = fixture_path("eu5b_events");
let opened = Engine::open_with_choice(&path, EmbedderChoice::Default)
.expect("EmbedderChoice::Default must succeed");
assert!(
!opened.report.embedder_events.is_empty(),
"Default open must surface loader events (downloads or cache hits)"
);
}
#[cfg(feature = "default-embedder")]
#[test]
fn open_report_embedder_download_ms_some_on_cold_open() {
skip_if_no_network!();
let (_dir, path) = fixture_path("eu5b_download_ms");
let opened = Engine::open_with_choice(&path, EmbedderChoice::Default)
.expect("EmbedderChoice::Default must succeed");
if let Some(ms) = opened.report.embedder_download_ms {
assert!(ms > 0, "download_ms must be > 0 on cold open");
}
}
#[test]
fn default_embedder_not_wired_error_variant_removed() {
fn _exhaustive_witness(err: EmbedderError) -> &'static str {
match err {
EmbedderError::Failed { .. } => "failed",
EmbedderError::Timeout => "timeout",
}
}
let _ = _exhaustive_witness;
}
#[derive(Clone, Debug)]
struct SimulatedBgeEmbedder {
identity: EmbedderIdentity,
}
impl Default for SimulatedBgeEmbedder {
fn default() -> Self {
Self {
identity: EmbedderIdentity::new(
"fathomdb-bge-small-en-v1.5",
"5c38ec7c405ec4b44b94cc5a9bb96e735b38267a",
384,
),
}
}
}
impl Embedder for SimulatedBgeEmbedder {
fn identity(&self) -> EmbedderIdentity {
self.identity.clone()
}
fn embed(&self, input: &str) -> Result<Vector, EmbedderError> {
let mut v = vec![0.0f32; 384];
let mut seed: u64 = 0;
for b in input.bytes() {
seed = seed.wrapping_mul(131).wrapping_add(u64::from(b));
}
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;
}
Ok(v)
}
}
fn open_with_simulated_bge(
name: &str,
) -> (TempDir, std::path::PathBuf, fathomdb_engine::OpenedEngine) {
let (dir, path) = fixture_path(name);
let opened = Engine::open_with_choice(
&path,
EmbedderChoice::Caller(Arc::new(SimulatedBgeEmbedder::default())),
)
.expect("open with simulated bge");
(dir, path, opened)
}
fn write_docs(engine: &Engine, count: u64) {
engine.configure_vector_kind_for_test("doc").expect("configure vector kind");
for i in 0..count {
let txt = format!("doc-{i}");
engine
.write_vector_for_test("doc", &txt)
.unwrap_or_else(|err| panic!("write_vector_for_test failed at i={i}: {err:?}"));
}
}
#[test]
fn mean_accumulator_pins_at_threshold_with_real_embedder() {
let (_dir, path, opened) = open_with_simulated_bge("eu5b_pin_threshold");
write_docs(&opened.engine, MEAN_VEC_PIN_THRESHOLD);
drop(opened);
let connection = Connection::open(&path).expect("reopen");
let mean_blob: Option<Vec<u8>> = connection
.query_row(
"SELECT mean_vec FROM _fathomdb_embedder_profiles WHERE profile = 'default'",
[],
|row| row.get::<_, Option<Vec<u8>>>(0),
)
.expect("mean_vec query");
let mean_blob = mean_blob.expect("mean_vec must be NOT NULL after threshold crossing");
assert_eq!(mean_blob.len(), 384 * 4, "mean_vec byte length must be 4 * dim");
}
#[test]
fn mean_accumulator_emits_mean_vec_pinned_event() {
let (_dir, _path, opened) = open_with_simulated_bge("eu5b_pin_event");
write_docs(&opened.engine, MEAN_VEC_PIN_THRESHOLD);
let events =
opened.engine.drain_mean_centering_events_for_test().expect("drain mean-centering events");
let mut saw_pin = false;
for ev in &events {
if let EmbedderEvent::MeanVecPinned { dim, doc_count } = ev.clone() {
assert_eq!(dim, 384u32);
assert!(doc_count >= 1u64, "doc_count must be >= 1 at pin commit");
saw_pin = true;
}
}
assert!(saw_pin, "MeanVecPinned event must be emitted at pin commit");
}
#[test]
fn requantize_pass_updates_prepin_rows() {
let (_dir, path, opened) = open_with_simulated_bge("eu5b_requantize");
write_docs(&opened.engine, MEAN_VEC_PIN_THRESHOLD);
drop(opened);
let connection = Connection::open(&path).expect("reopen");
let any_row: i64 = connection
.query_row("SELECT rowid FROM vector_default ORDER BY rowid LIMIT 1", [], |row| row.get(0))
.expect("at least one row exists");
let embedding: Vec<u8> = connection
.query_row("SELECT embedding FROM vector_default WHERE rowid = ?1", [any_row], |row| {
row.get::<_, Vec<u8>>(0)
})
.expect("read embedding f32");
let embedding_bin: Vec<u8> = connection
.query_row("SELECT embedding_bin FROM vector_default WHERE rowid = ?1", [any_row], |row| {
row.get::<_, Vec<u8>>(0)
})
.expect("read embedding_bin");
let mut uncentered_bits = Vec::with_capacity(embedding.len() / 4 / 8 + 1);
let mut acc: u8 = 0;
let mut count = 0;
for chunk in embedding.chunks_exact(4) {
let f = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
acc = (acc << 1) | (if f >= 0.0 { 1 } else { 0 });
count += 1;
if count == 8 {
uncentered_bits.push(acc);
acc = 0;
count = 0;
}
}
if count > 0 {
uncentered_bits.push(acc << (8 - count));
}
assert_eq!(
uncentered_bits.len(),
embedding_bin.len(),
"bit-packed length mismatch: uncentered={} stored={}",
uncentered_bits.len(),
embedding_bin.len()
);
assert!(
uncentered_bits != embedding_bin,
"stored sign-bits must differ from un-centered sign-bits after re-quantize"
);
}