// sequence_replay_buffer.mbt — Per-step continuous replay buffer with
// T-step rollout sampling for recurrent (LSTM/GRU) actor-critic RL.
//
// Stores per-step transitions of (s, a, r, s', done, t) in FIFO-wrapped
// flat Float arrays. On `sample_seq_batch(batch_size, seq_len, rng)`,
// picks `batch_size` starting indices, then for each start collects the
// next `seq_len` consecutive steps; if any step in the window has done=1,
// the rest of the window is zero-padded and a terminal flag is set so the
// caller can mask hidden-state bootstrapping.
//
// Layout conventions (flat row-major, project-wide):
//   states       : [capacity × state_dim]   row i at offset i*state_dim
//   actions      : [capacity × action_dim]
//   rewards      : [capacity]
//   next_states  : [capacity × state_dim]
//   dones        : [capacity]               0.0F / 1.0F
//
// Sample output shape:
//   obs_seq        : [batch_size × seq_len × state_dim]   flat row-major
//   action_seq     : [batch_size × seq_len × action_dim]
//   reward_seq     : [batch_size × seq_len]
//   next_obs_seq   : [batch_size × seq_len × state_dim]
//   done_seq       : [batch_size × seq_len]
//   terminal_mask  : [batch_size × seq_len]               1.0F if step
//                                                       follows a `done`,
//                                                       else 0.0F
//   hidden_init    : [batch_size × hidden_dim]            zeros
//
// API:
//   SequenceReplayBuffer::new(capacity, state_dim, action_dim)
//   push(state, action, reward, next_state, done)
//   len() -> Int
//   sample_seq_batch(batch_size, seq_len, rng)
//     -> (obs_seq, action_seq, reward_seq, next_obs_seq, done_seq,
//         terminal_mask, hidden_init)
//   reset() -> Unit                       (clear stored transitions)
//
// Reference: Hausknecht & Stone 2015 "Deep Recurrent Q-Learning for
// Partially Observable MDPs" (the LSTM-DQN that motivates this design).

///|
/// Per-step continuous replay buffer with sequence-rollout sampling.
pub struct SequenceReplayBuffer {
  capacity : Int
  state_dim : Int
  action_dim : Int
  states : Array[Float]
  actions : Array[Float]
  rewards : Array[Float]
  next_states : Array[Float]
  dones : Array[Float]
  mut size : Int
  mut cursor : Int
}

///|
/// Construct a new sequence replay buffer with the given capacity and
/// observation / action dimensions. All internal arrays are
/// zero-initialised; `size` and `cursor` start at 0.
pub fn SequenceReplayBuffer::new(
  capacity : Int,
  state_dim : Int,
  action_dim : Int,
) -> SequenceReplayBuffer {
  let n = capacity
  let s_len = n * state_dim
  let a_len = n * action_dim
  {
    capacity: n,
    state_dim,
    action_dim,
    states: Array::make(s_len, 0.0F),
    actions: Array::make(a_len, 0.0F),
    rewards: Array::make(n, 0.0F),
    next_states: Array::make(s_len, 0.0F),
    dones: Array::make(n, 0.0F),
    size: 0,
    cursor: 0,
  }
}

///|
/// Number of transitions currently stored (clamped to capacity).
pub fn SequenceReplayBuffer::len(self : SequenceReplayBuffer) -> Int {
  self.size
}

///|
/// Clear the buffer without releasing storage. Both `size` and
/// `cursor` reset to 0; the underlying Float arrays keep their
/// capacity so the buffer can be re-used without re-allocation.
pub fn SequenceReplayBuffer::reset(self : SequenceReplayBuffer) -> Unit {
  self.size = 0
  self.cursor = 0
}

///|
/// Append a transition. Overwrites the oldest slot when the buffer
/// is full. The `done` flag marks a terminal transition; subsequent
/// sample windows crossing a terminal will zero-pad the rest and set
/// `terminal_mask = 1.0F` so the caller can break BPTT.
pub fn SequenceReplayBuffer::push(
  self : SequenceReplayBuffer,
  state : Array[Float],
  action : Array[Float],
  reward : Float,
  next_state : Array[Float],
  done : Bool,
) -> Unit {
  let i = self.cursor
  let s_off = i * self.state_dim
  let a_off = i * self.action_dim
  for k in 0..= self.capacity {
    self.cursor = 0
  }
  if self.size < self.capacity {
    self.size = self.size + 1
  }
}

