1use {
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
20fn 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
70fn 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
94fn 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
120fn 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; }
141
142 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 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 for &idx in &i_join[..opened] {
196 let _ = iters[idx].up();
197 }
198}
199
200fn 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
228fn 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
245pub 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 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 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 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 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 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}