use crate::bind::{
BoundExpr, BoundFrameBound, BoundOrderTerm, BoundResultColumn, BoundSelect, BoundWindow,
SourceRows,
};
use crate::dml::{BoundDelete, BoundInsert, BoundInsertSource, BoundUpdate, ColumnSource};
pub type Rewrite<'a> = &'a mut dyn FnMut(&mut BoundExpr);
pub fn rewrite_expr(expr: &mut BoundExpr, rewrite: Rewrite<'_>) {
rewrite(expr);
for child in expr.children_mut() {
rewrite_expr(child, rewrite);
}
if let Some(block) = expr.block_mut() {
rewrite_select(block, rewrite);
}
}
fn rewrite_option(expr: Option<&mut BoundExpr>, rewrite: Rewrite<'_>) {
if let Some(expr) = expr {
rewrite_expr(expr, rewrite);
}
}
fn rewrite_columns(columns: &mut [BoundResultColumn], rewrite: Rewrite<'_>) {
for column in columns {
rewrite_expr(&mut column.expr, rewrite);
}
}
fn rewrite_order(terms: &mut [BoundOrderTerm], rewrite: Rewrite<'_>) {
for term in terms {
rewrite_expr(&mut term.expr, rewrite);
}
}
fn rewrite_bound(bound: &mut BoundFrameBound, rewrite: Rewrite<'_>) {
match bound {
BoundFrameBound::Preceding(expr) | BoundFrameBound::Following(expr) => {
rewrite_expr(expr, rewrite)
}
BoundFrameBound::UnboundedPreceding
| BoundFrameBound::CurrentRow
| BoundFrameBound::UnboundedFollowing => {}
}
}
fn rewrite_window(window: &mut BoundWindow, rewrite: Rewrite<'_>) {
for argument in &mut window.arguments {
rewrite_expr(argument, rewrite);
}
rewrite_option(window.filter.as_mut(), rewrite);
for term in &mut window.partition_by {
rewrite_expr(term, rewrite);
}
rewrite_order(&mut window.order_by, rewrite);
rewrite_bound(&mut window.start, rewrite);
rewrite_bound(&mut window.end, rewrite);
}
pub fn rewrite_select(select: &mut BoundSelect, rewrite: Rewrite<'_>) {
for source in &mut select.sources {
rewrite_option(source.constraint.as_mut(), rewrite);
match &mut source.rows {
SourceRows::Table | SourceRows::RecursiveSelf { .. } => {}
SourceRows::Subquery(block) => rewrite_select(block, rewrite),
SourceRows::Recursive(body) => {
for (_, arm) in body.seeds.iter_mut().chain(body.steps.iter_mut()) {
rewrite_select(arm, rewrite);
}
}
}
}
rewrite_option(select.filter.as_mut(), rewrite);
for term in &mut select.group_by {
rewrite_expr(term, rewrite);
}
rewrite_option(select.having.as_mut(), rewrite);
rewrite_columns(&mut select.columns, rewrite);
rewrite_order(&mut select.order_by, rewrite);
rewrite_option(select.limit.as_mut(), rewrite);
rewrite_option(select.offset.as_mut(), rewrite);
for aggregate in &mut select.aggregates {
for argument in &mut aggregate.arguments {
rewrite_expr(argument, rewrite);
}
}
for row in &mut select.values {
for value in row {
rewrite_expr(value, rewrite);
}
}
for (_, arm) in &mut select.compounds {
rewrite_select(arm, rewrite);
}
for window in &mut select.windows {
rewrite_window(window, rewrite);
}
}
fn rewrite_source(source: &mut ColumnSource, rewrite: Rewrite<'_>) {
match source {
ColumnSource::Row(_) => {}
ColumnSource::Expr(expr) | ColumnSource::Generated(expr) => rewrite_expr(expr, rewrite),
}
}
pub fn rewrite_insert(statement: &mut BoundInsert, rewrite: Rewrite<'_>) {
for column in &mut statement.columns {
rewrite_source(column, rewrite);
}
if let Some(rowid) = statement.rowid.as_mut() {
rewrite_source(rowid, rewrite);
}
match &mut statement.source {
BoundInsertSource::Values(rows) => {
for row in rows {
for value in row {
rewrite_expr(value, rewrite);
}
}
}
BoundInsertSource::Select(select) => rewrite_select(select, rewrite),
}
for check in &mut statement.checks {
rewrite_expr(&mut check.expr, rewrite);
}
for upsert in &mut statement.upsert {
for assignment in &mut upsert.assignments {
rewrite_expr(&mut assignment.value, rewrite);
}
rewrite_option(upsert.filter.as_mut(), rewrite);
}
rewrite_columns(&mut statement.returning, rewrite);
}
pub fn rewrite_update(statement: &mut BoundUpdate, rewrite: Rewrite<'_>) {
for assignment in &mut statement.assignments {
rewrite_expr(&mut assignment.value, rewrite);
}
rewrite_option(statement.filter.as_mut(), rewrite);
for check in &mut statement.checks {
rewrite_expr(&mut check.expr, rewrite);
}
rewrite_columns(&mut statement.returning, rewrite);
rewrite_order(&mut statement.order_by, rewrite);
rewrite_option(statement.limit.as_mut(), rewrite);
rewrite_option(statement.offset.as_mut(), rewrite);
if let Some(rows) = statement.view_rows.as_mut() {
rewrite_select(rows, rewrite);
}
}
pub fn rewrite_delete(statement: &mut BoundDelete, rewrite: Rewrite<'_>) {
rewrite_option(statement.filter.as_mut(), rewrite);
rewrite_columns(&mut statement.returning, rewrite);
rewrite_order(&mut statement.order_by, rewrite);
rewrite_option(statement.limit.as_mut(), rewrite);
rewrite_option(statement.offset.as_mut(), rewrite);
if let Some(rows) = statement.view_rows.as_mut() {
rewrite_select(rows, rewrite);
}
}