///|
fn classification_alpha_candidates(
  node : ClassificationNode,
  values : Array[Double],
) -> (Double, Int) {
  match node {
    ClassificationLeaf(_, _, samples, impurity) =>
      (samples.to_double() * impurity, 1)
    ClassificationBranch(_, _, left, right, samples, impurity, _) => {
      let (left_risk, left_leaves) = classification_alpha_candidates(
        left, values,
      )
      let (right_risk, right_leaves) = classification_alpha_candidates(
        right, values,
      )
      let subtree_risk = left_risk + right_risk
      let leaves = left_leaves + right_leaves
      let collapsed_risk = samples.to_double() * impurity
      let raw = (collapsed_risk - subtree_risk) / (leaves - 1).to_double()
      values.push(if raw > 0.0 { raw } else { 0.0 })
      (subtree_risk, leaves)
    }
  }
}

///|
fn regression_alpha_candidates(
  node : RegressionNode,
  values : Array[Double],
) -> (Double, Int) {
  match node {
    RegressionLeaf(_, samples, variance) => (samples.to_double() * variance, 1)
    RegressionBranch(_, _, left, right, samples, variance, _) => {
      let (left_risk, left_leaves) = regression_alpha_candidates(left, values)
      let (right_risk, right_leaves) = regression_alpha_candidates(
        right, values,
      )
      let subtree_risk = left_risk + right_risk
      let leaves = left_leaves + right_leaves
      let collapsed_risk = samples.to_double() * variance
      let raw = (collapsed_risk - subtree_risk) / (leaves - 1).to_double()
      values.push(if raw > 0.0 { raw } else { 0.0 })
      (subtree_risk, leaves)
    }
  }
}

///|
fn sorted_unique_alphas(values : Array[Double]) -> Array[Double] {
  values.sort_by(fn(left, right) {
    if left < right {
      -1
    } else if left > right {
      1
    } else {
      0
    }
  })
  let unique : Array[Double] = []
  for value in values {
    if value > 0.0 &&
      (
        unique.is_empty() ||
        value - unique[unique.length() - 1] > 0.000000000001
      ) {
      unique.push(value)
    }
  }
  unique
}

///|
fn classification_root_risk(node : ClassificationNode) -> Double {
  match node {
    ClassificationLeaf(_, _, samples, impurity) =>
      samples.to_double() * impurity
    ClassificationBranch(_, _, _, _, samples, impurity, _) =>
      samples.to_double() * impurity
  }
}

///|
fn regression_root_risk(node : RegressionNode) -> Double {
  match node {
    RegressionLeaf(_, samples, variance) => samples.to_double() * variance
    RegressionBranch(_, _, _, _, samples, variance, _) =>
      samples.to_double() * variance
  }
}

///|
/// Generates distinct models at increasing cost-complexity strengths.
pub fn ClassificationTree::pruning_path(
  self : ClassificationTree,
) -> Array[ClassificationPruningStep] {
  let steps : Array[ClassificationPruningStep] = [
    {
      alpha: 0.0,
      model: self,
      node_count: self.node_count(),
      leaf_count: self.leaf_count(),
    },
  ]
  if self.node_count() == 1 {
    return steps
  }
  let candidates : Array[Double] = []
  ignore(classification_alpha_candidates(self.root, candidates))
  let alphas = sorted_unique_alphas(candidates)
  for alpha in alphas {
    let model = self.prune(alpha).unwrap()
    let node_count = model.node_count()
    if node_count < steps[steps.length() - 1].node_count {
      steps.push({ alpha, model, node_count, leaf_count: model.leaf_count() })
    }
  }
  if steps[steps.length() - 1].node_count > 1 {
    let previous_alpha = steps[steps.length() - 1].alpha
    let alpha = previous_alpha + classification_root_risk(self.root) + 1.0
    let model = self.prune(alpha).unwrap()
    steps.push({ alpha, model, node_count: 1, leaf_count: 1 })
  }
  steps
}

///|
/// Generates distinct models at increasing cost-complexity strengths.
pub fn RegressionTree::pruning_path(
  self : RegressionTree,
) -> Array[RegressionPruningStep] {
  let steps : Array[RegressionPruningStep] = [
    {
      alpha: 0.0,
      model: self,
      node_count: self.node_count(),
      leaf_count: self.leaf_count(),
    },
  ]
  if self.node_count() == 1 {
    return steps
  }
  let candidates : Array[Double] = []
  ignore(regression_alpha_candidates(self.root, candidates))
  let alphas = sorted_unique_alphas(candidates)
  for alpha in alphas {
    let model = self.prune(alpha).unwrap()
    let node_count = model.node_count()
    if node_count < steps[steps.length() - 1].node_count {
      steps.push({ alpha, model, node_count, leaf_count: model.leaf_count() })
    }
  }
  if steps[steps.length() - 1].node_count > 1 {
    let previous_alpha = steps[steps.length() - 1].alpha
    let alpha = previous_alpha + regression_root_risk(self.root) + 1.0
    let model = self.prune(alpha).unwrap()
    steps.push({ alpha, model, node_count: 1, leaf_count: 1 })
  }
  steps
}