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_runtime::hash::Hash128;
11use reifydb_transaction::transaction::Transaction;
12use reifydb_value::{
13 fragment::Fragment,
14 reifydb_assertions,
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 mut result_rows = Vec::new();
106 let mut result_row_numbers: Vec<RowNumber> = Vec::new();
107
108 let right_col_indices: Vec<usize> = common_columns.iter().map(|(_, _, ri)| *ri).collect();
109 let mut hash_buf = Vec::with_capacity(256);
110 let mut hash_table: HashMap<Hash128, Vec<usize>> = HashMap::new();
111 let right_rows = right_columns.row_count();
112 for j in 0..right_rows {
113 if let Some(h) = compute_join_hash(&right_columns, &right_col_indices, j, &mut hash_buf) {
114 hash_table.entry(h).or_default().push(j);
115 }
116 }
117
118 let left_col_indices: Vec<usize> = common_columns.iter().map(|(_, li, _)| *li).collect();
119
120 for i in 0..left_rows {
121 let left_row = left_columns.get_row(i);
122 let mut matched = false;
123
124 let candidates = compute_join_hash(&left_columns, &left_col_indices, i, &mut hash_buf)
125 .and_then(|h| hash_table.get(&h));
126
127 if let Some(indices) = candidates {
128 for &j in indices {
129 let right_row = right_columns.get_row(j);
130
131 let all_match = common_columns.iter().all(|(_, left_idx, right_idx)| {
132 left_row[*left_idx] == right_row[*right_idx]
133 });
134
135 if all_match {
136 let mut combined = left_row.clone();
137 for (idx, value) in right_row.iter().enumerate() {
138 if !excluded_right_cols.contains(&idx) {
139 combined.push(value.clone());
140 }
141 }
142 result_rows.push(combined);
143 matched = true;
144 if !left_row_numbers.is_empty() {
145 result_row_numbers.push(left_row_numbers[i]);
146 }
147 }
148 }
149 }
150
151 if !matched && matches!(self.join_type, JoinType::Left) {
152 let mut combined = left_row.clone();
153
154 let undefined_count = right_columns.len() - excluded_right_cols.len();
155 combined.extend(vec![Value::none(); undefined_count]);
156 result_rows.push(combined);
157 if !left_row_numbers.is_empty() {
158 result_row_numbers.push(left_row_numbers[i]);
159 }
160 }
161 }
162
163 let names_refs: Vec<&str> = resolved.qualified_names.iter().map(|s| s.as_str()).collect();
164 let columns = if result_row_numbers.is_empty() {
165 Columns::from_rows(&names_refs, &result_rows)
166 } else {
167 Columns::from_rows(&names_refs, &result_rows).with_row_numbers(result_row_numbers)
168 };
169
170 self.headers = Some(ColumnHeaders::from_columns(&columns));
171 Ok(Some(columns))
172 }
173
174 fn headers(&self) -> Option<ColumnHeaders> {
175 self.headers.clone()
176 }
177}