use std::borrow::Cow;
use std::ops::Bound;
use std::sync::Arc;
use common::future::stream::{self, Yielder};
use futures::StreamExt;
use surrealdb_datastore::values::inline_cache::{CacheValue, RefCacheEntry};
use super::common::{
evaluate_bound_key, extract_record_ids_into, resolve_record_batch, resolve_version_stamp,
};
use crate::catalog::providers::TableProvider;
use crate::catalog::{DatabaseId, NamespaceId};
use crate::exec::permission::{
PhysicalPermission, should_check_perms, validate_record_user_access,
};
use crate::exec::{
AccessMode, ContextLevel, ControlFlowExt, ExecOperator, ExecutionContext, FlowResult,
OperatorMetrics, PhysicalExpr, ValueBatch, ValueBatchStream, buffer_stream, monitor_stream,
};
use crate::expr::ControlFlow;
use crate::iam::Action;
use crate::key::schema::{
DbRoot, RefCacheKey, ReferenceForeignFieldPrefix, ReferenceForeignTablePrefix,
ReferenceIdPrefix, ReferenceKey,
};
use crate::key::{KVKeyDecode, TypedRange};
use crate::kvs::{CachePolicy, Direction};
use crate::val::{RecordId, RecordIdKey, TableName};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ReferenceScanOutput {
#[default]
RecordId,
FullRecord,
}
#[derive(Debug, Clone)]
pub struct ReferenceScan {
pub(crate) input: Arc<dyn ExecOperator>,
pub(crate) referencing_table: Option<TableName>,
pub(crate) referencing_field: Option<String>,
pub(crate) output_mode: ReferenceScanOutput,
pub(crate) range_start: Bound<Arc<dyn PhysicalExpr>>,
pub(crate) range_end: Bound<Arc<dyn PhysicalExpr>>,
pub(crate) version: Option<Arc<dyn PhysicalExpr>>,
pub(crate) metrics: Arc<OperatorMetrics>,
}
impl ReferenceScan {
pub(crate) fn new(
input: Arc<dyn ExecOperator>,
referencing_table: Option<TableName>,
referencing_field: Option<String>,
output_mode: ReferenceScanOutput,
range_start: Bound<Arc<dyn PhysicalExpr>>,
range_end: Bound<Arc<dyn PhysicalExpr>>,
version: Option<Arc<dyn PhysicalExpr>>,
) -> Self {
Self {
input,
referencing_table,
referencing_field,
output_mode,
range_start,
range_end,
version,
metrics: Arc::new(OperatorMetrics::new()),
}
}
}
impl ExecOperator for ReferenceScan {
fn name(&self) -> &'static str {
"ReferenceScan"
}
fn attrs(&self) -> Vec<(String, String)> {
let mut attrs = vec![(
"table".to_string(),
self.referencing_table
.as_ref()
.map(|t| t.as_str().to_string())
.unwrap_or_else(|| "?".to_string()),
)];
if let Some(field) = &self.referencing_field {
attrs.push(("field".to_string(), field.clone()));
}
if self.output_mode == ReferenceScanOutput::FullRecord {
attrs.push(("output".to_string(), "full_record".to_string()));
}
if !matches!(self.range_start, Bound::Unbounded)
|| !matches!(self.range_end, Bound::Unbounded)
{
attrs.push(("range".to_string(), "bounded".to_string()));
}
attrs
}
fn required_context(&self) -> ContextLevel {
self.input.required_context().max(ContextLevel::Database)
}
fn access_mode(&self) -> AccessMode {
let mut mode = self.input.access_mode();
if let Some(ref version) = self.version {
mode = mode.combine(version.access_mode());
}
mode
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn children(&self) -> Vec<&Arc<dyn ExecOperator>> {
vec![&self.input]
}
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 mut input_stream = buffer_stream(
self.input.execute(ctx)?,
self.input.access_mode(),
self.input.cardinality_hint(),
ctx.root().ctx.config.exec.operator_buffer_size,
);
let scan_args = Arc::new(RefScanArgs {
table: self.referencing_table.clone(),
field: self.referencing_field.clone(),
range_start: self.range_start.clone(),
range_end: self.range_end.clone(),
});
let output_mode = self.output_mode;
let scan_batch_size = ctx.root().ctx.config.exec.scan_batch_size;
let ctx = ctx.clone();
let fetch_full = output_mode == ReferenceScanOutput::FullRecord;
let version_expr = self.version.clone();
let metrics = Arc::clone(&self.metrics);
let record_metrics = metrics.is_enabled();
let stream = stream::try_async_stream(async move |mut yielder: Yielder<_>| {
let txn = ctx.txn();
let ns_id = db_ctx.ns_ctx.ns.namespace_id;
let db_id = db_ctx.db.database_id;
let mut perm_cache: std::collections::HashMap<
surrealdb_strand::TableName,
PhysicalPermission,
> = std::collections::HashMap::new();
let version: Option<u64> = resolve_version_stamp(&ctx, version_expr.as_ref()).await?;
let cache_eligible = version.is_none()
&& matches!(scan_args.range_start, Bound::Unbounded)
&& matches!(scan_args.range_end, Bound::Unbounded);
let cache_filter_prefix: Option<Vec<u8>> = match (&scan_args.table, cache_eligible) {
(Some(table), true) => {
let mut prefix =
storekey::encode_vec(table).map_err(anyhow::Error::from_boxed)?;
if let Some(field) = &scan_args.field {
prefix.extend(
storekey::encode_vec(field.as_str())
.map_err(anyhow::Error::from_boxed)?,
);
}
Some(prefix)
}
_ => None,
};
let mut rid_batch: Vec<RecordId> = Vec::with_capacity(scan_batch_size);
while let Some(batch_result) = input_stream.next().await {
let batch = batch_result?;
let target_rids: Vec<RecordId> = batch
.into_iter()
.flat_map(|v| {
let mut rids = Vec::new();
extract_record_ids_into(v, &mut rids);
rids
})
.collect();
let mut cache_map: std::collections::HashMap<usize, Vec<RefCacheEntry>> =
std::collections::HashMap::new();
if cache_eligible {
let mut table_capped: std::collections::HashMap<TableName, bool> =
std::collections::HashMap::new();
let mut slots: Vec<usize> = Vec::new();
let mut keys: Vec<RefCacheKey> = Vec::new();
for (tgt_idx, rid) in target_rids.iter().enumerate() {
let capped = match table_capped.get(&rid.table) {
Some(capped) => *capped,
None => {
let tb = txn
.get_tb(ns_id, db_id, &rid.table, None)
.await
.context("Failed to resolve the referenced record's table")?;
let capped =
crate::idx::inline_cache::effective_refs_cap(tb.as_deref())
.is_some();
table_capped.insert(rid.table.clone(), capped);
capped
}
};
if !capped {
continue;
}
slots.push(tgt_idx);
keys.push(RefCacheKey {
ns: ns_id,
db: db_id,
tb: Cow::Borrowed(&rid.table),
id: Cow::Borrowed(&rid.key),
});
}
if !keys.is_empty() {
let values = txn
.get_many_key(keys, None)
.await
.context("Failed to read the inline reference caches")?;
let mut hits: u64 = 0;
let mut misses: u64 = 0;
for (tgt_idx, value) in slots.into_iter().zip(values) {
match value {
Some(CacheValue::Live(entries)) => {
hits += 1;
cache_map.insert(tgt_idx, entries);
}
_ => misses += 1,
}
}
if record_metrics {
metrics.add_cache_hits(hits);
metrics.add_cache_misses(misses);
}
}
}
let mut scanned: Vec<Option<TargetRefs>> = Vec::new();
scanned.resize_with(target_rids.len(), || None);
#[cfg(not(target_family = "wasm"))]
let mut fanout_engaged: Option<bool> = None;
let mut win_start = 0usize;
while win_start < target_rids.len() {
let win_end = (win_start + REFERENCE_FANOUT_CONCURRENCY).min(target_rids.len());
#[cfg(not(target_family = "wasm"))]
if fanout_engaged == Some(true) {
let missed: Vec<usize> = (win_start..win_end)
.filter(|tgt_idx| !cache_map.contains_key(tgt_idx))
.collect();
if missed.len() > 1 {
let mut pending: std::collections::VecDeque<_> = missed
.into_iter()
.map(|tgt_idx| {
let ctx = ctx.clone();
let txn = Arc::clone(&txn);
let rid = target_rids[tgt_idx].clone();
let scan_args = Arc::clone(&scan_args);
async move {
let refs = async {
let range = compute_ref_key_range(
ns_id,
db_id,
&rid,
scan_args.table.as_ref(),
scan_args.field.as_deref(),
&scan_args.range_start,
&scan_args.range_end,
&ctx,
)
.await?;
scan_target_refs(&txn, range, version, scan_batch_size)
.await
}
.await;
(tgt_idx, refs)
}
})
.collect();
let mut set = tokio::task::JoinSet::new();
let mut first_err: Option<(usize, ControlFlow)> = None;
let mut join_err: Option<ControlFlow> = None;
loop {
if first_err.is_none() && join_err.is_none() {
while set.len() < REFERENCE_FANOUT_CONCURRENCY {
let Some(task) = pending.pop_front() else {
break;
};
set.spawn(task);
}
}
let Some(joined) = set.join_next().await else {
break;
};
match joined {
Ok((tgt_idx, Ok(refs))) => scanned[tgt_idx] = Some(refs),
Ok((tgt_idx, Err(err))) => {
if first_err
.as_ref()
.is_none_or(|(lowest, _)| tgt_idx < *lowest)
{
first_err = Some((tgt_idx, err));
}
}
Err(e) => {
join_err.get_or_insert_with(|| {
ControlFlow::Err(anyhow::anyhow!(
"reference fan-out task failed: {e}"
))
});
}
}
}
if let Some((_, err)) = first_err {
return Err(err);
}
if let Some(err) = join_err {
return Err(err);
}
}
}
#[cfg(not(target_family = "wasm"))]
let win_timer = fanout_engaged.is_none().then(web_time::Instant::now);
for tgt_idx in win_start..win_end {
let rid = &target_rids[tgt_idx];
if let Some(entries) = cache_map.get(&tgt_idx) {
for entry in entries {
if let Some(prefix) = &cache_filter_prefix
&& !entry.reference.starts_with(prefix)
{
continue;
}
let (ft, ff, fk) =
storekey::decode_borrow::<(TableName, String, RecordIdKey)>(
&entry.reference,
)
.map_err(|e| {
anyhow::anyhow!("Failed to decode a cached reference: {e}")
})?;
if scan_args.table.as_ref().is_some_and(|t| *t != ft) {
continue;
}
if scan_args.field.as_deref().is_some_and(|f| f != ff) {
continue;
}
rid_batch.push(RecordId {
table: ft,
key: fk,
});
if rid_batch.len() >= scan_batch_size {
let values = resolve_record_batch(
&ctx,
&txn,
ns_id,
db_id,
&rid_batch,
fetch_full,
check_perms,
version,
CachePolicy::ReadWrite,
&mut perm_cache,
)
.await?;
yielder.emit(ValueBatch::new(values)).await;
rid_batch.clear();
}
}
continue;
}
let refs = match scanned[tgt_idx].take() {
Some(refs) => refs,
None => {
let range = compute_ref_key_range(
ns_id,
db_id,
rid,
scan_args.table.as_ref(),
scan_args.field.as_deref(),
&scan_args.range_start,
&scan_args.range_end,
&ctx,
)
.await?;
scan_target_refs(&txn, range, version, scan_batch_size).await?
}
};
for found in refs.rids {
rid_batch.push(found);
if rid_batch.len() >= scan_batch_size {
let values = resolve_record_batch(
&ctx,
&txn,
ns_id,
db_id,
&rid_batch,
fetch_full,
check_perms,
version,
CachePolicy::ReadWrite,
&mut perm_cache,
)
.await?;
yielder.emit(ValueBatch::new(values)).await;
rid_batch.clear();
}
}
let Some(range) = refs.resume else {
continue;
};
let mut cursor = txn
.open_keys_cursor(range, Direction::Forward, 0, version)
.await
.context("Failed to open reference cursor")?;
loop {
let mut decode_err: Option<anyhow::Error> = None;
let stats = cursor
.for_each(crate::kvs::NORMAL_BATCH_SIZE, &mut |key| {
match ReferenceKey::decode_key(key) {
Ok(decoded) => {
rid_batch.push(RecordId {
table: decoded.foreign_table.into_owned(),
key: decoded.foreign_key.into_owned(),
});
Ok(std::ops::ControlFlow::Continue(()))
}
Err(e) => {
decode_err = Some(e);
Ok(std::ops::ControlFlow::Break(()))
}
}
})
.await
.context("Failed to scan reference")?;
if let Some(e) = decode_err {
return Err(e).context("Failed to decode ref key")?;
}
if stats.rows == 0 {
break;
}
if rid_batch.len() >= scan_batch_size {
let values = resolve_record_batch(
&ctx,
&txn,
ns_id,
db_id,
&rid_batch,
fetch_full,
check_perms,
version,
CachePolicy::ReadWrite,
&mut perm_cache,
)
.await?;
yielder.emit(ValueBatch::new(values)).await;
rid_batch.clear();
}
}
}
#[cfg(not(target_family = "wasm"))]
if let Some(started) = win_timer {
let per_target =
started.elapsed().as_nanos() as u64 / (win_end - win_start) as u64;
fanout_engaged = Some(per_target >= REFERENCE_SPAWN_PER_TARGET_NANOS);
}
win_start = win_end;
}
}
if !rid_batch.is_empty() {
let values = resolve_record_batch(
&ctx,
&txn,
ns_id,
db_id,
&rid_batch,
fetch_full,
check_perms,
version,
CachePolicy::ReadWrite,
&mut perm_cache,
)
.await?;
yielder.emit(ValueBatch::new(values)).await;
}
Ok(())
});
Ok(monitor_stream(Box::pin(stream), "ReferenceScan", &self.metrics))
}
}
#[allow(clippy::too_many_arguments)]
async fn compute_ref_key_range(
ns_id: NamespaceId,
db_id: DatabaseId,
rid: &RecordId,
referencing_table: Option<&TableName>,
referencing_field: Option<&str>,
range_start: &Bound<Arc<dyn PhysicalExpr>>,
range_end: &Bound<Arc<dyn PhysicalExpr>>,
ctx: &ExecutionContext,
) -> Result<TypedRange<()>, ControlFlow> {
let has_range =
!matches!(range_start, Bound::Unbounded) || !matches!(range_end, Bound::Unbounded);
let prefix = DbRoot {
ns: ns_id,
db: db_id,
};
if has_range {
let table = referencing_table.ok_or_else(|| {
anyhow::anyhow!("Range-bounded reference scans require a referencing table")
})?;
let field = referencing_field.ok_or_else(|| {
anyhow::anyhow!(
"Cannot scan a specific range of record references without a referencing field"
)
})?;
let start = match range_start {
Bound::Included(x) => Bound::Included(Cow::Owned(evaluate_bound_key(x, ctx).await?)),
Bound::Excluded(x) => Bound::Excluded(Cow::Owned(evaluate_bound_key(x, ctx).await?)),
Bound::Unbounded => Bound::Unbounded,
};
let end = match range_end {
Bound::Included(x) => Bound::Included(Cow::Owned(evaluate_bound_key(x, ctx).await?)),
Bound::Excluded(x) => Bound::Excluded(Cow::Owned(evaluate_bound_key(x, ctx).await?)),
Bound::Unbounded => Bound::Unbounded,
};
Ok(ReferenceForeignFieldPrefix {
ns: prefix.ns,
db: prefix.db,
tb: Cow::Borrowed(&rid.table),
id: Cow::Borrowed(&rid.key),
foreign_table: Cow::Borrowed(table),
foreign_field: Cow::Borrowed(field),
}
.range_where((start, end))?)
} else if let Some(table) = referencing_table {
if let Some(field) = referencing_field {
Ok(ReferenceForeignFieldPrefix {
ns: prefix.ns,
db: prefix.db,
tb: Cow::Borrowed(&rid.table),
id: Cow::Borrowed(&rid.key),
foreign_table: Cow::Borrowed(table),
foreign_field: Cow::Borrowed(field),
}
.range()?)
} else {
Ok(ReferenceForeignTablePrefix {
ns: prefix.ns,
db: prefix.db,
tb: Cow::Borrowed(&rid.table),
id: Cow::Borrowed(&rid.key),
foreign_table: Cow::Borrowed(table),
}
.range()?)
}
} else {
Ok(ReferenceIdPrefix {
ns: prefix.ns,
db: prefix.db,
tb: Cow::Borrowed(&rid.table),
id: Cow::Borrowed(&rid.key),
}
.range()?)
}
}
const REFERENCE_FANOUT_CONCURRENCY: usize = 16;
#[cfg(not(target_family = "wasm"))]
const REFERENCE_SPAWN_PER_TARGET_NANOS: u64 = 5_000;
struct RefScanArgs {
table: Option<TableName>,
field: Option<String>,
range_start: Bound<Arc<dyn PhysicalExpr>>,
range_end: Bound<Arc<dyn PhysicalExpr>>,
}
struct TargetRefs {
rids: Vec<RecordId>,
resume: Option<TypedRange<()>>,
}
async fn scan_target_refs(
txn: &crate::kvs::Transaction,
range: TypedRange<()>,
version: Option<u64>,
cap: usize,
) -> Result<TargetRefs, ControlFlow> {
use crate::key::Resumable;
let full = range.clone();
let mut cursor = txn
.open_keys_cursor(range, Direction::Forward, 0, version)
.await
.context("Failed to open reference cursor")?;
let mut rids = Vec::new();
let mut resume: Option<TypedRange<()>> = None;
loop {
let mut decode_err: Option<anyhow::Error> = None;
let stats = cursor
.for_each(crate::kvs::NORMAL_BATCH_SIZE, &mut |key| match ReferenceKey::decode_key(key)
{
Ok(decoded) => {
rids.push(RecordId {
table: decoded.foreign_table.into_owned(),
key: decoded.foreign_key.into_owned(),
});
if rids.len() >= cap {
resume = Some(full.clone().resume_after(key, Direction::Forward));
Ok(std::ops::ControlFlow::Break(()))
} else {
Ok(std::ops::ControlFlow::Continue(()))
}
}
Err(e) => {
decode_err = Some(e);
Ok(std::ops::ControlFlow::Break(()))
}
})
.await
.context("Failed to scan reference")?;
if let Some(e) = decode_err {
return Err(e).context("Failed to decode ref key")?;
}
if resume.is_some() || stats.rows == 0 {
break;
}
}
Ok(TargetRefs {
rids,
resume,
})
}
#[cfg(all(feature = "kv-mem", not(target_family = "wasm")))]
#[cfg(test)]
mod window_tests {
use surrealdb_cnf::ConfigMap;
use super::*;
use crate::exec::operators::test_util::{TestDb, ValuesOperator, collect};
use crate::val::Value;
fn rid(table: &str, key: impl std::fmt::Display) -> Value {
Value::RecordId(RecordId::new(table.into(), key.to_string()))
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn results_are_exact_across_window_and_cap_boundaries() {
const TARGETS: usize = 40;
const HUB_REFS: usize = 10;
let mut setup = String::from(
"DEFINE TABLE person SCHEMALESS; DEFINE TABLE post SCHEMALESS; \
DEFINE FIELD author ON post TYPE record<person> REFERENCE;",
);
setup.push_str("CREATE person:p0;");
for j in 0..HUB_REFS {
setup.push_str(&format!("CREATE post:h{j} SET author = person:p0;"));
}
for i in 1..TARGETS {
setup.push_str(&format!(
"CREATE person:p{i}; CREATE post:q{i} SET author = person:p{i};"
));
}
let db = TestDb::new_with_config(
&setup,
ConfigMap::empty().with_key_value("scan_batch_size", "4"),
)
.await;
let ctx = db.exec_ctx().await;
let scan = ReferenceScan::new(
ValuesOperator::new((0..TARGETS).map(|i| rid("person", format!("p{i}"))).collect()),
Some("post".into()),
Some("author".to_string()),
ReferenceScanOutput::RecordId,
Bound::Unbounded,
Bound::Unbounded,
None,
);
let op: Arc<dyn ExecOperator> = Arc::new(scan);
let mut out = collect(&op, &ctx).await;
out.sort();
let mut expected: Vec<Value> = (0..HUB_REFS)
.map(|j| rid("post", format!("h{j}")))
.chain((1..TARGETS).map(|i| rid("post", format!("q{i}"))))
.collect();
expected.sort();
assert_eq!(out, expected);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::exec::operators::CurrentValueSource;
#[test]
fn test_reference_scan_attrs() {
let scan = ReferenceScan::new(
Arc::new(CurrentValueSource::new()),
Some("post".into()),
Some("author".to_string()),
ReferenceScanOutput::RecordId,
Bound::Unbounded,
Bound::Unbounded,
None,
);
assert_eq!(scan.name(), "ReferenceScan");
let attrs = scan.attrs();
assert!(attrs.iter().any(|(k, v)| k == "table" && v == "post"));
assert!(attrs.iter().any(|(k, v)| k == "field" && v == "author"));
}
}