diff --git a/pkg/lock/lock.go b/pkg/lock/lock.go index f525f781000..c504f107886 100644 --- a/pkg/lock/lock.go +++ b/pkg/lock/lock.go @@ -49,7 +49,7 @@ type longOperationLocker interface { type longOperation struct { lock longOperationLocker mu sync.Mutex - tick *time.Ticker + done chan struct{} timeout time.Duration } @@ -57,23 +57,26 @@ func (l *longOperation) Lock() error { if err := l.lock.Lock(); err != nil { return err } - l.tick = time.NewTicker(l.timeout / 3) + l.mu.Lock() + done := make(chan struct{}) + l.done = done + l.mu.Unlock() go func() { - defer l.mu.Unlock() + tick := time.NewTicker(l.timeout / 3) + defer tick.Stop() for { - l.mu.Lock() - if l.tick == nil { - return - } - ch := l.tick.C - l.mu.Unlock() - <-ch - l.mu.Lock() - if l.tick == nil { + select { + case <-done: return + case <-tick.C: + l.mu.Lock() + if l.done != done { + l.mu.Unlock() + return + } + l.lock.Extend() + l.mu.Unlock() } - l.lock.Extend() - l.mu.Unlock() } }() return nil @@ -82,9 +85,9 @@ func (l *longOperation) Lock() error { func (l *longOperation) Unlock() { l.mu.Lock() defer l.mu.Unlock() - if l.tick != nil { - l.tick.Stop() - l.tick = nil + if l.done != nil { + close(l.done) + l.done = nil } l.lock.Unlock() } diff --git a/pkg/lock/long_operation_shutdown_test.go b/pkg/lock/long_operation_shutdown_test.go new file mode 100644 index 00000000000..4106ac9a15a --- /dev/null +++ b/pkg/lock/long_operation_shutdown_test.go @@ -0,0 +1,61 @@ +package lock + +import ( + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +type shutdownLeaseStub struct { + renewals atomic.Int32 + unlocked atomic.Bool + renewalsAfterUnlock atomic.Int32 +} + +func (s *shutdownLeaseStub) Lock() error { s.unlocked.Store(false); return nil } +func (s *shutdownLeaseStub) Unlock() { s.unlocked.Store(true) } +func (s *shutdownLeaseStub) Extend() { + s.renewals.Add(1) + if s.unlocked.Load() { + s.renewalsAfterUnlock.Add(1) + } +} + +func TestLongOperationStopsRenewingOnUnlock(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &shutdownLeaseStub{} + l := &longOperation{lock: s, timeout: 3 * time.Second} + for range 2 { + require.NoError(t, l.Lock()) + synctest.Wait() + before := s.renewals.Load() + time.Sleep(time.Second) + synctest.Wait() + require.Equal(t, before+1, s.renewals.Load()) + + l.Unlock() + synctest.Wait() + time.Sleep(time.Second) + synctest.Wait() + require.Equal(t, before+1, s.renewals.Load()) + } + }) +} + +func TestLongOperationDoesNotRenewAfterUnlock(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &shutdownLeaseStub{} + l := &longOperation{lock: s, timeout: 3 * time.Second} + for range 100 { + require.NoError(t, l.Lock()) + synctest.Wait() + time.Sleep(time.Second) + l.Unlock() + synctest.Wait() + } + require.Zero(t, s.renewalsAfterUnlock.Load()) + }) +} diff --git a/pkg/lock/simple_mem.go b/pkg/lock/simple_mem.go index 9d0366d6b04..5f9669ccfce 100644 --- a/pkg/lock/simple_mem.go +++ b/pkg/lock/simple_mem.go @@ -20,13 +20,9 @@ func (i *InMemoryLockGetter) ReadWrite(_ prefixer.Prefixer, name string) ErrorRW return lock.(*memLock) } -// LongOperation returns a lock suitable for long operations. It will refresh -// the lock in redis to avoid its automatic expiration. +// LongOperation returns an in-memory lock, which does not expire. func (i *InMemoryLockGetter) LongOperation(db prefixer.Prefixer, name string) ErrorLocker { - return &longOperation{ - lock: i.ReadWrite(db, name).(*memLock), - timeout: LockTimeout, - } + return i.ReadWrite(db, name) } type memLock struct { @@ -35,6 +31,5 @@ type memLock struct { func (ml *memLock) Lock() error { ml.RWMutex.Lock(); return nil } func (ml *memLock) RLock() error { ml.RWMutex.RLock(); return nil } -func (ml *memLock) Extend() {} func (ml *memLock) Unlock() { ml.RWMutex.Unlock() } func (ml *memLock) RUnlock() { ml.RWMutex.RUnlock() }