diff --git a/libs/persistent_range_query/src/lib.rs b/libs/persistent_range_query/src/lib.rs index 102755bd21..03c9b27377 100644 --- a/libs/persistent_range_query/src/lib.rs +++ b/libs/persistent_range_query/src/lib.rs @@ -2,6 +2,7 @@ use std::ops::Range; pub mod naive; pub mod ops; +pub mod segment_tree; /// Should be a monoid: /// * Identity element: for all a: combine(new_for_empty_range(), a) = combine(a, new_for_empty_range()) = a diff --git a/libs/persistent_range_query/src/segment_tree.rs b/libs/persistent_range_query/src/segment_tree.rs new file mode 100644 index 0000000000..ef63e513db --- /dev/null +++ b/libs/persistent_range_query/src/segment_tree.rs @@ -0,0 +1,255 @@ +//! # Segment Tree +//! It is a competitive programming folklore data structure. Do not confuse with the interval tree. + +use crate::{LazyRangeInitializer, PersistentVecStorage, RangeQueryResult, VecReadableVersion}; +use std::ops::Range; +use std::rc::Rc; + +pub trait MidpointableKey: Clone + Ord + Sized { + fn midpoint(range: &Range) -> Self; +} + +pub trait RangeModification: Clone + crate::RangeModification {} + +// TODO: use trait alias when stabilized +impl, Key> RangeModification for T {} + +#[derive(Debug)] +struct Node, Key> { + result: Modification::Result, + modify_children: Modification, + left: Option>, + right: Option>, +} + +// Manual implementation because we don't need `Key: Clone` for this, unlike with `derive`. +impl, Key> Clone for Node { + fn clone(&self) -> Self { + Node { + result: self.result.clone(), + modify_children: self.modify_children.clone(), + left: self.left.clone(), + right: self.right.clone(), + } + } +} + +impl, Key> Node { + fn new>( + range: &Range, + initializer: &Initializer, + ) -> Self { + Node { + result: initializer.get(range), + modify_children: Modification::no_op(), + left: None, + right: None, + } + } + + pub fn apply(&mut self, modification: &Modification, range: &Range) { + modification.apply(&mut self.result, range); + Modification::compose(modification, &mut self.modify_children); + if self.modify_children.is_reinitialization() { + self.left = None; + self.right = None; + } + } + + pub fn force_children>( + &mut self, + initializer: &Initializer, + range_left: &Range, + range_right: &Range, + ) { + let left = Rc::make_mut( + self.left + .get_or_insert_with(|| Rc::new(Node::new(&range_left, initializer))), + ); + let right = Rc::make_mut( + self.right + .get_or_insert_with(|| Rc::new(Node::new(&range_right, initializer))), + ); + left.apply(&self.modify_children, &range_left); + right.apply(&self.modify_children, &range_right); + self.modify_children = Modification::no_op(); + } + + pub fn recalculate_from_children(&mut self, range_left: &Range, range_right: &Range) { + assert!(self.modify_children.is_no_op()); + assert!(self.left.is_some()); + assert!(self.right.is_some()); + self.result = Modification::Result::combine( + &self.left.as_ref().unwrap().result, + &range_left, + &self.right.as_ref().unwrap().result, + &range_right, + ); + } +} + +fn split_range(range: &Range) -> (Range, Range) { + let range_left = range.start.clone()..MidpointableKey::midpoint(range); + let range_right = range_left.end.clone()..range.end.clone(); + (range_left, range_right) +} + +pub struct PersistentSegmentTreeVersion< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: Clone, +> { + root: Rc>, + all_keys: Range, + initializer: Rc, +} + +// Manual implementation because we don't need `Key: Clone` for this, unlike with `derive`. +impl< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: Clone, + > Clone for PersistentSegmentTreeVersion +{ + fn clone(&self) -> Self { + Self { + root: self.root.clone(), + all_keys: self.all_keys.clone(), + initializer: self.initializer.clone(), + } + } +} + +fn get< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, +>( + node: &mut Rc>, + node_keys: &Range, + initializer: &Initializer, + keys: &Range, +) -> Modification::Result { + if node_keys.end <= keys.start || keys.end <= node_keys.start { + return Modification::Result::new_for_empty_range(); + } + if keys.start <= node_keys.start && node_keys.end <= keys.end { + return node.result.clone(); + } + let node = Rc::make_mut(node); + let (left_keys, right_keys) = split_range(node_keys); + node.force_children(initializer, &left_keys, &right_keys); + let mut result = get(node.left.as_mut().unwrap(), &left_keys, initializer, keys); + Modification::Result::add( + &mut result, + &left_keys, + &get(node.right.as_mut().unwrap(), &right_keys, initializer, keys), + &right_keys, + ); + result +} + +fn modify< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, +>( + node: &mut Rc>, + node_keys: &Range, + initializer: &Initializer, + keys: &Range, + modification: &Modification, +) { + if modification.is_no_op() || node_keys.end <= keys.start || keys.end <= node_keys.start { + return; + } + let node = Rc::make_mut(node); + if keys.start <= node_keys.start && node_keys.end <= keys.end { + node.apply(modification, node_keys); + return; + } + let (left_keys, right_keys) = split_range(node_keys); + node.force_children(initializer, &left_keys, &right_keys); + modify( + node.left.as_mut().unwrap(), + &left_keys, + initializer, + keys, + &modification, + ); + modify( + node.right.as_mut().unwrap(), + &right_keys, + initializer, + keys, + &modification, + ); + node.recalculate_from_children(&left_keys, &right_keys); +} + +impl< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, + > VecReadableVersion + for PersistentSegmentTreeVersion +{ + fn get(&self, keys: Range) -> Modification::Result { + get( + &mut self.root.clone(), // TODO: do not always force a branch + &self.all_keys, + self.initializer.as_ref(), + &keys, + ) + } +} + +pub struct PersistentSegmentTree< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, +>(PersistentSegmentTreeVersion); + +impl< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, + > VecReadableVersion + for PersistentSegmentTree +{ + fn get(&self, keys: Range) -> Modification::Result { + self.0.get(keys) + } +} + +impl< + Modification: RangeModification, + Initializer: LazyRangeInitializer, + Key: MidpointableKey, + > PersistentVecStorage + for PersistentSegmentTree +{ + fn new(all_keys: Range, initializer: Initializer) -> Self { + PersistentSegmentTree(PersistentSegmentTreeVersion { + root: Rc::new(Node::new(&all_keys, &initializer)), + all_keys: all_keys, + initializer: Rc::new(initializer), + }) + } + + type FrozenVersion = PersistentSegmentTreeVersion; + + fn modify(&mut self, keys: Range, modification: Modification) { + modify( + &mut self.0.root, // TODO: do not always force a branch + &self.0.all_keys, + self.0.initializer.as_ref(), + &keys, + &modification, + ) + } + + fn freeze(&mut self) -> Self::FrozenVersion { + self.0.clone() + } +} diff --git a/libs/persistent_range_query/tests/rsq_test.rs b/libs/persistent_range_query/tests/rsq_test.rs index 196c1c1f10..ee02ed7b1d 100644 --- a/libs/persistent_range_query/tests/rsq_test.rs +++ b/libs/persistent_range_query/tests/rsq_test.rs @@ -1,10 +1,11 @@ use persistent_range_query::naive::*; use persistent_range_query::ops::rsq::*; use persistent_range_query::ops::SameElementsInitializer; +use persistent_range_query::segment_tree::{MidpointableKey, PersistentSegmentTree}; use persistent_range_query::{PersistentVecStorage, VecReadableVersion}; use std::ops::Range; -#[derive(Copy, Clone, Debug)] +#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd)] struct K(u8); impl IndexableKey for K { @@ -23,10 +24,16 @@ impl SumOfSameElements for i32 { } } -#[test] -fn test_naive() { - let mut s: NaiveVecStorage, _, _> = - NaiveVecStorage::new(K(0)..K(12), SameElementsInitializer::new(0i32)); +impl MidpointableKey for K { + fn midpoint(range: &Range) -> Self { + K(range.start.0 + (range.end.0 - range.start.0) / 2) + } +} + +fn test_storage< + S: PersistentVecStorage, SameElementsInitializer, K>, +>() { + let mut s = S::new(K(0)..K(12), SameElementsInitializer::new(0i32)); assert_eq!(*s.get(K(0)..K(12)).sum(), 0); s.modify(K(2)..K(5), AddAssignModification::Add(3)); @@ -42,3 +49,13 @@ fn test_naive() { assert_eq!(*s.get(K(4)..K(6)).sum(), 12 + 12); assert_eq!(*s_old.get(K(4)..K(6)).sum(), 3); } + +#[test] +fn test_naive() { + test_storage::>(); +} + +#[test] +fn test_segment_tree() { + test_storage::>(); +}