From f050e052940a30c533f1affe9131ca88cd784470 Mon Sep 17 00:00:00 2001 From: andig Date: Mon, 21 Mar 2022 13:09:07 +0100 Subject: [PATCH] Fix a race condition in waiter --- util/waiter.go | 13 ++++++++++--- util/waiter_test.go | 14 ++++++++++---- 2 files changed, 20 insertions(+), 7 deletions(-) diff --git a/util/waiter.go b/util/waiter.go index 61047dd05..d842716ab 100644 --- a/util/waiter.go +++ b/util/waiter.go @@ -27,7 +27,7 @@ func NewWaiter(timeout time.Duration, logInitialWait func()) *Waiter { } // Update is called when client has received data. Update resets the timeout counter. -// It is client responsibility to ensure that the waiter is not locked when Update is called. +// Waiter MUST be locked when calling Update. func (p *Waiter) Update() { p.updated = time.Now() p.cond.Broadcast() @@ -43,18 +43,25 @@ func (p *Waiter) Overdue() time.Duration { go func() { defer close(c) + p.Lock() // establish lock once go routine has started for p.updated.IsZero() { p.cond.Wait() } }() + // release lock so external updates can occur + p.Unlock() + select { case <-c: // initial value received, lock established case <-time.After(waitInitialTimeout): - p.Update() // unblock the sync.Cond + // establish lock per contract of `Update()`` + p.Lock() + p.Update() // unblock the sync.Cond + p.Unlock() <-c // wait for goroutine, re-establish lock - p.updated = time.Time{} // reset updated to initial value missing + p.updated = time.Time{} // reset to "initial value missing" return waitInitialTimeout } } diff --git a/util/waiter_test.go b/util/waiter_test.go index 182b54c59..7ff903b9a 100644 --- a/util/waiter_test.go +++ b/util/waiter_test.go @@ -19,15 +19,16 @@ func TestWaiterInitialUpdateInTime(t *testing.T) { go func() { time.Sleep(testTimeout / 2) + w.Lock() w.Update() + w.Unlock() }() w.Lock() - defer w.Unlock() - if elapsed := w.Overdue(); elapsed != 0 { t.Errorf("expected %v, got %v", 0, elapsed) } + w.Unlock() } } @@ -36,21 +37,24 @@ func TestWaiterInitialUpdateNotReceived(t *testing.T) { w := NewWaiter(timeout, func() {}) w.Lock() - defer w.Unlock() - if elapsed := w.Overdue(); elapsed != waitInitialTimeout { t.Errorf("expected %v, got %v", waitInitialTimeout, elapsed) } + w.Unlock() } } func TestWaiterUpdateInTime(t *testing.T) { w := NewWaiter(testTimeout, func() {}) + w.Lock() w.Update() + w.Unlock() go func() { time.Sleep(testTimeout / 2) + w.Lock() w.Update() + w.Unlock() }() w.Lock() @@ -63,7 +67,9 @@ func TestWaiterUpdateInTime(t *testing.T) { func TestWaiterUpdateNotReceived(t *testing.T) { w := NewWaiter(testTimeout, func() {}) + w.Lock() w.Update() + w.Unlock() time.Sleep(2 * testTimeout)