// Ports of the sqlglot/transforms.py (and related) functions used by the base Generator.

///|
/// Some dialects only allow CTEs to be defined at the top-level. This moves all
/// nested CTEs to the top level (Python `transforms.move_ctes_to_top_level`).
pub fn move_ctes_to_top_level(expression : Expr) -> Expr {
  let mut top_level_with = expression.arg("with_")
  let it = expression.find_all([With])
  while it.next() is Some(inner_with) {
    match inner_with.parent {
      Some(p) if physical_equal(p, expression) => continue
      _ => ()
    }
    match top_level_with {
      None => {
        let w = inner_with.pop()
        top_level_with = Some(w)
        expression.set("with_", w)
      }
      Some(tlw) => {
        if inner_with.has("recursive") {
          tlw.set("recursive", true)
        }
        let parent_cte = inner_with.find_ancestor([CTE])
        inner_with.pop() |> ignore
        let existing = tlw.expressions()
        match parent_cte {
          Some(pc) => {
            let mut i = 0
            while i < existing.length() && !(existing[i] == pc) {
              i += 1
            }
            if i >= existing.length() {
              abort("ValueError: CTE is not in list")
            }
            let new_exprs = existing[:i].to_owned() +
              inner_with.expressions() +
              existing[i:].to_owned()
            tlw.set("expressions", new_exprs)
          }
          None => tlw.set("expressions", existing + inner_with.expressions())
        }
      }
    }
  }
  expression
}

///|
/// Python `optimizer.canonicalize.ensure_bools`.
fn canonicalize_ensure_bools(
  expression : Expr,
  replace_func : (Expr) -> Unit,
) -> Expr {
  if expression.kind.is_a(Connector) {
    match expression.arg("this") {
      Some(l) => replace_func(l)
      None => ()
    }
    match expression.arg("expression") {
      Some(r) => replace_func(r)
      None => ()
    }
  } else if expression.kind.is_a(Not) {
    match expression.this() {
      Some(t) => replace_func(t)
      None => ()
    }
  } else if expression.kind.is_a(If) &&
    !(match expression.parent {
      Some(p) => p.kind.is_a(Case) && p.has("this")
      None => false
    }) {
    match expression.this() {
      Some(t) => replace_func(t)
      None => ()
    }
  } else if expression.kind.is_any([Where, Having]) {
    match expression.this() {
      Some(t) => replace_func(t)
      None => ()
    }
  }
  expression
}

///|
/// Converts numeric values used in conditions into explicit boolean expressions
/// (Python `transforms.ensure_bools`).
pub fn ensure_bools(expression : Expr) -> Expr {
  let numeric : Array[DType] = [DType::UNKNOWN]
  for d in dtype_numeric_types {
    numeric.push(d)
  }
  fn ensure_bool(node : Expr) -> Unit {
    if node.is_number() ||
      (!node.kind.is_a(SubqueryPredicate) && node.is_type(numeric)) ||
      (node.kind.is_a(Column) && node.type_ is None) {
      node.replace(Some(exp_neq(node, literal_int(0)))) |> ignore
    }
  }

  let it = expression.walk()
  while it.next() is Some(node) {
    canonicalize_ensure_bools(node, ensure_bool) |> ignore
  }
  expression
}

///|
/// Python `optimizer.scope.walk_in_scope`: visits all nodes in the syntax tree,
/// stopping at nodes that start child scopes.
pub fn gen_walk_in_scope(expression : Expr, out : Array[Expr]) -> Unit {
  let stack : Array[Expr] = [expression]
  while stack.pop() is Some(node) {
    out.push(node)
    if !physical_equal(node, expression) && node.kind.is_any([CTE, Query]) {
      let parent_kind_is = fn(kinds : Array[Kind]) {
        match node.parent {
          Some(p) => p.kind.is_any(kinds)
          None => false
        }
      }
      let is_derived_table = node.kind.is_a(Subquery) &&
        (
          node.alias() != "" ||
          (match node.this() {
            Some(t) => t.kind.is_any([Select, SetOperation])
            None => false
          })
        )
      if node.kind.is_a(CTE) ||
        (parent_kind_is([From, Join]) && is_derived_table) ||
        parent_kind_is([UDTF]) ||
        node.kind.is_any([Select, SetOperation]) {
        if node.kind.is_any([Subquery, UDTF]) {
          for key in ["joins", "laterals", "pivots"] {
            for arg in node.list(key) {
              gen_walk_in_scope(arg, out)
            }
          }
        }
        continue
      }
    }
    let values = []
    for _, v in node.args {
      values.push(v)
    }
    for i = values.length() - 1; i >= 0; i = i - 1 {
      match values[i] {
        List(l) =>
          for j = l.length() - 1; j >= 0; j = j - 1 {
            match l[j] {
              Node(e) => stack.push(e)
              _ => ()
            }
          }
        Node(e) => stack.push(e)
        _ => ()
      }
    }
  }
}

///|
/// Python `optimizer.scope.find_all_in_scope`.
pub fn gen_find_all_in_scope(
  expression : Expr,
  kinds : ArrayView[Kind],
) -> Array[Expr] {
  let all = []
  gen_walk_in_scope(expression, all)
  all.filter(n => n.kind.is_any(kinds))
}