matdb 0.1.0

An experimental embedded SQL-like DBMS
Documentation
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);
        }

        // TODO: i *think* the only way there can be a table reference here is if it is created, but i'm not sure
        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);
        }
    }
}