// Port of sqlglot/typing/__init__.py: the base dialect's expression metadata, and the
// registry of per-dialect metadata (Python `Dialect.EXPRESSION_METADATA` / `COERCES_TO`).
///|
fn returns_all(
m : ExprMetadata,
kinds : Array[@core.Kind],
dtype : @core.DType,
) -> Unit {
for k in kinds {
m[k] = Returns(D(dtype))
}
}
///|
fn annotate_all(
m : ExprMetadata,
kinds : Array[@core.Kind],
f : (TypeAnnotator, @core.Expr) -> Unit raise @core.SqlglotError,
) -> Unit {
for k in kinds {
m[k] = Annotator(f)
}
}
///|
/// Python `subclasses(exp.__name__, classes)`: every expression class that is a subclass
/// of one of `bases` (including the bases themselves).
fn subclasses(bases : Array[@core.Kind]) -> Array[@core.Kind] {
let out = []
for b in bases {
for k in @core.subclasses_of(b) {
if !out.contains(k) {
out.push(k)
}
}
if !b.is_trait() && !out.contains(b) {
out.push(b)
}
}
out
}
///|
let timestamp_expressions : Array[@core.Kind] = [
CurrentTimestamp, StrToTime, TimeStrToTime, TimestampAdd, TimestampSub, UnixToTime,
]
///|
fn build_base_expression_metadata() -> ExprMetadata {
let m : ExprMetadata = {}
annotate_all(m, subclasses([Binary]), (s, e) => s.annotate_binary(e))
annotate_all(m, subclasses([Unary, Alias, IgnoreNulls, RespectNulls]), (s, e) => s.annotate_unary(
e,
))
returns_all(
m,
[
ApproxDistinct, ArraySize, CountIf, DenseRank, Int64, Ntile, Rank, RowNumber,
UnixSeconds, UnixMicros, UnixMillis,
],
BIGINT,
)
returns_all(m, [FromBase32, FromBase64], BINARY)
returns_all(
m,
[
All, Any, Between, Boolean, Contains, EndsWith, Exists, In, IsInf, IsNan, LogicalAnd,
LogicalOr, StartsWith,
],
BOOLEAN,
)
returns_all(
m,
[
CurrentDate, Date, DateFromParts, DateStrToDate, DiToDate, LastDay, StrToDate,
TimeStrToDate, TsOrDsToDate,
],
DATE,
)
returns_all(m, [CurrentDatetime, Datetime, DatetimeAdd, DatetimeSub], DATETIME)
returns_all(
m,
[
Asin, Asinh, Acos, CovarPop, CovarSamp, Acosh, ApproxQuantile, Atan, Atanh, Avg,
Cbrt, Cos, Cosh, Cot, Degrees, Exp, Kurtosis, Ln, Log, Pi, Pow, PercentileCont, Quantile,
Radians, Round, SafeDivide, Sin, Sinh, Sqrt, Stddev, StddevPop, StddevSamp, Rand,
Tan, Tanh, ToDouble, CumeDist, PercentRank, Variance, VariancePop, Skewness,
],
DOUBLE,
)
returns_all(
m,
[
Ascii, BitLength, Ceil, DatetimeDiff, DayOfMonth, DayOfWeek, DayOfYear, Floor, Getbit,
Hour, TimestampDiff, TimeDiff, Unicode, DateToDi, Levenshtein, Length, Sign, StrPosition,
TsOrDiToDi, Quarter, UnixDate,
],
INT,
)
returns_all(
m,
[Interval, JustifyDays, JustifyHours, JustifyInterval, MakeInterval],
INTERVAL,
)
returns_all(m, [ParseJSON], JSON)
returns_all(m, [CurrentTime, Localtime, Time, TimeAdd, TimeSub], TIME)
returns_all(m, [TimestampLtzFromParts], TIMESTAMPLTZ)
returns_all(m, [CurrentTimestampLTZ, TimestampTzFromParts], TIMESTAMPTZ)
returns_all(m, timestamp_expressions, TIMESTAMP)
returns_all(
m,
[Day, DayOfWeekIso, Month, Week, WeekOfYear, Year, YearOfWeek, YearOfWeekIso],
TINYINT,
)
returns_all(
m,
[
ArrayToString, Concat, ConcatWs, Chr, CurrentCatalog, CurrentRole, CurrentSchema,
CurrentVersion, CurrentUser, Dayname, DateToDateStr, DPipe, GroupConcat, Initcap,
Lower, MD5, Monthname, RawString, Repeat, SHA, SHA2, SessionUser, Space, String, Substring,
TimeToStr, TimeToTimeStr, Trim, ToBase32, ToBase64, Translate, TsOrDsToDateStr, Typeof,
UnixToStr, UnixToTimeStr, Upper,
],
VARCHAR,
)
annotate_all(
m,
[
Abs, AnyValue, ArrayConcatAgg, ArrayReverse, ArraySlice, Filter, FirstValue, HavingMax,
LastValue, Limit, NthValue, Order, SortArray, Window,
],
(s, e) => s.annotate_by_args(e, [Key("this")]),
)
annotate_all(
m,
[ArrayConcat, Coalesce, Greatest, Least, Max, Min],
(s, e) => s.annotate_by_args(e, [Key("this"), Key("expressions")]),
)
annotate_all(m, [ArrayFirst, ArrayLast], (s, e) => s.annotate_by_array_element(
e,
))
m[Anonymous] = Annotator((s, e) => s.set_type(
e,
Some(T(s.schema.get_udf_type(e))),
))
annotate_all(m, [DateAdd, DateSub, DateTrunc], (s, e) => s.annotate_timeunit(e))
annotate_all(m, [Cast, TryCast], (s, e) => s.set_type_expr(e, e.arg("to")))
annotate_all(m, [Map, VarMap], (s, e) => s.annotate_map(e))
m[Array] = Annotator((s, e) => s.annotate_by_args(
e,
[Key("expressions")],
array=true,
))
m[ArrayAgg] = Annotator((s, e) => s.annotate_by_args(e, [Key("this")], array=true))
m[Bracket] = Annotator((s, e) => s.annotate_bracket(e))
m[Case] = Annotator((s, e) => {
let args = e.list("ifs").map(if_expr => ArgRef::Node(if_expr.arg("true").unwrap()))
args.push(Key("default"))
s.annotate_by_args(e, args)
})
m[Count] = Annotator((s, e) => s.set_dtype(
e,
if e.has("big_int") {
BIGINT
} else {
INT
},
))
m[DateDiff] = Annotator((s, e) => s.set_dtype(
e,
if e.has("big_int") {
BIGINT
} else {
INT
},
))
m[DataType] = Annotator((_, _) => ())
m[Div] = Annotator((s, e) => s.annotate_div(e))
m[Distinct] = Annotator((s, e) => s.annotate_by_args(e, [Key("expressions")]))
m[Dot] = Annotator((s, e) => s.annotate_dot(e))
m[Explode] = Annotator((s, e) => s.annotate_explode(e))
m[Extract] = Annotator((s, e) => s.annotate_extract(e))
m[HexString] = Annotator((s, e) => s.set_dtype(
e,
if e.has("is_integer") {
BIGINT
} else {
BINARY
},
))
m[GenerateSeries] = Annotator((s, e) => s.annotate_by_args(
e,
[Key("start"), Key("end"), Key("step")],
array=true,
))
m[GenerateDateArray] = Annotator((s, e) => s.set_type_expr(
e,
Some(@core.datatype_from_str("ARRAY")),
))
m[GenerateTimestampArray] = Annotator((s, e) => s.set_type_expr(
e,
Some(@core.datatype_from_str("ARRAY")),
))
m[If] = Annotator((s, e) => s.annotate_by_args(e, [Key("true"), Key("false")]))
m[Lag] = Annotator((s, e) => s.annotate_by_args(e, [Key("this"), Key("default")]))
m[Lead] = Annotator((s, e) => s.annotate_by_args(e, [Key("this"), Key("default")]))
m[Literal] = Annotator((s, e) => s.annotate_literal(e))
m[Null] = Returns(D(NULL))
m[Nullif] = Annotator((s, e) => s.annotate_by_args(e, [
Key("this"),
Key("expression"),
]))
m[PropertyEQ] = Annotator((s, e) => s.annotate_by_args(e, [Key("expression")]))
m[Struct] = Annotator((s, e) => s.annotate_struct(e))
m[Sum] = Annotator((s, e) => s.annotate_by_args(
e,
[Key("this"), Key("expressions")],
promote=true,
))
m[Timestamp] = Annotator((s, e) => s.set_dtype(
e,
if e.has("with_tz") {
TIMESTAMPTZ
} else {
TIMESTAMP
},
))
m[ToMap] = Annotator((s, e) => s.annotate_to_map(e))
m[Unnest] = Annotator((s, e) => s.annotate_unnest(e))
m[WithinGroup] = Annotator((s, e) => s.annotate_within_group(e))
m[Subquery] = Annotator((s, e) => s.annotate_subquery(e))
m
}
///|
/// The base dialect's expression metadata (Python `sqlglot.typing.EXPRESSION_METADATA`).
pub let base_expression_metadata : ExprMetadata = build_base_expression_metadata()
///|
/// Builders of per-dialect expression metadata, keyed by dialect name.
let expression_metadata_builders : Map[String, () -> ExprMetadata] = {}
///|
let expression_metadata_cache : Map[String, ExprMetadata] = {}
///|
/// Per-dialect `COERCES_TO`, keyed by dialect name.
let coerces_to_builders : Map[
String,
() -> Map[@core.DType, @set.Set[@core.DType]],
] = {}
///|
let coerces_to_cache : Map[String, Map[@core.DType, @set.Set[@core.DType]]] = {}
///|
/// Registers the expression metadata of a dialect (by name). Dialects that don't register
/// metadata inherit the metadata of their parent dialect.
pub fn register_expression_metadata(
dialect_name : String,
build : () -> ExprMetadata,
) -> Unit {
expression_metadata_builders[dialect_name] = build
expression_metadata_cache.remove(dialect_name)
}
///|
/// Registers the `COERCES_TO` mapping of a dialect (by name).
pub fn register_coerces_to(
dialect_name : String,
build : () -> Map[@core.DType, @set.Set[@core.DType]],
) -> Unit {
coerces_to_builders[dialect_name] = build
coerces_to_cache.remove(dialect_name)
}
///|
/// The expression metadata of `dialect` (Python `dialect.EXPRESSION_METADATA`).
pub fn dialect_expression_metadata(dialect : @core.Dialect) -> ExprMetadata {
ensure_typing_registered()
let mut d : @core.Dialect? = Some(dialect)
while d is Some(x) {
match expression_metadata_cache.get(x.name) {
Some(m) => return m
None => ()
}
match expression_metadata_builders.get(x.name) {
Some(b) => {
let m = b()
expression_metadata_cache[x.name] = m
return m
}
None => ()
}
d = x.parent
}
base_expression_metadata
}
///|
/// The `COERCES_TO` of `dialect` (empty when the dialect uses the default one).
pub fn dialect_coerces_to(
dialect : @core.Dialect,
) -> Map[@core.DType, @set.Set[@core.DType]] {
ensure_typing_registered()
let mut d : @core.Dialect? = Some(dialect)
while d is Some(x) {
match coerces_to_cache.get(x.name) {
Some(m) => return m
None => ()
}
match coerces_to_builders.get(x.name) {
Some(b) => {
let m = b()
coerces_to_cache[x.name] = m
return m
}
None => ()
}
d = x.parent
}
{}
}
///|
let typing_registered : Ref[Bool] = Ref(false)
///|
/// Registers the built-in per-dialect typing rules (see typing_*.mbt).
fn ensure_typing_registered() -> Unit {
if typing_registered.val {
return
}
typing_registered.val = true
register_dialect_typings()
}
///|
/// Copies the metadata of a parent and applies updates (Python `{**parent, ...}`).
pub fn extend_metadata(parent : ExprMetadata) -> ExprMetadata {
parent.copy()
}