// n_step_buffer.mbt — Fixed-window buffer for n-step return computation.
//
// An N-step transition packs together the first state and action of a window
// of `n` consecutive environment transitions, with the discounted sum of all
// rewards and the (n+1)-th state as the lookahead. When the lookahead corresponds
// to a terminal state, the bootstrap term is zero (done = true).
//
// This buffer holds up to `n` consecutive (s, a, r, s', done) tuples in a
// FIFO ring. After each push, if it contains at least `n` tuples, the caller
// can call `make_n_step_transition(...)` to extract the n-step transition whose
// root is `n` pushes back.

///|
/// Per-step transition tuple stored in the n-step buffer.
pub struct StepRecord {
  s : Int
  a : Int
  r : Float
  s_next : Int
  done : Bool
}

///|
/// Circular buffer of recent step records.
pub struct NStepBuffer {
  n : Int
  records : Array[StepRecord]
  mut count : Int
  mut head : Int // index of oldest record
}

///|
/// Create an empty N-step buffer holding at most `n` records.
pub fn NStepBuffer::new(n : Int) -> NStepBuffer {
  if n < 1 {
    abort("NStepBuffer::new: n must be >= 1")
  }
  let recs : Array[StepRecord] = []
  for _k in 0.. Int {
  self.count
}

///|
/// Window size (capacity).
pub fn NStepBuffer::capacity(self : NStepBuffer) -> Int {
  self.n
}

///|
/// Push a new step record, overwriting the oldest if the buffer is full.
/// Returns true if the buffer now has `n` records (i.e., a complete n-step
/// window is available for extraction).
pub fn NStepBuffer::push(
  self : NStepBuffer,
  s : Int,
  a : Int,
  r : Float,
  s_next : Int,
  done : Bool,
) -> Bool {
  // Write into the slot at (head + count) % n (oldest if full, else tail).
  let write_at = if self.count < self.n {
    self.head + self.count
  } else {
    self.head
  }
  // Modulo write_at into [0, n).
  let write_idx = write_at - (write_at / self.n) * self.n
  self.records[write_idx] = { s, a, r, s_next, done }
  if self.count < self.n {
    self.count = self.count + 1
  } else {
    // Buffer full: advance head to drop the oldest.
    self.head = self.head + 1
    let _ = self.head - (self.head / self.n) * self.n
    self.head = self.head - (self.head / self.n) * self.n
  }
  // Return whether we now hold a full window.
  self.count == self.n
}

///|
/// Read the i-th record in chronological order (0 = oldest, count-1 = newest).
/// Returns a copy of the record; the caller cannot mutate the buffer via it.
pub fn NStepBuffer::at(self : NStepBuffer, i : Int) -> StepRecord {
  let idx = self.head + i
  let wrapped = idx - (idx / self.n) * self.n
  self.records[wrapped]
}

///|
/// Empty the buffer (used at episode boundaries when we flush remaining
/// transitions to the replay buffer with shorter effective horizons).
pub fn NStepBuffer::reset(self : NStepBuffer) -> Unit {
  self.count = 0
  self.head = 0
}

///|
/// Compute the n-step return for the transition rooted at record index `root`
/// (0 = oldest in the buffer).
///
/// Given records R_root, R_root+1, ..., R_root+n-1, the n-step return is
///
///     G_n = Σ_{k=0..n-1} γ^k · r_{root+k}
///
/// If any intermediate step in the window is terminal (`done = true`), the
/// sum is truncated at the terminal step and the bootstrap is suppressed (the
/// Q-value of s_{terminal+1} is treated as zero).
///
/// Also returns the effective horizon `h` (number of rewards included, 1..=n)
/// and the terminal flag — the caller uses these to decide whether to use the
/// standard n-step bootstrap (h == n, !done) or the truncated form (h < n).
pub fn NStepBuffer::n_step_return(
  self : NStepBuffer,
  root : Int,
  gamma : Float,
) -> (Float, Int, Int, Int, Bool) {
  // Returns: (G_n, h_eff, s_root, s_lookahead, truncated_at_terminal)
  let s_root = self.records[(self.head + root) - ((self.head + root) / self.n) * self.n].s
  let mut g = 0.0F
  let mut disc = 1.0F
  let mut truncated = false
  let mut h_eff = self.n
  let mut s_lookahead = s_root
  // Only iterate over the actual records currently in the buffer (count),
  // not the underlying capacity (which may contain stale zero-initialised
  // records when the buffer is partial).
  let upper = if self.count < self.n { self.count } else { self.n }
  for k in 0.. upper {
    h_eff = upper
  }
  (g, h_eff, s_root, s_lookahead, truncated)
}