mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-10 03:56:01 -04:00
[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:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user