// Port of sqlglot/optimizer/pushdown_projections.py and journal.py.
///|
/// One recorded argument mutation: (node, arg_key, value before the mutation).
pub type Journal = Array[(@core.Expr, String, @core.Value?)]
///|
/// Records the current value of `node.args[arg_key]` so `revert` can restore it.
pub fn record(journal : Journal, node : @core.Expr, arg_key : String) -> Unit {
let value = match node.get(arg_key) {
Some(List(l)) => Some(@core.Value::List(l.copy()))
v => v
}
journal.push((node, arg_key, value))
}
///|
/// Restores every argument recorded from `start` onwards, newest first.
pub fn revert(journal : Journal, start? : Int = 0) -> Unit {
let mut i = journal.length() - 1
while i >= start {
let (node, arg_key, value) = journal[i]
node.set(arg_key, value)
i -= 1
}
while journal.length() > start {
journal.pop() |> ignore
}
}
///|
/// Remove unused projections and CTEs while preserving all outermost outputs.
pub fn pushdown_projections(
expression : @core.Expr,
journal? : Journal,
) -> @core.Expr raise @core.SqlglotError {
let reachability = projection_reachability(expression, whole_query=true)
prune_projections(reachability, 0, journal?)
expression
}
///|
/// Which root outputs reach each scope and output column through dependency edges.
pub struct ProjectionReachability {
scopes : Array[Scope]
/// scope id -> bitset of root outputs requiring the scope
live : Map[Int, Int64]
/// scope id -> per output column bitsets
selections : Map[Int, Array[Int64]]
is_agg : Map[Int, Bool]
group_by_ordinals : Map[Int, Array[(@core.Expr, @core.Expr)]]
set_names : Map[Int, Array[String]]
}
///|
priv struct DepNode {
id : Int
mut required_by : Int64
dependencies : Array[DepNode]
}
///|
let dep_node_counter : Ref[Int] = Ref(0)
///|
fn DepNode::new() -> DepNode {
dep_node_counter.val += 1
{ id: dep_node_counter.val, required_by: 0L, dependencies: [] }
}
///|
fn empty_reachability(scopes : Array[Scope]) -> ProjectionReachability {
{
scopes,
live: {},
selections: {},
is_agg: {},
group_by_ordinals: {},
set_names: {},
}
}
///|
/// Find which scopes and projections each outermost output reaches through dependencies.
pub fn projection_reachability(
expression : @core.Expr,
whole_query? : Bool = false,
) -> ProjectionReachability raise @core.SqlglotError {
if !whole_query && !expression.kind.is_a(Query) {
raise @core.OptimizeError("projection_reachability requires a query")
}
let scopes = traverse_scope(expression)
if scopes.is_empty() {
if whole_query {
return empty_reachability([])
}
raise @core.OptimizeError("projection_reachability requires a query scope")
}
if whole_query && scopes.length() == 1 {
let scope = scopes[0]
let query = scope.expression
if query.kind.is_a(Select) {
if query.is_star() {
raise @core.OptimizeError(
"projection_reachability requires star-free selections",
)
}
let windows = query.list("windows")
let r = empty_reachability(scopes)
r.live[scope.id] = 1L
r.selections[scope.id] = query.selects().map(_ => 1L)
r.is_agg[scope.id] = query
.selects()
.iter()
.any(s => projection_has_aggregate(s, windows))
r.group_by_ordinals[scope.id] = group_by_ordinal_refs(query)
return r
}
}
let scope_nodes : Map[Int, DepNode] = {}
let output_names : Map[Int, Array[String]] = {}
let output_nodes : Map[Int, Array[DepNode]] = {}
let outputs_by_name : Map[Int, Map[String, Array[DepNode]]] = {}
let owners : Map[Int, DepNode] = {}
let scopes_by_expression : Map[Int, Scope] = {}
let is_agg : Map[Int, Bool] = {}
let group_by_ordinals : Map[Int, Array[(@core.Expr, @core.Expr)]] = {}
for scope in scopes {
let query = scope.expression
if query.kind.is_a(Select) && query.is_star() {
raise @core.OptimizeError(
"projection_reachability requires star-free selections",
)
}
scopes_by_expression[query.uid] = scope
let scope_node = DepNode::new()
scope_nodes[scope.id] = scope_node
owners[query.uid] = scope_node
if query.kind.is_a(SetOperation) {
let left = scope.set_operation_scopes[0]
let right = scope.set_operation_scopes[1]
output_names[scope.id] = if query.has("by_name") {
dedup_strings(output_names[left.id] + output_names[right.id])
} else {
output_names[left.id]
}
} else {
output_names[scope.id] = if query.kind.is_a(Selectable) {
query.selects().map(s => s.alias_or_name())
} else {
[]
}
}
let outputs = output_names[scope.id].map(_ => DepNode::new())
output_nodes[scope.id] = outputs
let by_name : Map[String, Array[DepNode]] = {}
outputs_by_name[scope.id] = by_name
for i, name in output_names[scope.id] {
if !by_name.contains(name) {
by_name[name] = []
}
by_name[name].push(outputs[i])
outputs[i].dependencies.push(scope_node)
}
if query.kind.is_a(Select) {
let selects = query.selects()
for i in 0..<@core.min_int(selects.length(), outputs.length()) {
owners[selects[i].uid] = outputs[i]
}
}
}
fn owner(node : @core.Expr) -> DepNode {
let mut node = node
while !owners.contains(node.uid) {
node = node.parent.unwrap()
}
owners[node.uid]
}
fn by_name_get(scope_id : Int, name : String) -> Array[DepNode] {
match outputs_by_name[scope_id].get(name) {
Some(l) => l
None => []
}
}
for scope in scopes {
let query = scope.expression
let scope_node = scope_nodes[scope.id]
let outputs = output_nodes[scope.id]
let order = query.arg("order")
let mut keep_all = query.has("distinct") ||
query.kind.is_any([Intersect, Except]) ||
is_self_referencing_cte(scope) ||
!query.kind.is_any([Select, SetOperation])
for child in scope.subquery_scopes {
let anchor = child.expression.parent.unwrap()
let o = owner(anchor)
o.dependencies.push(scope_nodes[child.id])
for n in output_nodes[child.id] {
o.dependencies.push(n)
}
}
if query.kind.is_a(Subquery) {
for child in scope.derived_table_scopes {
scope_node.dependencies.push(scope_nodes[child.id])
for n in output_nodes[child.id] {
scope_node.dependencies.push(n)
}
}
}
if query.kind.is_a(SetOperation) {
let left = scope.set_operation_scopes[0]
let right = scope.set_operation_scopes[1]
scope_node.dependencies.push(scope_nodes[left.id])
scope_node.dependencies.push(scope_nodes[right.id])
let by_name = query.has("by_name")
if query.text("kind") != "" ||
query.text("side") != "" ||
(by_name && !scope.outer_columns.is_empty()) {
keep_all = true
}
if !by_name &&
output_nodes[left.id].length() != output_nodes[right.id].length() {
raise @core.OptimizeError(
"Invalid set operation due to column mismatch: \{expr_sql(query)}.",
)
}
for branch in [left, right] {
for i, output_node in output_nodes[branch.id] {
let targets = if by_name {
by_name_get(scope.id, output_names[branch.id][i])
} else {
match outputs.get(i) {
Some(o) => [o]
None => []
}
}
for output in targets {
output.dependencies.push(output_node)
output_node.dependencies.push(output)
}
}
}
}
if keep_all {
for o in outputs {
scope_node.dependencies.push(o)
}
} else {
match order {
Some(o) => {
let mut max_ordinal = 0
for ordered in o.expressions() {
match ordered.this() {
Some(t) if t.kind.is_a(Literal) && t.is_int() =>
match t.to_py_int() {
Some(v) => max_ordinal = @core.max_int(max_ordinal, v.to_int())
None => ()
}
_ => ()
}
}
for i in 0..<@core.min_int(max_ordinal, outputs.length()) {
scope_node.dependencies.push(outputs[i])
}
}
None => ()
}
}
for name in output_column_refs(query, !query.kind.is_a(Select)) {
for n in by_name_get(scope.id, name) {
scope_node.dependencies.push(n)
}
}
for i in 0..<@core.min_int(scope.outer_columns.length(), outputs.length()) {
scope_node.dependencies.push(outputs[i])
}
if query.kind.is_a(Select) {
let windows = query.list("windows")
let group_all = is_implicit_group_by_all(query)
group_by_ordinals[scope.id] = group_by_ordinal_refs(query)
let ordinals : @set.Set[Int] = @set.new()
for r in group_by_ordinals[scope.id] {
ordinals.add(r.1.uid)
}
let mut first_aggregate : DepNode? = None
let non_aggregates = []
let mut has_grouping_key = false
let selects = query.selects()
for i in 0..<@core.min_int(selects.length(), outputs.length()) {
let selection = selects[i]
let output_node = outputs[i]
let (aggregate, has_column, has_srf) = projection_properties(
selection, windows,
)
if aggregate {
if first_aggregate is None {
first_aggregate = Some(output_node)
}
} else {
non_aggregates.push(output_node)
}
if group_all && !aggregate && has_column {
has_grouping_key = true
}
if ordinals.contains(selection.uid) || (group_all && !aggregate) || has_srf {
scope_node.dependencies.push(output_node)
}
}
is_agg[scope.id] = first_aggregate is Some(_)
match first_aggregate {
Some(fa) if !query.has("group") || (group_all && !has_grouping_key) =>
for n in non_aggregates {
n.dependencies.push(fa)
}
_ => ()
}
}
for r in scope.references() {
let (name, reference) = r
let source = match scope.sources.get(name) {
Some(ScopeSource(s)) => s
_ => continue
}
let source = scopes_by_expression[source.expression.uid]
scope_node.dependencies.push(scope_nodes[source.id])
let source_outputs = output_nodes[source.id]
let first = if source.expression.kind.is_a(Selectable) {
source.expression.selects().get(0)
} else {
None
}
if scope.semi_or_anti_join_tables().contains(name) ||
scope.scans_all_subscope_columns() ||
!scope.pivots().is_empty() ||
(match first {
Some(f) => f.kind.is_a(QueryTransform)
None => false
}) {
for n in source_outputs {
scope_node.dependencies.push(n)
}
}
for i in 0..<@core.min_int(
reference.alias_column_names().length(),
source_outputs.length(),
) {
scope_node.dependencies.push(source_outputs[i])
}
}
for col in scope.columns() {
let key = if col.table_name() != "" { col.table_name() } else { col.name() }
match scope.sources.get(key) {
Some(ScopeSource(s)) => {
let source = scopes_by_expression[s.expression.uid]
let o = owner(col)
let deps = if col.table_name() != "" {
by_name_get(source.id, col.name())
} else {
output_nodes[source.id]
}
for n in deps {
o.dependencies.push(n)
}
}
_ => ()
}
}
for table_column in scope.table_columns() {
match scope.sources.get(table_column.name()) {
Some(ScopeSource(s)) => {
let source = scopes_by_expression[s.expression.uid]
let o = owner(table_column)
for n in output_nodes[source.id] {
o.dependencies.push(n)
}
}
_ => ()
}
}
}
let roots = if expression.kind.is_a(Query) {
[scopes[scopes.length() - 1]]
} else {
scopes.filter(s => match s.parent {
Some(p) => !scope_nodes.contains(p.id)
None => true
})
}
let pending : @deque.Deque[DepNode] = @deque.new()
let queued : @set.Set[Int] = @set.new()
for root in roots {
let root_outputs = output_nodes[root.id]
let all_roots = if whole_query {
1L
} else {
(1L << root_outputs.length()) - 1L
}
scope_nodes[root.id].required_by = all_roots
for i, output_node in root_outputs {
output_node.required_by = if whole_query { all_roots } else { 1L << i }
}
pending.push_back(scope_nodes[root.id])
queued.add(scope_nodes[root.id].id)
for o in root_outputs {
pending.push_back(o)
queued.add(o.id)
}
}
while pending.pop_front() is Some(node) {
queued.remove(node.id)
for dependency in node.dependencies {
let required_by = dependency.required_by | node.required_by
if required_by != dependency.required_by {
dependency.required_by = required_by
if !queued.contains(dependency.id) {
queued.add(dependency.id)
pending.push_back(dependency)
}
}
}
}
let r = empty_reachability(scopes)
// `named_selects` of a set operation is that of its left operand: memoized by
// expression so that a long left-deep chain isn't walked once per set operation
let set_names_by_expr : Map[Int, Array[String]] = {}
for scope in scopes {
r.live[scope.id] = scope_nodes[scope.id].required_by
r.selections[scope.id] = output_nodes[scope.id].map(o => o.required_by)
let query = scope.expression
if query.kind.is_a(SetOperation) {
let names = match query.this().map(t => t.unnest()) {
Some(left) if left.kind.is_a(SetOperation) =>
match set_names_by_expr.get(left.uid) {
Some(n) => n
None => query.named_selects()
}
_ => query.named_selects()
}
set_names_by_expr[query.uid] = names
r.set_names[scope.id] = names
}
}
for k, v in is_agg {
r.is_agg[k] = v
}
for k, v in group_by_ordinals {
r.group_by_ordinals[k] = v
}
r
}
///|
/// Prune the analyzed tree to the scopes and projections reachable from output `root`.
pub fn prune_projections(
reachability : ProjectionReachability,
root : Int,
journal? : Journal,
remove_ctes? : Bool = true,
) -> Unit {
let bit = 1L << root
let scopes = reachability.scopes.copy()
scopes.rev_in_place()
for scope in scopes {
if (reachability.live[scope.id] & bit) == 0L {
if remove_ctes && scope.is_cte() {
match scope.expression.parent {
Some(cte_node) if cte_node.kind.is_a(CTE) => {
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
}
_ => ()
}
}
_ => ()
}
}
continue
}
let expression = scope.expression
if !expression.kind.is_a(Select) {
continue
}
let selects = expression.selects()
let sels = reachability.selections[scope.id]
let mut subset = []
for i in 0..<@core.min_int(selects.length(), sels.length()) {
if (sels[i] & bit) != 0L {
subset.push(selects[i])
}
}
if subset.length() == selects.length() {
continue
}
let ordinal_refs = reachability.group_by_ordinals.get(scope.id).unwrap_or([])
if subset.is_empty() {
let agg = reachability.is_agg.get(scope.id).unwrap_or(false)
let placeholder = default_selection(agg)
let mut ancestor = scope
while ancestor.is_set_operation() && ancestor.parent is Some(p) {
ancestor = p
let names = reachability.set_names.get(ancestor.id).unwrap_or([])
let asels = reachability.selections.get(ancestor.id).unwrap_or([])
let mut retained_name : String? = None
for i in 0..<@core.min_int(names.length(), asels.length()) {
if (asels[i] & bit) != 0L {
retained_name = Some(names[i])
break
}
}
match retained_name {
Some(name) => {
placeholder.set(
"this",
if agg {
@core.mk1(Max, @core.null_())
} else {
@core.null_()
},
)
placeholder.set("alias", @core.to_identifier(name, quoted=true))
break
}
None => ()
}
}
subset = [placeholder]
}
match journal {
Some(j) => record(j, expression, "expressions")
None => ()
}
expression.set("expressions", subset)
if !ordinal_refs.is_empty() {
let new_pos : Map[Int, Int] = {}
for i, selection in subset {
new_pos[selection.uid] = i + 1
}
for r in ordinal_refs {
let (node, old_selection) = r
match new_pos.get(old_selection.uid) {
Some(pos) =>
if node.to_py_int() != Some(pos.to_int64()) {
match journal {
Some(j) => record(j, node, "this")
None => ()
}
node.set("this", pos.to_string())
}
None => ()
}
}
}
}
}
///|
/// Whether a projection aggregates rows, reads columns, or contains a set-returning function.
fn projection_properties(
selection : @core.Expr,
windows : Array[@core.Expr],
) -> (Bool, Bool, Bool) {
let mut has_aggregate_or_window = false
let mut has_column = false
let mut has_srf = false
for node in find_all_in_scope(selection, [
AggFunc, Window, Column, Anonymous, UDTF, ExplodingGenerateSeries,
]) {
if node.kind.is_any([AggFunc, Window]) {
has_aggregate_or_window = true
}
if node.kind.is_a(Column) {
has_column = true
}
if node.kind.is_any([Anonymous, UDTF, ExplodingGenerateSeries]) {
has_srf = true
}
}
let aggregate = has_aggregate_or_window &&
projection_has_aggregate(selection, windows)
(aggregate, has_column, has_srf)
}
///|
fn output_column_refs(expression : @core.Expr, scoped : Bool) -> Array[String] {
let refs = []
for arg in ["order", "sort", "distribute", "cluster"] {
match expression.arg(arg) {
Some(node) => {
let columns = if scoped {
find_all_in_scope(node, [Column]).collect()
} else {
node.find_all([Column]).collect()
}
for c in columns {
if c.table_name() == "" && !refs.contains(c.name()) {
refs.push(c.name())
}
}
}
None => ()
}
}
refs
}
///|
fn is_self_referencing_cte(scope : Scope) -> Bool {
match scope.expression.parent {
Some(cte) if cte.kind.is_a(CTE) =>
match cte.parent {
Some(w) if w.kind.is_a(With) && w.has("recursive") =>
scope.expression
.find_all([Table])
.any(table => table.db() == "" && table.name() == cte.alias())
_ => false
}
_ => false
}
}
///|
/// Selection to use if the selection list is empty.
pub fn default_selection(is_agg : Bool) -> @core.Expr {
let e = if is_agg {
@core.mk1(Max, @core.literal_int(1))
} else {
@core.literal_int(1)
}
@core.alias_(e, "_", copy=false)
}
///|
fn is_implicit_group_by_all(select : @core.Expr) -> Bool {
match select.arg("group") {
Some(group) if group.has("all") =>
!(group.has("expressions") ||
group.has("cube") ||
group.has("rollup") ||
group.has("grouping_sets"))
_ => false
}
}
///|
fn group_by_ordinal_refs(
select : @core.Expr,
) -> Array[(@core.Expr, @core.Expr)] {
let group = match select.arg("group") {
Some(g) => g
None => return []
}
let selects = select.selects()
let n = selects.length()
let refs = []
fn collect(nodes : Array[@core.Expr]) -> Unit {
for node in nodes {
if node.kind.is_any([Cube, GroupingSets, Paren, Rollup, Tuple]) {
collect(node.iter_expressions())
} else if node.is_int() && node.kind.is_a(Literal) {
match node.to_py_int() {
Some(p) => {
let pos = p.to_int()
if 1 <= pos && pos <= n {
refs.push((node, selects[pos - 1]))
}
}
None => ()
}
}
}
}
collect(group.iter_expressions())
refs
}
///|
let window_has_aggregate_key : String = "window_has_aggregate"
///|
/// Port of `optimizer.helpers.projection_has_aggregate`.
pub fn projection_has_aggregate(
projection : @core.Expr,
windows : Array[@core.Expr],
) -> Bool {
let windowed_aggregates : @set.Set[Int] = @set.new()
let remaining_windows : Map[String, @core.Expr] = {}
for w in windows {
remaining_windows[w.name()] = w
}
for node in walk_in_scope(projection) {
if node.kind.is_a(Window) {
let mut target = node.this()
while target is Some(t) && !t.kind.is_a(Func) {
target = t.this()
}
match target {
Some(t) if t.kind.is_a(AggFunc) => windowed_aggregates.add(t.uid)
_ => ()
}
let mut name = node.alias()
while name != "" {
let window = match remaining_windows.get(name) {
Some(w) => w
None => break
}
remaining_windows.remove(name)
let has_aggregate = match window.meta_get(window_has_aggregate_key) {
Some(Bool(b)) => b
_ => {
let b = find_in_scope(window, [AggFunc]) is Some(_)
window.get_meta()[window_has_aggregate_key] = Bool(b)
b
}
}
if has_aggregate {
return true
}
name = window.alias()
}
} else if node.kind.is_a(AggFunc) && !windowed_aggregates.contains(node.uid) {
return true
}
}
false
}