Skip to main content

competitive/data_structure/
treap.rs

1use super::{
2    Allocator, BoxAllocator, LazyMapMonoid, MonoidAct, Xorshift,
3    binary_search_tree::{
4        BstDataAccess, BstDataMutRef, BstNode, BstNodeId, BstNodeIdManager, BstRoot, BstSeeker,
5        BstSpec, EqualSide,
6        data::{self, LazyMapElement, MonoidActElement},
7        node::WithParent,
8        seeker::{SeekByAccCond, SeekByKey, SeekByRaccCond},
9        split::{Split, Split3},
10    },
11};
12use std::{
13    borrow::Borrow,
14    cmp::Ordering,
15    fmt::{self, Debug},
16    marker::PhantomData,
17    mem::ManuallyDrop,
18    ops::{DerefMut, RangeBounds},
19};
20
21type TreapRoot<M, L> = BstRoot<TreapSpec<M, L>>;
22type TreapNode<M, L> = BstNode<TreapData<M, L>, WithParent<TreapData<M, L>>>;
23
24pub struct TreapSpec<M, L> {
25    _marker: PhantomData<(M, L)>,
26}
27
28pub struct TreapData<M, L>
29where
30    M: MonoidAct<Key: Ord>,
31    L: LazyMapMonoid,
32{
33    priority: u64,
34    key: MonoidActElement<M>,
35    value: LazyMapElement<L>,
36}
37
38impl<M, L> Debug for TreapData<M, L>
39where
40    M: MonoidAct<Key: Ord + Debug, Act: Debug>,
41    L: LazyMapMonoid<Key: Debug, Agg: Debug, Act: Debug>,
42{
43    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
44        f.debug_struct("TreapData")
45            .field("priority", &self.priority)
46            .field("key", &self.key)
47            .field("value", &self.value)
48            .finish()
49    }
50}
51
52impl<M, L> BstDataAccess<data::marker::Key> for TreapData<M, L>
53where
54    M: MonoidAct<Key: Ord>,
55    L: LazyMapMonoid,
56{
57    type Value = M::Key;
58
59    fn bst_data(&self) -> &Self::Value {
60        &self.key.key
61    }
62
63    fn bst_data_mut(&mut self) -> &mut Self::Value {
64        &mut self.key.key
65    }
66}
67
68impl<M, L> BstDataAccess<data::marker::MonoidAct> for TreapData<M, L>
69where
70    M: MonoidAct<Key: Ord>,
71    L: LazyMapMonoid,
72{
73    type Value = MonoidActElement<M>;
74
75    fn bst_data(&self) -> &Self::Value {
76        &self.key
77    }
78
79    fn bst_data_mut(&mut self) -> &mut Self::Value {
80        &mut self.key
81    }
82}
83
84impl<M, L> BstDataAccess<data::marker::LazyMap> for TreapData<M, L>
85where
86    M: MonoidAct<Key: Ord>,
87    L: LazyMapMonoid,
88{
89    type Value = LazyMapElement<L>;
90
91    fn bst_data(&self) -> &Self::Value {
92        &self.value
93    }
94
95    fn bst_data_mut(&mut self) -> &mut Self::Value {
96        &mut self.value
97    }
98}
99
100impl<M, L> BstSpec for TreapSpec<M, L>
101where
102    M: MonoidAct<Key: Ord>,
103    L: LazyMapMonoid,
104{
105    type Parent = WithParent<Self::Data>;
106    type Data = TreapData<M, L>;
107
108    fn top_down(mut node: BstDataMutRef<'_, Self>) {
109        MonoidActElement::<M>::top_down(node.reborrow_datamut());
110        LazyMapElement::<L>::top_down(node.reborrow_datamut());
111    }
112
113    fn bottom_up(mut node: BstDataMutRef<'_, Self>) {
114        LazyMapElement::<L>::bottom_up(node.reborrow_datamut());
115    }
116
117    fn merge(
118        left: Option<TreapRoot<M, L>>,
119        right: Option<TreapRoot<M, L>>,
120    ) -> Option<TreapRoot<M, L>> {
121        match (left, right) {
122            (None, None) => None,
123            (None, Some(node)) | (Some(node), None) => Some(node),
124            (Some(mut left), Some(mut right)) => unsafe {
125                if left.reborrow().into_data().priority > right.reborrow().into_data().priority {
126                    TreapSpec::top_down(left.borrow_datamut());
127                    let lr = left.borrow_mut().right().take();
128                    let lr = Self::merge(lr, Some(right)).unwrap_unchecked();
129                    left.borrow_mut().right().set(lr);
130                    TreapSpec::bottom_up(left.borrow_datamut());
131                    Some(left)
132                } else {
133                    TreapSpec::top_down(right.borrow_datamut());
134                    let rl = right.borrow_mut().left().take();
135                    let rl = Self::merge(Some(left), rl).unwrap_unchecked();
136                    right.borrow_mut().left().set(rl);
137                    TreapSpec::bottom_up(right.borrow_datamut());
138                    Some(right)
139                }
140            },
141        }
142    }
143
144    fn split<Seeker>(
145        node: Option<TreapRoot<M, L>>,
146        mut seeker: Seeker,
147        equal_side: EqualSide,
148    ) -> (Option<TreapRoot<M, L>>, Option<TreapRoot<M, L>>)
149    where
150        Seeker: BstSeeker<Spec = Self>,
151    {
152        match node {
153            None => (None, None),
154            Some(mut node) => {
155                Self::top_down(node.borrow_datamut());
156                if equal_side.goes_left(seeker.bst_seek(node.reborrow())) {
157                    unsafe {
158                        let right = node.borrow_mut().right().take();
159                        let (l, r) = Self::split(right, seeker, equal_side);
160                        if let Some(l) = l {
161                            node.borrow_mut().right().set(l);
162                        }
163                        Self::bottom_up(node.borrow_datamut());
164                        (Some(node), r)
165                    }
166                } else {
167                    unsafe {
168                        let left = node.borrow_mut().left().take();
169                        let (l, r) = Self::split(left, seeker, equal_side);
170                        if let Some(r) = r {
171                            node.borrow_mut().left().set(r);
172                        }
173                        Self::bottom_up(node.borrow_datamut());
174                        (l, Some(node))
175                    }
176                }
177            }
178        }
179    }
180}
181
182impl<M, L> TreapSpec<M, L>
183where
184    M: MonoidAct<Key: Ord>,
185    L: LazyMapMonoid,
186{
187    pub fn merge_ordered(
188        left: Option<TreapRoot<M, L>>,
189        right: Option<TreapRoot<M, L>>,
190    ) -> Option<TreapRoot<M, L>> {
191        match (left, right) {
192            (None, None) => None,
193            (None, Some(node)) | (Some(node), None) => Some(node),
194            (Some(mut left), Some(mut right)) => unsafe {
195                if left.reborrow().into_data().priority > right.reborrow().into_data().priority {
196                    Self::top_down(left.borrow_datamut());
197                    let key = &left.reborrow().into_data().key.key;
198                    let (rl, rr) = Self::split(Some(right), SeekByKey::new(key), EqualSide::Right);
199                    let ll = left.borrow_mut().left().take();
200                    let lr = left.borrow_mut().right().take();
201                    if let Some(l) = Self::merge_ordered(ll, rl) {
202                        left.borrow_mut().left().set(l);
203                    }
204                    if let Some(r) = Self::merge_ordered(lr, rr) {
205                        left.borrow_mut().right().set(r);
206                    }
207                    Self::bottom_up(left.borrow_datamut());
208                    Some(left)
209                } else {
210                    Self::top_down(right.borrow_datamut());
211                    let key = &right.reborrow().into_data().key.key;
212                    let (ll, lr) = Self::split(Some(left), SeekByKey::new(key), EqualSide::Right);
213                    let rl = right.borrow_mut().left().take();
214                    let rr = right.borrow_mut().right().take();
215                    if let Some(l) = Self::merge_ordered(ll, rl) {
216                        right.borrow_mut().left().set(l);
217                    }
218                    if let Some(r) = Self::merge_ordered(lr, rr) {
219                        right.borrow_mut().right().set(r);
220                    }
221                    Self::bottom_up(right.borrow_datamut());
222                    Some(right)
223                }
224            },
225        }
226    }
227}
228
229pub struct Treap<M, L, A = BoxAllocator<TreapNode<M, L>>>
230where
231    M: MonoidAct<Key: Ord>,
232    L: LazyMapMonoid,
233    A: Allocator<TreapNode<M, L>>,
234{
235    root: Option<TreapRoot<M, L>>,
236    node_id_manager: BstNodeIdManager<TreapSpec<M, L>>,
237    rng: Xorshift,
238    allocator: ManuallyDrop<A>,
239    _marker: PhantomData<(M, L)>,
240}
241
242impl<M, L, A> Default for Treap<M, L, A>
243where
244    M: MonoidAct<Key: Ord>,
245    L: LazyMapMonoid,
246    A: Allocator<TreapNode<M, L>> + Default,
247{
248    fn default() -> Self {
249        Self {
250            root: None,
251            node_id_manager: Default::default(),
252            rng: Xorshift::new(),
253            allocator: ManuallyDrop::new(A::default()),
254            _marker: PhantomData,
255        }
256    }
257}
258
259impl<M, L, A> Drop for Treap<M, L, A>
260where
261    M: MonoidAct<Key: Ord>,
262    L: LazyMapMonoid,
263    A: Allocator<TreapNode<M, L>>,
264{
265    fn drop(&mut self) {
266        unsafe {
267            if let Some(root) = self.root.take() {
268                root.into_dying().drop_all(self.allocator.deref_mut());
269            }
270            ManuallyDrop::drop(&mut self.allocator);
271        }
272    }
273}
274
275impl<M, L> Treap<M, L>
276where
277    M: MonoidAct<Key: Ord>,
278    L: LazyMapMonoid,
279{
280    pub fn new() -> Self {
281        Self::default()
282    }
283}
284
285impl<M, L, A> Treap<M, L, A>
286where
287    M: MonoidAct<Key: Ord>,
288    L: LazyMapMonoid,
289    A: Allocator<TreapNode<M, L>>,
290{
291    pub fn len(&self) -> usize {
292        self.node_id_manager.len()
293    }
294
295    pub fn is_empty(&self) -> bool {
296        self.node_id_manager.is_empty()
297    }
298
299    pub fn clear(&mut self) {
300        unsafe {
301            if let Some(root) = self.root.take() {
302                root.into_dying().drop_all(self.allocator.deref_mut());
303            }
304            self.node_id_manager.clear();
305        }
306    }
307
308    pub fn get(&mut self, node_id: BstNodeId<TreapSpec<M, L>>) -> Option<(&M::Key, &L::Key)> {
309        if !self.node_id_manager.contains(&node_id) {
310            return None;
311        }
312        unsafe {
313            WithParent::resolve_top_down::<TreapSpec<M, L>>(
314                node_id.reborrow_datamut(&mut self.root),
315            );
316            let data = node_id.reborrow(&self.root).into_data();
317            Some((&data.key.key, &data.value.key))
318        }
319    }
320
321    pub fn change(
322        &mut self,
323        node_id: BstNodeId<TreapSpec<M, L>>,
324        f: impl FnOnce(&mut L::Key),
325    ) -> bool {
326        if !self.node_id_manager.contains(&node_id) {
327            return false;
328        }
329        unsafe {
330            WithParent::resolve_top_down::<TreapSpec<M, L>>(
331                node_id.reborrow_datamut(&mut self.root),
332            );
333            let data = node_id.reborrow_datamut(&mut self.root).into_data_mut();
334            f(&mut data.value.key);
335            WithParent::resolve_bottom_up::<TreapSpec<M, L>>(
336                node_id.reborrow_datamut(&mut self.root),
337            );
338        }
339        true
340    }
341
342    pub fn change_key_value(
343        &mut self,
344        node_id: BstNodeId<TreapSpec<M, L>>,
345        f: impl FnOnce(&mut M::Key, &mut L::Key),
346    ) -> bool {
347        if !self.node_id_manager.contains(&node_id) {
348            return false;
349        }
350        unsafe {
351            WithParent::resolve_top_down::<TreapSpec<M, L>>(
352                node_id.reborrow_datamut(&mut self.root),
353            );
354            let mut node = if WithParent::is_root(node_id.reborrow(&self.root)) {
355                WithParent::remove_root(&mut self.root).unwrap_unchecked()
356            } else {
357                WithParent::remove_not_root(node_id.reborrow_mut(&mut self.root))
358            };
359            let data = node.borrow_datamut().into_data_mut();
360            f(&mut data.key.key, &mut data.value.key);
361            self.root = TreapSpec::merge_ordered(self.root.take(), Some(node));
362            true
363        }
364    }
365
366    pub fn insert(&mut self, key: M::Key, value: L::Key) -> BstNodeId<TreapSpec<M, L>> {
367        let (left, right) =
368            TreapSpec::split(self.root.take(), SeekByKey::new(&key), EqualSide::Right);
369        let data = TreapData {
370            priority: self.rng.rand64(),
371            key: MonoidActElement::from_key(key),
372            value: LazyMapElement::from_key(value),
373        };
374        let node = BstRoot::from_data(data, self.allocator.deref_mut());
375        let node_id = self.node_id_manager.register(&node);
376        self.root = TreapSpec::merge(TreapSpec::merge(left, Some(node)), right);
377        node_id
378    }
379
380    pub fn remove(&mut self, node_id: BstNodeId<TreapSpec<M, L>>) -> Option<(M::Key, L::Key)> {
381        if !self.node_id_manager.contains(&node_id) {
382            return None;
383        }
384        unsafe {
385            WithParent::resolve_top_down::<TreapSpec<M, L>>(
386                node_id.reborrow_datamut(&mut self.root),
387            );
388            let node = if WithParent::is_root(node_id.reborrow(&self.root)) {
389                WithParent::remove_root(&mut self.root).unwrap_unchecked()
390            } else {
391                WithParent::remove_not_root(node_id.reborrow_mut(&mut self.root))
392            };
393            self.node_id_manager.unregister(node_id);
394            let data = node.into_dying().into_data(self.allocator.deref_mut());
395            Some((data.key.key, data.value.key))
396        }
397    }
398
399    pub fn range_by_key<Q, R>(&mut self, range: R) -> TreapSplit3<'_, M, L>
400    where
401        M: MonoidAct<Key: Borrow<Q>>,
402        Q: Ord + ?Sized,
403        R: RangeBounds<Q>,
404    {
405        let split3 = Split3::seek_by_key(&mut self.root, range);
406        TreapSplit3 {
407            split3,
408            key_updated: false,
409        }
410    }
411
412    pub fn find_by_key<Q>(&mut self, key: &Q) -> Option<BstNodeId<TreapSpec<M, L>>>
413    where
414        M: MonoidAct<Key: Borrow<Q>>,
415        Q: Ord + ?Sized,
416    {
417        let split = Split::new(
418            &mut self.root,
419            SeekByKey::<TreapSpec<M, L>, M::Key, Q>::new(key),
420            EqualSide::Right,
421        );
422        let node = split.right()?.leftmost();
423        matches!(node.into_data().key.key.borrow().cmp(key), Ordering::Equal)
424            .then(|| self.node_id_manager.registered_node_id(node))
425            .flatten()
426    }
427
428    pub fn find_by_acc_cond<F>(&mut self, f: F) -> Option<BstNodeId<TreapSpec<M, L>>>
429    where
430        F: FnMut(&L::Agg) -> bool,
431    {
432        let split = Split::new(
433            &mut self.root,
434            SeekByAccCond::<TreapSpec<M, L>, L, F>::new(f),
435            EqualSide::Right,
436        );
437        let node = split.right()?.leftmost();
438        self.node_id_manager.registered_node_id(node)
439    }
440
441    pub fn find_by_racc_cond<F>(&mut self, f: F) -> Option<BstNodeId<TreapSpec<M, L>>>
442    where
443        F: FnMut(&L::Agg) -> bool,
444    {
445        let split = Split::new(
446            &mut self.root,
447            SeekByRaccCond::<TreapSpec<M, L>, L, F>::new(f),
448            EqualSide::Left,
449        );
450        let node = split.left()?.rightmost();
451        self.node_id_manager.registered_node_id(node)
452    }
453}
454
455pub struct TreapSplit3<'a, M, L>
456where
457    M: MonoidAct<Key: Ord>,
458    L: LazyMapMonoid,
459{
460    split3: Split3<'a, TreapSpec<M, L>>,
461    key_updated: bool,
462}
463
464impl<'a, M, L> TreapSplit3<'a, M, L>
465where
466    M: MonoidAct<Key: Ord>,
467    L: LazyMapMonoid,
468{
469    pub fn fold(&self) -> L::Agg {
470        if let Some(node) = self.split3.mid() {
471            node.reborrow().into_data().value.agg.clone()
472        } else {
473            L::agg_unit()
474        }
475    }
476
477    pub fn update_key(&mut self, act: M::Act) {
478        if let Some(node) = self.split3.mid_datamut() {
479            MonoidActElement::<M>::update_act(node, &act);
480            self.key_updated = true;
481        }
482    }
483
484    pub fn update_value(&mut self, act: L::Act) {
485        if let Some(node) = self.split3.mid_datamut() {
486            LazyMapElement::<L>::update_act(node, &act);
487        }
488    }
489}
490
491impl<'a, M, L> Drop for TreapSplit3<'a, M, L>
492where
493    M: MonoidAct<Key: Ord>,
494    L: LazyMapMonoid,
495{
496    fn drop(&mut self) {
497        if self.key_updated {
498            self.split3.manually_merge(TreapSpec::merge_ordered);
499        }
500    }
501}
502
503#[cfg(test)]
504mod tests {
505    use super::*;
506    use crate::algebra::{
507        AdditiveOperation, EmptyAct, FlattenAct, RangeMaxRangeAdd, RangeSumRangeAdd,
508    };
509
510    #[test]
511    fn test_treap() {
512        const A: i64 = 100;
513        let mut rng = Xorshift::default();
514        let mut treap = Treap::<FlattenAct<AdditiveOperation<i64>>, RangeMaxRangeAdd<i64>>::new();
515        let mut node_ids = vec![];
516        let mut data = vec![];
517        for _ in 0..10000 {
518            let (l, r) = loop {
519                let l = rng.random(-A..=A);
520                let r = rng.random(-A..=A);
521                if l <= r {
522                    break (l, r);
523                }
524            };
525            assert_eq!(data.len(), treap.len());
526            assert_eq!(data.is_empty(), treap.is_empty());
527            match rng.random(0..8) {
528                0 => {
529                    let key = rng.random(-A..=A);
530                    let value = rng.random(-A..=A);
531                    let k = data.partition_point(|(k, _)| *k < key);
532                    data.insert(k, (key, value));
533                    node_ids.insert(k, treap.insert(key, value));
534                }
535                1 => {
536                    if !data.is_empty() {
537                        let k = rng.random(0..data.len());
538                        let expected = data.remove(k);
539                        let result = treap.remove(node_ids.remove(k)).unwrap();
540                        assert_eq!(expected, result);
541                    }
542                }
543                2 => {
544                    let expected: i64 = data
545                        .iter()
546                        .filter(|(k, _)| (l..r).contains(k))
547                        .map(|(_, v)| *v)
548                        .max()
549                        .unwrap_or(i64::MIN);
550                    let result = treap.range_by_key(l..r).fold();
551                    assert_eq!(expected, result);
552                }
553                3 => {
554                    let add = rng.random(-A..=A);
555                    for (k, v) in data.iter_mut() {
556                        if (l..r).contains(k) {
557                            *v += add;
558                        }
559                    }
560                    treap.range_by_key(l..r).update_value(add);
561                }
562                4 => {
563                    let add = rng.random(-A..=A);
564                    for (k, _) in data.iter_mut() {
565                        if (l..r).contains(k) {
566                            *k += add;
567                        }
568                    }
569                    treap.range_by_key(l..r).update_key(add);
570                }
571                5 => {
572                    if !data.is_empty() {
573                        let k = rng.random(0..data.len());
574                        let expected = data[k];
575                        let result = treap.get(node_ids[k]).unwrap();
576                        assert_eq!(expected, (*result.0, *result.1));
577                    }
578                }
579                6 => {
580                    if !data.is_empty() {
581                        let k = rng.random(0..data.len());
582                        let x = rng.random(-A..=A);
583                        data[k].1 = x;
584                        treap.change(node_ids[k], |value| *value = x);
585                    }
586                }
587                _ => {
588                    if !data.is_empty() {
589                        let k = rng.random(0..data.len());
590                        let nk = rng.random(-A..=A);
591                        let nv = rng.random(-A..=A);
592                        data[k].0 = nk;
593                        data[k].1 = nv;
594                        treap.change_key_value(node_ids[k], |key, value| {
595                            *key = nk;
596                            *value = nv;
597                        });
598                    }
599                }
600            }
601        }
602
603        let mut treap = Treap::<EmptyAct<i64>, RangeSumRangeAdd<i64>>::new();
604        let mut node_ids = vec![];
605        let mut data = vec![];
606        for _ in 0..10000 {
607            let (l, r) = loop {
608                let l = rng.random(-A..=A);
609                let r = rng.random(-A..=A);
610                if l <= r {
611                    break (l, r);
612                }
613            };
614            assert_eq!(data.len(), treap.len());
615            assert_eq!(data.is_empty(), treap.is_empty());
616            match rng.random(0..10) {
617                0 => {
618                    let key = rng.random(-A..=A);
619                    let value = rng.random(1..=A);
620                    let k = data.partition_point(|(k, _)| *k < key);
621                    data.insert(k, (key, value));
622                    node_ids.insert(k, treap.insert(key, value));
623                }
624                1 => {
625                    if !data.is_empty() {
626                        let k = rng.random(0..data.len());
627                        let expected = data.remove(k);
628                        let result = treap.remove(node_ids.remove(k)).unwrap();
629                        assert_eq!(expected, result);
630                    }
631                }
632                2 => {
633                    let expected: i64 = data
634                        .iter()
635                        .filter(|(k, _)| (l..r).contains(k))
636                        .map(|(_, v)| *v)
637                        .sum();
638                    let result = treap.range_by_key(l..r).fold().0;
639                    assert_eq!(expected, result);
640                }
641                3 => {
642                    let add = rng.random(1..=A);
643                    for (k, v) in data.iter_mut() {
644                        if (l..r).contains(k) {
645                            *v += add;
646                        }
647                    }
648                    treap.range_by_key(l..r).update_value(add);
649                }
650                5 => {
651                    if !data.is_empty() {
652                        let k = rng.random(0..data.len());
653                        let expected = data[k];
654                        let result = treap.get(node_ids[k]).unwrap();
655                        assert_eq!(expected, (*result.0, *result.1));
656                    }
657                }
658                6 => {
659                    if !data.is_empty() {
660                        let k = rng.random(0..data.len());
661                        let x = rng.random(1..=A);
662                        data[k].1 = x;
663                        treap.change(node_ids[k], |value| *value = x);
664                    }
665                }
666                7 => {
667                    let key = rng.random(-A..=A);
668                    let expected = data.iter().find(|(k, _)| *k == key).cloned();
669                    let result = treap.find_by_key(&key).map(|id| treap.get(id).unwrap());
670                    assert_eq!(expected, result.map(|(k, v)| (*k, *v)));
671                }
672                8 => {
673                    let s = rng.random(0..=A);
674                    let mut acc = 0;
675                    let expected = data.iter().find_map(|(k, v)| {
676                        acc += *v;
677                        if acc >= s { Some((*k, *v)) } else { None }
678                    });
679                    let result = treap
680                        .find_by_acc_cond(|agg| agg.0 >= s)
681                        .map(|id| treap.get(id).unwrap());
682                    assert_eq!(expected, result.map(|(k, v)| (*k, *v)));
683                }
684                _ => {
685                    let s = rng.random(0..=A);
686                    let mut acc = 0;
687                    let expected = data.iter().rev().find_map(|(k, v)| {
688                        acc += *v;
689                        if acc >= s { Some((*k, *v)) } else { None }
690                    });
691                    let result = treap
692                        .find_by_racc_cond(|agg| agg.0 >= s)
693                        .map(|id| treap.get(id).unwrap());
694                    assert_eq!(expected, result.map(|(k, v)| (*k, *v)));
695                }
696            }
697        }
698    }
699}