///| Sampling controls applied after a model produces logits. Keeping this

///| logic separate from proposal verification makes draft and target policies

///|
/// explicit, reproducible, and independently testable.
pub enum SamplingError {
  InvalidTopK(Int)
  InvalidTopP(Double)
  InvalidMinP(Double)
  InvalidDistribution
  ProbabilityFailure
  InvalidTemperature
} derive(Eq, Debug)

///| A conservative decoding policy. `top_k = 0` means no rank cutoff, while

///|
/// `top_p = 1.0` and `min_p = 0.0` mean their respective filters are off.
pub struct SamplingConfig {
  temperature : Double
  top_k : Int
  top_p : Double
  min_p : Double
}

///|
pub fn SamplingConfig::new(
  temperature? : Double = 1.0,
  top_k? : Int = 0,
  top_p? : Double = 1.0,
  min_p? : Double = 0.0,
) -> Result[SamplingConfig, SamplingError] {
  if !finite_number(temperature) || temperature <= 0.0 {
    return Err(InvalidTemperature)
  }
  if top_k < 0 {
    return Err(InvalidTopK(top_k))
  }
  if !finite_number(top_p) || top_p <= 0.0 || top_p > 1.0 {
    return Err(InvalidTopP(top_p))
  }
  if !finite_number(min_p) || min_p < 0.0 || min_p > 1.0 {
    return Err(InvalidMinP(min_p))
  }
  Ok({ temperature, top_k, top_p, min_p })
}

///|
/// Default is pure temperature sampling with no vocabulary truncation.
pub fn SamplingConfig::default() -> SamplingConfig {
  { temperature: 1.0, top_k: 0, top_p: 1.0, min_p: 0.0 }
}

///|
pub fn SamplingConfig::temperature(self : SamplingConfig) -> Double {
  self.temperature
}

///|
pub fn SamplingConfig::top_k(self : SamplingConfig) -> Int {
  self.top_k
}

///| Return vocabulary indices in descending probability order. Ties retain

///|
/// their lower token id, which is useful when test fixtures need stability.
pub fn ranked_indices(
  distribution : Array[Double],
) -> Result[Array[Int], SamplingError] {
  if distribution.length() == 0 {
    return Err(InvalidDistribution)
  }
  let indices : Array[Int] = []
  for index in 0.. distribution[candidate] {
        indices.insert(position, index)
        inserted = true
        break
      }
    }
    if !inserted {
      indices.push(index)
    }
  }
  Ok(indices)
}

///| Normalize non-negative weights after one or more filters have masked the

///| vocabulary. It rejects an all-zero mask rather than silently sampling an

///|
/// impossible token.
pub fn normalize_weights(
  weights : Array[Double],
) -> Result[Array[Double], SamplingError] {
  if weights.length() == 0 {
    return Err(InvalidDistribution)
  }
  let mut total = 0.0
  for weight in weights {
    if !finite_number(weight) || weight < 0.0 {
      return Err(InvalidDistribution)
    }
    total = total + weight
  }
  if !finite_number(total) || total <= 0.0 {
    return Err(InvalidDistribution)
  }
  let normalized : Array[Double] = []
  for weight in weights {
    normalized.push(weight / total)
  }
  Ok(normalized)
}

///|
/// Keep the `limit` most likely tokens. `limit = 0` keeps every token.
pub fn filter_top_k(
  distribution : Array[Double],
  limit : Int,
) -> Result[Array[Double], SamplingError] {
  if limit < 0 {
    return Err(InvalidTopK(limit))
  }
  let ranking = match ranked_indices(distribution) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  if limit == 0 || limit >= ranking.length() {
    return normalize_weights(distribution)
  }
  let filtered : Array[Double] = []
  for _ in distribution {
    filtered.push(0.0)
  }
  for position in 0.. Result[Array[Double], SamplingError] {
  if threshold <= 0.0 || threshold > 1.0 {
    return Err(InvalidTopP(threshold))
  }
  let normalized = match normalize_weights(distribution) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let ranking = match ranked_indices(normalized) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let filtered : Array[Double] = []
  for _ in normalized {
    filtered.push(0.0)
  }
  let mut cumulative = 0.0
  for index in ranking {
    filtered[index] = normalized[index]
    cumulative = cumulative + normalized[index]
    if cumulative >= threshold {
      break
    }
  }
  normalize_weights(filtered)
}

///| Min-p filtering uses the most likely token as an adaptive reference:

///|
/// every retained probability is at least `threshold * max_probability`.
pub fn filter_min_p(
  distribution : Array[Double],
  threshold : Double,
) -> Result[Array[Double], SamplingError] {
  if threshold < 0.0 || threshold > 1.0 {
    return Err(InvalidMinP(threshold))
  }
  let normalized = match normalize_weights(distribution) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let mut maximum = normalized[0]
  for value in normalized {
    if value > maximum {
      maximum = value
    }
  }
  let filtered : Array[Double] = []
  for value in normalized {
    filtered.push(if value >= maximum * threshold { value } else { 0.0 })
  }
  normalize_weights(filtered)
}

///| Transform logits into a filtered categorical distribution. Filter order is

///| intentionally fixed: temperature, top-k, top-p, then min-p. This makes a

///|
/// configuration portable across a draft worker and a target worker.
pub fn distribution_for_sampling(
  logits : Array[Double],
  config : SamplingConfig,
) -> Result[Array[Double], SamplingError] {
  let initial = match softmax(logits, temperature=config.temperature) {
    Ok(value) => value
    Err(_) => return Err(ProbabilityFailure)
  }
  let after_k = match filter_top_k(initial, config.top_k) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let after_p = match filter_top_p(after_k, config.top_p) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  filter_min_p(after_p, config.min_p)
}

///| Sample a token from filtered logits with a caller-supplied random number.

///| The caller owns randomness so command-line runs and regression tests can

///|
/// be replayed bit-for-bit.
pub fn sample_logits(
  logits : Array[Double],
  config : SamplingConfig,
  unit_interval : Double,
) -> Result[Int, SamplingError] {
  let distribution = match distribution_for_sampling(logits, config) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  match sample_categorical(distribution, unit_interval) {
    Ok(token) => Ok(token)
    Err(_) => Err(ProbabilityFailure)
  }
}