1#![allow(clippy::doc_markdown)]
4
5use alloc::collections::BTreeMap;
44use alloc::string::String;
45use alloc::vec::Vec;
46
47use spg_sql::ast::{ColumnName, Expr, FromClause, FromJoin, JoinKind, SelectStatement, TableRef};
48
49use crate::selectivity;
50use crate::statistics::Statistics;
51use spg_storage::Catalog;
52
53pub const FULL_ENUM_MAX: usize = 4;
57
58pub fn choose_order_for_test(
67 stmt: &SelectStatement,
68 catalog: &Catalog,
69 stats: &Statistics,
70) -> Option<Vec<usize>> {
71 let mut clone = stmt.clone();
72 choose_order_inner(&mut clone, catalog, stats)
73}
74
75fn choose_order_inner(
76 stmt: &mut SelectStatement,
77 catalog: &Catalog,
78 stats: &Statistics,
79) -> Option<Vec<usize>> {
80 let from = stmt.from.as_mut()?;
81 if from.joins.is_empty() {
82 return None;
83 }
84 if from
85 .joins
86 .iter()
87 .any(|j| !matches!(j.kind, JoinKind::Inner))
88 {
89 return None;
90 }
91 let mut tables: Vec<TableRef> = Vec::with_capacity(1 + from.joins.len());
92 tables.push(from.primary.clone());
93 for j in &from.joins {
94 tables.push(j.table.clone());
95 }
96 let n = tables.len();
97 let mut alias_to_idx: BTreeMap<String, usize> = BTreeMap::new();
98 for (i, t) in tables.iter().enumerate() {
99 let key = t.alias.clone().unwrap_or_else(|| t.name.clone());
100 alias_to_idx.insert(key, i);
101 if t.alias.is_some() {
102 alias_to_idx.entry(t.name.clone()).or_insert(i);
103 }
104 }
105 let mut edges: Vec<Edge> = Vec::new();
106 for j in &from.joins {
107 let on = j.on.as_ref()?;
108 for sub in split_and_conjunctions(on) {
109 let mut endpoint_set: Vec<usize> = Vec::new();
110 if !collect_referenced_tables(sub, &alias_to_idx, &mut endpoint_set) {
111 return None;
112 }
113 endpoint_set.sort_unstable();
114 endpoint_set.dedup();
115 edges.push(Edge {
116 endpoints: endpoint_set,
117 predicate: sub.clone(),
118 selectivity: estimate_edge_selectivity(sub, &tables, catalog, stats),
119 });
120 }
121 }
122 let mut sizes: Vec<u64> = Vec::with_capacity(n);
123 for t in &tables {
124 let table = catalog.get(&t.name)?;
125 sizes.push(table.rows().len() as u64);
126 }
127 Some(if n <= FULL_ENUM_MAX {
128 best_order_brute(n, &sizes, &edges)
129 } else {
130 best_order_greedy(n, &sizes, &edges)
131 })
132}
133
134pub fn reorder_joins(stmt: &mut SelectStatement, catalog: &Catalog, stats: &Statistics) {
135 reorder_joins_with(stmt, catalog, stats, false);
136}
137
138pub fn reorder_joins_with(
144 stmt: &mut SelectStatement,
145 catalog: &Catalog,
146 stats: &Statistics,
147 plan_deterministic: bool,
148) {
149 REORDER_INNER_RUN_TRIED.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
156 if plan_deterministic {
157 return;
158 }
159 let Some(from) = stmt.from.as_mut() else {
160 return;
161 };
162 if from.joins.is_empty() {
163 return;
164 }
165 let split = from
175 .joins
176 .iter()
177 .position(|j| !matches!(j.kind, JoinKind::Inner))
178 .unwrap_or(from.joins.len());
179 if split == 0 {
180 return;
184 }
185 if stats.is_empty() {
193 return;
194 }
195 let mut tables: Vec<TableRef> = Vec::with_capacity(1 + split);
199 tables.push(from.primary.clone());
200 for j in &from.joins[..split] {
201 tables.push(j.table.clone());
202 }
203 let n = tables.len();
204 let mut alias_to_idx: BTreeMap<String, usize> = BTreeMap::new();
207 for (i, t) in tables.iter().enumerate() {
208 let key = t.alias.clone().unwrap_or_else(|| t.name.clone());
209 alias_to_idx.insert(key, i);
210 if t.alias.is_some() {
214 alias_to_idx.entry(t.name.clone()).or_insert(i);
215 }
216 }
217 let mut edges: Vec<Edge> = Vec::new();
219 for j in &from.joins[..split] {
220 let Some(on) = j.on.as_ref() else {
221 return;
224 };
225 for sub in split_and_conjunctions(on) {
226 let mut endpoint_set: Vec<usize> = Vec::new();
227 if !collect_referenced_tables(sub, &alias_to_idx, &mut endpoint_set) {
228 return;
229 }
230 endpoint_set.sort_unstable();
231 endpoint_set.dedup();
232 edges.push(Edge {
233 endpoints: endpoint_set,
234 predicate: sub.clone(),
235 selectivity: estimate_edge_selectivity(sub, &tables, catalog, stats),
236 });
237 }
238 }
239 let mut sizes: Vec<u64> = Vec::with_capacity(n);
243 for t in &tables {
244 let Some(table) = catalog.get(&t.name) else {
245 return;
246 };
247 sizes.push(table.rows().len() as u64);
248 }
249 let order: Vec<usize> = if n <= FULL_ENUM_MAX {
251 best_order_brute(n, &sizes, &edges)
252 } else {
253 best_order_greedy(n, &sizes, &edges)
254 };
255 if order.iter().enumerate().all(|(i, &j)| i == j) {
257 return;
258 }
259 REORDER_INNER_RUN_FIRED.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
263 rewrite_from_with_trailing(from, &tables, &edges, &order, split);
264}
265
266pub static REORDER_INNER_RUN_TRIED: core::sync::atomic::AtomicU64 =
270 core::sync::atomic::AtomicU64::new(0);
271pub static REORDER_INNER_RUN_FIRED: core::sync::atomic::AtomicU64 =
272 core::sync::atomic::AtomicU64::new(0);
273
274struct Edge {
275 endpoints: Vec<usize>,
277 predicate: Expr,
280 selectivity: f64,
284}
285
286pub(crate) fn split_and_conjunctions(expr: &Expr) -> Vec<&Expr> {
291 use spg_sql::ast::BinOp;
292 let mut out: Vec<&Expr> = Vec::new();
293 let mut stack: Vec<&Expr> = alloc::vec![expr];
294 while let Some(e) = stack.pop() {
295 if let Expr::Binary {
296 op: BinOp::And,
297 lhs,
298 rhs,
299 } = e
300 {
301 stack.push(rhs);
302 stack.push(lhs);
303 } else {
304 out.push(e);
305 }
306 }
307 out
308}
309
310fn collect_referenced_tables(
311 expr: &Expr,
312 alias_to_idx: &BTreeMap<String, usize>,
313 out: &mut Vec<usize>,
314) -> bool {
315 match expr {
316 Expr::Column(ColumnName {
317 qualifier: Some(q), ..
318 }) => {
319 if let Some(&i) = alias_to_idx.get(q) {
320 out.push(i);
321 true
322 } else {
323 false
324 }
325 }
326 Expr::Column(_) => {
327 false
330 }
331 Expr::Literal(_) | Expr::Placeholder(_) => true,
332 Expr::Binary { lhs, rhs, .. } => {
333 collect_referenced_tables(lhs, alias_to_idx, out)
334 && collect_referenced_tables(rhs, alias_to_idx, out)
335 }
336 Expr::Unary { expr, .. } => collect_referenced_tables(expr, alias_to_idx, out),
337 Expr::FunctionCall { args, .. } => args
338 .iter()
339 .all(|a| collect_referenced_tables(a, alias_to_idx, out)),
340 Expr::Cast { expr, .. } | Expr::IsNull { expr, .. } => {
341 collect_referenced_tables(expr, alias_to_idx, out)
342 }
343 Expr::Like {
344 expr: e, pattern, ..
345 } => {
346 collect_referenced_tables(e, alias_to_idx, out)
347 && collect_referenced_tables(pattern, alias_to_idx, out)
348 }
349 _ => false,
351 }
352}
353
354fn estimate_edge_selectivity(
359 on: &Expr,
360 tables: &[TableRef],
361 catalog: &Catalog,
362 stats: &Statistics,
363) -> f64 {
364 use spg_sql::ast::BinOp;
365 let Expr::Binary {
366 op: BinOp::Eq,
367 lhs,
368 rhs,
369 } = on
370 else {
371 return selectivity::DEFAULT_RANGE;
372 };
373 let lhs_col = column_ref(lhs);
374 let rhs_col = column_ref(rhs);
375 let (Some(lhs_col), Some(rhs_col)) = (lhs_col, rhs_col) else {
376 return selectivity::DEFAULT_RANGE;
377 };
378 let lhs_distinct = column_n_distinct(&lhs_col, tables, catalog, stats);
379 let rhs_distinct = column_n_distinct(&rhs_col, tables, catalog, stats);
380 let max_distinct = lhs_distinct.max(rhs_distinct).max(1);
381 1.0 / max_distinct as f64
382}
383
384fn column_ref(expr: &Expr) -> Option<(Option<String>, String)> {
385 if let Expr::Column(ColumnName { qualifier, name }) = expr {
386 Some((qualifier.clone(), name.clone()))
387 } else {
388 None
389 }
390}
391
392fn column_n_distinct(
393 col: &(Option<String>, String),
394 tables: &[TableRef],
395 catalog: &Catalog,
396 stats: &Statistics,
397) -> u64 {
398 let Some(alias) = col.0.as_ref() else {
399 return 0;
400 };
401 let Some(table_name) = tables
402 .iter()
403 .find(|t| t.alias.as_deref() == Some(alias.as_str()) || t.name == *alias)
404 .map(|t| t.name.clone())
405 else {
406 return 0;
407 };
408 if let Some(s) = stats.get(&table_name, &col.1) {
409 return s.n_distinct.max(1);
410 }
411 catalog
412 .get(&table_name)
413 .map_or(1, |t| (t.rows().len() as u64).max(1))
414}
415
416fn best_order_brute(n: usize, sizes: &[u64], edges: &[Edge]) -> Vec<usize> {
417 let mut indices: Vec<usize> = (0..n).collect();
418 let mut best_cost = f64::INFINITY;
419 let mut best_order = indices.clone();
420 permute(&mut indices, 0, &mut |perm| {
421 let c = plan_cost(perm, sizes, edges);
422 if c < best_cost {
423 best_cost = c;
424 best_order = perm.to_vec();
425 }
426 });
427 best_order
428}
429
430fn permute<F: FnMut(&[usize])>(arr: &mut Vec<usize>, k: usize, visit: &mut F) {
431 if k >= arr.len() {
432 visit(arr);
433 return;
434 }
435 for i in k..arr.len() {
436 arr.swap(i, k);
437 permute(arr, k + 1, visit);
438 arr.swap(i, k);
439 }
440}
441
442fn best_order_greedy(n: usize, sizes: &[u64], edges: &[Edge]) -> Vec<usize> {
443 let mut chosen: Vec<usize> = Vec::with_capacity(n);
445 let mut remaining: Vec<usize> = (0..n).collect();
446 let &first = remaining.iter().min_by_key(|&&i| sizes[i]).expect("n > 0");
447 chosen.push(first);
448 remaining.retain(|&x| x != first);
449 while !remaining.is_empty() {
450 let mut best_cand = remaining[0];
454 let mut best_cost = f64::INFINITY;
455 for &cand in &remaining {
456 let mut probe = chosen.clone();
457 probe.push(cand);
458 let c = plan_cost(&probe, sizes, edges);
459 if c < best_cost {
460 best_cost = c;
461 best_cand = cand;
462 }
463 }
464 chosen.push(best_cand);
465 remaining.retain(|&x| x != best_cand);
466 }
467 chosen
468}
469
470fn plan_cost(order: &[usize], sizes: &[u64], edges: &[Edge]) -> f64 {
474 let mut running = sizes[order[0]] as f64;
476 let mut cost = 0.0_f64;
477 let mut in_prefix: Vec<bool> = alloc::vec![false; sizes.len()];
478 in_prefix[order[0]] = true;
479 for &table_idx in &order[1..] {
480 let right = sizes[table_idx] as f64;
481 cost += running * right;
484 in_prefix[table_idx] = true;
485 let mut step_output = running * right;
486 for edge in edges {
489 if edge.endpoints.iter().all(|&e| in_prefix[e]) {
490 if edge.endpoints.contains(&table_idx) {
495 step_output *= edge.selectivity;
496 }
497 }
498 }
499 running = step_output.max(1.0);
500 }
501 cost
502}
503
504fn rewrite_from(from: &mut FromClause, tables: &[TableRef], edges: &[Edge], order: &[usize]) {
505 rewrite_from_with_trailing(from, tables, edges, order, from.joins.len());
506}
507
508fn rewrite_from_with_trailing(
515 from: &mut FromClause,
516 tables: &[TableRef],
517 edges: &[Edge],
518 order: &[usize],
519 split: usize,
520) {
521 let trailing: alloc::vec::Vec<FromJoin> = from.joins[split..].to_vec();
522 from.primary = tables[order[0]].clone();
523 from.joins.clear();
524 let mut in_prefix: Vec<bool> = alloc::vec![false; tables.len()];
525 in_prefix[order[0]] = true;
526 let mut edges_used: Vec<bool> = alloc::vec![false; edges.len()];
527 for &table_idx in &order[1..] {
528 in_prefix[table_idx] = true;
529 let mut combined: Option<Expr> = None;
533 for (ei, edge) in edges.iter().enumerate() {
534 if edges_used[ei] {
535 continue;
536 }
537 if edge.endpoints.contains(&table_idx) && edge.endpoints.iter().all(|&e| in_prefix[e]) {
538 edges_used[ei] = true;
539 combined = Some(match combined {
540 None => edge.predicate.clone(),
541 Some(prev) => Expr::Binary {
542 op: spg_sql::ast::BinOp::And,
543 lhs: alloc::boxed::Box::new(prev),
544 rhs: alloc::boxed::Box::new(edge.predicate.clone()),
545 },
546 });
547 }
548 }
549 let on = combined.unwrap_or_else(|| Expr::Literal(spg_sql::ast::Literal::Bool(true)));
554 from.joins.push(FromJoin {
555 kind: JoinKind::Inner,
556 table: tables[table_idx].clone(),
557 on: Some(on),
558 using_cols: None,
559 natural: false,
560 });
561 }
562 from.joins.extend(trailing);
564}
565
566pub(crate) fn drive_from(stmt: &mut SelectStatement, driver_alias: &str) -> bool {
585 let Some(from) = stmt.from.as_mut() else {
586 return false;
587 };
588 let mut tables: Vec<TableRef> = Vec::with_capacity(1 + from.joins.len());
592 tables.push(from.primary.clone());
593 for j in &from.joins {
594 tables.push(j.table.clone());
595 }
596 let mut alias_to_idx: BTreeMap<String, usize> = BTreeMap::new();
597 for (i, t) in tables.iter().enumerate() {
598 let key = t.alias.clone().unwrap_or_else(|| t.name.clone());
599 alias_to_idx.insert(key, i);
600 if t.alias.is_some() {
601 alias_to_idx.entry(t.name.clone()).or_insert(i);
602 }
603 }
604 let Some(&driver_idx) = alias_to_idx.get(driver_alias) else {
605 return false;
606 };
607 if from.joins.is_empty() {
608 return driver_idx == 0;
611 }
612 if driver_idx == 0 {
613 return true; }
615 if from
616 .joins
617 .iter()
618 .any(|j| !matches!(j.kind, JoinKind::Inner))
619 {
620 return false; }
622 let mut edges: Vec<Edge> = Vec::new();
626 for j in &from.joins {
627 let Some(on) = j.on.as_ref() else {
628 return false;
629 };
630 for sub in split_and_conjunctions(on) {
631 let mut endpoint_set: Vec<usize> = Vec::new();
632 if !collect_referenced_tables(sub, &alias_to_idx, &mut endpoint_set) {
633 return false;
634 }
635 endpoint_set.sort_unstable();
636 endpoint_set.dedup();
637 edges.push(Edge {
638 endpoints: endpoint_set,
639 predicate: sub.clone(),
640 selectivity: 0.0,
641 });
642 }
643 }
644 let n = tables.len();
649 let mut order: Vec<usize> = alloc::vec![driver_idx];
650 let mut included: Vec<bool> = alloc::vec![false; n];
651 included[driver_idx] = true;
652 loop {
653 let mut progressed = false;
654 for ti in 0..n {
655 if included[ti] {
656 continue;
657 }
658 let connects = edges.iter().any(|e| {
659 e.endpoints.contains(&ti) && e.endpoints.iter().all(|&x| x == ti || included[x])
660 });
661 if connects {
662 order.push(ti);
663 included[ti] = true;
664 progressed = true;
665 }
666 }
667 if !progressed {
668 break;
669 }
670 }
671 for ti in 0..n {
672 if !included[ti] {
673 order.push(ti);
674 }
675 }
676 rewrite_from(from, &tables, &edges, &order);
677 true
678}
679
680#[cfg(test)]
681mod tests {
682 use super::*;
683 use spg_sql::parser;
684
685 #[test]
686 fn no_joins_is_noop() {
687 let mut stmt = match parser::parse_statement("SELECT * FROM users").unwrap() {
688 spg_sql::ast::Statement::Select(s) => s,
689 _ => panic!(),
690 };
691 let cat = Catalog::new();
692 let stats = Statistics::new();
693 let snap = stmt.clone();
694 reorder_joins(&mut stmt, &cat, &stats);
695 assert_eq!(stmt, snap);
696 }
697
698 #[test]
699 fn five_table_star_picks_fact_first() {
700 let mut e = crate::Engine::new();
704 e.execute("CREATE TABLE fact (id INT NOT NULL, k1 INT NOT NULL, k2 INT NOT NULL, k3 INT NOT NULL, k4 INT NOT NULL)").unwrap();
705 for tag in ["big1", "big2", "big3", "big4"] {
706 e.execute(&alloc::format!("CREATE TABLE {tag} (k INT NOT NULL)"))
707 .unwrap();
708 }
709 for i in 0..3 {
710 e.execute(&alloc::format!(
711 "INSERT INTO fact VALUES ({i}, {i}, {i}, {i}, {i})"
712 ))
713 .unwrap();
714 }
715 for tag in ["big1", "big2", "big3", "big4"] {
716 for i in 0..40 {
717 e.execute(&alloc::format!("INSERT INTO {tag} VALUES ({i})"))
718 .unwrap();
719 }
720 }
721 e.execute("ANALYZE").unwrap();
722 let stmt = e.prepare(
723 "SELECT fact.id FROM big1 \
724 INNER JOIN big2 ON 1 = 1 \
725 INNER JOIN big3 ON 1 = 1 \
726 INNER JOIN big4 ON 1 = 1 \
727 INNER JOIN fact ON fact.k1 = big1.k AND fact.k2 = big2.k AND fact.k3 = big3.k AND fact.k4 = big4.k",
728 )
729 .unwrap();
730 let spg_sql::ast::Statement::Select(sel) = stmt else {
731 panic!()
732 };
733 let from = sel.from.unwrap();
734 assert_eq!(
735 from.primary.name, "fact",
736 "reorder must put fact first; got primary={:?}",
737 from.primary.name
738 );
739 }
740
741 #[test]
742 fn left_join_is_skipped() {
743 let mut stmt = match parser::parse_statement(
745 "SELECT * FROM a LEFT JOIN b ON a.id = b.id LEFT JOIN c ON b.id = c.id",
746 )
747 .unwrap()
748 {
749 spg_sql::ast::Statement::Select(s) => s,
750 _ => panic!(),
751 };
752 let cat = Catalog::new();
753 let stats = Statistics::new();
754 let snap = stmt.clone();
755 reorder_joins(&mut stmt, &cat, &stats);
756 assert_eq!(stmt, snap);
757 }
758
759 fn parse_select(sql: &str) -> SelectStatement {
760 match parser::parse_statement(sql).unwrap() {
761 spg_sql::ast::Statement::Select(s) => s,
762 _ => panic!(),
763 }
764 }
765
766 #[test]
767 fn drive_from_promotes_named_table_to_primary() {
768 let mut s = parse_select(
771 "SELECT e2.category FROM email_analysis e2 \
772 INNER JOIN messages m2 ON e2.message_id = m2.id \
773 WHERE m2.thread_id = 'th-5'",
774 );
775 assert!(drive_from(&mut s, "m2"));
776 let from = s.from.as_ref().unwrap();
777 assert_eq!(from.primary.alias.as_deref(), Some("m2"));
778 assert_eq!(from.primary.name, "messages");
779 assert_eq!(from.joins.len(), 1);
780 assert_eq!(from.joins[0].table.alias.as_deref(), Some("e2"));
781 assert!(from.joins[0].on.is_some());
783 }
784
785 #[test]
786 fn drive_from_noop_when_already_primary() {
787 let mut s = parse_select(
788 "SELECT m2.id FROM messages m2 INNER JOIN email_analysis e2 ON e2.message_id = m2.id",
789 );
790 let snap = s.clone();
791 assert!(drive_from(&mut s, "m2"));
792 assert_eq!(s, snap, "already-driving promotion must not mutate");
793 }
794
795 #[test]
796 fn drive_from_refuses_left_join() {
797 let mut s = parse_select(
799 "SELECT e2.id FROM email_analysis e2 LEFT JOIN messages m2 ON e2.message_id = m2.id",
800 );
801 let snap = s.clone();
802 assert!(!drive_from(&mut s, "m2"));
803 assert_eq!(s, snap, "refused promotion must not mutate");
804 }
805
806 #[test]
807 fn drive_from_unknown_alias_is_false() {
808 let mut s = parse_select(
809 "SELECT e2.id FROM email_analysis e2 INNER JOIN messages m2 ON e2.message_id = m2.id",
810 );
811 assert!(!drive_from(&mut s, "nope"));
812 }
813}