///|
fn classification_leaf_from_subtree(
  node : ClassificationNode,
  class_count : Int,
  impurity : Double,
) -> ClassificationNode {
  let weighted = Array::make(class_count, 0.0)
  let sample_count = collect_classification_probabilities(node, weighted)
  let probabilities = weighted.map(fn(value) {
    value / sample_count.to_double()
  })
  let mut prediction = 0
  for class_index = 1; class_index < class_count; class_index = class_index + 1 {
    if probabilities[class_index] > probabilities[prediction] {
      prediction = class_index
    }
  }
  ClassificationLeaf(prediction, probabilities, sample_count, impurity)
}

///|
fn collect_classification_probabilities(
  node : ClassificationNode,
  weighted : Array[Double],
) -> Int {
  match node {
    ClassificationLeaf(_, probabilities, samples, _) => {
      for class_index = 0
          class_index < probabilities.length()
          class_index = class_index + 1 {
        weighted[class_index] = weighted[class_index] +
          probabilities[class_index] * samples.to_double()
      }
      samples
    }
    ClassificationBranch(_, _, left, right, _, _, _) =>
      collect_classification_probabilities(left, weighted) +
      collect_classification_probabilities(right, weighted)
  }
}

///|
fn prune_classification_node(
  node : ClassificationNode,
  class_count : Int,
  alpha : Double,
) -> (ClassificationNode, Double, Int) {
  match node {
    ClassificationLeaf(_, _, samples, impurity) =>
      (node, samples.to_double() * impurity, 1)
    ClassificationBranch(
      feature,
      threshold,
      left,
      right,
      samples,
      impurity,
      gain
    ) => {
      let (pruned_left, left_risk, left_leaves) = prune_classification_node(
        left, class_count, alpha,
      )
      let (pruned_right, right_risk, right_leaves) = prune_classification_node(
        right, class_count, alpha,
      )
      let subtree_risk = left_risk + right_risk
      let leaf_total = left_leaves + right_leaves
      let collapsed_risk = samples.to_double() * impurity
      let branch = ClassificationBranch(
        feature, threshold, pruned_left, pruned_right, samples, impurity, gain,
      )
      if collapsed_risk + alpha <= subtree_risk + alpha * leaf_total.to_double() {
        (
          classification_leaf_from_subtree(branch, class_count, impurity),
          collapsed_risk,
          1,
        )
      } else {
        (branch, subtree_risk, leaf_total)
      }
    }
  }
}

///|
fn regression_leaf_from_subtree(
  node : RegressionNode,
  variance : Double,
) -> RegressionNode {
  let (weighted_sum, samples) = collect_regression_predictions(node)
  RegressionLeaf(weighted_sum / samples.to_double(), samples, variance)
}

///|
fn collect_regression_predictions(node : RegressionNode) -> (Double, Int) {
  match node {
    RegressionLeaf(prediction, samples, _) =>
      (prediction * samples.to_double(), samples)
    RegressionBranch(_, _, left, right, _, _, _) => {
      let (left_sum, left_samples) = collect_regression_predictions(left)
      let (right_sum, right_samples) = collect_regression_predictions(right)
      (left_sum + right_sum, left_samples + right_samples)
    }
  }
}

///|
fn prune_regression_node(
  node : RegressionNode,
  alpha : Double,
) -> (RegressionNode, Double, Int) {
  match node {
    RegressionLeaf(_, samples, variance) =>
      (node, samples.to_double() * variance, 1)
    RegressionBranch(feature, threshold, left, right, samples, variance, gain) => {
      let (pruned_left, left_risk, left_leaves) = prune_regression_node(
        left, alpha,
      )
      let (pruned_right, right_risk, right_leaves) = prune_regression_node(
        right, alpha,
      )
      let subtree_risk = left_risk + right_risk
      let leaf_total = left_leaves + right_leaves
      let collapsed_risk = samples.to_double() * variance
      let branch = RegressionBranch(
        feature, threshold, pruned_left, pruned_right, samples, variance, gain,
      )
      if collapsed_risk + alpha <= subtree_risk + alpha * leaf_total.to_double() {
        (regression_leaf_from_subtree(branch, variance), collapsed_risk, 1)
      } else {
        (branch, subtree_risk, leaf_total)
      }
    }
  }
}

///|
/// Returns a pruned copy using leaf-risk plus alpha times leaf count.
pub fn ClassificationTree::prune(
  self : ClassificationTree,
  alpha : Double,
) -> Result[ClassificationTree, TreeError] {
  if !finite_number(alpha) || alpha < 0.0 {
    return Err(InvalidPruningAlpha(alpha))
  }
  if alpha == 0.0 {
    return Ok(self)
  }
  let (root, _, _) = prune_classification_node(
    self.root,
    self.class_total,
    alpha,
  )
  Ok({ ..self, root, })
}

///|
/// Returns a pruned copy using leaf-risk plus alpha times leaf count.
pub fn RegressionTree::prune(
  self : RegressionTree,
  alpha : Double,
) -> Result[RegressionTree, TreeError] {
  if !finite_number(alpha) || alpha < 0.0 {
    return Err(InvalidPruningAlpha(alpha))
  }
  if alpha == 0.0 {
    return Ok(self)
  }
  let (root, _, _) = prune_regression_node(self.root, alpha)
  Ok({ ..self, root, })
}