use super::ast::{SelectStatement, SimpleWhenBranch, SqlExpression, WhenBranch};
pub fn map_children(
expr: SqlExpression,
mut f: impl FnMut(SqlExpression) -> SqlExpression,
) -> SqlExpression {
map_children_crossing(expr, &mut f, |f, e| f(e), |_, stmt| stmt)
}
pub fn map_children_crossing<C>(
expr: SqlExpression,
ctx: &mut C,
mut f: impl FnMut(&mut C, SqlExpression) -> SqlExpression,
mut f_stmt: impl FnMut(&mut C, Box<SelectStatement>) -> Box<SelectStatement>,
) -> SqlExpression {
match expr {
e @ (SqlExpression::Column(_)
| SqlExpression::StringLiteral(_)
| SqlExpression::NumberLiteral(_)
| SqlExpression::BooleanLiteral(_)
| SqlExpression::Null
| SqlExpression::DateTimeConstructor { .. }
| SqlExpression::DateTimeToday { .. }) => e,
SqlExpression::ScalarSubquery { query } => SqlExpression::ScalarSubquery {
query: f_stmt(ctx, query),
},
SqlExpression::MethodCall {
object,
method,
args,
} => SqlExpression::MethodCall {
object,
method,
args: args.into_iter().map(|e| f(&mut *ctx, e)).collect(),
},
SqlExpression::ChainedMethodCall { base, method, args } => {
SqlExpression::ChainedMethodCall {
base: Box::new(f(&mut *ctx, *base)),
method,
args: args.into_iter().map(|e| f(&mut *ctx, e)).collect(),
}
}
SqlExpression::FunctionCall {
name,
args,
distinct,
} => SqlExpression::FunctionCall {
name,
args: args.into_iter().map(|e| f(&mut *ctx, e)).collect(),
distinct,
},
SqlExpression::WindowFunction {
name,
args,
mut window_spec,
} => {
let args = args.into_iter().map(|e| f(&mut *ctx, e)).collect();
for item in &mut window_spec.order_by {
let taken = std::mem::replace(&mut item.expr, SqlExpression::Null);
item.expr = f(&mut *ctx, taken);
}
SqlExpression::WindowFunction {
name,
args,
window_spec,
}
}
SqlExpression::BinaryOp { left, op, right } => SqlExpression::BinaryOp {
left: Box::new(f(&mut *ctx, *left)),
op,
right: Box::new(f(&mut *ctx, *right)),
},
SqlExpression::InList { expr, values } => SqlExpression::InList {
expr: Box::new(f(&mut *ctx, *expr)),
values: values.into_iter().map(|e| f(&mut *ctx, e)).collect(),
},
SqlExpression::NotInList { expr, values } => SqlExpression::NotInList {
expr: Box::new(f(&mut *ctx, *expr)),
values: values.into_iter().map(|e| f(&mut *ctx, e)).collect(),
},
SqlExpression::Between { expr, lower, upper } => SqlExpression::Between {
expr: Box::new(f(&mut *ctx, *expr)),
lower: Box::new(f(&mut *ctx, *lower)),
upper: Box::new(f(&mut *ctx, *upper)),
},
SqlExpression::Not { expr } => SqlExpression::Not {
expr: Box::new(f(&mut *ctx, *expr)),
},
SqlExpression::CaseExpression {
when_branches,
else_branch,
} => SqlExpression::CaseExpression {
when_branches: when_branches
.into_iter()
.map(|b| WhenBranch {
condition: Box::new(f(&mut *ctx, *b.condition)),
result: Box::new(f(&mut *ctx, *b.result)),
})
.collect(),
else_branch: else_branch.map(|e| Box::new(f(&mut *ctx, *e))),
},
SqlExpression::SimpleCaseExpression {
expr,
when_branches,
else_branch,
} => SqlExpression::SimpleCaseExpression {
expr: Box::new(f(&mut *ctx, *expr)),
when_branches: when_branches
.into_iter()
.map(|b| SimpleWhenBranch {
value: Box::new(f(&mut *ctx, *b.value)),
result: Box::new(f(&mut *ctx, *b.result)),
})
.collect(),
else_branch: else_branch.map(|e| Box::new(f(&mut *ctx, *e))),
},
SqlExpression::Unnest { column, delimiter } => SqlExpression::Unnest {
column: Box::new(f(&mut *ctx, *column)),
delimiter,
},
SqlExpression::InSubquery { expr, subquery } => SqlExpression::InSubquery {
expr: Box::new(f(&mut *ctx, *expr)),
subquery: f_stmt(ctx, subquery),
},
SqlExpression::NotInSubquery { expr, subquery } => SqlExpression::NotInSubquery {
expr: Box::new(f(&mut *ctx, *expr)),
subquery: f_stmt(ctx, subquery),
},
SqlExpression::InSubqueryTuple { exprs, subquery } => SqlExpression::InSubqueryTuple {
exprs: exprs.into_iter().map(|e| f(&mut *ctx, e)).collect(),
subquery: f_stmt(ctx, subquery),
},
SqlExpression::NotInSubqueryTuple { exprs, subquery } => {
SqlExpression::NotInSubqueryTuple {
exprs: exprs.into_iter().map(|e| f(&mut *ctx, e)).collect(),
subquery: f_stmt(ctx, subquery),
}
}
}
}
pub fn visit_children<'a>(expr: &'a SqlExpression, mut f: impl FnMut(&'a SqlExpression)) {
visit_children_crossing(expr, &mut f, |f, e| f(e), |_, _| {});
}
pub fn visit_children_crossing<'a, C>(
expr: &'a SqlExpression,
ctx: &mut C,
mut f: impl FnMut(&mut C, &'a SqlExpression),
mut f_stmt: impl FnMut(&mut C, &'a SelectStatement),
) {
match expr {
SqlExpression::Column(_)
| SqlExpression::StringLiteral(_)
| SqlExpression::NumberLiteral(_)
| SqlExpression::BooleanLiteral(_)
| SqlExpression::Null
| SqlExpression::DateTimeConstructor { .. }
| SqlExpression::DateTimeToday { .. } => {}
SqlExpression::ScalarSubquery { query } => f_stmt(ctx, query),
SqlExpression::MethodCall { args, .. } | SqlExpression::FunctionCall { args, .. } => {
args.iter().for_each(|e| f(&mut *ctx, e));
}
SqlExpression::ChainedMethodCall { base, args, .. } => {
f(&mut *ctx, base);
args.iter().for_each(|e| f(&mut *ctx, e));
}
SqlExpression::WindowFunction {
args, window_spec, ..
} => {
args.iter().for_each(|e| f(&mut *ctx, e));
window_spec
.order_by
.iter()
.for_each(|item| f(&mut *ctx, &item.expr));
}
SqlExpression::BinaryOp { left, right, .. } => {
f(&mut *ctx, left);
f(&mut *ctx, right);
}
SqlExpression::InList { expr, values } | SqlExpression::NotInList { expr, values } => {
f(&mut *ctx, expr);
values.iter().for_each(|e| f(&mut *ctx, e));
}
SqlExpression::Between { expr, lower, upper } => {
f(&mut *ctx, expr);
f(&mut *ctx, lower);
f(&mut *ctx, upper);
}
SqlExpression::Not { expr } | SqlExpression::Unnest { column: expr, .. } => {
f(&mut *ctx, expr)
}
SqlExpression::CaseExpression {
when_branches,
else_branch,
} => {
for branch in when_branches {
f(&mut *ctx, &branch.condition);
f(&mut *ctx, &branch.result);
}
if let Some(e) = else_branch {
f(&mut *ctx, e);
}
}
SqlExpression::SimpleCaseExpression {
expr,
when_branches,
else_branch,
} => {
f(&mut *ctx, expr);
for branch in when_branches {
f(&mut *ctx, &branch.value);
f(&mut *ctx, &branch.result);
}
if let Some(e) = else_branch {
f(&mut *ctx, e);
}
}
SqlExpression::InSubquery { expr, subquery }
| SqlExpression::NotInSubquery { expr, subquery } => {
f(&mut *ctx, expr);
f_stmt(ctx, subquery);
}
SqlExpression::InSubqueryTuple { exprs, subquery }
| SqlExpression::NotInSubqueryTuple { exprs, subquery } => {
exprs.iter().for_each(|e| f(&mut *ctx, e));
f_stmt(ctx, subquery);
}
}
}
pub fn visit_all<'a>(expr: &'a SqlExpression, f: &mut impl FnMut(&'a SqlExpression)) {
f(expr);
visit_children(expr, |child| visit_all(child, f));
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sql::parser::ast::ColumnRef;
use crate::sql::recursive_parser::Parser;
fn expr_of(select_expr: &str) -> SqlExpression {
let sql = format!("SELECT {select_expr} FROM t");
let stmt = Parser::new(&sql)
.parse()
.unwrap_or_else(|e| panic!("{sql} should parse: {e}"));
stmt.select_items
.into_iter()
.find_map(|item| match item {
crate::sql::parser::ast::SelectItem::Expression { expr, .. } => Some(expr),
_ => None,
})
.expect("expected a projected expression")
}
fn columns(expr: &SqlExpression) -> Vec<String> {
let mut found = Vec::new();
visit_all(expr, &mut |e| {
if let SqlExpression::Column(c) = e {
found.push(c.name.clone());
}
});
found
}
fn rename_all(expr: SqlExpression, to: &str) -> SqlExpression {
match expr {
SqlExpression::Column(_) => SqlExpression::Column(ColumnRef::unquoted(to.to_string())),
other => map_children(other, |e| rename_all(e, to)),
}
}
#[test]
fn visit_all_reaches_case_branches() {
let expr = expr_of("CASE WHEN a > 1 THEN b ELSE c END");
let mut found = columns(&expr);
found.sort();
assert_eq!(found, vec!["a", "b", "c"]);
}
#[test]
fn visit_all_reaches_nested_function_args() {
let expr = expr_of("UPPER(TRIM(name))");
assert_eq!(columns(&expr), vec!["name"]);
}
#[test]
fn visit_all_reaches_between_operands() {
let expr = expr_of("x BETWEEN lo AND hi");
assert_eq!(columns(&expr), vec!["x", "lo", "hi"]);
}
#[test]
fn visit_all_reaches_window_order_by() {
let expr = expr_of("ROW_NUMBER() OVER (ORDER BY created_at)");
assert!(
columns(&expr).contains(&"created_at".to_string()),
"window ORDER BY expressions must be reachable"
);
}
#[test]
fn walkers_do_not_cross_into_subqueries() {
let expr = expr_of("(SELECT MAX(inner_col) FROM other)");
assert!(
matches!(expr, SqlExpression::ScalarSubquery { .. }),
"expected a scalar subquery"
);
assert!(
columns(&expr).is_empty(),
"must not descend into a subquery's own scope"
);
let stmt = Parser::new("SELECT a FROM t WHERE outer_col IN (SELECT x FROM other)")
.parse()
.expect("should parse");
let cond = &stmt.where_clause.expect("where clause").conditions[0].expr;
assert_eq!(columns(cond), vec!["outer_col"]);
}
#[test]
fn crossing_walkers_do_reach_into_subqueries() {
let expr = expr_of("(SELECT MAX(inner_col) FROM other)");
let mut statements_seen = 0;
visit_children_crossing(&expr, &mut statements_seen, |_, _| {}, |n, _| *n += 1);
assert_eq!(
statements_seen, 1,
"a scalar subquery's statement must be reachable when crossing"
);
let rewritten = map_children_crossing(
expr,
&mut (),
|_, e| e,
|_, mut stmt| {
stmt.limit = Some(1);
stmt
},
);
match rewritten {
SqlExpression::ScalarSubquery { query } => assert_eq!(query.limit, Some(1)),
other => panic!("expected a scalar subquery, got {other:?}"),
}
}
#[test]
fn crossing_reaches_tuple_subquery_operands_and_statement() {
let stmt = Parser::new("SELECT a FROM t WHERE (a, b) IN (SELECT x, y FROM u)")
.parse()
.expect("should parse");
let cond = &stmt.where_clause.expect("where clause").conditions[0].expr;
let mut ctx = (Vec::<String>::new(), 0);
visit_children_crossing(
cond,
&mut ctx,
|ctx, e| {
if let SqlExpression::Column(c) = e {
ctx.0.push(c.name.clone());
}
},
|ctx, _| ctx.1 += 1,
);
assert_eq!(ctx.0, vec!["a", "b"], "same-scope operands must be visited");
assert_eq!(ctx.1, 1, "the subquery statement must be visited too");
}
#[test]
fn map_children_rewrites_nested_expressions() {
let expr = expr_of("CASE WHEN a > 1 THEN UPPER(b) ELSE c END");
let renamed = rename_all(expr, "z");
assert_eq!(columns(&renamed), vec!["z", "z", "z"]);
}
#[test]
fn map_children_rewrites_window_order_by() {
let expr = expr_of("ROW_NUMBER() OVER (ORDER BY created_at)");
let renamed = rename_all(expr, "z");
assert_eq!(columns(&renamed), vec!["z"]);
}
#[test]
fn map_children_leaves_leaves_alone() {
let expr = expr_of("42");
let mapped = map_children(expr, |_| panic!("a literal has no children"));
assert!(matches!(mapped, SqlExpression::NumberLiteral(ref n) if n == "42"));
}
}