// Port of sqlglot/optimizer/normalize.py.

///|
/// Rewrite the AST into conjunctive normal form (or disjunctive normal form if `dnf`).
pub fn normalize(
  expression : @core.Expr,
  dnf? : Bool = false,
  max_distance? : Int = 128,
) -> @core.Expr raise @core.SqlglotError {
  let simplifier = Simplifier::new(annotate_new_expressions=false)
  let mut expression = expression
  let nodes = expression.walk(prune=e => e.kind.is_a(Connector)).collect()
  for node in nodes {
    if !node.kind.is_a(Connector) {
      continue
    }
    if normalized(node, dnf~) {
      continue
    }
    let root = physical_equal(node, expression)
    let original = node.copy()
    node.transform(
      e => Some(simplifier.rewrite_between(e) catch { _ => e }),
      copy=false,
    )
    |> ignore
    let distance = normalization_distance(node, dnf~, max=max_distance)
    if distance > max_distance {
      return expression
    }
    let result = while_changing(node, e => distributive_law(
      e, dnf, max_distance, simplifier,
    )) catch {
      @core.OptimizeError(_) => {
        node.replace(Some(original)) |> ignore
        if root {
          return original
        }
        return expression
      }
      e => raise e
    }
    let node = node.replace(Some(result)).unwrap()
    if root {
      expression = node
    }
  }
  expression
}

///|
/// Applies `func` until the expression's hash stops changing.
fn while_changing(
  expression : @core.Expr,
  func : (@core.Expr) -> @core.Expr raise @core.SqlglotError,
) -> @core.Expr raise @core.SqlglotError {
  let mut expression = expression
  for ;; {
    let start_hash = expression.hash()
    expression = func(expression)
    if expression.hash() == start_hash {
      break
    }
  }
  expression
}

///|
/// Checks whether a given expression is in a normal form of interest.
pub fn normalized(expression : @core.Expr, dnf? : Bool = false) -> Bool {
  let (ancestor, root) : (@core.Kind, @core.Kind) = if dnf {
    (And, Or)
  } else {
    (Or, And)
  }
  // (one memoized upward search for all the connectors of a long chain)
  let ancestors = AncestorCache::new([ancestor])
  !find_all_in_scope(expression, [root]).any(connector => {
    ancestors.find(connector) is Some(_)
  })
}

///|
/// The difference in the number of predicates between a given expression and its
/// normalized form.
pub fn normalization_distance(
  expression : @core.Expr,
  dnf? : Bool = false,
  max? : Int = 2147483647,
) -> Int {
  let total = Ref(-(expression.find_all([Connector]).count() + 1))
  predicate_lengths(expression, dnf, max, 0, length => {
    total.val += length
    total.val <= max
  })
  |> ignore
  total.val
}

///|
/// Emits the predicate lengths when expanded to normalized form; `emit` returns false
/// to stop early. Returns false if stopped.
fn predicate_lengths(
  expression : @core.Expr,
  dnf : Bool,
  max : Int,
  depth : Int,
  emit : (Int) -> Bool,
) -> Bool {
  if depth > max {
    return emit(depth)
  }
  let expression = expression.unnest()
  if !expression.kind.is_a(Connector) {
    return emit(1)
  }
  let depth = depth + 1
  let left = expression.this_()
  let right = expression.expression_()
  let distributes = if dnf {
    expression.kind.is_a(And)
  } else {
    expression.kind.is_a(Or)
  }
  if distributes {
    predicate_lengths(left, dnf, max, depth, a => predicate_lengths(
      right,
      dnf,
      max,
      depth,
      b => emit(a + b),
    ))
  } else {
    predicate_lengths(left, dnf, max, depth, emit) &&
    predicate_lengths(right, dnf, max, depth, emit)
  }
}

///|
/// Python `exp.replace_children(expression, fun)`.
fn replace_children(
  expression : @core.Expr,
  fun : (@core.Expr) -> @core.Expr raise @core.SqlglotError,
) -> Unit raise @core.SqlglotError {
  for k, v in expression.args.copy() {
    match v {
      List(l) => {
        let new_nodes : Array[@core.Value] = []
        for cn in l {
          match cn {
            Node(e) => new_nodes.push(Node(fun(e)))
            other => new_nodes.push(other)
          }
        }
        expression.set(k, @core.Value::List(new_nodes))
      }
      Node(e) => expression.set(k, fun(e))
      _ => ()
    }
  }
}

///|
/// x OR (y AND z) -> (x OR y) AND (x OR z)
fn distributive_law(
  expression : @core.Expr,
  dnf : Bool,
  max_distance : Int,
  simplifier : Simplifier,
) -> @core.Expr raise @core.SqlglotError {
  if normalized(expression, dnf~) {
    return expression
  }
  let distance = normalization_distance(expression, dnf~, max=max_distance)
  if distance > max_distance {
    raise @core.OptimizeError(
      "Normalization distance \{distance} exceeds max \{max_distance}",
    )
  }
  replace_children(expression, e => distributive_law(
    e, dnf, max_distance, simplifier,
  ))
  let (to_exp, from_exp) : (@core.Kind, @core.Kind) = if dnf {
    (Or, And)
  } else {
    (And, Or)
  }
  if expression.kind.is_a(from_exp) {
    let operands = expression.unnest_operands()
    let a = operands[0]
    let b = operands[1]
    if a.kind.is_a(to_exp) && b.kind.is_a(to_exp) {
      if a.find_all([Connector]).count() > b.find_all([Connector]).count() {
        return distribute(a, b, from_exp, to_exp, simplifier)
      }
      return distribute(b, a, from_exp, to_exp, simplifier)
    }
    if a.kind.is_a(to_exp) {
      return distribute(b, a, from_exp, to_exp, simplifier)
    }
    if b.kind.is_a(to_exp) {
      return distribute(a, b, from_exp, to_exp, simplifier)
    }
  }
  expression
}

///|
fn combine2(kind : @core.Kind, a : @core.Expr, b : @core.Expr, copy : Bool) -> @core.Expr {
  let a = if copy { a.copy() } else { a }
  let b = if copy { b.copy() } else { b }
  @core.combine_conditions([a, b], kind, copy=false)
}

///|
fn distribute(
  a : @core.Expr,
  b : @core.Expr,
  from_kind : @core.Kind,
  to_kind : @core.Kind,
  simplifier : Simplifier,
) -> @core.Expr raise @core.SqlglotError {
  if a.kind.is_a(Connector) && a.kind.is_a(b.kind) {
    replace_children(a, c => combine2(
      to_kind,
      simplifier.uniq_sort(
        flatten_connector(combine2(from_kind, c, b.this_(), true)),
        true,
      ),
      simplifier.uniq_sort(
        flatten_connector(combine2(from_kind, c, b.expression_(), true)),
        true,
      ),
      false,
    ))
    return a
  }
  combine2(
    to_kind,
    simplifier.uniq_sort(
      flatten_connector(combine2(from_kind, a, b.this_(), true)),
      true,
    ),
    simplifier.uniq_sort(
      flatten_connector(combine2(from_kind, a, b.expression_(), true)),
      true,
    ),
    false,
  )
}