matdb 0.1.0

An experimental embedded SQL-like DBMS
Documentation
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)"),
        // TODO: union when !=
        // FIXME: this causes stack overflow when (filter true (join ?b ?c)) or such :)
        rewrite!("transpose-filter-into-join-right"; "(filter ?a (join ?b ?c))" => "(join ?b (filter ?a ?c))"),
        // rewrite!("transpose-filter-into-join-left"; "(filter ?a (join ?b ?c))" => "(join (filter ?a ?b) ?c)" if references_only_subset("?a", "?b")),
        rewrite!("swap-join"; "(join ?a ?b)" => "(join ?b ?a)" if dont_reference_each_other("?a", "?b")),
        // TODO: figure out if this works once you implement proper cost estimation
        //  FIXME: it doesn't, it stack overflows lmao
        // rewrite!("gt-to-gteq-and-gt";"(> ?a ?b)" => "(and (>= ?a ?b) (> ?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)));

                    // the left tail of the expression list
                    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]))));

                    // the last expression
                    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(),
    }
}