Revert "[client] Fix race between WG watcher initial handshake read and endpoint creation (#6626)"

Restore the pre-#6626 WG watcher shape
(single EnableWgWatcher reading the baseline inside the goroutine, no
PrepareInitialHandshake, no ctx-recheck before onDisconnected) to test whether
#6626 causes the missed handshake progression and subsequent reconnect loop.

Kept the WGW-DIAG logs.

This reverts commit 06839a4731.
This commit is contained in:
riccardom
2026-07-07 08:41:05 +02:00
parent 2f3bf5bb16
commit dd20a4076b
3 changed files with 29 additions and 49 deletions
+8 -10
View File
@@ -803,17 +803,15 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
} }
func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) { func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) {
if !conn.wgWatcher.PrepareInitialHandshake() { if !conn.wgWatcher.IsEnabled() {
return wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx)
conn.wgWatcherCancel = wgWatcherCancel
conn.wgWatcherWg.Add(1)
go func() {
defer conn.wgWatcherWg.Done()
conn.wgWatcher.EnableWgWatcher(wgWatcherCtx, enabledTime, conn.onWGDisconnected, conn.onWGHandshakeSuccess)
}()
} }
wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx)
conn.wgWatcherCancel = wgWatcherCancel
conn.wgWatcherWg.Add(1)
go func() {
defer conn.wgWatcherWg.Done()
conn.wgWatcher.EnableWgWatcher(wgWatcherCtx, enabledTime, conn.onWGDisconnected, conn.onWGHandshakeSuccess)
}()
} }
func (conn *Conn) disableWgWatcherIfNeeded() { func (conn *Conn) disableWgWatcherIfNeeded() {
+21 -29
View File
@@ -31,9 +31,7 @@ type WGWatcher struct {
stateDump *stateDump stateDump *stateDump
enabled bool enabled bool
muEnabled sync.Mutex muEnabled sync.RWMutex
// initialHandshake is not thread-safe; never call PrepareInitialHandshake and EnableWgWatcher concurrently.
initialHandshake time.Time
resetCh chan struct{} resetCh chan struct{}
} }
@@ -48,39 +46,40 @@ func NewWGWatcher(log *log.Entry, wgIfaceStater WGInterfaceStater, peerKey strin
} }
} }
// PrepareInitialHandshake reserves the watcher and reads the peer's current WireGuard // EnableWgWatcher starts the WireGuard watcher. If it is already enabled, it will return immediately and do nothing.
// handshake time. It must be called before the peer is (re)configured on the WireGuard // The watcher runs until ctx is cancelled. Caller is responsible for context lifecycle management.
// interface, so the captured baseline reflects the state prior to this connection attempt // NOTE: reverted to the pre-#6626 shape for bisecting the NHN issue.
// instead of racing with that configuration. Returns ok=false if the watcher is already func (w *WGWatcher) EnableWgWatcher(ctx context.Context, enabledTime time.Time, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time)) {
// running, in which case EnableWgWatcher must not be called.
func (w *WGWatcher) PrepareInitialHandshake() (ok bool) {
w.muEnabled.Lock() w.muEnabled.Lock()
if w.enabled { if w.enabled {
w.muEnabled.Unlock() w.muEnabled.Unlock()
return false return
} }
w.log.Debugf("enable WireGuard watcher") w.log.Debugf("enable WireGuard watcher")
w.enabled = true w.enabled = true
w.muEnabled.Unlock() w.muEnabled.Unlock()
handshake, _ := w.wgState() initialHandshake, err := w.wgState()
w.initialHandshake = handshake if err != nil {
w.log.Warnf("PSK-DIAG: watcher baseline handshake=%v (zero=%v)", handshake, handshake.IsZero()) w.log.Warnf("failed to read initial wg stats: %v", err)
return true }
} w.log.Warnf("PSK-DIAG: watcher baseline handshake=%v (zero=%v) [pre-6626 revert]", initialHandshake, initialHandshake.IsZero())
// EnableWgWatcher runs the WireGuard watcher loop using the handshake baseline captured by w.periodicHandshakeCheck(ctx, onDisconnectedFn, onHandshakeSuccessFn, enabledTime, initialHandshake)
// PrepareInitialHandshake. The watcher runs until ctx is cancelled. Caller is responsible
// for context lifecycle management.
func (w *WGWatcher) EnableWgWatcher(ctx context.Context, enabledTime time.Time, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time)) {
w.periodicHandshakeCheck(ctx, onDisconnectedFn, onHandshakeSuccessFn, enabledTime, w.initialHandshake)
w.muEnabled.Lock() w.muEnabled.Lock()
w.enabled = false w.enabled = false
w.muEnabled.Unlock() w.muEnabled.Unlock()
} }
// IsEnabled returns true if the WireGuard watcher is currently enabled
func (w *WGWatcher) IsEnabled() bool {
w.muEnabled.RLock()
defer w.muEnabled.RUnlock()
return w.enabled
}
// Reset signals the watcher that the WireGuard peer has been reset and a new // Reset signals the watcher that the WireGuard peer has been reset and a new
// handshake is expected. This restarts the handshake timeout from scratch. // handshake is expected. This restarts the handshake timeout from scratch.
func (w *WGWatcher) Reset() { func (w *WGWatcher) Reset() {
@@ -106,21 +105,14 @@ func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn
w.log.Warnf("WGW-DIAG: check fire id=%d lastHandshake=%v", enabledTime.UnixNano(), lastHandshake) w.log.Warnf("WGW-DIAG: check fire id=%d lastHandshake=%v", enabledTime.UnixNano(), lastHandshake)
handshake, ok := w.handshakeCheck(lastHandshake) handshake, ok := w.handshakeCheck(lastHandshake)
if !ok { if !ok {
// #6626 race check: a superseded/cancelled watcher must not tear w.log.Warnf("WGW-DIAG: check failed -> firing onDisconnected (TEARDOWN, pre-6626 no ctx-recheck) id=%d", enabledTime.UnixNano())
// down a now-healthy connection. Log which branch we take so a
// bundle shows whether teardowns fire on live vs cancelled ctx.
if ctx.Err() != nil {
w.log.Warnf("WGW-DIAG: check failed but ctx cancelled -> standing down, NO teardown id=%d", enabledTime.UnixNano())
return
}
w.log.Warnf("WGW-DIAG: check failed, ctx live -> firing onDisconnected (TEARDOWN) id=%d", enabledTime.UnixNano())
onDisconnectedFn() onDisconnectedFn()
return return
} }
if lastHandshake.IsZero() { if lastHandshake.IsZero() {
elapsed := calcElapsed(enabledTime, *handshake) elapsed := calcElapsed(enabledTime, *handshake)
w.log.Infof("first wg handshake detected within: %.2fsec, (%s)", elapsed, handshake) w.log.Infof("first wg handshake detected within: %.2fsec, (%s)", elapsed, handshake)
if onHandshakeSuccessFn != nil && ctx.Err() == nil { if onHandshakeSuccessFn != nil {
onHandshakeSuccessFn(*handshake) onHandshakeSuccessFn(*handshake)
} }
} }
-10
View File
@@ -7,7 +7,6 @@ import (
"time" "time"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/iface/configurer" "github.com/netbirdio/netbird/client/iface/configurer"
) )
@@ -35,9 +34,6 @@ func TestWGWatcher_EnableWgWatcher(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
defer cancel() defer cancel()
ok := watcher.PrepareInitialHandshake()
require.True(t, ok, "watcher should not be enabled yet")
onDisconnected := make(chan struct{}, 1) onDisconnected := make(chan struct{}, 1)
go watcher.EnableWgWatcher(ctx, time.Now(), func() { go watcher.EnableWgWatcher(ctx, time.Now(), func() {
mlog.Infof("onDisconnectedFn") mlog.Infof("onDisconnectedFn")
@@ -66,9 +62,6 @@ func TestWGWatcher_ReEnable(t *testing.T) {
watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{})) watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{}))
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
ok := watcher.PrepareInitialHandshake()
require.True(t, ok, "watcher should not be enabled yet")
wg := &sync.WaitGroup{} wg := &sync.WaitGroup{}
wg.Add(1) wg.Add(1)
go func() { go func() {
@@ -83,9 +76,6 @@ func TestWGWatcher_ReEnable(t *testing.T) {
ctx, cancel = context.WithCancel(context.Background()) ctx, cancel = context.WithCancel(context.Background())
defer cancel() defer cancel()
ok = watcher.PrepareInitialHandshake()
require.True(t, ok, "watcher should be re-enabled after the previous run stopped")
onDisconnected := make(chan struct{}, 1) onDisconnected := make(chan struct{}, 1)
go watcher.EnableWgWatcher(ctx, time.Now(), func() { go watcher.EnableWgWatcher(ctx, time.Now(), func() {
onDisconnected <- struct{}{} onDisconnected <- struct{}{}