// Port of sqlglot/optimizer/canonicalize.py.

///|
/// Python `exp.replace_tree`: replaces the tree with the results of `fun` on each node,
/// leaves first; new nodes are traversed too.
pub fn replace_tree(
  expression : @core.Expr,
  fun : (@core.Expr) -> @core.Expr raise @core.SqlglotError,
  prune? : (@core.Expr) -> Bool,
) -> @core.Expr raise @core.SqlglotError {
  let stack = expression.dfs(prune?).collect()
  let mut new_node = expression
  while stack.pop() is Some(node) {
    new_node = fun(node)
    if !physical_equal(new_node, node) {
      node.replace(Some(new_node)) |> ignore
      stack.push(new_node)
    }
  }
  new_node
}

///|
let canonicalize_kinds : Array[@core.Kind] = [
  Add, Date, TsOrDsToDate, Timestamp, Sub, EQ, NEQ, GT, GTE, LT, LTE, NullSafeEQ, NullSafeNEQ,
  Between, Extract, DateAdd, DateSub, DateTrunc, DateDiff, Cast, Connector, Not, If, Where,
  Having, Ordered,
]

///|
let coercible_date_ops : Array[@core.Kind] = [
  Add, Sub, EQ, NEQ, GT, GTE, LT, LTE, NullSafeEQ, NullSafeNEQ,
]

///|
/// Converts a sql expression into a standard form.
pub fn canonicalize(
  expression : @core.Expr,
  dialect? : @core.Dialect,
) -> @core.Expr raise @core.SqlglotError {
  let dialect = get_dialect(dialect)
  replace_tree(expression, e => {
    if !e.kind.is_any(canonicalize_kinds) {
      return e
    }
    let mut e = add_text_to_concat(e)
    e = replace_date_funcs(e, dialect)
    e = coerce_type(e, dialect.cfg.promote_to_inferred_datetime_type)
    e = remove_redundant_casts(e)
    e = canonicalize_ensure_bools(e, replace_int_predicate)
    e = remove_ascending_order(e)
    e
  })
}

///|
fn type_in(e : @core.Expr, set : Array[@core.DType]) -> Bool {
  match e.get_type() {
    Some(t) =>
      match t.datatype_this() {
        Some(d) => set.contains(d)
        None => false
      }
    None => false
  }
}

///|
fn add_text_to_concat(node : @core.Expr) -> @core.Expr {
  if node.kind.is_a(Add) && type_in(node, @core.dtype_text_types) {
    return @core.mk(Concat, [
      ("expressions", [node.this_(), node.expression_()]),
      ("coalesce", false),
    ])
  }
  node
}

///|
fn replace_date_funcs(
  node : @core.Expr,
  dialect : @core.Dialect,
) -> @core.Expr raise @core.SqlglotError {
  if node.kind.is_any([Date, TsOrDsToDate]) &&
    node.expressions().is_empty() &&
    !node.has("zone") &&
    req(node.this(), "is_string").is_string() &&
    is_iso_date(node.this_().name()) {
    return @core.exp_cast(node.this_(), DATE)
  }
  if node.kind.is_a(Timestamp) && !node.has("zone") {
    let node = if node.get_type() is None {
      annotate_types(node, dialect~)
    } else {
      node
    }
    return match node.get_type() {
      Some(t) => @core.exp_cast_to(node.this_(), t)
      None => @core.exp_cast(node.this_(), TIMESTAMP)
    }
  }
  node
}

///|
fn coerce_type(
  node : @core.Expr,
  promote_to_inferred_datetime_type : Bool,
) -> @core.Expr {
  if node.kind.is_any(coercible_date_ops) {
    coerce_date_args(
      node.this_(),
      node.expression_(),
      promote_to_inferred_datetime_type,
    )
  } else if node.kind.is_a(Between) {
    coerce_date_args(
      node.this_(),
      node.arg("low").unwrap(),
      promote_to_inferred_datetime_type,
    )
  } else if node.kind.is_a(Extract) &&
    !node.expression_().is_type(@core.dtype_temporal_types) {
    replace_cast(node.expression_(), @core.datatype_of(DATETIME))
  } else if node.kind.is_any([DateAdd, DateSub, DateTrunc]) {
    coerce_timeunit_arg(node.this_(), node.arg("unit")) |> ignore
  } else if node.kind.is_a(DateDiff) {
    for e in [node.this_(), node.expression_()] {
      if !type_in(e, @core.dtype_temporal_types) {
        e.replace(Some(@core.exp_cast(e.copy(), DATETIME))) |> ignore
      }
    }
  }
  node
}

