///|
pub enum StopReason {
  Limit
  EndToken
  Observer
} derive(Eq, Debug)

///|
pub fn StopReason::limit() -> StopReason {
  Limit
}

///|
pub fn StopReason::end_token() -> StopReason {
  EndToken
}

///|
pub fn StopReason::observer() -> StopReason {
  Observer
}

///|
pub struct Generation {
  tokens : Array[Int]
  steps : Array[Step]
  reason : StopReason
} derive(Eq, Debug)

///|
pub fn Generation::tokens(self : Generation) -> Array[Int] {
  self.tokens.copy()
}

///|
pub fn Generation::steps(self : Generation) -> Array[Step] {
  self.steps.copy()
}

///|
pub fn Generation::reason(self : Generation) -> StopReason {
  self.reason
}

///|
/// Drive an autoregressive model callback until an end token or step limit.
/// The callback receives a copy of the generated prefix and must return the
/// next logit row. The RNG is consumed only after a model row is available.
/// Any error preserves successful prior sampler steps and RNG draws.
pub fn generate(
  sampler : Sampler,
  rng : ParkMiller,
  model : (Array[Int]) -> Result[Array[Double], SamplingError],
  limit : Int,
  end_token : Int?,
) -> Result[Generation, SamplingError] {
  generate_from_prompt(sampler, rng, model, [], limit, end_token)
}

///|
/// Generate after a pre-existing prompt. The callback sees prompt followed by
/// generated tokens; the returned token list contains only new tokens.
pub fn generate_from_prompt(
  sampler : Sampler,
  rng : ParkMiller,
  model : (Array[Int]) -> Result[Array[Double], SamplingError],
  prompt : Array[Int],
  limit : Int,
  end_token : Int?,
) -> Result[Generation, SamplingError] {
  let steps : Array[Step] = []
  let result = match
    generate_stream(sampler, rng, model, prompt, limit, end_token, fn(step) {
      steps.push(step)
      true
    }) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  Ok({ tokens: result.tokens(), steps, reason: result.reason(), })
}

///|
pub struct StreamGeneration {
  tokens : Array[Int]
  reason : StopReason
} derive(Eq, Debug)

///|
pub fn StreamGeneration::tokens(self : StreamGeneration) -> Array[Int] {
  self.tokens.copy()
}

///|
pub fn StreamGeneration::count(self : StreamGeneration) -> Int {
  self.tokens.length()
}

///|
pub fn StreamGeneration::reason(self : StreamGeneration) -> StopReason {
  self.reason
}

///|
/// Stream each successful decision to an observer and avoid retaining Step
/// metadata. Return false from the observer to stop after that token. End
/// token detection takes priority when both conditions apply to one step.
pub fn generate_stream(
  sampler : Sampler,
  rng : ParkMiller,
  model : (Array[Int]) -> Result[Array[Double], SamplingError],
  prompt : Array[Int],
  limit : Int,
  end_token : Int?,
  observer : (Step) -> Bool,
) -> Result[StreamGeneration, SamplingError] {
  if limit < 0 {
    return Err(InvalidParameter("generation limit must be non-negative"))
  }
  match end_token {
    Some(token) =>
      if token < 0 {
        return Err(InvalidParameter("end token must be non-negative"))
      }
    None => ()
  }
  for token in prompt {
    if token < 0 {
      return Err(InvalidParameter("prompt token must be non-negative"))
    }
  }
  let tokens : Array[Int] = []
  for _ in 0.. value
      Err(error) => return Err(error)
    }
    let step = match sampler.sample_with_rng(logits, rng) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    tokens.push(step.token())
    let keep_going = observer(step)
    match end_token {
      Some(token) =>
        if step.token() == token {
          return Ok({ tokens, reason: EndToken, })
        }
      None => ()
    }
    if !keep_going {
      return Ok({ tokens, reason: Observer, })
    }
  }
  Ok({ tokens, reason: Limit, })
}