///|
fn class_position(classes : Array[Int], label : Int) -> Int {
  for index, value in classes {
    if value == label {
      return index
    }
  }
  -1
}

///|
fn safe_rate(numerator : Int, denominator : Int) -> Double {
  if denominator == 0 {
    0.0
  } else {
    numerator.to_double() / denominator.to_double()
  }
}

///|
fn harmonic_f1(precision : Double, recall : Double) -> Double {
  if precision + recall == 0.0 {
    0.0
  } else {
    2.0 * precision * recall / (precision + recall)
  }
}

///|
/// Computes a confusion matrix and common single-label classification metrics.
pub fn classification_metrics(
  actual : Array[Int],
  predicted : Array[Int],
) -> Result[ClassificationMetrics, SvmError] {
  if actual.is_empty() {
    return Err(EmptyMetricInput)
  }
  if actual.length() != predicted.length() {
    return Err(MetricLengthMismatch(actual.length(), predicted.length()))
  }
  let combined = actual.copy()
  for label in predicted {
    combined.push(label)
  }
  let classes = sorted_distinct_labels(combined)
  let size = classes.length()
  let matrix : Array[Array[Int]] = Array::makei(size, fn(_) {
    Array::make(size, 0)
  })
  for index = 0; index < actual.length(); index = index + 1 {
    let row = class_position(classes, actual[index])
    let column = class_position(classes, predicted[index])
    if row < 0 {
      return Err(UnknownClassLabel(actual[index]))
    }
    if column < 0 {
      return Err(UnknownClassLabel(predicted[index]))
    }
    matrix[row][column] = matrix[row][column] + 1
  }
  let per_class : Array[ClassMetric] = []
  let mut correct = 0
  let mut macro_precision = 0.0
  let mut macro_recall = 0.0
  let mut macro_f1 = 0.0
  let mut weighted_precision = 0.0
  let mut weighted_recall = 0.0
  let mut weighted_f1 = 0.0
  for class_index = 0; class_index < size; class_index = class_index + 1 {
    let true_positive = matrix[class_index][class_index]
    correct = correct + true_positive
    let mut support = 0
    let mut predicted_count = 0
    for other = 0; other < size; other = other + 1 {
      support = support + matrix[class_index][other]
      predicted_count = predicted_count + matrix[other][class_index]
    }
    let false_positive = predicted_count - true_positive
    let false_negative = support - true_positive
    let precision = safe_rate(true_positive, predicted_count)
    let recall = safe_rate(true_positive, support)
    let f1 = harmonic_f1(precision, recall)
    per_class.push({
      class_label: classes[class_index],
      true_positive_count: true_positive,
      false_positive_count: false_positive,
      false_negative_count: false_negative,
      class_support: support,
      precision_value: precision,
      recall_value: recall,
      f1_value: f1,
    })
    macro_precision = macro_precision + precision
    macro_recall = macro_recall + recall
    macro_f1 = macro_f1 + f1
    let fraction = support.to_double() / actual.length().to_double()
    weighted_precision = weighted_precision + fraction * precision
    weighted_recall = weighted_recall + fraction * recall
    weighted_f1 = weighted_f1 + fraction * f1
  }
  macro_precision = macro_precision / size.to_double()
  macro_recall = macro_recall / size.to_double()
  macro_f1 = macro_f1 / size.to_double()
  let accuracy = correct.to_double() / actual.length().to_double()
  Ok({
    ordered_classes: classes,
    confusion_counts: matrix,
    class_metrics: per_class,
    total_observations: actual.length(),
    accuracy_value: accuracy,
    error_rate_value: 1.0 - accuracy,
    balanced_accuracy_value: macro_recall,
    macro_precision_value: macro_precision,
    macro_recall_value: macro_recall,
    macro_f1_value: macro_f1,
    micro_precision_value: accuracy,
    micro_recall_value: accuracy,
    micro_f1_value: accuracy,
    weighted_precision_value: weighted_precision,
    weighted_recall_value: weighted_recall,
    weighted_f1_value: weighted_f1,
  })
}

///|
pub fn ClassMetric::label(self : ClassMetric) -> Int {
  self.class_label
}

///|
pub fn ClassMetric::true_positives(self : ClassMetric) -> Int {
  self.true_positive_count
}

///|
pub fn ClassMetric::false_positives(self : ClassMetric) -> Int {
  self.false_positive_count
}

///|
pub fn ClassMetric::false_negatives(self : ClassMetric) -> Int {
  self.false_negative_count
}

///|
pub fn ClassMetric::support(self : ClassMetric) -> Int {
  self.class_support
}

///|
pub fn ClassMetric::precision(self : ClassMetric) -> Double {
  self.precision_value
}

///|
pub fn ClassMetric::recall(self : ClassMetric) -> Double {
  self.recall_value
}

///|
pub fn ClassMetric::f1(self : ClassMetric) -> Double {
  self.f1_value
}

///|
pub fn ClassificationMetrics::classes(
  self : ClassificationMetrics,
) -> Array[Int] {
  self.ordered_classes.copy()
}

///|
pub fn ClassificationMetrics::confusion_matrix(
  self : ClassificationMetrics,
) -> Array[Array[Int]] {
  let copied : Array[Array[Int]] = []
  for row in self.confusion_counts {
    copied.push(row.copy())
  }
  copied
}

///|
pub fn ClassificationMetrics::per_class(
  self : ClassificationMetrics,
) -> Array[ClassMetric] {
  self.class_metrics.copy()
}

///|
pub fn ClassificationMetrics::observation_count(
  self : ClassificationMetrics,
) -> Int {
  self.total_observations
}

///|
pub fn ClassificationMetrics::accuracy(self : ClassificationMetrics) -> Double {
  self.accuracy_value
}

///|
pub fn ClassificationMetrics::error_rate(
  self : ClassificationMetrics,
) -> Double {
  self.error_rate_value
}

///|
pub fn ClassificationMetrics::balanced_accuracy(
  self : ClassificationMetrics,
) -> Double {
  self.balanced_accuracy_value
}

///|
pub fn ClassificationMetrics::macro_precision(
  self : ClassificationMetrics,
) -> Double {
  self.macro_precision_value
}

///|
pub fn ClassificationMetrics::macro_recall(
  self : ClassificationMetrics,
) -> Double {
  self.macro_recall_value
}

///|
pub fn ClassificationMetrics::macro_f1(self : ClassificationMetrics) -> Double {
  self.macro_f1_value
}

///|
pub fn ClassificationMetrics::micro_precision(
  self : ClassificationMetrics,
) -> Double {
  self.micro_precision_value
}

///|
pub fn ClassificationMetrics::micro_recall(
  self : ClassificationMetrics,
) -> Double {
  self.micro_recall_value
}

///|
pub fn ClassificationMetrics::micro_f1(self : ClassificationMetrics) -> Double {
  self.micro_f1_value
}

///|
pub fn ClassificationMetrics::weighted_precision(
  self : ClassificationMetrics,
) -> Double {
  self.weighted_precision_value
}

///|
pub fn ClassificationMetrics::weighted_recall(
  self : ClassificationMetrics,
) -> Double {
  self.weighted_recall_value
}

///|
pub fn ClassificationMetrics::weighted_f1(
  self : ClassificationMetrics,
) -> Double {
  self.weighted_f1_value
}