reifydb_engine/vm/volcano/join/
natural.rs1use std::collections::{HashMap, HashSet};
5
6use reifydb_core::{
7 common::JoinType,
8 value::column::{columns::Columns, headers::ColumnHeaders},
9};
10use reifydb_transaction::transaction::Transaction;
11use reifydb_value::{
12 fragment::Fragment,
13 reifydb_assertions,
14 util::hash::Hash128,
15 value::{Value, row_number::RowNumber},
16};
17use tracing::instrument;
18
19use super::common::{JoinContext, compute_join_hash, load_and_merge_all, resolve_column_names};
20use crate::{
21 Result,
22 vm::volcano::query::{QueryContext, QueryNode},
23};
24
25pub struct NaturalJoinNode {
26 left: Box<dyn QueryNode>,
27 right: Box<dyn QueryNode>,
28 join_type: JoinType,
29 alias: Option<Fragment>,
30 headers: Option<ColumnHeaders>,
31 context: JoinContext,
32}
33
34impl NaturalJoinNode {
35 pub(crate) fn new(
36 left: Box<dyn QueryNode>,
37 right: Box<dyn QueryNode>,
38 join_type: JoinType,
39 alias: Option<Fragment>,
40 ) -> Self {
41 Self {
42 left,
43 right,
44 join_type,
45 alias,
46 headers: None,
47 context: JoinContext::new(),
48 }
49 }
50
51 fn find_common_columns(left_columns: &Columns, right_columns: &Columns) -> Vec<(String, usize, usize)> {
52 let mut common_columns = Vec::new();
53
54 for (left_idx, left_col) in left_columns.iter().enumerate() {
55 for (right_idx, right_col) in right_columns.iter().enumerate() {
56 if left_col.name() == right_col.name() {
57 common_columns.push((left_col.name().text().to_string(), left_idx, right_idx));
58 }
59 }
60 }
61
62 common_columns
63 }
64}
65
66impl QueryNode for NaturalJoinNode {
67 #[instrument(name = "volcano::join::natural::initialize", level = "trace", skip_all)]
68 fn initialize<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &QueryContext) -> Result<()> {
69 self.context.set(ctx);
70 self.left.initialize(rx, ctx)?;
71 self.right.initialize(rx, ctx)?;
72 Ok(())
73 }
74
75 #[instrument(name = "volcano::join::natural::next", level = "trace", skip_all)]
76 fn next<'a>(&mut self, rx: &mut Transaction<'a>, ctx: &mut QueryContext) -> Result<Option<Columns>> {
77 reifydb_assertions! {
78 assert!(self.context.is_initialized(), "NaturalJoinNode::next() called before initialize()");
79 }
80
81 if self.headers.is_some() {
82 return Ok(None);
83 }
84
85 let left_columns = load_and_merge_all(&mut self.left, rx, ctx)?;
86 let right_columns = load_and_merge_all(&mut self.right, rx, ctx)?;
87
88 let left_rows = left_columns.row_count();
89 let left_row_numbers = left_columns.row_numbers().to_vec();
90
91 let common_columns = Self::find_common_columns(&left_columns, &right_columns);
92
93 if common_columns.is_empty() {
94 return Ok(None);
95 }
96
97 let excluded_right_cols: HashSet<usize> =
98 common_columns.iter().map(|(_, _, right_idx)| *right_idx).collect();
99
100 let excluded_indices: Vec<usize> = excluded_right_cols.iter().copied().collect();
101
102 let resolved =
103 resolve_column_names(&left_columns, &right_columns, &self.alias, Some(&excluded_indices));
104
105 let right_col_indices: Vec<usize> = common_columns.iter().map(|(_, _, ri)| *ri).collect();
106 let mut hash_buf = Vec::with_capacity(256);
107 let hash_table = Self::build(&right_columns, &right_col_indices, &mut hash_buf);
108
109 let left_col_indices: Vec<usize> = common_columns.iter().map(|(_, li, _)| *li).collect();
110
111 let (result_rows, result_row_numbers) = self.probe(
112 &left_columns,
113 &right_columns,
114 &hash_table,
115 &common_columns,
116 &excluded_right_cols,
117 &left_col_indices,
118 &left_row_numbers,
119 left_rows,
120 &mut hash_buf,
121 );
122
123 let columns = Self::materialize(&resolved.qualified_names, result_rows, result_row_numbers);
124
125 self.headers = Some(ColumnHeaders::from_columns(&columns));
126 Ok(Some(columns))
127 }
128
129 fn headers(&self) -> Option<ColumnHeaders> {
130 self.headers.clone()
131 }
132}
133
134impl NaturalJoinNode {
135 #[instrument(level = "trace", skip_all, name = "volcano::join::natural::build")]
136 fn build(
137 right_columns: &Columns,
138 right_col_indices: &[usize],
139 hash_buf: &mut Vec<u8>,
140 ) -> HashMap<Hash128, Vec<usize>> {
141 let mut hash_table: HashMap<Hash128, Vec<usize>> = HashMap::new();
142 let right_rows = right_columns.row_count();
143 for j in 0..right_rows {
144 if let Some(h) = compute_join_hash(right_columns, right_col_indices, j, hash_buf) {
145 hash_table.entry(h).or_default().push(j);
146 }
147 }
148 hash_table
149 }
150
151 #[allow(clippy::too_many_arguments)]
152 #[instrument(level = "trace", skip_all, name = "volcano::join::natural::probe")]
153 fn probe(
154 &self,
155 left_columns: &Columns,
156 right_columns: &Columns,
157 hash_table: &HashMap<Hash128, Vec<usize>>,
158 common_columns: &[(String, usize, usize)],
159 excluded_right_cols: &HashSet<usize>,
160 left_col_indices: &[usize],
161 left_row_numbers: &[RowNumber],
162 left_rows: usize,
163 hash_buf: &mut Vec<u8>,
164 ) -> (Vec<Vec<Value>>, Vec<RowNumber>) {
165 let mut result_rows = Vec::new();
166 let mut result_row_numbers: Vec<RowNumber> = Vec::new();
167
168 for i in 0..left_rows {
169 let left_row = left_columns.get_row(i);
170 let mut matched = false;
171
172 let candidates = compute_join_hash(left_columns, left_col_indices, i, hash_buf)
173 .and_then(|h| hash_table.get(&h));
174
175 if let Some(indices) = candidates {
176 for &j in indices {
177 let right_row = right_columns.get_row(j);
178
179 let all_match = common_columns.iter().all(|(_, left_idx, right_idx)| {
180 left_row[*left_idx] == right_row[*right_idx]
181 });
182
183 if all_match {
184 let mut combined = left_row.clone();
185 for (idx, value) in right_row.iter().enumerate() {
186 if !excluded_right_cols.contains(&idx) {
187 combined.push(value.clone());
188 }
189 }
190 result_rows.push(combined);
191 matched = true;
192 if !left_row_numbers.is_empty() {
193 result_row_numbers.push(left_row_numbers[i]);
194 }
195 }
196 }
197 }
198
199 if !matched && matches!(self.join_type, JoinType::Left) {
200 let mut combined = left_row.clone();
201
202 let undefined_count = right_columns.len() - excluded_right_cols.len();
203 combined.extend(vec![Value::none(); undefined_count]);
204 result_rows.push(combined);
205 if !left_row_numbers.is_empty() {
206 result_row_numbers.push(left_row_numbers[i]);
207 }
208 }
209 }
210
211 (result_rows, result_row_numbers)
212 }
213
214 #[instrument(level = "trace", skip_all, name = "volcano::join::natural::materialize")]
215 fn materialize(
216 qualified_names: &[String],
217 result_rows: Vec<Vec<Value>>,
218 result_row_numbers: Vec<RowNumber>,
219 ) -> Columns {
220 let names_refs: Vec<&str> = qualified_names.iter().map(|s| s.as_str()).collect();
221 if result_row_numbers.is_empty() {
222 Columns::from_rows(&names_refs, &result_rows)
223 } else {
224 Columns::from_rows(&names_refs, &result_rows).with_row_numbers(result_row_numbers)
225 }
226 }
227}