// prioritized_replay.mbt — Prioritized Experience Replay buffer (Schaul 2016).
//
// Like `dqn.mbt`'s `ReplayBuffer`, this stores transitions (s, a, r, s', done)
// and supports mini-batch sampling. The difference is that each transition has
// a *priority* p_i, and transitions are sampled with probability
//
//     P(i) = p_i^α / Σ_j p_j^α
//
// where α ∈ [0, 1] controls how much prioritization is used (α = 0 → uniform).
//
// Sampling corrects for the bias introduced by non-uniform sampling with
// importance-sampling weights
//
//     w_i = (N · P(i))^{-β}   /   max_j w_j
//
// where β ∈ [0, 1] anneals from 0 (no correction) toward 1 (full correction).
//
// Priorities are stored in a `SumTree` so cumulative-sum sampling is O(log N)
// and batch priority updates are O(log N) each.

///|
/// Prioritized replay buffer with proportional priorities and importance-sampling
/// weights. New transitions are inserted at priority = `max_priority` so that
/// every transition is guaranteed to be sampled at least once.
pub struct PrioritizedReplayBuffer {
  capacity : Int
  states : Array[Int]
  actions : Array[Int]
  rewards : Array[Float]
  next_states : Array[Int]
  dones : Array[Bool]
  tree : SumTree
  mut max_priority : Float
  alpha : Float
  beta : Float
  epsilon : Float
  mut size : Int
  mut cursor : Int
}

///|
/// Construct a new PER buffer.
///
/// * `capacity`     — number of transitions to hold
/// * `alpha`        — prioritization exponent (0 = uniform, 1 = fully prioritized)
/// * `beta`         — importance-sampling exponent (anneal 0 → 1 during training)
/// * `eps`          — small constant added to |δ| before exponentiating, ensures
///                    every transition is sampleable even with zero TD error
pub fn PrioritizedReplayBuffer::new(
  capacity : Int,
  alpha : Float,
  beta : Float,
  eps : Float,
) -> PrioritizedReplayBuffer {
  let tree = SumTree::new(capacity)
  {
    capacity,
    states: Array::make(capacity, 0),
    actions: Array::make(capacity, 0),
    rewards: Array::make(capacity, 0.0F),
    next_states: Array::make(capacity, 0),
    dones: Array::make(capacity, false),
    tree,
    max_priority: 1.0F,
    alpha,
    beta,
    epsilon: eps,
    size: 0,
    cursor: 0,
  }
}

///|
/// Current number of stored transitions.
pub fn PrioritizedReplayBuffer::len(self : PrioritizedReplayBuffer) -> Int {
  self.size
}

///|
/// Underlying slot capacity.
pub fn PrioritizedReplayBuffer::capacity(self : PrioritizedReplayBuffer) -> Int {
  self.capacity
}

///|
/// Total priority mass (= root of the sum tree). Used to normalise per-step
/// sampling segments and to detect "all priorities zero" edge cases.
pub fn PrioritizedReplayBuffer::total_priority(
  self : PrioritizedReplayBuffer,
) -> Float {
  self.tree.total()
}

///|
/// Maximum priority observed so far. New transitions are inserted at this
/// priority so they are guaranteed to be sampled at least once.
pub fn PrioritizedReplayBuffer::max_priority(
  self : PrioritizedReplayBuffer,
) -> Float {
  self.max_priority
}

///|
/// Append a transition. New transitions get priority = max_priority (Schmidt's
/// trick: ensures P(i) > 0 even when the initial |δ| is zero). Overwrites
/// oldest slot when the buffer is full (FIFO ring).
pub fn PrioritizedReplayBuffer::push(
  self : PrioritizedReplayBuffer,
  s : Int,
  a : Int,
  r : Float,
  s_next : Int,
  done : Bool,
) -> Unit {
  // Choose slot: while buffer is not full, append; else overwrite oldest.
  let slot = if self.size < self.capacity {
    self.size
  } else {
    self.cursor
  }
  // Write data.
  self.states[slot] = s
  self.actions[slot] = a
  self.rewards[slot] = r
  self.next_states[slot] = s_next
  self.dones[slot] = done
  // Write priority.
  let p = self.max_priority
  if self.size < self.capacity {
    let tree_idx = self.tree.add(p)
    let _ = tree_idx
    self.size = self.size + 1
  } else {
    // Overwrite priority at slot.
    self.tree.set_at(slot, p)
  }
  // Advance cursor.
  self.cursor = self.cursor + 1
  if self.cursor >= self.capacity {
    self.cursor = 0
  }
}

