use sqlparser::ast::{Expr, Select, SelectItem, SetExpr, Statement, TableFactor};
use sqlparser::dialect::PostgreSqlDialect;
use sqlparser::parser::Parser;
use std::collections::HashMap;
#[derive(Debug, Default)]
struct JoinGraph {
aliases: HashMap<String, String>,
}
impl JoinGraph {
fn resolve(&self, name: &str) -> String {
self.aliases
.get(name)
.cloned()
.unwrap_or_else(|| name.to_string())
}
}
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 are not plain relations".to_string()),
}
}
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::default();
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 let Some(alias) = alias {
graph.aliases.insert(alias, name);
}
}
}
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 }
})
}
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 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);
}
}