diff --git a/src/data_structures/segment_tree.rs b/src/data_structures/segment_tree.rs index f569381967e..5b5985712f8 100644 --- a/src/data_structures/segment_tree.rs +++ b/src/data_structures/segment_tree.rs @@ -80,12 +80,13 @@ where let mut left = range.start + self.size; let mut right = range.end + self.size; - let mut result = None; + let mut left_result = None; + let mut right_result = None; // Iterate through the segment tree to accumulate results while left < right { if left % 2 == 1 { - result = Some(match result { + left_result = Some(match left_result { None => self.nodes[left], Some(old) => (self.merge_fn)(old, self.nodes[left]), }); @@ -93,16 +94,20 @@ where } if right % 2 == 1 { right -= 1; - result = Some(match result { + right_result = Some(match right_result { None => self.nodes[right], - Some(old) => (self.merge_fn)(old, self.nodes[right]), + Some(old) => (self.merge_fn)(self.nodes[right], old), }); } left /= 2; right /= 2; } - Ok(result) + Ok(match (left_result, right_result) { + (Some(left), Some(right)) => Some((self.merge_fn)(left, right)), + (Some(result), None) | (None, Some(result)) => Some(result), + (None, None) => None, + }) } /// Updates the value at the specified index in the segment tree. @@ -139,8 +144,142 @@ where #[cfg(test)] mod tests { use super::*; + use std::cell::Cell; use std::cmp::{max, min}; + #[derive(Debug, Default, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] + struct Affine { + multiplier: i64, + offset: i64, + } + + impl Affine { + fn then(self, next: Self) -> Self { + Self { + multiplier: next.multiplier * self.multiplier, + offset: next.multiplier * self.offset + next.offset, + } + } + } + + fn sample_affine_transforms(size: usize) -> Vec { + (0..size) + .map(|index| Affine { + multiplier: 1 + (index % 3) as i64, + offset: (index % 5) as i64 - 2, + }) + .collect() + } + + fn assert_queries_match_slice( + tree: &SegmentTree Affine>, + values: &[Affine], + ) { + for start in 0..values.len() { + for end in start..=values.len() { + let expected = values[start..end].iter().copied().reduce(Affine::then); + assert_eq!( + tree.query(start..end), + Ok(expected), + "range {start}..{end}, values {values:?}" + ); + } + } + } + + #[test] + fn test_noncommutative_segments() { + let values = [ + Affine { + multiplier: 2, + offset: 1, + }, + Affine { + multiplier: 3, + offset: 4, + }, + Affine { + multiplier: 5, + offset: 2, + }, + Affine { + multiplier: 7, + offset: 3, + }, + ]; + let tree = SegmentTree::from_vec(&values, Affine::then); + assert_eq!( + tree.query(0..3), + Ok(Some(Affine { + multiplier: 30, + offset: 37, + })) + ); + assert_eq!( + tree.query(1..4), + Ok(Some(Affine { + multiplier: 105, + offset: 157, + })) + ); + let values = sample_affine_transforms(8); + let tree = SegmentTree::from_vec(&values, Affine::then); + let expected = values[1..7].iter().copied().reduce(Affine::then); + assert_eq!(tree.query(1..7), Ok(expected)); + } + + #[test] + fn test_noncommutative_all_ranges() { + for size in [1, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17] { + let values = sample_affine_transforms(size); + let tree = SegmentTree::from_vec(&values, Affine::then); + assert_queries_match_slice(&tree, &values); + } + } + + #[test] + fn test_noncommutative_updates() { + for size in [3, 4, 5, 8, 9, 17] { + let mut values = sample_affine_transforms(size); + let mut tree = SegmentTree::from_vec(&values, Affine::then); + for index in [0, size / 2, size - 1, 0] { + let replacement = Affine { + multiplier: 2, + offset: -3, + }; + values[index] = replacement; + assert_eq!(tree.update(index, replacement), Ok(())); + assert_queries_match_slice(&tree, &values); + assert_eq!(tree.update(index, replacement), Ok(())); + assert_queries_match_slice(&tree, &values); + } + } + } + + #[test] + fn test_noncommutative_empty_and_singleton_ranges() { + for size in [1, 4, 5, 8] { + let values = sample_affine_transforms(size); + let calls = Cell::new(0); + let tree = SegmentTree::from_vec(&values, |left: Affine, right| { + calls.set(calls.get() + 1); + left.then(right) + }); + let calls_after_build = calls.get(); + assert_eq!( + tree.query(values.len()..values.len()), + Err(SegmentTreeError::InvalidRange) + ); + assert_eq!(calls.get(), calls_after_build); + for (index, value) in values.iter().enumerate() { + assert_eq!(tree.query(index..index), Ok(None)); + assert_eq!(calls.get(), calls_after_build); + assert_eq!(tree.query(index..index + 1), Ok(Some(*value))); + assert_eq!(calls.get(), calls_after_build); + } + } + } + #[test] fn test_min_segments() { let vec = vec![-30, 2, -4, 7, 3, -5, 6, 11, -20, 9, 14, 15, 5, 2, -8];