///|
/// Apply an allow-mask. This is useful when a grammar or host application
/// restricts the next token vocabulary. At least one allowed finite logit is
/// required. The input array is not mutated.
pub fn mask_logits(
  logits : Array[Double],
  allowed : Array[Bool],
) -> Result[Array[Double], SamplingError] {
  if logits.length() != allowed.length() {
    return Err(InvalidParameter("mask length must match vocabulary"))
  }
  match check_logits(logits) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let output : Array[Double] = []
  for i in 0.. Ok(output)
    Err(error) => Err(error)
  }
}

///|
/// Mask a sparse list of forbidden token ids without building a full Boolean
/// mask in the caller. Repeated ids have the same effect as one occurrence.
pub fn block_token_ids(
  logits : Array[Double],
  blocked : Array[Int],
) -> Result[Array[Double], SamplingError] {
  let allowed = Array::make(logits.length(), true)
  for token in blocked {
    if token < 0 || token >= logits.length() {
      return Err(InvalidParameter("blocked token outside vocabulary"))
    }
    allowed[token] = false
  }
  mask_logits(logits, allowed)
}

///|
/// Keep only a sparse set of allowed token ids. A valid nonempty support must
/// remain after applying the allowlist and any masks already in the logits.
pub fn allow_token_ids(
  logits : Array[Double],
  allowed_tokens : Array[Int],
) -> Result[Array[Double], SamplingError] {
  let allowed = Array::make(logits.length(), false)
  for token in allowed_tokens {
    if token < 0 || token >= logits.length() {
      return Err(InvalidParameter("allowed token outside vocabulary"))
    }
    allowed[token] = true
  }
  mask_logits(logits, allowed)
}

///|
/// Add a per-token logit bias. Negative infinity in the base row remains
/// masked, and every bias must be finite. A bias of zero leaves the row alone.
pub fn bias_logits(
  logits : Array[Double],
  biases : Array[Double],
) -> Result[Array[Double], SamplingError] {
  if logits.length() != biases.length() {
    return Err(InvalidParameter("bias length must match vocabulary"))
  }
  match check_logits(logits) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let output : Array[Double] = []
  for i in 0.. Result[Array[Double], SamplingError] {
  if !finite(penalty) || penalty < 1.0 {
    return Err(InvalidParameter("repetition penalty must be at least one"))
  }
  match check_logits(logits) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let output = logits.copy()
  let seen = Array::make(logits.length(), false)
  for token in history {
    if token < 0 || token >= logits.length() {
      return Err(InvalidParameter("history token outside vocabulary"))
    }
    if !seen[token] {
      seen[token] = true
      let value = output[token]
      if finite(value) {
        output[token] = if value < 0.0 {
          value * penalty
        } else {
          value / penalty
        }
      }
    }
  }
  Ok(output)
}

///|
/// Subtract an additive presence cost once per seen token and an additive
/// frequency cost for every occurrence. Costs may be negative to reward
/// repetition. This is a preprocessing option outside Mirostat feedback.
pub fn penalize_counts(
  logits : Array[Double],
  history : Array[Int],
  presence : Double,
  frequency : Double,
) -> Result[Array[Double], SamplingError] {
  if !finite(presence) || !finite(frequency) {
    return Err(InvalidParameter("presence and frequency costs must be finite"))
  }
  match check_logits(logits) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let counts = Array::make(logits.length(), 0)
  for token in history {
    if token < 0 || token >= logits.length() {
      return Err(InvalidParameter("history token outside vocabulary"))
    }
    if counts[token] == 2147483647 {
      return Err(NumericalFailure)
    }
    counts[token] = counts[token] + 1
  }
  let output = logits.copy()
  for token in 0.. 0 && finite(output[token]) {
      let adjusted = output[token] -
        presence -
        frequency * counts[token].to_double()
      if !finite(adjusted) {
        return Err(NumericalFailure)
      }
      output[token] = adjusted
    }
  }
  Ok(output)
}