use std::borrow::Cow;
use std::ops::Bound;
use std::sync::Arc;
use common::future::stream::{Yielder, try_async_stream};
use tracing::instrument;
use crate::catalog::{DatabaseId, Error, Index, NamespaceId, table_select_permission};
use crate::err::EngineError;
use crate::exec::operators::scan::index_count::sum_index_count_deltas;
use crate::exec::permission::{
PhysicalPermission, check_permission_for_value, 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::{ControlFlow, ControlFlowExt};
use crate::iam::Action;
use crate::key::schema::{RecordKey, RecordPrefix};
use crate::key::{KVKeyDecode, KVValue, RawRange};
use crate::val::{Number, Object, RecordIdKey, RecordIdKeyRange, TableName, Value};
#[derive(Debug, Clone)]
pub struct CountScan {
pub(crate) source: Arc<dyn PhysicalExpr>,
pub(crate) version: Option<Arc<dyn PhysicalExpr>>,
pub(crate) field_names: Vec<String>,
pub(crate) metrics: Arc<OperatorMetrics>,
}
impl CountScan {
pub(crate) fn new(
source: Arc<dyn PhysicalExpr>,
version: Option<Arc<dyn PhysicalExpr>>,
field_names: Vec<String>,
) -> Self {
debug_assert!(!field_names.is_empty(), "CountScan requires at least one field name");
Self {
source,
version,
field_names,
metrics: Arc::new(OperatorMetrics::new()),
}
}
}
impl ExecOperator for CountScan {
fn name(&self) -> &'static str {
"CountScan"
}
fn attrs(&self) -> Vec<(String, String)> {
vec![("source".to_string(), self.source.to_sql())]
}
fn required_context(&self) -> ContextLevel {
let exprs_ctx = [Some(&self.source), self.version.as_ref()]
.into_iter()
.flatten()
.map(|e| e.required_context())
.max()
.unwrap_or(ContextLevel::Root);
exprs_ctx.max(ContextLevel::Database)
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn expressions(&self) -> Vec<(&str, &Arc<dyn PhysicalExpr>)> {
let mut exprs = vec![("source", &self.source)];
if let Some(ref version) = self.version {
exprs.push(("version", version));
}
exprs
}
fn access_mode(&self) -> AccessMode {
let version_mode =
self.version.as_ref().map(|e| e.access_mode()).unwrap_or(AccessMode::ReadOnly);
self.source.access_mode().combine(version_mode)
}
fn cardinality_hint(&self) -> CardinalityHint {
CardinalityHint::AtMostOne
}
#[instrument(name = "CountScan::execute", level = "trace", skip_all)]
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 source_expr = Arc::clone(&self.source);
let version = self.version.clone();
let field_names = self.field_names.clone();
let ctx = ctx.clone();
let stream = try_async_stream(async move |mut yielder: Yielder<_>| {
let db_ctx = ctx.database().context("CountScan requires database context")?;
let txn = ctx.txn();
let ns = Arc::clone(&db_ctx.ns_ctx.ns);
let db = Arc::clone(&db_ctx.db);
let version: Option<u64> = match &version {
Some(expr) => {
let eval_ctx = EvalContext::from_exec_ctx(&ctx);
let v = expr.evaluate(eval_ctx).await?;
Some(
v.cast_to::<crate::val::Datetime>()
.map_err(|e| anyhow::anyhow!("{e}"))?
.to_version_stamp(txn.timestamp_impl().as_ref())?,
)
}
None => ctx.version_stamp(),
};
let eval_ctx = EvalContext::from_exec_ctx(&ctx);
let table_value = source_expr.evaluate(eval_ctx).await?;
let (table_name, rid) = match table_value {
Value::Table(t) => (t, None),
Value::RecordId(rid) => (rid.table.clone(), Some(rid)),
_ => {
Err(ControlFlow::Err(anyhow::anyhow!(
"CountScan received a non-table source"
)))?;
unreachable!()
}
};
let table_def =
db_ctx.get_table_def(&table_name, version).await.context("Failed to get table")?;
if table_def.is_none() {
Err(ControlFlow::Err(anyhow::Error::new(Error::TbNotFound {
name: table_name.clone(),
})))?;
}
if table_def.as_deref().is_some_and(|def| {
crate::kvs::lightweight::lightweight_relation(&def.table_type).is_some()
}) {
Err(ControlFlow::Err(anyhow::anyhow!(
"a LIGHTWEIGHT relation cannot be scanned by this execution path; \
reference the table statically so the record-less scan can serve it"
)))?;
}
let select_permission = if check_perms {
let catalog_perm = table_select_permission(table_def.as_deref());
convert_permission_to_physical_runtime(catalog_perm, &ctx)
.await
.context("Failed to convert permission")?
} else {
PhysicalPermission::Allow
};
match select_permission {
PhysicalPermission::Deny => {
return Ok(());
}
PhysicalPermission::Conditional(_) => {
let count = count_with_perm_fallback(
&ctx,
ns.namespace_id,
db.database_id,
&table_name,
rid.as_ref(),
version,
&select_permission,
)
.await?;
yielder.emit(make_count_batch(count, &field_names)).await;
return Ok(());
}
PhysicalPermission::Allow => {
}
}
let count = if let Some(ref rid) = rid {
count_range(ns.namespace_id, db.database_id, &rid.table, &rid.key, &txn, version)
.await?
} else {
if let None = version
&& let Some(indexes) = db_ctx.get_table_indexes(&table_name, version).await.ok()
&& let Some(ix_def) =
indexes.iter().find(|ix| matches!(&ix.index, Index::Count(None)))
{
sum_index_count_deltas(
&ctx,
&txn,
ns.namespace_id,
db.database_id,
&table_name,
ix_def.index_id,
)
.await?
} else {
let range = RecordPrefix {
ns: ns.namespace_id,
db: db.database_id,
tb: Cow::Borrowed(&table_name),
}
.range()?;
txn.count(range, version).await.context("Failed to count table records")?
}
};
yielder.emit(make_count_batch(count, &field_names)).await;
Ok(())
});
Ok(monitor_stream(Box::pin(stream), "CountScan", &self.metrics))
}
}
fn make_count_batch(count: usize, field_names: &[String]) -> ValueBatch {
let mut obj = Object::default();
let count_val = Value::Number(Number::Int(count as i64));
for name in field_names {
obj.insert(name.clone(), count_val.clone());
}
ValueBatch::new(vec![Value::Object(obj)])
}
async fn count_range(
ns_id: NamespaceId,
db_id: DatabaseId,
table: &TableName,
key: &RecordIdKey,
txn: &crate::kvs::Transaction,
version: Option<u64>,
) -> Result<usize, ControlFlow> {
match key {
RecordIdKey::Range(range) => {
let range = record_key_range(ns_id, db_id, table, range)?;
txn.count(range, version).await.context("Failed to count range records")
}
_ => {
let record_key = RecordKey {
ns: ns_id,
db: db_id,
tb: Cow::Borrowed(table),
id: Cow::Borrowed(key),
};
let exists = txn
.exists_key(&record_key, version)
.await
.context("Failed to check record existence")?;
Ok(usize::from(exists))
}
}
}
pub(crate) fn record_key_range(
ns_id: NamespaceId,
db_id: DatabaseId,
table: &TableName,
range: &RecordIdKeyRange,
) -> Result<RawRange, ControlFlow> {
let start = match &range.start {
Bound::Unbounded => Bound::Unbounded,
Bound::Included(v) => Bound::Included(Cow::Borrowed(v)),
Bound::Excluded(v) => Bound::Excluded(Cow::Borrowed(v)),
};
let end = match &range.end {
Bound::Unbounded => Bound::Unbounded,
Bound::Included(v) => Bound::Included(Cow::Borrowed(v)),
Bound::Excluded(v) => Bound::Excluded(Cow::Borrowed(v)),
};
Ok(RecordPrefix {
ns: ns_id,
db: db_id,
tb: Cow::Borrowed(table),
}
.range_where((start, end))?)
}
async fn count_with_perm_fallback(
ctx: &ExecutionContext,
ns_id: NamespaceId,
db_id: DatabaseId,
table_name: &TableName,
rid: Option<&crate::val::RecordId>,
version: Option<u64>,
permission: &PhysicalPermission,
) -> Result<usize, ControlFlow> {
let txn = ctx.txn();
let range = if let Some(rid) = rid {
match &rid.key {
RecordIdKey::Range(range) => record_key_range(ns_id, db_id, &rid.table, range)?,
_ => {
let Some(value) =
crate::exec::operators::fetch::fetch_raw_record(ctx, rid, version).await?
else {
return Ok(0);
};
let allowed = check_perm_value(ctx, &value, permission).await?;
return Ok(usize::from(allowed));
}
}
} else {
RecordPrefix {
ns: ns_id,
db: db_id,
tb: Cow::Borrowed(table_name),
}
.range()?
};
let mut cursor = txn
.open_vals_cursor_raw(range, crate::kvs::Direction::Forward, 0, version)
.await
.context("Failed to open scan cursor")?;
let mut count = 0usize;
loop {
if ctx.cancellation().is_cancelled() {
return Err(ControlFlow::Err(anyhow::anyhow!(EngineError::QueryCancelled)));
}
let batch = cursor
.next_batch(crate::kvs::NORMAL_BATCH_SIZE)
.await
.context("Failed to scan record")?;
if batch.is_empty() {
break;
}
for (key, val) in &batch {
let decoded_key = RecordKey::decode_key(key).context("Failed to decode record key")?;
let rid_val = crate::val::RecordId {
table: decoded_key.tb.into_owned(),
key: decoded_key.id.into_owned(),
};
let record = crate::catalog::Record::kv_decode_value(val, rid_val)
.context("Failed to deserialize record")?;
let value = record.data;
let allowed = check_permission_for_value(permission, &value, None, ctx)
.await
.map_err(ControlFlow::Err)?;
if allowed {
count += 1;
}
}
}
Ok(count)
}
async fn check_perm_value(
ctx: &ExecutionContext,
value: &Value,
permission: &PhysicalPermission,
) -> Result<bool, ControlFlow> {
check_permission_for_value(permission, value, None, ctx).await.map_err(ControlFlow::Err)
}