Skip to main content

uqa_planner/
join_enumerator.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! DPccp join enumeration following Moerkotte and Neumann (2006).
8//!
9//! Enumerates connected-subgraph / complement pairs of the
10//! [`JoinGraph`] in canonical order: each connected subgraph S is
11//! formed by extending a smaller connected subgraph with an adjacent
12//! vertex whose index exceeds `min(S)`, ensuring each subgraph is
13//! emitted exactly once. Complexity is `O(3^n)` over the relation
14//! count; far below the `n!` of exhaustive enumeration. Falls back to
15//! a greedy `O(n^3)` heuristic when the graph has more than
16//! [`MAX_DP_RELATIONS`] relations.
17//!
18//! Internally relation subsets are encoded as `u64` bitmasks for
19//! O(1) hash-table lookup and set operations. Equijoins are costed as
20//! hash joins because that is the physical strategy available to the SQL
21//! execution pipeline. An index-join cost must never influence ordering unless
22//! the planner can prove that a compatible physical index join is executable.
23//!
24//! Returns a [`JoinPlan`] tree where each `Join` node records the
25//! `(left, right, edge, cost, cardinality)` tuple. Disconnected join
26//! graphs are handled by solving each connected component
27//! independently and cross-joining them in cardinality-ascending
28//! order.
29
30use std::collections::{BTreeMap, BTreeSet};
31
32use crate::cost_model::{CostEstimator, OperatorKind};
33use crate::join_graph::{JoinEdge, JoinGraph};
34
35/// Beyond this count, exact enumeration switches to the greedy fallback.
36pub const MAX_DP_RELATIONS: usize = 16;
37
38type StarLeaves = Vec<(usize, Vec<JoinEdge>)>;
39type StarShape = (usize, StarLeaves);
40
41#[derive(Clone, Copy)]
42struct StarState {
43    cardinality: f64,
44    cost: f64,
45    prev_mask: usize,
46    leaf_pos: usize,
47}
48
49/// A (sub)plan for joining a set of relations. `relations` is the bitmask
50/// of relation indices in the plan; `cardinality` and `cost` are the running
51/// estimates, and `left` / `right` / `join_edge` are populated for
52/// internal nodes.
53#[derive(Debug, Clone)]
54pub struct JoinPlan {
55    pub relations: u64,
56    pub cardinality: f64,
57    pub cost: f64,
58    pub left: Option<Box<JoinPlan>>,
59    pub right: Option<Box<JoinPlan>>,
60    pub join_edge: Option<JoinEdge>,
61    /// The join algorithm `_emit_csg_cmp_pair` picked for this node.
62    /// `None` for base relations and cross joins.
63    pub kind: Option<OperatorKind>,
64}
65
66impl JoinPlan {
67    /// Build a leaf plan for a single relation.
68    fn leaf(idx: usize, rows: f64, access_cost: f64) -> Self {
69        Self {
70            relations: 1u64 << idx,
71            cardinality: rows,
72            cost: access_cost,
73            left: None,
74            right: None,
75            join_edge: None,
76            kind: None,
77        }
78    }
79
80    /// Cardinality projected by this (sub)plan.
81    pub fn rows(&self) -> f64 {
82        self.cardinality
83    }
84
85    pub fn cost(&self) -> f64 {
86        self.cost
87    }
88}
89
90/// Run DPccp over `graph` and return the cheapest join plan over the
91/// full relation set. Returns `None` for an empty graph.
92pub fn enumerate_dpccp(graph: &JoinGraph) -> Option<JoinPlan> {
93    DPccp::new(graph).optimize()
94}
95
96/// Run DPccp with an explicit physical cost estimator.
97pub fn enumerate_dpccp_with_cost_estimator(
98    graph: &JoinGraph,
99    cost_estimator: CostEstimator,
100) -> Option<JoinPlan> {
101    DPccp::with_cost_estimator(graph, cost_estimator).optimize()
102}
103
104/// DPccp join-order optimiser. Public so callers that need the
105/// cancellation-friendly stages (`optimize`, `find_connected_components`)
106/// can drive them directly.
107pub struct DPccp<'g> {
108    graph: &'g JoinGraph,
109    dp: BTreeMap<u64, JoinPlan>,
110    all_mask: u64,
111    cost_estimator: CostEstimator,
112}
113
114impl<'g> DPccp<'g> {
115    pub fn new(graph: &'g JoinGraph) -> Self {
116        let all_mask = graph.full_set();
117        Self {
118            graph,
119            dp: BTreeMap::new(),
120            all_mask,
121            cost_estimator: CostEstimator::default(),
122        }
123    }
124
125    pub fn with_cost_estimator(graph: &'g JoinGraph, cost_estimator: CostEstimator) -> Self {
126        let all_mask = graph.full_set();
127        Self {
128            graph,
129            dp: BTreeMap::new(),
130            all_mask,
131            cost_estimator,
132        }
133    }
134
135    /// Find the optimal join plan for the full relation set. Falls
136    /// back to greedy for large queries. Returns `None` for empty
137    /// graphs.
138    pub fn optimize(mut self) -> Option<JoinPlan> {
139        let n = self.graph.relation_count();
140        if n == 0 {
141            return None;
142        }
143        if n == 1 {
144            return Some(JoinPlan::leaf(
145                0,
146                self.graph.cardinalities[0],
147                self.graph.access_costs[0],
148            ));
149        }
150        // Initialise base relations.
151        for i in 0..n {
152            self.dp.insert(
153                1u64 << i,
154                JoinPlan::leaf(i, self.graph.cardinalities[i], self.graph.access_costs[i]),
155            );
156        }
157        if n <= MAX_DP_RELATIONS {
158            if let Some(plan) = self.optimize_star(n) {
159                return Some(plan);
160            }
161        }
162        if n > MAX_DP_RELATIONS {
163            return self.greedy_optimize();
164        }
165        self.enumerate_csg_cmp_pairs(n);
166        if let Some(plan) = self.dp.get(&self.all_mask).cloned() {
167            return Some(plan);
168        }
169        // Disconnected: cross-join the components.
170        self.join_disconnected_components()
171    }
172
173    /// Enumerate every connected-subgraph/complement pair and feed it
174    /// through `enumerate_splits`.
175    fn enumerate_csg_cmp_pairs(&mut self, n: usize) {
176        let neighbors: Vec<Vec<usize>> = (0..n).map(|i| self.graph.neighbors(i)).collect();
177
178        // Keep connected subsets by their native u64 bitmask. This avoids
179        // narrowing a mask to `usize` merely to address a dense side table.
180        let mut connected = BTreeSet::new();
181        let mut prev_layer: Vec<u64> = Vec::with_capacity(n);
182        for i in 0..n {
183            let mask = 1u64 << i;
184            connected.insert(mask);
185            prev_layer.push(mask);
186        }
187
188        for _size in 2..=n {
189            let mut cur_layer: Vec<u64> = Vec::new();
190            for &s_mask in &prev_layer {
191                let Ok(min_node) = usize::try_from(s_mask.trailing_zeros()) else {
192                    return;
193                };
194                let mut node = 0usize;
195                let mut tmp = s_mask;
196                while tmp != 0 {
197                    if tmp & 1 == 1 {
198                        for &nb in &neighbors[node] {
199                            if nb > min_node && (s_mask & (1u64 << nb)) == 0 {
200                                let new_mask = s_mask | (1u64 << nb);
201                                if connected.insert(new_mask) {
202                                    cur_layer.push(new_mask);
203                                }
204                            }
205                        }
206                    }
207                    tmp >>= 1;
208                    node += 1;
209                }
210            }
211            for &subset_mask in &cur_layer {
212                self.enumerate_splits(subset_mask, &connected);
213            }
214            prev_layer = cur_layer;
215        }
216    }
217
218    fn optimize_star(&self, n: usize) -> Option<JoinPlan> {
219        let (centre, leaves) = self.star_shape(n)?;
220        let shift = u32::try_from(leaves.len()).ok()?;
221        let states = 1usize.checked_shl(shift)?;
222        let mut dp: Vec<Option<StarState>> = vec![None; states];
223        dp[0] = Some(StarState {
224            cardinality: self.graph.cardinalities[centre],
225            cost: self.graph.access_costs[centre],
226            prev_mask: 0,
227            leaf_pos: usize::MAX,
228        });
229
230        for mask in 0..states {
231            let Some(base) = dp[mask] else {
232                continue;
233            };
234            for (leaf_pos, (leaf_idx, edges)) in leaves.iter().enumerate() {
235                let bit = 1usize << leaf_pos;
236                if mask & bit != 0 {
237                    continue;
238                }
239                let leaf_cardinality = self.graph.cardinalities[*leaf_idx];
240                let mut cardinality = base.cardinality * leaf_cardinality;
241                for edge in edges {
242                    cardinality *= edge.selectivity;
243                }
244                let join_cost = self.join_cost(base.cardinality, leaf_cardinality).0;
245                let candidate = StarState {
246                    cardinality,
247                    cost: base.cost + self.graph.access_costs[*leaf_idx] + join_cost,
248                    prev_mask: mask,
249                    leaf_pos,
250                };
251                let next_mask = mask | bit;
252                let install = match &dp[next_mask] {
253                    Some(existing) => candidate.cost < existing.cost,
254                    None => true,
255                };
256                if install {
257                    dp[next_mask] = Some(candidate);
258                }
259            }
260        }
261
262        dp[states - 1]?;
263        let mut order: Vec<usize> = Vec::with_capacity(leaves.len());
264        let mut mask = states - 1;
265        while mask != 0 {
266            let state = dp[mask]?;
267            order.push(state.leaf_pos);
268            mask = state.prev_mask;
269        }
270        order.reverse();
271
272        let mut plan = JoinPlan::leaf(
273            centre,
274            self.graph.cardinalities[centre],
275            self.graph.access_costs[centre],
276        );
277        for leaf_pos in order {
278            let (leaf_idx, edges) = &leaves[leaf_pos];
279            let leaf = JoinPlan::leaf(
280                *leaf_idx,
281                self.graph.cardinalities[*leaf_idx],
282                self.graph.access_costs[*leaf_idx],
283            );
284            plan = self.join_plans(&plan, &leaf, edges);
285        }
286        Some(plan)
287    }
288
289    fn star_shape(&self, n: usize) -> Option<StarShape> {
290        if n < 3 || self.graph.edges.is_empty() {
291            return None;
292        }
293        let mut neighbor_masks = vec![0u64; n];
294        for edge in &self.graph.edges {
295            if edge.left.count_ones() != 1 || edge.right.count_ones() != 1 {
296                return None;
297            }
298            let left = usize::try_from(edge.left.trailing_zeros()).ok()?;
299            let right = usize::try_from(edge.right.trailing_zeros()).ok()?;
300            if left == right || left >= n || right >= n {
301                return None;
302            }
303            neighbor_masks[left] |= edge.right;
304            neighbor_masks[right] |= edge.left;
305        }
306        let candidates: Vec<usize> = neighbor_masks
307            .iter()
308            .enumerate()
309            .filter_map(|(idx, mask)| {
310                usize::try_from(mask.count_ones())
311                    .is_ok_and(|count| count == n - 1)
312                    .then_some(idx)
313            })
314            .collect();
315        if candidates.len() != 1 {
316            return None;
317        }
318        for centre in candidates {
319            let centre_mask = 1u64 << centre;
320            let mut by_leaf: BTreeMap<usize, Vec<JoinEdge>> = BTreeMap::new();
321            let mut valid = true;
322            for edge in &self.graph.edges {
323                let leaf_mask = if edge.left == centre_mask {
324                    edge.right
325                } else if edge.right == centre_mask {
326                    edge.left
327                } else {
328                    valid = false;
329                    break;
330                };
331                if leaf_mask == 0 || leaf_mask == centre_mask {
332                    valid = false;
333                    break;
334                }
335                let leaf = usize::try_from(leaf_mask.trailing_zeros()).ok()?;
336                by_leaf.entry(leaf).or_default().push(edge.clone());
337            }
338            if valid && by_leaf.len() == n - 1 {
339                return Some((centre, by_leaf.into_iter().collect()));
340            }
341        }
342        None
343    }
344
345    /// Enumerate every canonical split `(s1, s2)` of `subset_mask`
346    /// where `s1` contains the lowest set bit. Connectivity is
347    /// checked via the `connected` table; only pairs that survive get
348    /// fed through `emit_csg_cmp_pair`.
349    fn enumerate_splits(&mut self, subset_mask: u64, connected: &BTreeSet<u64>) {
350        let lowest_bit = subset_mask & subset_mask.wrapping_neg();
351        let rest = subset_mask ^ lowest_bit;
352
353        // Iterate proper non-empty submasks of `rest`. Each `sub | lowest_bit`
354        // forms a canonical S1 (containing the min element).
355        let mut sub_rest = rest.wrapping_sub(1) & rest;
356        while sub_rest != 0 {
357            let sub = sub_rest | lowest_bit;
358            let comp = subset_mask ^ sub;
359            if connected.contains(&sub) && connected.contains(&comp) {
360                if let (Some(plan1), Some(plan2)) =
361                    (self.dp.get(&sub).cloned(), self.dp.get(&comp).cloned())
362                {
363                    let edges = self
364                        .graph
365                        .edges_between(plan1.relations, plan2.relations)
366                        .into_iter()
367                        .cloned()
368                        .collect::<Vec<_>>();
369                    if !edges.is_empty() {
370                        self.emit_csg_cmp_pair(&plan1, &plan2, &edges, subset_mask);
371                    }
372                }
373            }
374            sub_rest = sub_rest.wrapping_sub(1) & rest;
375        }
376        // sub_rest == 0: S1 = {min element}, S2 = rest of subset.
377        if connected.contains(&rest) {
378            if let (Some(plan1), Some(plan2)) = (
379                self.dp.get(&lowest_bit).cloned(),
380                self.dp.get(&rest).cloned(),
381            ) {
382                let edges = self
383                    .graph
384                    .edges_between(plan1.relations, plan2.relations)
385                    .into_iter()
386                    .cloned()
387                    .collect::<Vec<_>>();
388                if !edges.is_empty() {
389                    self.emit_csg_cmp_pair(&plan1, &plan2, &edges, subset_mask);
390                }
391            }
392        }
393    }
394
395    /// Cost a candidate join and install the best variant in the DP table.
396    /// Cardinality is the cross-product times every edge's selectivity. The
397    /// physical SQL engine executes these equijoins as hash joins, so the
398    /// enumerator uses the same cost shape and records the executable kind.
399    fn emit_csg_cmp_pair(
400        &mut self,
401        plan1: &JoinPlan,
402        plan2: &JoinPlan,
403        edges: &[JoinEdge],
404        combined_mask: u64,
405    ) {
406        let candidate = self.join_plans(plan1, plan2, edges);
407        let install = match self.dp.get(&combined_mask) {
408            Some(existing) => candidate.cost < existing.cost,
409            None => true,
410        };
411        if install {
412            self.dp.insert(combined_mask, candidate);
413        }
414    }
415
416    fn join_plans(&self, plan1: &JoinPlan, plan2: &JoinPlan, edges: &[JoinEdge]) -> JoinPlan {
417        let mut cardinality = plan1.cardinality * plan2.cardinality;
418        for edge in edges {
419            cardinality *= edge.selectivity;
420        }
421        let c1 = plan1.cardinality;
422        let c2 = plan2.cardinality;
423        let (join_cost, kind) = self.join_cost(c1, c2);
424        JoinPlan {
425            relations: plan1.relations | plan2.relations,
426            cardinality,
427            cost: join_cost + plan1.cost + plan2.cost,
428            left: Some(Box::new(plan1.clone())),
429            right: Some(Box::new(plan2.clone())),
430            join_edge: edges.first().cloned(),
431            kind: Some(kind),
432        }
433    }
434
435    fn join_cost(&self, c1: f64, c2: f64) -> (f64, OperatorKind) {
436        let kind = OperatorKind::HashJoinInner;
437        (
438            self.cost_estimator.estimate_join(kind, c1, c2).total(),
439            kind,
440        )
441    }
442
443    fn cross_join_cost(&self, c1: f64, c2: f64) -> f64 {
444        self.cost_estimator
445            .estimate_join(OperatorKind::CrossJoin, c1, c2)
446            .total()
447    }
448
449    /// Cross-join every connected component in cardinality-ascending order.
450    fn join_disconnected_components(&mut self) -> Option<JoinPlan> {
451        let components = self.find_connected_components();
452        let mut component_plans: Vec<JoinPlan> = Vec::with_capacity(components.len());
453        for comp in &components {
454            if comp.len() == 1 {
455                let idx = *comp.first()?;
456                let plan = self.dp.get(&(1u64 << idx)).cloned()?;
457                component_plans.push(plan);
458                continue;
459            }
460            let mask: u64 = comp.iter().fold(0u64, |acc, i| acc | (1u64 << *i));
461            if let Some(plan) = self.dp.get(&mask).cloned() {
462                component_plans.push(plan);
463                continue;
464            }
465            // Component was not solved; recurse on a subgraph as a
466            // defensive fallback.
467            let original_indices: Vec<usize> = {
468                let mut v: Vec<usize> = comp.clone();
469                v.sort_unstable();
470                v
471            };
472            let sub_graph = self.build_subgraph(&original_indices)?;
473            let sub_plan =
474                DPccp::with_cost_estimator(&sub_graph, self.cost_estimator.clone()).optimize()?;
475            component_plans.push(remap_plan(&sub_plan, &original_indices));
476        }
477        component_plans.sort_by(|a, b| a.cardinality.total_cmp(&b.cardinality));
478        let mut iter = component_plans.into_iter();
479        let mut result = iter.next()?;
480        for plan in iter {
481            let combined = result.relations | plan.relations;
482            let cardinality = result.cardinality * plan.cardinality;
483            let cost = self.cross_join_cost(result.cardinality, plan.cardinality)
484                + result.cost
485                + plan.cost;
486            result = JoinPlan {
487                relations: combined,
488                cardinality,
489                cost,
490                left: Some(Box::new(result)),
491                right: Some(Box::new(plan)),
492                join_edge: None,
493                kind: Some(OperatorKind::CrossJoin),
494            };
495        }
496        Some(result)
497    }
498
499    /// Use BFS to enumerate the join graph's connected components.
500    fn find_connected_components(&self) -> Vec<Vec<usize>> {
501        let n = self.graph.relation_count();
502        let mut remaining: std::collections::BTreeSet<usize> = (0..n).collect();
503        let mut components: Vec<Vec<usize>> = Vec::new();
504        while let Some(&start) = remaining.iter().next() {
505            let mut visited: std::collections::BTreeSet<usize> = std::collections::BTreeSet::new();
506            visited.insert(start);
507            let mut stack: Vec<usize> = vec![start];
508            while let Some(node) = stack.pop() {
509                for nb in self.graph.neighbors(node) {
510                    if remaining.contains(&nb) && !visited.contains(&nb) {
511                        visited.insert(nb);
512                        stack.push(nb);
513                    }
514                }
515            }
516            for v in &visited {
517                remaining.remove(v);
518            }
519            components.push(visited.into_iter().collect());
520        }
521        components
522    }
523
524    /// Project a subgraph containing only `nodes` in their original indices.
525    /// Edge bitmasks are remapped to the dense `[0, k)` range used by the
526    /// recursive solve.
527    fn build_subgraph(&self, nodes: &[usize]) -> Option<JoinGraph> {
528        let mut sub = JoinGraph::new();
529        let mut index_map: BTreeMap<usize, usize> = BTreeMap::new();
530        for &old_idx in nodes {
531            let new_idx = sub.relations.len();
532            sub.relations.push(self.graph.relations[old_idx].clone());
533            sub.cardinalities.push(self.graph.cardinalities[old_idx]);
534            sub.access_costs.push(self.graph.access_costs[old_idx]);
535            index_map.insert(old_idx, new_idx);
536        }
537        for edge in &self.graph.edges {
538            let l_idx = usize::try_from(edge.left.trailing_zeros()).ok()?;
539            let r_idx = usize::try_from(edge.right.trailing_zeros()).ok()?;
540            if let (Some(&l_new), Some(&r_new)) = (index_map.get(&l_idx), index_map.get(&r_idx)) {
541                sub.edges.push(JoinEdge {
542                    left: 1_u64 << l_new,
543                    right: 1_u64 << r_new,
544                    selectivity: edge.selectivity,
545                });
546            }
547        }
548        Some(sub)
549    }
550
551    /// Greedy fallback for graphs with more than `MAX_DP_RELATIONS`
552    /// relations: at every step, pick the cheapest joinable pair until
553    /// only one plan remains. `O(n^3)`.
554    fn greedy_optimize(self) -> Option<JoinPlan> {
555        let mut active: BTreeMap<u64, JoinPlan> = self.dp.clone();
556        while active.len() > 1 {
557            let mut best_cost = f64::INFINITY;
558            let mut best_combined_mask: u64 = 0;
559            let mut best_plan: Option<JoinPlan> = None;
560            let items: Vec<(u64, JoinPlan)> = active.iter().map(|(k, v)| (*k, v.clone())).collect();
561            for i in 0..items.len() {
562                let (m1, ref p1) = items[i];
563                for (m2, p2) in items.iter().skip(i + 1) {
564                    let edges = self
565                        .graph
566                        .edges_between(p1.relations, p2.relations)
567                        .into_iter()
568                        .cloned()
569                        .collect::<Vec<_>>();
570                    if edges.is_empty() {
571                        continue;
572                    }
573                    let mut cardinality = p1.cardinality * p2.cardinality;
574                    for edge in &edges {
575                        cardinality *= edge.selectivity;
576                    }
577                    let (greedy_join_cost, kind) = self.join_cost(p1.cardinality, p2.cardinality);
578                    let cost = greedy_join_cost + p1.cost + p2.cost;
579                    if cost < best_cost {
580                        best_cost = cost;
581                        best_combined_mask = m1 | m2;
582                        best_plan = Some(JoinPlan {
583                            relations: p1.relations | p2.relations,
584                            cardinality,
585                            cost,
586                            left: Some(Box::new(p1.clone())),
587                            right: Some(Box::new(p2.clone())),
588                            join_edge: edges.first().cloned(),
589                            kind: Some(kind),
590                        });
591                    }
592                }
593            }
594            let Some(best_plan_unwrapped) = best_plan else {
595                // No more joinable edges; cross-join the rest in
596                // cardinality-ascending order.
597                let mut remaining: Vec<JoinPlan> = active.into_values().collect();
598                remaining.sort_by(|a, b| a.cardinality.total_cmp(&b.cardinality));
599                let mut iter = remaining.into_iter();
600                let mut result = iter.next()?;
601                for plan in iter {
602                    let combined = result.relations | plan.relations;
603                    let cardinality = result.cardinality * plan.cardinality;
604                    let cost = self.cross_join_cost(result.cardinality, plan.cardinality)
605                        + result.cost
606                        + plan.cost;
607                    result = JoinPlan {
608                        relations: combined,
609                        cardinality,
610                        cost,
611                        left: Some(Box::new(result)),
612                        right: Some(Box::new(plan)),
613                        join_edge: None,
614                        kind: Some(OperatorKind::CrossJoin),
615                    };
616                }
617                return Some(result);
618            };
619            // Drop every plan whose mask is fully contained in the
620            // newly merged mask, then insert the merged plan.
621            let drop: Vec<u64> = active
622                .keys()
623                .copied()
624                .filter(|rel_mask| rel_mask & best_combined_mask == *rel_mask)
625                .collect();
626            for k in drop {
627                active.remove(&k);
628            }
629            active.insert(best_combined_mask, best_plan_unwrapped);
630        }
631        active.into_values().next()
632    }
633}
634
635/// Remap relation indices in `plan` from a subgraph's dense range back to
636/// the parent graph's original indices.
637fn remap_plan(plan: &JoinPlan, original_indices: &[usize]) -> JoinPlan {
638    let new_relations = remap_mask(plan.relations, original_indices);
639    JoinPlan {
640        relations: new_relations,
641        cardinality: plan.cardinality,
642        cost: plan.cost,
643        left: plan
644            .left
645            .as_deref()
646            .map(|l| Box::new(remap_plan(l, original_indices))),
647        right: plan
648            .right
649            .as_deref()
650            .map(|r| Box::new(remap_plan(r, original_indices))),
651        join_edge: plan.join_edge.clone(),
652        kind: plan.kind,
653    }
654}
655
656fn remap_mask(mask: u64, original_indices: &[usize]) -> u64 {
657    let mut out = 0u64;
658    let mut m = mask;
659    let mut i = 0;
660    while m != 0 {
661        if m & 1 == 1 {
662            out |= 1u64 << original_indices[i];
663        }
664        m >>= 1;
665        i += 1;
666    }
667    out
668}
669
670#[cfg(test)]
671mod tests {
672    use super::*;
673    use crate::cost_model::CostCoefficients;
674
675    fn assert_executable_join_kinds(plan: &JoinPlan) {
676        match (&plan.left, &plan.right) {
677            (Some(left), Some(right)) => {
678                if plan.join_edge.is_some() {
679                    assert_eq!(plan.kind, Some(OperatorKind::HashJoinInner));
680                } else {
681                    assert_eq!(plan.kind, Some(OperatorKind::CrossJoin));
682                }
683                assert_executable_join_kinds(left);
684                assert_executable_join_kinds(right);
685            }
686            (None, None) => assert!(plan.kind.is_none()),
687            _ => panic!("join plan contains exactly one child: {plan:?}"),
688        }
689    }
690
691    #[test]
692    fn three_way_chain_picks_smallest_first() {
693        let mut g = JoinGraph::new();
694        let a = g.add_relation("a", 10.0).unwrap();
695        let b = g.add_relation("b", 100.0).unwrap();
696        let c = g.add_relation("c", 10_000.0).unwrap();
697        g.add_edge(a, b, 0.01).unwrap();
698        g.add_edge(b, c, 0.001).unwrap();
699        let plan = enumerate_dpccp(&g).unwrap();
700        assert_eq!(plan.relations, 0b111);
701        assert!(plan.left.is_some() && plan.right.is_some());
702        assert_executable_join_kinds(&plan);
703    }
704
705    #[test]
706    fn single_relation_returns_leaf() {
707        let mut g = JoinGraph::new();
708        g.add_relation("solo", 1.0).unwrap();
709        let plan = enumerate_dpccp(&g).unwrap();
710        assert!(plan.left.is_none());
711        assert!(plan.right.is_none());
712        assert_eq!(plan.relations, 0b1);
713    }
714
715    #[test]
716    fn explicit_cost_estimator_drives_join_cost() {
717        let mut graph = JoinGraph::new();
718        let left = graph.add_relation_with_cost("left", 10.0, 7.0).unwrap();
719        let right = graph.add_relation_with_cost("right", 100.0, 11.0).unwrap();
720        graph.add_edge(left, right, 0.5).unwrap();
721
722        let coefficients = CostCoefficients {
723            hashjoin_build_per_row: 2.0,
724            hashjoin_probe_per_row: 3.0,
725            ..CostCoefficients::default()
726        };
727        let estimator = CostEstimator::new(coefficients);
728        let expected_join_cost = estimator
729            .estimate_join(OperatorKind::HashJoinInner, 10.0, 100.0)
730            .total();
731
732        let plan = enumerate_dpccp_with_cost_estimator(&graph, estimator).unwrap();
733
734        assert_eq!(plan.cost, 7.0 + 11.0 + expected_join_cost);
735    }
736
737    #[test]
738    fn empty_graph_returns_none() {
739        let g = JoinGraph::new();
740        assert!(enumerate_dpccp(&g).is_none());
741    }
742
743    #[test]
744    fn disconnected_graph_cross_joins_components() {
745        let mut g = JoinGraph::new();
746        let a = g.add_relation("a", 50.0).unwrap();
747        let b = g.add_relation("b", 60.0).unwrap();
748        let c = g.add_relation("c", 70.0).unwrap();
749        g.add_edge(a, b, 0.5).unwrap();
750        // c is a disconnected component.
751        let _ = c;
752        let plan = enumerate_dpccp(&g).unwrap();
753        // The cross-join must cover every relation.
754        assert_eq!(plan.relations, 0b111);
755        assert_executable_join_kinds(&plan);
756    }
757
758    #[test]
759    fn star_query_picks_nested_plan() {
760        // centre as the centre, leaf_b/c/d as leaves: every connects to
761        // centre.
762        let mut g = JoinGraph::new();
763        let centre = g.add_relation("centre", 1_000.0).unwrap();
764        let leaf_b = g.add_relation("b", 10.0).unwrap();
765        let leaf_c = g.add_relation("c", 20.0).unwrap();
766        let leaf_d = g.add_relation("d", 30.0).unwrap();
767        g.add_edge(centre, leaf_b, 0.01).unwrap();
768        g.add_edge(centre, leaf_c, 0.01).unwrap();
769        g.add_edge(centre, leaf_d, 0.01).unwrap();
770        let plan = enumerate_dpccp(&g).unwrap();
771        assert_eq!(plan.relations, 0b1111);
772        assert!(plan.cost > 0.0);
773        assert_executable_join_kinds(&plan);
774    }
775
776    #[test]
777    fn greedy_fallback_kicks_in_above_threshold() {
778        // With > MAX_DP_RELATIONS we expect the greedy fallback to
779        // produce a plan that still covers every relation.
780        let mut g = JoinGraph::new();
781        let n = MAX_DP_RELATIONS + 2;
782        let mut prev: usize = 0;
783        for i in 0..n {
784            let idx = g.add_relation(format!("t{i}"), 100.0).unwrap();
785            if i > 0 {
786                g.add_edge(prev, idx, 0.05).unwrap();
787            }
788            prev = idx;
789        }
790        let plan = enumerate_dpccp(&g).unwrap();
791        assert_eq!(plan.relations.count_ones() as usize, n);
792        assert_executable_join_kinds(&plan);
793    }
794
795    #[test]
796    fn threshold_sized_star_uses_exact_plan() {
797        let mut g = JoinGraph::new();
798        let centre = g.add_relation("centre", 1_000.0).unwrap();
799        for i in 1..MAX_DP_RELATIONS {
800            let leaf = g
801                .add_relation(format!("leaf{i}"), 100.0 + i as f64)
802                .unwrap();
803            g.add_edge(centre, leaf, 0.01).unwrap();
804        }
805
806        let plan = enumerate_dpccp(&g).unwrap();
807
808        assert_eq!(plan.relations.count_ones() as usize, MAX_DP_RELATIONS);
809        assert_eq!(plan.relations, g.full_set());
810        assert!(plan.cost > 0.0);
811    }
812}