use rstest::rstest;
use sum_segment_tree::SumTree;
#[rstest]
#[case(0, 1)]
#[case(1, 1)]
#[case(2, 2)]
#[case(5, 5)]
#[case(8, 8)]
#[case(1000, 1000)]
#[case(20_000, 20_000)]
fn capacity_is_exact(#[case] requested: usize, #[case] expected: usize) {
let tree = SumTree::new(requested);
assert_eq!(tree.capacity(), expected);
assert!(tree.is_empty());
assert_eq!(tree.total(), 0.0);
}
#[rstest]
#[case(&[], 0.0)]
#[case(&[2.0], 2.0)]
#[case(&[1.0, 2.0, 3.0, 4.0], 10.0)]
#[case(&[0.5, 0.5, 0.5], 1.5)]
fn total_is_the_sum_of_priorities(#[case] priorities: &[f32], #[case] expected: f32) {
let mut tree = SumTree::new(4);
for &p in priorities {
tree.push(p);
}
assert_eq!(tree.total(), expected);
assert_eq!(tree.len(), priorities.len());
}
#[rstest]
#[case(0.0, 0)]
#[case(2.5, 0)]
#[case(5.5, 1)]
#[case(6.5, 2)]
#[case(8.0, 3)]
#[case(9.99, 3)]
fn get_locates_the_weighted_leaf(#[case] s: f32, #[case] expected_index: usize) {
let mut tree = SumTree::new(4);
for &p in &[5.0, 1.0, 1.0, 3.0] {
tree.push(p);
}
let (index, priority) = tree.get(s).unwrap();
assert_eq!(index, expected_index);
assert_eq!(priority, tree.priority(expected_index));
}
#[rstest]
#[case(-1.0, 0)]
#[case(0.0, 0)]
#[case(100.0, 3)]
fn get_clamps_out_of_range_samples(#[case] s: f32, #[case] expected_index: usize) {
let mut tree = SumTree::new(4);
for &p in &[5.0, 1.0, 1.0, 3.0] {
tree.push(p);
}
assert_eq!(tree.get(s).unwrap().0, expected_index);
}
#[test]
fn rebuild_keeps_sums_consistent() {
let mut tree = SumTree::new(8);
for &p in &[1.0, 2.0, 3.0, 4.0, 5.0] {
tree.push(p);
}
tree.update(1, 9.0);
tree.rebuild();
assert_eq!(tree.total(), 1.0 + 9.0 + 3.0 + 4.0 + 5.0);
assert_eq!(tree.get(0.5).unwrap().0, 0);
assert_eq!(tree.get(5.0).unwrap().0, 1);
}
#[rstest]
#[case(3)]
#[case(5)]
#[case(1000)]
#[case(20_000)]
fn non_power_of_two_capacity_keeps_indices_in_range(#[case] capacity: usize) {
let mut tree = SumTree::new(capacity);
assert_eq!(tree.capacity(), capacity);
let mut data = vec![0u32; tree.capacity()];
for k in 0..capacity {
let index = tree.push((k % 7) as f32 + 1.0);
assert!(index < capacity);
data[index] = k as u32;
}
let total = tree.total();
for step in 0..1000 {
let s = (step as f32 + 0.5) / 1000.0 * total;
let (index, _) = tree.get(s).unwrap();
assert!(index < tree.len());
let _ = data[index];
}
}
#[test]
fn get_on_empty_tree_is_none() {
let tree = SumTree::new(4);
assert_eq!(tree.get(0.0), None);
}
#[test]
fn update_changes_total_and_sampling() {
let mut tree = SumTree::new(4);
for &p in &[1.0, 1.0, 1.0, 1.0] {
tree.push(p);
}
assert_eq!(tree.total(), 4.0);
tree.update(2, 7.0);
assert_eq!(tree.total(), 10.0);
assert_eq!(tree.priority(2), 7.0);
assert_eq!(tree.get(5.0).unwrap().0, 2);
}
#[rstest]
#[case(2, &[1.0, 2.0, 3.0, 4.0, 5.0], &[5.0, 4.0], 9.0)]
#[case(4, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[5.0, 6.0, 3.0, 4.0], 18.0)]
fn push_overwrites_oldest_when_full(
#[case] capacity: usize,
#[case] priorities: &[f32],
#[case] expected_leaves: &[f32],
#[case] expected_total: f32,
) {
let mut tree = SumTree::new(capacity);
for &p in priorities {
tree.push(p);
}
assert!(tree.is_full());
assert_eq!(tree.total(), expected_total);
let leaves: Vec<f32> = (0..tree.capacity()).map(|i| tree.priority(i)).collect();
assert_eq!(leaves, expected_leaves);
}
#[test]
fn sampling_distribution_is_proportional() {
let mut tree = SumTree::new(4);
for &p in &[1.0, 3.0, 0.0, 6.0] {
tree.push(p);
}
let steps = 10_000;
let total = tree.total();
let mut counts = [0u32; 4];
for k in 0..steps {
let s = (k as f32 + 0.5) / steps as f32 * total;
counts[tree.get(s).unwrap().0] += 1;
}
let fraction = |i: usize| counts[i] as f32 / steps as f32;
assert!((fraction(0) - 0.1).abs() < 0.01);
assert!((fraction(1) - 0.3).abs() < 0.01);
assert_eq!(counts[2], 0);
assert!((fraction(3) - 0.6).abs() < 0.01);
}