use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, Weak};
use std::task::{Context, Poll, Waker};
use super::{Lock, LockError, LockManager};
#[derive(Default)]
struct LockState {
locked: bool,
waiters: VecDeque<Waker>,
}
struct LockRegistry {
locks: Weak<Mutex<HashMap<String, Arc<InMemoryLock>>>>,
key: String,
}
pub struct InMemoryLock {
state: Mutex<LockState>,
registry: Option<LockRegistry>,
}
impl InMemoryLock {
pub fn new() -> Self {
InMemoryLock {
state: Mutex::new(LockState::default()),
registry: None,
}
}
fn try_lock_core(&self) -> Result<bool, LockError> {
let mut state = self
.state
.lock()
.map_err(|err| LockError::Poisoned(err.to_string()))?;
if state.locked {
Ok(false)
} else {
state.locked = true;
Ok(true)
}
}
pub(crate) fn unlock_core(&self) -> Result<(), LockError> {
let woken = {
let mut state = self
.state
.lock()
.map_err(|err| LockError::Poisoned(err.to_string()))?;
if state.locked {
state.locked = false;
std::mem::take(&mut state.waiters)
} else {
VecDeque::new()
}
};
for waker in woken {
waker.wake();
}
self.evict_if_idle();
Ok(())
}
fn evict_if_idle(&self) {
let Some(registry) = &self.registry else {
return;
};
let Some(locks) = registry.locks.upgrade() else {
return;
};
let Ok(mut locks) = locks.lock() else {
return;
};
let Some(entry) = locks.get(®istry.key) else {
return;
};
if !std::ptr::eq(Arc::as_ptr(entry), self) || Arc::strong_count(entry) != 2 {
return;
}
let idle = self
.state
.lock()
.map(|state| !state.locked && state.waiters.is_empty())
.unwrap_or(false);
if idle {
locks.remove(®istry.key);
}
}
}
impl Default for InMemoryLock {
fn default() -> Self {
Self::new()
}
}
pub struct InMemoryLockFuture<'a> {
lock: &'a InMemoryLock,
}
impl Future for InMemoryLockFuture<'_> {
type Output = Result<(), LockError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut state = match self.lock.state.lock() {
Ok(state) => state,
Err(err) => return Poll::Ready(Err(LockError::Poisoned(err.to_string()))),
};
if !state.locked {
state.locked = true;
Poll::Ready(Ok(()))
} else {
if !state
.waiters
.iter()
.any(|waker| waker.will_wake(cx.waker()))
{
state.waiters.push_back(cx.waker().clone());
}
Poll::Pending
}
}
}
impl Lock for InMemoryLock {
fn lock(&self) -> impl Future<Output = Result<(), LockError>> + Send + '_ {
InMemoryLockFuture { lock: self }
}
async fn try_lock(&self) -> Result<bool, LockError> {
self.try_lock_core()
}
async fn unlock(&self) -> Result<(), LockError> {
self.unlock_core()
}
}
pub struct InMemoryLockManager {
locks: Arc<Mutex<HashMap<String, Arc<InMemoryLock>>>>,
}
impl InMemoryLockManager {
pub fn new() -> Self {
InMemoryLockManager {
locks: Arc::new(Mutex::new(HashMap::new())),
}
}
}
impl Default for InMemoryLockManager {
fn default() -> Self {
Self::new()
}
}
impl LockManager for InMemoryLockManager {
type Lock = InMemoryLock;
fn get_lock(&self, id: &str) -> Result<Arc<InMemoryLock>, LockError> {
let mut locks = self
.locks
.lock()
.map_err(|_| LockError::Poisoned("lock manager map poisoned".into()))?;
Ok(locks
.entry(id.to_string())
.or_insert_with(|| {
Arc::new(InMemoryLock {
state: Mutex::new(LockState::default()),
registry: Some(LockRegistry {
locks: Arc::downgrade(&self.locks),
key: id.to_string(),
}),
})
})
.clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
use std::time::Duration;
fn reentrant_waker(lock: Arc<InMemoryLock>) -> Waker {
unsafe fn clone(data: *const ()) -> RawWaker {
let arc = unsafe { Arc::from_raw(data as *const InMemoryLock) };
let cloned = Arc::clone(&arc);
std::mem::forget(arc);
RawWaker::new(Arc::into_raw(cloned) as *const (), &REENTRANT_VTABLE)
}
unsafe fn wake(data: *const ()) {
let arc = unsafe { Arc::from_raw(data as *const InMemoryLock) };
let _ = arc.try_lock_core(); }
unsafe fn wake_by_ref(data: *const ()) {
let arc = unsafe { Arc::from_raw(data as *const InMemoryLock) };
let _ = arc.try_lock_core();
std::mem::forget(arc);
}
unsafe fn drop_fn(data: *const ()) {
drop(unsafe { Arc::from_raw(data as *const InMemoryLock) });
}
static REENTRANT_VTABLE: RawWakerVTable =
RawWakerVTable::new(clone, wake, wake_by_ref, drop_fn);
let raw = RawWaker::new(Arc::into_raw(lock) as *const (), &REENTRANT_VTABLE);
unsafe { Waker::from_raw(raw) }
}
fn panicking_waker() -> Waker {
unsafe fn clone(_: *const ()) -> RawWaker {
RawWaker::new(std::ptr::null(), &PANIC_VTABLE)
}
unsafe fn wake(_: *const ()) {
panic!("waker panicked in wake()");
}
unsafe fn wake_by_ref(_: *const ()) {
panic!("waker panicked in wake_by_ref()");
}
unsafe fn drop_fn(_: *const ()) {}
static PANIC_VTABLE: RawWakerVTable =
RawWakerVTable::new(clone, wake, wake_by_ref, drop_fn);
unsafe { Waker::from_raw(RawWaker::new(std::ptr::null(), &PANIC_VTABLE)) }
}
fn park_waker(lock: &InMemoryLock, waker: &Waker) {
let mut cx = Context::from_waker(waker);
let mut fut = std::pin::pin!(lock.lock());
assert!(matches!(fut.as_mut().poll(&mut cx), Poll::Pending));
}
#[tokio::test]
async fn try_lock_reflects_state() {
let lock = InMemoryLock::new();
assert!(lock.try_lock().await.unwrap()); assert!(!lock.try_lock().await.unwrap()); lock.unlock().await.unwrap();
assert!(lock.try_lock().await.unwrap()); }
#[tokio::test]
async fn lock_resolves_immediately_when_free() {
let lock = InMemoryLock::new();
lock.lock().await.unwrap();
assert!(!lock.try_lock().await.unwrap()); lock.unlock().await.unwrap();
assert!(lock.try_lock().await.unwrap());
}
#[tokio::test]
async fn second_acquire_waits_until_unlock() {
let lock = Arc::new(InMemoryLock::new());
lock.lock().await.unwrap();
let order = Arc::new(AtomicUsize::new(0));
let waiter_lock = Arc::clone(&lock);
let waiter_order = Arc::clone(&order);
let waiter = tokio::spawn(async move {
waiter_lock.lock().await.unwrap();
waiter_order.fetch_add(1, Ordering::SeqCst)
});
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(
order.load(Ordering::SeqCst),
0,
"waiter must still be parked"
);
lock.unlock().await.unwrap();
let acquired_at = waiter.await.unwrap();
assert_eq!(acquired_at, 0, "waiter acquired exactly once after unlock");
assert!(!lock.try_lock().await.unwrap(), "waiter holds the lock");
}
#[test]
fn manager_returns_same_arc_per_key() {
let manager = InMemoryLockManager::new();
let a1 = manager.get_lock("agg-1").unwrap();
let a2 = manager.get_lock("agg-1").unwrap();
let b = manager.get_lock("agg-2").unwrap();
assert!(Arc::ptr_eq(&a1, &a2));
assert!(!Arc::ptr_eq(&a1, &b));
}
#[tokio::test]
async fn manager_evicts_idle_entry_on_unlock() {
let manager = InMemoryLockManager::new();
let lock = manager.get_lock("agg-1").unwrap();
lock.lock().await.unwrap();
assert_eq!(manager.locks.lock().unwrap().len(), 1);
lock.unlock().await.unwrap();
assert!(
manager.locks.lock().unwrap().is_empty(),
"an idle, otherwise-unreferenced entry is evicted on unlock"
);
let again = manager.get_lock("agg-1").unwrap();
assert!(again.try_lock().await.unwrap());
}
#[tokio::test]
async fn manager_keeps_entry_while_another_handle_is_held() {
let manager = InMemoryLockManager::new();
let a = manager.get_lock("agg-1").unwrap();
let b = manager.get_lock("agg-1").unwrap();
a.lock().await.unwrap();
a.unlock().await.unwrap();
let c = manager.get_lock("agg-1").unwrap();
assert!(Arc::ptr_eq(&b, &c));
assert_eq!(manager.locks.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn manager_keeps_entry_for_a_parked_waiter() {
let manager = InMemoryLockManager::new();
let holder = manager.get_lock("agg-1").unwrap();
holder.lock().await.unwrap();
let waiter_lock = manager.get_lock("agg-1").unwrap();
let waiter = tokio::spawn(async move {
waiter_lock.lock().await.unwrap();
waiter_lock
});
tokio::time::sleep(Duration::from_millis(20)).await;
holder.unlock().await.unwrap();
let waiter_lock = waiter.await.unwrap();
let same = manager.get_lock("agg-1").unwrap();
assert!(Arc::ptr_eq(&waiter_lock, &same));
assert!(!same.try_lock().await.unwrap(), "waiter holds the lock");
drop(holder);
drop(same);
waiter_lock.unlock().await.unwrap();
assert!(
manager.locks.lock().unwrap().is_empty(),
"the entry is evicted once the last holder unlocks"
);
}
#[tokio::test]
async fn distinct_keys_do_not_contend() {
let manager = InMemoryLockManager::new();
let a = manager.get_lock("agg-1").unwrap();
let b = manager.get_lock("agg-2").unwrap();
a.lock().await.unwrap();
b.lock().await.unwrap();
a.unlock().await.unwrap();
b.unlock().await.unwrap();
}
#[test]
fn unlock_does_not_deadlock_with_reentrant_waker() {
let lock = Arc::new(InMemoryLock::new());
assert!(lock.try_lock_core().unwrap()); park_waker(&lock, &reentrant_waker(Arc::clone(&lock)));
let (tx, rx) = mpsc::channel();
let unlock_lock = Arc::clone(&lock);
std::thread::spawn(move || {
let _ = tx.send(unlock_lock.unlock_core());
});
let result = rx
.recv_timeout(Duration::from_secs(2))
.expect("unlock deadlocked while waking a re-entrant waker");
result.expect("unlock should succeed");
}
#[test]
fn unlock_does_not_poison_when_a_waker_panics() {
let lock = Arc::new(InMemoryLock::new());
assert!(lock.try_lock_core().unwrap()); park_waker(&lock, &panicking_waker());
let unlock_lock = Arc::clone(&lock);
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = unlock_lock.unlock_core();
}))
.is_err();
assert!(panicked, "the panicking waker should unwind out of unlock");
assert!(
lock.try_lock_core().unwrap(),
"lock must remain usable after a waker panic"
);
}
}