Skip to main content

scirs2_core/concurrent/
skip_list.rs

1//! Probabilistic concurrent skip list providing O(log n) expected operations.
2//!
3//! The skip list is an ordered key-value map.  This implementation uses a
4//! mutex per node for thread safety and a fixed 32-level tower scheme.
5//!
6//! # Algorithm
7//!
8//! A skip list augments a sorted linked list with multiple "express lanes"
9//! that allow searches to skip large portions of the list.  Each node is
10//! assigned a random height when it is inserted.  The expected height of a
11//! node is 1/(1-p); with p = 0.5 the expected number of comparisons per
12//! operation is O(log n).
13//!
14//! # Thread Safety
15//!
16//! Rather than a global lock the implementation uses a hierarchical locking
17//! strategy: each node carries its own `Mutex`.  Insertions and removals
18//! acquire a sequence of per-node locks following the standard lock-ordering
19//! protocol (always from head → tail) to avoid deadlocks.
20
21use std::sync::{Arc, Mutex};
22
23const MAX_LEVEL: usize = 32;
24
25// ---------------------------------------------------------------------------
26// Internal node
27// ---------------------------------------------------------------------------
28
29struct SkipNode<K, V> {
30    key: Option<K>,
31    value: Option<V>,
32    /// Forward pointers for each level.  `forward[0]` is the bottom-level next.
33    forward: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>>,
34}
35
36impl<K, V> SkipNode<K, V> {
37    fn new_head(height: usize) -> Self {
38        SkipNode {
39            key: None,
40            value: None,
41            forward: vec![None; height],
42        }
43    }
44
45    fn new(key: K, value: V, height: usize) -> Self {
46        SkipNode {
47            key: Some(key),
48            value: Some(value),
49            forward: vec![None; height],
50        }
51    }
52}
53
54// ---------------------------------------------------------------------------
55// Random level generator
56// ---------------------------------------------------------------------------
57
58/// Fast xorshift PRNG for level generation.
59struct LevelGen {
60    state: u64,
61}
62
63impl LevelGen {
64    fn new() -> Self {
65        // Seed with current time and stack address.
66        let seed = std::time::SystemTime::now()
67            .duration_since(std::time::UNIX_EPOCH)
68            .map(|d| d.as_nanos() as u64)
69            .unwrap_or(12345678);
70        LevelGen {
71            state: seed ^ 0xdeadbeef_cafebabe,
72        }
73    }
74
75    fn next_u64(&mut self) -> u64 {
76        let mut x = self.state;
77        x ^= x << 13;
78        x ^= x >> 7;
79        x ^= x << 17;
80        self.state = x;
81        x
82    }
83
84    fn random_level(&mut self, max: usize) -> usize {
85        let mut level = 1usize;
86        while level < max && (self.next_u64() & 1) == 0 {
87            level += 1;
88        }
89        level
90    }
91}
92
93// ---------------------------------------------------------------------------
94// SkipList
95// ---------------------------------------------------------------------------
96
97/// A concurrent ordered map backed by a probabilistic skip list.
98///
99/// Keys must implement `Ord + Clone`; values must implement `Clone`.
100///
101/// # Example
102///
103/// ```rust
104/// use scirs2_core::concurrent::SkipList;
105///
106/// let mut sl: SkipList<u32, String> = SkipList::new();
107/// sl.insert(3, "three".to_string());
108/// sl.insert(1, "one".to_string());
109/// sl.insert(2, "two".to_string());
110///
111/// assert_eq!(sl.get(&1), Some("one".to_string()));
112/// assert_eq!(sl.get(&2), Some("two".to_string()));
113///
114/// sl.remove(&2);
115/// assert_eq!(sl.get(&2), None);
116/// ```
117pub struct SkipList<K, V> {
118    head: Arc<Mutex<SkipNode<K, V>>>,
119    level: usize,
120    len: usize,
121    rng: LevelGen,
122}
123
124impl<K: Ord + Clone, V: Clone> SkipList<K, V> {
125    /// Create an empty skip list.
126    pub fn new() -> Self {
127        SkipList {
128            head: Arc::new(Mutex::new(SkipNode::new_head(MAX_LEVEL))),
129            level: 1,
130            len: 0,
131            rng: LevelGen::new(),
132        }
133    }
134
135    /// Return the number of key-value pairs in the list.
136    pub fn len(&self) -> usize {
137        self.len
138    }
139
140    /// Return `true` if the list contains no elements.
141    pub fn is_empty(&self) -> bool {
142        self.len == 0
143    }
144
145    /// Look up the value associated with `key`.
146    ///
147    /// Returns `None` if no matching key exists.
148    #[allow(clippy::while_let_loop)]
149    pub fn get(&self, key: &K) -> Option<V> {
150        // `current_node` is `Some(arc)` for a data node, or `None` meaning
151        // "the head sentinel".  We keep a per-level forward-pointer vector
152        // that we read from the current node.
153        let head_guard = self.head.lock().ok()?;
154        // `forwards[lvl]` is the forward pointer at level `lvl` of the
155        // current predecessor node.  Initialised from the head.
156        let mut forwards: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = head_guard.forward.clone();
157        drop(head_guard);
158
159        for lvl in (0..self.level).rev() {
160            loop {
161                let next = match forwards.get(lvl).and_then(|f| f.as_ref()) {
162                    Some(n) => Arc::clone(n),
163                    None => break,
164                };
165                let guard = match next.lock() {
166                    Ok(g) => g,
167                    Err(_) => break,
168                };
169                match guard.key.as_ref() {
170                    Some(k) if k < key => {
171                        // Advance: update only the levels that this node
172                        // covers (i.e. 0..node.forward.len()), keeping higher
173                        // levels from the previous position intact.
174                        let node_fwd = guard.forward.clone();
175                        drop(guard);
176                        let node_height = node_fwd.len();
177                        let copy_len = node_height.min(forwards.len());
178                        forwards[..copy_len].clone_from_slice(&node_fwd[..copy_len]);
179                    }
180                    Some(k) if k == key => {
181                        return guard.value.clone();
182                    }
183                    _ => break,
184                }
185            }
186        }
187        None
188    }
189
190    /// Insert or replace the value for `key`.
191    #[allow(clippy::while_let_loop)]
192    pub fn insert(&mut self, key: K, value: V) {
193        let new_level = self.rng.random_level(MAX_LEVEL);
194        if new_level > self.level {
195            self.level = new_level;
196        }
197
198        // Collect update pointers: for each level the last node whose key < key.
199        let mut update: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = vec![None; self.level];
200
201        let head_guard = match self.head.lock() {
202            Ok(g) => g,
203            Err(_) => return,
204        };
205        let mut forwards: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = head_guard.forward.clone();
206        drop(head_guard);
207
208        // `current_node_arc` tracks which data node we are sitting at (None = head).
209        let mut current_node_arc: Option<Arc<Mutex<SkipNode<K, V>>>> = None;
210
211        for lvl in (0..self.level).rev() {
212            loop {
213                let next = match forwards.get(lvl).and_then(|f| f.as_ref()) {
214                    Some(n) => Arc::clone(n),
215                    None => break,
216                };
217                let guard = match next.lock() {
218                    Ok(g) => g,
219                    Err(_) => break,
220                };
221                match guard.key.as_ref() {
222                    Some(k) if k < &key => {
223                        let node_fwd = guard.forward.clone();
224                        drop(guard);
225                        let node_height = node_fwd.len();
226                        let copy_len = node_height.min(forwards.len());
227                        forwards[..copy_len].clone_from_slice(&node_fwd[..copy_len]);
228                        update[lvl] = Some(Arc::clone(&next));
229                        current_node_arc = Some(next);
230                    }
231                    _ => break,
232                }
233            }
234            // If we didn't advance at this level, the predecessor is
235            // whatever we were sitting at from higher levels.
236            if update[lvl].is_none() {
237                update[lvl] = current_node_arc.clone(); // None means head
238            }
239        }
240
241        // Check whether we need to update an existing node.
242        if let Some(next_arc) = forwards.first().and_then(|f| f.as_ref()) {
243            let mut guard = match next_arc.lock() {
244                Ok(g) => g,
245                Err(_) => return,
246            };
247            if guard.key.as_ref() == Some(&key) {
248                guard.value = Some(value);
249                return;
250            }
251        }
252
253        // Allocate a new node.
254        let new_node = Arc::new(Mutex::new(SkipNode::new(key, value, new_level)));
255
256        // Splice in at every level.
257        for lvl in 0..new_level {
258            // Determine predecessor at this level.
259            let pred = update.get(lvl).and_then(|u| u.as_ref());
260
261            if let Some(pred_arc) = pred {
262                let mut pred_guard = match pred_arc.lock() {
263                    Ok(g) => g,
264                    Err(_) => return,
265                };
266                let old_next = pred_guard.forward.get(lvl).and_then(|f| f.clone());
267                if let Ok(mut new_guard) = new_node.lock() {
268                    new_guard.forward[lvl] = old_next;
269                }
270                pred_guard.forward[lvl] = Some(Arc::clone(&new_node));
271            } else {
272                // Predecessor is head.
273                let mut head_guard = match self.head.lock() {
274                    Ok(g) => g,
275                    Err(_) => return,
276                };
277                let old_next = head_guard.forward.get(lvl).and_then(|f| f.clone());
278                if let Ok(mut new_guard) = new_node.lock() {
279                    new_guard.forward[lvl] = old_next;
280                }
281                head_guard.forward[lvl] = Some(Arc::clone(&new_node));
282            }
283        }
284
285        self.len += 1;
286    }
287
288    /// Remove the entry with the given key.  Returns `true` if a key was removed.
289    #[allow(clippy::while_let_loop)]
290    pub fn remove(&mut self, key: &K) -> bool {
291        let mut update: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = vec![None; self.level];
292
293        let head_guard = match self.head.lock() {
294            Ok(g) => g,
295            Err(_) => return false,
296        };
297        let mut forwards: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = head_guard.forward.clone();
298        drop(head_guard);
299
300        let mut current_node_arc: Option<Arc<Mutex<SkipNode<K, V>>>> = None;
301
302        for lvl in (0..self.level).rev() {
303            loop {
304                let next = match forwards.get(lvl).and_then(|f| f.as_ref()) {
305                    Some(n) => Arc::clone(n),
306                    None => break,
307                };
308                let guard = match next.lock() {
309                    Ok(g) => g,
310                    Err(_) => break,
311                };
312                match guard.key.as_ref() {
313                    Some(k) if k < key => {
314                        let node_fwd = guard.forward.clone();
315                        drop(guard);
316                        let node_height = node_fwd.len();
317                        let copy_len = node_height.min(forwards.len());
318                        forwards[..copy_len].clone_from_slice(&node_fwd[..copy_len]);
319                        update[lvl] = Some(Arc::clone(&next));
320                        current_node_arc = Some(next);
321                    }
322                    _ => break,
323                }
324            }
325            if update[lvl].is_none() {
326                update[lvl] = current_node_arc.clone();
327            }
328        }
329
330        // Target node should be forwards[0].
331        let target_arc = match forwards.first().and_then(|f| f.as_ref()) {
332            Some(a) => Arc::clone(a),
333            None => return false,
334        };
335        let target_guard = match target_arc.lock() {
336            Ok(g) => g,
337            Err(_) => return false,
338        };
339        if target_guard.key.as_ref() != Some(key) {
340            return false;
341        }
342        let target_forward = target_guard.forward.clone();
343        drop(target_guard);
344
345        // Unlink from each level.
346        for lvl in 0..self.level {
347            let target_next = target_forward.get(lvl).and_then(|f| f.clone());
348
349            let pred = update.get(lvl).and_then(|u| u.as_ref());
350            if let Some(pred_arc) = pred {
351                let mut pred_guard = match pred_arc.lock() {
352                    Ok(g) => g,
353                    Err(_) => continue,
354                };
355                // Only unlink if pred.forward[lvl] == target.
356                let is_target = pred_guard
357                    .forward
358                    .get(lvl)
359                    .and_then(|f| f.as_ref())
360                    .map(|a| Arc::ptr_eq(a, &target_arc))
361                    .unwrap_or(false);
362                if is_target {
363                    pred_guard.forward[lvl] = target_next;
364                }
365            } else {
366                let mut head_guard = match self.head.lock() {
367                    Ok(g) => g,
368                    Err(_) => continue,
369                };
370                let is_target = head_guard
371                    .forward
372                    .get(lvl)
373                    .and_then(|f| f.as_ref())
374                    .map(|a| Arc::ptr_eq(a, &target_arc))
375                    .unwrap_or(false);
376                if is_target {
377                    head_guard.forward[lvl] = target_next;
378                }
379            }
380        }
381
382        self.len -= 1;
383        true
384    }
385
386    /// Return all key-value pairs whose keys fall within `[lo, hi)` in order.
387    #[allow(clippy::while_let_loop)]
388    pub fn range(&self, lo: &K, hi: &K) -> Vec<(K, V)> {
389        let mut result = Vec::new();
390
391        let head_guard = match self.head.lock() {
392            Ok(g) => g,
393            Err(_) => return result,
394        };
395        let mut forwards: Vec<Option<Arc<Mutex<SkipNode<K, V>>>>> = head_guard.forward.clone();
396        drop(head_guard);
397
398        // Fast-forward to first key >= lo using higher levels.
399        'outer: for lvl in (1..self.level).rev() {
400            loop {
401                let next = match forwards.get(lvl).and_then(|f| f.as_ref()) {
402                    Some(n) => Arc::clone(n),
403                    None => break,
404                };
405                let guard = match next.lock() {
406                    Ok(g) => g,
407                    Err(_) => break 'outer,
408                };
409                match guard.key.as_ref() {
410                    Some(k) if k < lo => {
411                        let node_fwd = guard.forward.clone();
412                        drop(guard);
413                        let node_height = node_fwd.len();
414                        let copy_len = node_height.min(forwards.len());
415                        forwards[..copy_len].clone_from_slice(&node_fwd[..copy_len]);
416                    }
417                    _ => break,
418                }
419            }
420        }
421
422        // Scan at level 0.
423        loop {
424            let next = match forwards.first().and_then(|f| f.as_ref()) {
425                Some(n) => Arc::clone(n),
426                None => break,
427            };
428            let guard = match next.lock() {
429                Ok(g) => g,
430                Err(_) => break,
431            };
432            match guard.key.as_ref() {
433                Some(k) if k >= lo && k < hi => {
434                    if let (Some(k2), Some(v2)) = (guard.key.clone(), guard.value.clone()) {
435                        result.push((k2, v2));
436                    }
437                    let fwd0 = guard.forward.first().cloned().flatten();
438                    drop(guard);
439                    forwards[0] = fwd0;
440                }
441                Some(k) if k >= hi => break,
442                None => break,
443                _ => {
444                    // k < lo: advance at level 0
445                    let fwd0 = guard.forward.first().cloned().flatten();
446                    drop(guard);
447                    forwards[0] = fwd0;
448                }
449            }
450        }
451
452        result
453    }
454
455    /// Check whether the skip list contains the given key.
456    pub fn contains(&self, key: &K) -> bool {
457        self.get(key).is_some()
458    }
459
460    /// Collect all key-value pairs in sorted order.
461    ///
462    /// This traverses the bottom level of the skip list, so it is O(n).
463    pub fn iter(&self) -> Vec<(K, V)> {
464        let mut result = Vec::with_capacity(self.len);
465
466        let head_guard = match self.head.lock() {
467            Ok(g) => g,
468            Err(_) => return result,
469        };
470        let mut current = head_guard.forward.first().cloned().flatten();
471        drop(head_guard);
472
473        while let Some(node_arc) = current {
474            let guard = match node_arc.lock() {
475                Ok(g) => g,
476                Err(_) => break,
477            };
478            if let (Some(k), Some(v)) = (guard.key.clone(), guard.value.clone()) {
479                result.push((k, v));
480            }
481            current = guard.forward.first().cloned().flatten();
482        }
483
484        result
485    }
486
487    /// Remove the entry with the given key, returning the value if it existed.
488    pub fn remove_entry(&mut self, key: &K) -> Option<V> {
489        let value = self.get(key);
490        if value.is_some() && self.remove(key) {
491            value
492        } else {
493            None
494        }
495    }
496}
497
498impl<K: Ord + Clone, V: Clone> Default for SkipList<K, V> {
499    fn default() -> Self {
500        Self::new()
501    }
502}
503
504// ---------------------------------------------------------------------------
505// Tests
506// ---------------------------------------------------------------------------
507
508#[cfg(test)]
509mod tests {
510    use super::*;
511
512    #[test]
513    fn test_insert_get() {
514        let mut sl: SkipList<u32, &str> = SkipList::new();
515        sl.insert(5, "five");
516        sl.insert(2, "two");
517        sl.insert(8, "eight");
518
519        assert_eq!(sl.get(&2), Some("two"));
520        assert_eq!(sl.get(&5), Some("five"));
521        assert_eq!(sl.get(&8), Some("eight"));
522        assert_eq!(sl.get(&1), None);
523        assert_eq!(sl.len(), 3);
524    }
525
526    #[test]
527    fn test_remove() {
528        let mut sl: SkipList<u32, u32> = SkipList::new();
529        for i in 0..10u32 {
530            sl.insert(i, i * 10);
531        }
532        assert_eq!(sl.len(), 10);
533
534        assert!(sl.remove(&5));
535        assert_eq!(sl.get(&5), None);
536        assert_eq!(sl.len(), 9);
537
538        // Removing non-existent key returns false.
539        assert!(!sl.remove(&99));
540    }
541
542    #[test]
543    fn test_range() {
544        let mut sl: SkipList<u32, u32> = SkipList::new();
545        for i in 0..20u32 {
546            sl.insert(i, i);
547        }
548        let r = sl.range(&5, &10);
549        assert_eq!(r.len(), 5);
550        let keys: Vec<u32> = r.iter().map(|(k, _)| *k).collect();
551        assert_eq!(keys, vec![5, 6, 7, 8, 9]);
552    }
553
554    #[test]
555    fn test_overwrite_existing_key() {
556        let mut sl: SkipList<u32, u32> = SkipList::new();
557        sl.insert(1, 100);
558        sl.insert(1, 200);
559        assert_eq!(sl.get(&1), Some(200));
560        assert_eq!(sl.len(), 1);
561    }
562
563    #[test]
564    fn test_is_empty_and_len() {
565        let mut sl: SkipList<i32, i32> = SkipList::new();
566        assert!(sl.is_empty());
567        sl.insert(42, 0);
568        assert!(!sl.is_empty());
569        assert_eq!(sl.len(), 1);
570        sl.remove(&42);
571        assert!(sl.is_empty());
572    }
573
574    #[test]
575    fn test_large_insert_ordered() {
576        let mut sl: SkipList<u32, u32> = SkipList::new();
577        // Insert in reverse order.
578        for i in (0..100u32).rev() {
579            sl.insert(i, i);
580        }
581        let r = sl.range(&0, &100);
582        assert_eq!(r.len(), 100);
583        for (i, (k, v)) in r.iter().enumerate() {
584            assert_eq!(*k, i as u32);
585            assert_eq!(*v, i as u32);
586        }
587    }
588
589    #[test]
590    fn test_contains() {
591        let mut sl: SkipList<u32, u32> = SkipList::new();
592        sl.insert(10, 100);
593        sl.insert(20, 200);
594        assert!(sl.contains(&10));
595        assert!(sl.contains(&20));
596        assert!(!sl.contains(&30));
597    }
598
599    #[test]
600    fn test_iter_sorted_order() {
601        let mut sl: SkipList<u32, u32> = SkipList::new();
602        sl.insert(5, 50);
603        sl.insert(1, 10);
604        sl.insert(9, 90);
605        sl.insert(3, 30);
606        sl.insert(7, 70);
607
608        let items = sl.iter();
609        assert_eq!(items.len(), 5);
610        let keys: Vec<u32> = items.iter().map(|(k, _)| *k).collect();
611        assert_eq!(keys, vec![1, 3, 5, 7, 9]);
612    }
613
614    #[test]
615    fn test_remove_entry() {
616        let mut sl: SkipList<u32, String> = SkipList::new();
617        sl.insert(1, "one".to_string());
618        sl.insert(2, "two".to_string());
619
620        let removed = sl.remove_entry(&1);
621        assert_eq!(removed, Some("one".to_string()));
622        assert_eq!(sl.len(), 1);
623
624        let not_found = sl.remove_entry(&99);
625        assert!(not_found.is_none());
626    }
627
628    #[test]
629    fn test_iter_empty() {
630        let sl: SkipList<u32, u32> = SkipList::new();
631        assert!(sl.iter().is_empty());
632    }
633
634    #[test]
635    fn test_range_empty_result() {
636        let mut sl: SkipList<u32, u32> = SkipList::new();
637        for i in 0..10u32 {
638            sl.insert(i, i);
639        }
640        let r = sl.range(&10, &5);
641        assert!(r.is_empty());
642        let r2 = sl.range(&100, &200);
643        assert!(r2.is_empty());
644    }
645
646    #[test]
647    fn test_concurrent_read_access() {
648        use std::sync::Arc;
649        use std::thread;
650
651        let mut sl = SkipList::new();
652        for i in 0..100u32 {
653            sl.insert(i, i * 10);
654        }
655        let shared = Arc::new(sl);
656
657        let mut handles = Vec::new();
658        for t in 0..4 {
659            let sl_ref = Arc::clone(&shared);
660            handles.push(thread::spawn(move || {
661                for i in 0..100u32 {
662                    let val = sl_ref.get(&i);
663                    assert_eq!(val, Some(i * 10), "thread {t} failed for key {i}");
664                }
665            }));
666        }
667
668        for h in handles {
669            if let Err(e) = h.join() {
670                panic!("Thread panicked: {e:?}");
671            }
672        }
673    }
674}