sum-segment-tree 0.2.0

A fixed-capacity sum-segment tree for weighted sampling and prioritized experience replay.
Documentation
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());
}

// Leaves [5, 1, 1, 3] give cumulative intervals 0..5, 5..6, 6..7, 7..10.
// Boundaries fall to the earlier leaf.
#[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);

    // A side buffer sized to capacity must never be indexed out of bounds.
    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);
    // Leaf 2 now owns interval 2..9, so a mid sample lands there.
    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);
}