Skip to main content

kermit_algos/
hash_triejoin.rs

1//! Hash Trie Join — worst-case-optimal multi-way join over hash tries.
2//!
3//! Implements the probe phase from §3.2.3 of the SIGMOD 2020 paper
4//! "Combining Worst-Case Optimal and Traditional Binary Join Processing".
5//! Coordinates one [`HashTrieIterator`] per body predicate, descending in
6//! lockstep variable by variable. At each depth it picks the iterator with
7//! the smallest hash table (`argmin size`), iterates that table's hashes,
8//! and probes the rest via [`HashTrieIterator::lookup`]. At the leaf level
9//! it cross-products the participating iterators' tuple chains and
10//! verifies the actual join condition (hash collisions can produce false
11//! positives at any inner level).
12
13use {
14    crate::{join_algo::JoinAlgo, leapfrog_triejoin::global_attribute_order},
15    kermit_iters::{HashTrieIterable, HashTrieIterator},
16    kermit_parser::{JoinQuery, Term},
17    std::collections::HashMap,
18};
19
20/// Indexes the variables in a query for the hash-trie-join algorithm.
21///
22/// Like [`crate::leapfrog_triejoin::build_variable_index`]: head-first
23/// canonical indices fix the output column order, while the *descent* order is
24/// a valid global attribute order (shared [`global_attribute_order`] helper).
25/// The hash join descends each relation one physical column per depth, so the
26/// same subject-position-constant hazard applies.
27fn build_variable_index(query: &JoinQuery) -> (Vec<usize>, Vec<Vec<usize>>) {
28    let mut var_to_index: HashMap<String, usize> = HashMap::new();
29    let mut next_index: usize = 0;
30
31    let register_var = |name: &str, map: &mut HashMap<String, usize>, next: &mut usize| {
32        *map.entry(name.to_string()).or_insert_with(|| {
33            let idx = *next;
34            *next += 1;
35            idx
36        })
37    };
38
39    for t in &query.head.terms {
40        if let Term::Var(ref vname) = t {
41            let _ = register_var(vname, &mut var_to_index, &mut next_index);
42        }
43    }
44    for pred in &query.body {
45        for t in &pred.terms {
46            if let Term::Var(ref vname) = t {
47                let _ = register_var(vname, &mut var_to_index, &mut next_index);
48            }
49        }
50    }
51
52    let mut predicate_variables: Vec<Vec<usize>> = Vec::with_capacity(query.body.len());
53    for pred in &query.body {
54        let mut vars_for_pred: Vec<usize> = Vec::new();
55        for t in &pred.terms {
56            if let Term::Var(ref vname) = t {
57                if let Some(idx) = var_to_index.get(vname) {
58                    vars_for_pred.push(*idx);
59                }
60            }
61        }
62        predicate_variables.push(vars_for_pred);
63    }
64
65    let variable_ordering = global_attribute_order(var_to_index.len(), &predicate_variables);
66
67    (variable_ordering, predicate_variables)
68}
69
70/// Build `variable_to_iter_map[i] = predicate indices that carry the i-th
71/// variable in `variable_ordering`. Identical shape to LFTJ's inline
72/// construction.
73fn build_variable_to_iter_map(
74    variable_ordering: &[usize], predicate_variables: &[Vec<usize>],
75) -> Vec<Vec<usize>> {
76    variable_ordering
77        .iter()
78        .map(|v| {
79            predicate_variables
80                .iter()
81                .enumerate()
82                .filter_map(|(i, vars)| {
83                    if vars.contains(v) {
84                        Some(i)
85                    } else {
86                        None
87                    }
88                })
89                .collect()
90        })
91        .collect()
92}
93
94/// Verify a candidate result tuple's join condition and construct the
95/// output tuple in `variable_ordering` order. Returns `None` if any
96/// variable mentioned by two or more predicates has inconsistent values
97/// in the candidate (a hash false positive).
98///
99/// `candidate[k]` is one tuple from the k-th participating relation's
100/// leaf chain. `predicate_variables[k]` lists the variable indices carried
101/// by the k-th relation (in attribute order). `arity` is the number of
102/// distinct variables in the query.
103fn verify_and_construct(
104    candidate: &[&Vec<usize>], predicate_variables: &[Vec<usize>], arity: usize,
105) -> Option<Vec<usize>> {
106    let mut result: Vec<Option<usize>> = vec![None; arity];
107    for (rel_idx, tuple) in candidate.iter().enumerate() {
108        for (col, &var_idx) in predicate_variables[rel_idx].iter().enumerate() {
109            let v = tuple[col];
110            match result[var_idx] {
111                | None => result[var_idx] = Some(v),
112                | Some(existing) if existing != v => return None,
113                | _ => {},
114            }
115        }
116    }
117    Some(result.into_iter().map(Option::unwrap).collect())
118}
119
120/// Algorithm 3 from the paper. Recursively descends through attribute
121/// positions; emits all verified result tuples into `output`.
122///
123/// **Contract.** `enumerate(i, ...)` is called with each iterator in
124/// `variable_to_iter_map[i]` positioned at depth `i - 1` (or pre-root for
125/// `i == 0`). The function descends each one level (to depth `i`), scans
126/// at depth `i`, recurses on matches, then ascends back. This way the
127/// caller's stack is unchanged on return.
128fn enumerate<IT: HashTrieIterator>(
129    i: usize, arity: usize, iters: &mut [IT], predicate_variables: &[Vec<usize>],
130    variable_to_iter_map: &[Vec<usize>], output: &mut Vec<Vec<usize>>,
131) {
132    if i == arity {
133        emit_leaf(iters, predicate_variables, arity, output);
134        return;
135    }
136
137    let i_join = &variable_to_iter_map[i];
138    if i_join.is_empty() {
139        return; // no relation carries this variable
140    }
141
142    // Descend every participating iterator to depth i. Track how many we
143    // successfully opened so we can match `up` calls on early return.
144    let mut opened = 0;
145    let mut descend_ok = true;
146    for &idx in i_join {
147        if iters[idx].open() {
148            opened += 1;
149        } else {
150            descend_ok = false;
151            break;
152        }
153    }
154
155    if descend_ok {
156        let i_scan = *i_join
157            .iter()
158            .min_by_key(|&&idx| iters[idx].size())
159            .expect("i_join non-empty");
160
161        while !iters[i_scan].at_end() {
162            let h = iters[i_scan]
163                .key()
164                .expect("scan iterator not at end => key Some");
165
166            // Probe the other iterators in i_join for this hash.
167            let mut all_match = true;
168            for &idx in i_join {
169                if idx == i_scan {
170                    continue;
171                }
172                if !iters[idx].lookup(h) {
173                    all_match = false;
174                    break;
175                }
176            }
177
178            if all_match {
179                enumerate(
180                    i + 1,
181                    arity,
182                    iters,
183                    predicate_variables,
184                    variable_to_iter_map,
185                    output,
186                );
187            }
188
189            iters[i_scan].next();
190        }
191    }
192
193    // Ascend back to the parent depth, matching the descend count so
194    // partial-open failures are symmetric.
195    for &idx in &i_join[..opened] {
196        let _ = iters[idx].up();
197    }
198}
199
200/// Algorithm 3 lines 16–19. Cross-product the leaf chains of every
201/// iterator and emit each verified candidate.
202fn emit_leaf<IT: HashTrieIterator>(
203    iters: &[IT], predicate_variables: &[Vec<usize>], arity: usize, output: &mut Vec<Vec<usize>>,
204) {
205    let chains: Vec<&[Vec<usize>]> = iters
206        .iter()
207        .map(|it| {
208            it.leaf_tuples()
209                .expect("at leaf level for every participating iter")
210        })
211        .collect();
212    if chains.iter().any(|c| c.is_empty()) {
213        return;
214    }
215
216    let mut cursor: Vec<usize> = vec![0; chains.len()];
217    loop {
218        let candidate: Vec<&Vec<usize>> = chains.iter().zip(&cursor).map(|(c, &i)| &c[i]).collect();
219        if let Some(result) = verify_and_construct(&candidate, predicate_variables, arity) {
220            output.push(result);
221        }
222        if !advance_cursor(&mut cursor, &chains) {
223            break;
224        }
225    }
226}
227
228/// Advance a mixed-base cursor over the chain lengths. Returns `false`
229/// once every position has been exhausted (overflow off the high end).
230fn advance_cursor(cursor: &mut [usize], chains: &[&[Vec<usize>]]) -> bool {
231    let mut k = 0;
232    loop {
233        if k >= cursor.len() {
234            return false;
235        }
236        cursor[k] += 1;
237        if cursor[k] < chains[k].len() {
238            return true;
239        }
240        cursor[k] = 0;
241        k += 1;
242    }
243}
244
245/// Entry point for the hash-trie-join algorithm.
246///
247/// Implements [`JoinAlgo`] for any [`HashTrieIterable`] data structure.
248/// See the module docs for the algorithm overview.
249pub struct HashTriejoin {}
250
251impl<DS> JoinAlgo<DS> for HashTriejoin
252where
253    DS: HashTrieIterable,
254{
255    fn join_iter(
256        query: JoinQuery, datastructures: HashMap<String, &DS>,
257    ) -> impl Iterator<Item = Vec<usize>> {
258        let (variable_ordering, predicate_variables) = build_variable_index(&query);
259        let mut iters: Vec<_> = query
260            .body
261            .iter()
262            .map(|pred| {
263                datastructures
264                    .get(&pred.name)
265                    .expect("Missing datastructure for predicate name")
266                    .hash_trie_iter()
267            })
268            .collect();
269        let variable_to_iter_map =
270            build_variable_to_iter_map(&variable_ordering, &predicate_variables);
271        let mut output = Vec::new();
272        enumerate(
273            0,
274            variable_ordering.len(),
275            &mut iters,
276            &predicate_variables,
277            &variable_to_iter_map,
278            &mut output,
279        );
280        output.into_iter()
281    }
282}
283
284#[cfg(test)]
285mod tests {
286    use super::*;
287
288    #[test]
289    fn variable_index_triangle() {
290        let query: JoinQuery = "Q(X, Y, Z) :- R(X, Y), S(Y, Z), T(X, Z).".parse().unwrap();
291        let (ordering, predicate_vars) = build_variable_index(&query);
292        assert_eq!(ordering, vec![0, 1, 2]);
293        assert_eq!(predicate_vars, vec![vec![0, 1], vec![1, 2], vec![0, 2]]);
294    }
295
296    #[test]
297    fn variable_to_iter_map_triangle() {
298        let pred_vars = vec![vec![0, 1], vec![1, 2], vec![0, 2]];
299        let map = build_variable_to_iter_map(&[0, 1, 2], &pred_vars);
300        assert_eq!(map, vec![vec![0, 2], vec![0, 1], vec![1, 2]]);
301    }
302
303    #[test]
304    fn verify_constructs_when_shared_var_agrees() {
305        // R(X, Y), S(Y, Z) with X=1, Y=2, Z=3.
306        // predicate_variables = [[0, 1], [1, 2]]
307        // arity = 3
308        let r_tuple = vec![1, 2];
309        let s_tuple = vec![2, 3];
310        let candidate: Vec<&Vec<usize>> = vec![&r_tuple, &s_tuple];
311        let pv = vec![vec![0, 1], vec![1, 2]];
312        assert_eq!(
313            verify_and_construct(&candidate, &pv, 3),
314            Some(vec![1, 2, 3])
315        );
316    }
317
318    #[test]
319    fn verify_rejects_when_shared_var_disagrees() {
320        // R(X, Y), S(Y, Z) but Y differs in R (=2) vs S (=99)
321        let r_tuple = vec![1, 2];
322        let s_tuple = vec![99, 3];
323        let candidate: Vec<&Vec<usize>> = vec![&r_tuple, &s_tuple];
324        let pv = vec![vec![0, 1], vec![1, 2]];
325        assert_eq!(verify_and_construct(&candidate, &pv, 3), None);
326    }
327
328    #[test]
329    fn verify_passes_no_shared_var() {
330        // Two disjoint unary predicates R(X), S(Y).
331        let r_tuple = vec![1];
332        let s_tuple = vec![2];
333        let candidate: Vec<&Vec<usize>> = vec![&r_tuple, &s_tuple];
334        let pv = vec![vec![0], vec![1]];
335        assert_eq!(verify_and_construct(&candidate, &pv, 2), Some(vec![1, 2]));
336    }
337
338    #[test]
339    fn enumerate_unary_intersection() {
340        use kermit_ds::{HashTrie, Relation};
341        // Explicit `HashTrie` annotation pins the default `H = SipHashStrategy`
342        // since the local bindings escape into `Vec<_>` iter values that
343        // would otherwise leave `H` ambiguous.
344        let r: HashTrie = HashTrie::from_tuples(1.into(), vec![vec![1], vec![2], vec![3]]);
345        let s: HashTrie = HashTrie::from_tuples(1.into(), vec![vec![2], vec![3], vec![4]]);
346        let mut iters = vec![r.hash_trie_iter(), s.hash_trie_iter()];
347        // Inline the same setup the JoinAlgo entry point does.
348        let predicate_variables = vec![vec![0], vec![0]];
349        let variable_to_iter_map = vec![vec![0, 1]];
350        let mut output = Vec::new();
351        enumerate(
352            0,
353            1,
354            &mut iters,
355            &predicate_variables,
356            &variable_to_iter_map,
357            &mut output,
358        );
359        output.sort();
360        assert_eq!(output, vec![vec![2], vec![3]]);
361    }
362
363    #[test]
364    fn join_algo_unary_intersection() {
365        use kermit_ds::{HashTrie, Relation};
366        let r = HashTrie::from_tuples(1.into(), vec![vec![1], vec![2], vec![3]]);
367        let s = HashTrie::from_tuples(1.into(), vec![vec![2], vec![3], vec![4]]);
368        let query: JoinQuery = "Q(X) :- R(X), S(X).".parse().unwrap();
369        let mut ds: HashMap<String, &HashTrie> = HashMap::new();
370        ds.insert("R".to_string(), &r);
371        ds.insert("S".to_string(), &s);
372        let mut out: Vec<Vec<usize>> = HashTriejoin::join_iter(query, ds).collect();
373        out.sort();
374        assert_eq!(out, vec![vec![2], vec![3]]);
375    }
376
377    #[test]
378    fn join_algo_triangle() {
379        use kermit_ds::{HashTrie, Relation};
380        let r = HashTrie::from_tuples(2.into(), vec![vec![1, 2], vec![2, 3], vec![3, 1]]);
381        let s = HashTrie::from_tuples(2.into(), vec![vec![2, 3], vec![3, 1], vec![1, 2]]);
382        let t = HashTrie::from_tuples(2.into(), vec![vec![1, 3], vec![2, 1], vec![3, 2]]);
383        let query: JoinQuery = "Q(X, Y, Z) :- R(X, Y), S(Y, Z), T(X, Z).".parse().unwrap();
384        let mut ds: HashMap<String, &HashTrie> = HashMap::new();
385        ds.insert("R".to_string(), &r);
386        ds.insert("S".to_string(), &s);
387        ds.insert("T".to_string(), &t);
388        let mut out: Vec<Vec<usize>> = HashTriejoin::join_iter(query, ds).collect();
389        out.sort();
390        assert_eq!(out, vec![vec![1, 2, 3], vec![2, 3, 1], vec![3, 1, 2]]);
391    }
392}