use std::sync::Arc;
use anyhow::Result;
use crate::catalog::providers::CatalogProvider;
use crate::dbs::{NewPlannerStrategy, Session};
use crate::key::schema::RecordKey;
use crate::kvs::{Datastore, QueryRequest, TransactionType};
use crate::val::RecordIdKey;
#[tokio::test(flavor = "multi_thread")]
async fn test_diskann_filtered_knn_batches_record_fetches() -> Result<()> {
let ds = Datastore::builder().without_maintenance_tasks().build_with_path("memory").await?;
{
let tx = ds.transaction(TransactionType::Write).await?;
tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
}
let session = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let n = 500u32;
let cats = 20u32;
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec DISKANN DIMENSION 8 DIST EUCLIDEAN TYPE F32;\n",
);
for i in 0..n {
let mut v = String::new();
for j in 0..8u32 {
if j > 0 {
v.push_str(", ");
}
let f = ((i.wrapping_mul(7).wrapping_add(j.wrapping_mul(131))) % 1000) as f32 / 1000.0;
v.push_str(&format!("{f}f"));
}
setup.push_str(&format!("CREATE pts:{i} SET vec = [{v}], category = {};\n", i % cats));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
async fn run(
ds: &Arc<Datastore>,
session: &Session,
query: &str,
) -> Result<(usize, crate::observe::TransactionMetricsSnapshot)> {
let tx = Arc::new(ds.transaction(TransactionType::Read).await?);
let mut response =
ds.run(QueryRequest::new(query, session).with_transaction(Arc::clone(&tx))).await?;
let len = match response.remove(0).result? {
surrealdb_types::Value::Array(a) => a.len(),
_ => 0,
};
Ok((len, tx.metrics_snapshot_for_test()))
}
let selective = "SELECT id FROM pts \
WHERE vec <|10,400|> [0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f] AND category = 7;";
let (pending_len, pending_m) = run(&ds, &session, selective).await?;
eprintln!(
"PENDING ops_get={} keys_read={} value_bytes_read={} results={pending_len}",
pending_m.ops_get, pending_m.keys_read, pending_m.value_bytes_read
);
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
let (committed_len, committed_m) = run(&ds, &session, selective).await?;
eprintln!(
"COMMITTED ops_get={} keys_read={} value_bytes_read={} results={committed_len}",
committed_m.ops_get, committed_m.keys_read, committed_m.value_bytes_read
);
let nonselective = "SELECT id FROM pts \
WHERE vec <|5,400|> [0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f] AND category < 10;";
let (ns_len, ns_m) = run(&ds, &session, nonselective).await?;
eprintln!(
"NONSELECTIVE ops_get={} keys_read={} value_bytes_read={} results={ns_len}",
ns_m.ops_get, ns_m.keys_read, ns_m.value_bytes_read
);
assert_eq!(pending_len, 10, "pending filtered KNN should return K matches");
assert_eq!(committed_len, 10, "committed filtered KNN should return K matches");
assert!(
u64::from(pending_m.ops_get) * 4 < pending_m.keys_read * 3,
"pending path should batch: ops_get={} keys_read={}",
pending_m.ops_get,
pending_m.keys_read
);
assert!(
u64::from(committed_m.ops_get) * 4 < committed_m.keys_read * 3,
"committed path should batch: ops_get={} keys_read={}",
committed_m.ops_get,
committed_m.keys_read
);
assert_eq!(ns_len, 5, "non-selective filtered KNN should return K matches");
assert!(
ns_m.keys_read * 4 < committed_m.keys_read,
"windowed prefetch should bound non-selective over-fetch: \
non-selective keys_read={} vs selective keys_read={}",
ns_m.keys_read,
committed_m.keys_read
);
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn test_diskann_filtered_knn_skips_missing_record() -> Result<()> {
let ds = Datastore::builder().without_maintenance_tasks().build_with_path("memory").await?;
let db_def = {
let tx = ds.transaction(TransactionType::Write).await?;
let db = tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
db
};
let session = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let mut setup = String::from(
"DEFINE INDEX pt ON pts FIELDS point DISKANN DIMENSION 1 DIST EUCLIDEAN TYPE F32;\n",
);
for i in 1..=12u32 {
let cat = if i % 2 == 1 {
"a"
} else {
"b"
};
setup.push_str(&format!("CREATE pts:{i} SET point = [{}f], category = '{cat}';\n", i * 10));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
{
let tx = ds.transaction(TransactionType::Write).await?;
let tb = surrealdb_strand::TableName::from("pts");
let key = RecordKey {
ns: db_def.namespace_id,
db: db_def.database_id,
tb: std::borrow::Cow::Borrowed(&tb),
id: std::borrow::Cow::Owned(RecordIdKey::Number(1)),
};
tx.del_key(&key).await?;
tx.commit().await?;
}
let query = "SELECT VALUE vector::distance::knn() FROM pts \
WHERE point <|2,40|> [0f] AND category = 'a';";
let mut dists: Vec<f64> =
ds.execute(query, &session, None).await?.remove(0).result?.into_t::<Vec<f64>>()?;
dists.sort_by(f64::total_cmp);
assert_eq!(
dists,
vec![30.0, 50.0],
"missing pts:1 (dist 10) must be excluded and backfilled, got {dists:?}"
);
Ok(())
}
#[cfg(feature = "kv-surrealkv")]
#[tokio::test(flavor = "multi_thread")]
async fn diskann_versioned_filtered_knn_reads_each_candidate_once() -> Result<()> {
let dir = temp_dir::TempDir::new()?;
let path = format!("surrealkv://{}?versioned=true&retention=1h", dir.path().to_string_lossy());
let ds = Datastore::builder().without_maintenance_tasks().build_with_path(&path).await?;
{
let tx = ds.transaction(TransactionType::Write).await?;
tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
}
let session = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec DISKANN DIMENSION 8 DIST EUCLIDEAN TYPE F32;\n",
);
for i in 0..500u32 {
let mut v = String::new();
for j in 0..8u32 {
if j > 0 {
v.push_str(", ");
}
let f = ((i.wrapping_mul(7).wrapping_add(j.wrapping_mul(131))) % 1000) as f32 / 1000.0;
v.push_str(&format!("{f}f"));
}
setup.push_str(&format!("CREATE pts:{i} SET vec = [{v}], category = {};\n", i % 20));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
let mut response = ds.execute("RETURN <string> time::now();", &session, None).await?;
let surrealdb_types::Value::String(stamp) = response.remove(0).result? else {
panic!("Expected a datetime string");
};
async fn run(ds: &Arc<Datastore>, session: &Session, query: &str) -> Result<(usize, u64)> {
let tx = Arc::new(ds.transaction(TransactionType::Read).await?);
let mut response =
ds.run(QueryRequest::new(query, session).with_transaction(Arc::clone(&tx))).await?;
let len = match response.remove(0).result? {
surrealdb_types::Value::Array(a) => a.len(),
_ => 0,
};
Ok((len, tx.metrics_snapshot_for_test().keys_read))
}
const K: u64 = 10;
let query = format!(
"SELECT id FROM pts \
WHERE vec <|{K},400|> [0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f] AND category = 7"
);
let current = format!("{query};");
let versioned = format!("{query} VERSION d'{stamp}';");
run(&ds, &session, ¤t).await?;
run(&ds, &session, &versioned).await?;
let (current_len, current_keys) = run(&ds, &session, ¤t).await?;
let (versioned_len, versioned_keys) = run(&ds, &session, &versioned).await?;
eprintln!("CURRENT keys_read={current_keys} VERSIONED keys_read={versioned_keys}");
assert_eq!(current_len as u64, K, "filtered KNN should return K matches");
assert_eq!(versioned_len as u64, K, "versioned filtered KNN should return K matches");
assert!(
versioned_keys <= current_keys + 2 * K,
"a versioned filtered KNN should read each candidate record once: \
versioned keys_read={versioned_keys}, current keys_read={current_keys}"
);
ds.shutdown().await?;
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn diskann_selective_filter_finds_admitted_neighbours_past_the_search_list() -> Result<()> {
let ds = Datastore::builder().without_maintenance_tasks().build_with_path("memory").await?;
{
let tx = ds.transaction(TransactionType::Write).await?;
tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
}
let owner = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let record = Session::for_record(
"test",
"test",
"user",
crate::types::PublicValue::String("user:1".to_owned()),
)
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let mut setup = String::from(
"DEFINE TABLE pts SCHEMALESS PERMISSIONS FOR select WHERE category = 7 \
FOR create, update, delete NONE;\n\
DEFINE INDEX emb ON pts FIELDS vec DISKANN DIMENSION 8 DIST EUCLIDEAN TYPE F32;\n",
);
let mut state = 0x9E37_79B9_7F4A_7C15u64;
for i in 0..1000u32 {
let mut v = String::new();
for j in 0..8u32 {
if j > 0 {
v.push_str(", ");
}
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
let f = (state.wrapping_mul(0x2545_F491_4F6C_DD1D) >> 40) as f32 / (1u64 << 24) as f32;
v.push_str(&format!("{f}f"));
}
setup.push_str(&format!("CREATE pts:{i} SET vec = [{v}], category = {};\n", i % 100));
}
for response in ds.execute(&setup, &owner, None).await? {
response.result?;
}
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
async fn ids(ds: &Datastore, session: &Session, query: &str) -> Result<Vec<String>> {
let mut response = ds.execute(query, session, None).await?;
let surrealdb_types::Value::Array(rows) = response.remove(0).result? else {
panic!("Expected array result");
};
let mut ids: Vec<String> = rows.into_iter().map(|v| format!("{v:?}")).collect();
ids.sort();
Ok(ids)
}
let q = "[0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f]";
let truth = ids(
&ds,
&owner,
&format!("SELECT VALUE id FROM pts WHERE vec <|5,EUCLIDEAN|> {q} AND category = 7;"),
)
.await?;
assert_eq!(truth.len(), 5, "brute force should find K admitted rows");
let narrow =
ids(&ds, &record, &format!("SELECT VALUE id FROM pts WHERE vec <|5,10|> {q};")).await?;
assert_eq!(narrow.len(), 5, "a narrow search list should still return K visible rows");
let permitted =
ids(&ds, &record, &format!("SELECT VALUE id FROM pts WHERE vec <|5,100|> {q};")).await?;
assert_eq!(
permitted, truth,
"a record user's bare KNN should return its K nearest visible rows"
);
let conditioned = ids(
&ds,
&owner,
&format!("SELECT VALUE id FROM pts WHERE vec <|5,100|> {q} AND category = 7;"),
)
.await?;
assert_eq!(conditioned, truth, "a residual condition should return its K nearest matches");
Ok(())
}