use std::collections::HashSet;
use egg::{Analysis, EGraph, Id, Subst};
use crate::{
ast::{BinOp, Expr},
eval::Context,
value::Value,
};
use super::{EPlanNode, IndexRef, Null, TableRef};
fn created_references_inner(
set: &mut HashSet<String>,
egraph: &EGraph<EPlanNode, ConstantFolding>,
id: Id,
) {
egraph[id].nodes.iter().for_each(|n| match n {
EPlanNode::SeqScan(id) => {
created_references_inner(set, egraph, *id);
}
EPlanNode::TableScanEq(items)
| EPlanNode::TableScanGtEq(items)
| EPlanNode::TableScanGt(items)
| EPlanNode::TableScanLtEq(items)
| EPlanNode::TableScanLt(items)
| EPlanNode::Select(items) => {
items
.iter()
.for_each(|id| created_references_inner(set, egraph, *id));
}
EPlanNode::FilterScan([a, b])
| EPlanNode::Delete([a, b])
| EPlanNode::Update([a, b])
| EPlanNode::Join([a, b])
| EPlanNode::And([a, b])
| EPlanNode::Or([a, b])
| EPlanNode::Eq([a, b])
| EPlanNode::NEq([a, b])
| EPlanNode::GtEq([a, b])
| EPlanNode::Gt([a, b])
| EPlanNode::LtEq([a, b])
| EPlanNode::Lt([a, b])
| EPlanNode::Add([a, b])
| EPlanNode::Sub([a, b])
| EPlanNode::Mul([a, b])
| EPlanNode::Div([a, b])
| EPlanNode::Mod([a, b]) => {
created_references_inner(set, egraph, *a);
created_references_inner(set, egraph, *b);
}
EPlanNode::Not([a]) => {
created_references_inner(set, egraph, *a);
}
EPlanNode::TableRef(TableRef(_, table)) | EPlanNode::IndexRef(IndexRef(_, table, _)) => {
set.insert(table.to_string());
}
EPlanNode::NoneScan
| EPlanNode::Null(_)
| EPlanNode::NamedColumn(_)
| EPlanNode::Binding(_)
| EPlanNode::Int(_)
| EPlanNode::Bool(_)
| EPlanNode::String(_) => {}
});
}
fn created_references(egraph: &EGraph<EPlanNode, ConstantFolding>, id: Id) -> HashSet<String> {
let mut set = HashSet::new();
created_references_inner(&mut set, egraph, id);
set
}
fn references_tables_inner(
set: &mut HashSet<String>,
egraph: &EGraph<EPlanNode, ConstantFolding>,
id: Id,
) {
egraph[id].nodes.iter().for_each(|n| match n {
EPlanNode::SeqScan(id) => {
references_tables_inner(set, egraph, *id);
}
EPlanNode::NamedColumn(column_ref) => {
set.insert(column_ref.0.clone());
}
EPlanNode::NoneScan
| EPlanNode::Int(_)
| EPlanNode::Bool(_)
| EPlanNode::String(_)
| EPlanNode::Binding(_)
| EPlanNode::Null(_) => {}
EPlanNode::TableScanEq(items)
| EPlanNode::TableScanGt(items)
| EPlanNode::TableScanGtEq(items)
| EPlanNode::TableScanLt(items)
| EPlanNode::TableScanLtEq(items)
| EPlanNode::Select(items) => {
items
.iter()
.for_each(|id| created_references_inner(set, egraph, *id));
}
EPlanNode::FilterScan([a, b])
| EPlanNode::Delete([a, b])
| EPlanNode::Update([a, b])
| EPlanNode::Join([a, b])
| EPlanNode::And([a, b])
| EPlanNode::Or([a, b])
| EPlanNode::Eq([a, b])
| EPlanNode::NEq([a, b])
| EPlanNode::Gt([a, b])
| EPlanNode::GtEq([a, b])
| EPlanNode::Lt([a, b])
| EPlanNode::LtEq([a, b])
| EPlanNode::Add([a, b])
| EPlanNode::Sub([a, b])
| EPlanNode::Mul([a, b])
| EPlanNode::Div([a, b])
| EPlanNode::Mod([a, b]) => {
references_tables_inner(set, egraph, *a);
references_tables_inner(set, egraph, *b);
}
EPlanNode::Not([a]) => {
created_references_inner(set, egraph, *a);
}
EPlanNode::TableRef(TableRef(_, table)) | EPlanNode::IndexRef(IndexRef(_, table, _)) => {
set.insert(table.to_string());
}
});
}
fn references_tables(egraph: &EGraph<EPlanNode, ConstantFolding>, id: Id) -> HashSet<String> {
let mut set = HashSet::new();
references_tables_inner(&mut set, egraph, id);
set
}
pub fn dont_reference_each_other(
var1: &'static str,
var2: &'static str,
) -> impl Fn(&mut EGraph<EPlanNode, ConstantFolding>, Id, &Subst) -> bool {
let var1 = var1.parse().unwrap();
let var2 = var2.parse().unwrap();
move |egraph, _, subst| {
let lhs_created = created_references(egraph, subst[var1]);
let rhs_created = created_references(egraph, subst[var2]);
let lhs_referenced = references_tables(egraph, subst[var1]);
let rhs_referenced = references_tables(egraph, subst[var2]);
lhs_referenced.is_disjoint(&rhs_created) && rhs_referenced.is_disjoint(&lhs_created)
}
}
#[derive(Default)]
pub struct ConstantFolding;
impl Analysis<EPlanNode> for ConstantFolding {
type Data = Option<Value>;
fn make(egraph: &egg::EGraph<EPlanNode, Self>, enode: &EPlanNode) -> Self::Data {
let x = |i: &Id| egraph[*i].data.clone();
match enode {
EPlanNode::Null(_) => Some(Value::Null),
EPlanNode::Int(i) => Some(Value::Int(*i)),
EPlanNode::Bool(b) => Some(Value::Bool(*b)),
EPlanNode::String(s) => Some(Value::String(s.to_string())),
EPlanNode::Eq([a, b]) => Some(Value::Bool(x(a)? == x(b)?)),
EPlanNode::NEq([a, b]) => Some(Value::Bool(x(a)? != x(b)?)),
EPlanNode::Gt([a, b]) => Some(Value::Bool(x(a)? > x(b)?)),
EPlanNode::GtEq([a, b]) => Some(Value::Bool(x(a)? >= x(b)?)),
EPlanNode::Lt([a, b]) => Some(Value::Bool(x(a)? < x(b)?)),
EPlanNode::LtEq([a, b]) => Some(Value::Bool(x(a)? <= x(b)?)),
EPlanNode::Not([a]) => match x(a)? {
Value::Bool(a) => Some(Value::Bool(!a)),
_ => None,
},
EPlanNode::And([a, b]) => match (x(a)?, x(b)?) {
(Value::Bool(a), Value::Bool(b)) => Some(Value::Bool(a && b)),
_ => None,
},
EPlanNode::Or([a, b]) => match (x(a)?, x(b)?) {
(Value::Bool(a), Value::Bool(b)) => Some(Value::Bool(a || b)),
_ => None,
},
EPlanNode::Add([a, b]) => Context::new(Vec::new())
.eval(&Expr::Bin(
Box::new(Expr::Literal(x(a)?)),
BinOp::Add,
Box::new(Expr::Literal(x(b)?)),
))
.ok(),
EPlanNode::Sub([a, b]) => Context::new(Vec::new())
.eval(&Expr::Bin(
Box::new(Expr::Literal(x(a)?)),
BinOp::Sub,
Box::new(Expr::Literal(x(b)?)),
))
.ok(),
EPlanNode::Mul([a, b]) => Context::new(Vec::new())
.eval(&Expr::Bin(
Box::new(Expr::Literal(x(a)?)),
BinOp::Mul,
Box::new(Expr::Literal(x(b)?)),
))
.ok(),
EPlanNode::Div([a, b]) => Context::new(Vec::new())
.eval(&Expr::Bin(
Box::new(Expr::Literal(x(a)?)),
BinOp::Div,
Box::new(Expr::Literal(x(b)?)),
))
.ok(),
EPlanNode::Mod([a, b]) => Context::new(Vec::new())
.eval(&Expr::Bin(
Box::new(Expr::Literal(x(a)?)),
BinOp::Mod,
Box::new(Expr::Literal(x(b)?)),
))
.ok(),
EPlanNode::SeqScan(_)
| EPlanNode::NoneScan
| EPlanNode::FilterScan(_)
| EPlanNode::TableScanEq(_)
| EPlanNode::TableScanGt(_)
| EPlanNode::TableScanGtEq(_)
| EPlanNode::TableScanLt(_)
| EPlanNode::TableScanLtEq(_)
| EPlanNode::Select(_)
| EPlanNode::Delete(_)
| EPlanNode::Update(_)
| EPlanNode::Join(_)
| EPlanNode::NamedColumn(_)
| EPlanNode::Binding(_)
| EPlanNode::TableRef(_)
| EPlanNode::IndexRef(_) => None,
}
}
fn merge(&mut self, a: &mut Self::Data, b: Self::Data) -> egg::DidMerge {
egg::merge_max(a, b)
}
fn modify(egraph: &mut egg::EGraph<EPlanNode, Self>, id: Id) {
if let Some(v) = &egraph[id].data {
let added = match v {
Value::Null => egraph.add(EPlanNode::Null(Null)),
Value::Bool(b) => egraph.add(EPlanNode::Bool(*b)),
Value::Int(i) => egraph.add(EPlanNode::Int(*i)),
Value::String(s) => egraph.add(EPlanNode::String(s.to_string())),
};
egraph.union(id, added);
}
}
}