///|
/// Sample a mini-batch of `batch_size` sequences, each of length
/// `seq_len`. Returns 7 parallel arrays. The starting index for each
/// sequence is drawn uniformly from `[0, size - seq_len)` (clamped).
///
/// Padding behaviour: if a sequence window crosses a `done` boundary,
/// the corresponding `obs_seq` / `action_seq` slots are zero-filled
/// and `terminal_mask[t]` is set to 1.0F. The caller is responsible
/// for resetting hidden state at terminal positions.
pub fn SequenceReplayBuffer::sample_seq_batch(
  self : SequenceReplayBuffer,
  batch_size : Int,
  seq_len : Int,
  rng : Xoshiro,
) -> (Array[Float], Array[Float], Array[Float], Array[Float], Array[Float], Array[Float], Array[Float]) {
  let max_start = self.size - seq_len
  let safe_max = if max_start < 1 {
    0
  } else {
    max_start
  }
  let obs_seq : Array[Float] = Array::make(batch_size * seq_len * self.state_dim, 0.0F)
  let action_seq : Array[Float] = Array::make(batch_size * seq_len * self.action_dim, 0.0F)
  let reward_seq : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let next_obs_seq : Array[Float] = Array::make(batch_size * seq_len * self.state_dim, 0.0F)
  let done_seq : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let terminal_mask : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let hidden_init : Array[Float] = Array::make(batch_size, 0.0F)
  for k in 0.. size),
    // force start_idx = 0 so that every step in the window falls
    // past the end of the buffer and gets zero-padded.
    let mut start_idx = 0
    if safe_max > 0 {
      let (z, _) = box_muller(rng)
      let u_pos = if z < 0.0 {
        -z
      } else {
        z
      }
      start_idx = Float::from_double(u_pos * Double::from_int(safe_max)).to_int()
      if start_idx < 0 {
        start_idx = 0
      } else if start_idx > safe_max {
        start_idx = safe_max
      }
    }
    // Walk forward seq_len steps; break if any earlier step was
    // terminal (subsequent slots get zero-padded + mask set).
    let mut crossed_done = false
    for t in 0..= self.size {
        // Already past the end of the buffer or terminal:
        // leave the slot at 0 (zero-init) and mark the mask.
        terminal_mask[k * seq_len + t] = 1.0F
        done_seq[k * seq_len + t] = 0.0F
      } else {
        let buf_off = global_idx * self.state_dim
        for d in 0.. 0.5F {
          // This slot is the terminal one. Mark mask=1 for the
          // NEXT slot (caller resets hidden state there) but keep
          // the current obs/next_obs visible.
          crossed_done = true
          if t + 1 < seq_len {
            terminal_mask[k * seq_len + t + 1] = 1.0F
          }
        }
      }
    }
  }
  (obs_seq, action_seq, reward_seq, next_obs_seq, done_seq, terminal_mask, hidden_init)
}

// ----------------------------------------------------------------------------
// Deterministic sampling helper.
//
// `sample_seq_batch_at` is identical to `sample_seq_batch` except the start
// index for each batch element is taken from the caller-supplied
// `start_indices` array (length `batch_size`) instead of being drawn from
// an RNG. This is useful for:
//   1. Reproducible / regression tests where a specific window must be
//      observed (the RNG-based path uses box-muller + |z| which produces
//      a half-normal distribution concentrated near 0, so hitting a
//      specific `start_idx` requires many seeded draws).
//   2. Callers that want bit-exact reproducibility across runs (e.g.
//      curriculum learning, evaluation rollouts).
//
// Preconditions:
//   - `start_indices.length() == batch_size`
//   - For each k: `0 <= start_indices[k] <= max_start` where
//     `max_start = size - seq_len` (clamped to 0 when buffer is empty).
//     Out-of-range values are clamped, matching `sample_seq_batch`'s
//     post-RNG clamp behaviour.
///|
/// Same as `sample_seq_batch` but takes explicit start indices (one per
/// batch element) instead of drawing from an RNG. See block comment
/// above for rationale. Returns the same 7-tuple as `sample_seq_batch`.
pub fn SequenceReplayBuffer::sample_seq_batch_at(
  self : SequenceReplayBuffer,
  batch_size : Int,
  seq_len : Int,
  start_indices : Array[Int],
) -> (Array[Float], Array[Float], Array[Float], Array[Float], Array[Float], Array[Float], Array[Float]) {
  let max_start = self.size - seq_len
  let safe_max = if max_start < 1 {
    0
  } else {
    max_start
  }
  let obs_seq : Array[Float] = Array::make(batch_size * seq_len * self.state_dim, 0.0F)
  let action_seq : Array[Float] = Array::make(batch_size * seq_len * self.action_dim, 0.0F)
  let reward_seq : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let next_obs_seq : Array[Float] = Array::make(batch_size * seq_len * self.state_dim, 0.0F)
  let done_seq : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let terminal_mask : Array[Float] = Array::make(batch_size * seq_len, 0.0F)
  let hidden_init : Array[Float] = Array::make(batch_size, 0.0F)
  for k in 0.. safe_max {
        start_idx = safe_max
      }
    }
    // Walk forward seq_len steps; break if any earlier step was
    // terminal (subsequent slots get zero-padded + mask set).
    let mut crossed_done = false
    for t in 0..= self.size {
        // Already past the end of the buffer or terminal:
        // leave the slot at 0 (zero-init) and mark the mask.
        terminal_mask[k * seq_len + t] = 1.0F
        done_seq[k * seq_len + t] = 0.0F
      } else {
        let buf_off = global_idx * self.state_dim
        for d in 0.. 0.5F {
          // This slot is the terminal one. Mark mask=1 for the
          // NEXT slot (caller resets hidden state there) but keep
          // the current obs/next_obs visible.
          crossed_done = true
          if t + 1 < seq_len {
            terminal_mask[k * seq_len + t + 1] = 1.0F
          }
        }
      }
    }
  }
  (obs_seq, action_seq, reward_seq, next_obs_seq, done_seq, terminal_mask, hidden_init)
}