[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.
This commit is contained in:
Zoltán Papp
2026-08-05 16:14:11 +02:00
parent fab3a42967
commit 47df3c3ef0
3 changed files with 40 additions and 19 deletions
+30 -11
View File
@@ -155,6 +155,9 @@ type Conn struct {
wgWatcher *wg_watcher.WGWatcher wgWatcher *wg_watcher.WGWatcher
wgWatcherWg sync.WaitGroup wgWatcherWg sync.WaitGroup
wgWatcherCancel context.CancelFunc 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 // wgTimeouts counts consecutive WireGuard handshake timeouts without a
// successful handshake in between. Owned by the event loop. // successful handshake in between. Owned by the event loop.
wgTimeouts int wgTimeouts int
@@ -206,7 +209,6 @@ func NewConn(config ConnConfig, services ServiceDependencies) (*Conn, error) {
statusICE: NewAtomicStatus(), statusICE: NewAtomicStatus(),
dumpState: dumpState, dumpState: dumpState,
endpointUpdater: NewEndpointUpdater(connLog, config.WgConfig, config.IsController()), endpointUpdater: NewEndpointUpdater(connLog, config.WgConfig, config.IsController()),
wgWatcher: wg_watcher.NewWGWatcher(connLog, config.WgConfig.WgInterface, config.Key, dumpState),
metricsRecorder: services.MetricsRecorder, metricsRecorder: services.MetricsRecorder,
} }
@@ -449,7 +451,7 @@ func (conn *Conn) handleEvent(ev event) {
case evRelayDialDone: case evRelayDialDone:
conn.handleRelayDialDone() conn.handleRelayDialDone()
case evWGTimeout: case evWGTimeout:
conn.handleWGTimeout() conn.handleWGTimeout(e.gen)
case evWGHandshake: case evWGHandshake:
conn.handleWGHandshakeSuccess(e.when) conn.handleWGHandshakeSuccess(e.when)
case evWGCheckOK: case evWGCheckOK:
@@ -867,11 +869,16 @@ func (conn *Conn) handleRelayDisconnected() {
// handleWGTimeout closes the active connection after a WireGuard handshake // handleWGTimeout closes the active connection after a WireGuard handshake
// timeout so the guard can trigger a reconnection. // timeout so the guard can trigger a reconnection.
func (conn *Conn) handleWGTimeout() { func (conn *Conn) handleWGTimeout(gen uint64) {
if conn.ctx.Err() != nil { if conn.ctx.Err() != nil {
return 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") conn.Log.Warnf("WireGuard handshake timeout detected, closing current connection")
// Close the active connection based on current priority // Close the active connection based on current priority
@@ -945,8 +952,8 @@ func (conn *Conn) onGuardEvent() {
conn.post(evGuardTick{}) conn.post(evGuardTick{})
} }
func (conn *Conn) onWGDisconnected() { func (conn *Conn) onWGDisconnected(gen uint64) {
conn.post(evWGTimeout{}) conn.post(evWGTimeout{gen: gen})
} }
func (conn *Conn) onWGHandshakeSuccess(when time.Time) { 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) { func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) {
if !conn.wgWatcher.PrepareInitialHandshake() { if conn.wgWatcher != nil {
return 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) wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx)
conn.wgWatcher = watcher
conn.wgWatcherCancel = wgWatcherCancel conn.wgWatcherCancel = wgWatcherCancel
conn.wgWatcherWg.Add(1) conn.wgWatcherWg.Add(1)
go func() { go func() {
defer conn.wgWatcherWg.Done() 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() { func (conn *Conn) disableWgWatcherIfNeeded() {
if conn.currentConnPriority == worker.None && conn.wgWatcherCancel != nil { if conn.currentConnPriority != worker.None || conn.wgWatcher == nil {
conn.wgWatcherCancel() return
conn.wgWatcherCancel = nil
} }
conn.wgWatcherCancel()
conn.wgWatcher = nil
conn.wgWatcherCancel = nil
} }
func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) { func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) {
@@ -1144,7 +1161,9 @@ func (conn *Conn) resetEndpoint() {
return return
} }
conn.Log.Infof("reset wg endpoint") conn.Log.Infof("reset wg endpoint")
conn.wgWatcher.Reset() if conn.wgWatcher != nil {
conn.wgWatcher.Reset()
}
if err := conn.endpointUpdater.RemoveEndpointAddress(); err != nil { if err := conn.endpointUpdater.RemoveEndpointAddress(); err != nil {
conn.Log.Warnf("failed to remove endpoint address before update: %v", err) conn.Log.Warnf("failed to remove endpoint address before update: %v", err)
} }
+7 -7
View File
@@ -289,20 +289,20 @@ func TestConn_onWGDisconnected_EscalatesToRosenpassReset(t *testing.T) {
conn := newWGTimeoutTestConn(true, &disconnected) conn := newWGTimeoutTestConn(true, &disconnected)
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
conn.handleWGTimeout() conn.handleWGTimeout(conn.wgWatcherGen)
} }
assert.Empty(t, disconnected, "escalation must not fire below the threshold") 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, assert.Equal(t, []string{conn.config.WgConfig.RemoteKey}, disconnected,
"reaching the threshold must report the peer disconnected once") "reaching the threshold must report the peer disconnected once")
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
conn.handleWGTimeout() conn.handleWGTimeout(conn.wgWatcherGen)
} }
assert.Len(t, disconnected, 1, "escalation must restart counting after firing") 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") 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) conn := newWGTimeoutTestConn(true, &disconnected)
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
conn.handleWGTimeout() conn.handleWGTimeout(conn.wgWatcherGen)
} }
conn.handleWGCheckSuccess() conn.handleWGCheckSuccess()
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ { for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
conn.handleWGTimeout() conn.handleWGTimeout(conn.wgWatcherGen)
} }
assert.Empty(t, disconnected, "handshake success must reset the timeout count") 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) conn := newWGTimeoutTestConn(false, &disconnected)
for i := 0; i < wgTimeoutEscalationThreshold*3; i++ { for i := 0; i < wgTimeoutEscalationThreshold*3; i++ {
conn.handleWGTimeout() conn.handleWGTimeout(conn.wgWatcherGen)
} }
assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections") assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections")
} }
+3 -1
View File
@@ -54,7 +54,9 @@ type evRelayDown struct{}
// successfully or not, so the loop may dispatch a pending offer. // successfully or not, so the loop may dispatch a pending offer.
type evRelayDialDone struct{} type evRelayDialDone struct{}
type evWGTimeout struct{} type evWGTimeout struct {
gen uint64
}
// evWGHandshake reports the first WireGuard handshake of the current watcher run. // evWGHandshake reports the first WireGuard handshake of the current watcher run.
type evWGHandshake struct { type evWGHandshake struct {