// Port of the optimizer-dependent transforms of sqlglot/transforms.py:
// `explode_projection_to_unnest` (needs `Scope`) and `eliminate_join_marks` (needs
// `traverse_scope` and `normalize`). The other transforms live in core's transforms.mbt.
///|
/// Python `exp.func(name, *args)` (copies the arguments and validates the result).
fn transforms_func(
name : String,
args : Array[@core.Expr],
) -> @core.Expr raise @core.SqlglotError {
let converted = args.map(a => a.copy())
let function = @core.func_(name, converted)
for error_message in function.error_messages(nargs=converted.length()) {
raise @core.ValueError(error_message)
}
function
}
///|
/// Python `exp.column(col, table=table)` where both parts are strings or identifiers.
fn transforms_column(
col : @core.Value,
table : @core.Value,
) -> @core.Expr raise @core.SqlglotError {
@core.mk(Column, [
("this", @core.to_identifier_value(col)),
("table", @core.to_identifier_value(table)),
])
}
///|
/// Python `exp.alias_(expression, alias, table=[...columns])`.
fn transforms_alias_table(
expression : @core.Expr,
alias_name : String,
columns : Array[@core.Value],
) -> @core.Expr raise @core.SqlglotError {
let e = expression.copy()
let table_alias = @core.mk1(TableAlias, @core.to_identifier(alias_name))
e.set("alias", table_alias)
for column in columns {
table_alias.append("columns", @core.to_identifier_value(column))
}
e
}
///|
/// Python `expressions.index(value)` (structural equality) for an expression list.
fn transforms_index_of(
expressions : Array[@core.Expr],
value : @core.Expr,
) -> Int raise @core.SqlglotError {
match expressions.search(value) {
Some(i) => i
None => raise @core.ValueError("\{value.kind.name()} is not in list")
}
}
///|
/// Convert explode/posexplode projections into unnests.
pub fn explode_projection_to_unnest(
index_offset? : Int = 0,
unnest_map? : Bool = false,
) -> @core.Transform {
fn(expression : @core.Expr) raise @core.SqlglotError {
if expression.kind == Select {
let taken_select_names : @set.Set[String] = @set.Set::new()
for name in expression.named_selects() {
taken_select_names.add(name)
}
let taken_source_names : @set.Set[String] = @set.Set::new()
for reference in Scope::new(expression).references() {
taken_source_names.add(reference.0)
}
fn new_name(names : @set.Set[String], name : String) -> String {
let name = @core.find_new_name(n => names.contains(n), name)
names.add(name)
name
}
let arrays : Array[@core.Expr] = []
let series_alias = new_name(taken_select_names, "pos")
let series = transforms_alias_table(
@core.mk(Unnest, [
(
"expressions",
[
@core.mk(GenerateSeries, [
("start", @core.literal_int(index_offset)),
]),
],
),
]),
new_name(taken_source_names, "_u"),
[Str(series_alias)],
)
// we use a snapshot here because expression.selects is mutated inside the loop
for select in expression.selects() {
let mut explode = match select.find([Explode]) {
Some(e) => e
None => continue
}
if unnest_map &&
explode.kind == Explode &&
explode.this_().is_type([MAP]) &&
(physical_equal(select, explode) || select.kind.is_a(Aliases)) {
let (map_key_alias, map_value_alias) : (@core.Value, @core.Value) = if select.kind.is_a(
Aliases,
) {
match select.expressions() {
[k, v] => (Node(k), Node(v))
aliases =>
raise @core.ValueError(
if aliases.length() < 2 {
"not enough values to unpack (expected 2, got \{aliases.length()})"
} else {
"too many values to unpack (expected 2)"
},
)
}
} else {
let k = new_name(taken_select_names, "key")
let v = new_name(taken_select_names, "value")
(Str(k), Str(v))
}
let map_unnest_source = new_name(taken_source_names, "_u")
let map_key_select = select
.replace(
Some(
@core.exp_alias(
transforms_column(map_key_alias, Str(map_unnest_source)),
map_key_alias,
),
),
)
.unwrap()
let expressions = expression.expressions()
expressions.insert(
transforms_index_of(expressions, map_key_select) + 1,
@core.exp_alias(
transforms_column(map_value_alias, Str(map_unnest_source)),
map_value_alias,
),
)
expression.set("expressions", expressions)
let unnest = transforms_alias_table(
@core.mk(Unnest, [("expressions", [explode.this_().copy()])]),
map_unnest_source,
[map_key_alias, map_value_alias],
)
if expression.has("from_") {
expression.join_(unnest, kind="CROSS", copy=false) |> ignore
} else {
expression.from_(unnest, copy=false) |> ignore
}
continue
}
let mut pos_alias : @core.Value = Str("")
let mut explode_alias : @core.Value = Str("")
let aliased = if select.kind.is_a(Alias) {
explode_alias = select.get("alias").unwrap_or(Str(""))
select
} else if select.kind.is_a(Aliases) {
let aliases = select.expressions()
pos_alias = Node(aliases[0])
explode_alias = Node(aliases[1])
select
.replace(Some(@core.alias_(select.this_(), "", copy=false)))
.unwrap()
} else {
let aliased = select.replace(Some(@core.alias_(select, ""))).unwrap()
explode = aliased.find([Explode]).unwrap()
aliased
}
let is_posexplode = explode.kind.is_a(Posexplode)
let mut explode_arg = explode.this_()
if explode.kind.is_a(ExplodeOuter) {
let bracket = @core.mk(Bracket, [
("this", explode_arg.copy()),
("expressions", [@core.literal_int(0)]),
])
bracket.set("safe", true)
bracket.set("offset", true)
explode_arg = transforms_func("IF", [
@core.exp_eq(
transforms_func("ARRAY_SIZE", [
transforms_func("COALESCE", [explode_arg, @core.mk0(Array)]),
]),
@core.literal_int(0),
),
@core.array_([bracket], copy=false),
explode_arg,
])
}
// This ensures that we won't use [POS]EXPLODE's argument as a new selection
if explode_arg.kind.is_a(Column) {
taken_select_names.add(explode_arg.output_name())
}
let unnest_source_alias = new_name(taken_source_names, "_u")
if !explode_alias.truthy() {
explode_alias = Str(new_name(taken_select_names, "col"))
if is_posexplode {
pos_alias = Str(new_name(taken_select_names, "pos"))
}
}
if !pos_alias.truthy() {
pos_alias = Str(new_name(taken_select_names, "pos"))
}
aliased.set("alias", @core.to_identifier_value(explode_alias))
let series_table_alias : @core.Value = Node(
series.arg("alias").unwrap().this_(),
)
let column = @core.mk(If, [
(
"this",
@core.exp_eq(
transforms_column(Str(series_alias), series_table_alias),
transforms_column(pos_alias, Str(unnest_source_alias)),
),
),
("true", transforms_column(explode_alias, Str(unnest_source_alias))),
])
explode.replace(Some(column)) |> ignore
if is_posexplode {
let expressions = expression.expressions()
expressions.insert(
transforms_index_of(expressions, aliased) + 1,
@core.exp_alias(
@core.mk(If, [
(
"this",
@core.exp_eq(
transforms_column(Str(series_alias), series_table_alias),
transforms_column(pos_alias, Str(unnest_source_alias)),
),
),
("true", transforms_column(pos_alias, Str(unnest_source_alias))),
]),
pos_alias,
),
)
expression.set("expressions", expressions)
}
if arrays.is_empty() {
if expression.has("from_") {
expression.join_(series, kind="CROSS", copy=false) |> ignore
} else {
expression.from_(series, copy=false) |> ignore
}
}
let mut size = @core.mk1(ArraySize, explode_arg.copy())
arrays.push(size)
// trino doesn't support left join unnest with on conditions
// if it did, this would be much simpler
expression.join_(
transforms_alias_table(
@core.mk(Unnest, [
("expressions", [explode_arg.copy()]),
("offset", @core.to_identifier_value(pos_alias)),
]),
unnest_source_alias,
[explode_alias],
),
kind="CROSS",
copy=false,
)
|> ignore
if index_offset != 1 {
size = @core.exp_binop(Sub, size, @core.literal_int(1))
}
expression.where_(
[
@core.exp_or([
@core.exp_eq(
transforms_column(Str(series_alias), series_table_alias),
transforms_column(pos_alias, Str(unnest_source_alias)),
),
@core.exp_and([
@core.exp_binop(
GT,
transforms_column(Str(series_alias), series_table_alias),
size,
),
@core.exp_eq(
transforms_column(pos_alias, Str(unnest_source_alias)),
size,
),
]),
]),
],
copy=false,
)
|> ignore
}
if !arrays.is_empty() {
let mut end = @core.mk(Greatest, [
("this", arrays[0]),
("expressions", arrays[1:].to_array()),
])
if index_offset != 1 {
end = @core.exp_binop(Sub, end, @core.literal_int(1 - index_offset))
}
series.expressions()[0].set("end", end)
}
}
expression
}
}
///|
/// Python `assert condition, message` (raised as a `ValueError` carrying the
/// `AssertionError` message, since `SqlglotError` has no assertion variant).
fn transforms_assert(
condition : Bool,
message : String,
) -> Unit raise @core.SqlglotError {
if !condition {
raise @core.ValueError("AssertionError: \{message}")
}
}
///|
/// Remove Oracle-style `(+)` join marks by converting them into explicit LEFT JOINs.
///
/// See https://docs.oracle.com/cd/B19306_01/server.102/b14200/queries006.htm#sthref3178
///
/// 1. You cannot specify the (+) operator in a query block that also contains FROM clause
/// join syntax.
/// 2. The (+) operator can appear only in the WHERE clause or, in the context of
/// left-correlation (that is, when specifying the TABLE clause) in the FROM clause, and can
/// be applied only to a column of a table or view.
///
/// The (+) operator does not produce an outer join if you specify one table in the outer query
/// and the other table in an inner query. A WHERE condition containing the (+) operator cannot
/// be combined with another condition using the OR logical operator, cannot use the IN
/// comparison condition and cannot compare a marked column with a subquery.
pub fn eliminate_join_marks(
expression : @core.Expr,
) -> @core.Expr raise @core.SqlglotError {
// we go in reverse to check the main query for left correlation
let scopes = traverse_scope(expression)
for i = scopes.length() - 1; i >= 0; i = i - 1 {
let scope = scopes[i]
let query = scope.expression
let where_ = match query.arg("where") {
Some(w) => w
None => continue
}
let joins = query.list("joins")
if !where_.find_all([Column]).any(c => c.has("join_mark")) {
continue
}
// knockout: we do not support left correlation (see point 2)
transforms_assert(
!scope.is_correlated_subquery(),
"Correlated queries are not supported",
)
// make sure we have AND of ORs to have clear join terms
let where_ = normalize(where_.this_())
transforms_assert(normalized(where_), "Cannot normalize JOIN predicates")
// {name: list of join AND conditions}, in insertion order
let joins_ons : Map[String, Array[@core.Expr]] = {}
let conds = if where_.kind.is_a(And) {
where_.flatten().collect()
} else {
[where_]
}
for cond in conds {
let join_cols = cond
.find_all([Column])
.filter(col => col.has("join_mark"))
.collect()
let left_join_table : Array[String] = []
for col in join_cols {
let table = col.table_name()
if !left_join_table.contains(table) {
left_join_table.push(table)
}
}
if left_join_table.is_empty() {
continue
}
transforms_assert(
!(left_join_table.length() > 1),
"Cannot combine JOIN predicates from different tables",
)
for col in join_cols {
col.set("join_mark", false)
}
match joins_ons.get(left_join_table[0]) {
Some(l) => l.push(cond)
None => joins_ons[left_join_table[0]] = [cond]
}
}
let old_joins : Map[String, @core.Expr] = {}
for join in joins {
old_joins[join.alias_or_name()] = join
}
let new_joins : Map[String, @core.Expr] = {}
let query_from = match query.arg("from_") {
Some(f) => f
None => raise @core.ValueError("KeyError: 'from_'")
}
for table, predicates in joins_ons {
let join_what = old_joins.get(table).unwrap_or(query_from).this_().copy()
new_joins[join_what.alias_or_name()] = @core.mk(Join, [
("this", join_what),
("on", @core.exp_and(predicates)),
("kind", "LEFT"),
])
for p in predicates {
while p.parent is Some(pp) && pp.kind == Paren {
pp.replace(Some(p)) |> ignore
}
let parent = p.parent
p.pop() |> ignore
match parent {
Some(parent) if parent.kind.is_a(Binary) =>
match parent.arg("this") {
None => parent.replace(parent.arg("expression")) |> ignore
Some(left) => parent.replace(Some(left)) |> ignore
}
Some(parent) if parent.kind.is_a(Where) => parent.pop() |> ignore
_ => ()
}
}
}
if new_joins.contains(query_from.alias_or_name()) {
// Python takes the first element of the set difference `old_joins.keys() -
// new_joins.keys()`, whose order is unspecified; we use insertion order.
let only_old_joins = old_joins
.keys()
.filter(k => !new_joins.contains(k))
.collect()
transforms_assert(
only_old_joins.length() >= 1,
"Cannot determine which table to use in the new FROM clause",
)
let new_from_name = only_old_joins[0]
query.set("from_", @core.mk1(From, old_joins[new_from_name].this()))
}
if !new_joins.is_empty() {
// preserve any other joins
for n, j in old_joins {
if !new_joins.contains(n) && n != query.arg("from_").unwrap().name() {
if j.text("kind").is_empty() {
j.set("kind", "CROSS")
}
new_joins[n] = j
}
}
query.set("joins", new_joins.values().collect())
}
}
expression
}