///|
fn validate_metric_lengths(
  actual_count : Int,
  predicted_count : Int,
) -> Result[Unit, TreeError] {
  if actual_count != predicted_count {
    return Err(MetricLengthMismatch(actual_count, predicted_count))
  }
  if actual_count == 0 {
    return Err(EmptyDataset)
  }
  Ok(())
}

///|
pub fn classification_accuracy(
  actual : Array[Int],
  predicted : Array[Int],
) -> Result[Double, TreeError] {
  match validate_metric_lengths(actual.length(), predicted.length()) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let mut correct = 0
  for index = 0; index < actual.length(); index = index + 1 {
    if actual[index] == predicted[index] {
      correct = correct + 1
    }
  }
  Ok(correct.to_double() / actual.length().to_double())
}

///|
/// Rows are actual classes and columns are predicted classes.
pub fn confusion_matrix(
  actual : Array[Int],
  predicted : Array[Int],
  class_count : Int,
) -> Result[Array[Array[Int]], TreeError] {
  if class_count <= 0 {
    return Err(InvalidClassCount(class_count))
  }
  match validate_metric_lengths(actual.length(), predicted.length()) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let matrix = Array::makei(class_count, fn(_) { Array::make(class_count, 0) })
  for index = 0; index < actual.length(); index = index + 1 {
    if actual[index] < 0 || actual[index] >= class_count {
      return Err(InvalidClassLabel(index, actual[index], class_count))
    }
    if predicted[index] < 0 || predicted[index] >= class_count {
      return Err(InvalidClassLabel(index, predicted[index], class_count))
    }
    matrix[actual[index]][predicted[index]] = matrix[actual[index]][predicted[index]] +
      1
  }
  Ok(matrix)
}

///|
pub fn mean_squared_error(
  actual : Array[Double],
  predicted : Array[Double],
) -> Result[Double, TreeError] {
  match validate_metric_lengths(actual.length(), predicted.length()) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let mut sum = 0.0
  for index = 0; index < actual.length(); index = index + 1 {
    let difference = actual[index] - predicted[index]
    sum = sum + difference * difference
  }
  Ok(sum / actual.length().to_double())
}

///|
pub fn mean_absolute_error(
  actual : Array[Double],
  predicted : Array[Double],
) -> Result[Double, TreeError] {
  match validate_metric_lengths(actual.length(), predicted.length()) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let mut sum = 0.0
  for index = 0; index < actual.length(); index = index + 1 {
    let difference = actual[index] - predicted[index]
    let absolute = if difference < 0.0 { -difference } else { difference }
    sum = sum + absolute
  }
  Ok(sum / actual.length().to_double())
}

///|
fn square_root(value : Double) -> Double {
  if value <= 0.0 {
    return 0.0
  }
  let mut estimate = if value < 1.0 { 1.0 } else { value }
  let mut iteration = 0
  while iteration < 32 {
    estimate = (estimate + value / estimate) / 2.0
    iteration = iteration + 1
  }
  estimate
}

///|
pub fn root_mean_squared_error(
  actual : Array[Double],
  predicted : Array[Double],
) -> Result[Double, TreeError] {
  match mean_squared_error(actual, predicted) {
    Ok(value) => Ok(square_root(value))
    Err(error) => Err(error)
  }
}

///|
/// Returns 1 for an exact constant-target prediction and 0 otherwise.
pub fn r_squared(
  actual : Array[Double],
  predicted : Array[Double],
) -> Result[Double, TreeError] {
  match validate_metric_lengths(actual.length(), predicted.length()) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let mean = mean_value(actual)
  let mut residual = 0.0
  let mut total = 0.0
  for index = 0; index < actual.length(); index = index + 1 {
    let prediction_difference = actual[index] - predicted[index]
    residual = residual + prediction_difference * prediction_difference
    let mean_difference = actual[index] - mean
    total = total + mean_difference * mean_difference
  }
  if total <= 0.000000000001 {
    return Ok(if residual <= 0.000000000001 { 1.0 } else { 0.0 })
  }
  Ok(1.0 - residual / total)
}