use std::ops::Bound;
use std::sync::Arc;
use futures::StreamExt;
use super::common::{
evaluate_bound_key, extract_record_ids_into, resolve_record_batch, resolve_version_stamp,
};
use crate::catalog::{DatabaseId, NamespaceId};
use crate::exec::permission::{PhysicalPermission, should_check_perms};
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::idx::planner::ScanDirection;
use crate::kvs::CachePolicy;
use crate::val::{RecordId, 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();
let check_perms = should_check_perms(&db_ctx, Action::View)?;
let input_stream = buffer_stream(
self.input.execute(ctx)?,
self.input.access_mode(),
self.input.cardinality_hint(),
ctx.root().ctx.config.operator_buffer_size,
);
let referencing_table = self.referencing_table.clone();
let referencing_field = self.referencing_field.clone();
let output_mode = self.output_mode;
let range_start = self.range_start.clone();
let range_end = self.range_end.clone();
let scan_batch_size = ctx.root().ctx.config.scan_batch_size;
let ctx = ctx.clone();
let fetch_full = output_mode == ReferenceScanOutput::FullRecord;
let version_expr = self.version.clone();
let stream = async_stream::try_stream! {
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<
crate::val::TableName,
PhysicalPermission,
> = std::collections::HashMap::new();
let version: Option<u64> = resolve_version_stamp(&ctx, version_expr.as_ref()).await?;
futures::pin_mut!(input_stream);
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.values
.into_iter()
.flat_map(|v| {
let mut rids = Vec::new();
extract_record_ids_into(v, &mut rids);
rids
})
.collect();
for rid in &target_rids {
let (beg, end) = compute_ref_key_range(
ns_id, db_id, rid,
referencing_table.as_ref(),
referencing_field.as_deref(),
&range_start, &range_end,
&ctx,
).await?;
let mut cursor = txn
.open_keys_cursor(beg..end, ScanDirection::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 crate::key::r#ref::Ref::decode_key(key) {
Ok(decoded) => {
rid_batch.push(RecordId {
table: decoded.ft.into_owned(),
key: decoded.fk.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?;
yield ValueBatch { values };
rid_batch.clear();
}
}
}
}
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?;
yield ValueBatch { values };
}
};
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<(Vec<u8>, Vec<u8>), ControlFlow> {
let has_range =
!matches!(range_start, Bound::Unbounded) || !matches!(range_end, Bound::Unbounded);
if has_range {
let table = referencing_table
.context("Range-bounded reference scans require a referencing table")?;
let field = referencing_field.context(
"Cannot scan a specific range of record references without a referencing field",
)?;
let beg = eval_ref_bound(ns_id, db_id, rid, table, field, range_start, true, ctx).await?;
let end = eval_ref_bound(ns_id, db_id, rid, table, field, range_end, false, ctx).await?;
Ok((beg, end))
} else if let Some(table) = referencing_table {
if let Some(field) = referencing_field {
let beg = crate::key::r#ref::ffprefix(ns_id, db_id, &rid.table, &rid.key, table, field)
.context("Failed to create field prefix")?;
let end = crate::key::r#ref::ffsuffix(ns_id, db_id, &rid.table, &rid.key, table, field)
.context("Failed to create field suffix")?;
Ok((beg, end))
} else {
let beg = crate::key::r#ref::ftprefix(ns_id, db_id, &rid.table, &rid.key, table)
.context("Failed to create table prefix")?;
let end = crate::key::r#ref::ftsuffix(ns_id, db_id, &rid.table, &rid.key, table)
.context("Failed to create table suffix")?;
Ok((beg, end))
}
} else {
let beg = crate::key::r#ref::prefix(ns_id, db_id, &rid.table, &rid.key)
.context("Failed to create wildcard prefix")?;
let end = crate::key::r#ref::suffix(ns_id, db_id, &rid.table, &rid.key)
.context("Failed to create wildcard suffix")?;
Ok((beg, end))
}
}
#[allow(clippy::too_many_arguments)]
async fn eval_ref_bound(
ns_id: NamespaceId,
db_id: DatabaseId,
rid: &RecordId,
table: &TableName,
field: &str,
bound: &Bound<Arc<dyn PhysicalExpr>>,
is_start: bool,
ctx: &ExecutionContext,
) -> Result<Vec<u8>, ControlFlow> {
match bound {
Bound::Unbounded => {
if is_start {
crate::key::r#ref::ffprefix(ns_id, db_id, &rid.table, &rid.key, table, field)
.context("Failed to create field prefix")
} else {
crate::key::r#ref::ffsuffix(ns_id, db_id, &rid.table, &rid.key, table, field)
.context("Failed to create field suffix")
}
}
Bound::Included(expr) => {
let fk = evaluate_bound_key(expr, ctx).await?;
if is_start {
crate::key::r#ref::refprefix(ns_id, db_id, &rid.table, &rid.key, table, field, &fk)
.context("Failed to create range key")
} else {
crate::key::r#ref::refsuffix(ns_id, db_id, &rid.table, &rid.key, table, field, &fk)
.context("Failed to create range key")
}
}
Bound::Excluded(expr) => {
let fk = evaluate_bound_key(expr, ctx).await?;
if is_start {
crate::key::r#ref::refsuffix(ns_id, db_id, &rid.table, &rid.key, table, field, &fk)
.context("Failed to create range key")
} else {
crate::key::r#ref::refprefix(ns_id, db_id, &rid.table, &rid.key, table, field, &fk)
.context("Failed to create range key")
}
}
}
}
#[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"));
}
}