// Port of sqlglot/optimizer/optimize_joins.py, eliminate_ctes.py and eliminate_joins.py.
///|
/// Python `helper.tsort`: topological sort of a DAG (name -> dependencies).
pub fn tsort(
dag : Array[(String, Array[String])],
) -> Array[String] raise @core.SqlglotError {
let nodes : Array[(String, Array[String])] = dag.map(kv => (kv.0, kv.1.copy()))
let present = fn(n : String) { nodes.iter().any(kv => kv.0 == n) }
for kv in nodes.copy() {
for dep in kv.1 {
if !present(dep) {
nodes.push((dep, []))
}
}
}
let result = []
while nodes.length() > 0 {
let current = nodes.filter(kv => kv.1.is_empty()).map(kv => kv.0)
if current.is_empty() {
raise @core.ValueError("Cycle error")
}
let remaining = nodes.filter(kv => !current.contains(kv.0))
nodes.clear()
for kv in remaining {
nodes.push((kv.0, kv.1.filter(d => !current.contains(d))))
}
for c in sorted_strings(dedup_strings(current)) {
result.push(c)
}
}
result
}
///|
fn other_table_names(join : @core.Expr) -> Array[String] {
match join.arg("on") {
Some(on) => @core.column_table_names(on, exclude=join.alias_or_name())
None => []
}
}
///|
fn is_reorderable(joins : Array[@core.Expr]) -> Bool {
!joins.iter().any(j => j.text("side") != "")
}
///|
/// Removes cross joins if possible and reorder joins based on predicate dependencies.
pub fn optimize_joins(
expression : @core.Expr,
) -> @core.Expr raise @core.SqlglotError {
for select in expression.find_all([Select]).collect() {
let joins = select.list("joins")
if !is_reorderable(joins) {
continue
}
let references : Map[String, Array[@core.Expr]] = {}
let cross_joins : Array[(String, @core.Expr)] = []
for join in joins {
let tables = other_table_names(join)
if !tables.is_empty() {
for table in tables {
if !references.contains(table) {
references[table] = []
}
references[table].push(join)
}
} else {
cross_joins.push((join.alias_or_name(), join))
}
}
for cj in cross_joins {
let (name, join) = cj
for dep in references.get(name).unwrap_or([]) {
if @core.py_upper(dep.text("kind")) == "ANTI" {
continue
}
let on = dep.arg("on").unwrap()
if on.kind.is_a(And) {
if other_table_names(dep).length() < 2 {
continue
}
let it = on.flatten()
while it.next() is Some(predicate) {
if @core.column_table_names(predicate).contains(name) {
predicate.replace(Some(@core.true_())) |> ignore
let combined = match join.arg("on") {
Some(existing) =>
@core.combine_conditions(
[existing, predicate],
And,
copy=false,
)
None => predicate
}
join.set("on", combined)
if @core.py_upper(join.text("kind")) == "CROSS" {
join.set("kind", @core.null_arg)
}
}
}
}
}
}
}
let expression = reorder_joins(expression)
normalize_joins(expression)
}
///|
/// Reorder joins by topological sort order based on predicate references.
pub fn reorder_joins(
expression : @core.Expr,
) -> @core.Expr raise @core.SqlglotError {
for from_ in expression.find_all([From]).collect() {
let parent = match from_.parent {
Some(p) => p
None => raise @core.OptimizeError("FROM clause without parent expression")
}
let joins = parent.list("joins")
if !is_reorderable(joins) {
continue
}
let joins_by_name : Map[String, @core.Expr] = {}
for join in joins {
joins_by_name[join.alias_or_name()] = join
}
let dag = []
for name, join in joins_by_name {
dag.push((name, other_table_names(join)))
}
let from_name = from_.alias_or_name()
let ordered = tsort(dag)
.filter(name => name != from_name && joins_by_name.contains(name))
.map(name => joins_by_name[name])
parent.set("joins", ordered)
}
expression
}
///|
/// Remove INNER and OUTER from joins as they are optional.
fn normalize_joins(expression : @core.Expr) -> @core.Expr {
for join in expression.find_all([Join]).collect() {
if !["on", "side", "kind", "using", "method"].iter().any(k => join.has(k)) {
join.set("kind", "CROSS")
}
let kind = @core.py_upper(join.text("kind"))
if kind == "CROSS" {
join.set("on", @core.null_arg)
} else {
if kind == "INNER" || kind == "OUTER" {
join.set("kind", @core.null_arg)
}
if !join.has("on") && !join.has("using") {
join.set("on", @core.true_())
}
}
}
expression
}
///|
/// Remove unused CTEs from an expression.
pub fn eliminate_ctes(
expression : @core.Expr,
journal? : Journal,
) -> @core.Expr raise @core.SqlglotError {
match build_scope(expression) {
Some(root) => {
let ref_count = root.ref_count()
let scopes = root.traverse()
scopes.rev_in_place()
for scope in scopes {
if scope.is_cte() {
let count = ref_count.get(scope.key()).unwrap_or(0)
if count <= 0 {
let cte_node = match scope.expression.parent {
Some(c) => c
None => continue
}
let with_node = cte_node.parent
match (journal, with_node) {
(Some(j), Some(w)) => record(j, w, "expressions")
_ => ()
}
cte_node.pop() |> ignore
match with_node {
Some(w) if w.expressions().is_empty() => {
match (journal, w.parent) {
(Some(j), Some(p)) => record(j, p, "with_")
_ => ()
}
w.pop() |> ignore
}
_ => ()
}
for _, v in scope.selected_sources() {
match v.1 {
ScopeSource(s) =>
ref_count[s.key()] = ref_count.get(s.key()).unwrap_or(0) - 1
_ => ()
}
}
}
}
}
}
None => ()
}
expression
}
///|
/// Remove unused joins from an expression.
pub fn eliminate_joins(
expression : @core.Expr,
) -> @core.Expr raise @core.SqlglotError {
for scope in traverse_scope(expression) {
let joins = scope.expression.list("joins")
if joins.is_empty() {
continue
}
if !scope.unqualified_columns().is_empty() {
continue
}
let reversed = joins.copy()
reversed.rev_in_place()
for join in reversed {
if is_semi_or_anti_join(join) {
continue
}
let alias = join.alias_or_name()
if should_eliminate_join(scope, join, alias) {
join.pop() |> ignore
scope.remove_source(alias)
}
}
}
expression
}
///|
fn should_eliminate_join(scope : Scope, join : @core.Expr, alias : String) -> Bool {
match scope.sources.get(alias) {
Some(ScopeSource(inner)) =>
!join_is_used(scope, join, alias) &&
((@core.py_upper(join.text("side")) == "LEFT" &&
is_joined_on_all_unique_outputs(inner, join)) ||
(!join.has("on") && has_single_output_row(inner)))
_ => false
}
}
///|
fn join_is_used(scope : Scope, join : @core.Expr, alias : String) -> Bool {
let on_ids : @set.Set[Int] = @set.new()
match join.arg("on") {
Some(on) => for c in on.find_all([Column]) { on_ids.add(c.uid) }
None => ()
}
scope.source_columns(alias).iter().any(c => !on_ids.contains(c.uid))
}
///|
fn is_joined_on_all_unique_outputs(scope : Scope, join : @core.Expr) -> Bool {
let unique_outputs = unique_outputs(scope)
if unique_outputs.is_empty() {
return false
}
let (_, join_keys, _) = join_condition(join)
let names = join_keys.map(c => c.name())
unique_outputs.iter().all(o => names.contains(o))
}
///|
fn unique_outputs(scope : Scope) -> Array[String] {
let expr = scope.expression
if expr.get("distinct") is Some(_) {
return dedup_strings(expr.named_selects())
}
match expr.arg("group") {
Some(group) => {
let grouped_expressions = expr_set(group.expressions())
let grouped_outputs = []
let unique = []
for select in expr.selects() {
let output = select.unalias()
if grouped_expressions.contains(output) {
if !grouped_outputs.contains(output) {
grouped_outputs.push(output)
}
if !unique.contains(select.alias_or_name()) {
unique.push(select.alias_or_name())
}
}
}
if grouped_expressions.iter().all(g => grouped_outputs.contains(g)) {
return unique
}
return []
}
None => ()
}
if has_single_output_row(scope) {
return dedup_strings(expr.named_selects())
}
[]
}
///|
fn has_single_output_row(scope : Scope) -> Bool {
let e = scope.expression
e.kind.is_a(Select) &&
(e.selects().iter().all(s => s.unalias().kind.is_a(AggFunc)) ||
is_limit_1(scope) ||
!e.has("from_"))
}
///|
fn is_limit_1(scope : Scope) -> Bool {
match scope.expression.arg("limit") {
Some(limit) =>
match limit.expression() {
Some(e) => e.get("this") is Some(Str("1"))
None => false
}
None => false
}
}
///|
/// Extract the join condition: (source keys, join keys, remaining predicate).
pub fn join_condition(
join : @core.Expr,
) -> (Array[@core.Expr], Array[@core.Expr], @core.Expr) {
let name = join.alias_or_name()
let mut on = match join.arg("on") {
Some(o) => o.copy()
None => @core.true_()
}
let source_key = []
let join_key = []
fn extract_condition(condition : @core.Expr) {
let operands = condition.unnest_operands()
let left = operands[0]
let right = operands[1]
let left_tables = @core.column_table_names(left)
let right_tables = @core.column_table_names(right)
if left_tables.contains(name) && !right_tables.contains(name) {
join_key.push(left)
source_key.push(right)
condition.replace(Some(@core.true_())) |> ignore
} else if right_tables.contains(name) && !left_tables.contains(name) {
join_key.push(right)
source_key.push(left)
condition.replace(Some(@core.true_())) |> ignore
}
}
if normalized(on) {
if !on.kind.is_a(And) {
on = @core.and_([on, @core.true_()], copy=false)
}
let it = on.flatten()
while it.next() is Some(condition) {
if condition.kind.is_a(EQ) {
extract_condition(condition)
}
}
} else if normalized(on, dnf=true) {
let mut conditions : Array[@core.Expr] = []
for condition in on.flatten().collect() {
let parts = condition.flatten().filter(p => p.kind.is_a(EQ)).collect()
if conditions.is_empty() {
conditions = parts
} else {
let temp = []
for p in parts {
let cs = conditions.filter(c => p == c)
if !cs.is_empty() {
temp.push(p)
for c in cs {
temp.push(c)
}
}
}
conditions = temp
}
}
for condition in conditions {
extract_condition(condition)
}
}
(source_key, join_key, on)
}