mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-11 16:09:07 +02: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
|
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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
Reference in New Issue
Block a user