Skip to main content

reifydb_engine/vm/volcano/join/
natural.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use 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}