Skip to main content

competitive/data_structure/
splay_tree.rs

1use super::{
2    Allocator, MemoryPool,
3    binary_search_tree::{
4        BstDataAccess, BstDataMutRef, BstNode, BstRoot, BstSeeker, BstSpec, EqualSide, data,
5        node::WithNoParent,
6        seeker::{SeekByKey, SeekBySize},
7        split::Split3,
8    },
9    splay_operations,
10};
11use std::{
12    borrow::Borrow,
13    cmp::Ordering,
14    fmt::{self, Debug},
15    iter::FusedIterator,
16    marker::PhantomData,
17    mem::{ManuallyDrop, replace},
18    ops::{DerefMut, RangeBounds},
19    ptr::NonNull,
20};
21
22type SplayTreeRoot<K, V> = BstRoot<SplayTreeSpec<K, V>>;
23type SplayTreeNode<K, V> = BstNode<SplayTreeData<K, V>>;
24
25pub struct SplayTreeSpec<K, V> {
26    _marker: PhantomData<fn() -> (K, V)>,
27}
28
29pub struct SplayTreeData<K, V> {
30    key: K,
31    value: V,
32    size: usize,
33}
34
35impl<K, V> Debug for SplayTreeData<K, V>
36where
37    K: Debug,
38    V: Debug,
39{
40    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
41        f.debug_struct("SplayTreeData")
42            .field("key", &self.key)
43            .field("value", &self.value)
44            .field("size", &self.size)
45            .finish()
46    }
47}
48
49impl<K, V> BstDataAccess<data::marker::Key> for SplayTreeData<K, V> {
50    type Value = K;
51
52    fn bst_data(&self) -> &Self::Value {
53        &self.key
54    }
55
56    fn bst_data_mut(&mut self) -> &mut Self::Value {
57        &mut self.key
58    }
59}
60
61impl<K, V> BstDataAccess<data::marker::Size> for SplayTreeData<K, V> {
62    type Value = usize;
63
64    fn bst_data(&self) -> &Self::Value {
65        &self.size
66    }
67
68    fn bst_data_mut(&mut self) -> &mut Self::Value {
69        &mut self.size
70    }
71}
72
73impl<K, V> BstSpec for SplayTreeSpec<K, V> {
74    type Parent = WithNoParent<Self::Data>;
75    type Data = SplayTreeData<K, V>;
76
77    #[inline]
78    fn bottom_up(mut node: BstDataMutRef<'_, Self>) {
79        let left = node
80            .reborrow()
81            .left()
82            .descend()
83            .map(|node| node.into_data().size)
84            .unwrap_or_default();
85        let right = node
86            .reborrow()
87            .right()
88            .descend()
89            .map(|node| node.into_data().size)
90            .unwrap_or_default();
91        node.data_mut().size = left + right + 1;
92    }
93
94    #[inline]
95    fn merge(
96        left: Option<SplayTreeRoot<K, V>>,
97        right: Option<SplayTreeRoot<K, V>>,
98    ) -> Option<SplayTreeRoot<K, V>> {
99        splay_operations::merge(left, right)
100    }
101
102    #[inline]
103    fn split<Seeker>(
104        node: Option<SplayTreeRoot<K, V>>,
105        seeker: Seeker,
106        equal_side: EqualSide,
107    ) -> (Option<SplayTreeRoot<K, V>>, Option<SplayTreeRoot<K, V>>)
108    where
109        Seeker: BstSeeker<Spec = Self>,
110    {
111        splay_operations::split(node, seeker, equal_side)
112    }
113}
114
115pub struct SplayTree<K, V, A = MemoryPool<SplayTreeNode<K, V>>>
116where
117    A: Allocator<SplayTreeNode<K, V>>,
118{
119    root: Option<SplayTreeRoot<K, V>>,
120    length: usize,
121    allocator: ManuallyDrop<A>,
122}
123
124impl<K, V, A> Debug for SplayTree<K, V, A>
125where
126    K: Debug,
127    V: Debug,
128    A: Allocator<SplayTreeNode<K, V>>,
129{
130    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
131        f.debug_struct("SplayTree")
132            .field("length", &self.length)
133            .finish_non_exhaustive()
134    }
135}
136
137impl<K, V, A> Default for SplayTree<K, V, A>
138where
139    A: Allocator<SplayTreeNode<K, V>> + Default,
140{
141    fn default() -> Self {
142        Self {
143            root: None,
144            length: 0,
145            allocator: ManuallyDrop::new(A::default()),
146        }
147    }
148}
149
150impl<K, V, A> Drop for SplayTree<K, V, A>
151where
152    A: Allocator<SplayTreeNode<K, V>>,
153{
154    fn drop(&mut self) {
155        unsafe {
156            if let Some(root) = self.root.take() {
157                root.into_dying().drop_all(self.allocator.deref_mut());
158            }
159            ManuallyDrop::drop(&mut self.allocator);
160        }
161    }
162}
163
164impl<K, V> SplayTree<K, V> {
165    pub fn new() -> Self {
166        Self::default()
167    }
168
169    pub fn with_capacity(capacity: usize) -> Self {
170        Self {
171            root: None,
172            length: 0,
173            allocator: ManuallyDrop::new(MemoryPool::with_capacity(capacity)),
174        }
175    }
176}
177
178impl<K, V, A> SplayTree<K, V, A>
179where
180    A: Allocator<SplayTreeNode<K, V>>,
181{
182    #[inline]
183    fn splay<Seeker>(&mut self, seeker: Seeker) -> Option<Ordering>
184    where
185        Seeker: BstSeeker<Spec = SplayTreeSpec<K, V>>,
186    {
187        let (ordering, root) = splay_operations::splay(self.root.take()?, seeker);
188        self.root = Some(root);
189        Some(ordering)
190    }
191
192    fn splay_by_key<Q>(&mut self, key: &Q) -> Option<Ordering>
193    where
194        K: Borrow<Q>,
195        Q: Ord + ?Sized,
196    {
197        self.splay(SeekByKey::new(key))
198    }
199
200    fn splay_by_size(&mut self, index: usize) -> Option<Ordering> {
201        self.splay(SeekBySize::new(index))
202    }
203
204    pub fn get<Q>(&mut self, key: &Q) -> Option<&V>
205    where
206        K: Borrow<Q>,
207        Q: Ord + ?Sized,
208    {
209        self.get_key_value(key).map(|(_, value)| value)
210    }
211
212    pub fn get_key_value<Q>(&mut self, key: &Q) -> Option<(&K, &V)>
213    where
214        K: Borrow<Q>,
215        Q: Ord + ?Sized,
216    {
217        matches!(self.splay_by_key(key)?, Ordering::Equal).then(|| {
218            let data = self.root.as_ref().unwrap().reborrow().into_data();
219            (&data.key, &data.value)
220        })
221    }
222
223    pub fn get_key_value_at(&mut self, index: usize) -> Option<(&K, &V)> {
224        if index >= self.length {
225            return None;
226        }
227        self.splay_by_size(index);
228        let data = self.root.as_ref()?.reborrow().into_data();
229        Some((&data.key, &data.value))
230    }
231
232    pub fn insert(&mut self, key: K, value: V) -> Option<V>
233    where
234        K: Ord,
235    {
236        let ordering = self.splay_by_key(&key);
237        if matches!(ordering, Some(Ordering::Equal)) {
238            return Some(replace(
239                &mut self
240                    .root
241                    .as_mut()
242                    .unwrap()
243                    .borrow_datamut()
244                    .data_mut()
245                    .value,
246                value,
247            ));
248        }
249        let mut node = BstRoot::from_data(
250            SplayTreeData {
251                key,
252                value,
253                size: 1,
254            },
255            self.allocator.deref_mut(),
256        );
257        if let Some(mut root) = self.root.take() {
258            match ordering.unwrap() {
259                Ordering::Greater => {
260                    let left = unsafe { root.borrow_mut().left_mut().take() };
261                    if let Some(left) = left {
262                        unsafe { node.borrow_mut().left_mut().set(left) };
263                    }
264                    SplayTreeSpec::bottom_up(root.borrow_datamut());
265                    unsafe { node.borrow_mut().right_mut().set(root) };
266                }
267                Ordering::Less => {
268                    let right = unsafe { root.borrow_mut().right_mut().take() };
269                    if let Some(right) = right {
270                        unsafe { node.borrow_mut().right_mut().set(right) };
271                    }
272                    SplayTreeSpec::bottom_up(root.borrow_datamut());
273                    unsafe { node.borrow_mut().left_mut().set(root) };
274                }
275                Ordering::Equal => unreachable!(),
276            }
277            SplayTreeSpec::bottom_up(node.borrow_datamut());
278        }
279        self.root = Some(node);
280        self.length += 1;
281        None
282    }
283
284    pub fn remove<Q>(&mut self, key: &Q) -> Option<V>
285    where
286        K: Borrow<Q>,
287        Q: Ord + ?Sized,
288    {
289        if !matches!(self.splay_by_key(key)?, Ordering::Equal) {
290            return None;
291        }
292        Some(self.remove_root().1)
293    }
294
295    pub fn remove_at(&mut self, index: usize) -> Option<(K, V)> {
296        if index >= self.length {
297            return None;
298        }
299        self.splay_by_size(index);
300        Some(self.remove_root())
301    }
302
303    fn remove_root(&mut self) -> (K, V) {
304        let mut node = self.root.take().unwrap();
305        let left = unsafe { node.borrow_mut().left_mut().take() };
306        let right = unsafe { node.borrow_mut().right_mut().take() };
307        self.root = SplayTreeSpec::merge(left, right);
308        self.length -= 1;
309        let data = unsafe { node.into_dying().into_data(self.allocator.deref_mut()) };
310        (data.key, data.value)
311    }
312
313    pub fn len(&self) -> usize {
314        self.length
315    }
316
317    pub fn is_empty(&self) -> bool {
318        self.length == 0
319    }
320
321    pub fn iter(&mut self) -> Iter<'_, K, V> {
322        Iter::new(Split3::seek_by_size(&mut self.root, ..))
323    }
324
325    pub fn range<Q, R>(&mut self, range: R) -> Iter<'_, K, V>
326    where
327        K: Borrow<Q>,
328        Q: Ord + ?Sized,
329        R: RangeBounds<Q>,
330    {
331        Iter::new(Split3::seek_by_key(&mut self.root, range))
332    }
333
334    pub fn range_at<R>(&mut self, range: R) -> Iter<'_, K, V>
335    where
336        R: RangeBounds<usize>,
337    {
338        Iter::new(Split3::seek_by_size(&mut self.root, range))
339    }
340}
341
342pub struct Iter<'a, K, V> {
343    split: Split3<'a, SplayTreeSpec<K, V>>,
344    front: Vec<NonNull<SplayTreeNode<K, V>>>,
345    back: Vec<NonNull<SplayTreeNode<K, V>>>,
346    remaining: usize,
347}
348
349impl<K, V> Debug for Iter<'_, K, V> {
350    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
351        f.debug_struct("Iter")
352            .field("remaining", &self.remaining)
353            .finish_non_exhaustive()
354    }
355}
356
357impl<'a, K, V> Iter<'a, K, V> {
358    fn new(split: Split3<'a, SplayTreeSpec<K, V>>) -> Self {
359        let remaining = split
360            .mid()
361            .map(|node| node.into_data().size)
362            .unwrap_or_default();
363        let mut iter = Self {
364            split,
365            front: vec![],
366            back: vec![],
367            remaining,
368        };
369        if let Some(root) = iter.split.mid() {
370            Self::push_left(root.node, &mut iter.front);
371            Self::push_right(root.node, &mut iter.back);
372        }
373        iter
374    }
375
376    fn push_left(
377        mut node: NonNull<SplayTreeNode<K, V>>,
378        stack: &mut Vec<NonNull<SplayTreeNode<K, V>>>,
379    ) {
380        loop {
381            stack.push(node);
382            let Some(left) = (unsafe { node.as_ref().child[0] }) else {
383                break;
384            };
385            node = left;
386        }
387    }
388
389    fn push_right(
390        mut node: NonNull<SplayTreeNode<K, V>>,
391        stack: &mut Vec<NonNull<SplayTreeNode<K, V>>>,
392    ) {
393        loop {
394            stack.push(node);
395            let Some(right) = (unsafe { node.as_ref().child[1] }) else {
396                break;
397            };
398            node = right;
399        }
400    }
401}
402
403impl<K, V> Iterator for Iter<'_, K, V>
404where
405    K: Clone,
406    V: Clone,
407{
408    type Item = (K, V);
409
410    fn next(&mut self) -> Option<Self::Item> {
411        if self.remaining == 0 {
412            return None;
413        }
414        let node = self.front.pop().unwrap();
415        if let Some(right) = unsafe { node.as_ref().child[1] } {
416            Self::push_left(right, &mut self.front);
417        }
418        self.remaining -= 1;
419        let data = unsafe { &node.as_ref().data };
420        Some((data.key.clone(), data.value.clone()))
421    }
422
423    fn last(mut self) -> Option<Self::Item> {
424        self.next_back()
425    }
426
427    fn min(mut self) -> Option<Self::Item> {
428        self.next()
429    }
430
431    fn max(mut self) -> Option<Self::Item> {
432        self.next_back()
433    }
434
435    fn size_hint(&self) -> (usize, Option<usize>) {
436        (self.remaining, Some(self.remaining))
437    }
438}
439
440impl<K, V> DoubleEndedIterator for Iter<'_, K, V>
441where
442    K: Clone,
443    V: Clone,
444{
445    fn next_back(&mut self) -> Option<Self::Item> {
446        if self.remaining == 0 {
447            return None;
448        }
449        let node = self.back.pop().unwrap();
450        if let Some(left) = unsafe { node.as_ref().child[0] } {
451            Self::push_right(left, &mut self.back);
452        }
453        self.remaining -= 1;
454        let data = unsafe { &node.as_ref().data };
455        Some((data.key.clone(), data.value.clone()))
456    }
457}
458
459impl<K, V> ExactSizeIterator for Iter<'_, K, V>
460where
461    K: Clone,
462    V: Clone,
463{
464}
465
466impl<K, V> FusedIterator for Iter<'_, K, V>
467where
468    K: Clone,
469    V: Clone,
470{
471}
472
473#[cfg(test)]
474mod tests {
475    use super::*;
476    use crate::tools::Xorshift;
477    use std::collections::BTreeSet;
478    use std::{
479        cell::RefCell,
480        collections::{BTreeMap, VecDeque},
481        ops::Bound,
482    };
483
484    #[test]
485    fn test_splay_tree() {
486        const Q: usize = 30_000;
487        const A: u64 = 500;
488        let mut tree = SplayTree::new();
489        let mut map = BTreeMap::new();
490        let mut rng = Xorshift::default();
491        for key in 0..A {
492            map.insert(key, key as usize);
493            tree.insert(key, key as usize);
494        }
495        for value in 0..Q {
496            let key = rng.rand(A);
497            match rng.rand(5) {
498                0 => assert_eq!(map.remove(&key), tree.remove(&key)),
499                1 => assert_eq!(map.insert(key, value), tree.insert(key, value)),
500                2 => assert_eq!(map.get_key_value(&key), tree.get_key_value(&key)),
501                3 => {
502                    let index = rng.rand((map.len() + 1) as u64) as usize;
503                    assert_eq!(map.iter().nth(index), tree.get_key_value_at(index));
504                }
505                _ => {
506                    let index = rng.rand((map.len() + 1) as u64) as usize;
507                    let key = map.iter().nth(index).map(|(&key, _)| key);
508                    assert_eq!(
509                        key.and_then(|key| map.remove_entry(&key)),
510                        tree.remove_at(index)
511                    );
512                }
513            }
514            assert_eq!(map.len(), tree.len());
515            assert_eq!(map.is_empty(), tree.is_empty());
516            let expected = map
517                .iter()
518                .map(|(&key, &value)| (key, value))
519                .collect::<Vec<_>>();
520            assert_eq!(tree.iter().collect::<Vec<_>>(), expected);
521
522            let key_range = {
523                let left = rng.rand(A + 1);
524                let right = rng.rand(A + 1);
525                let (left, right) = (left.min(right), left.max(right));
526                let start = match rng.rand(3) {
527                    0 => Bound::Included(left),
528                    1 => Bound::Excluded(left),
529                    _ => Bound::Unbounded,
530                };
531                let end = match rng.rand(3) {
532                    0 => Bound::Included(right),
533                    1 if start == Bound::Excluded(right) => Bound::Included(right),
534                    1 => Bound::Excluded(right),
535                    _ => Bound::Unbounded,
536                };
537                (start, end)
538            };
539            assert_eq!(
540                tree.range(key_range).collect::<Vec<_>>(),
541                map.range(key_range)
542                    .map(|(&key, &value)| (key, value))
543                    .collect::<Vec<_>>()
544            );
545
546            let index_range = {
547                let left = rng.rand((map.len() + 1) as u64) as usize;
548                let right = rng.rand((map.len() + 1) as u64) as usize;
549                let (left, right) = (left.min(right), left.max(right));
550                let start = match rng.rand(3) {
551                    0 => Bound::Included(left),
552                    1 => Bound::Excluded(left),
553                    _ => Bound::Unbounded,
554                };
555                let end = match rng.rand(3) {
556                    0 => Bound::Included(right),
557                    1 if start == Bound::Excluded(right) => Bound::Included(right),
558                    1 => Bound::Excluded(right),
559                    _ => Bound::Unbounded,
560                };
561                (start, end)
562            };
563            let left = match index_range.0 {
564                Bound::Included(index) => index,
565                Bound::Excluded(index) => (index + 1).min(expected.len()),
566                Bound::Unbounded => 0,
567            };
568            let right = match index_range.1 {
569                Bound::Included(index) => (index + 1).min(expected.len()),
570                Bound::Excluded(index) => index,
571                Bound::Unbounded => expected.len(),
572            };
573            assert_eq!(
574                tree.range_at(index_range).collect::<Vec<_>>(),
575                expected[left..right].to_vec()
576            );
577            assert_eq!(tree.iter().last(), expected.last().copied());
578            assert_eq!(tree.iter().min(), expected.first().copied());
579            assert_eq!(tree.iter().max(), expected.last().copied());
580
581            let mut iter = tree.iter();
582            let mut expected = VecDeque::from(expected);
583            while !expected.is_empty() {
584                if rng.rand(2) == 0 {
585                    assert_eq!(iter.next(), expected.pop_front());
586                } else {
587                    assert_eq!(iter.next_back(), expected.pop_back());
588                }
589            }
590            assert_eq!(iter.next(), None);
591            assert_eq!(iter.next_back(), None);
592        }
593    }
594
595    #[test]
596    fn test_drop() {
597        #[derive(Debug)]
598        struct CheckDrop;
599        thread_local! {
600            static COUNT: RefCell<usize> = const { RefCell::new(0) };
601        }
602        impl Drop for CheckDrop {
603            fn drop(&mut self) {
604                COUNT.with(|count| *count.borrow_mut() += 1);
605            }
606        }
607
608        let mut rng = Xorshift::default();
609        for _ in 0..100 {
610            COUNT.with(|count| *count.borrow_mut() = 0);
611            let mut inserted = 0;
612            let mut expected = BTreeSet::new();
613            {
614                let mut tree = SplayTree::new();
615                for _ in 0..1000 {
616                    let key = rng.random(0..=100);
617                    if rng.random(0..2) == 0 {
618                        tree.insert(key, CheckDrop);
619                        expected.insert(key);
620                        inserted += 1;
621                    } else {
622                        tree.remove(&key);
623                        expected.remove(&key);
624                    }
625                    assert_eq!(
626                        COUNT.with(|count| *count.borrow()),
627                        inserted - expected.len()
628                    );
629                }
630            }
631            assert_eq!(COUNT.with(|count| *count.borrow()), inserted);
632        }
633    }
634}