///|
/// Update priorities for a list of (slot_index, |δ|+ε) pairs after a training
/// step. `slot_indices` is the array returned by `sample`. `td_errors` may be
/// signed (the absolute value is taken internally).
///
/// Bumps `max_priority` upward if any new |δ|+ε exceeds the current max.
pub fn PrioritizedReplayBuffer::update_priorities(
  self : PrioritizedReplayBuffer,
  slot_indices : Array[Int],
  td_errors : Array[Float],
) -> Unit {
  let n = slot_indices.length()
  for k in 0.. self.max_priority {
      self.max_priority = p
    }
  }
}

///|
/// Proportional-priority sample. Returns:
///
///   (states, actions, rewards, next_states, dones,
///    slot_indices, is_weights, p_i_per_index)
///
/// where:
///
/// * `slot_indices` is the leaf slot index for each sampled transition (for use
///   with `update_priorities`).
/// * `is_weights` is the importance-sampling weight normalised so the maximum
///   weight in the batch equals 1.0.
/// * `p_i_per_index` is the raw sampling probability P(i) used for the weight,
///   exposed so tests can verify sampling bias.
pub fn PrioritizedReplayBuffer::sample(
  self : PrioritizedReplayBuffer,
  batch_size : Int,
  rng : Xoshiro,
) -> (
  Array[Int],
  Array[Int],
  Array[Float],
  Array[Int],
  Array[Bool],
  Array[Int],
  Array[Float],
  Array[Float],
) {
  let states : Array[Int] = Array::make(batch_size, 0)
  let actions : Array[Int] = Array::make(batch_size, 0)
  let rewards : Array[Float] = Array::make(batch_size, 0.0F)
  let next_states : Array[Int] = Array::make(batch_size, 0)
  let dones : Array[Bool] = Array::make(batch_size, false)
  let slot_indices : Array[Int] = Array::make(batch_size, 0)
  let is_weights : Array[Float] = Array::make(batch_size, 0.0F)
  let probs : Array[Float] = Array::make(batch_size, 0.0F)
  let total = self.tree.total()
  if total <= 0.0F || self.size == 0 {
    // Empty buffer or all priorities zero: leave arrays at their default
    // values. Callers must check `len()` before sampling.
    let _ = rng
    return (
      states, actions, rewards, next_states, dones, slot_indices, is_weights, probs,
    )
  }
  // Standard PER: partition [0, total] into batch_size equal segments.
  let seg = total / Float::from_int(batch_size)
  // First pass: sample one leaf per segment, collect p_i.
  for k in 0..= total {
      target = total - 1.0e-6F
    }
    if target < 0.0F {
      target = 0.0F
    }
    let (tree_idx, p_i) = self.tree.find_prefix(target)
    let slot = self.tree.tree_to_leaf(tree_idx)
    // Clamp slot into [0, size).
    let safe_slot = if slot < 0 {
      0
    } else if slot >= self.size {
      self.size - 1
    } else {
      slot
    }
    states[k] = self.states[safe_slot]
    actions[k] = self.actions[safe_slot]
    rewards[k] = self.rewards[safe_slot]
    next_states[k] = self.next_states[safe_slot]
    dones[k] = self.dones[safe_slot]
    slot_indices[k] = safe_slot
    // P(i) = p_i / total
    probs[k] = p_i / total
  }
  // Second pass: compute IS weights via Double arithmetic for stability.
  let raw_weights : Array[Float] = Array::make(batch_size, 0.0F)
  let n_float = Float::from_int(if self.size > 0 { self.size } else { 1 })
  for k in 0.. max_w {
      max_w = raw_weights[k]
    }
  }
  let max_w_safe = if max_w < 1.0e-12F { 1.0e-12F } else { max_w }
  for k in 0..