From 47df3c3ef0eb2f36120c1e4c8b98ebfea188db64 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Zolt=C3=A1n=20Papp?= Date: Tue, 14 Jul 2026 11:16:15 +0200 Subject: [PATCH] [client] Adapt WGWatcher to per-instance model after #6664 rebase The rebase carried #6664's WGWatcher changes into our wg_watcher package (single-shot, no enabled flag). Adapt conn.go to match: create a fresh watcher per connection attempt in enableWgWatcherIfNeeded, drop it in disableWgWatcherIfNeeded, nil-guard resetEndpoint. Guard stale WG timeouts on the event loop instead of #6664's conn.mu recheck: onWGDisconnected only checked watcherCtx on the watcher goroutine, racing the loop that cancels it and processes the timeout. A loop-owned wgWatcherGen tags evWGTimeout; handleWGTimeout drops events from a superseded generation, so the check and the teardown happen atomically on the single loop. --- client/internal/peer/conn.go | 41 ++++++++++++++++++++++--------- client/internal/peer/conn_test.go | 14 +++++------ client/internal/peer/event.go | 4 ++- 3 files changed, 40 insertions(+), 19 deletions(-) diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index d8f436f32..ee0a42299 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -155,6 +155,9 @@ type Conn struct { wgWatcher *wg_watcher.WGWatcher wgWatcherWg sync.WaitGroup wgWatcherCancel context.CancelFunc + // wgWatcherGen identifies the current watcher generation; a WG timeout event + // carrying an older generation is dropped by the loop. Owned by the event loop. + wgWatcherGen uint64 // wgTimeouts counts consecutive WireGuard handshake timeouts without a // successful handshake in between. Owned by the event loop. wgTimeouts int @@ -206,7 +209,6 @@ func NewConn(config ConnConfig, services ServiceDependencies) (*Conn, error) { statusICE: NewAtomicStatus(), dumpState: dumpState, endpointUpdater: NewEndpointUpdater(connLog, config.WgConfig, config.IsController()), - wgWatcher: wg_watcher.NewWGWatcher(connLog, config.WgConfig.WgInterface, config.Key, dumpState), metricsRecorder: services.MetricsRecorder, } @@ -449,7 +451,7 @@ func (conn *Conn) handleEvent(ev event) { case evRelayDialDone: conn.handleRelayDialDone() case evWGTimeout: - conn.handleWGTimeout() + conn.handleWGTimeout(e.gen) case evWGHandshake: conn.handleWGHandshakeSuccess(e.when) case evWGCheckOK: @@ -867,11 +869,16 @@ func (conn *Conn) handleRelayDisconnected() { // handleWGTimeout closes the active connection after a WireGuard handshake // timeout so the guard can trigger a reconnection. -func (conn *Conn) handleWGTimeout() { +func (conn *Conn) handleWGTimeout(gen uint64) { if conn.ctx.Err() != nil { return } + if gen != conn.wgWatcherGen { + conn.Log.Debugf("ignore WG timeout from superseded watcher generation %d (current %d)", gen, conn.wgWatcherGen) + return + } + conn.Log.Warnf("WireGuard handshake timeout detected, closing current connection") // Close the active connection based on current priority @@ -945,8 +952,8 @@ func (conn *Conn) onGuardEvent() { conn.post(evGuardTick{}) } -func (conn *Conn) onWGDisconnected() { - conn.post(evWGTimeout{}) +func (conn *Conn) onWGDisconnected(gen uint64) { + conn.post(evWGTimeout{gen: gen}) } func (conn *Conn) onWGHandshakeSuccess(when time.Time) { @@ -1104,24 +1111,34 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { } func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) { - if !conn.wgWatcher.PrepareInitialHandshake() { + if conn.wgWatcher != nil { return } + conn.wgWatcherGen++ + gen := conn.wgWatcherGen + + watcher := wg_watcher.NewWGWatcher(conn.Log, conn.config.WgConfig.WgInterface, conn.config.Key, conn.dumpState) + watcher.PrepareInitialHandshake() + wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx) + conn.wgWatcher = watcher conn.wgWatcherCancel = wgWatcherCancel conn.wgWatcherWg.Add(1) go func() { defer conn.wgWatcherWg.Done() - conn.wgWatcher.EnableWgWatcher(wgWatcherCtx, enabledTime, conn.onWGDisconnected, conn.onWGHandshakeSuccess, conn.onWGCheckSuccess) + onDisconnected := func() { conn.onWGDisconnected(gen) } + watcher.EnableWgWatcher(wgWatcherCtx, enabledTime, onDisconnected, conn.onWGHandshakeSuccess, conn.onWGCheckSuccess) }() } func (conn *Conn) disableWgWatcherIfNeeded() { - if conn.currentConnPriority == worker.None && conn.wgWatcherCancel != nil { - conn.wgWatcherCancel() - conn.wgWatcherCancel = nil + if conn.currentConnPriority != worker.None || conn.wgWatcher == nil { + return } + conn.wgWatcherCancel() + conn.wgWatcher = nil + conn.wgWatcherCancel = nil } func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) { @@ -1144,7 +1161,9 @@ func (conn *Conn) resetEndpoint() { return } conn.Log.Infof("reset wg endpoint") - conn.wgWatcher.Reset() + if conn.wgWatcher != nil { + conn.wgWatcher.Reset() + } if err := conn.endpointUpdater.RemoveEndpointAddress(); err != nil { conn.Log.Warnf("failed to remove endpoint address before update: %v", err) } diff --git a/client/internal/peer/conn_test.go b/client/internal/peer/conn_test.go index 39bbf175f..5ed94a1be 100644 --- a/client/internal/peer/conn_test.go +++ b/client/internal/peer/conn_test.go @@ -289,20 +289,20 @@ func TestConn_onWGDisconnected_EscalatesToRosenpassReset(t *testing.T) { conn := newWGTimeoutTestConn(true, &disconnected) for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) } assert.Empty(t, disconnected, "escalation must not fire below the threshold") - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) assert.Equal(t, []string{conn.config.WgConfig.RemoteKey}, disconnected, "reaching the threshold must report the peer disconnected once") for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) } assert.Len(t, disconnected, 1, "escalation must restart counting after firing") - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) assert.Len(t, disconnected, 2, "continued timeouts must escalate again") } @@ -314,12 +314,12 @@ func TestConn_onWGDisconnected_CheckSuccessResetsEscalation(t *testing.T) { conn := newWGTimeoutTestConn(true, &disconnected) for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) } conn.handleWGCheckSuccess() for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) } assert.Empty(t, disconnected, "handshake success must reset the timeout count") } @@ -332,7 +332,7 @@ func TestConn_onWGDisconnected_NoEscalationWithoutRosenpass(t *testing.T) { conn := newWGTimeoutTestConn(false, &disconnected) for i := 0; i < wgTimeoutEscalationThreshold*3; i++ { - conn.handleWGTimeout() + conn.handleWGTimeout(conn.wgWatcherGen) } assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections") } diff --git a/client/internal/peer/event.go b/client/internal/peer/event.go index 8dc1888b2..6bfc5f8d4 100644 --- a/client/internal/peer/event.go +++ b/client/internal/peer/event.go @@ -54,7 +54,9 @@ type evRelayDown struct{} // successfully or not, so the loop may dispatch a pending offer. type evRelayDialDone struct{} -type evWGTimeout struct{} +type evWGTimeout struct { + gen uint64 +} // evWGHandshake reports the first WireGuard handshake of the current watcher run. type evWGHandshake struct {