diff --git a/crates/lock/src/fast_lock/shard.rs b/crates/lock/src/fast_lock/shard.rs index e6cb39485..460d6640e 100644 --- a/crates/lock/src/fast_lock/shard.rs +++ b/crates/lock/src/fast_lock/shard.rs @@ -41,6 +41,42 @@ pub struct LockShard { active_guards: parking_lot::Mutex>, } +/// Cancellation-safe waiter counter ticket. +/// +/// Ensures waiting counters are decremented even if the waiting future +/// is cancelled/dropped before the normal post-await path runs. +struct WaiterCounterGuard { + state: Arc, + mode: LockMode, + incremented: bool, +} + +impl WaiterCounterGuard { + fn new(state: Arc, mode: LockMode) -> Self { + let incremented = match mode { + LockMode::Shared => state.atomic_state.inc_readers_waiting(), + LockMode::Exclusive => state.atomic_state.inc_writers_waiting(), + }; + Self { + state, + mode, + incremented, + } + } +} + +impl Drop for WaiterCounterGuard { + fn drop(&mut self) { + if !self.incremented { + return; + } + match self.mode { + LockMode::Shared => self.state.atomic_state.dec_readers_waiting(), + LockMode::Exclusive => self.state.atomic_state.dec_writers_waiting(), + } + } +} + impl LockShard { pub fn new(shard_id: usize) -> Self { Self { @@ -184,16 +220,12 @@ impl LockShard { // If we've exhausted quick retries or have little time left, use notification wait let wait_result = match request.mode { LockMode::Shared => { - state.atomic_state.inc_readers_waiting(); - let result = timeout(remaining, state.optimized_notify.wait_for_read()).await; - state.atomic_state.dec_readers_waiting(); - result + let _waiter_guard = WaiterCounterGuard::new(state.clone(), LockMode::Shared); + timeout(remaining, state.optimized_notify.wait_for_read()).await } LockMode::Exclusive => { - state.atomic_state.inc_writers_waiting(); - let result = timeout(remaining, state.optimized_notify.wait_for_write()).await; - state.atomic_state.dec_writers_waiting(); - result + let _waiter_guard = WaiterCounterGuard::new(state.clone(), LockMode::Exclusive); + timeout(remaining, state.optimized_notify.wait_for_write()).await } }; @@ -784,4 +816,70 @@ mod tests { let lock_info = shard.get_lock_info(&obj1_key); assert!(lock_info.is_some(), "obj1 should still be locked by blocking_owner"); } + + #[tokio::test] + async fn test_exclusive_waiter_abort_does_not_block_following_shared_lock() { + let shard = Arc::new(LockShard::new(0)); + let key = ObjectKey::new("bucket", "abort-waiter-key"); + + let owner1: Arc = Arc::from("writer-owner-1"); + let owner2: Arc = Arc::from("writer-owner-2"); + let reader_owner: Arc = Arc::from("reader-owner"); + + let hold_writer = ObjectLockRequest { + key: key.clone(), + mode: LockMode::Exclusive, + owner: owner1.clone(), + acquire_timeout: Duration::from_secs(1), + lock_timeout: Duration::from_secs(30), + priority: LockPriority::Normal, + }; + + assert!(shard.acquire_lock(&hold_writer).await.is_ok()); + + let contended_writer = ObjectLockRequest { + key: key.clone(), + mode: LockMode::Exclusive, + owner: owner2.clone(), + acquire_timeout: Duration::from_secs(5), + lock_timeout: Duration::from_secs(30), + priority: LockPriority::Normal, + }; + + let shard_for_waiter = shard.clone(); + let waiter_handle = tokio::spawn(async move { shard_for_waiter.acquire_lock(&contended_writer).await }); + + // Ensure we actually enter slow-path wait registration before aborting. + tokio::time::timeout(Duration::from_secs(3), async { + loop { + if let Some(state) = shard.objects.read().get(&key).cloned() + && state.atomic_state.writers_waiting_count() > 0 + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("timed out waiting for contended writer to register as waiting"); + waiter_handle.abort(); + let _ = waiter_handle.await; + + assert!(shard.release_lock(&key, &owner1, LockMode::Exclusive)); + + let followup_reader = ObjectLockRequest { + key: key.clone(), + mode: LockMode::Shared, + owner: reader_owner.clone(), + acquire_timeout: Duration::from_millis(200), + lock_timeout: Duration::from_secs(30), + priority: LockPriority::Normal, + }; + + assert!( + shard.acquire_lock(&followup_reader).await.is_ok(), + "shared lock should succeed after writer waiter task is aborted" + ); + assert!(shard.release_lock(&key, &reader_owner, LockMode::Shared)); + } } diff --git a/crates/lock/src/fast_lock/state.rs b/crates/lock/src/fast_lock/state.rs index 2896ead00..847b02c11 100644 --- a/crates/lock/src/fast_lock/state.rs +++ b/crates/lock/src/fast_lock/state.rs @@ -165,13 +165,13 @@ impl AtomicLockState { } /// Increment waiting readers count - pub fn inc_readers_waiting(&self) { + pub fn inc_readers_waiting(&self) -> bool { loop { let current = self.state.load(Ordering::Acquire); let waiting = self.readers_waiting(current); if waiting == 0xFFFF { - break; // Max waiting readers + return false; // Max waiting readers } let new_state = current + (1 << READERS_WAITING_SHIFT); @@ -181,7 +181,7 @@ impl AtomicLockState { .compare_exchange_weak(current, new_state, Ordering::AcqRel, Ordering::Relaxed) .is_ok() { - break; + return true; } } } @@ -209,13 +209,13 @@ impl AtomicLockState { } /// Increment waiting writers count - pub fn inc_writers_waiting(&self) { + pub fn inc_writers_waiting(&self) -> bool { loop { let current = self.state.load(Ordering::Acquire); let waiting = self.writers_waiting(current); if waiting == 0xFFFF { - break; // Max waiting writers + return false; // Max waiting writers } let new_state = current + (1 << WRITERS_WAITING_SHIFT); @@ -225,7 +225,7 @@ impl AtomicLockState { .compare_exchange_weak(current, new_state, Ordering::AcqRel, Ordering::Relaxed) .is_ok() { - break; + return true; } } } @@ -288,6 +288,12 @@ impl AtomicLockState { fn writers_waiting(&self, state: u64) -> u16 { ((state & WRITERS_WAITING_MASK) >> WRITERS_WAITING_SHIFT) as u16 } + + #[cfg(test)] + pub fn writers_waiting_count(&self) -> u16 { + let state = self.state.load(Ordering::Acquire); + self.writers_waiting(state) + } } /// Object lock state with version support - optimized memory layout