radixdb_executor/binding/
source.rs1use std::sync::Arc;
4
5use radixdb_core::{Error, Result};
6use radixdb_functions::FunctionRegistry;
7use radixdb_sql::ast::{Expression, JoinTableSource, SelectStatement, Statement};
8use radixdb_storage::mvcc::engine::MVCCEngine;
9use radixdb_storage::mvcc::ViewDefinition;
10use radixdb_storage::traits::Engine;
11use rustc_hash::{FxHashMap, FxHashSet};
12
13use crate::context::ExecutionContext;
14use crate::dispatch::cache::QueryCache;
15use crate::utils::extract_base_column_name;
16
17const MAX_VIEW_BINDING_DEPTH: usize = 32;
18
19#[doc(hidden)]
22pub trait SourceBindingHost {
23 type BindingCache: Default;
24
25 fn source_binding_engine(&self) -> &MVCCEngine;
26
27 fn source_binding_functions(&self) -> &FunctionRegistry;
28
29 fn source_binding_query_cache(&self) -> &QueryCache<Self::BindingCache>;
30
31 fn source_binding_view(&self, name_lower: &str) -> Result<Option<Arc<ViewDefinition>>>;
32}
33
34#[doc(hidden)]
36pub trait SourceBindingExt: SourceBindingHost {
37 fn parse_view_statement(&self, view_query: &str) -> Result<Arc<Statement>> {
38 if let Some(cached) = self.source_binding_query_cache().get(view_query) {
39 if !matches!(cached.statement(), Statement::Select(_)) {
40 return Err(Error::InvalidArgument(
41 "View definition is not a SELECT statement".to_string(),
42 ));
43 }
44 return Ok(cached.statement);
45 }
46
47 let mut statements = radixdb_sql::parse_sql(view_query).map_err(|error| {
48 Error::InvalidArgument(format!("Failed to parse view query: {error}"))
49 })?;
50 if statements.len() != 1 {
51 return Err(Error::InvalidArgument(
52 "View definition must contain exactly one statement".to_string(),
53 ));
54 }
55 let statement = statements
56 .pop()
57 .expect("single statement length was checked");
58 if !matches!(statement, Statement::Select(_)) {
59 return Err(Error::InvalidArgument(
60 "View definition is not a SELECT statement".to_string(),
61 ));
62 }
63
64 Ok(self
65 .source_binding_query_cache()
66 .put(view_query, Arc::new(statement), false, 0)
67 .statement)
68 }
69
70 fn collect_select_binding_columns(
71 &self,
72 statement: &SelectStatement,
73 context: &ExecutionContext,
74 columns: &mut FxHashMap<String, usize>,
75 depth: usize,
76 ) -> Result<()> {
77 let has_star = statement.columns.iter().any(|expression| {
78 matches!(
79 expression,
80 Expression::Star(_) | Expression::QualifiedStar(_)
81 )
82 });
83 if has_star {
84 if let Some(source) = &statement.table_expr {
85 self.collect_join_binding_columns(source, context, columns, depth + 1)?;
86 }
87 }
88 for expression in &statement.columns {
89 match expression {
90 Expression::Aliased(aliased) => {
91 add_bound_join_column(columns, &aliased.alias.value)
92 }
93 Expression::Identifier(identifier) => {
94 add_bound_join_column(columns, &identifier.value)
95 }
96 Expression::QualifiedIdentifier(identifier) => {
97 add_bound_join_column(columns, &identifier.name.value)
98 }
99 Expression::Star(_) | Expression::QualifiedStar(_) => {}
100 _ => {}
101 }
102 }
103 Ok(())
104 }
105
106 fn collect_join_binding_columns(
107 &self,
108 expression: &Expression,
109 context: &ExecutionContext,
110 columns: &mut FxHashMap<String, usize>,
111 depth: usize,
112 ) -> Result<()> {
113 if depth > MAX_VIEW_BINDING_DEPTH {
114 return Ok(());
115 }
116 match expression {
117 Expression::TableSource(source) => {
118 let name = source.name.value_lower.as_str();
119 if let Some((cte_columns, _, _)) = context.get_cte_by_lower(name) {
120 for column in cte_columns.iter() {
121 add_bound_join_column(columns, column);
122 }
123 } else if let Some(view) = self.source_binding_view(name)? {
124 let statement = self.parse_view_statement(&view.query)?;
125 if let Statement::Select(select) = statement.as_ref() {
126 self.collect_select_binding_columns(select, context, columns, depth + 1)?;
127 }
128 } else {
129 let schema = self.source_binding_engine().get_table_schema(name)?;
130 for column in &schema.columns {
131 add_bound_join_column(columns, &column.name);
132 }
133 }
134 }
135 Expression::JoinSource(join) => {
136 self.collect_join_binding_columns(&join.left, context, columns, depth + 1)?;
137 self.collect_join_binding_columns(&join.right, context, columns, depth + 1)?;
138 }
139 Expression::Aliased(aliased) => {
140 self.collect_join_binding_columns(&aliased.expression, context, columns, depth + 1)?
141 }
142 Expression::SubquerySource(source) => {
143 self.collect_select_binding_columns(&source.subquery, context, columns, depth + 1)?
144 }
145 Expression::CteReference(source) => {
146 if let Some((cte_columns, _, _)) =
147 context.get_cte_by_lower(&source.name.value_lower)
148 {
149 for column in cte_columns.iter() {
150 add_bound_join_column(columns, column);
151 }
152 }
153 }
154 Expression::FunctionTableSource(source) => {
155 if source.column_aliases.is_empty() {
156 if let Some(function) = self
157 .source_binding_functions()
158 .get_tvf(source.function.value.as_str())
159 {
160 for column in function.column_names() {
161 add_bound_join_column(columns, &column);
162 }
163 }
164 } else {
165 for column in &source.column_aliases {
166 add_bound_join_column(columns, &column.value);
167 }
168 }
169 }
170 Expression::ValuesSource(source) => {
171 if source.column_aliases.is_empty() {
172 if let Some(first_row) = source.rows.first() {
173 for index in 0..first_row.len() {
174 add_bound_join_column(columns, &format!("column{}", index + 1));
175 }
176 }
177 } else {
178 for column in &source.column_aliases {
179 add_bound_join_column(columns, &column.value);
180 }
181 }
182 }
183 _ => {}
184 }
185 Ok(())
186 }
187
188 fn validate_join_statement_bindings(
189 &self,
190 statement: &SelectStatement,
191 join_source: &JoinTableSource,
192 context: &ExecutionContext,
193 ) -> Result<()> {
194 let mut left_columns = FxHashMap::default();
195 let mut right_columns = FxHashMap::default();
196 self.collect_join_binding_columns(&join_source.left, context, &mut left_columns, 0)?;
197 self.collect_join_binding_columns(&join_source.right, context, &mut right_columns, 0)?;
198 validate_join_output_bindings(statement, join_source, &left_columns, &right_columns)
199 }
200}
201
202impl<T: SourceBindingHost + ?Sized> SourceBindingExt for T {}
203
204fn add_bound_join_column(columns: &mut FxHashMap<String, usize>, name: &str) {
205 let base = extract_base_column_name(name).to_lowercase();
206 let count = columns.entry(base).or_insert(0);
207 *count = count.saturating_add(1);
208}
209
210#[doc(hidden)]
211pub fn collect_unqualified_join_columns(expression: &Expression, columns: &mut FxHashSet<String>) {
212 match expression {
213 Expression::Identifier(identifier) => {
214 columns.insert(identifier.value_lower.to_string());
215 }
216 Expression::Infix(infix) => {
217 collect_unqualified_join_columns(&infix.left, columns);
218 collect_unqualified_join_columns(&infix.right, columns);
219 }
220 Expression::Prefix(prefix) => {
221 collect_unqualified_join_columns(&prefix.right, columns);
222 }
223 Expression::In(value) => {
224 collect_unqualified_join_columns(&value.left, columns);
225 match value.right.as_ref() {
226 Expression::ExpressionList(list) => {
227 for expression in &list.expressions {
228 collect_unqualified_join_columns(expression, columns);
229 }
230 }
231 Expression::List(list) => {
232 for expression in &list.elements {
233 collect_unqualified_join_columns(expression, columns);
234 }
235 }
236 other => collect_unqualified_join_columns(other, columns),
237 }
238 }
239 Expression::Between(value) => {
240 collect_unqualified_join_columns(&value.expr, columns);
241 collect_unqualified_join_columns(&value.lower, columns);
242 collect_unqualified_join_columns(&value.upper, columns);
243 }
244 Expression::Like(value) => {
245 collect_unqualified_join_columns(&value.left, columns);
246 collect_unqualified_join_columns(&value.pattern, columns);
247 if let Some(escape) = &value.escape {
248 collect_unqualified_join_columns(escape, columns);
249 }
250 }
251 Expression::FunctionCall(function) => {
252 for argument in &function.arguments {
253 collect_unqualified_join_columns(argument, columns);
254 }
255 if let Some(filter) = &function.filter {
256 collect_unqualified_join_columns(filter, columns);
257 }
258 }
259 Expression::Aliased(aliased) => {
260 collect_unqualified_join_columns(&aliased.expression, columns);
261 }
262 Expression::Cast(cast) => collect_unqualified_join_columns(&cast.expr, columns),
263 Expression::Case(value) => {
264 if let Some(expression) = &value.value {
265 collect_unqualified_join_columns(expression, columns);
266 }
267 for clause in &value.when_clauses {
268 collect_unqualified_join_columns(&clause.condition, columns);
269 collect_unqualified_join_columns(&clause.then_result, columns);
270 }
271 if let Some(expression) = &value.else_value {
272 collect_unqualified_join_columns(expression, columns);
273 }
274 }
275 _ => {}
276 }
277}
278
279fn validate_unqualified_join_expression(
280 expression: &Expression,
281 join_source: &JoinTableSource,
282 left_columns: &FxHashMap<String, usize>,
283 right_columns: &FxHashMap<String, usize>,
284) -> Result<()> {
285 let mut columns = FxHashSet::default();
286 collect_unqualified_join_columns(expression, &mut columns);
287 for column in columns {
288 let matches = left_columns
289 .get(&column)
290 .copied()
291 .unwrap_or(0)
292 .saturating_add(right_columns.get(&column).copied().unwrap_or(0));
293 if matches > 1
294 && !join_column_is_coalesced(join_source, &column, left_columns, right_columns)
295 {
296 return Err(Error::AmbiguousColumn(column));
297 }
298 }
299 Ok(())
300}
301
302#[doc(hidden)]
303pub fn join_column_is_coalesced(
304 join_source: &JoinTableSource,
305 column: &str,
306 left_columns: &FxHashMap<String, usize>,
307 right_columns: &FxHashMap<String, usize>,
308) -> bool {
309 if left_columns.get(column).copied() != Some(1) || right_columns.get(column).copied() != Some(1)
310 {
311 return false;
312 }
313 join_source
314 .join_type
315 .to_ascii_uppercase()
316 .contains("NATURAL")
317 || join_source
318 .using_columns
319 .iter()
320 .any(|using_column| using_column.value_lower.eq_ignore_ascii_case(column))
321}
322
323fn validate_join_output_bindings(
324 statement: &SelectStatement,
325 join_source: &JoinTableSource,
326 left_columns: &FxHashMap<String, usize>,
327 right_columns: &FxHashMap<String, usize>,
328) -> Result<()> {
329 for expression in &statement.columns {
330 validate_unqualified_join_expression(expression, join_source, left_columns, right_columns)?;
331 }
332
333 let mut output_labels = FxHashMap::default();
334 for column in &statement.columns {
335 let label = match column {
336 Expression::Aliased(aliased) => Some(aliased.alias.value_lower.as_str()),
337 Expression::Identifier(identifier) => Some(identifier.value_lower.as_str()),
338 Expression::QualifiedIdentifier(identifier) => {
339 Some(identifier.name.value_lower.as_str())
340 }
341 _ => None,
342 };
343 if let Some(label) = label {
344 let count = output_labels.entry(label.to_string()).or_insert(0usize);
345 *count = count.saturating_add(1);
346 }
347 }
348
349 for order in &statement.order_by {
350 let unique_output_label = match &order.expression {
351 Expression::Identifier(identifier) => {
352 output_labels.get(identifier.value_lower.as_str()).copied() == Some(1)
353 }
354 _ => false,
355 };
356 if !unique_output_label {
357 validate_unqualified_join_expression(
358 &order.expression,
359 join_source,
360 left_columns,
361 right_columns,
362 )?;
363 }
364 }
365 Ok(())
366}
367
368#[cfg(test)]
369mod tests {
370 use super::*;
371
372 fn select_and_join(sql: &str) -> (SelectStatement, JoinTableSource) {
373 let mut statements = radixdb_sql::parse_sql(sql).unwrap();
374 let Statement::Select(select) = statements.pop().unwrap() else {
375 panic!("expected SELECT");
376 };
377 let Some(Expression::JoinSource(join)) = select.table_expr.as_deref() else {
378 panic!("expected JOIN source");
379 };
380 (select.clone(), join.as_ref().clone())
381 }
382
383 #[test]
384 fn rejects_ambiguous_unqualified_projection() {
385 let (select, join) =
386 select_and_join("SELECT id FROM left_t JOIN right_t ON left_t.id = right_t.id");
387 let left = FxHashMap::from_iter([("id".to_string(), 1)]);
388 let right = FxHashMap::from_iter([("id".to_string(), 1)]);
389
390 assert!(matches!(
391 validate_join_output_bindings(&select, &join, &left, &right),
392 Err(Error::AmbiguousColumn(column)) if column == "id"
393 ));
394 }
395
396 #[test]
397 fn unique_select_alias_disambiguates_order_by() {
398 let (select, join) = select_and_join(
399 "SELECT left_t.id AS selected_id FROM left_t JOIN right_t ON left_t.id = right_t.id ORDER BY selected_id",
400 );
401 let left = FxHashMap::from_iter([("id".to_string(), 1)]);
402 let right = FxHashMap::from_iter([("id".to_string(), 1)]);
403
404 validate_join_output_bindings(&select, &join, &left, &right).unwrap();
405 }
406}