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}