///|
fn log2_positive(value : Double) -> Double {
  if value <= 0.0 {
    return 0.0
  }
  let mut scaled = value
  let mut exponent = 0
  while scaled < 0.5 {
    scaled = scaled * 2.0
    exponent = exponent - 1
  }
  while scaled >= 1.0 {
    scaled = scaled / 2.0
    exponent = exponent + 1
  }
  let y = (scaled - 1.0) / (scaled + 1.0)
  let y_squared = y * y
  let mut power = y
  let mut denominator = 1
  let mut series = 0.0
  let mut term_index = 0
  while term_index < 24 {
    series = series + power / denominator.to_double()
    power = power * y_squared
    denominator = denominator + 2
    term_index = term_index + 1
  }
  let natural_log = 2.0 * series
  exponent.to_double() + natural_log / 0.6931471805599453
}

///|
fn impurity_from_counts(
  counts : Array[Int],
  sample_count : Int,
  criterion : ClassificationCriterion,
) -> Double {
  if sample_count == 0 {
    return 0.0
  }
  match criterion {
    Gini => {
      let mut result = 1.0
      for count in counts {
        if count > 0 {
          let probability = count.to_double() / sample_count.to_double()
          result = result - probability * probability
        }
      }
      result
    }
    Entropy => {
      let mut result = 0.0
      for count in counts {
        if count > 0 {
          let probability = count.to_double() / sample_count.to_double()
          result = result - probability * log2_positive(probability)
        }
      }
      result
    }
  }
}

///|
pub fn gini_impurity(labels : Array[Int], class_count : Int) -> Double {
  if labels.is_empty() || class_count <= 0 {
    return 0.0
  }
  let counts = Array::make(class_count, 0)
  for label in labels {
    if label >= 0 && label < class_count {
      counts[label] = counts[label] + 1
    }
  }
  impurity_from_counts(counts, labels.length(), Gini)
}

///|
pub fn entropy_impurity(labels : Array[Int], class_count : Int) -> Double {
  if labels.is_empty() || class_count <= 0 {
    return 0.0
  }
  let counts = Array::make(class_count, 0)
  for label in labels {
    if label >= 0 && label < class_count {
      counts[label] = counts[label] + 1
    }
  }
  impurity_from_counts(counts, labels.length(), Entropy)
}

///|
fn validate_impurity_labels(
  labels : Array[Int],
  class_count : Int,
) -> Result[Unit, TreeError] {
  if class_count <= 0 {
    return Err(InvalidClassCount(class_count))
  }
  for index, label in labels {
    if label < 0 || label >= class_count {
      return Err(InvalidClassLabel(index, label, class_count))
    }
  }
  Ok(())
}

///|
/// Validates labels before calculating Gini impurity.
pub fn checked_gini_impurity(
  labels : Array[Int],
  class_count : Int,
) -> Result[Double, TreeError] {
  match validate_impurity_labels(labels, class_count) {
    Err(error) => Err(error)
    Ok(_) => Ok(gini_impurity(labels, class_count))
  }
}

///|
/// Validates labels before calculating entropy impurity.
pub fn checked_entropy_impurity(
  labels : Array[Int],
  class_count : Int,
) -> Result[Double, TreeError] {
  match validate_impurity_labels(labels, class_count) {
    Err(error) => Err(error)
    Ok(_) => Ok(entropy_impurity(labels, class_count))
  }
}

///|
fn compare_feature_indices(
  dataset : ClassificationDataset,
  feature_index : Int,
  left : Int,
  right : Int,
) -> Int {
  let left_value = dataset.rows[left][feature_index]
  let right_value = dataset.rows[right][feature_index]
  if left_value < right_value {
    -1
  } else if left_value > right_value {
    1
  } else if left < right {
    -1
  } else if left > right {
    1
  } else {
    0
  }
}