///|
fn remove_redundant_casts(expression : @core.Expr) -> @core.Expr {
  if expression.kind.is_a(Cast) {
    match (expression.this_().get_type(), expression.arg("to")) {
      (Some(t), Some(to)) if to == t => return expression.this_()
      _ => ()
    }
  }
  if expression.kind.is_any([Date, TsOrDsToDate]) {
    match expression.this_().get_type() {
      Some(t) if t.datatype_this() == Some(DATE) && t.expressions().is_empty() =>
        return expression.this_()
      _ => ()
    }
  }
  expression
}

///|
fn canonicalize_ensure_bools(
  expression : @core.Expr,
  replace_func : (@core.Expr) -> Unit,
) -> @core.Expr {
  if expression.kind.is_a(Connector) {
    replace_func(expression.this_())
    replace_func(expression.expression_())
  } else if expression.kind.is_a(Not) {
    replace_func(expression.this_())
  } else if expression.kind.is_a(If) &&
    !(match expression.parent {
      Some(p) => p.kind.is_a(Case) && p.has("this")
      None => false
    }) {
    replace_func(expression.this_())
  } else if expression.kind.is_any([Where, Having]) {
    replace_func(expression.this_())
  }
  expression
}

///|
fn remove_ascending_order(expression : @core.Expr) -> @core.Expr {
  if expression.kind.is_a(Ordered) && expression.get("desc") is Some(Bool(false)) {
    expression.set("desc", @core.null_arg)
  }
  expression
}

///|
fn coerce_date_args(
  a : @core.Expr,
  b : @core.Expr,
  promote_to_inferred_datetime_type : Bool,
) -> Unit {
  for perm in [(a, b), (b, a)] {
    let (a0, b) = perm
    let mut a = a0
    if b.kind.is_a(Interval) {
      a = coerce_timeunit_arg(a, b.arg("unit"))
    }
    let a_type = match a.get_type() {
      Some(t) => t
      None => continue
    }
    let a_this = match a_type.datatype_this() {
      Some(d) if @core.dtype_temporal_types.contains(d) => d
      _ => continue
    }
    if !type_in(b, @core.dtype_text_types) {
      continue
    }
    let target_type = if promote_to_inferred_datetime_type {
      let b_type = if b.is_string() {
        let date_text = b.name()
        if is_iso_date(date_text) {
          @core.DType::DATE
        } else if is_iso_datetime(date_text) {
          DATETIME
        } else {
          a_this
        }
      } else {
        DATETIME
      }
      match default_coerces_to.get(a_this) {
        Some(s) if s.contains(b_type) => @core.datatype_of(b_type)
        _ => a_type
      }
    } else {
      a_type
    }
    if target_type != a_type {
      replace_cast(a, target_type)
    }
    replace_cast(b, target_type)
  }
}

///|
fn coerce_timeunit_arg(arg : @core.Expr, unit : @core.Expr?) -> @core.Expr {
  let t = match arg.get_type() {
    Some(t) => t
    None => return arg
  }
  let this = t.datatype_this()
  match this {
    Some(d) if @core.dtype_text_types.contains(d) => {
      let date_text = arg.name()
      let is_iso_date_ = is_iso_date(date_text)
      if is_iso_date_ && is_date_unit(unit) {
        return arg.replace(Some(@core.exp_cast(arg.copy(), DATE))).unwrap()
      }
      if is_iso_date_ || is_iso_datetime(date_text) {
        return arg.replace(Some(@core.exp_cast(arg.copy(), DATETIME))).unwrap()
      }
    }
    Some(DATE) if !is_date_unit(unit) =>
      return arg.replace(Some(@core.exp_cast(arg.copy(), DATETIME))).unwrap()
    _ => ()
  }
  arg
}

///|
fn replace_cast(node : @core.Expr, to : @core.Expr) -> Unit {
  node.replace(Some(@core.exp_cast_to(node.copy(), to))) |> ignore
}

///|
fn replace_int_predicate(expression : @core.Expr) -> Unit {
  if expression.kind.is_a(Coalesce) {
    for child in expression.iter_expressions() {
      replace_int_predicate(child)
    }
  } else if type_in(expression, @core.dtype_integer_types) {
    expression.replace(Some(@core.exp_neq(expression, @core.literal_int(0))))
    |> ignore
  }
}