pub mod const_fold;
pub mod distributive_or;
pub mod join_filter_or;
pub mod like;
pub mod unnest_conjunction;
use const_fold::ConstFold;
use distributive_or::DistributiveOrRewrite;
use glaredb_error::Result;
use join_filter_or::JoinFilterOrRewrite;
use like::LikeRewrite;
use unnest_conjunction::UnnestConjunctionRewrite;
use super::OptimizeRule;
use crate::expr::{self, Expression};
use crate::logical::binder::bind_context::BindContext;
use crate::logical::operator::{LogicalNode, LogicalOperator};
pub trait ExpressionRewriteRule {
fn rewrite(expression: Expression) -> Result<Expression>;
}
#[derive(Debug)]
pub struct ExpressionRewriter;
impl OptimizeRule for ExpressionRewriter {
fn optimize(
&mut self,
_bind_context: &mut BindContext,
plan: LogicalOperator,
) -> Result<LogicalOperator> {
let mut plan = match plan {
LogicalOperator::Project(mut project) => {
project.node.projections = Self::apply_rewrites_all(project.node.projections)?;
LogicalOperator::Project(project)
}
LogicalOperator::Filter(mut filter) => {
filter.node.filter = Self::apply_rewrites(filter.node.filter)?;
filter.node.filter = JoinFilterOrRewrite::rewrite(filter.node.filter)?; LogicalOperator::Filter(filter)
}
LogicalOperator::ArbitraryJoin(mut join) => {
join.node.condition = Self::apply_rewrites(join.node.condition)?;
join.node.condition = JoinFilterOrRewrite::rewrite(join.node.condition)?; LogicalOperator::ArbitraryJoin(join)
}
mut other => {
other.for_each_expr_mut(|expr| {
let mut orig = std::mem::replace(expr, expr::lit(83).into());
orig = Self::apply_rewrites(orig)?;
*expr = orig;
Ok(())
})?;
other
}
};
plan.modify_replace_children(&mut |child| self.optimize(_bind_context, child))?;
Ok(plan)
}
}
impl ExpressionRewriter {
pub fn apply_rewrites_all(exprs: Vec<Expression>) -> Result<Vec<Expression>> {
exprs
.into_iter()
.map(Self::apply_rewrites)
.collect::<Result<Vec<_>>>()
}
pub fn apply_rewrites(expr: Expression) -> Result<Expression> {
let expr = LikeRewrite::rewrite(expr)?; let expr = ConstFold::rewrite(expr)?;
let expr = UnnestConjunctionRewrite::rewrite(expr)?;
let expr = DistributiveOrRewrite::rewrite(expr)?;
Ok(expr)
}
}