Skip to main content

radixdb_executor/binding/
source.rs

1//! CTE, view, table-source, and JOIN column binding.
2
3use 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/// Composition state required by source binding while the root owns the
20/// concrete database object.
21#[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/// Schema-only source and column binding owned by the executor crate.
35#[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}