use std::env;
use std::sync::Mutex;
use std::time::Instant;
static PG_LOCK: Mutex<()> = Mutex::new(());
fn pg_lock() -> std::sync::MutexGuard<'static, ()> {
PG_LOCK.lock().unwrap_or_else(|e| e.into_inner())
}
const VEC_DIM: usize = 384;
const TOL: f64 = 1e-5;
fn pg_url() -> String {
env::var("LEANKG_PG_URL")
.unwrap_or_else(|_| "postgresql://postgres:postgres@localhost:5433/leankg".to_string())
}
fn pgvector(v: &[f32]) -> String {
format!(
"[{}]",
v.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",")
)
}
fn random_unit_vector(seed: u64, dim: usize) -> Vec<f32> {
let mut state = seed;
let mut next = || {
state ^= state << 7;
state ^= state >> 9;
state.wrapping_mul(0x9E37_79B9_7F4A_7C15)
};
let mut v: Vec<f32> = Vec::with_capacity(dim);
for _ in 0..dim {
let u1 = ((next() >> 11) as f64 / (1u64 << 53) as f64).max(f64::EPSILON);
let u2 = (next() >> 11) as f64 / (1u64 << 53) as f64;
let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
v.push(z as f32);
}
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
for x in &mut v {
*x /= norm;
}
v
}
fn cosine_dist(a: &[f32], b: &[f32]) -> f64 {
1.0 - a
.iter()
.zip(b)
.map(|(x, y)| (*x as f64) * (*y as f64))
.sum::<f64>()
}
fn brute_force_topk(names: &[String], vecs: &[Vec<f32>], q: &[f32], k: usize) -> Vec<String> {
let mut scored: Vec<(&str, f64)> = names
.iter()
.zip(vecs)
.map(|(n, v)| (n.as_str(), cosine_dist(v, q)))
.collect();
scored.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap().then_with(|| a.0.cmp(b.0)));
scored
.into_iter()
.take(k)
.map(|(n, _)| n.to_string())
.collect()
}
fn load_vectors(client: &mut postgres::Client, names: &[String], vecs: &[Vec<f32>]) {
client
.batch_execute("DROP TABLE IF EXISTS embedding_vectors")
.unwrap();
client
.batch_execute(
"CREATE TABLE embedding_vectors (qualified_name TEXT PRIMARY KEY, vec vector(384))",
)
.unwrap();
client
.batch_execute(
"CREATE INDEX embedding_vectors_vec_hnsw_idx \
ON embedding_vectors USING hnsw (vec vector_cosine_ops) \
WITH (m = 16, ef_construction = 200)",
)
.unwrap();
let mut tx = client.transaction().unwrap();
{
let stmt = tx
.prepare(
"INSERT INTO embedding_vectors (qualified_name, vec) VALUES ($1, $2::text::vector)",
)
.unwrap();
for (n, v) in names.iter().zip(vecs) {
tx.execute(&stmt, &[&n.as_str(), &pgvector(v)]).unwrap();
}
}
tx.commit().unwrap();
}
fn pg_hnsw_topk(
client: &mut postgres::Client,
q: &[f32],
k: usize,
ef: usize,
) -> Vec<(String, f64)> {
let mut tx = client.transaction().unwrap();
let set_sql = format!("SET LOCAL hnsw.ef_search = {ef}");
tx.batch_execute(&set_sql).unwrap();
let rows = tx
.query(
"SELECT vec <-> $1::text::vector AS dist, qualified_name
FROM embedding_vectors
ORDER BY vec <-> $1::text::vector
LIMIT $2::int8",
&[&pgvector(q), &(k as i64)],
)
.unwrap();
let out: Vec<(String, f64)> = rows
.iter()
.map(|r| (r.get::<_, String>(1), r.get::<_, f64>(0)))
.collect();
tx.commit().unwrap();
out
}
#[test]
fn test_cosine_distance_identity() {
let v = random_unit_vector(0xC0FFEE, VEC_DIM);
assert!(cosine_dist(&v, &v) < TOL);
}
#[test]
fn test_pgvector_roundtrip() {
let v: Vec<f32> = (0..VEC_DIM).map(|i| (i as f32) * 0.001).collect();
assert_eq!(
pgvector(&v),
format!(
"[{}]",
v.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",")
)
);
}
#[test]
#[ignore = "requires the leankg-pg-phase0 container (localhost:5433)"]
fn phase4_vector_round_trip_top_k_matches_brute_force() {
let _guard = pg_lock();
let mut client = postgres::Client::connect(&pg_url(), postgres::NoTls).unwrap();
let n = 50usize;
let k = 5usize;
let names: Vec<String> = (0..n).map(|i| format!("v{i:04}")).collect();
let vecs: Vec<Vec<f32>> = (0..n)
.map(|i| random_unit_vector(0xA1B2_0000 + i as u64, VEC_DIM))
.collect();
let q = random_unit_vector(0xDEAD_BEEF, VEC_DIM);
load_vectors(&mut client, &names, &vecs);
let brute = brute_force_topk(&names, &vecs, &q, k);
let hnsw = pg_hnsw_topk(&mut client, &q, k, 100);
let hnsw_names: Vec<String> = hnsw.iter().map(|(n, _)| n.clone()).collect();
assert_eq!(
hnsw_names, brute,
"pgvector top-k differs from brute force on 50 vectors"
);
}
#[test]
#[ignore = "requires the leankg-pg-phase0 container (localhost:5433)"]
fn phase4_set_local_hnsw_ef_search_takes_effect() {
let _guard = pg_lock();
let mut client = postgres::Client::connect(&pg_url(), postgres::NoTls).unwrap();
let names: Vec<String> = (0..32).map(|i| format!("ef{i:03}")).collect();
let vecs: Vec<Vec<f32>> = (0..32)
.map(|i| random_unit_vector(i as u64, VEC_DIM))
.collect();
load_vectors(&mut client, &names, &vecs);
let q = random_unit_vector(0xFEED_FACE, VEC_DIM);
let small = pg_hnsw_topk(&mut client, &q, 5, 10);
let big = pg_hnsw_topk(&mut client, &q, 5, 200);
assert_eq!(small.len(), 5, "small-ef query must return k rows");
assert_eq!(big.len(), 5, "big-ef query must return k rows");
let post = pg_hnsw_topk(&mut client, &q, 5, 50);
assert_eq!(post.len(), 5, "post-commit GUC must not poison next query");
}
#[test]
#[ignore = "requires the leankg-pg-phase0 container (localhost:5433)"]
fn phase4_batched_upsert_throughput_smoke() {
let _guard = pg_lock();
let mut client = postgres::Client::connect(&pg_url(), postgres::NoTls).unwrap();
client
.batch_execute("DROP TABLE IF EXISTS embedding_vectors")
.unwrap();
client
.batch_execute(
"CREATE TABLE embedding_vectors (qualified_name TEXT PRIMARY KEY, vec vector(384))",
)
.unwrap();
let n = 1_000usize;
let names: Vec<String> = (0..n).map(|i| format!("up{i:05}")).collect();
let vecs: Vec<Vec<f32>> = (0..n)
.map(|i| random_unit_vector(i as u64, VEC_DIM))
.collect();
let t = Instant::now();
let mut tx = client.transaction().unwrap();
{
let stmt = tx
.prepare(
"INSERT INTO embedding_vectors (qualified_name, vec) \
VALUES ($1, $2::text::vector) \
ON CONFLICT (qualified_name) DO UPDATE SET vec = EXCLUDED.vec",
)
.unwrap();
for (name, v) in names.iter().zip(&vecs) {
tx.execute(&stmt, &[&name.as_str(), &pgvector(v)]).unwrap();
}
}
tx.commit().unwrap();
let elapsed_ms = t.elapsed().as_millis();
let v_per_s = if elapsed_ms > 0 {
(n as f64) / (elapsed_ms as f64 / 1000.0)
} else {
f64::INFINITY
};
println!(
"[phase4] batched upsert (1000 vectors, single tx): {elapsed_ms} ms -> {v_per_s:.0} v/s"
);
assert_eq!(
client
.query_one("SELECT count(*) FROM embedding_vectors", &[])
.unwrap()
.get::<_, i64>(0),
n as i64,
"all rows must land"
);
assert!(
elapsed_ms < 30_000,
"1000-vector single-tx upsert exceeded 30s: {elapsed_ms}ms"
);
}
#[test]
#[ignore = "requires the leankg-pg-phase0 container (localhost:5433)"]
fn phase4_has_any_proxies_correctly() {
let _guard = pg_lock();
let mut client = postgres::Client::connect(&pg_url(), postgres::NoTls).unwrap();
client
.batch_execute("DROP TABLE IF EXISTS embedding_vectors")
.unwrap();
client
.batch_execute(
"CREATE TABLE embedding_vectors (qualified_name TEXT PRIMARY KEY, vec vector(384))",
)
.unwrap();
let empty: i64 = client
.query_one("SELECT count(*) FROM embedding_vectors", &[])
.unwrap()
.get(0);
assert_eq!(empty, 0, "freshly-dropped table must be empty");
let v = random_unit_vector(0xABCDEF12, VEC_DIM);
client
.execute(
"INSERT INTO embedding_vectors (qualified_name, vec) VALUES ($1, $2::text::vector)",
&[&"probe", &pgvector(&v)],
)
.unwrap();
let present: bool = client
.query_one(
"SELECT EXISTS(SELECT 1 FROM embedding_vectors LIMIT 1)",
&[],
)
.unwrap()
.get(0);
assert!(present, "EXISTS probe must report true after insert");
}