diff --git a/klock.go b/klock.go index ac47526..d6a5395 100644 --- a/klock.go +++ b/klock.go @@ -185,6 +185,7 @@ func (kl *KeyLock) TryLock(key string) bool { // LockWithTimeout 尝试在给定的时间内为指定的键获取写锁。 // 如果在超时前成功获取锁,则返回 true;否则返回 false。 func (kl *KeyLock) LockWithTimeout(key string, timeout time.Duration) bool { + deadline := time.Now().Add(timeout) le := kl.prepareLock(key) // 立即尝试一次,以避免在锁可用时产生不必要的延迟。 if le.rw.TryLock() { @@ -197,25 +198,33 @@ func (kl *KeyLock) LockWithTimeout(key string, timeout time.Duration) bool { kl.cancelLock(key, le) } }() - timer := time.NewTimer(timeout) - defer timer.Stop() pollInterval := kl.config.InitialPollInterval maxPollInterval := kl.config.MaxPollInterval for { - select { - case <-timer.C: + remaining := time.Until(deadline) + if remaining <= 0 { return false - default: - time.Sleep(pollInterval) - if le.rw.TryLock() { - kl.commitLock(key, le) - acquired = true - return true - } - pollInterval *= 2 - if pollInterval > maxPollInterval { - pollInterval = maxPollInterval + } + sleepInterval := pollInterval + if sleepInterval > remaining { + sleepInterval = remaining + } + time.Sleep(sleepInterval) + if time.Until(deadline) <= 0 { + return false + } + if le.rw.TryLock() { + if time.Until(deadline) <= 0 { + le.rw.Unlock() + return false } + kl.commitLock(key, le) + acquired = true + return true + } + pollInterval *= 2 + if pollInterval > maxPollInterval { + pollInterval = maxPollInterval } } } @@ -244,6 +253,7 @@ func (kl *KeyLock) TryRLock(key string) bool { // RLockWithTimeout 尝试在给定的时间内为指定的键获取读锁。 // 如果在超时前成功获取锁,则返回 true;否则返回 false。 func (kl *KeyLock) RLockWithTimeout(key string, timeout time.Duration) bool { + deadline := time.Now().Add(timeout) le := kl.prepareLock(key) // 立即尝试一次 if le.rw.TryRLock() { @@ -256,25 +266,33 @@ func (kl *KeyLock) RLockWithTimeout(key string, timeout time.Duration) bool { kl.cancelLock(key, le) } }() - timer := time.NewTimer(timeout) - defer timer.Stop() pollInterval := kl.config.InitialPollInterval maxPollInterval := kl.config.MaxPollInterval for { - select { - case <-timer.C: + remaining := time.Until(deadline) + if remaining <= 0 { return false - default: - time.Sleep(pollInterval) - if le.rw.TryRLock() { - kl.commitLock(key, le) - acquired = true - return true - } - pollInterval *= 2 - if pollInterval > maxPollInterval { - pollInterval = maxPollInterval + } + sleepInterval := pollInterval + if sleepInterval > remaining { + sleepInterval = remaining + } + time.Sleep(sleepInterval) + if time.Until(deadline) <= 0 { + return false + } + if le.rw.TryRLock() { + if time.Until(deadline) <= 0 { + le.rw.RUnlock() + return false } + kl.commitLock(key, le) + acquired = true + return true + } + pollInterval *= 2 + if pollInterval > maxPollInterval { + pollInterval = maxPollInterval } } } diff --git a/klock_test.go b/klock_test.go index 3114f0f..a5a7b5c 100644 --- a/klock_test.go +++ b/klock_test.go @@ -250,6 +250,30 @@ func TestLockWithTimeout(t *testing.T) { } } +// TestLockWithTimeoutDoesNotAcquireAfterDeadline 测试写锁在超时后被释放时不会再被获取。 +func TestLockWithTimeoutDoesNotAcquireAfterDeadline(t *testing.T) { + kl := NewWithConfig(Config{ + MaxPollInterval: 50 * time.Millisecond, + InitialPollInterval: 50 * time.Millisecond, + }) + key := "test_key" + released := make(chan struct{}) + kl.Lock(key) + go func() { + time.Sleep(20 * time.Millisecond) + kl.Unlock(key) + close(released) + }() + acquired := kl.LockWithTimeout(key, 10*time.Millisecond) + if acquired { + kl.Unlock(key) + } + <-released + if acquired { + t.Fatal("LockWithTimeout 不应该在超时后获取锁") + } +} + // TestRLockWithTimeout 测试带超时的读锁获取功能。 func TestRLockWithTimeout(t *testing.T) { kl := New() @@ -270,6 +294,30 @@ func TestRLockWithTimeout(t *testing.T) { } } +// TestRLockWithTimeoutDoesNotAcquireAfterDeadline 测试读锁在超时后被释放时不会再被获取。 +func TestRLockWithTimeoutDoesNotAcquireAfterDeadline(t *testing.T) { + kl := NewWithConfig(Config{ + MaxPollInterval: 50 * time.Millisecond, + InitialPollInterval: 50 * time.Millisecond, + }) + key := "test_key" + released := make(chan struct{}) + kl.Lock(key) + go func() { + time.Sleep(20 * time.Millisecond) + kl.Unlock(key) + close(released) + }() + acquired := kl.RLockWithTimeout(key, 10*time.Millisecond) + if acquired { + kl.RUnlock(key) + } + <-released + if acquired { + t.Fatal("RLockWithTimeout 不应该在超时后获取锁") + } +} + // TestPanicOnUnlockOfUnlockedKey 测试对未锁定的键执行 Unlock 操作是否会引发 panic。 func TestPanicOnUnlockOfUnlockedKey(t *testing.T) { defer func() {