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 {
tables: Vec<String>,
edges: Vec<JoinEdge>,
col_map: HashMap<String, Vec<(String, String)>>,
}
type Outputs = Vec<(Option<String>, Vec<(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((tables, edges, outputs)) = resolve_body(&cte.query.body, &resolved) {
let declared = &cte.alias.columns;
let mut col_map: HashMap<String, Vec<(String, String)>> = HashMap::new();
for (i, (out_name, sources)) in outputs.into_iter().enumerate() {
let name = declared.get(i).map(|d| d.value.clone()).or(out_name);
if let Some(name) = name
&& !sources.is_empty()
{
col_map.entry(name).or_default().extend(sources);
}
}
resolved.insert(
name,
ResolvedCte {
tables,
edges,
col_map,
},
);
}
}
resolved
}
fn resolve_body(
body: &SetExpr,
ctes: &HashMap<String, ResolvedCte>,
) -> Option<(Vec<String>, Vec<JoinEdge>, Outputs)> {
match body {
SetExpr::Select(select) => {
if select.from.is_empty() {
return None;
}
let mut graph = JoinGraph::new();
for twj in &select.from {
build_graph_from_table_with_joins(twj, &mut graph, ctes).ok()?;
}
if let Some(where_expr) = &select.selection {
extract_implicit_joins(where_expr, &mut graph);
}
let single_from = (select.from.len() == 1 && select.from[0].joins.is_empty())
.then(|| extract_table_info(&select.from[0].relation).ok())
.flatten()
.map(|(name, alias)| alias.unwrap_or(name));
let outputs = select
.projection
.iter()
.map(|item| {
let (name, expr) = match item {
SelectItem::UnnamedExpr(e) => (expr_bare_column(e), e),
SelectItem::ExprWithAlias { expr, alias } => {
(Some(alias.value.clone()), expr)
}
_ => return (None, Vec::new()),
};
let sources = match expr {
Expr::Identifier(ident) => single_from
.as_deref()
.map(|q| column_sources(q, &ident.value, &graph))
.unwrap_or_default(),
e => extract_col_ref(e, &graph),
};
(name, sources)
})
.collect();
let mut tables: Vec<String> = graph.tables.into_iter().collect();
tables.sort();
Some((tables, graph.edges, outputs))
}
SetExpr::SetOperation {
left, right, op, ..
} => {
if !matches!(op, SetOperator::Union) {
return None;
}
let (mut tables, mut edges, left_out) = resolve_body(left, ctes)?;
let (right_tables, right_edges, right_out) = resolve_body(right, ctes)?;
if left_out.len() != right_out.len() {
return None;
}
tables.extend(right_tables);
tables.sort();
tables.dedup();
edges.extend(right_edges);
let outputs = left_out
.into_iter()
.zip(right_out)
.map(|((name, mut sources), (_, more))| {
sources.extend(more);
(name, sources)
})
.collect();
Some((tables, edges, outputs))
}
SetExpr::Query(q) => resolve_body(&q.body, ctes),
_ => None,
}
}
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 => {
match column_sources(&parts[0].value, &parts[1].value, graph).as_slice() {
[single] => Some(single.clone()),
_ => None,
}
}
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.extend(resolved.tables.iter().cloned());
for edge in &resolved.edges {
graph.add_edge(edge.clone());
}
graph.cte_refs.insert(name.to_string(), resolved.clone());
if let Some(a) = alias {
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 => {
let lefts = extract_col_ref(left, graph);
let rights = extract_col_ref(right, graph);
for (lt, lc) in &lefts {
for (rt, rc) in &rights {
if lt != rt {
graph.add_edge(JoinEdge {
left_table: lt.clone(),
left_col: lc.clone(),
right_table: rt.clone(),
right_col: rc.clone(),
});
}
}
}
}
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) -> Vec<(String, String)> {
match expr {
Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
column_sources(&parts[0].value, &parts[1].value, graph)
}
Expr::Nested(inner) => extract_col_ref(inner, graph),
_ => Vec::new(),
}
}
fn column_sources(qualifier: &str, column: &str, graph: &JoinGraph) -> Vec<(String, String)> {
if let Some(cte) = graph.cte_refs.get(qualifier) {
return cte.col_map.get(column).cloned().unwrap_or_default();
}
vec![(graph.resolve(qualifier), column.to_string())]
}
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 plain_select_graph(select_sql: &str) -> Option<(Box<Select>, JoinGraph)> {
let stmts = Parser::new(&PostgreSqlDialect {})
.try_with_sql(select_sql)
.ok()?
.parse_statements()
.ok()?;
let Some(Statement::Query(query)) = stmts.into_iter().next() else {
return None;
};
if query.with.is_some() {
return None;
}
let SetExpr::Select(select) = *query.body else {
return None;
};
let mut graph = JoinGraph::new();
for twj in &select.from {
build_graph_from_table_with_joins(twj, &mut graph, &HashMap::new()).ok()?;
}
if let Some(where_expr) = &select.selection {
extract_implicit_joins(where_expr, &mut graph);
}
Some((select, graph))
}
#[must_use]
pub fn table_qualifier(select_sql: &str, table: &str) -> Option<String> {
let (select, _) = plain_select_graph(select_sql)?;
let mut qualifiers = Vec::new();
for twj in &select.from {
let factors = std::iter::once(&twj.relation).chain(twj.joins.iter().map(|j| &j.relation));
for factor in factors {
let (name, alias) = extract_table_info(factor).ok()?;
if name == table {
qualifiers.push(alias.unwrap_or(name));
}
}
}
match qualifiers.as_slice() {
[one] => Some(one.clone()),
_ => None,
}
}
#[must_use]
pub fn output_column_for(select_sql: &str, table: &str, column: &str) -> Option<String> {
let (select, graph) = plain_select_graph(select_sql)?;
projected_output(&select, &graph, table, column)
}
fn projected_output(
select: &Select,
graph: &JoinGraph,
table: &str,
column: &str,
) -> Option<String> {
select.projection.iter().find_map(|item| {
let (expr, name) = match item {
SelectItem::ExprWithAlias { expr, alias } => (expr, Some(alias.value.clone())),
SelectItem::UnnamedExpr(expr) => (expr, expr_bare_column(expr)),
_ => return None,
};
let projects = match expr {
Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
graph.resolve(&parts[0].value) == table && parts[1].value == column
}
Expr::Identifier(ident) => ident.value == column,
_ => false,
};
if projects { name } else { None }
})
}
pub fn embed_lookup_columns(
select_sql: &str,
entities: &[String],
) -> Result<Vec<(String, Option<String>)>, String> {
let referenced = referenced_entities(select_sql, entities)?;
if referenced.is_empty() {
return Ok(Vec::new());
}
let parsed = plain_select_graph(select_sql);
Ok(referenced
.into_iter()
.map(|entity| {
let column = parsed
.as_ref()
.and_then(|(select, graph)| pk_join_output(select, graph, &entity));
(entity, column)
})
.collect())
}
fn referenced_entities(select_sql: &str, entities: &[String]) -> Result<Vec<String>, String> {
use sqlparser::tokenizer::{Token, Tokenizer};
let tokens = Tokenizer::new(&PostgreSqlDialect {}, select_sql)
.tokenize()
.map_err(|e| format!("SQL tokenize error: {e}"))?;
let words: HashSet<String> = tokens
.into_iter()
.filter_map(|t| match t {
Token::Word(w) if w.quote_style.is_none() => Some(w.value.to_lowercase()),
Token::Word(w) => Some(w.value),
_ => None,
})
.collect();
Ok(entities
.iter()
.filter(|e| words.contains(&format!("v_{e}")) || words.contains(&format!("tv_{e}")))
.cloned()
.collect())
}
fn pk_join_output(select: &Select, graph: &JoinGraph, entity: &str) -> Option<String> {
let pk = format!("pk_{entity}");
let relations = [format!("v_{entity}"), format!("tv_{entity}")];
let is_pk = |table: &String, col: &String| relations.contains(table) && *col == pk;
graph.edges.iter().find_map(|e| {
let (table, col) = if is_pk(&e.left_table, &e.left_col) {
(&e.right_table, &e.right_col)
} else if is_pk(&e.right_table, &e.right_col) {
(&e.left_table, &e.left_col)
} else {
return None;
};
projected_output(select, graph, table, col)
})
}
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"
));
}
fn sources(paths: &[JoinPath]) -> Vec<String> {
let mut v: Vec<String> = paths.iter().map(|p| p.source_table.clone()).collect();
v.sort();
v
}
#[test]
fn cte_chain_resolves_through_an_earlier_cte() {
let sql = "WITH a AS (SELECT item_id, label FROM tb_item_i18n WHERE locale = 'fr'), \
b AS (SELECT item_id, upper(label) AS label FROM a) \
SELECT i.pk_item, i.id, jsonb_build_object('label', b.label) AS data \
FROM tb_item i LEFT JOIN b ON b.item_id = i.pk_item";
let paths = extract_join_paths(sql, "tb_item").unwrap();
assert_eq!(sources(&paths), ["tb_item_i18n"]);
assert_eq!(paths[0].initial_col, "item_id");
assert!(paths[0].steps.is_empty());
}
#[test]
fn multi_base_cte_body_yields_a_path_per_base_table() {
let sql = "WITH l AS (SELECT t.item_id, t.label, loc.code \
FROM tb_item_i18n t JOIN tb_locale loc ON loc.pk_locale = t.fk_locale) \
SELECT i.pk_item, i.id, jsonb_build_object('label', l.label, 'code', l.code) AS data \
FROM tb_item i JOIN l ON l.item_id = i.pk_item";
let paths = extract_join_paths(sql, "tb_item").unwrap();
assert_eq!(sources(&paths), ["tb_item_i18n", "tb_locale"]);
let locale = paths
.iter()
.find(|p| p.source_table == "tb_locale")
.unwrap();
assert_eq!(locale.initial_col, "pk_locale");
assert_eq!(
locale.steps,
[JoinStep {
table_name: "tb_item_i18n".into(),
lookup_col: "fk_locale".into(),
carry_col: "item_id".into(),
}]
);
}
#[test]
fn union_bodied_cte_yields_a_path_per_branch() {
let sql = "WITH n AS (SELECT item_id, label FROM tb_item_i18n \
UNION ALL SELECT item_id, label FROM tb_item_default) \
SELECT i.pk_item, i.id, jsonb_build_object('label', n.label) AS data \
FROM tb_item i JOIN n ON n.item_id = i.pk_item";
let paths = extract_join_paths(sql, "tb_item").unwrap();
assert_eq!(sources(&paths), ["tb_item_default", "tb_item_i18n"]);
assert!(
paths
.iter()
.all(|p| p.initial_col == "item_id" && p.steps.is_empty())
);
}
#[test]
fn computed_cte_join_column_yields_no_path() {
let sql = "WITH c AS (SELECT item_id + 0 AS item_id, label FROM tb_item_i18n) \
SELECT i.pk_item, i.id, jsonb_build_object('label', c.label) AS data \
FROM tb_item i JOIN c ON c.item_id = i.pk_item";
let paths = extract_join_paths(sql, "tb_item").unwrap();
assert!(paths.is_empty(), "{paths:?}");
}
fn aggregates() -> Vec<String> {
vec!["user_summary".to_string(), "tag_count".to_string()]
}
#[test]
fn embed_lookup_through_left_join_on_own_pk() {
let sql = "SELECT u.pk_user, u.id, jsonb_build_object('s', s.data) AS data \
FROM tb_user u LEFT JOIN v_user_summary s ON s.pk_user_summary = u.pk_user";
assert_eq!(
embed_lookup_columns(sql, &aggregates()).unwrap(),
vec![("user_summary".to_string(), Some("pk_user".to_string()))]
);
}
#[test]
fn embed_lookup_uses_the_output_alias_and_tv_relation() {
let sql = "SELECT p.pk_post, p.fk_author AS author, c.data AS tags \
FROM tb_post p JOIN tv_tag_count c ON p.fk_author = c.pk_tag_count";
assert_eq!(
embed_lookup_columns(sql, &aggregates()).unwrap(),
vec![("tag_count".to_string(), Some("author".to_string()))]
);
}
#[test]
fn embed_lookup_through_where_equality() {
let sql = "SELECT u.pk_user, s.data FROM tb_user u, v_user_summary s \
WHERE u.pk_user = s.pk_user_summary";
assert_eq!(
embed_lookup_columns(sql, &aggregates()).unwrap(),
vec![("user_summary".to_string(), Some("pk_user".to_string()))]
);
}
#[test]
fn embed_lookup_none_when_join_column_not_projected() {
let sql = "SELECT p.pk_post, s.data FROM tb_post p \
JOIN v_user_summary s ON s.pk_user_summary = p.fk_user";
assert_eq!(
embed_lookup_columns(sql, &aggregates()).unwrap(),
vec![("user_summary".to_string(), None)]
);
}
#[test]
fn embed_lookup_none_inside_a_subquery() {
let sql = "SELECT u.pk_user, (SELECT s.data FROM v_user_summary s \
WHERE s.pk_user_summary = u.pk_user) AS data FROM tb_user u";
assert_eq!(
embed_lookup_columns(sql, &aggregates()).unwrap(),
vec![("user_summary".to_string(), None)]
);
}
#[test]
fn embed_lookup_ignores_unreferenced_entities() {
let sql = "SELECT u.pk_user, u.data FROM tb_user u JOIN v_user_summary_extra x \
ON x.pk_user = u.pk_user";
assert!(embed_lookup_columns(sql, &aggregates()).unwrap().is_empty());
}
#[test]
fn table_qualifier_is_the_alias_or_the_name() {
let sql = "SELECT p.pk_post FROM tb_post p JOIN tb_user u ON u.pk_user = p.fk_user";
assert_eq!(table_qualifier(sql, "tb_user"), Some("u".to_string()));
let sql = "SELECT pk_post FROM tb_post JOIN tb_user ON tb_user.pk_user = tb_post.fk_user";
assert_eq!(table_qualifier(sql, "tb_user"), Some("tb_user".to_string()));
}
#[test]
fn table_qualifier_none_when_read_twice_or_not_at_all() {
let sql = "SELECT p.pk_post FROM tb_post p JOIN tb_user a ON a.pk_user = p.fk_author \
JOIN tb_user e ON e.pk_user = p.fk_editor";
assert_eq!(table_qualifier(sql, "tb_user"), None);
assert_eq!(table_qualifier(sql, "tb_tag"), None);
}
#[test]
fn output_column_for_follows_aliases() {
let sql = "SELECT p.pk_post, p.fk_user AS author, u.name FROM tb_post p \
JOIN tb_user u ON u.pk_user = p.fk_user";
assert_eq!(
output_column_for(sql, "tb_post", "fk_user"),
Some("author".to_string())
);
assert_eq!(
output_column_for(sql, "tb_user", "name"),
Some("name".to_string())
);
assert_eq!(output_column_for(sql, "tb_user", "pk_user"), None);
}
}