mod common;
use common::pgwire_harness::TestServer;
async fn create_vector_collection(server: &TestServer, name: &str, dim: usize) {
server
.exec(&format!("CREATE COLLECTION {name}"))
.await
.unwrap();
server
.exec(&format!(
"CREATE VECTOR INDEX idx_{name}_emb ON {name} METRIC l2 DIM {dim}"
))
.await
.unwrap();
}
async fn insert_vector(server: &TestServer, coll: &str, id: &str, v: &[f32]) {
let arr = v
.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",");
server
.exec(&format!(
"INSERT INTO {coll} (id, embedding) VALUES ('{id}', ARRAY[{arr}])"
))
.await
.unwrap();
}
async fn ranked_ids(server: &TestServer, coll: &str, query: &[f32], k: usize) -> Vec<String> {
let arr = query
.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",");
let rows = server
.query_rows(&format!(
"SELECT id FROM {coll} \
ORDER BY vector_distance(embedding, ARRAY[{arr}]) \
LIMIT {k}"
))
.await
.unwrap();
rows.into_iter().map(|r| r[0].clone()).collect()
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn insert_close_vector_visible_in_txn_ranked_then_commit_persists() {
let server = TestServer::start().await;
create_vector_collection(&server, "vec_ov_close", 3).await;
insert_vector(&server, "vec_ov_close", "base_far", &[10.0, 10.0, 10.0]).await;
server.exec("BEGIN").await.unwrap();
insert_vector(&server, "vec_ov_close", "staged_close", &[0.0, 0.0, 0.1]).await;
let query = [0.0, 0.0, 0.0];
let in_txn = ranked_ids(&server, "vec_ov_close", &query, 2).await;
assert_eq!(
in_txn.len(),
2,
"expected both rows visible in-tx: {in_txn:?}"
);
assert_eq!(
in_txn[0], "staged_close",
"staged vector close to the query must rank first before COMMIT: {in_txn:?}"
);
server.client.simple_query("COMMIT").await.unwrap();
let after_commit = ranked_ids(&server, "vec_ov_close", &query, 2).await;
assert_eq!(
after_commit.len(),
2,
"both rows remain searchable after COMMIT: {after_commit:?}"
);
assert_eq!(
after_commit[0], "staged_close",
"post-commit vector search must project the committed row's PK \
'staged_close' (nearest to the query), not its surrogate: {after_commit:?}"
);
assert_eq!(
after_commit[1], "base_far",
"the base row must also project its PK 'base_far', proving surrogate → PK \
resolution holds for the whole document+vector-index class: {after_commit:?}"
);
let persisted = server
.query_rows("SELECT id FROM vec_ov_close WHERE id = 'staged_close'")
.await
.unwrap();
assert_eq!(
persisted.len(),
1,
"committed insert must persist under its PK: {persisted:?}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn insert_visible_in_txn_then_rollback_excludes() {
let server = TestServer::start().await;
create_vector_collection(&server, "vec_ov_rb", 3).await;
insert_vector(&server, "vec_ov_rb", "base1", &[5.0, 5.0, 5.0]).await;
server.exec("BEGIN").await.unwrap();
insert_vector(&server, "vec_ov_rb", "staged_rb", &[0.0, 0.0, 0.0]).await;
let query = [0.0, 0.0, 0.0];
let in_txn = ranked_ids(&server, "vec_ov_rb", &query, 5).await;
assert!(
in_txn.iter().any(|id| id == "staged_rb"),
"in-tx search must include the staged insert: {in_txn:?}"
);
server.client.simple_query("ROLLBACK").await.unwrap();
let after_rollback = ranked_ids(&server, "vec_ov_rb", &query, 5).await;
assert!(
after_rollback.iter().all(|id| id != "staged_rb"),
"ROLLBACK must leave no durable trace of the staged insert: {after_rollback:?}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn far_staged_vector_ranks_last_and_needs_larger_k() {
let server = TestServer::start().await;
create_vector_collection(&server, "vec_ov_far", 3).await;
for i in 0..3 {
insert_vector(
&server,
"vec_ov_far",
&format!("near{i}"),
&[i as f32 * 0.1, 0.0, 0.0],
)
.await;
}
server.exec("BEGIN").await.unwrap();
insert_vector(&server, "vec_ov_far", "staged_far", &[100.0, 100.0, 100.0]).await;
let query = [0.0, 0.0, 0.0];
let small_k = ranked_ids(&server, "vec_ov_far", &query, 3).await;
assert_eq!(small_k.len(), 3);
assert!(
small_k.iter().all(|id| id != "staged_far"),
"far staged vector must not displace closer base rows at k=3: {small_k:?}"
);
let large_k = ranked_ids(&server, "vec_ov_far", &query, 4).await;
assert_eq!(large_k.len(), 4, "expected 4 rows at k=4: {large_k:?}");
assert_eq!(
large_k[3], "staged_far",
"far staged vector must rank last at k=4: {large_k:?}"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn base_and_staged_vectors_interleave_by_true_distance() {
let server = TestServer::start().await;
create_vector_collection(&server, "vec_ov_mix", 3).await;
insert_vector(&server, "vec_ov_mix", "base_1", &[1.0, 0.0, 0.0]).await;
insert_vector(&server, "vec_ov_mix", "base_5", &[5.0, 0.0, 0.0]).await;
insert_vector(&server, "vec_ov_mix", "base_9", &[9.0, 0.0, 0.0]).await;
server.exec("BEGIN").await.unwrap();
insert_vector(&server, "vec_ov_mix", "staged_3", &[3.0, 0.0, 0.0]).await;
insert_vector(&server, "vec_ov_mix", "staged_7", &[7.0, 0.0, 0.0]).await;
let query = [0.0, 0.0, 0.0];
let ranked = ranked_ids(&server, "vec_ov_mix", &query, 5).await;
assert_eq!(ranked.len(), 5, "expected all five rows: {ranked:?}");
assert_eq!(
ranked[1], "staged_3",
"staged vector at distance 3 must rank second (between base 1 and 5): {ranked:?}"
);
assert_eq!(
ranked[3], "staged_7",
"staged vector at distance 7 must rank fourth (between base 5 and 9): {ranked:?}"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}