fix(sql): resolve CTE references by query scope
This commit is contained in:
parent
f7741a55f4
commit
72083b476f
|
|
@ -1,4 +1,5 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashSet;
|
||||
use std::sync::LazyLock;
|
||||
|
||||
use regex::Regex;
|
||||
|
|
@ -79,6 +80,7 @@ struct Analyzer {
|
|||
columns: Vec<SqlColumnReference>,
|
||||
scopes: Vec<SqlReferenceScope>,
|
||||
scope_stack: Vec<usize>,
|
||||
cte_scope_stack: Vec<HashSet<String>>,
|
||||
next_scope_id: usize,
|
||||
is_sqlserver: bool,
|
||||
}
|
||||
|
|
@ -221,7 +223,9 @@ impl Analyzer {
|
|||
self.next_scope_id += 1;
|
||||
self.scopes.push(SqlReferenceScope { id: scope_id, parent_id });
|
||||
self.scope_stack.push(scope_id);
|
||||
self.cte_scope_stack.push(HashSet::new());
|
||||
self.visit_query(query);
|
||||
self.cte_scope_stack.pop();
|
||||
self.scope_stack.pop();
|
||||
}
|
||||
|
||||
|
|
@ -233,9 +237,37 @@ impl Analyzer {
|
|||
self.scope_stack.last().copied()
|
||||
}
|
||||
|
||||
fn add_visible_cte(&mut self, ident: &Ident) {
|
||||
let key = self.cte_name_key(ident);
|
||||
if let Some(visible_ctes) = self.cte_scope_stack.last_mut() {
|
||||
visible_ctes.insert(key);
|
||||
}
|
||||
}
|
||||
|
||||
fn is_visible_cte(&self, name: &ObjectName) -> bool {
|
||||
if name.0.len() != 1 {
|
||||
return false;
|
||||
}
|
||||
let Some(ident) = name.0.first().and_then(ObjectNamePart::as_ident) else {
|
||||
return false;
|
||||
};
|
||||
let key = self.cte_name_key(ident);
|
||||
self.cte_scope_stack.iter().rev().any(|visible_ctes| visible_ctes.contains(&key))
|
||||
}
|
||||
|
||||
fn cte_name_key(&self, ident: &Ident) -> String {
|
||||
if self.is_sqlserver || ident.quote_style.is_none() {
|
||||
ident.value.to_ascii_lowercase()
|
||||
} else {
|
||||
ident.value.clone()
|
||||
}
|
||||
}
|
||||
|
||||
fn visit_query(&mut self, query: &Query) {
|
||||
if let Some(with) = &query.with {
|
||||
for cte in &with.cte_tables {
|
||||
// Add each name before its body: recursive/self and earlier CTEs are visible, later CTEs are not.
|
||||
self.add_visible_cte(&cte.alias.name);
|
||||
self.visit_child_query(&cte.query);
|
||||
}
|
||||
}
|
||||
|
|
@ -362,7 +394,8 @@ impl Analyzer {
|
|||
fn visit_table_factor(&mut self, factor: &TableFactor) {
|
||||
match factor {
|
||||
TableFactor::Table { name, alias, args, .. } => {
|
||||
if args.is_none() {
|
||||
// Qualified names remain physical objects even when their final component matches a visible CTE.
|
||||
if args.is_none() && !self.is_visible_cte(name) {
|
||||
if let Some(table) = table_reference_from_name(
|
||||
name,
|
||||
alias.as_ref().map(|a| a.name.value.clone()),
|
||||
|
|
@ -426,7 +459,7 @@ impl Analyzer {
|
|||
}
|
||||
Expr::InSubquery { expr, subquery, .. } => {
|
||||
self.visit_expr(expr);
|
||||
self.visit_query(subquery);
|
||||
self.visit_child_query(subquery);
|
||||
}
|
||||
Expr::InUnnest { expr, array_expr, .. } => {
|
||||
self.visit_expr(expr);
|
||||
|
|
|
|||
|
|
@ -33,6 +33,68 @@ fn extracts_nested_query_scopes_for_correlated_subqueries() {
|
|||
assert_eq!(columns, vec![(Some("aa"), "house_id", 0), (None, "HOUSE_ID", 1), (Some("aa"), "HOUSE_ID", 1)]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extracts_in_subquery_in_a_child_scope() {
|
||||
let sql = "select u.id from users u where u.id in (select o.user_id from orders o)";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
let tables: Vec<_> = analysis.tables.iter().map(|table| (table.name.as_str(), table.scope_id)).collect();
|
||||
assert_eq!(tables, vec![("users", 0), ("orders", 1)]);
|
||||
|
||||
let scopes: Vec<_> = analysis.scopes.iter().map(|scope| (scope.id, scope.parent_id)).collect();
|
||||
assert_eq!(scopes, vec![(0, None), (1, Some(0))]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlserver_single_cte_is_not_reported_as_a_physical_table() {
|
||||
let sql = "WITH SalesCte AS (SELECT * FROM dbo.sales) SELECT * FROM salescte";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
let tables: Vec<_> = analysis.tables.iter().map(|table| (table.schema.as_deref(), table.name.as_str())).collect();
|
||||
assert_eq!(tables, vec![(Some("dbo"), "sales")]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlserver_recursive_cte_can_reference_itself() {
|
||||
let sql = "WITH numbers AS (SELECT 1 AS value UNION ALL SELECT value + 1 FROM numbers WHERE value < 10) SELECT * FROM numbers";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
assert!(analysis.tables.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlserver_ctes_only_hide_names_after_they_are_declared() {
|
||||
let sql =
|
||||
"WITH first_cte AS (SELECT * FROM later_cte), later_cte AS (SELECT * FROM first_cte) SELECT * FROM later_cte";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
let tables: Vec<_> = analysis.tables.iter().map(|table| table.name.as_str()).collect();
|
||||
assert_eq!(tables, vec!["later_cte"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sqlserver_qualified_table_is_not_hidden_by_same_named_cte() {
|
||||
let sql = "WITH employees AS (SELECT * FROM dbo.employees) SELECT * FROM employees";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
assert_eq!(analysis.tables.len(), 1);
|
||||
assert_eq!(analysis.tables[0].schema.as_deref(), Some("dbo"));
|
||||
assert_eq!(analysis.tables[0].name, "employees");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nested_queries_inherit_and_shadow_cte_names() {
|
||||
let sql = "WITH source AS (SELECT * FROM dbo.outer_source) SELECT * FROM source WHERE EXISTS (WITH source AS (SELECT * FROM dbo.inner_source) SELECT * FROM source) AND EXISTS (SELECT * FROM source)";
|
||||
let analysis = analyze_sql_references(sql, Some("sqlserver")).unwrap();
|
||||
|
||||
let tables: Vec<_> =
|
||||
analysis.tables.iter().map(|table| (table.schema.as_deref(), table.name.as_str(), table.scope_id)).collect();
|
||||
assert_eq!(tables, vec![(Some("dbo"), "outer_source", 1), (Some("dbo"), "inner_source", 3)]);
|
||||
|
||||
let scopes: Vec<_> = analysis.scopes.iter().map(|scope| (scope.id, scope.parent_id)).collect();
|
||||
assert_eq!(scopes, vec![(0, None), (1, Some(0)), (2, Some(0)), (3, Some(2)), (4, Some(0))]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extracts_unqualified_columns_from_single_table_select() {
|
||||
let analysis = analyze_sql_references("select missing, id from users", Some("postgres")).unwrap();
|
||||
|
|
|
|||
Loading…
Reference in New Issue