use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use tokio::sync::{Mutex, OwnedMutexGuard};
#[derive(Debug)]
pub struct KeyedLock {
shards: Box<[Arc<Mutex<()>>]>,
mask: usize,
}
pub type KeyGuard = OwnedMutexGuard<()>;
impl KeyedLock {
#[must_use]
pub fn new(shards: usize) -> Self {
let count = shards.max(1).next_power_of_two();
let shards = (0..count).map(|_| Arc::new(Mutex::new(()))).collect();
Self {
shards,
mask: count - 1,
}
}
fn shard_for(&self, key: &str) -> &Arc<Mutex<()>> {
let mut hasher = DefaultHasher::new();
key.hash(&mut hasher);
let idx = (hasher.finish() as usize) & self.mask;
&self.shards[idx]
}
pub async fn lock(&self, key: &str) -> KeyGuard {
Arc::clone(self.shard_for(key)).lock_owned().await
}
#[must_use]
pub fn try_lock(&self, key: &str) -> Option<KeyGuard> {
Arc::clone(self.shard_for(key)).try_lock_owned().ok()
}
}
impl Default for KeyedLock {
fn default() -> Self {
Self::new(1024)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn same_key_serializes() {
let lock = Arc::new(KeyedLock::new(64));
let counter = Arc::new(AtomicUsize::new(0));
let max_seen = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::with_capacity(16);
for _ in 0..16 {
let lock = Arc::clone(&lock);
let counter = Arc::clone(&counter);
let max_seen = Arc::clone(&max_seen);
handles.push(tokio::spawn(async move {
let _g = lock.lock("hot-key").await;
let inside = counter.fetch_add(1, Ordering::SeqCst) + 1;
max_seen.fetch_max(inside, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(5)).await;
counter.fetch_sub(1, Ordering::SeqCst);
}));
}
for h in handles {
h.await.unwrap();
}
assert_eq!(max_seen.load(Ordering::SeqCst), 1);
}
}