fix(sql): resolve CTE references by query scope

This commit is contained in:
t8y2 2026-07-23 00:15:24 +08:00
parent f7741a55f4
commit 72083b476f
2 changed files with 97 additions and 2 deletions

View File

@ -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);

View File

@ -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();