use sqlparser::ast::{
BinaryOperator, Expr, JoinConstraint, JoinOperator, Select, SelectItem, SetExpr, Statement,
TableFactor, TableWithJoins,
};
use sqlparser::dialect::PostgreSqlDialect;
use sqlparser::parser::Parser;
use std::collections::{HashMap, HashSet};
#[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)>>,
}
#[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);
}
}
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 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 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_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 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);
}
}