///|
/// - Does: Runs the structural pre-pass used before common-subexpression elimination.
/// - Input: Any `Expr`.
/// - Returns: One rewritten `Expr`.
/// - Limits: Only the implemented subtraction-shape normalizations are applied.
pub fn sub_pre(expr : Expr) -> Expr {
  rewrite_bottom_up_cseopts(expr, rewrite_sub_pre)
}

///|
/// - Does: Runs the cleanup pass used after common-subexpression elimination.
/// - Input: Any `Expr`.
/// - Returns: One rewritten `Expr`.
/// - Limits: Only the implemented sign-normalization cleanup is applied.
pub fn sub_post(expr : Expr) -> Expr {
  rewrite_bottom_up_cseopts(expr, rewrite_sub_post)
}

///|
fn rewrite_bottom_up_cseopts(expr : Expr, rule : (Expr) -> Expr) -> Expr {
  let rewritten = @symcore.map_children(expr, child => {
    rewrite_bottom_up_cseopts(child, rule)
  })
  rule(rewritten)
}

///|
fn rewrite_sub_pre(expr : Expr) -> Expr {
  match expr {
    Expr::Add([a, b]) => {
      let (sa, pa) = split_sign(a)
      let (sb, pb) = split_sign(b)
      if sa < 0 && sb > 0 {
        @symcore.mul([int(-1), @symcore.add([pa, @symcore.mul([int(-1), pb])])])
      } else if sb < 0 && sa > 0 {
        @symcore.mul([int(-1), @symcore.add([pb, @symcore.mul([int(-1), pa])])])
      } else {
        expr
      }
    }
    _ => expr
  }
}

///|
fn rewrite_sub_post(expr : Expr) -> Expr {
  match expr {
    Expr::Mul(args) => {
      if args.length() < 2 {
        return expr
      }
      let filtered : Array[Expr] = Array::new()
      let mut neg_count = 0
      for arg in args {
        match arg {
          Expr::Number(n) if n.is_one() => ()
          Expr::Number(n) if n.compare(@symnum.BigRational::from_int(-1)) == 0 =>
            neg_count = neg_count + 1
          _ => filtered.push(arg)
        }
      }
      let out = @symcore.mul(filtered)
      if neg_count % 2 == 1 {
        @symcore.mul([int(-1), out])
      } else {
        out
      }
    }
    _ => expr
  }
}

///|
fn split_sign(expr : Expr) -> (Int, Expr) {
  match expr {
    Expr::Number(n) =>
      if n.compare(@symnum.BigRational::zero()) < 0 {
        (-1, @symcore.Expr::Number(n.neg_r()))
      } else {
        (1, expr)
      }
    Expr::Mul(args) => {
      let mut sign = 1
      let factors : Array[Expr] = Array::new()
      for arg in args {
        match arg {
          Expr::Number(n) if n.compare(@symnum.BigRational::zero()) < 0 => {
            sign = -sign
            factors.push(@symcore.Expr::Number(n.neg_r()))
          }
          _ => factors.push(arg)
        }
      }
      (sign, @symcore.mul(factors))
    }
    _ => (1, expr)
  }
}