// Port of sqlglot/optimizer/pushdown_predicates.py.
///|
fn dialect_is_a(dialect : @core.Dialect, names : Array[String]) -> Bool {
let mut d : @core.Dialect? = Some(dialect)
while d is Some(x) {
if names.contains(x.name) {
return true
}
d = x.parent
}
false
}
///|
fn sorted_strings(xs : Array[String]) -> Array[String] {
let out = xs.copy()
out.sort_by(py_str_cmp)
out
}
///|
/// `Join.on(predicate, copy=False)`
fn join_on(join : @core.Expr, predicate : @core.Expr) -> Unit {
let node = match join.arg("on") {
Some(existing) => @core.and_([existing, predicate], copy=false)
None => predicate
}
join.set("on", node)
if @core.py_upper(join.text("kind")) == "CROSS" {
join.set("kind", @core.null_arg)
}
}
///|
/// Rewrite the AST to pushdown predicates in FROMS and JOINS.
pub fn pushdown_predicates(
expression : @core.Expr,
dialect? : @core.Dialect,
) -> @core.Expr raise @core.SqlglotError {
let root = build_scope(expression)
let dialect = get_dialect(dialect)
let unnest_requires_cross_join = dialect_is_a(dialect, ["athena", "presto"])
match root {
Some(root) => {
let scope_ref_count = root.ref_count()
let scopes = root.traverse()
scopes.rev_in_place()
for scope in scopes {
let select = scope.expression
let joins = select.list("joins")
match select.arg("where") {
Some(where_) => {
let join_index : Map[String, Int] = {}
for i, join in joins {
join_index[join.alias_or_name()] = i
}
let mut last_null_extending = -1
for i, join in joins {
let side = @core.py_upper(join.text("side"))
if side == "RIGHT" || side == "FULL" {
last_null_extending = i
}
}
let mut pushdown_allowed = true
let reachable : Map[String, (@core.Expr, Source)] = {}
for k, v in scope.selected_sources() {
let (node, _) = v
let position = match node.find_ancestor([Join, From]) {
Some(p) if p.kind.is_a(Join) => {
if node.kind.is_a(Unnest) && unnest_requires_cross_join {
pushdown_allowed = false
break
}
join_index.get(p.alias_or_name()).unwrap_or(-1)
}
_ => -1
}
if position >= last_null_extending {
reachable[k] = v
}
}
if pushdown_allowed {
pushdown(
where_.this(),
reachable,
scope_ref_count,
dialect,
Some(join_index),
)
}
}
None => ()
}
for join in joins {
let name = join.alias_or_name()
let side = @core.py_upper(join.text("side"))
if side == "RIGHT" || side == "FULL" {
continue
}
match scope.selected_sources().get(name) {
Some(v) => {
let sources : Map[String, (@core.Expr, Source)] = {}
sources[name] = v
pushdown(join.arg("on"), sources, scope_ref_count, dialect, None)
}
None => ()
}
}
}
}
None => ()
}
expression
}
///|
fn pushdown(
condition : @core.Expr?,
sources : Map[String, (@core.Expr, Source)],
scope_ref_count : Map[Int, Int],
dialect : @core.Dialect,
join_index : Map[String, Int]?,
) -> Unit raise @core.SqlglotError {
let condition = match condition {
Some(c) => c
None => return
}
let condition = condition
.replace(Some(simplify(condition, dialect~)))
.unwrap()
let cnf_like = normalized(condition) || !normalized(condition, dnf=true)
let predicates = if condition.kind.is_a(if cnf_like { And } else { Or }) {
condition.flatten().collect()
} else {
[condition]
}
if cnf_like {
pushdown_cnf(predicates, sources, scope_ref_count, join_index)
} else {
pushdown_dnf(predicates, sources, scope_ref_count, join_index)
}
}
///|
fn pushdown_cnf(
predicates : Array[@core.Expr],
sources : Map[String, (@core.Expr, Source)],
scope_ref_count : Map[Int, Int],
join_index : Map[String, Int]?,
) -> Unit raise @core.SqlglotError {
// The predicates are the operands of one connector: their JOIN/WHERE ancestor is
// found once for the whole chain (pushing a predicate replaces it by TRUE, which
// doesn't move the others).
let clauses = AncestorCache::new([Join, Where])
for predicate in predicates {
for
_, node in nodes_for_predicate(
predicate,
sources,
scope_ref_count,
clauses~,
) {
if node.kind.is_a(Join) {
let name = node.alias_or_name()
let predicate_tables = @core.column_table_names(predicate, exclude=name)
match join_index {
Some(ji) if !ji.is_empty() => {
let this_index = match ji.get(name) {
Some(i) => i
None => raise @core.OptimizeError("KeyError: \{name}")
}
if predicate_tables.iter().all(t => ji.get(t).unwrap_or(-1) < this_index) {
predicate.replace(Some(@core.true_())) |> ignore
join_on(node, predicate)
break
}
}
_ => ()
}
}
if node.kind.is_a(Select) {
predicate.replace(Some(@core.true_())) |> ignore
let inner_predicate = replace_aliases(node, predicate)
if find_in_scope(inner_predicate, [AggFunc]) is Some(_) {
node.having_([inner_predicate], copy=false) |> ignore
} else {
node.where_([inner_predicate], copy=false) |> ignore
}
}
}
}
}
///|
fn pushdown_dnf(
predicates : Array[@core.Expr],
sources : Map[String, (@core.Expr, Source)],
scope_ref_count : Map[Int, Int],
join_index : Map[String, Int]?,
) -> Unit raise @core.SqlglotError {
let pushdown_tables : Array[String] = []
for a in predicates {
let mut a_tables = @core.column_table_names(a)
for b in predicates {
let bt = @core.column_table_names(b)
a_tables = a_tables.filter(t => bt.contains(t))
}
for t in a_tables {
if !pushdown_tables.contains(t) {
pushdown_tables.push(t)
}
}
}
let conditions : Map[String, @core.Expr] = {}
for table in sorted_strings(pushdown_tables) {
let mut nodes : Map[String, @core.Expr] = {}
for predicate in predicates {
nodes = nodes_for_predicate(predicate, sources, scope_ref_count)
if !nodes.contains(table) {
continue
}
conditions[table] = match conditions.get(table) {
Some(c) => @core.or_([c, predicate])
None => predicate
}
}
for name, node in nodes {
let predicate = match conditions.get(name) {
Some(p) => p
None => continue
}
if node.kind.is_a(Join) {
match join_index {
Some(ji) if !ji.is_empty() => {
let this_index = match ji.get(name) {
Some(i) => i
None => raise @core.OptimizeError("KeyError: \{name}")
}
let predicate_tables = @core.column_table_names(predicate, exclude=name)
if !predicate_tables.iter().all(t => ji.get(t).unwrap_or(-1) < this_index) {
continue
}
}
_ => ()
}
join_on(node, predicate)
} else if node.kind.is_a(Select) {
let inner_predicate = replace_aliases(node, predicate)
if find_in_scope(inner_predicate, [AggFunc]) is Some(_) {
node.having_([inner_predicate], copy=false) |> ignore
} else {
node.where_([inner_predicate], copy=false) |> ignore
}
}
}
}
}
///|
fn nodes_for_predicate(
predicate : @core.Expr,
sources : Map[String, (@core.Expr, Source)],
scope_ref_count : Map[Int, Int],
clauses? : AncestorCache,
) -> Map[String, @core.Expr] raise @core.SqlglotError {
let nodes : Map[String, @core.Expr] = {}
let tables = @core.column_table_names(predicate)
let clause = match clauses {
Some(c) => c.find(predicate)
None => predicate.find_ancestor([Join, Where])
}
let where_condition = match clause {
Some(a) => a.kind.is_a(Where)
None => false
}
for table in sorted_strings(tables) {
let (node0, source) = match sources.get(table) {
Some((n, s)) => (Some(n), Some(s))
None => (None, None)
}
let mut node = node0
if node is Some(n) && where_condition {
node = n.find_ancestor([Join, From])
}
match (node, source) {
(Some(n), Some(ScopeSource(s))) if n.kind.is_a(From) => {
let parent = match s.parent {
Some(p) => p
None => raise @core.ValueError("Source node has no parent")
}
match parent.expression.arg("with_") {
Some(w) if w.has("recursive") => return {}
_ => ()
}
node = Some(s.expression)
}
_ => ()
}
match node {
Some(n) if n.kind.is_a(Join) => {
let side = @core.py_upper(n.text("side"))
if side != "" {
let pushable = match source {
Some(ScopeSource(s)) if side == "RIGHT" => Some(s)
_ => None
}
match pushable {
None => return {}
Some(s) => node = Some(s.expression)
}
} else {
nodes[table] = n
}
}
_ => ()
}
match node {
Some(n) if n.kind.is_a(Select) && tables.length() == 1 => {
let has_window_expression = n
.selects()
.iter()
.any(s => find_in_scope(s, [Window]) is Some(_))
let ref_count = match source {
Some(s) => scope_ref_count.get(s.key()).unwrap_or(0)
None => 0
}
if !n.has("group") &&
ref_count < 2 &&
!has_window_expression &&
!n.has("limit") &&
!n.has("offset") &&
!n.has("qualify") {
nodes[table] = n
}
}
_ => ()
}
}
nodes
}
///|
fn is_operator_expression(e : @core.Expr) -> Bool {
e.kind.is_any([Binary, Unary, Predicate])
}
///|
fn replace_aliases(source : @core.Expr, predicate : @core.Expr) -> @core.Expr {
let aliases : Map[String, @core.Expr] = {}
for select in source.selects() {
if select.kind.is_a(Alias) {
aliases[select.alias()] = select.this_()
} else {
aliases[select.name()] = select
}
}
predicate.transform(column => {
if column.kind.is_a(Column) && aliases.contains(column.name()) {
let mut replaced = aliases[column.name()].copy()
if is_operator_expression(replaced) &&
(match column.parent {
Some(p) => is_operator_expression(p)
None => false
}) {
replaced = @core.paren(replaced)
}
return Some(replaced)
}
Some(column)
})
}