use crate::query_plan::pipeline::ASTTransformer;
use crate::sql::parser::ast::{SelectItem, SelectStatement, SqlExpression};
use crate::sql::parser::walk;
use anyhow::Result;
use std::collections::HashMap;
use tracing::debug;
pub struct WhereAliasExpander {
expansions: usize,
}
impl WhereAliasExpander {
pub fn new() -> Self {
Self { expansions: 0 }
}
fn extract_aliases(select_items: &[SelectItem]) -> HashMap<String, SqlExpression> {
let mut aliases = HashMap::new();
for item in select_items {
if let SelectItem::Expression { expr, alias, .. } = item {
if !alias.is_empty() {
aliases.insert(alias.clone(), expr.clone());
debug!("Found SELECT alias: {} -> {:?}", alias, expr);
}
}
}
aliases
}
fn expand_expression(
expr: &SqlExpression,
aliases: &HashMap<String, SqlExpression>,
) -> (SqlExpression, bool) {
match expr {
SqlExpression::Column(col_ref) => {
if col_ref.table_prefix.is_none() {
if let Some(alias_expr) = aliases.get(&col_ref.name) {
debug!(
"Expanding alias '{}' in WHERE to: {:?}",
col_ref.name, alias_expr
);
return (alias_expr.clone(), true);
}
}
(expr.clone(), false)
}
SqlExpression::MethodCall {
object,
method,
args,
} => {
let mut expanded = false;
let new_args: Vec<SqlExpression> = args
.iter()
.map(|arg| {
let (new_arg, arg_expanded) = Self::expand_expression(arg, aliases);
expanded = expanded || arg_expanded;
new_arg
})
.collect();
let mut new_object = object.clone();
if let Some(SqlExpression::Column(col_ref)) = aliases.get(object) {
if col_ref.table_prefix.is_none() {
debug!(
"Expanding alias '{}' in WHERE method call to column '{}'",
object, col_ref.name
);
new_object = col_ref.name.clone();
expanded = true;
}
}
(
SqlExpression::MethodCall {
object: new_object,
method: method.clone(),
args: new_args,
},
expanded,
)
}
other => {
let mut expanded = false;
let new_expr = walk::map_children(other.clone(), |child| {
let (new_child, child_expanded) = Self::expand_expression(&child, aliases);
expanded = expanded || child_expanded;
new_child
});
(new_expr, expanded)
}
}
}
fn expand_where_clause(
&mut self,
where_clause: &mut crate::sql::parser::ast::WhereClause,
aliases: &HashMap<String, SqlExpression>,
) -> bool {
let mut any_expanded = false;
for condition in &mut where_clause.conditions {
let (new_expr, expanded) = Self::expand_expression(&condition.expr, aliases);
if expanded {
condition.expr = new_expr;
any_expanded = true;
self.expansions += 1;
}
}
any_expanded
}
}
impl Default for WhereAliasExpander {
fn default() -> Self {
Self::new()
}
}
impl ASTTransformer for WhereAliasExpander {
fn name(&self) -> &str {
"WhereAliasExpander"
}
fn description(&self) -> &str {
"Expands SELECT aliases in WHERE clauses to their full expressions"
}
fn transform(&mut self, mut stmt: SelectStatement) -> Result<SelectStatement> {
if stmt.where_clause.is_none() {
return Ok(stmt);
}
let aliases = Self::extract_aliases(&stmt.select_items);
if aliases.is_empty() {
return Ok(stmt);
}
if let Some(ref mut where_clause) = stmt.where_clause {
let expanded = self.expand_where_clause(where_clause, &aliases);
if expanded {
debug!(
"Expanded {} alias reference(s) in WHERE clause",
self.expansions
);
}
}
Ok(stmt)
}
fn begin(&mut self) -> Result<()> {
self.expansions = 0;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sql::parser::ast::{ColumnRef, Condition, QuoteStyle, WhereClause};
#[test]
fn test_extract_aliases() {
let double_a_expr = SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef {
name: "a".to_string(),
quote_style: QuoteStyle::None,
table_prefix: None,
})),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
};
let select_items = vec![SelectItem::Expression {
expr: double_a_expr.clone(),
alias: "double_a".to_string(),
leading_comments: vec![],
trailing_comment: None,
}];
let aliases = WhereAliasExpander::extract_aliases(&select_items);
assert_eq!(aliases.len(), 1);
assert!(aliases.contains_key("double_a"));
}
#[test]
fn test_expand_simple_column_reference() {
let aliases = HashMap::from([(
"double_a".to_string(),
SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted("a".to_string()))),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
},
)]);
let expr = SqlExpression::Column(ColumnRef::unquoted("double_a".to_string()));
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(changed);
assert!(matches!(expanded, SqlExpression::BinaryOp { .. }));
}
#[test]
fn test_expand_in_binary_op() {
let aliases = HashMap::from([(
"double_a".to_string(),
SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted("a".to_string()))),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
},
)]);
let expr = SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted(
"double_a".to_string(),
))),
op: ">".to_string(),
right: Box::new(SqlExpression::NumberLiteral("10".to_string())),
};
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(changed);
if let SqlExpression::BinaryOp { left, op, right } = expanded {
assert_eq!(op, ">");
assert!(matches!(left.as_ref(), SqlExpression::BinaryOp { .. }));
assert!(matches!(
right.as_ref(),
SqlExpression::NumberLiteral(s) if s == "10"
));
} else {
panic!("Expected BinaryOp");
}
}
#[test]
fn test_transform_with_no_where() {
let mut transformer = WhereAliasExpander::new();
let stmt = SelectStatement {
where_clause: None,
..Default::default()
};
let result = transformer.transform(stmt);
assert!(result.is_ok());
}
#[test]
fn test_transform_expands_alias() {
let mut transformer = WhereAliasExpander::new();
let double_a_expr = SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted("a".to_string()))),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
};
let stmt = SelectStatement {
select_items: vec![SelectItem::Expression {
expr: double_a_expr.clone(),
alias: "double_a".to_string(),
leading_comments: vec![],
trailing_comment: None,
}],
where_clause: Some(WhereClause {
conditions: vec![Condition {
expr: SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted(
"double_a".to_string(),
))),
op: ">".to_string(),
right: Box::new(SqlExpression::NumberLiteral("10".to_string())),
},
connector: None,
}],
}),
..Default::default()
};
let result = transformer.transform(stmt).unwrap();
if let Some(where_clause) = &result.where_clause {
if let SqlExpression::BinaryOp { left, .. } = &where_clause.conditions[0].expr {
assert!(matches!(left.as_ref(), SqlExpression::BinaryOp { .. }));
} else {
panic!("Expected BinaryOp in WHERE");
}
} else {
panic!("Expected WHERE clause");
}
assert_eq!(transformer.expansions, 1);
}
#[test]
fn test_expand_alias_in_method_call_receiver() {
let aliases = HashMap::from([(
"name".to_string(),
SqlExpression::Column(ColumnRef {
name: "name.common".to_string(),
quote_style: QuoteStyle::DoubleQuotes,
table_prefix: None,
}),
)]);
let expr = SqlExpression::MethodCall {
object: "name".to_string(),
method: "Contains".to_string(),
args: vec![SqlExpression::StringLiteral("united".to_string())],
};
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(changed);
match expanded {
SqlExpression::MethodCall { object, method, .. } => {
assert_eq!(object, "name.common");
assert_eq!(method, "Contains");
}
other => panic!("Expected MethodCall, got {other:?}"),
}
}
#[test]
fn test_does_not_expand_method_call_for_nonalias() {
let aliases = HashMap::from([(
"name".to_string(),
SqlExpression::Column(ColumnRef::unquoted("name.common".to_string())),
)]);
let expr = SqlExpression::MethodCall {
object: "capital".to_string(),
method: "Contains".to_string(),
args: vec![SqlExpression::StringLiteral("x".to_string())],
};
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(!changed);
assert!(matches!(
expanded,
SqlExpression::MethodCall { object, .. } if object == "capital"
));
}
#[test]
fn test_expands_alias_on_in_subquery_lhs_not_body() {
let double = SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted("price".into()))),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
};
let aliases = HashMap::from([("dbl".to_string(), double.clone())]);
let body = SelectStatement {
where_clause: Some(WhereClause {
conditions: vec![Condition {
expr: SqlExpression::Column(ColumnRef::unquoted("dbl".into())),
connector: None,
}],
}),
..Default::default()
};
let expr = SqlExpression::InSubquery {
expr: Box::new(SqlExpression::Column(ColumnRef::unquoted("dbl".into()))),
subquery: Box::new(body.clone()),
};
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(changed, "the LHS alias should have been expanded");
match expanded {
SqlExpression::InSubquery { expr, subquery } => {
assert!(matches!(expr.as_ref(), SqlExpression::BinaryOp { .. }));
let inner = &subquery.where_clause.as_ref().unwrap().conditions[0].expr;
assert!(
matches!(inner, SqlExpression::Column(c) if c.name == "dbl"),
"subquery body must not be touched (different scope), got {inner:?}"
);
}
other => panic!("expected InSubquery, got {other:?}"),
}
}
#[test]
fn test_does_not_expand_table_prefixed_columns() {
let aliases = HashMap::from([(
"double_a".to_string(),
SqlExpression::BinaryOp {
left: Box::new(SqlExpression::Column(ColumnRef::unquoted("a".to_string()))),
op: "*".to_string(),
right: Box::new(SqlExpression::NumberLiteral("2".to_string())),
},
)]);
let expr = SqlExpression::Column(ColumnRef {
name: "double_a".to_string(),
quote_style: QuoteStyle::None,
table_prefix: Some("t".to_string()),
});
let (expanded, changed) = WhereAliasExpander::expand_expression(&expr, &aliases);
assert!(!changed);
assert!(matches!(expanded, SqlExpression::Column(_)));
}
}