use std::sync::Arc;
use anyhow::Result;
use test_log::test;
use crate::catalog::providers::{CatalogProvider, TableProvider};
use crate::dbs::Session;
use crate::idx::IndexKeyBase;
use crate::idx::trees::hnsw::HnswState;
use crate::kvs::{Datastore, QueryRequest, TransactionType};
#[test(tokio::test(flavor = "multi_thread"))]
async fn test_hnsw_inner_product_smoke() -> Result<()> {
let ds = Datastore::new("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");
let sql = "
DEFINE INDEX hnsw_pts ON pts FIELDS point HNSW DIMENSION 2 DIST INNER_PRODUCT TYPE F32 EFC 100 M 12;
CREATE pts:1 SET point = [1f, 0f];
CREATE pts:2 SET point = [2f, 0f];
CREATE pts:3 SET point = [0f, 1f];
";
for response in ds.execute(sql, &session, None).await? {
response.result?;
}
let mut response =
ds.execute("SELECT id FROM pts WHERE point <|2,40|> [1f, 0f];", &session, None).await?;
let result = response.remove(0).result?;
let surrealdb_types::Value::Array(result) = result else {
panic!("Expected array result");
};
assert_eq!(result.len(), 2);
Ok(())
}
#[test(tokio::test(flavor = "multi_thread"))]
async fn test_hnsw_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(crate::dbs::NewPlannerStrategy::AllReadOnlyStatements);
let n = 500u32;
let cats = 20u32;
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec HNSW DIMENSION 8 DIST EUCLIDEAN TYPE F32 EFC 200 M 12;\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?;
}
let query = "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;";
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 (pending_len, pending_m) = run(&ds, &session, query).await?;
eprintln!(
"PENDING ops_get={} keys_read={} results={pending_len}",
pending_m.ops_get, pending_m.keys_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, query).await?;
eprintln!(
"COMMITTED ops_get={} keys_read={} results={committed_len}",
committed_m.ops_get, committed_m.keys_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
);
Ok(())
}
#[test(tokio::test(flavor = "multi_thread"))]
async fn hnsw_filtered_knn_spans_pending_materialisation_batches() -> 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(crate::dbs::NewPlannerStrategy::AllReadOnlyStatements);
let n = 1100u32;
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec HNSW DIMENSION 2 DIST EUCLIDEAN TYPE F32 EFC 100 M 12;\n",
);
for i in 1..=n {
setup
.push_str(&format!("CREATE pts:{i} SET vec = [{i}f, {i}f], category = {};\n", i % 100));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
async fn ids(ds: &Arc<Datastore>, session: &Session, query: &str) -> Result<Vec<i64>> {
let mut response = ds.execute(query, session, None).await?;
let surrealdb_types::Value::Array(rows) = response.remove(0).result? else {
panic!("expected an array result");
};
let mut ids: Vec<i64> = rows
.iter()
.map(|row| match row {
surrealdb_types::Value::Object(o) => match o.get("id") {
Some(surrealdb_types::Value::RecordId(rid)) => match &rid.key {
surrealdb_types::RecordIdKey::Number(n) => *n,
other => panic!("unexpected record key: {other:?}"),
},
other => panic!("unexpected id field: {other:?}"),
},
other => panic!("unexpected row: {other:?}"),
})
.collect();
ids.sort();
Ok(ids)
}
assert_eq!(
ids(&ds, &session, "SELECT id FROM pts WHERE vec <|5,400|> [0f, 0f] AND category = 7;")
.await?,
vec![7, 107, 207, 307, 407]
);
assert_eq!(
ids(
&ds,
&session,
"SELECT id FROM pts WHERE vec <|3,400|> [1025f, 1025f] AND category = 25;"
)
.await?,
vec![825, 925, 1025]
);
Ok(())
}
#[test(tokio::test(flavor = "multi_thread"))]
async fn hnsw_blocking_define_index_compacts_pending_vectors() -> Result<()> {
let ds = Datastore::new("memory").await?;
let db = {
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");
let sql = "
CREATE pts:1 SET point = [1f, 2f];
CREATE pts:2 SET point = [2f, 3f];
CREATE pts:3 SET point = [3f, 4f];
DEFINE INDEX hnsw_pts ON pts FIELDS point HNSW DIMENSION 2 DIST EUCLIDEAN TYPE F32 EFC 100 M 12;
";
for response in ds.execute(sql, &session, None).await? {
response.result?;
}
let tx = ds.transaction(TransactionType::Read).await?;
let tb = "pts".into();
let ix =
tx.get_tb_index(db.namespace_id, db.database_id, &tb, "hnsw_pts", None).await?.unwrap();
let ikb = IndexKeyBase::new(db.namespace_id, db.database_id, tb, ix.index_id);
let pending_records = tx.getr(ikb.new_hr_range()?, None).await?;
let pending_appends = tx.getr_raw(ikb.new_hp_range()?, None).await?;
let state: HnswState = tx.get_key(&ikb.new_hs_key(), None).await?.unwrap();
tx.cancel().await?;
assert!(pending_records.is_empty());
assert!(pending_appends.is_empty());
assert_eq!(state.next_element_id, 3);
Ok(())
}
#[test(tokio::test(flavor = "multi_thread"))]
async fn hnsw_query_reads_record_keyed_pending_vectors() -> Result<()> {
let ds = Datastore::new("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");
let sql = "
DEFINE INDEX hnsw_pts ON pts FIELDS point HNSW DIMENSION 2 DIST EUCLIDEAN TYPE F32 EFC 100 M 12;
CREATE pts:1 SET point = [1f, 2f];
CREATE pts:2 SET point = [2f, 3f];
CREATE pts:3 SET point = [3f, 4f];
";
for response in ds.execute(sql, &session, None).await? {
response.result?;
}
let mut response =
ds.execute("SELECT id FROM pts WHERE point <|2,40|> [1f, 2f];", &session, None).await?;
let result = response.remove(0).result?;
let surrealdb_types::Value::Array(result) = result else {
panic!("Expected array result");
};
assert_eq!(result.len(), 2);
Ok(())
}
#[cfg(feature = "kv-surrealkv")]
#[test(tokio::test(flavor = "multi_thread"))]
async fn hnsw_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(crate::dbs::NewPlannerStrategy::AllReadOnlyStatements);
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec HNSW DIMENSION 8 DIST EUCLIDEAN TYPE F32 EFC 200 M 12;\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(())
}