// sum_tree.mbt — Segmented sum-tree for proportional prioritized sampling.
//
// Reference: Schaul et al. 2016, "Prioritized Experience Replay" (PER), ICLR.
//
// A sum-tree is a complete binary tree stored in a flat array:
//   - Indices [0..n_leaves)        hold the *leaf* priorities p_i (data slots)
//   - Indices [n_leaves..2*n-1)    hold internal nodes, each = sum(children)
//   - Index 0                      is the root, holding Σ p_i (total priority).
//
// Sampling in O(log N):
//   1. Draw u ~ Uniform(0, total)
//   3. Walk down from root: at each internal node, compare u with left child's
//      sum; if u <= left, descend left and keep u; else descend right with
//      u -= left_sum. The leaf reached is the sampled transition.
//
// Update is O(log N): propagate the priority delta up from the leaf to root.
//
// Tree layout (capacity = 4 leaves):
//
//        tree[0] = total
//       /        \
//   tree[1]      tree[2]
//   /   \       /    \
// t[3]  t[4]  t[5]  t[6]
//  p_0  p_1   p_2   p_3
//
// Capacity must be a power of two for a perfectly balanced tree, but the code
// below accepts any positive capacity by treating `tree` of size `2 * cap - 1`
// with explicit child index math.

///|
/// Flat-array sum tree. Holds up to `capacity` non-negative priorities.
pub struct SumTree {
  capacity : Int
  tree : Array[Float]
  mut data_count : Int
}

///|
/// Create an empty sum tree. All priorities start at 0.
pub fn SumTree::new(capacity : Int) -> SumTree {
  // Need 2*capacity - 1 slots: (capacity - 1) internal + capacity leaves.
  let n_total = if capacity <= 0 { 1 } else { 2 * capacity - 1 }
  let tree : Array[Float] = Array::make(n_total, 0.0F)
  { capacity, tree, data_count: 0 }
}

///|
/// Number of priority slots (capacity in leaves).
pub fn SumTree::len(self : SumTree) -> Int {
  self.data_count
}

///|
/// Underlying slot capacity (in leaves). Slots are filled lazily as data is added.
pub fn SumTree::capacity(self : SumTree) -> Int {
  self.capacity
}

///|
/// Total priority sum = root node.
pub fn SumTree::total(self : SumTree) -> Float {
  if self.tree.length() == 0 {
    0.0F
  } else {
    self.tree[0]
  }
}

///|
/// Get priority at a leaf slot index (NOT a tree index). Leaf i is stored at
/// tree index `capacity - 1 + i` in our layout.
pub fn SumTree::get_leaf(self : SumTree, slot : Int) -> Float {
  let ti = self.leaf_to_tree(slot)
  if ti < 0 || ti >= self.tree.length() {
    0.0F
  } else {
    self.tree[ti]
  }
}

///|
/// Internal helper: leaf slot → tree index.
fn SumTree::leaf_to_tree(self : SumTree, slot : Int) -> Int {
  self.capacity - 1 + slot
}

///|
/// Internal helper: tree index → leaf slot.
fn SumTree::tree_to_leaf(self : SumTree, tree_idx : Int) -> Int {
  tree_idx - (self.capacity - 1)
}

///|
/// Propagate a priority delta from a tree index up to the root.
fn SumTree::propagate(self : SumTree, tree_idx_in : Int, delta : Float) -> Unit {
  let mut tree_idx = tree_idx_in
  while tree_idx > 0 {
    let parent = (tree_idx - 1) / 2
    self.tree[parent] = self.tree[parent] + delta
    tree_idx = parent
  }
}

///|
/// Set priority at a tree index, propagating the delta up to the root.
/// Caller computes delta as (new - old) — this method does NOT read the
/// existing value, which keeps the cost O(log N) without an extra read.
pub fn SumTree::set_priority(
  self : SumTree,
  tree_idx : Int,
  new_priority : Float,
) -> Unit {
  let old = self.tree[tree_idx]
  let delta = new_priority - old
  self.tree[tree_idx] = new_priority
  if delta != 0.0F {
    self.propagate(tree_idx, delta)
  }
}

///|
/// Add a new priority to the next empty leaf slot (FIFO order). Returns the
/// tree index of the leaf, or -1 if the tree is at full capacity.
///
/// On overwrite (when data_count == capacity), the call is a no-op and
/// returns -1. Use `set_at(slot, ...)` to overwrite an existing slot.
pub fn SumTree::add(self : SumTree, priority : Float) -> Int {
  if self.data_count >= self.capacity {
    return -1
  }
  let slot = self.data_count
  let tree_idx = self.leaf_to_tree(slot)
  let old = self.tree[tree_idx]
  let delta = priority - old
  self.tree[tree_idx] = priority
  if delta != 0.0F {
    self.propagate(tree_idx, delta)
  }
  self.data_count = self.data_count + 1
  tree_idx
}

///|
/// Set priority at an existing leaf slot (used to update after TD-error updates).
pub fn SumTree::set_at(
  self : SumTree,
  slot : Int,
  new_priority : Float,
) -> Unit {
  let tree_idx = self.leaf_to_tree(slot)
  self.set_priority(tree_idx, new_priority)
}

///|
/// Cumulative-sum lookup: find the leaf whose prefix sum covers `target_sum`.
///
/// Algorithm (Schaul 2016, Algorithm 1, simplified):
///   parent = root (idx 0)
///
///   loop:
///     left  = 2*parent + 1
///     right = 2*parent + 2
///     if left >= len(tree):  // parent is a leaf
///       return (parent, tree[parent])
///     if target_sum <= tree[left]:
///       parent = left
///     else:
///       target_sum -= tree[left]
///       parent = right
///
/// Returns (tree_idx, priority_at_that_leaf). If `target_sum` exceeds the
/// total, the rightmost leaf is returned.
pub fn SumTree::find_prefix(self : SumTree, target_sum_in : Float) -> (Int, Float) {
  let n = self.tree.length()
  if n == 0 || self.data_count == 0 {
    return (0, 0.0F)
  }
  let mut parent = 0
  let mut target_sum = target_sum_in
  let total = self.tree[0]
  if target_sum > total {
    target_sum = total
  }
  if target_sum < 0.0F {
    target_sum = 0.0F
  }
  // Walk down until parent points at a leaf.
  while true {
    let left = 2 * parent + 1
    if left >= n {
      // parent is a leaf.
      return (parent, self.tree[parent])
    }
    let right = left + 1
    let left_sum = self.tree[left]
    if target_sum <= left_sum {
      parent = left
    } else {
      target_sum = target_sum - left_sum
      parent = right
    }
  }
  // Unreachable; the loop always returns from inside.
  (0, 0.0F)
}

///|
/// Batch cumulative-sum lookup: draw `n_segments` uniform segments over [0, sum)
/// and return one leaf per segment. Segments partition [0, total] evenly, so
/// over many calls each leaf's selection probability approaches p_i / Σp_j.
pub fn SumTree::find_prefix_batch(
  self : SumTree,
  segment_values : Array[Float],
) -> Array[(Int, Float)] {
  let n = segment_values.length()
  let out : Array[(Int, Float)] = Array::make(n, (0, 0.0F))
  for k in 0..