use sqlparser::ast::{
BinaryOperator, Distinct, Expr, JoinConstraint, JoinOperator, Select, SelectItem, SetExpr,
SetOperator, Statement, TableFactor, TableWithJoins,
};
use sqlparser::dialect::PostgreSqlDialect;
use sqlparser::parser::Parser;
use std::collections::{HashMap, HashSet, VecDeque};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct JoinPath {
pub source_table: String,
pub initial_col: String,
pub steps: Vec<JoinStep>,
pub root_join_col: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct JoinStep {
pub table_name: String,
pub lookup_col: String,
pub carry_col: String,
}
#[derive(Debug, Clone)]
struct JoinEdge {
left_table: String,
left_col: String,
right_table: String,
right_col: String,
}
#[derive(Debug, Clone)]
struct ResolvedCte {
base_table: String,
col_map: HashMap<String, String>,
}
#[derive(Debug)]
struct JoinGraph {
tables: HashSet<String>,
edges: Vec<JoinEdge>,
aliases: HashMap<String, String>,
cte_refs: HashMap<String, ResolvedCte>,
}
impl JoinGraph {
fn new() -> Self {
Self {
tables: HashSet::new(),
edges: Vec::new(),
aliases: HashMap::new(),
cte_refs: HashMap::new(),
}
}
fn add_table(&mut self, name: &str, alias: Option<&str>) {
self.tables.insert(name.to_string());
if let Some(a) = alias {
self.aliases.insert(a.to_string(), name.to_string());
}
}
fn resolve(&self, name: &str) -> String {
self.aliases
.get(name)
.cloned()
.unwrap_or_else(|| name.to_string())
}
fn add_edge(&mut self, edge: JoinEdge) {
self.edges.push(edge);
}
fn edges_between(&self, a: &str, b: &str) -> Vec<&JoinEdge> {
self.edges
.iter()
.filter(|e| {
(e.left_table == a && e.right_table == b)
|| (e.left_table == b && e.right_table == a)
})
.collect()
}
fn neighbors(&self, table: &str) -> HashSet<String> {
let mut result = HashSet::new();
for edge in &self.edges {
if edge.left_table == table {
result.insert(edge.right_table.clone());
} else if edge.right_table == table {
result.insert(edge.left_table.clone());
}
}
result
}
}
pub fn extract_join_paths(select_sql: &str, root_table: &str) -> Result<Vec<JoinPath>, String> {
let dialect = PostgreSqlDialect {};
let stmts = Parser::new(&dialect)
.try_with_sql(select_sql)
.map_err(|e| format!("SQL init error: {e}"))?
.parse_statements()
.map_err(|e| format!("SQL parse error: {e}"))?;
let stmt = stmts
.into_iter()
.next()
.ok_or_else(|| "Empty SQL statement".to_string())?;
let query = match stmt {
Statement::Query(q) => q,
_ => return Err("Only SELECT queries are supported".to_string()),
};
let entity = root_table.strip_prefix("tb_").unwrap_or(root_table);
let pk_col = format!("pk_{entity}");
let ctes = match &query.with {
Some(with) if !with.recursive => resolve_ctes(with),
_ => HashMap::new(),
};
let pk_position = leftmost_select(&query.body).and_then(|s| pk_output_position(s, &pk_col));
let mut paths = Vec::new();
collect_paths_from_setexpr(
&query.body,
root_table,
&pk_col,
pk_position,
false,
&ctes,
&mut paths,
)?;
Ok(dedup_join_paths(paths))
}
fn resolve_ctes(with: &sqlparser::ast::With) -> HashMap<String, ResolvedCte> {
let mut resolved: HashMap<String, ResolvedCte> = HashMap::new();
for cte in &with.cte_tables {
let name = cte.alias.name.value.clone();
if let Some(rc) = resolve_single_cte(cte, &resolved) {
resolved.insert(name, rc);
}
}
resolved
}
fn resolve_single_cte(
cte: &sqlparser::ast::Cte,
already_resolved: &HashMap<String, ResolvedCte>,
) -> Option<ResolvedCte> {
let SetExpr::Select(select) = &*cte.query.body else {
return None; };
if select.from.len() != 1 {
return None;
}
let twj = &select.from[0];
if !twj.joins.is_empty() {
return None; }
let (from_name, _) = extract_table_info(&twj.relation).ok()?;
if already_resolved.contains_key(&from_name) {
return None; }
let declared = &cte.alias.columns;
let mut col_map = HashMap::new();
for (i, item) in select.projection.iter().enumerate() {
let (out_name, src_col) = match item {
SelectItem::UnnamedExpr(e) => match expr_bare_column(e) {
Some(col) => (col.clone(), col),
None => continue, },
SelectItem::ExprWithAlias { expr, alias } => match expr_bare_column(expr) {
Some(col) => (alias.value.clone(), col),
None => continue,
},
_ => continue,
};
let out_name = declared.get(i).map_or(out_name, |ident| ident.value.clone());
col_map.insert(out_name, src_col);
}
Some(ResolvedCte {
base_table: from_name,
col_map,
})
}
pub fn has_recursive_cte(select_sql: &str) -> bool {
let dialect = PostgreSqlDialect {};
let Ok(mut parser) = Parser::new(&dialect).try_with_sql(select_sql) else {
return false;
};
let Ok(stmts) = parser.parse_statements() else {
return false;
};
matches!(
stmts.into_iter().next(),
Some(Statement::Query(q)) if q.with.as_ref().is_some_and(|w| w.recursive)
)
}
fn collect_paths_from_setexpr(
body: &SetExpr,
root_table: &str,
pk_col: &str,
pk_position: Option<usize>,
in_set_op: bool,
ctes: &HashMap<String, ResolvedCte>,
out: &mut Vec<JoinPath>,
) -> Result<(), String> {
match body {
SetExpr::Select(select) => {
out.extend(paths_for_select(
select,
root_table,
pk_col,
pk_position,
in_set_op,
ctes,
)?);
Ok(())
}
SetExpr::SetOperation {
left, right, op, ..
} => {
if !matches!(op, SetOperator::Union) {
return Err(format!(
"{op:?} set operations are not supported for cascade paths"
));
}
collect_paths_from_setexpr(left, root_table, pk_col, pk_position, true, ctes, out)?;
collect_paths_from_setexpr(right, root_table, pk_col, pk_position, true, ctes, out)?;
Ok(())
}
SetExpr::Query(q) => {
collect_paths_from_setexpr(&q.body, root_table, pk_col, pk_position, in_set_op, ctes, out)
}
_ => Err("Unsupported query body for cascade paths".to_string()),
}
}
fn paths_for_select(
select: &Select,
root_table: &str,
pk_col: &str,
pk_position: Option<usize>,
in_set_op: bool,
ctes: &HashMap<String, ResolvedCte>,
) -> Result<Vec<JoinPath>, String> {
if select.from.is_empty() {
return Ok(vec![]);
}
let mut graph = JoinGraph::new();
for table_with_joins in &select.from {
build_graph_from_table_with_joins(table_with_joins, &mut graph, ctes)?;
}
if let Some(where_expr) = &select.selection {
extract_implicit_joins(where_expr, &mut graph);
}
if graph.tables.contains(root_table) {
return Ok(build_paths_to_root(&graph, root_table));
}
if !in_set_op {
return Ok(vec![]);
}
let Some(pk_pos) = pk_position else {
return Ok(vec![]);
};
let Some((pk_table, pk_source_col)) = branch_pk_provider(select, pk_pos, &graph) else {
return Ok(vec![]);
};
graph.add_table(root_table, None);
graph.add_edge(JoinEdge {
left_table: pk_table,
left_col: pk_source_col,
right_table: root_table.to_string(),
right_col: pk_col.to_string(),
});
Ok(build_paths_to_root(&graph, root_table))
}
fn build_paths_to_root(graph: &JoinGraph, root_table: &str) -> Vec<JoinPath> {
graph
.tables
.iter()
.filter(|t| *t != root_table)
.filter_map(|leaf| build_path(graph, leaf, root_table))
.collect()
}
fn branch_pk_provider(
select: &Select,
pk_position: usize,
graph: &JoinGraph,
) -> Option<(String, String)> {
let expr = match select.projection.get(pk_position)? {
SelectItem::UnnamedExpr(e) => e,
SelectItem::ExprWithAlias { expr, .. } => expr,
_ => return None,
};
match expr {
Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
let table = graph.resolve(&parts[0].value);
Some((table, parts[1].value.clone()))
}
Expr::Identifier(ident) => {
if graph.tables.len() == 1 {
graph
.tables
.iter()
.next()
.map(|t| (t.clone(), ident.value.clone()))
} else {
None
}
}
_ => None,
}
}
fn leftmost_select(body: &SetExpr) -> Option<&Select> {
match body {
SetExpr::Select(s) => Some(s),
SetExpr::SetOperation { left, .. } => leftmost_select(left),
SetExpr::Query(q) => leftmost_select(&q.body),
_ => None,
}
}
fn pk_output_position(select: &Select, pk_col: &str) -> Option<usize> {
select.projection.iter().position(|item| {
let name = match item {
SelectItem::ExprWithAlias { alias, .. } => Some(alias.value.clone()),
SelectItem::UnnamedExpr(e) => expr_bare_column(e),
_ => None,
};
name.as_deref() == Some(pk_col)
})
}
fn dedup_join_paths(paths: Vec<JoinPath>) -> Vec<JoinPath> {
let mut unique: Vec<JoinPath> = Vec::with_capacity(paths.len());
for p in paths {
if !unique.contains(&p) {
unique.push(p);
}
}
unique
}
fn build_graph_from_table_with_joins(
twj: &TableWithJoins,
graph: &mut JoinGraph,
ctes: &HashMap<String, ResolvedCte>,
) -> Result<(), String> {
let (root_name, root_alias) = extract_table_info(&twj.relation)?;
add_table_or_cte(&root_name, root_alias.as_deref(), ctes, graph);
for join in &twj.joins {
let (right_name, right_alias) = extract_table_info(&join.relation)?;
add_table_or_cte(&right_name, right_alias.as_deref(), ctes, graph);
let constraint = match &join.join_operator {
JoinOperator::Inner(c)
| JoinOperator::LeftOuter(c)
| JoinOperator::RightOuter(c)
| JoinOperator::FullOuter(c) => c,
_ => continue, };
match constraint {
JoinConstraint::On(expr) => {
extract_equalities(expr, graph);
}
JoinConstraint::Using(cols) => {
for col in cols {
let col_name = col.value.clone();
graph.add_edge(JoinEdge {
left_table: root_name.clone(),
left_col: col_name.clone(),
right_table: right_name.clone(),
right_col: col_name,
});
}
}
JoinConstraint::Natural | JoinConstraint::None => {}
}
}
Ok(())
}
fn add_table_or_cte(
name: &str,
alias: Option<&str>,
ctes: &HashMap<String, ResolvedCte>,
graph: &mut JoinGraph,
) {
if let Some(resolved) = ctes.get(name) {
graph.tables.insert(resolved.base_table.clone());
graph
.aliases
.insert(name.to_string(), resolved.base_table.clone());
graph.cte_refs.insert(name.to_string(), resolved.clone());
if let Some(a) = alias {
graph
.aliases
.insert(a.to_string(), resolved.base_table.clone());
graph.cte_refs.insert(a.to_string(), resolved.clone());
}
} else {
graph.add_table(name, alias);
}
}
fn extract_table_info(factor: &TableFactor) -> Result<(String, Option<String>), String> {
match factor {
TableFactor::Table { name, alias, .. } => {
let table_name = name
.0
.last()
.map(|i| i.value.clone())
.ok_or("Invalid table name")?;
let alias_name = alias.as_ref().map(|a| a.name.value.clone());
Ok((table_name, alias_name))
}
_ => {
Err("Subqueries and derived tables not supported in cascade path analysis".to_string())
}
}
}
fn extract_equalities(expr: &Expr, graph: &mut JoinGraph) {
match expr {
Expr::BinaryOp { left, op, right } => match op {
BinaryOperator::Eq => {
if let (Some(left_ref), Some(right_ref)) =
(extract_col_ref(left, graph), extract_col_ref(right, graph))
&& left_ref.0 != right_ref.0
{
graph.add_edge(JoinEdge {
left_table: left_ref.0,
left_col: left_ref.1,
right_table: right_ref.0,
right_col: right_ref.1,
});
}
}
BinaryOperator::And => {
extract_equalities(left, graph);
extract_equalities(right, graph);
}
_ => {}
},
Expr::Nested(inner) => extract_equalities(inner, graph),
_ => {}
}
}
fn extract_col_ref(expr: &Expr, graph: &JoinGraph) -> Option<(String, String)> {
match expr {
Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
let qualifier = &parts[0].value;
let column = &parts[1].value;
if let Some(cte) = graph.cte_refs.get(qualifier) {
let base_col = cte.col_map.get(column)?;
return Some((cte.base_table.clone(), base_col.clone()));
}
Some((graph.resolve(qualifier), column.clone()))
}
_ => None,
}
}
fn extract_implicit_joins(expr: &Expr, graph: &mut JoinGraph) {
extract_equalities(expr, graph);
}
fn build_path(graph: &JoinGraph, leaf: &str, root: &str) -> Option<JoinPath> {
let mut queue = VecDeque::new();
let mut visited = HashSet::new();
let mut parent_map: HashMap<String, String> = HashMap::new();
queue.push_back(leaf.to_string());
visited.insert(leaf.to_string());
while let Some(current) = queue.pop_front() {
if current == root {
break;
}
for neighbor in graph.neighbors(¤t) {
if !visited.contains(&neighbor) {
visited.insert(neighbor.clone());
parent_map.insert(neighbor.clone(), current.clone());
queue.push_back(neighbor);
}
}
}
if !visited.contains(root) {
return None; }
let mut chain = vec![root.to_string()];
let mut current = root.to_string();
loop {
if current == leaf {
break;
}
let prev = parent_map.get(¤t)?;
chain.push(prev.clone());
current = prev.clone();
}
if chain.len() < 2 {
return None;
}
chain.reverse();
let first_edge = graph.edges_between(&chain[0], &chain[1]);
let first_edge = first_edge.first()?;
let initial_col = if first_edge.left_table == chain[0] {
first_edge.left_col.clone()
} else {
first_edge.right_col.clone()
};
let mut steps = Vec::new();
for i in 1..chain.len() - 1 {
let table = &chain[i];
let incoming_edge = graph.edges_between(&chain[i - 1], table);
let incoming = incoming_edge.first()?;
let lookup_col = if incoming.left_table == *table {
incoming.left_col.clone()
} else {
incoming.right_col.clone()
};
let outgoing_edge = graph.edges_between(table, &chain[i + 1]);
let outgoing = outgoing_edge.first()?;
let carry_col = if outgoing.left_table == *table {
outgoing.left_col.clone()
} else {
outgoing.right_col.clone()
};
steps.push(JoinStep {
table_name: table.clone(),
lookup_col,
carry_col,
});
}
let last_idx = chain.len() - 1;
let root_edge = graph.edges_between(&chain[last_idx - 1], &chain[last_idx]);
let root_edge = root_edge.first()?;
let root_join_col = if root_edge.left_table == chain[last_idx] {
root_edge.left_col.clone()
} else {
root_edge.right_col.clone()
};
Some(JoinPath {
source_table: chain[0].clone(),
initial_col,
steps,
root_join_col,
})
}
pub fn extract_distinct_on_output_keys(select_sql: &str) -> Result<Vec<String>, String> {
let dialect = PostgreSqlDialect {};
let stmts = Parser::new(&dialect)
.try_with_sql(select_sql)
.map_err(|e| format!("SQL init error: {e}"))?
.parse_statements()
.map_err(|e| format!("SQL parse error: {e}"))?;
let Some(Statement::Query(query)) = stmts.into_iter().next() else {
return Ok(vec![]);
};
let SetExpr::Select(select) = *query.body else {
return Ok(vec![]);
};
let on_exprs = match &select.distinct {
Some(Distinct::On(exprs)) => exprs,
_ => return Ok(vec![]), };
let mut output_keys = Vec::with_capacity(on_exprs.len());
for on_expr in on_exprs {
let name = resolve_projection_output(on_expr, &select.projection).ok_or_else(|| {
format!(
"DISTINCT ON key `{on_expr}` is not projected under a resolvable output \
column; project it explicitly (e.g. `{on_expr} AS pk_<entity>`)"
)
})?;
output_keys.push(name);
}
Ok(output_keys)
}
fn resolve_projection_output(on_expr: &Expr, projection: &[SelectItem]) -> Option<String> {
let on_col = expr_bare_column(on_expr);
let matches =
|expr: &Expr| expr == on_expr || (on_col.is_some() && expr_bare_column(expr) == on_col);
for item in projection {
match item {
SelectItem::ExprWithAlias { expr, alias } if matches(expr) => {
return Some(alias.value.clone());
}
SelectItem::UnnamedExpr(expr) if matches(expr) => {
return expr_bare_column(expr);
}
_ => {} }
}
None
}
fn expr_bare_column(expr: &Expr) -> Option<String> {
match expr {
Expr::Identifier(ident) => Some(ident.value.clone()),
Expr::CompoundIdentifier(parts) => parts.last().map(|i| i.value.clone()),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_single_join() {
let sql = "SELECT o.pk_order, g.name \
FROM tb_order o \
JOIN tb_group g ON g.fk_order = o.pk_order";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 1);
let path = &paths[0];
assert_eq!(path.source_table, "tb_group");
assert_eq!(path.initial_col, "fk_order");
assert!(
path.steps.is_empty(),
"single hop should have no intermediate steps"
);
assert_eq!(
path.root_join_col, "pk_order",
"root_join_col should be the root's PK (child→parent FK)"
);
}
#[test]
fn test_two_hop_chain() {
let sql = "SELECT o.pk_order, g.name, i.value \
FROM tb_order o \
JOIN tb_group g ON g.fk_order = o.pk_order \
JOIN tb_item i ON i.fk_group = g.pk_group";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 2);
let group_path = paths.iter().find(|p| p.source_table == "tb_group").unwrap();
assert_eq!(group_path.initial_col, "fk_order");
assert!(group_path.steps.is_empty());
assert_eq!(group_path.root_join_col, "pk_order");
let item_path = paths.iter().find(|p| p.source_table == "tb_item").unwrap();
assert_eq!(item_path.initial_col, "fk_group");
assert_eq!(item_path.steps.len(), 1);
let hop = &item_path.steps[0];
assert_eq!(hop.table_name, "tb_group");
assert_eq!(hop.lookup_col, "pk_group");
assert_eq!(hop.carry_col, "fk_order");
}
#[test]
fn test_three_hop_chain() {
let sql = "SELECT o.pk_order, g.name, i.value, d.detail \
FROM tb_order o \
JOIN tb_group g ON g.fk_order = o.pk_order \
JOIN tb_item i ON i.fk_group = g.pk_group \
JOIN tb_detail d ON d.fk_item = i.pk_item";
let paths = extract_join_paths(sql, "tb_order").unwrap();
let detail_path = paths
.iter()
.find(|p| p.source_table == "tb_detail")
.unwrap();
assert_eq!(detail_path.initial_col, "fk_item");
assert_eq!(detail_path.steps.len(), 2);
assert_eq!(detail_path.steps[0].table_name, "tb_item");
assert_eq!(detail_path.steps[0].lookup_col, "pk_item");
assert_eq!(detail_path.steps[0].carry_col, "fk_group");
assert_eq!(detail_path.steps[1].table_name, "tb_group");
assert_eq!(detail_path.steps[1].lookup_col, "pk_group");
assert_eq!(detail_path.steps[1].carry_col, "fk_order");
}
#[test]
fn test_aliases_resolved() {
let sql = "SELECT 1 FROM tb_order o JOIN tb_group g ON g.fk_order = o.pk_order";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 1);
assert_eq!(paths[0].source_table, "tb_group");
assert_eq!(paths[0].initial_col, "fk_order");
}
#[test]
fn test_left_join() {
let sql = "SELECT 1 FROM tb_order o LEFT JOIN tb_group g ON g.fk_order = o.pk_order";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 1);
assert_eq!(paths[0].source_table, "tb_group");
}
#[test]
fn test_implicit_join() {
let sql = "SELECT 1 FROM tb_order o, tb_group g WHERE g.fk_order = o.pk_order";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 1);
assert_eq!(paths[0].source_table, "tb_group");
assert_eq!(paths[0].initial_col, "fk_order");
}
#[test]
fn test_subquery_returns_error() {
let sql = "SELECT 1 FROM (SELECT * FROM tb_a) sub";
let result = extract_join_paths(sql, "tb_a");
assert!(result.is_err());
}
#[test]
fn test_root_not_found_returns_empty() {
let sql = "SELECT 1 FROM tb_order o JOIN tb_group g ON g.fk_order = o.pk_order";
let paths = extract_join_paths(sql, "tb_nonexistent").unwrap();
assert!(paths.is_empty());
}
#[test]
fn test_lookup_table_join_root_holds_fk() {
let sql = "SELECT o.pk_order, o.fk_currency, c.iso_code \
FROM tb_order o \
LEFT JOIN tb_currency c ON o.fk_currency = c.pk_currency";
let paths = extract_join_paths(sql, "tb_order").unwrap();
assert_eq!(paths.len(), 1);
let path = &paths[0];
assert_eq!(path.source_table, "tb_currency");
assert_eq!(path.initial_col, "pk_currency");
assert!(
path.steps.is_empty(),
"direct lookup join has no intermediate steps"
);
assert_eq!(
path.root_join_col, "fk_currency",
"root_join_col must be the FK on the root table, not its PK"
);
}
#[test]
fn test_distinct_on_output_aliased() {
let sql = "SELECT DISTINCT ON (c.id_contract) c.id_contract AS pk_contract, c.id \
FROM tb_contract c ORDER BY c.id_contract, c.version_no DESC";
let keys = extract_distinct_on_output_keys(sql).unwrap();
assert_eq!(keys, vec!["pk_contract"]);
}
#[test]
fn test_distinct_on_output_unaliased() {
let sql = "SELECT DISTINCT ON (c.id) c.id, c.pk_contract FROM tb_contract c";
let keys = extract_distinct_on_output_keys(sql).unwrap();
assert_eq!(keys, vec!["id"]);
}
#[test]
fn test_distinct_on_output_qualifier_mismatch() {
let sql =
"SELECT DISTINCT ON (id_contract) c.id_contract AS pk_contract FROM tb_contract c";
let keys = extract_distinct_on_output_keys(sql).unwrap();
assert_eq!(keys, vec!["pk_contract"]);
}
#[test]
fn test_distinct_on_output_with_cte_preamble() {
let sql = "WITH x AS (SELECT 1) \
SELECT DISTINCT ON (c.id_contract) c.id_contract AS pk_contract \
FROM tb_contract c";
let keys = extract_distinct_on_output_keys(sql).unwrap();
assert_eq!(keys, vec!["pk_contract"]);
}
#[test]
fn test_distinct_on_output_none_when_no_distinct_on() {
assert!(
extract_distinct_on_output_keys("SELECT a, b FROM t")
.unwrap()
.is_empty()
);
assert!(
extract_distinct_on_output_keys("SELECT DISTINCT a FROM t")
.unwrap()
.is_empty()
);
}
#[test]
fn test_distinct_on_output_unresolvable_expression_key() {
let sql = "SELECT DISTINCT ON (lower(c.code)) c.id FROM tb_contract c";
assert!(extract_distinct_on_output_keys(sql).is_err());
}
#[test]
fn test_union_cross_table_paths() {
let sql = "SELECT pk_task, id, title FROM tb_task \
UNION ALL \
SELECT pk_task, id, title FROM tb_task_archive";
let paths = extract_join_paths(sql, "tb_task").unwrap();
let archive = paths
.iter()
.find(|p| p.source_table == "tb_task_archive")
.expect("expected a cascade path for the UNION archive branch");
assert_eq!(archive.initial_col, "pk_task");
assert!(archive.steps.is_empty());
assert_eq!(archive.root_join_col, "pk_task");
}
#[test]
fn test_union_aliased_pk_position() {
let sql = "SELECT pk_task, id FROM tb_task \
UNION ALL \
SELECT arch_id, id FROM tb_task_archive";
let paths = extract_join_paths(sql, "tb_task").unwrap();
let archive = paths
.iter()
.find(|p| p.source_table == "tb_task_archive")
.expect("expected a cascade path for the aliased-position branch");
assert_eq!(archive.initial_col, "arch_id");
assert_eq!(archive.root_join_col, "pk_task");
}
#[test]
fn test_intersect_rejected() {
let sql = "SELECT pk_task FROM tb_task INTERSECT SELECT pk_task FROM tb_task_archive";
assert!(extract_join_paths(sql, "tb_task").is_err());
}
#[test]
fn test_cte_single_base_resolves() {
let sql = "WITH cust_orders AS (\
SELECT fk_customer, count(*) AS n FROM tb_order GROUP BY fk_customer\
) \
SELECT c.pk_customer, c.id FROM tb_customer c \
LEFT JOIN cust_orders co ON co.fk_customer = c.pk_customer";
let paths = extract_join_paths(sql, "tb_customer").unwrap();
let order = paths
.iter()
.find(|p| p.source_table == "tb_order")
.expect("CTE-wrapped tb_order should resolve to a cascade path");
assert_eq!(order.initial_col, "fk_customer");
assert!(order.steps.is_empty());
assert_eq!(order.root_join_col, "pk_customer");
}
#[test]
fn test_cte_computed_join_column_unresolvable() {
let sql = "WITH s AS (SELECT fk_customer, count(*) AS c FROM tb_order GROUP BY fk_customer) \
SELECT cust.pk_customer FROM tb_customer cust JOIN s ON s.c = cust.pk_customer";
let paths = extract_join_paths(sql, "tb_customer").unwrap();
assert!(
paths.iter().all(|p| p.source_table != "tb_order"),
"a join on a CTE aggregate must not produce a cascade path"
);
}
#[test]
fn test_recursive_cte_detected() {
assert!(has_recursive_cte(
"WITH RECURSIVE t AS (SELECT 1) SELECT pk_x FROM tb_x"
));
assert!(!has_recursive_cte("SELECT pk_x FROM tb_x"));
assert!(!has_recursive_cte(
"WITH t AS (SELECT 1) SELECT pk_x FROM tb_x"
));
}
}