// Port of sqlglot/optimizer/resolver.py.
///|
/// Helper for resolving columns.
pub struct Resolver {
scope : Scope
schema : MappingSchema
dialect : @core.Dialect
priv mut source_columns_ : Map[String, Array[String]]?
priv mut unambiguous_columns_ : Map[String, String]?
priv mut all_columns_ : @set.Set[String]?
priv infer_schema : Bool
priv get_source_columns_cache : Map[(String, Bool), Array[String]]
priv column_type_from_scope_cache : Map[(Int, String), @core.Expr?]
}
///|
pub fn Resolver::new(
scope : Scope,
schema : MappingSchema,
infer_schema? : Bool = true,
) -> Resolver {
{
scope,
schema,
dialect: schema.dialect,
source_columns_: None,
unambiguous_columns_: None,
all_columns_: None,
infer_schema,
get_source_columns_cache: {},
column_type_from_scope_cache: {},
}
}
///|
/// Get the table for a column name.
pub fn Resolver::get_table_by_name(
self : Resolver,
column_name : String,
) -> @core.Expr? raise @core.SqlglotError {
self.get_table_impl(column_name, None)
}
///|
/// Get the table for a column.
pub fn Resolver::get_table(
self : Resolver,
column : @core.Expr,
) -> @core.Expr? raise @core.SqlglotError {
self.get_table_impl(column.name(), Some(column))
}
///|
fn Resolver::get_table_impl(
self : Resolver,
column_name : String,
column : @core.Expr?,
) -> @core.Expr? raise @core.SqlglotError {
let join_context = match column {
Some(c) if c.kind.is_a(Column) => self.get_column_join_context(c)
_ => None
}
let is_semi_or_anti = match join_context {
Some(j) => is_semi_or_anti_join(j)
None => false
}
let mut table_name : String? = if is_semi_or_anti {
None
} else {
self.get_table_name_from_sources(column_name, None)
}
if table_name is None && join_context is Some(jc) {
table_name = self.get_table_name_from_sources(
column_name,
Some(self.get_available_source_columns(jc)),
) catch {
@core.OptimizeError(_) => None
e => raise e
}
}
if table_name is None && self.infer_schema {
let sources_without_schema = []
for source, columns in self.get_all_source_columns() {
if columns.is_empty() || columns.contains("*") {
sources_without_schema.push(source)
}
}
if sources_without_schema.length() == 1 {
table_name = Some(sources_without_schema[0])
}
}
let table_name = match table_name {
Some(t) => t
None => return None
}
let selected = self.scope.selected_sources()
match selected.get(table_name) {
None => Some(@core.to_identifier(table_name))
Some((node, _)) => {
let mut node = node
if node.kind.is_a(Query) {
while node.alias() != table_name && node.parent is Some(p) {
node = p
}
}
match node.arg("alias") {
Some(node_alias) =>
match node_alias.this() {
Some(t) => Some(t.copy())
None => Some(@core.to_identifier(node_alias.name()))
}
None => Some(@core.to_identifier(table_name))
}
}
}
}
///|
/// Resolvers for the outer scopes a correlated subquery can reference, innermost first.
pub fn Resolver::outer_resolvers(self : Resolver) -> Array[Resolver] {
let out = []
let mut scope = self.scope
while scope.can_be_correlated && scope.parent is Some(p) {
scope = p
out.push(Resolver::new(scope, self.schema, infer_schema=self.infer_schema))
}
out
}
///|
/// Whether some source's columns can't be determined.
pub fn Resolver::has_unknown_sources(
self : Resolver,
) -> Bool raise @core.SqlglotError {
for _, columns in self.get_all_source_columns() {
if columns.is_empty() || columns.contains("*") {
return true
}
}
false
}
///|
/// All available columns of all sources in this scope.
pub fn Resolver::all_columns(
self : Resolver,
) -> @set.Set[String] raise @core.SqlglotError {
match self.all_columns_ {
Some(c) => c
None => {
let s : @set.Set[String] = @set.new()
for _, columns in self.get_all_source_columns() {
for c in columns {
s.add(c)
}
}
self.all_columns_ = Some(s)
s
}
}
}
///|
fn dedup_strings(xs : Array[String]) -> Array[String] {
let seen : @set.Set[String] = @set.new()
let out = []
for x in xs {
if !seen.contains(x) {
seen.add(x)
out.push(x)
}
}
out
}
///|
pub fn Resolver::get_source_columns_from_set_op(
self : Resolver,
expression : @core.Expr,
) -> Array[String] raise @core.SqlglotError {
if expression.kind.is_a(Select) {
return expression.named_selects()
}
if expression.kind.is_a(Subquery) {
return self.get_source_columns_from_set_op(expression.unnest())
}
if !expression.kind.is_a(SetOperation) {
raise @core.OptimizeError("Unknown set operation: \{expr_sql(expression)}")
}
let set_op = expression
let on_column_list = set_op.list("on")
if !on_column_list.is_empty() {
on_column_list.map(c => c.name())
} else {
let side = @core.py_upper(set_op.text("side"))
let kind = @core.py_upper(set_op.text("kind"))
if side != "" || kind != "" {
let left = self.get_source_columns_from_set_op(set_op.this_())
let right = self.get_source_columns_from_set_op(set_op.expression_())
if side == "LEFT" {
left
} else if side == "FULL" {
dedup_strings(left + right)
} else if kind == "INNER" {
// dict keys intersection: Python set semantics, order of the left operand
let r : @set.Set[String] = @set.from_array(right)
dedup_strings(left).filter(x => r.contains(x))
} else {
// Python leaves `columns` unbound here
raise @core.OptimizeError("Unknown set operation: \{expr_sql(expression)}")
}
} else {
set_op.named_selects()
}
}
}
///|
/// Resolve the source columns for a given source `name`.
pub fn Resolver::get_source_columns(
self : Resolver,
name : String,
only_visible? : Bool = false,
) -> Array[String] raise @core.SqlglotError {
let cache_key = (name, only_visible)
match self.get_source_columns_cache.get(cache_key) {
Some(c) => return c
None => ()
}
let mut source = match self.scope.sources.get(name) {
Some(s) => s
None => raise @core.OptimizeError("Unknown table: \{name}")
}
match source {
TableSource(t) if t.db() == "" &&
t.has("pivots") &&
self.scope.cte_sources.contains(t.name()) =>
source = self.scope.cte_sources[t.name()]
_ => ()
}
let mut columns : Array[String] = match source {
TableSource(t) => self.schema.column_names(t, only_visible~)
ScopeSource(s) if s.expression.kind.is_any([Values, Unnest, Lateral]) => {
let source_expr = s.expression
let mut columns = source_expr.named_selects()
if self.dialect.cfg.unnest_column_only && source_expr.kind.is_a(Unnest) {
if source_expr.get_type() is None ||
type_is(source_expr.get_type(), [UNKNOWN]) {
match source_expr.expressions().get(0) {
Some(unnest_expr) if unnest_expr.kind.is_a(Column) &&
self.scope.parent is Some(parent) => {
let col_type = self.get_unnest_column_type(unnest_expr, parent)
match col_type {
Some(ct) =>
if ct.is_type([ARRAY]) {
let element_types = ct.expressions()
if !element_types.is_empty() {
source_expr.set_type(Some(element_types[0].copy()))
}
} else {
source_expr.set_type(Some(ct.copy()))
}
None => ()
}
}
_ => ()
}
}
columns = columns + struct_field_names(source_expr.get_type())
} else if source_expr.kind.is_a(Lateral) &&
(match source_expr.this() {
Some(t) => t.kind.is_a(Explode)
None => false
}) {
let explode_col = source_expr.this_().this()
match explode_col {
Some(ec) if ec.kind.is_a(Column) &&
ec.table_name() != "" &&
s.parent is Some(sp) => {
let col_type = self.get_unnest_column_type(ec, sp)
columns = columns + struct_field_names(col_type)
}
_ => ()
}
} else if source_expr.kind.is_a(Lateral) &&
(match source_expr.this() {
Some(t) => t.kind.is_a(Query)
None => false
}) {
columns = named_selects_of(source_expr.this_())
}
columns
}
ScopeSource(s) if s.expression.kind.is_a(SetOperation) =>
self.get_source_columns_from_set_op(s.expression)
ScopeSource(s) => {
let selectable = s.expression
match selects_of(selectable).get(0) {
Some(select) if select.kind.is_a(QueryTransform) =>
match select.arg("schema") {
Some(schema) => schema.expressions().map(c => c.name())
None => ["key", "value"]
}
_ => named_selects_of(selectable)
}
}
}
let column_aliases = match self.scope.selected_sources().get(name) {
Some((node, _)) => node.alias_column_names()
None => []
}
if !column_aliases.is_empty() {
let n = @core.max_int(columns.length(), column_aliases.length())
let out = []
for i in 0.. out.push(a)
_ =>
match columns.get(i) {
Some(c) => out.push(c)
None => out.push("")
}
}
}
columns = out
}
self.get_source_columns_cache[cache_key] = columns
columns
}
///|
fn Resolver::get_all_source_columns(
self : Resolver,
) -> Map[String, Array[String]] raise @core.SqlglotError {
match self.source_columns_ {
Some(s) => return s
None => ()
}
let result : Map[String, Array[String]] = {}
for source_name, _ in self.scope.selected_sources() {
result[source_name] = self.get_source_columns(source_name)
}
for source_name, _ in self.scope.lateral_sources {
result[source_name] = self.get_source_columns(source_name)
}
self.source_columns_ = Some(result)
result
}
///|
fn Resolver::get_table_name_from_sources(
self : Resolver,
column_name : String,
source_columns : Map[String, Array[String]]?,
) -> String? raise @core.SqlglotError {
let unambiguous_columns = match source_columns {
Some(sc) if !sc.is_empty() => self.get_unambiguous_columns(sc)
_ =>
match self.unambiguous_columns_ {
Some(u) => u
None => {
let u = self.get_unambiguous_columns(self.get_all_source_columns())
self.unambiguous_columns_ = Some(u)
u
}
}
}
unambiguous_columns.get(column_name)
}
///|
fn Resolver::get_column_join_context(
self : Resolver,
column : @core.Expr,
) -> @core.Expr? {
let e = self.scope.expression
if !e.has("joins") || e.has("laterals") || e.has("pivots") {
return None
}
match column.find_ancestor([Join, Select]) {
Some(j) if j.kind.is_a(Join) => {
let join_name = j.alias_or_name()
if self.scope.selected_sources_or_empty().contains(join_name) ||
self.scope.semi_or_anti_join_tables().contains(join_name) {
Some(j)
} else {
None
}
}
_ => None
}
}
///|
fn Resolver::get_available_source_columns(
self : Resolver,
join_ancestor : @core.Expr,
) -> Map[String, Array[String]] raise @core.SqlglotError {
let e = self.scope.expression
let from_name = match e.arg("from_") {
Some(f) => f.alias_or_name()
None => raise @core.OptimizeError("KeyError: from_")
}
let available : Map[String, Array[String]] = {}
available[from_name] = self.get_source_columns(from_name)
let joins = e.list("joins")
let upto = match join_ancestor.index {
Some(i) => i + 1
None => 0
}
for i in 0..<@core.min_int(upto, joins.length()) {
let join = joins[i]
available[join.alias_or_name()] = self.get_source_columns(join.alias_or_name())
}
available
}
///|
fn Resolver::get_unambiguous_columns(
self : Resolver,
source_columns : Map[String, Array[String]],
) -> Map[String, String] {
if source_columns.is_empty() {
return {}
}
let pairs = source_columns.to_array()
let (first_table, first_columns) = pairs[0]
let unambiguous_columns : Map[String, String] = {}
for col in first_columns {
unambiguous_columns[col] = first_table
}
if pairs.length() == 1 {
return unambiguous_columns
}
let unnest_original_aliases : Map[String, String] = {}
if self.dialect.cfg.unnest_column_only {
for source_name, source in self.scope.sources {
match source.expression() {
Some(se) if se.kind.is_a(Unnest) =>
match se.arg("alias") {
Some(alias_arg) => {
let cols = alias_arg.list("columns")
if !cols.is_empty() {
unnest_original_aliases[cols[0].name()] = source_name
}
}
None => ()
}
_ => ()
}
}
}
let all_columns : @set.Set[String] = @set.new()
for c in unambiguous_columns.keys() {
all_columns.add(c)
}
for i in 1.. all_columns.contains(c))
let ambiguous_set : @set.Set[String] = @set.from_array(ambiguous)
for c in columns {
all_columns.add(c)
}
for column in ambiguous {
match unnest_original_aliases.get(column) {
Some(s) => {
unambiguous_columns[column] = s
continue
}
None => ()
}
unambiguous_columns.remove(column)
}
for column in unique {
if !ambiguous_set.contains(column) {
unambiguous_columns[column] = table
}
}
}
unambiguous_columns
}
///|
fn struct_field_names(col_type : @core.Expr?) -> Array[String] {
let mut col_type = col_type
if type_is(col_type, [ARRAY]) {
col_type = col_type.unwrap().expressions().get(0)
}
match col_type {
Some(ct) if ct.is_type([STRUCT]) => ct.expressions().map(k => k.name())
_ => []
}
}
///|
fn Resolver::get_unnest_column_type(
self : Resolver,
column : @core.Expr,
scope : Scope,
) -> @core.Expr? raise @core.SqlglotError {
let table_name = if column.table_name() != "" {
column.table_name()
} else {
let parent_resolver = Resolver::new(
scope,
self.schema,
infer_schema=self.infer_schema,
)
match parent_resolver.get_table(column) {
Some(t) => t.name()
None => return None
}
}
match scope.sources.get(table_name) {
Some(source) => self.get_column_type_from_scope(source, column)
None => None
}
}
///|
/// The number of calls of `Resolver::get_column_type_from_scope` (Python
/// `Resolver._get_column_type_from_scope`) so far; tests use it to check that the
/// trace is memoized (Python's test patches the method to count its calls).
pub let column_type_trace_calls : Ref[Int] = Ref(0)
///|
fn Resolver::get_column_type_from_scope(
self : Resolver,
source : Source,
column : @core.Expr,
) -> @core.Expr? raise @core.SqlglotError {
column_type_trace_calls.val += 1
let source_id = match source {
TableSource(t) => t.uid * 2
ScopeSource(s) => s.id * 2 + 1
}
let cache_key = (source_id, column.name())
match self.column_type_from_scope_cache.get(cache_key) {
Some(r) => return r
None => ()
}
let mut result : @core.Expr? = None
match source {
TableSource(t) => {
let col_type = self.schema.get_column_type(t, column)
if !col_type.is_type([UNKNOWN]) {
result = Some(col_type)
}
}
ScopeSource(s) =>
for _, nested_source in s.sources {
let nested_type = self.get_column_type_from_scope(nested_source, column)
match nested_type {
Some(nt) if !nt.is_type([UNKNOWN]) => {
result = Some(nt)
break
}
_ => ()
}
}
}
self.column_type_from_scope_cache[cache_key] = result
result
}