use std::{collections::BTreeMap, iter::once, str::FromStr};
use anyhow::Result;
use egg::{rewrite, ENodeOrVar, Id, Pattern, RecExpr, Rewrite, Var};
use crate::ast::{BinOp, Column, Expr};
use crate::Db;
use super::TableRef;
use super::{super::plan::compile::compile_expr, IndexRef};
use super::{
analysis::{dont_reference_each_other, ConstantFolding},
EPlanNode,
};
pub fn rewrite_rules(
db: &mut Db,
tx_id: u64,
tables: &BTreeMap<String, String>,
) -> Result<Vec<Rewrite<EPlanNode, ConstantFolding>>> {
let mut rules = vec![
rewrite!("filter-false"; "(filter FALSE ?a)" => "nonescan"),
rewrite!("filter-true"; "(filter TRUE ?a)" => "?a"),
rewrite!("eq-commutate"; "(= ?a ?b)" => "(= ?b ?a)"),
rewrite!("and-true"; "(and TRUE ?a)" => "?a"),
rewrite!("and-false"; "(and FALSE ?b)" => "FALSE"),
rewrite!("and-commutate"; "(and ?a ?b)" => "(and ?b ?a)"),
rewrite!("filter-commutate"; "(filter ?a (filter ?b ?c))" => "(filter ?b (filter ?a ?c))"),
rewrite!("split-filter"; "(filter (and ?a ?b) ?c)" => "(filter ?a (filter ?b ?c))"),
rewrite!("carry-over-add-to-sub"; "(= (+ ?a ?b) ?c)" => "(= ?a (- ?c ?b))"),
rewrite!("carry-over-sub-to-add"; "(= (- ?a ?b) ?c)" => "(= ?a (+ ?c ?b))"),
rewrite!("collapse-not-eq-to-neq"; "(not (= ?a ?b))" => "(!= ?a ?b)"),
rewrite!("collapse-not-gt-to-lte"; "(not (> ?a ?b))" => "(<= ?a ?b)"),
rewrite!("collapse-not-ge-to-lt"; "(not (>= ?a ?b))" => "(< ?a ?b)"),
rewrite!("collapse-not-lt-to-gte"; "(not (< ?a ?b))" => "(>= ?a ?b)"),
rewrite!("collapse-not-le-to-gt"; "(not (<= ?a ?b))" => "(> ?a ?b)"),
rewrite!("transpose-filter-into-join-right"; "(filter ?a (join ?b ?c))" => "(join ?b (filter ?a ?c))"),
rewrite!("swap-join"; "(join ?a ?b)" => "(join ?b ?a)" if dont_reference_each_other("?a", "?b")),
];
for (a, t) in tables {
let header = db.get_table_schema(tx_id, &t)?;
for (op, node, scan) in [
(
BinOp::Eq,
EPlanNode::Eq as fn([Id; 2]) -> EPlanNode,
EPlanNode::TableScanEq as fn(Vec<Id>) -> EPlanNode,
),
(
BinOp::GtEq,
EPlanNode::GtEq as fn([Id; 2]) -> EPlanNode,
EPlanNode::TableScanGtEq as fn(Vec<Id>) -> EPlanNode,
),
(
BinOp::Gt,
EPlanNode::Gt as fn([Id; 2]) -> EPlanNode,
EPlanNode::TableScanGt as fn(Vec<Id>) -> EPlanNode,
),
(
BinOp::LtEq,
EPlanNode::LtEq as fn([Id; 2]) -> EPlanNode,
EPlanNode::TableScanLtEq as fn(Vec<Id>) -> EPlanNode,
),
(
BinOp::Lt,
EPlanNode::Lt as fn([Id; 2]) -> EPlanNode,
EPlanNode::TableScanLt as fn(Vec<Id>) -> EPlanNode,
),
] {
rules.append(&mut index_rule(
a,
a,
tx_id,
op,
&header.primary_key,
node,
scan,
|table_name| EPlanNode::TableRef(TableRef(tx_id, table_name)),
));
for index in &header.indexes {
rules.append(&mut index_rule(
&index.name(&header.name),
a,
tx_id,
op,
&index.exprs,
node,
scan,
|index_name| EPlanNode::IndexRef(IndexRef(tx_id, a.to_string(), index_name)),
));
}
}
}
Ok(rules)
}
fn index_rule<F: Fn(String) -> EPlanNode>(
name_in_ctx: &String,
table_name: &String,
tx_id: u64,
op: BinOp,
exprs: &[Expr],
node: fn([Id; 2]) -> EPlanNode,
scan: fn(Vec<Id>) -> EPlanNode,
r: F,
) -> Vec<Rewrite<EPlanNode, ConstantFolding>> {
let mut v = Vec::new();
for i in 1..exprs.len() + 1 {
let expr_slice = &exprs[..i];
let expr_tail = expr_slice.iter().rev().skip(1).rev();
v.push(
Rewrite::new(
format!("filter-to-index-{table_name}-{name_in_ctx}-{i}-{op}"),
{
let mut expr: RecExpr<ENodeOrVar<EPlanNode>> = RecExpr::default();
let table = expr.add(ENodeOrVar::ENode(EPlanNode::TableRef(TableRef(
tx_id,
table_name.to_string(),
))));
let seq_scan = expr.add(ENodeOrVar::ENode(EPlanNode::SeqScan(table)));
let tail = expr_tail
.clone()
.enumerate()
.map(|(n, e)| {
let lhs = compile_expr(
&mut expr,
&rename_expr(&e, table_name),
ENodeOrVar::ENode,
);
let rhs =
expr.add(ENodeOrVar::Var(Var::from_str(&format!("?{n}")).unwrap()));
let cond = expr.add(ENodeOrVar::ENode(EPlanNode::Eq([lhs, rhs])));
cond
})
.collect::<Vec<_>>()
.into_iter()
.reduce(|a, b| expr.add(ENodeOrVar::ENode(EPlanNode::And([a, b]))));
let head = compile_expr(
&mut expr,
&rename_expr(&expr_slice.last().unwrap(), table_name),
ENodeOrVar::ENode,
);
let rhs = expr.add(ENodeOrVar::Var(Var::from_str("?t").unwrap()));
let cond = expr.add(ENodeOrVar::ENode(node([head, rhs])));
let whole = if let Some(t) = tail {
expr.add(ENodeOrVar::ENode(EPlanNode::And([t, cond])))
} else {
cond
};
let _filter =
expr.add(ENodeOrVar::ENode(EPlanNode::FilterScan([whole, seq_scan])));
Pattern::new(expr)
},
{
let mut expr: RecExpr<ENodeOrVar<EPlanNode>> = RecExpr::default();
let table = expr.add(ENodeOrVar::ENode(r(name_in_ctx.to_string())));
let items = once(table)
.chain(
expr_tail
.enumerate()
.map(|(n, _)| {
expr.add(ENodeOrVar::Var(
Var::from_str(&format!("?{n}")).unwrap(),
))
})
.collect::<Vec<_>>(),
)
.chain(once(
expr.add(ENodeOrVar::Var(Var::from_str("?t").unwrap())),
))
.collect();
let _primary_key_scan = expr.add(ENodeOrVar::ENode(scan(items)));
Pattern::new(expr)
},
)
.unwrap(),
);
}
v
}
fn rename_expr(e: &Expr, name: &str) -> Expr {
match e {
Expr::Bin(expr, op, expr1) => Expr::Bin(
Box::new(rename_expr(expr, name)),
*op,
Box::new(rename_expr(expr1, name)),
),
Expr::Column(Column(_l, r)) => Expr::Column(Column(name.to_string(), r.to_string())),
e => e.clone(),
}
}