1use std::collections::hash_map::DefaultHasher;
17use std::hash::{Hash, Hasher};
18use std::sync::Arc;
19
20use tokio::sync::{Mutex, OwnedMutexGuard};
21
22#[derive(Debug)]
24pub struct KeyedLock {
25 shards: Box<[Arc<Mutex<()>>]>,
26 mask: usize,
27}
28
29pub type KeyGuard = OwnedMutexGuard<()>;
32
33impl KeyedLock {
34 #[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 pub async fn lock(&self, key: &str) -> KeyGuard {
55 Arc::clone(self.shard_for(key)).lock_owned().await
56 }
57
58 #[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 assert_eq!(max_seen.load(Ordering::SeqCst), 1);
104 }
105}