use crate::ir::*;
use crate::types::*;
#[derive(Debug)]
pub enum OptimizationRule {
PushDownPredicates,
JoinOrderOptimization,
EliminateUnnecessaryProjections,
ConstantFolding,
IndexSelection,
}
#[derive(Debug, Clone)]
pub struct CostEstimate {
pub cardinality: f64,
pub cost: f64,
pub selectivity: f64,
}
#[derive(Debug)]
struct JoinPlan {
relations: Vec<LogicalOp>,
cost: f64,
cardinality: f64,
}
#[derive(Debug)]
pub struct QueryOptimizer {
rules: Vec<OptimizationRule>,
statistics: CostStatistics,
}
#[derive(Debug)]
pub struct CostStatistics {
default_selectivity: f64,
index_selectivity: f64,
join_cost_factor: f64,
}
impl QueryOptimizer {
pub fn new() -> Self {
Self {
rules: vec![
OptimizationRule::PushDownPredicates,
OptimizationRule::JoinOrderOptimization,
OptimizationRule::EliminateUnnecessaryProjections,
OptimizationRule::ConstantFolding,
OptimizationRule::IndexSelection,
],
statistics: CostStatistics {
default_selectivity: 0.1, index_selectivity: 0.01, join_cost_factor: 0.5, },
}
}
pub fn optimize(&self, plan: &PlanIR, catalog: &Catalog) -> PlanIR {
let mut optimized = plan.clone();
for rule in &self.rules {
optimized = self.apply_rule(optimized, rule, catalog);
}
optimized
}
fn apply_rule(&self, plan: PlanIR, rule: &OptimizationRule, catalog: &Catalog) -> PlanIR {
match rule {
OptimizationRule::PushDownPredicates => {
self.push_down_predicates(plan)
}
OptimizationRule::JoinOrderOptimization => {
self.optimize_join_order(plan, catalog)
}
OptimizationRule::EliminateUnnecessaryProjections => {
self.eliminate_unnecessary_projections(plan)
}
OptimizationRule::ConstantFolding => {
self.constant_folding(plan)
}
OptimizationRule::IndexSelection => {
self.select_indexes(plan, catalog)
}
}
}
fn push_down_predicates(&self, plan: PlanIR) -> PlanIR {
let optimized_plan = self.push_down_predicates_op(&plan.plan);
PlanIR {
plan: optimized_plan,
limit: plan.limit,
}
}
fn push_down_predicates_op(&self, op: &LogicalOp) -> LogicalOp {
match op {
LogicalOp::Filter { pred, input } => {
match input.as_ref() {
LogicalOp::Join { left, right, on } => {
let (left_pred, right_pred, remaining_pred) =
self.split_predicate_for_join(pred, on);
let new_left = if let Some(lp) = left_pred {
Box::new(LogicalOp::Filter {
pred: lp,
input: left.clone(),
})
} else {
left.clone()
};
let new_right = if let Some(rp) = right_pred {
Box::new(LogicalOp::Filter {
pred: rp,
input: right.clone(),
})
} else {
right.clone()
};
if let Some(rem_pred) = remaining_pred {
LogicalOp::Filter {
pred: rem_pred,
input: Box::new(LogicalOp::Join {
left: new_left,
right: new_right,
on: on.clone(),
}),
}
} else {
LogicalOp::Join {
left: new_left,
right: new_right,
on: on.clone(),
}
}
}
_ => {
LogicalOp::Filter {
pred: pred.clone(),
input: Box::new(self.push_down_predicates_op(input)),
}
}
}
}
LogicalOp::Join { left, right, on } => {
LogicalOp::Join {
left: Box::new(self.push_down_predicates_op(left)),
right: Box::new(self.push_down_predicates_op(right)),
on: on.clone(),
}
}
LogicalOp::Project { cols, input } => {
LogicalOp::Project {
cols: cols.clone(),
input: Box::new(self.push_down_predicates_op(input)),
}
}
_ => op.clone(),
}
}
fn split_predicate_for_join(&self, pred: &Predicate, join_keys: &[String])
-> (Option<Predicate>, Option<Predicate>, Option<Predicate>) {
match pred {
Predicate::And { and } => {
let mut left_preds = Vec::new();
let mut right_preds = Vec::new();
let mut remaining = Vec::new();
for p in and {
let (l, r, rem) = self.split_predicate_for_join(p, join_keys);
if let Some(lp) = l { left_preds.push(lp); }
if let Some(rp) = r { right_preds.push(rp); }
if let Some(rem_p) = rem { remaining.push(rem_p); }
}
let left = if left_preds.is_empty() {
None
} else if left_preds.len() == 1 {
Some(left_preds.into_iter().next().unwrap())
} else {
Some(Predicate::And { and: left_preds })
};
let right = if right_preds.is_empty() {
None
} else if right_preds.len() == 1 {
Some(right_preds.into_iter().next().unwrap())
} else {
Some(Predicate::And { and: right_preds })
};
let rem = if remaining.is_empty() {
None
} else if remaining.len() == 1 {
Some(remaining.into_iter().next().unwrap())
} else {
Some(Predicate::And { and: remaining })
};
(left, right, rem)
}
Predicate::Eq { eq } if eq.len() == 2 => {
let left_vars = self.extract_variables(&eq[0]);
let right_vars = self.extract_variables(&eq[1]);
if self.contains_join_key(&left_vars, join_keys) &&
self.contains_join_key(&right_vars, join_keys) {
(None, None, Some(pred.clone()))
} else if self.contains_join_key(&left_vars, join_keys) {
(Some(pred.clone()), None, None)
} else if self.contains_join_key(&right_vars, join_keys) {
(None, Some(pred.clone()), None)
} else {
(None, None, Some(pred.clone()))
}
}
_ => (None, None, Some(pred.clone())),
}
}
fn extract_variables(&self, expr: &Expr) -> Vec<String> {
match expr {
Expr::Var(v) => vec![v.clone()],
Expr::Fn { args, .. } => args.iter()
.flat_map(|arg| self.extract_variables(arg))
.collect(),
_ => Vec::new(),
}
}
fn contains_join_key(&self, vars: &[String], join_keys: &[String]) -> bool {
vars.iter().any(|v| join_keys.contains(v))
}
fn optimize_join_order(&self, plan: PlanIR, catalog: &Catalog) -> PlanIR {
let optimized_plan = self.optimize_join_order_op(&plan.plan, catalog);
PlanIR {
plan: optimized_plan,
limit: plan.limit,
}
}
fn optimize_join_order_op(&self, op: &LogicalOp, catalog: &Catalog) -> LogicalOp {
match op {
LogicalOp::Join { left, right, on } => {
let mut relations = Vec::new();
self.collect_relations(left, &mut relations);
self.collect_relations(right, &mut relations);
if relations.len() > 2 {
self.optimize_join_order_dp(&relations, on, catalog)
} else {
let left_cost = self.estimate_cost_detailed(left, catalog);
let right_cost = self.estimate_cost_detailed(right, catalog);
if left_cost.cost > right_cost.cost {
LogicalOp::Join {
left: Box::new(self.optimize_join_order_op(right, catalog)),
right: Box::new(self.optimize_join_order_op(left, catalog)),
on: on.clone(),
}
} else {
LogicalOp::Join {
left: Box::new(self.optimize_join_order_op(left, catalog)),
right: Box::new(self.optimize_join_order_op(right, catalog)),
on: on.clone(),
}
}
}
}
_ => op.clone(),
}
}
fn collect_relations(&self, op: &LogicalOp, relations: &mut Vec<LogicalOp>) {
match op {
LogicalOp::Join { left, right, .. } => {
self.collect_relations(left, relations);
self.collect_relations(right, relations);
}
_ => relations.push(op.clone()),
}
}
fn optimize_join_order_dp(&self, relations: &[LogicalOp], join_keys: &[String], catalog: &Catalog) -> LogicalOp {
let n = relations.len();
let mut dp = vec![vec![None; n]; n];
let mut costs = vec![vec![f64::INFINITY; n]; n];
let mut cardinalities = vec![vec![0.0; n]; n];
for i in 0..n {
let cost_est = self.estimate_cost_detailed(&relations[i], catalog);
costs[i][i] = cost_est.cost;
cardinalities[i][i] = cost_est.cardinality;
dp[i][i] = Some(relations[i].clone());
}
for len in 2..=n {
for i in 0..=n-len {
let j = i + len - 1;
costs[i][j] = f64::INFINITY;
for k in i..j {
let left_cost = costs[i][k];
let right_cost = costs[k+1][j];
let left_card = cardinalities[i][k];
let right_card = cardinalities[k+1][j];
let join_cost = self.calculate_join_cost(left_card, right_card, join_keys);
let total_cost = left_cost + right_cost + join_cost;
if total_cost < costs[i][j] {
costs[i][j] = total_cost;
cardinalities[i][j] = self.estimate_join_cardinality(left_card, right_card, join_keys);
if let (Some(left_plan), Some(right_plan)) = (&dp[i][k], &dp[k+1][j]) {
dp[i][j] = Some(LogicalOp::Join {
left: Box::new(left_plan.clone()),
right: Box::new(right_plan.clone()),
on: join_keys.to_vec(),
});
}
}
}
}
}
dp[0][n-1].clone().unwrap_or_else(|| relations[0].clone())
}
fn estimate_cost_detailed(&self, op: &LogicalOp, catalog: &Catalog) -> CostEstimate {
match op {
LogicalOp::NodeScan { label, props, .. } => {
let base_cardinality = catalog.get_label(label)
.map(|_| 1000.0) .unwrap_or(100.0);
let selectivity = if props.is_some() {
self.statistics.index_selectivity
} else {
1.0
};
CostEstimate {
cardinality: base_cardinality * selectivity,
cost: base_cardinality * selectivity * 10.0, selectivity,
}
}
LogicalOp::IndexScan { .. } => {
CostEstimate {
cardinality: 10.0, cost: 5.0, selectivity: self.statistics.index_selectivity,
}
}
LogicalOp::Filter { input, pred } => {
let input_cost = self.estimate_cost_detailed(input, catalog);
let filter_selectivity = self.estimate_filter_selectivity(pred);
CostEstimate {
cardinality: input_cost.cardinality * filter_selectivity,
cost: input_cost.cost + (input_cost.cardinality * filter_selectivity * 2.0), selectivity: input_cost.selectivity * filter_selectivity,
}
}
LogicalOp::Join { left, right, on } => {
let left_cost = self.estimate_cost_detailed(left, catalog);
let right_cost = self.estimate_cost_detailed(right, catalog);
let join_card = self.estimate_join_cardinality(left_cost.cardinality, right_cost.cardinality, on);
let join_cost = self.calculate_join_cost(left_cost.cardinality, right_cost.cardinality, on);
CostEstimate {
cardinality: join_card,
cost: left_cost.cost + right_cost.cost + join_cost,
selectivity: (left_cost.selectivity + right_cost.selectivity) / 2.0,
}
}
_ => CostEstimate {
cardinality: 100.0,
cost: 10.0,
selectivity: self.statistics.default_selectivity,
},
}
}
fn estimate_filter_selectivity(&self, pred: &Predicate) -> f64 {
match pred {
Predicate::Eq { .. } => self.statistics.index_selectivity, Predicate::Gt { .. } | Predicate::Lt { .. } | Predicate::Ge { .. } | Predicate::Le { .. } => 0.3, Predicate::And { and } => {
and.iter().map(|p| self.estimate_filter_selectivity(p)).product()
}
Predicate::Or { or } => {
let sum: f64 = or.iter().map(|p| self.estimate_filter_selectivity(p)).sum();
sum.min(1.0)
}
_ => self.statistics.default_selectivity,
}
}
fn estimate_join_cardinality(&self, left_card: f64, right_card: f64, join_keys: &[String]) -> f64 {
if join_keys.is_empty() {
left_card * right_card
} else {
(left_card * right_card * self.statistics.default_selectivity).max(left_card.max(right_card))
}
}
fn calculate_join_cost(&self, left_card: f64, right_card: f64, _join_keys: &[String]) -> f64 {
left_card * right_card * self.statistics.join_cost_factor
}
fn estimate_cost(&self, op: &LogicalOp, catalog: &Catalog) -> f64 {
self.estimate_cost_detailed(op, catalog).cost
}
fn eliminate_unnecessary_projections(&self, plan: PlanIR) -> PlanIR {
let optimized_plan = self.eliminate_unnecessary_projections_op(&plan.plan);
PlanIR {
plan: optimized_plan,
limit: plan.limit,
}
}
fn eliminate_unnecessary_projections_op(&self, op: &LogicalOp) -> LogicalOp {
match op {
LogicalOp::Project { cols, input } => {
match input.as_ref() {
LogicalOp::Project { cols: inner_cols, input: inner_input } => {
let merged_cols = cols.iter()
.filter(|col| inner_cols.contains(col))
.cloned()
.collect();
LogicalOp::Project {
cols: merged_cols,
input: inner_input.clone(),
}
}
_ => LogicalOp::Project {
cols: cols.clone(),
input: Box::new(self.eliminate_unnecessary_projections_op(input)),
}
}
}
_ => op.clone(),
}
}
fn constant_folding(&self, plan: PlanIR) -> PlanIR {
plan
}
fn select_indexes(&self, plan: PlanIR, catalog: &Catalog) -> PlanIR {
let optimized_plan = self.select_indexes_op(&plan.plan, catalog);
PlanIR {
plan: optimized_plan,
limit: plan.limit,
}
}
fn select_indexes_op(&self, op: &LogicalOp, catalog: &Catalog) -> LogicalOp {
match op {
LogicalOp::Filter { pred, input } => {
match input.as_ref() {
LogicalOp::NodeScan { label, as_, props: _ } => {
if let Some(index) = self.find_best_index(catalog, label, pred) {
LogicalOp::Filter {
pred: pred.clone(),
input: Box::new(LogicalOp::IndexScan {
label: label.clone(),
as_: as_.clone(),
index: index.name,
value: self.extract_index_value(pred, &index.properties[0]),
}),
}
} else {
LogicalOp::Filter {
pred: pred.clone(),
input: Box::new(self.select_indexes_op(input, catalog)),
}
}
}
_ => LogicalOp::Filter {
pred: pred.clone(),
input: Box::new(self.select_indexes_op(input, catalog)),
}
}
}
_ => op.clone(),
}
}
fn find_best_index(&self, catalog: &Catalog, label: &Label, pred: &Predicate) -> Option<IndexDef> {
catalog.indexes.iter()
.filter(|idx| &idx.label == label)
.find(|idx| self.can_use_index(pred, &idx.properties[0]))
.cloned()
}
fn can_use_index(&self, pred: &Predicate, prop: &PropertyKey) -> bool {
match pred {
Predicate::Eq { eq } if eq.len() == 2 => {
let left_vars = self.extract_variables(&eq[0]);
let right_vars = self.extract_variables(&eq[1]);
left_vars.contains(prop) || right_vars.contains(prop)
}
_ => false,
}
}
fn extract_index_value(&self, pred: &Predicate, prop: &PropertyKey) -> Value {
match pred {
Predicate::Eq { eq } if eq.len() == 2 => {
if let Expr::Var(var) = &eq[0] {
if var == prop {
if let Expr::Const(val) = &eq[1] {
return val.clone();
}
}
}
if let Expr::Var(var) = &eq[1] {
if var == prop {
if let Expr::Const(val) = &eq[0] {
return val.clone();
}
}
}
}
_ => {}
}
Value::Null
}
}