use std::collections::HashMap;
use std::convert::Infallible;
use std::sync::Arc;
use super::super::Planner;
use super::super::util::{apply_literal_part, try_literal_to_value};
use crate::catalog::providers::TableProvider;
use crate::catalog::{DatabaseId, NamespaceId};
use crate::expr::visit::{MutVisitor, Visit, VisitMut, Visitor};
use crate::expr::{Cond, Expr, Idiom, Part};
use crate::kvs::Transaction;
use crate::val::{RecordId, RecordIdKey, TableName, Value};
const FOLD_FETCH_BUDGET: u32 = 32;
impl<'ctx> Planner<'ctx> {
pub(crate) async fn fold_constant_record_idioms(&self, cond: &mut Cond) {
let Some(txn) = self.txn.as_ref() else {
return;
};
let (Some(ns_name), Some(db_name)) = (self.ns.as_deref(), self.db.as_deref()) else {
return;
};
if self.should_check_perms_for_view(ns_name, db_name) {
return;
}
let mut collector = CandidateCollector {
idioms: Vec::new(),
};
let _ = collector.visit_expr(&cond.0);
if collector.idioms.is_empty() {
return;
}
let Some((ns_id, db_id)) = self.ns_db_ids().await else {
return;
};
let mut folded: HashMap<Idiom, Option<Expr>> = HashMap::new();
let mut computed_cache: HashMap<TableName, bool> = HashMap::new();
let mut fetch_budget = FOLD_FETCH_BUDGET;
for (idiom, rid) in collector.idioms {
if folded.contains_key(&idiom) {
continue;
}
let value = fold_record_walk(
txn,
ns_id,
db_id,
rid,
&idiom,
&mut computed_cache,
&mut fetch_budget,
)
.await;
folded.insert(idiom, value.and_then(|v| literal_if_round_trips(&v)));
}
if !folded.values().any(Option::is_some) {
return;
}
let _ = FoldApplier {
folded: &folded,
}
.visit_mut_expr(&mut cond.0);
}
}
fn literal_if_round_trips(value: &Value) -> Option<Expr> {
let expr = value.clone().into_literal();
match &expr {
Expr::Literal(lit) if try_literal_to_value(lit).as_ref() == Some(value) => Some(expr),
_ => None,
}
}
async fn fold_record_walk(
txn: &Arc<Transaction>,
ns_id: NamespaceId,
db_id: DatabaseId,
start: RecordId,
idiom: &Idiom,
computed_cache: &mut HashMap<TableName, bool>,
fetch_budget: &mut u32,
) -> Option<Value> {
let mut pivot = Value::RecordId(start);
for part in idiom.0.iter().skip(1) {
let Part::Field(name) = part else {
pivot = apply_literal_part(pivot, part)?;
continue;
};
let name = name.as_str();
pivot = match pivot {
Value::RecordId(rid) => {
if let RecordIdKey::Object(obj) = &rid.key
&& let Some(component) = obj.get(name)
{
component.clone()
} else {
if *fetch_budget == 0 {
tracing::debug!(
table = %rid.table,
"plan-time record fetch budget exhausted in \
fold_constant_record_idioms; leaving record idiom unfolded",
);
return None;
}
*fetch_budget -= 1;
if table_has_computed_fields(txn, ns_id, db_id, &rid.table, computed_cache)
.await?
{
return None;
}
let record =
match txn.get_record(ns_id, db_id, &rid.table, &rid.key, None).await {
Ok(record) => record,
Err(e) => {
tracing::debug!(
table = %rid.table,
error = %e,
"plan-time record read failed in \
fold_constant_record_idioms; leaving record idiom unfolded",
);
return None;
}
};
match &record.data {
Value::Object(obj) => obj.get(name).cloned().unwrap_or(Value::None),
_ => Value::None,
}
}
}
Value::Object(obj) => obj.get(name).cloned().unwrap_or(Value::None),
Value::Array(_) | Value::Geometry(_) => return None,
_ => Value::None,
};
}
Some(pivot)
}
async fn table_has_computed_fields(
txn: &Arc<Transaction>,
ns_id: NamespaceId,
db_id: DatabaseId,
table: &TableName,
cache: &mut HashMap<TableName, bool>,
) -> Option<bool> {
if let Some(has) = cache.get(table) {
return Some(*has);
}
let fields = match txn.all_tb_fields(ns_id, db_id, table, None).await {
Ok(fields) => fields,
Err(e) => {
tracing::debug!(
table = %table,
error = %e,
"plan-time field list failed in fold_constant_record_idioms; \
leaving record idiom unfolded",
);
return None;
}
};
let has = fields.iter().any(|fd| fd.computed.is_some());
cache.insert(table.clone(), has);
Some(has)
}
fn candidate_record_root(idiom: &Idiom) -> Option<RecordId> {
let mut parts = idiom.0.iter();
let Some(Part::Start(Expr::Literal(lit))) = parts.next() else {
return None;
};
if !matches!(parts.next(), Some(Part::Field(_))) {
return None;
}
let foldable = |p: &Part| match p {
Part::Field(_) | Part::First | Part::Last => true,
Part::Value(index) => matches!(index, Expr::Literal(_)),
_ => false,
};
if !parts.all(foldable) {
return None;
}
match try_literal_to_value(lit) {
Some(Value::RecordId(rid)) => Some(rid),
_ => None,
}
}
fn is_fold_boundary(expr: &Expr) -> bool {
matches!(
expr,
Expr::Select(_)
| Expr::Create(_)
| Expr::Update(_)
| Expr::Upsert(_)
| Expr::Delete(_)
| Expr::Relate(_)
| Expr::Insert(_)
| Expr::Define(_)
| Expr::Remove(_)
| Expr::Rebuild(_)
| Expr::Alter(_)
| Expr::Info(_)
| Expr::Foreach(_)
| Expr::Let(_)
| Expr::Block(_)
| Expr::Closure(_)
| Expr::Sleep(_)
)
}
struct CandidateCollector {
idioms: Vec<(Idiom, RecordId)>,
}
impl Visitor for CandidateCollector {
type Error = Infallible;
fn visit_expr(&mut self, expr: &Expr) -> Result<(), Self::Error> {
if is_fold_boundary(expr) {
return Ok(());
}
if let Expr::Idiom(idiom) = expr
&& let Some(rid) = candidate_record_root(idiom)
{
self.idioms.push((idiom.clone(), rid));
return Ok(());
}
expr.visit(self)
}
}
struct FoldApplier<'a> {
folded: &'a HashMap<Idiom, Option<Expr>>,
}
impl MutVisitor for FoldApplier<'_> {
type Error = Infallible;
fn visit_mut_expr(&mut self, expr: &mut Expr) -> Result<(), Self::Error> {
if is_fold_boundary(expr) {
return Ok(());
}
if let Expr::Idiom(idiom) = expr
&& let Some(Some(folded)) = self.folded.get(idiom)
{
*expr = folded.clone();
return Ok(());
}
expr.visit_mut(self)
}
}