Fix race in port-forward (#418)
This commit is contained in:
parent
5691835f67
commit
e20a16ee8e
|
|
@ -37,3 +37,20 @@ func (cmap *ConcurrentMap[T]) Delete(key string) {
|
||||||
|
|
||||||
delete(cmap.nonConcurrentMap, key)
|
delete(cmap.nonConcurrentMap, key)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cmap *ConcurrentMap[T]) DeleteIf(key string, predicate func(T) bool) bool {
|
||||||
|
cmap.mtx.Lock()
|
||||||
|
defer cmap.mtx.Unlock()
|
||||||
|
|
||||||
|
value, ok := cmap.nonConcurrentMap[key]
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !predicate(value) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(cmap.nonConcurrentMap, key)
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,34 @@
|
||||||
|
package concurrentmap
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDeleteIf(t *testing.T) {
|
||||||
|
cmap := NewConcurrentMap[int]()
|
||||||
|
cmap.Store("a", 1)
|
||||||
|
|
||||||
|
deleted := cmap.DeleteIf("a", func(value int) bool {
|
||||||
|
return value == 1
|
||||||
|
})
|
||||||
|
require.True(t, deleted)
|
||||||
|
|
||||||
|
_, ok := cmap.Load("a")
|
||||||
|
require.False(t, ok)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDeleteIfPredicateFalse(t *testing.T) {
|
||||||
|
cmap := NewConcurrentMap[int]()
|
||||||
|
cmap.Store("a", 1)
|
||||||
|
|
||||||
|
deleted := cmap.DeleteIf("a", func(value int) bool {
|
||||||
|
return value == 2
|
||||||
|
})
|
||||||
|
require.False(t, deleted)
|
||||||
|
|
||||||
|
value, ok := cmap.Load("a")
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, 1, value)
|
||||||
|
}
|
||||||
|
|
@ -31,16 +31,19 @@ func NewNotifier(logger *zap.SugaredLogger) *Notifier {
|
||||||
func (watcher *Notifier) Register(ctx context.Context, worker string) (chan *rpc.WatchInstruction, func()) {
|
func (watcher *Notifier) Register(ctx context.Context, worker string) (chan *rpc.WatchInstruction, func()) {
|
||||||
subCtx, cancel := context.WithCancel(ctx)
|
subCtx, cancel := context.WithCancel(ctx)
|
||||||
workerCh := make(chan *rpc.WatchInstruction)
|
workerCh := make(chan *rpc.WatchInstruction)
|
||||||
|
slot := &WorkerSlot{
|
||||||
watcher.logger.Debugf("registering worker %s", worker)
|
|
||||||
watcher.workers.Store(worker, &WorkerSlot{
|
|
||||||
ctx: subCtx,
|
ctx: subCtx,
|
||||||
ch: workerCh,
|
ch: workerCh,
|
||||||
})
|
}
|
||||||
|
|
||||||
|
watcher.logger.Debugf("registering worker %s", worker)
|
||||||
|
watcher.workers.Store(worker, slot)
|
||||||
|
|
||||||
return workerCh, func() {
|
return workerCh, func() {
|
||||||
watcher.logger.Debugf("deleting worker %s", worker)
|
watcher.logger.Debugf("deleting worker %s", worker)
|
||||||
watcher.workers.Delete(worker)
|
watcher.workers.DeleteIf(worker, func(current *WorkerSlot) bool {
|
||||||
|
return current == slot
|
||||||
|
})
|
||||||
cancel()
|
cancel()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -47,3 +47,36 @@ func TestNotifier(t *testing.T) {
|
||||||
|
|
||||||
wg.Wait()
|
wg.Wait()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestNotifierReRegisterKeepsNewestSlot(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
watcher := notifier.NewNotifier(zap.NewNop().Sugar())
|
||||||
|
|
||||||
|
const worker = "worker-a"
|
||||||
|
|
||||||
|
_, staleCancel := watcher.Register(ctx, worker)
|
||||||
|
newestCh, newestCancel := watcher.Register(ctx, worker)
|
||||||
|
defer newestCancel()
|
||||||
|
|
||||||
|
// Simulate stale connection cleanup arriving after the worker has already re-registered.
|
||||||
|
staleCancel()
|
||||||
|
|
||||||
|
notifyCtx, notifyCancel := context.WithTimeout(ctx, 300*time.Millisecond)
|
||||||
|
defer notifyCancel()
|
||||||
|
|
||||||
|
notifyErrCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
notifyErrCh <- watcher.Notify(notifyCtx, worker, nil)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-newestCh:
|
||||||
|
case err := <-notifyErrCh:
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Fatal("notify returned before delivering message to newest registration")
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("timed out waiting for notify delivery")
|
||||||
|
}
|
||||||
|
|
||||||
|
require.NoError(t, <-notifyErrCh)
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue