// Port of sqlglot/optimizer/merge_subqueries.py.
///|
/// Caps the number of nodes copied by merges (relative to the statement size).
pub struct CopyBudget {
expression : @core.Expr
max_copy_factor : Int?
min_copy_budget : Int
mut remaining : Int?
}
///|
pub fn CopyBudget::new(
expression : @core.Expr,
max_copy_factor? : Int? = Some(8),
min_copy_budget? : Int = 1000,
) -> CopyBudget {
{ expression, max_copy_factor, min_copy_budget, remaining: None }
}
///|
/// Charges the copies needed to merge `inner_scope` into `outer_scope`, if they fit.
pub fn CopyBudget::consume(
self : CopyBudget,
outer_scope : Scope,
inner_scope : Scope,
alias : String,
) -> Bool {
let factor = match self.max_copy_factor {
Some(f) => f
None => return true
}
if self.remaining is None {
let size = self.expression.walk().count()
self.remaining = Some(@core.max_int(factor * size, self.min_copy_budget))
}
let remaining = self.remaining.unwrap()
let references : Map[String, Int] = {}
for c in outer_scope.columns() {
if c.table_name() == alias {
references[c.name()] = references.get_or_default(c.name(), 0) + 1
}
}
let mut copies = 0
for projection in inner_scope.expression.expressions() {
let count = references.get_or_default(projection.alias_or_name(), 0) - 1
if count > 0 {
for _ in projection.unalias().walk() {
copies += count
if copies > remaining {
return false
}
}
}
}
self.remaining = Some(remaining - copies)
true
}
///|
/// Rewrite the AST to merge derived tables into the outer query.
pub fn merge_subqueries(
expression : @core.Expr,
leave_tables_isolated? : Bool = false,
max_copy_factor? : Int? = Some(8),
min_copy_budget? : Int = 1000,
) -> @core.Expr raise @core.SqlglotError {
let mut scopes = traverse_scope(expression)
let copy_budget = CopyBudget::new(expression, max_copy_factor~, min_copy_budget~)
let (expression, merged_ctes) = merge_ctes(
expression,
leave_tables_isolated~,
scopes~,
copy_budget~,
)
if merged_ctes {
scopes = traverse_scope(expression)
}
merge_derived_tables(expression, leave_tables_isolated~, scopes~, copy_budget~)
}
///|
let unmergable_args : Array[String] = {
let keep = ["expressions", "from_", "joins", "where", "order", "hint"]
let out = []
for kv in @core.Kind::Select.arg_types() {
if !keep.contains(kv.0) {
out.push(kv.0)
}
}
out
}
///|
pub fn merge_ctes(
expression : @core.Expr,
leave_tables_isolated? : Bool = false,
scopes? : Array[Scope],
copy_budget? : CopyBudget,
) -> (@core.Expr, Bool) raise @core.SqlglotError {
let copy_budget = match copy_budget {
Some(c) => c
None => CopyBudget::new(expression)
}
let scopes = match scopes {
Some(s) => s
None => traverse_scope(expression)
}
let cte_selections : Map[Int, Array[(Scope, Scope, @core.Expr)]] = {}
for outer_scope in scopes {
for _, v in outer_scope.selected_sources() {
let (table, source) = v
match source {
ScopeSource(inner_scope) if inner_scope.is_cte() => {
if !cte_selections.contains(inner_scope.id) {
cte_selections[inner_scope.id] = []
}
cte_selections[inner_scope.id].push((outer_scope, inner_scope, table))
}
_ => ()
}
}
}
let mut merged = false
let singular = []
for _, v in cte_selections {
if v.length() == 1 {
singular.push(v[0])
}
}
for sel in singular {
let (outer_scope, inner_scope, table) = sel
let from_or_join = match table.find_ancestor([From, Join]) {
Some(f) if f.kind.is_any([From, Join]) => f
_ => continue
}
let alias = table.alias_or_name()
if mergeable(outer_scope, inner_scope, leave_tables_isolated, from_or_join) &&
copy_budget.consume(outer_scope, inner_scope, alias) {
rename_inner_sources(outer_scope, inner_scope, alias)
merge_from(outer_scope, inner_scope, table, alias)
merge_expressions(outer_scope, inner_scope, alias)
merge_order(outer_scope, inner_scope)
merge_joins(outer_scope, inner_scope, from_or_join)
merge_where(outer_scope, inner_scope, from_or_join)
merge_hints(outer_scope, inner_scope)
pop_cte(inner_scope)
outer_scope.clear_cache()
merged = true
}
}
(expression, merged)
}
///|
pub fn merge_derived_tables(
expression : @core.Expr,
leave_tables_isolated? : Bool = false,
scopes? : Array[Scope],
copy_budget? : CopyBudget,
) -> @core.Expr raise @core.SqlglotError {
let copy_budget = match copy_budget {
Some(c) => c
None => CopyBudget::new(expression)
}
let scopes = match scopes {
Some(s) => s
None => traverse_scope(expression)
}
for outer_scope in scopes {
for subquery in outer_scope.derived_tables() {
let from_or_join = match subquery.find_ancestor([From, Join]) {
Some(f) if f.kind.is_any([From, Join]) => f
_ => continue
}
let alias = subquery.alias_or_name()
let inner_scope = match outer_scope.sources.get(alias) {
Some(ScopeSource(s)) => s
Some(_) => continue
None => raise @core.OptimizeError("KeyError: \{alias}")
}
if mergeable(outer_scope, inner_scope, leave_tables_isolated, from_or_join) &&
copy_budget.consume(outer_scope, inner_scope, alias) {
rename_inner_sources(outer_scope, inner_scope, alias)
merge_from(outer_scope, inner_scope, subquery, alias)
merge_expressions(outer_scope, inner_scope, alias)
merge_order(outer_scope, inner_scope)
merge_joins(outer_scope, inner_scope, from_or_join)
merge_where(outer_scope, inner_scope, from_or_join)
merge_hints(outer_scope, inner_scope)
outer_scope.clear_cache()
}
}
}
expression
}
///|
fn side_in(join : @core.Expr, sides : Array[String]) -> Bool {
sides.contains(@core.py_upper(join.text("side")))
}
///|
fn mergeable(
outer_scope : Scope,
inner_scope : Scope,
leave_tables_isolated : Bool,
from_or_join : @core.Expr,
) -> Bool raise @core.SqlglotError {
let inner_select = inner_scope.expression.unnest()
let outer = outer_scope.expression
let inner_name = from_or_join.alias_or_name()
let is_join = from_or_join.kind.is_a(Join)
if !outer.kind.is_a(Select) ||
outer.is_star() ||
!inner_select.kind.is_a(Select) ||
unmergable_args.iter().any(k => inner_select.has(k)) ||
inner_select.get("from_") is None ||
!outer_scope.pivots().is_empty() ||
(leave_tables_isolated && outer_scope.selected_sources().length() > 1) ||
(is_join && inner_select.has("joins")) ||
(is_join &&
inner_select.has("where") &&
side_in(from_or_join, ["FULL", "LEFT", "RIGHT"])) ||
(from_or_join.kind.is_a(From) &&
inner_select.has("where") &&
outer.list("joins").iter().any(j => side_in(j, ["FULL", "RIGHT"]))) ||
(inner_select.has("order") && outer_scope.is_set_operation()) ||
(match inner_select.expressions().get(0) {
Some(e) => e.kind.is_a(QueryTransform)
None => false
}) {
return false
}
let window_aliases = []
let number_literal_aliases = []
let projections : Map[String, @core.Expr] = {}
for s in inner_select.selects() {
let name = s.alias_or_name()
projections[name] = s
if s.unalias().is_number() && !number_literal_aliases.contains(name) {
number_literal_aliases.push(name)
}
for node in s.walk() {
if node.kind.is_any([
AggFunc, Select, Anonymous, UDTF, ExplodingGenerateSeries,
]) {
return false
}
if node.kind.is_a(Window) && !window_aliases.contains(name) {
window_aliases.push(name)
}
}
}
// _outer_select_joins_on_inner_select_join
let joins_on_inner_join = if !is_join {
false
} else {
match from_or_join.arg("on") {
None => false
Some(on) => {
let selections = on
.find_all([Column])
.filter(c => c.table_name() == inner_name)
.map(c => c.name())
.collect()
match inner_scope.expression.arg("from_") {
None => false
Some(inner_from) => {
let inner_from_table = inner_from.alias_or_name()
let mut found = false
for selection in selections {
let p = match projections.get(selection) {
Some(p) => p
None => raise @core.OptimizeError("KeyError: \{selection}")
}
if p.find_all([Column]).any(col => col.table_name() != inner_from_table) {
found = true
break
}
}
found
}
}
}
}
}
// _window_projection_blocks_merge
let window_blocks = if window_aliases.is_empty() {
false
} else if outer.has("where") || outer.has("joins") {
true
} else {
outer_scope
.columns()
.iter()
.any(column => column.table_name() == inner_name &&
window_aliases.contains(column.name()) &&
column.find_ancestor([Group, Order, Having, AggFunc]) is Some(_))
}
// _literal_group_unmergeable
let literal_group = match outer.arg("group") {
None => false
Some(_) if number_literal_aliases.is_empty() => false
Some(group) => {
let grouped = []
let top_level_ids : @set.Set[Int] = @set.new()
for e in group.expressions() {
top_level_ids.add(e.unnest().uid)
}
let mut blocked = false
for col in group.find_all([Column]).collect() {
if col.table_name() != inner_name ||
!number_literal_aliases.contains(col.name()) {
continue
}
if !top_level_ids.contains(col.uid) {
blocked = true
break
}
grouped.push(col.name())
}
if blocked {
true
} else if grouped.is_empty() {
false
} else {
let projected = []
for s in outer.selects() {
let unaliased = s.unalias()
if unaliased.kind.is_a(Column) && unaliased.table_name() == inner_name {
projected.push(unaliased.name())
}
}
!grouped.iter().all(g => projected.contains(g))
}
}
}
// _literal_in_order_by
let literal_order = match outer.arg("order") {
None => false
Some(order) =>
order
.expressions()
.iter()
.any(o => {
let key = o.this_().unnest()
key.kind.is_a(Column) &&
key.table_name() == inner_name &&
number_literal_aliases.contains(key.name())
})
}
// _is_recursive
let is_recursive = if inner_scope.is_cte() {
let cte = inner_scope.expression.parent
let mut node = outer.parent
let mut found = false
while node is Some(n) {
match cte {
Some(c) if physical_equal(n, c) => {
found = true
break
}
_ => ()
}
node = n.parent
}
found
} else {
false
}
!joins_on_inner_join &&
!window_blocks &&
!literal_group &&
!literal_order &&
!is_recursive
}
///|
/// Renames any sources in the inner query that conflict with names in the outer query.
fn rename_inner_sources(
outer_scope : Scope,
inner_scope : Scope,
alias : String,
) -> Unit raise @core.SqlglotError {
let inner_taken = inner_scope.selected_sources().keys().collect()
let outer_taken = outer_scope.selected_sources().keys().collect()
let conflicts = outer_taken.filter(n => inner_taken.contains(n) && n != alias)
let taken = dedup_strings(outer_taken + inner_taken)
for conflict in conflicts {
let new_name = @core.find_new_name(n => taken.contains(n), conflict)
let (source, _) = inner_scope.selected_sources()[conflict]
let new_alias = @core.to_identifier(new_name)
if source.kind.is_a(Table) && source.alias() != "" {
source.set("alias", @core.mk1(TableAlias, new_alias))
} else if source.kind.is_a(Table) {
source.replace(Some(@core.alias_expr(source, Some(new_alias)))) |> ignore
} else if parent_is(source, [Subquery]) {
source.parent.unwrap().set("alias", @core.mk1(TableAlias, new_alias))
}
for column in inner_scope.source_columns(conflict) {
column.set("table", @core.to_identifier(new_name))
}
inner_scope.rename_source(Some(conflict), new_name)
}
}
///|
fn merge_from(
outer_scope : Scope,
inner_scope : Scope,
node_to_replace : @core.Expr,
alias : String,
) -> Unit raise @core.SqlglotError {
let new_subquery = inner_scope.expression.arg("from_").unwrap().this_()
new_subquery.set("joins", node_to_replace.get("joins"))
node_to_replace.replace(Some(new_subquery)) |> ignore
for join_hint in outer_scope.join_hints() {
for table in join_hint.find_all([Table]).collect() {
if table.alias_or_name() == node_to_replace.alias_or_name() {
table.set("this", @core.to_identifier(new_subquery.alias_or_name()))
}
}
}
outer_scope.remove_source(alias)
match inner_scope.sources.get(new_subquery.alias_or_name()) {
Some(s) => outer_scope.add_source(new_subquery.alias_or_name(), s)
None =>
raise @core.OptimizeError("KeyError: \{new_subquery.alias_or_name()}")
}
}
///|
fn merge_joins(
outer_scope : Scope,
inner_scope : Scope,
from_or_join : @core.Expr,
) -> Unit raise @core.SqlglotError {
let new_joins = []
for join in inner_scope.expression.list("joins") {
new_joins.push(join)
match inner_scope.sources.get(join.alias_or_name()) {
Some(s) => outer_scope.add_source(join.alias_or_name(), s)
None => raise @core.OptimizeError("KeyError: \{join.alias_or_name()}")
}
}
if !new_joins.is_empty() {
let outer_joins = outer_scope.expression.list("joins")
let position = if from_or_join.kind.is_a(From) {
0
} else {
let mut idx = -1
for i, j in outer_joins {
if j == from_or_join {
idx = i
break
}
}
if idx < 0 {
raise @core.ValueError("join is not in list")
}
idx + 1
}
for i, j in new_joins {
outer_joins.insert(position + i, j)
}
outer_scope.expression.set("joins", outer_joins)
}
}
///|
fn merge_expressions(
outer_scope : Scope,
inner_scope : Scope,
alias : String,
) -> Unit {
let outer_columns : Map[String, Array[@core.Expr]] = {}
for column in outer_scope.columns() {
if column.table_name() == alias {
if !outer_columns.contains(column.name()) {
outer_columns[column.name()] = []
}
outer_columns[column.name()].push(column)
}
}
let group = outer_scope.expression.arg("group")
for expression in inner_scope.expression.expressions() {
let projection_name = expression.alias_or_name()
if projection_name == "" {
continue
}
let columns_to_replace = outer_columns.get(projection_name).unwrap_or([])
if columns_to_replace.is_empty() {
continue
}
let expression = expression.unalias()
let must_wrap_expression = !expression.kind.is_any([
Column, EQ, Func, NEQ, Paren,
])
let is_number = expression.is_number()
let last = columns_to_replace.length() - 1
let mut group_ordinal = 0
if is_number && outer_scope.expression.kind.is_a(Select) {
for j, s in outer_scope.expression.selects() {
let unaliased = s.unalias()
if unaliased.kind.is_a(Column) &&
unaliased.table_name() == alias &&
unaliased.name() == projection_name {
group_ordinal = j + 1
break
}
}
}
for i, column in columns_to_replace {
let parent = column.parent
if is_number {
match group {
Some(g) => {
let mut item : @core.Expr? = None
for e in g.expressions() {
if physical_equal(e.unnest(), column) {
item = Some(e)
break
}
}
match item {
Some(it) => {
it.replace(Some(lit_num(group_ordinal))) |> ignore
continue
}
None => ()
}
}
None => ()
}
}
let mut replacement = if i < last { expression.copy() } else { expression }
if (match parent {
Some(p) => p.kind.is_any([Unary, Binary])
None => false
}) &&
must_wrap_expression {
replacement = @core.paren(replacement, copy=false)
}
if (match parent {
Some(p) => p.kind.is_a(Select)
None => false
}) &&
column.name() != expression.name() {
replacement = @core.alias_(replacement, column.name(), copy=false)
}
column.replace(Some(replacement)) |> ignore
}
}
}
///|
fn merge_where(
outer_scope : Scope,
inner_scope : Scope,
from_or_join : @core.Expr,
) -> Unit raise @core.SqlglotError {
let where_ = match inner_scope.expression.arg("where") {
Some(w) => w
None => return
}
let cond = match where_.this() {
Some(c) => c
None => return
}
let expression = outer_scope.expression
if from_or_join.kind.is_a(Join) {
let sources = []
match expression.arg("from_") {
Some(f) => sources.push(f.alias_or_name())
None => ()
}
for join in expression.list("joins") {
let source = join.alias_or_name()
sources.push(source)
if source == from_or_join.alias_or_name() {
break
}
}
if @core.column_table_names(cond).iter().all(t => sources.contains(t)) {
join_on(from_or_join, cond)
from_or_join.set("on", from_or_join.get("on"))
return
}
}
expression.where_([cond], copy=false) |> ignore
}
///|
fn merge_order(outer_scope : Scope, inner_scope : Scope) -> Unit {
let inner_order = match inner_scope.expression.arg("order") {
Some(o) => o
None => return
}
let outer = outer_scope.expression
if ["group", "distinct", "having", "order"].iter().any(a => outer.has(a)) ||
outer_scope.selected_sources_or_empty().length() != 1 ||
outer.expressions().iter().any(e => e.find([AggFunc]) is Some(_)) {
return
}
outer.set("order", inner_order)
}
///|
fn merge_hints(outer_scope : Scope, inner_scope : Scope) -> Unit {
let inner_hint = match inner_scope.expression.arg("hint") {
Some(h) => h
None => return
}
match outer_scope.expression.arg("hint") {
Some(outer_hint) =>
for h in inner_hint.expressions() {
outer_hint.append("expressions", h)
}
None => outer_scope.expression.set("hint", inner_hint)
}
}
///|
fn pop_cte(inner_scope : Scope) -> Unit {
let cte = match inner_scope.expression.parent {
Some(c) => c
None => return
}
let with_ = match cte.parent {
Some(w) => w
None => return
}
if with_.expressions().length() == 1 {
with_.pop() |> ignore
} else {
cte.pop() |> ignore
}
}