// Module-level helpers of sqlglot/optimizer/simplify.py.
///|
let final_key : String = "final"
///|
/// Marks that an expression should not be simplified.
fn is_final(e : @core.Expr) -> Bool {
e.meta_bool(final_key)
}
///|
fn is_simplifiable(e : @core.Expr) -> Bool {
e.kind.is_any([Binary, Func, Lambda, Predicate, Unary])
}
///|
/// A AND (B AND C) -> A AND B AND C
fn flatten_connector(expression : @core.Expr) -> @core.Expr {
if expression.kind.is_a(Connector) {
for node in expression.iter_expressions() {
let child = node.unnest()
if child.kind == expression.kind {
node.replace(Some(child)) |> ignore
}
}
}
expression
}
///|
/// Propagate constants for conjunctions in DNF.
fn propagate_constants(expression : @core.Expr, root : Bool) -> @core.Expr {
if expression.kind.is_a(And) &&
(root || !expression.same_parent()) &&
normalized(expression, dnf=true) {
let constant_mapping : Array[(@core.Expr, @core.Expr, @core.Expr)] = []
for expr in walk_in_scope(expression, prune=n => n.kind.is_a(If)) {
if expr.kind.is_a(EQ) {
match (expr.this(), expr.expression()) {
(Some(l), Some(r)) =>
if l.kind.is_a(Column) &&
r.kind.is_a(Literal) &&
l.meta_get("nonnull") is Some(Bool(true)) {
// dict semantics: later equal keys overwrite the value, keep position
let mut found = false
for i, kv in constant_mapping {
if kv.0 == l {
constant_mapping[i] = (kv.0, l, r)
found = true
break
}
}
if !found {
constant_mapping.push((l, l, r))
}
}
_ => ()
}
}
}
if !constant_mapping.is_empty() {
for column in find_all_in_scope(expression, [Column]).collect() {
let parent = column.parent
for kv in constant_mapping {
if kv.0 == column {
let is_null_check = match parent {
Some(p) =>
p.kind.is_a(Is) &&
(match p.expression() {
Some(x) => x.kind.is_a(Null)
None => false
})
None => false
}
if !physical_equal(column, kv.1) && !is_null_check {
column.replace(Some(kv.2.copy())) |> ignore
}
break
}
}
}
}
}
expression
}
///|
fn is_number_expr(e : @core.Expr) -> Bool {
e.is_number()
}
///|
fn is_interval_expr(e : @core.Expr) -> Bool {
e.kind.is_a(Interval) && extract_interval(e) is Some(_)
}
///|
fn is_nonnull_constant(e : @core.Expr) -> Bool {
e.kind.is_any([Literal, Boolean]) || is_date_literal(e)
}
///|
fn is_constant_expr(e : @core.Expr) -> Bool {
let expr = if e.kind.is_a(Neg) {
match e.this() {
Some(t) => t
None => e
}
} else {
e
}
expr.kind.is_any([Literal, Boolean, Null]) || is_date_literal(expr)
}
///|
fn always_true(e : @core.Expr?) -> Bool {
match e {
Some(e) =>
(e.kind.is_a(Boolean) && e.has("this")) ||
(e.kind.is_a(Literal) && e.is_number() && !is_zero(Some(e)))
None => false
}
}
///|
fn always_false(e : @core.Expr?) -> Bool {
is_false(e) || is_null(e) || is_zero(e)
}
///|
fn is_zero(e : @core.Expr?) -> Bool {
match e {
Some(e) if e.kind.is_a(Literal) =>
match expr_to_pynum(e) {
Some(n) => n.is_zero()
None => false
}
_ => false
}
}
///|
fn is_false(e : @core.Expr?) -> Bool {
match e {
Some(e) => e.kind == Boolean && !e.has("this")
None => false
}
}
///|
fn is_null(e : @core.Expr?) -> Bool {
match e {
Some(e) => e.kind == Null
None => false
}
}
///|
fn boolean_literal(b : Bool) -> @core.Expr {
if b {
@core.true_()
} else {
@core.false_()
}
}
///|
/// Python `eval_boolean` given a comparison result (`cmp`) and equality (`eq`).
fn eval_boolean_cmp(expression : @core.Expr, cmp : () -> Int raise @core.SqlglotError, eq : () -> Bool) -> @core.Expr? raise @core.SqlglotError {
let k = expression.kind
if k.is_any([EQ, Is]) {
if k.is_a(Is) && expression.has("negate") {
return Some(boolean_literal(!eq()))
}
return Some(boolean_literal(eq()))
}
if k.is_a(NEQ) {
return Some(boolean_literal(!eq()))
}
if k.is_a(GT) {
return Some(boolean_literal(cmp() > 0))
}
if k.is_a(GTE) {
return Some(boolean_literal(cmp() >= 0))
}
if k.is_a(LT) {
return Some(boolean_literal(cmp() < 0))
}
if k.is_a(LTE) {
return Some(boolean_literal(cmp() <= 0))
}
None
}
///|
/// A value that `cast_value` operates on: a string, or a date/datetime.
priv enum DateValue {
DVStr(String)
DVDate(PyDT)
}
///|
fn cast_as_date(value : DateValue) -> PyDT? {
match value {
DVDate(d) => Some(d.to_date())
DVStr(s) => parse_py_datetime(s).map(d => d.to_date())
}
}
///|
fn cast_as_datetime(value : DateValue) -> PyDT? {
match value {
DVDate(d) => Some(d.to_datetime())
DVStr(s) => parse_py_datetime(s)
}
}
///|
fn cast_value(value : DateValue, to : @core.Expr) -> PyDT? {
match value {
DVStr("") => return None
_ => ()
}
if to.is_type([DATE]) {
return cast_as_date(value)
}
if to.is_type(@core.dtype_temporal_types) {
return cast_as_datetime(value)
}
None
}
///|
fn extract_date(cast : @core.Expr) -> PyDT? {
let to = if cast.kind.is_a(Cast) {
match cast.arg("to") {
Some(t) => t
None => return None
}
} else if cast.kind.is_a(TsOrDsToDate) && !cast.has("format") {
@core.datatype_of(DATE)
} else {
return None
}
let this = match cast.this() {
Some(t) => t
None => return None
}
let value = if this.kind.is_a(Literal) {
DVStr(this.name())
} else if this.kind.is_any([Cast, TsOrDsToDate]) {
match extract_date(this) {
Some(d) => DVDate(d)
None => return None
}
} else {
return None
}
cast_value(value, to)
}
///|
fn is_date_literal(e : @core.Expr) -> Bool {
extract_date(e) is Some(_)
}
///|
fn extract_interval(expression : @core.Expr) -> RelDelta? {
let this = match expression.this() {
Some(t) => t
None => return None
}
let n = match this.kind {
Literal =>
if this.is_number() {
match expr_to_pynum(this) {
Some(PInt(i)) => i
Some(PDec(d)) =>
// int(Decimal) truncates
{
let v = d.coeff
let truncated = if d.exp >= 0 {
v * pow10(d.exp)
} else {
v / pow10(-d.exp)
}
let t = truncated
if d.neg {
-t
} else {
t
}
}
None => return None
}
} else {
match parse_py_int(this.text("this")) {
Some(i) => i
None => return None
}
}
Neg =>
match expr_to_pynum(this) {
Some(PInt(i)) => i
_ => return None
}
_ => return None
}
let unit = @core.py_lower(expression.text("unit"))
interval_delta(unit, n~)
}
///|
fn is_exact_interval_move(
op : @core.Expr,
literal : @core.Expr,
interval : @core.Expr,
) -> Bool raise @core.SqlglotError {
let mut delta = match extract_interval(interval) {
Some(d) => d
None => return false
}
let value = match extract_date(literal) {
Some(v) => v
None => return false
}
if !delta.moves_months() {
return true
}
if op.kind.is_any([Sub, DateSub, DatetimeSub]) {
delta = delta.neg()
}
let moved = add_reldelta(value, delta.neg())
let back = add_reldelta(moved, delta)
let next = add_reldelta(moved.add_timedelta(1, 0L, 0L), delta)
pydt_eq(back, value) && !pydt_eq(next, value)
}
///|
fn extract_type(expressions : Array[@core.Expr]) -> @core.Expr? {
let mut target_type = None
for expression in expressions {
target_type = if expression.kind.is_a(Cast) {
expression.arg("to")
} else {
expression.get_type()
}
if target_type is Some(_) {
break
}
}
target_type
}
///|
fn date_literal(date : PyDT, target_type : @core.Expr?) -> @core.Expr {
let to = match target_type {
Some(t) if t.is_type(@core.dtype_temporal_types) => t.copy()
_ => @core.datatype_of(if date.is_datetime { DATETIME } else { DATE })
}
let cast = @core.mk(Cast, [
("this", @core.literal_string(date.to_py_string())),
("to", to),
])
cast.set_type(Some(to))
cast
}
///|
priv suberror UnsupportedUnit {
UnsupportedUnit
}
///|
fn interval_or_raise(unit : String) -> RelDelta raise UnsupportedUnit {
match interval_delta(unit) {
Some(d) => d
None => raise UnsupportedUnit
}
}
///|
fn floor_or_raise(
d : PyDT,
unit : String,
dialect : @core.Dialect,
) -> PyDT raise UnsupportedUnit {
match datetime_floor(d, unit, dialect.cfg.week_offset) {
Some(r) => r
None => raise UnsupportedUnit
}
}
///|
fn trunc_unit(unit : @core.Expr, dialect : @core.Dialect) -> String raise UnsupportedUnit {
if unit.kind.is_a(WeekStart) {
let dow = match @core.py_upper(unit.name()) {
"MONDAY" => 1
"TUESDAY" => 2
"WEDNESDAY" => 3
"THURSDAY" => 4
"FRIDAY" => 5
"SATURDAY" => 6
"SUNDAY" => 7
_ => -1
}
let offset = dialect.cfg.week_offset
let expected = (offset % 7 + 7) % 7 + 1
if dow != expected {
raise UnsupportedUnit
}
return "week"
}
@core.py_lower(unit.name())
}
///|
fn date_ceil(
d : PyDT,
unit : String,
dialect : @core.Dialect,
) -> PyDT raise UnsupportedUnit {
let floor = floor_or_raise(d, unit, dialect)
if pydt_eq(floor, d) {
return d
}
add_reldelta(floor, interval_or_raise(unit)) catch {
_ => raise UnsupportedUnit
}
}
///|
fn datetrunc_range(
date : PyDT,
unit : String,
dialect : @core.Dialect,
) -> (PyDT, PyDT)? raise UnsupportedUnit {
let floor = floor_or_raise(date, unit, dialect)
if !pydt_eq(date, floor) {
return None
}
let upper = add_reldelta(floor, interval_or_raise(unit)) catch {
_ => raise UnsupportedUnit
}
Some((floor, upper))
}
///|
fn ge_(l : @core.Expr, r : @core.Expr) -> @core.Expr {
binop(GTE, l, r)
}
///|
fn lt_(l : @core.Expr, r : @core.Expr) -> @core.Expr {
binop(LT, l, r)
}
///|
/// Python `Expr._binop(klass, other)`: copies both operands, wrapping binaries in parens.
fn binop(kind : @core.Kind, a : @core.Expr, b : @core.Expr) -> @core.Expr {
let mut this = a.copy()
let mut other = b.copy()
if !this.kind.is_a(kind) && !other.kind.is_a(kind) {
if this.kind.is_a(Binary) {
this = @core.paren(this)
}
if other.kind.is_a(Binary) {
other = @core.paren(other)
}
}
@core.mk2(kind, this, other)
}
///|
fn datetrunc_eq_expression(
left : @core.Expr,
drange : (PyDT, PyDT),
target_type : @core.Expr?,
) -> @core.Expr {
@core.and_(
[
ge_(left, date_literal(drange.0, target_type)),
lt_(left, date_literal(drange.1, target_type)),
],
copy=false,
)
}
///|
fn parenthesize_nested_connector(
expression : @core.Expr,
parent : @core.Expr?,
) -> @core.Expr {
if expression.kind.is_a(Connector) &&
(match parent {
Some(p) =>
p.kind.is_a(Not) || (p.kind.is_a(Connector) && p.kind != expression.kind)
None => false
}) {
return @core.paren(expression)
}
expression
}
///|
/// Python `helper.merge_ranges`.
fn merge_ranges(
ranges : Array[(PyDT, PyDT)],
) -> Array[(PyDT, PyDT)] raise @core.SqlglotError {
if ranges.is_empty() {
return []
}
let sorted = ranges.copy()
// insertion sort with Python tuple ordering
for i in 1.. 0 {
let a = sorted[j - 1]
let b = sorted[j]
let c = pydt_cmp(a.0, b.0)
let greater = c > 0 || (c == 0 && pydt_cmp(a.1, b.1) > 0)
if !greater {
break
}
sorted[j - 1] = b
sorted[j] = a
j -= 1
}
}
let merged = [sorted[0]]
for i in 1.. 0 { end } else { last_end }
merged[merged.length() - 1] = (last_start, m)
} else {
merged.push((start, end))
}
}
merged
}