Skip to main content

amalgam/
locking.rs

1//! Single-flight locking for cache-stampede protection.
2//!
3//! FusionCache guarantees that, per key, only one factory runs at a time; other
4//! callers await that result. We implement this with a fixed bank of sharded
5//! async mutexes: a key always maps to the same shard, so concurrent calls for
6//! the same key serialize on the same lock. Distinct keys that happen to share a
7//! shard may serialize too, but this only affects throughput, never
8//! correctness — after acquiring the lock each caller re-checks its *own* key.
9//!
10//! Sharding (rather than a per-key map) keeps memory bounded and side-steps the
11//! notoriously race-prone "remove the lock entry when the last waiter leaves"
12//! cleanup. The guard is an [`OwnedMutexGuard`] so it is `'static + Send` and can
13//! be moved into a spawned background task (needed for background factory
14//! completion and eager refresh).
15
16use std::collections::hash_map::DefaultHasher;
17use std::hash::{Hash, Hasher};
18use std::sync::Arc;
19
20use tokio::sync::{Mutex, OwnedMutexGuard};
21
22/// A bank of sharded async mutexes providing per-key single-flight.
23#[derive(Debug)]
24pub struct KeyedLock {
25    shards: Box<[Arc<Mutex<()>>]>,
26    mask: usize,
27}
28
29/// The guard returned by acquiring a key's lock. Releasing it (on drop) lets the
30/// next waiter for that shard proceed.
31pub type KeyGuard = OwnedMutexGuard<()>;
32
33impl KeyedLock {
34    /// Creates a lock bank with at least `shards` shards (rounded up to a power
35    /// of two for fast masking).
36    #[must_use]
37    pub fn new(shards: usize) -> Self {
38        let count = shards.max(1).next_power_of_two();
39        let shards = (0..count).map(|_| Arc::new(Mutex::new(()))).collect();
40        Self {
41            shards,
42            mask: count - 1,
43        }
44    }
45
46    fn shard_for(&self, key: &str) -> &Arc<Mutex<()>> {
47        let mut hasher = DefaultHasher::new();
48        key.hash(&mut hasher);
49        let idx = (hasher.finish() as usize) & self.mask;
50        &self.shards[idx]
51    }
52
53    /// Acquires the lock for `key`, waiting if necessary.
54    pub async fn lock(&self, key: &str) -> KeyGuard {
55        Arc::clone(self.shard_for(key)).lock_owned().await
56    }
57
58    /// Tries to acquire the lock for `key` without waiting.
59    ///
60    /// Returns `None` if another caller currently holds the shard — used by
61    /// non-blocking paths (eager refresh) that must not stall the caller.
62    #[must_use]
63    pub fn try_lock(&self, key: &str) -> Option<KeyGuard> {
64        Arc::clone(self.shard_for(key)).try_lock_owned().ok()
65    }
66}
67
68impl Default for KeyedLock {
69    fn default() -> Self {
70        Self::new(1024)
71    }
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77    use std::sync::atomic::{AtomicUsize, Ordering};
78    use std::time::Duration;
79
80    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
81    async fn same_key_serializes() {
82        let lock = Arc::new(KeyedLock::new(64));
83        let counter = Arc::new(AtomicUsize::new(0));
84        let max_seen = Arc::new(AtomicUsize::new(0));
85
86        let mut handles = Vec::with_capacity(16);
87        for _ in 0..16 {
88            let lock = Arc::clone(&lock);
89            let counter = Arc::clone(&counter);
90            let max_seen = Arc::clone(&max_seen);
91            handles.push(tokio::spawn(async move {
92                let _g = lock.lock("hot-key").await;
93                let inside = counter.fetch_add(1, Ordering::SeqCst) + 1;
94                max_seen.fetch_max(inside, Ordering::SeqCst);
95                tokio::time::sleep(Duration::from_millis(5)).await;
96                counter.fetch_sub(1, Ordering::SeqCst);
97            }));
98        }
99        for h in handles {
100            h.await.unwrap();
101        }
102        // Never more than one holder of the same key at once.
103        assert_eq!(max_seen.load(Ordering::SeqCst), 1);
104    }
105}