///|
/// Numerically stable probability primitives for decoding algorithms.
pub enum ProbabilityError {
  EmptyLogits
  InvalidTemperature
  InvalidProbability
  InvalidSample
} derive(Eq, Debug)

///|
/// Also rejects infinities: infinity minus itself is NaN.
fn finite_number(value : Double) -> Bool {
  value - value == 0.0
}

///|
fn valid_uniform(value : Double) -> Bool {
  finite_number(value) && value >= 0.0 && value < 1.0
}

///|
pub fn log_sum_exp(values : Array[Double]) -> Result[Double, ProbabilityError] {
  if values.length() == 0 {
    return Err(EmptyLogits)
  }
  let mut maximum = values[0]
  for value in values {
    if !finite_number(value) {
      return Err(InvalidProbability)
    }
    if value > maximum {
      maximum = value
    }
  }
  let mut total = 0.0
  for value in values {
    total = total + @math.exp(value - maximum)
  }
  Ok(maximum + @math.ln(total))
}

///|
pub fn softmax(
  logits : Array[Double],
  temperature? : Double = 1.0,
) -> Result[Array[Double], ProbabilityError] {
  if logits.length() == 0 {
    return Err(EmptyLogits)
  }
  if !finite_number(temperature) || temperature <= 0.0 {
    return Err(InvalidTemperature)
  }
  let scaled : Array[Double] = []
  for value in logits {
    if !finite_number(value) || !finite_number(value / temperature) {
      return Err(InvalidProbability)
    }
    scaled.push(value / temperature)
  }
  let mut maximum = scaled[0]
  for value in scaled {
    maximum = maximum.max(value)
  }
  let output : Array[Double] = []
  let mut total = 0.0
  for value in scaled {
    let weight = @math.exp(value - maximum)
    output.push(weight)
    total = total + weight
  }
  for i in 0.. Double {
  if index < 0 || index >= distribution.length() {
    0.0
  } else {
    distribution[index]
  }
}

///|
pub fn accept_probability(
  target : Double,
  draft : Double,
) -> Result[Double, ProbabilityError] {
  if !finite_number(target) ||
    !finite_number(draft) ||
    target < 0.0 ||
    target > 1.0 ||
    draft <= 0.0 ||
    draft > 1.0 {
    return Err(InvalidProbability)
  }
  Ok((target / draft).min(1.0))
}

///|
pub fn residual_distribution(
  target : Array[Double],
  draft : Array[Double],
) -> Result[Array[Double], ProbabilityError] {
  if !distribution_is_valid(target) ||
    !distribution_is_valid(draft) ||
    target.length() != draft.length() {
    return Err(InvalidProbability)
  }
  let raw : Array[Double] = []
  let mut total = 0.0
  for index in 0.. Result[Int, ProbabilityError] {
  if !valid_uniform(unit_interval) {
    return Err(InvalidSample)
  }
  if !distribution_is_valid(distribution) {
    return Err(InvalidProbability)
  }
  let mut cumulative = 0.0
  let mut last_positive = 0
  for index in 0.. 0.0 {
      last_positive = index
    }
    cumulative = cumulative + value
    if unit_interval < cumulative {
      return Ok(index)
    }
  }
  // Rounding near one must never select a zero-probability tail token.
  Ok(last_positive)
}