surrealdb-core 3.3.1

A scalable, distributed, collaborative, document-graph database, for the realtime web
//! HNSW behaviour observed through the query driver.

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(())
}

/// A filtered KNN search evaluates the residual `WHERE` against each visited
/// candidate's record, and fetches those records in batches: one multi-get per
/// neighbourhood for the committed graph (`search_with_filter`), and one per
/// materialisation batch of the pending queue (`search_pendings`). Both paths
/// therefore issue far fewer KV *get operations* than the records they read —
/// `ops_get` well below `keys_read`, where one fetch per candidate would make
/// the two roughly equal.
///
/// The assertions are structural — K matches are found and `ops_get` sits well
/// below `keys_read` — so they hold for any graph; no fixed build seed is
/// needed, which keeps the test free of a process-global `set_var` (the seed
/// is exercised out-of-process by the benchmark harness instead).
#[test(tokio::test(flavor = "multi_thread"))]
async fn test_hnsw_filtered_knn_batches_record_fetches() -> Result<()> {
	// Compaction happens only where this test calls it: the assertions compare
	// the pending set against the compacted graph, which a background
	// compactor would fold away before the first phase reads it.
	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);

	// 500 deterministic 8-d points with a selective `category` (1-in-20),
	// plus an HNSW index. A selective filter makes the search visit many
	// candidates before finding K matches — the case batching helps.
	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;";

	// Run the query on an owned read transaction so we can read its KV
	// metrics, returning (result_count, metrics).
	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()))
	}

	// Before compaction the data is in the pending set, searched by
	// `search_pendings` (batched here too).
	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
	);

	// Compact pending updates into the committed graph, then query again —
	// now `search_with_filter` (per-neighbourhood batching) handles it.
	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
	);

	// Both paths return the same K matching records...
	assert_eq!(pending_len, 10, "pending filtered KNN should return K matches");
	assert_eq!(committed_len, 10, "committed filtered KNN should return K matches");
	// ...and both batch their record fetches: one-per-get fetching gives
	// `ops_get` ~ `keys_read`; batching pulls `ops_get` well below it.
	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(())
}

/// A pending queue deeper than one materialisation batch is scored in several
/// batches, but the top-K is taken across the whole queue. The second query's
/// nearest neighbour is the queue's last record and its runners-up are far
/// earlier ones, so a result kept per batch — or a residual `WHERE` whose
/// per-candidate verdicts did not survive the batch that produced them —
/// returns a different set.
#[test(tokio::test(flavor = "multi_thread"))]
async fn hnsw_filtered_knn_spans_pending_materialisation_batches() -> Result<()> {
	// The queue must still be uncompacted when the queries run.
	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);

	// Record `i` sits at (i, i) and carries `category = i % 100`, so a
	// category selects an evenly spread 1-in-100 subset of the queue. The
	// count is deliberately above the pending materialisation batch size, so
	// the scan needs more than one batch.
	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)
	}

	// Winners at the head of the queue, wholly inside the first batch.
	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]
	);
	// Winners straddling the batch boundary: pts:1025 is read in the second
	// batch, its two runners-up in the first.
	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(())
}

/// A filtered KNN at a `VERSION` reads each visited candidate's record once.
/// Versioned reads bypass the transaction record cache, so the records the
/// per-neighbourhood prefetch reads must be the ones the filter evaluates; a
/// second, per-candidate read would double the record keys the query reads.
/// The same query unversioned reads each record once through the cache. The
/// versioned run additionally re-reads its K results when it materialises
/// them and reads its catalog at the version, so it may exceed the
/// unversioned count by at most 2K keys — far below the hundreds a second
/// read of every visited candidate adds.
#[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());
	// Compaction happens only where this test calls it.
	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}';");
	// The first search loads the graph into the process-wide HNSW cache; warm
	// it so both measured runs read only candidate records and doc-id maps.
	run(&ds, &session, &current).await?;
	run(&ds, &session, &versioned).await?;
	let (current_len, current_keys) = run(&ds, &session, &current).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(())
}