use rudb_common::{Field, LogicalType, Result, Value};
use rudb_functions::FILE_ROW_NUMBER;
use std::collections::HashMap;
use rudb_plan::{ColumnBinding, Expr, Node, NodeRef, Plan, Slice, SortKey};
use crate::pass::{Context, Pass};
use crate::walk;
pub const WORTH_FETCHING: u64 = 1024;
pub const WORTH_DEFERRING: usize = 8;
#[derive(Debug, Clone, Copy)]
pub struct LateMaterialization;
impl Pass for LateMaterialization {
fn name(&self) -> &'static str {
"late_materialization"
}
fn run(&self, plan: &mut Plan, _context: &Context) -> Result<()> {
defer(plan);
Ok(())
}
}
pub fn defer(plan: &mut Plan) {
let mut deferred = false;
let root = rewrite(plan, plan.root(), &mut deferred);
if !deferred {
return;
}
plan.set_root(root);
crate::columns::prune(plan);
}
fn rewrite(plan: &mut Plan, at: NodeRef, deferred: &mut bool) -> NodeRef {
let children = plan.node(at).children();
let rebuilt: Vec<NodeRef> =
children.into_iter().flatten().map(|child| rewrite(plan, child, deferred)).collect();
let mut here = at;
let moved = children.into_iter().flatten().zip(&rebuilt).any(|(was, &now)| was != now);
if moved {
let mut node = plan.node(at).clone();
replace_children(&mut node, &rebuilt);
here = plan.add_node(node);
}
match fetch(plan, here) {
Some(above) => {
*deferred = true;
above
}
None => here,
}
}
fn replace_children(node: &mut Node, children: &[NodeRef]) {
match node {
Node::Filter { input, .. }
| Node::Project { input, .. }
| Node::Aggregate { input, .. }
| Node::Sort { input, .. }
| Node::Limit { input, .. }
| Node::TopN { input, .. }
| Node::Fetch { input, .. }
| Node::Distinct { input, .. } => *input = children[0],
Node::Join { left, right, .. }
| Node::CrossProduct { left, right }
| Node::SetOp { left, right, .. } => {
*left = children[0];
*right = children[1];
}
Node::Get { .. } | Node::Dummy | Node::Values { .. } | Node::TableFunction { .. } => {}
}
}
fn fetch(plan: &mut Plan, at: NodeRef) -> Option<NodeRef> {
let Node::TopN { input, keys, count, offset } = *plan.node(at) else { return None };
if count.saturating_add(offset) > WORTH_FETCHING {
return None;
}
let Node::Project { input: under, index, exprs, names } = *plan.node(input) else {
return None;
};
let held: Vec<_> = plan.expr_list(exprs).to_vec();
let labels: Vec<_> = plan.name_list(names).to_vec();
let ordering: Vec<SortKey> = plan.sort_key_list(keys).to_vec();
if held.len() < ordering.len() + WORTH_DEFERRING {
return None;
}
let mut wanted = Vec::with_capacity(ordering.len());
for key in &ordering {
match *plan.expr(key.expr) {
Expr::Column(binding) if binding.table == index => {
let at = binding.column as usize;
wanted.push((*held.get(at)?, *labels.get(at)?));
}
_ => return None,
}
}
let chain = chain(plan, under)?;
let scan = *chain.last()?;
let columns = file_columns(plan, &chain, &held);
let deferred_projects = if columns.is_none() {
let scan_index = file_index(plan, scan)?;
let projects = projects(plan, &chain, index, &held, &labels);
replayable(plan, scan_index, &projects).then_some((scan_index, projects))
} else {
None
};
if columns.is_none() && deferred_projects.is_none() {
return None;
}
let columns = match columns {
Some(columns) => columns,
None => {
(0..file_width(plan, scan)?).map(|at| u32::try_from(at).ok()).collect::<Option<_>>()?
}
};
let mut carried = number(plan, scan)?;
for &node in chain.iter().rev().skip(1) {
carried = carry(plan, node, carried);
}
let row = plan.add_expr(Expr::Column(carried), LogicalType::BigInt);
let narrow = narrow(plan, under, &wanted, row);
let above = top(plan, narrow, &ordering, count, offset);
let deferred = fields(plan, scan, &columns);
let args = match *plan.node(scan) {
Node::TableFunction { args, .. } => args,
_ => return None,
};
let ordinal = plan.add_expr(
Expr::Column(ColumnBinding::new(narrow_index(plan, narrow), wanted.len() as u32)),
LogicalType::BigInt,
);
let fetched_index = if deferred_projects.is_some() { fresh(plan) } else { index };
let fetched = plan.add_node(Node::Fetch {
input: above,
index: fetched_index,
args,
columns: deferred,
row: ordinal,
});
match deferred_projects {
Some((scan_index, projects)) => {
Some(replay(plan, fetched, fetched_index, scan_index, projects))
}
None => Some(fetched),
}
}
fn file_index(plan: &Plan, scan: NodeRef) -> Option<u32> {
match *plan.node(scan) {
Node::TableFunction { index, .. } => Some(index),
_ => None,
}
}
fn file_width(plan: &Plan, scan: NodeRef) -> Option<usize> {
match *plan.node(scan) {
Node::TableFunction { columns, .. } => Some(
plan.field_list(columns).iter().filter(|field| field.name != FILE_ROW_NUMBER).count(),
),
_ => None,
}
}
fn projects(
plan: &Plan,
chain: &[NodeRef],
outer_index: u32,
outer_exprs: &[u32],
outer_names: &[u32],
) -> Vec<(u32, Vec<u32>, Vec<u32>)> {
let mut found = Vec::new();
for &node in chain.iter().rev() {
if let Node::Project { index, exprs, names, .. } = *plan.node(node) {
found.push((index, plan.expr_list(exprs).to_vec(), plan.name_list(names).to_vec()));
}
}
found.push((outer_index, outer_exprs.to_vec(), outer_names.to_vec()));
found
}
fn replayable(plan: &Plan, scan_index: u32, projects: &[(u32, Vec<u32>, Vec<u32>)]) -> bool {
let mut tables = std::collections::HashSet::from([scan_index]);
for (index, exprs, _) in projects {
for &expr in exprs {
let mut missing = false;
walk::columns(plan, expr, &mut |binding| missing |= !tables.contains(&binding.table));
if missing {
return false;
}
}
tables.insert(*index);
}
true
}
fn replay(
plan: &mut Plan,
mut input: NodeRef,
fetched_index: u32,
scan_index: u32,
projects: Vec<(u32, Vec<u32>, Vec<u32>)>,
) -> NodeRef {
let mut tables = HashMap::from([(scan_index, fetched_index)]);
let count = projects.len();
for (at, (old_index, exprs, names)) in projects.into_iter().enumerate() {
let rewritten: Vec<u32> =
exprs.into_iter().map(|expr| rebase(plan, expr, &tables)).collect();
let exprs = plan.add_expr_list(&rewritten);
let names = plan.add_name_list(&names);
let index = if at + 1 == count { old_index } else { fresh(plan) };
input = plan.add_node(Node::Project { input, index, exprs, names });
tables.insert(old_index, index);
}
input
}
fn rebase(plan: &mut Plan, expr: u32, tables: &HashMap<u32, u32>) -> u32 {
if let Expr::Column(binding) = *plan.expr(expr) {
let table = tables.get(&binding.table).copied().unwrap_or(binding.table);
return plan.add_expr(
Expr::Column(ColumnBinding::new(table, binding.column)),
plan.expr_type(expr).clone(),
);
}
walk::rebuild(plan, expr, &mut |plan, child| rebase(plan, child, tables))
}
fn narrow_index(plan: &Plan, node: NodeRef) -> u32 {
plan.node(node).table_index().unwrap_or(0)
}
fn narrow(plan: &mut Plan, input: NodeRef, wanted: &[(u32, u32)], row: u32) -> NodeRef {
let index = fresh(plan);
let mut exprs: Vec<u32> = wanted.iter().map(|&(expr, _)| expr).collect();
let mut names: Vec<u32> = wanted.iter().map(|&(_, name)| name).collect();
exprs.push(row);
names.push(plan.intern(FILE_ROW_NUMBER));
let exprs = plan.add_expr_list(&exprs);
let names = plan.add_name_list(&names);
plan.add_node(Node::Project { input, index, exprs, names })
}
fn top(plan: &mut Plan, input: NodeRef, ordering: &[SortKey], count: u64, offset: u64) -> NodeRef {
let index = narrow_index(plan, input);
let mut keys = Vec::with_capacity(ordering.len());
for (at, key) in ordering.iter().enumerate() {
let column = u32::try_from(at).unwrap_or(u32::MAX);
let ty = plan.expr_type(key.expr).clone();
let expr = plan.add_expr(Expr::Column(ColumnBinding::new(index, column)), ty);
keys.push(SortKey { expr, descending: key.descending, nulls_first: key.nulls_first });
}
let keys = plan.add_sort_keys(&keys);
plan.add_node(Node::TopN { input, keys, count, offset })
}
fn fields(plan: &mut Plan, scan: NodeRef, columns: &[u32]) -> Slice {
let held = match *plan.node(scan) {
Node::TableFunction { columns, .. } => plan.field_list(columns).to_vec(),
_ => Vec::new(),
};
let wanted: Vec<Field> =
columns.iter().filter_map(|&at| held.get(at as usize).cloned()).collect();
plan.add_fields(&wanted)
}
fn chain(plan: &Plan, at: NodeRef) -> Option<Vec<NodeRef>> {
let mut found = vec![at];
let mut node = at;
loop {
match *plan.node(node) {
Node::TableFunction { .. } => return Some(found),
Node::Project { input, .. }
| Node::Filter { input, .. }
| Node::Sort { input, .. }
| Node::Limit { input, .. }
| Node::TopN { input, .. }
| Node::Distinct { input, .. } => {
node = input;
found.push(node);
}
_ => return None,
}
}
}
fn file_columns(plan: &Plan, chain: &[NodeRef], exprs: &[u32]) -> Option<Vec<u32>> {
let mut carried = Vec::with_capacity(exprs.len());
for &expr in exprs {
match *plan.expr(expr) {
Expr::Column(binding) => carried.push(binding),
_ => return None,
}
}
for &node in chain {
match *plan.node(node) {
Node::Project { index, exprs, .. } => {
let held = plan.expr_list(exprs);
let mut next = Vec::with_capacity(carried.len());
for binding in &carried {
if binding.table != index {
return None;
}
match *plan.expr(*held.get(binding.column as usize)?) {
Expr::Column(below) => next.push(below),
_ => return None,
}
}
carried = next;
}
Node::TableFunction { index, .. } => {
if carried.iter().any(|binding| binding.table != index) {
return None;
}
return Some(carried.into_iter().map(|binding| binding.column).collect());
}
_ => {}
}
}
None
}
fn number(plan: &mut Plan, scan: NodeRef) -> Option<ColumnBinding> {
let Node::TableFunction { index, function, args, options, settings, columns } =
*plan.node(scan)
else {
return None;
};
if plan.string(function) != "read_parquet" || plan.expr_list(args).len() != 1 {
return None;
}
let mut fields = plan.field_list(columns).to_vec();
if fields.iter().any(|field| field.name == FILE_ROW_NUMBER) {
return None;
}
let at = u32::try_from(fields.len()).ok()?;
fields.push(Field::required(FILE_ROW_NUMBER.to_string(), LogicalType::BigInt));
let widened = plan.add_fields(&fields);
let mut names = plan.name_list(options).to_vec();
let mut values = plan.expr_list(settings).to_vec();
names.push(plan.intern(FILE_ROW_NUMBER));
values.push(plan.add_constant(Value::Boolean(true)));
let named = plan.add_name_list(&names);
let given = plan.add_expr_list(&values);
match plan.node_mut(scan) {
Node::TableFunction { columns, options, settings, .. } => {
*columns = widened;
*options = named;
*settings = given;
}
_ => return None,
}
Some(ColumnBinding::new(index, at))
}
fn carry(plan: &mut Plan, node: NodeRef, below: ColumnBinding) -> ColumnBinding {
let Node::Project { index, exprs, names, .. } = *plan.node(node) else { return below };
let mut held = plan.expr_list(exprs).to_vec();
let mut labels = plan.name_list(names).to_vec();
let at = u32::try_from(held.len()).unwrap_or(u32::MAX);
held.push(plan.add_expr(Expr::Column(below), LogicalType::BigInt));
labels.push(plan.intern(FILE_ROW_NUMBER));
let widened = plan.add_expr_list(&held);
let renamed = plan.add_name_list(&labels);
match plan.node_mut(node) {
Node::Project { exprs, names, .. } => {
*exprs = widened;
*names = renamed;
}
_ => return below,
}
ColumnBinding::new(index, at)
}
fn fresh(plan: &Plan) -> u32 {
let mut next = 0;
for at in 0..plan.node_count() {
let node = plan.node(u32::try_from(at).unwrap_or(u32::MAX));
if let Some(index) = node.table_index() {
next = next.max(index + 1);
}
}
next
}
#[cfg(test)]
mod tests {
use rudb_plan::Plan;
use super::defer;
fn wide(columns: &[&str], keys: &str, extra: &str) -> String {
let schema: Vec<String> = columns.iter().map(|name| format!("{name}::INTEGER")).collect();
let project: Vec<String> = columns
.iter()
.enumerate()
.map(|(at, name)| format!("#1.{at}::INTEGER AS {name}"))
.collect();
let above: Vec<String> = columns
.iter()
.enumerate()
.map(|(at, name)| format!("#3.{at}::INTEGER AS {name}"))
.collect();
format!(
"Project #9 [{}]\n TopN 10 offset 0 [{keys}]\n Project #3 [{}]\n{extra} \
TableFunction read_parquet args=['hits.parquet'::VARCHAR] #1 [{}]\n",
above.join(", "),
project.join(", "),
schema.join(", ")
)
}
fn deferred(text: &str) -> String {
let mut plan =
Plan::parse(text).unwrap_or_else(|error| panic!("{text} did not parse: {error}"));
defer(&mut plan);
plan.validate().unwrap_or_else(|error| panic!("{text} did not stay valid: {error}"));
plan.to_string()
}
const TEN: [&str; 10] = ["a", "b", "c", "d", "e", "f", "g", "h", "i", "j"];
#[test]
fn a_top_n_over_a_wide_scan_keeps_the_ordering_column_and_fetches_the_rest() {
let out = deferred(&wide(&TEN, "#3.1::INTEGER ASC NULLS LAST", ""));
assert!(out.contains("Fetch args=['hits.parquet'::VARCHAR]"), "{out}");
assert!(out.contains("file_row_number=TRUE"), "{out}");
assert!(out.contains("#1 [b::INTEGER, file_row_number::BIGINT]"), "{out}");
}
#[test]
fn the_answer_still_has_the_columns_the_query_asked_for() {
let out = deferred(&wide(&TEN, "#3.0::INTEGER ASC NULLS LAST", ""));
let fetched = out.lines().find(|line| line.contains("Fetch")).unwrap_or_default();
for name in TEN {
assert!(fetched.contains(&format!("{name}::INTEGER")), "{name} missing from {out}");
}
}
#[test]
fn a_filter_between_the_scan_and_the_top_n_comes_along() {
let text = wide(&TEN, "#3.0::INTEGER ASC NULLS LAST", "").replace(
" TableFunction",
" Filter (#1.9::INTEGER = 5::INTEGER)::BOOLEAN\n TableFunction",
);
let out = deferred(&text);
assert!(out.contains("Fetch args="), "{out}");
assert!(out.contains("#1 [a::INTEGER, j::INTEGER, file_row_number::BIGINT]"), "{out}");
}
#[test]
fn computed_file_columns_are_replayed_after_the_fetch() {
let text = wide(&TEN, "#3.1::BIGINT ASC NULLS LAST", "")
.replace("#1.1::INTEGER AS b", "CAST(#1.1::INTEGER)::BIGINT AS b");
let out = deferred(&text);
assert!(out.contains("Fetch args="), "{out}");
assert!(out.contains("CAST(#"), "{out}");
assert!(out.contains("[b::INTEGER, file_row_number::BIGINT]"), "{out}");
assert!(out.lines().next().is_some_and(|line| line.starts_with("Project #9")), "{out}");
}
#[test]
fn a_scan_that_is_not_much_wider_than_the_ordering_is_left_alone() {
let text = wide(&["a", "b", "c"], "#3.0::INTEGER ASC NULLS LAST", "");
assert!(!deferred(&text).contains("Fetch"), "{text}");
}
#[test]
fn a_limit_too_large_to_be_worth_fetching_for_is_left_alone() {
let text = wide(&TEN, "#3.0::INTEGER ASC NULLS LAST", "").replace("TopN 10", "TopN 100000");
assert!(!deferred(&text).contains("Fetch"), "{text}");
}
#[test]
fn an_ordering_that_is_not_a_bare_column_is_left_alone() {
let text = wide(&TEN, "CAST(#3.0::INTEGER)::BIGINT ASC NULLS LAST", "");
assert!(!deferred(&text).contains("Fetch"), "{text}");
}
#[test]
fn a_top_n_over_a_base_table_is_left_alone_because_a_table_has_no_ordinals() {
let text = wide(&TEN, "#3.0::INTEGER ASC NULLS LAST", "").replace(
"TableFunction read_parquet args=['hits.parquet'::VARCHAR] #1",
"Get memory.main.t AS t #1",
);
assert!(!deferred(&text).contains("Fetch"), "{text}");
}
#[test]
fn running_it_twice_is_running_it_once() {
let once = deferred(&wide(&TEN, "#3.0::INTEGER ASC NULLS LAST", ""));
assert_eq!(deferred(&once), once);
}
}