use std::collections::{HashSet, VecDeque};
use std::sync::Arc;
use common::future::stream::{self, Yielder};
use reblessive::TreeStack;
use roaring::RoaringTreemap;
use surrealdb_types::ToSql;
use super::bitmap::BitmapNode;
use super::common::{fetch_and_filter_records_batch, resolve_version_stamp};
use super::pipeline::{ScanPipeline, build_field_state};
use super::resolved::ResolvedTableContext;
use crate::catalog::{DatabaseId, Distance, Error, Index, NamespaceId, VectorType};
use crate::exec::index::access_path::IndexRef;
use crate::exec::operators::{KnnTopKHeap, check_cancelled, extract_vector};
use crate::exec::permission::{
PhysicalPermission, PhysicalTableSelect, convert_permission_to_physical_runtime,
should_check_perms, validate_record_user_access,
};
use crate::exec::{
AccessMode, CardinalityHint, ContextLevel, EvalContext, ExecOperator, ExecutionContext,
FlowResult, OperatorMetrics, PhysicalExpr, ValueBatch, ValueBatchStream, monitor_stream,
};
use crate::expr::{Cond, ControlFlow, ControlFlowExt, Idiom};
use crate::iam::Action;
use crate::idx::docids::TableDocIds;
use crate::idx::trees::gate::{CachedTableSelect, CandidateFetchCounter};
use crate::idx::trees::vector::{score_raw_vector, typed_query_vector};
use crate::idx::trees::{KnnCondFilter, KnnIteratorResult};
use crate::kvs::{CachePolicy, Transaction};
use crate::legacy::knn::LegacyCondition;
use crate::val::{Number, Object, RecordId, TableName, Value};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub(crate) enum PrefilterTier {
Fallback = 1,
Exact = 2,
Graph = 3,
GraphUnboosted = 4,
}
impl PrefilterTier {
pub(crate) fn label(code: u8) -> Option<&'static str> {
match code {
1 => Some("fallback"),
2 => Some("exact"),
3 => Some("graph"),
4 => Some("graph_unboosted"),
_ => None,
}
}
}
pub(crate) fn choose_prefilter_tier(allow_len: u64, t_exact: u64, t_boost: u64) -> PrefilterTier {
if allow_len <= t_exact {
return PrefilterTier::Exact;
}
if allow_len <= t_boost {
PrefilterTier::Graph
} else {
PrefilterTier::GraphUnboosted
}
}
pub(crate) fn boosted_ef(ef: u32, boost: u32, ef_max: u32) -> u32 {
ef.saturating_mul(boost).min(ef_max).max(ef)
}
#[derive(Debug, Clone)]
pub(crate) struct KnnPrefilter {
pub(crate) node: Arc<BitmapNode>,
pub(crate) node_dyn: Arc<dyn ExecOperator>,
pub(crate) residual: Option<Cond>,
pub(crate) residual_phys: Option<Arc<dyn PhysicalExpr>>,
pub(crate) uncovered_matches: bool,
}
#[derive(Debug)]
pub struct KnnScan {
pub index_ref: IndexRef,
pub vector: Vec<Number>,
pub k: u32,
pub ef: u32,
pub table_name: surrealdb_strand::TableName,
pub(crate) version: Option<Arc<dyn PhysicalExpr>>,
pub(crate) resolved: Option<ResolvedTableContext>,
pub(crate) metrics: Arc<OperatorMetrics>,
pub(crate) knn_context: Option<Arc<crate::exec::function::KnnContext>>,
pub(crate) residual_cond: Option<Cond>,
pub(crate) prefilter: Option<KnnPrefilter>,
pub(crate) matches_condition: Option<Arc<dyn PhysicalExpr>>,
pub(crate) tier: Arc<std::sync::atomic::AtomicU8>,
pub(crate) records_unread: Arc<std::sync::atomic::AtomicBool>,
pub(crate) fields_read_as_stored: bool,
pub(crate) needed_fields: Option<Option<HashSet<String>>>,
}
impl KnnScan {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
index_ref: IndexRef,
vector: Vec<Number>,
k: u32,
ef: u32,
table_name: surrealdb_strand::TableName,
version: Option<Arc<dyn PhysicalExpr>>,
knn_context: Option<Arc<crate::exec::function::KnnContext>>,
residual_cond: Option<Cond>,
prefilter: Option<KnnPrefilter>,
needed_fields: Option<Option<HashSet<String>>>,
) -> Self {
Self {
index_ref,
vector,
k,
ef,
table_name,
version,
resolved: None,
metrics: Arc::new(OperatorMetrics::new()),
knn_context,
residual_cond,
prefilter,
matches_condition: None,
tier: Arc::new(std::sync::atomic::AtomicU8::new(0)),
records_unread: Arc::new(std::sync::atomic::AtomicBool::new(false)),
fields_read_as_stored: false,
needed_fields,
}
}
pub(crate) fn with_resolved(mut self, resolved: ResolvedTableContext) -> Self {
self.resolved = Some(resolved);
self
}
pub(crate) fn with_fields_read_as_stored(mut self, read_as_stored: bool) -> Self {
self.fields_read_as_stored = read_as_stored;
self
}
pub(crate) fn with_matches_condition(mut self, phys: Option<Arc<dyn PhysicalExpr>>) -> Self {
self.matches_condition = phys;
self
}
}
impl KnnScan {
pub(crate) const NAME: &'static str = "KnnScan";
}
impl ExecOperator for KnnScan {
fn name(&self) -> &'static str {
Self::NAME
}
fn attrs(&self) -> Vec<(String, String)> {
let mut attrs = vec![
("index".to_string(), self.index_ref.name.to_string()),
("k".to_string(), self.k.to_string()),
("ef".to_string(), self.ef.to_string()),
("dimension".to_string(), self.vector.len().to_string()),
];
let shown_cond = match &self.prefilter {
Some(p) => p.residual.as_ref(),
None => self.residual_cond.as_ref(),
};
if let Some(cond) = shown_cond {
attrs.push(("predicate".to_string(), cond.0.to_sql()));
}
if let Some(label) =
PrefilterTier::label(self.tier.load(std::sync::atomic::Ordering::Relaxed))
{
attrs.push(("prefilter_tier".to_string(), label.to_string()));
}
if self.records_unread.load(std::sync::atomic::Ordering::Relaxed) {
attrs.push(("records".to_string(), "not read".to_string()));
}
attrs
}
fn children(&self) -> Vec<&Arc<dyn ExecOperator>> {
match &self.prefilter {
Some(p) => vec![&p.node_dyn],
None => vec![],
}
}
fn required_context(&self) -> ContextLevel {
ContextLevel::Database
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadOnly
}
fn cardinality_hint(&self) -> CardinalityHint {
CardinalityHint::Bounded(self.k as usize)
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn execute(&self, ctx: &ExecutionContext) -> FlowResult<ValueBatchStream> {
let db_ctx = ctx.database()?.clone();
validate_record_user_access(&db_ctx)?;
let check_perms = should_check_perms(&db_ctx, Action::View)?;
let index_ref = self.index_ref.clone();
let vector = self.vector.clone();
let k = self.k;
let ef = self.ef;
let table_name = self.table_name.clone();
let version_expr = self.version.clone();
let knn_context = self.knn_context.clone();
let residual_cond = self.residual_cond.clone();
let matches_condition = self.matches_condition.clone();
let prefilter = self.prefilter.clone();
let tier_cell = Arc::clone(&self.tier);
let records_unread = Arc::clone(&self.records_unread);
let fields_read_as_stored = self.fields_read_as_stored;
let op_metrics = Arc::clone(&self.metrics);
let resolved = self.resolved.clone();
let needed_fields = self.needed_fields.clone();
let ctx = ctx.clone();
let stream = stream::try_async_stream(async move |mut yielder: Yielder<_>| {
let db_ctx = ctx.database().context("KnnScan requires database context")?;
let ns = Arc::clone(&db_ctx.ns_ctx.ns);
let db = Arc::clone(&db_ctx.db);
let txn = ctx.txn();
let version: Option<u64> = resolve_version_stamp(&ctx, version_expr.as_ref()).await?;
let root = ctx.root();
let frozen_ctx = &root.ctx;
let (select_permission, table_id) = if let Some(ref res) = resolved {
let perm = res.select_permission(check_perms);
(perm, res.table_def.table_id)
} else {
let table_def = db_ctx
.get_table_def(&table_name, version)
.await
.context("Failed to get table")?;
let table_def = match table_def {
Some(def) => def,
None => {
Err(ControlFlow::Err(anyhow::Error::new(Error::TbNotFound {
name: table_name.clone(),
})))?;
unreachable!()
}
};
let select_permission = if check_perms {
convert_permission_to_physical_runtime(&table_def.permissions.select, &ctx)
.await
.context("Failed to convert permission")?
} else {
PhysicalPermission::Allow
};
(select_permission, table_def.table_id)
};
if matches!(select_permission, PhysicalPermission::Deny) {
return Ok(());
}
let field_state = match &needed_fields {
Some(nf) => {
if let Some(ref res) = resolved {
res.field_state_for_projection(nf.as_ref())
} else {
build_field_state(&ctx, &table_name, check_perms, nf.as_ref()).await?
}
}
None => super::pipeline::FieldState::empty(),
};
let index_def = index_ref.definition();
let ikb = crate::idx::IndexKeyBase::new(
ns.namespace_id,
db.database_id,
index_def.table_name.clone(),
index_def.index_id,
);
let mut tier: Option<PrefilterTier> = None;
let mut allow_list: Option<RoaringTreemap> = None;
let mut in_traversal_cond: Option<Cond> = residual_cond;
if let Some(pf) = prefilter.as_ref().filter(|_| index_def.uses_shared_doc_ids()) {
if version.is_some() {
tier = Some(PrefilterTier::Fallback);
} else {
match pf
.node
.build_allowlist(&ctx, &table_name, *surrealdb_cnf::BITMAP_BRANCH_BUDGET)
.await?
{
Some(bm) => {
let t = choose_prefilter_tier(
bm.len(),
*surrealdb_cnf::KNN_PREFILTER_EXACT_THRESHOLD,
*surrealdb_cnf::KNN_PREFILTER_EF_BOOST_THRESHOLD,
);
in_traversal_cond = pf.residual.clone();
allow_list = Some(bm);
tier = Some(t);
}
None => {
tier = Some(PrefilterTier::Fallback);
}
}
}
}
let uncovered_matches =
prefilter.as_ref().is_some_and(|p| p.uncovered_matches) && allow_list.is_some();
let matches_phys = if allow_list.is_none() || uncovered_matches {
matches_condition
} else {
None
};
if matches_phys.is_some() {
in_traversal_cond = None;
}
tier_cell.store(tier.map_or(0, |t| t as u8), std::sync::atomic::Ordering::Relaxed);
let search_read_records;
let knn_results = if matches!(tier, Some(PrefilterTier::Exact)) {
search_read_records = true;
let allow = allow_list.as_ref().expect("exact tier implies an allow-list");
let (distance, vector_type, dimension) = match &index_def.index {
Index::Hnsw(p) => (p.distance.clone(), p.vector_type, p.dimension as usize),
#[cfg(diskann)]
Index::DiskAnn(p) => (p.distance.clone(), p.vector_type, p.dimension as usize),
#[cfg(not(diskann))]
Index::DiskAnn(_) => {
Err(ControlFlow::Err(anyhow::anyhow!(
"DISKANN indexes require a 64-bit, non-WASM platform"
)))?;
unreachable!()
}
_ => {
Err(ControlFlow::Err(anyhow::anyhow!(
"Index '{}' is not an ANN index",
index_def.name
)))?;
unreachable!()
}
};
let field = match index_def.cols.first() {
Some(f) => f.clone(),
None => {
Err(ControlFlow::Err(anyhow::anyhow!(
"Index '{}' has no indexed column",
index_def.name
)))?;
unreachable!()
}
};
exact_prefiltered_knn(
&ctx,
&txn,
ns.namespace_id,
db.database_id,
&table_name,
&field,
&distance,
vector_type,
dimension,
&vector,
k as usize,
allow,
matches_phys
.as_ref()
.or_else(|| prefilter.as_ref().and_then(|p| p.residual_phys.as_ref())),
&select_permission,
check_perms,
)
.await?
} else {
let (ef_eff, search_allow) = match tier {
Some(PrefilterTier::Graph) => (
boosted_ef(
ef,
*surrealdb_cnf::KNN_PREFILTER_EF_BOOST,
*surrealdb_cnf::KNN_PREFILTER_EF_MAX,
),
allow_list.as_ref(),
),
Some(PrefilterTier::GraphUnboosted) => (ef, allow_list.as_ref()),
_ => (ef, None),
};
let needs_perm_gate = !matches!(select_permission, PhysicalPermission::Allow);
let select_gate = || {
CachedTableSelect::Gate(Arc::new(PhysicalTableSelect::new(
select_permission.clone(),
ctx.clone(),
)))
};
let fetch_counter = || Arc::clone(&op_metrics) as Arc<dyn CandidateFetchCounter>;
let cond_filter = if let Some(phys) = &matches_phys {
Some(KnnCondFilter {
select_gate: select_gate(),
cond: Some(Arc::new(crate::legacy::knn::PhysicalCondition::new(
Arc::clone(phys),
ctx.clone(),
))),
metrics: Some(fetch_counter()),
version,
})
} else {
match (&in_traversal_cond, ctx.options()) {
(Some(cond), Some(opt)) => Some(KnnCondFilter {
select_gate: select_gate(),
cond: Some(Arc::new(LegacyCondition::new(
frozen_ctx,
opt,
Arc::new(cond.clone()),
))),
metrics: Some(fetch_counter()),
version,
}),
(None, _) if needs_perm_gate => Some(KnnCondFilter {
select_gate: select_gate(),
cond: None,
metrics: Some(fetch_counter()),
version,
}),
_ => None,
}
};
search_read_records = cond_filter.is_some();
match &index_def.index {
Index::Hnsw(hnsw_params) => {
let hnsw_index = frozen_ctx
.get_index_stores()
.get_index_hnsw(
frozen_ctx,
table_id,
&ikb,
hnsw_params,
index_def.format_version,
)
.await
.context("Failed to get HNSW index")?;
let mut stack = TreeStack::new();
stack
.enter(|stk| {
let hnsw_index = &hnsw_index;
let vector = &vector;
async move {
hnsw_index
.knn_search(
frozen_ctx,
stk,
vector,
k as usize,
ef_eff as usize,
cond_filter,
search_allow,
)
.await
}
})
.finish()
.await
.context("HNSW KNN search failed")?
}
#[cfg(diskann)]
Index::DiskAnn(diskann_params) => {
let diskann_index = frozen_ctx
.get_index_stores()
.get_index_diskann(
index_def.format_version,
table_id,
&ikb,
diskann_params,
)
.await
.context("Failed to get DiskANN index")?;
diskann_index
.check_state()
.await
.context("Failed to check DiskANN index state")?;
let mut stack = TreeStack::new();
stack
.enter(|stk| {
let diskann_index = &diskann_index;
let vector = &vector;
async move {
diskann_index
.knn_search(
frozen_ctx,
stk,
vector,
k as usize,
ef_eff as usize,
cond_filter,
search_allow,
)
.await
}
})
.finish()
.await
.context("DiskANN KNN search failed")?
}
#[cfg(not(diskann))]
Index::DiskAnn(_) => {
Err(ControlFlow::Err(anyhow::anyhow!(
"DISKANN indexes require a 64-bit, non-WASM platform"
)))?;
unreachable!()
}
_ => {
Err(ControlFlow::Err(anyhow::anyhow!(
"Index '{}' is not an ANN index",
index_def.name
)))?;
unreachable!()
}
}
};
let mut rids = Vec::with_capacity(knn_results.len());
if let Some(ref knn_ctx) = knn_context {
for (rid, distance, _) in &knn_results {
knn_ctx.insert(rid.as_ref().clone(), Number::Float(*distance)).await;
rids.push(rid.as_ref().clone());
}
} else {
for (rid, _, _) in &knn_results {
rids.push(rid.as_ref().clone());
}
}
if emits_record_ids_only(
needed_fields.as_ref(),
&select_permission,
version,
fields_read_as_stored,
) {
records_unread.store(!search_read_records, std::sync::atomic::Ordering::Relaxed);
let values: Vec<Value> = rids.into_iter().map(record_id_row).collect();
if !values.is_empty() {
yielder.emit(ValueBatch::new(values)).await;
}
return Ok(());
}
let mut pipeline = ScanPipeline::new(
PhysicalPermission::Allow,
None,
field_state,
check_perms,
None,
0,
);
let mut values = fetch_and_filter_records_batch(
&ctx,
&txn,
ns.namespace_id,
db.database_id,
&rids,
&select_permission,
check_perms,
version,
CachePolicy::ReadWrite,
)
.await?;
pipeline.process_batch(&mut values, &ctx).await?;
if !values.is_empty() {
yielder.emit(ValueBatch::new(values)).await;
}
Ok(())
});
Ok(monitor_stream(Box::pin(stream), "KnnScan", &self.metrics))
}
}
fn emits_record_ids_only(
needed_fields: Option<&Option<HashSet<String>>>,
select_permission: &PhysicalPermission,
version: Option<u64>,
fields_read_as_stored: bool,
) -> bool {
matches!(needed_fields, Some(Some(needed)) if needed.iter().all(|f| f == "id"))
&& fields_read_as_stored
&& version.is_none()
&& matches!(select_permission, PhysicalPermission::Allow)
}
fn record_id_row(rid: RecordId) -> Value {
let mut row = Object::default();
row.insert("id", Value::RecordId(rid));
Value::Object(row)
}
const EXACT_TIER_BATCH_SIZE: usize = 1000;
#[expect(clippy::too_many_arguments)]
async fn exact_prefiltered_knn(
ctx: &ExecutionContext,
txn: &Arc<Transaction>,
ns: NamespaceId,
db: DatabaseId,
table_name: &TableName,
field: &Idiom,
distance: &Distance,
vector_type: VectorType,
dimension: usize,
query_vector: &[Number],
k: usize,
allow: &RoaringTreemap,
residual_phys: Option<&Arc<dyn PhysicalExpr>>,
select_permission: &PhysicalPermission,
check_perms: bool,
) -> std::result::Result<VecDeque<KnnIteratorResult>, ControlFlow> {
let query = typed_query_vector(vector_type, dimension, query_vector)
.context("Invalid KNN query vector")?;
let doc_ids = TableDocIds::new(ns, db, table_name.clone());
let mut heap: KnnTopKHeap<Arc<crate::val::RecordId>> = KnnTopKHeap::new(k);
let mut iter = allow.iter();
loop {
check_cancelled(ctx)?;
let chunk: Vec<u64> = iter.by_ref().take(EXACT_TIER_BATCH_SIZE).collect();
if chunk.is_empty() {
break;
}
let keys = doc_ids
.get_record_ids_batch(txn.as_ref(), &chunk)
.await
.context("Failed to resolve doc-IDs to record IDs")?;
let mut rids = Vec::with_capacity(keys.len());
for key in keys.into_iter().flatten() {
rids.push(crate::val::RecordId {
table: table_name.clone(),
key,
});
}
let mut values = fetch_and_filter_records_batch(
ctx,
txn,
ns,
db,
&rids,
select_permission,
check_perms,
None,
CachePolicy::ReadWrite,
)
.await?;
if let Some(phys) = residual_phys {
let eval_ctx = EvalContext::from_exec_ctx(ctx);
let results = phys.evaluate_batch(eval_ctx, &values).await?;
let mut kept = Vec::with_capacity(values.len());
for (value, res) in values.into_iter().zip(results) {
if res.is_truthy() {
kept.push(value);
}
}
values = kept;
}
for value in values {
let Some(record_vec) = extract_vector(&value, field) else {
continue;
};
let Some(dist) =
score_raw_vector(distance, vector_type, dimension, &query, &record_vec)
else {
continue;
};
let dist = Number::Float(dist);
let Value::Object(ref obj) = value else {
continue;
};
let Some(Value::RecordId(rid)) = obj.get("id") else {
continue;
};
heap.offer(dist, Arc::new(rid.clone()));
}
}
let mut res = VecDeque::with_capacity(k);
for (dist, rid) in heap.into_sorted_nearest_first() {
res.push_back((rid, dist.to_float(), None));
}
Ok(res)
}
#[cfg(test)]
mod tests {
use super::{PrefilterTier, boosted_ef, choose_prefilter_tier};
#[test]
fn prefilter_tier_decision_table() {
use PrefilterTier::*;
const T_EXACT: u64 = 2_000;
const T_BOOST: u64 = 100_000;
let choose = |len| choose_prefilter_tier(len, T_EXACT, T_BOOST);
assert_eq!(choose(0), Exact);
assert_eq!(choose(T_EXACT), Exact);
assert_eq!(choose(T_EXACT + 1), Graph);
assert_eq!(choose(T_BOOST), Graph);
assert_eq!(choose(T_BOOST + 1), GraphUnboosted);
assert_eq!(choose(u64::MAX), GraphUnboosted);
}
#[test]
fn boosted_ef_bounds() {
assert_eq!(boosted_ef(40, 4, 1_024), 160);
assert_eq!(boosted_ef(400, 4, 1_024), 1_024);
assert_eq!(boosted_ef(2_000, 4, 1_024), 2_000);
assert_eq!(boosted_ef(40, 0, 1_024), 40);
assert_eq!(boosted_ef(u32::MAX, 2, u32::MAX), u32::MAX);
}
#[test]
fn prefilter_tier_labels() {
assert_eq!(PrefilterTier::label(0), None);
assert_eq!(PrefilterTier::label(PrefilterTier::Fallback as u8), Some("fallback"));
assert_eq!(PrefilterTier::label(PrefilterTier::Exact as u8), Some("exact"));
assert_eq!(PrefilterTier::label(PrefilterTier::Graph as u8), Some("graph"));
assert_eq!(
PrefilterTier::label(PrefilterTier::GraphUnboosted as u8),
Some("graph_unboosted")
);
assert_eq!(PrefilterTier::label(5), None);
}
}