diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index 09a4e8b02..6733acf92 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -841,6 +841,8 @@ func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) { return } + conn.endpointUpdater.EnableKeepAlive() + watcher := NewWGWatcher(conn.Log, conn.config.WgConfig.WgInterface, conn.config.Key, conn.dumpState) watcher.PrepareInitialHandshake() @@ -888,6 +890,7 @@ func (conn *Conn) resetEndpoint() { return } conn.Log.Infof("reset wg endpoint") + conn.endpointUpdater.EnableKeepAlive() if conn.wgWatcher != nil { conn.wgWatcher.Reset() } @@ -945,7 +948,12 @@ func (conn *Conn) onWGHandshakeSuccess(when time.Time) { func (conn *Conn) onWGCheckSuccess() { conn.mu.Lock() conn.wgTimeouts = 0 + presharedKey := conn.presharedKey(conn.rosenpassRemoteKey) conn.mu.Unlock() + + if err := conn.endpointUpdater.DisableKeepAlive(presharedKey); err != nil { + conn.Log.Warnf("failed to disable WireGuard keepalive: %v", err) + } } // recordConnectionMetrics records connection stage timestamps as metrics diff --git a/client/internal/peer/conn_test.go b/client/internal/peer/conn_test.go index 49979ea83..7ed0362b9 100644 --- a/client/internal/peer/conn_test.go +++ b/client/internal/peer/conn_test.go @@ -316,11 +316,16 @@ func newWGTimeoutTestConn(rosenpassEnabled bool, disconnected *[]string) *Conn { cfg.RosenpassConfig = RosenpassConfig{PubKey: []byte("dummykey")} } + connLog := log.WithField("peer", cfg.Key) + endpointUpdater := NewEndpointUpdater(connLog, cfg.WgConfig, false) + endpointUpdater.keepAlive = 0 + conn := &Conn{ - ctx: context.Background(), - config: cfg, - Log: log.WithField("peer", cfg.Key), - metricsStages: &MetricsStages{}, + ctx: context.Background(), + config: cfg, + Log: connLog, + metricsStages: &MetricsStages{}, + endpointUpdater: endpointUpdater, } conn.SetOnDisconnected(func(remotePeer string) { *disconnected = append(*disconnected, remotePeer) diff --git a/client/internal/peer/endpoint.go b/client/internal/peer/endpoint.go index 9ba1efb6e..559012de3 100644 --- a/client/internal/peer/endpoint.go +++ b/client/internal/peer/endpoint.go @@ -20,10 +20,11 @@ type EndpointUpdater struct { wgConfig WgConfig initiator bool - // mu protects cancelFunc + // mu protects cancelFunc and keepAlive mu sync.Mutex cancelFunc func() updateWg sync.WaitGroup + keepAlive time.Duration } func NewEndpointUpdater(log *logrus.Entry, wgConfig WgConfig, initiator bool) *EndpointUpdater { @@ -31,6 +32,7 @@ func NewEndpointUpdater(log *logrus.Entry, wgConfig WgConfig, initiator bool) *E log: log, wgConfig: wgConfig, initiator: initiator, + keepAlive: defaultWgKeepAlive, } } @@ -73,6 +75,31 @@ func (e *EndpointUpdater) RemoveEndpointAddress() error { return e.wgConfig.WgInterface.RemoveEndpointAddress(e.wgConfig.RemoteKey) } +func (e *EndpointUpdater) DisableKeepAlive(presharedKey *wgtypes.Key) error { + if !isWgKeepAliveDisabled() { + return nil + } + + e.mu.Lock() + defer e.mu.Unlock() + + if e.keepAlive == 0 { + return nil + } + + e.waitForCloseTheDelayedUpdate() + e.keepAlive = 0 + e.log.Debugf("disable WireGuard persistent keepalive") + return e.updateWireGuardPeer(nil, presharedKey) +} + +func (e *EndpointUpdater) EnableKeepAlive() { + e.mu.Lock() + defer e.mu.Unlock() + + e.keepAlive = defaultWgKeepAlive +} + func (e *EndpointUpdater) configureAsInitiator(addr *net.UDPAddr, presharedKey *wgtypes.Key) error { if err := e.updateWireGuardPeer(addr, presharedKey); err != nil { return err @@ -127,7 +154,7 @@ func (e *EndpointUpdater) updateWireGuardPeer(endpoint *net.UDPAddr, presharedKe return e.wgConfig.WgInterface.UpdatePeer( e.wgConfig.RemoteKey, e.wgConfig.AllowedIps, - defaultWgKeepAlive, + e.keepAlive, endpoint, presharedKey, ) diff --git a/client/internal/peer/env.go b/client/internal/peer/env.go index ed6a3af53..b48fdf9fb 100644 --- a/client/internal/peer/env.go +++ b/client/internal/peer/env.go @@ -7,8 +7,9 @@ import ( ) const ( - EnvKeyNBForceRelay = "NB_FORCE_RELAY" - EnvKeyNBHomeRelayServers = "NB_HOME_RELAY_SERVERS" + EnvKeyNBForceRelay = "NB_FORCE_RELAY" + EnvKeyNBHomeRelayServers = "NB_HOME_RELAY_SERVERS" + EnvKeyNBDisableWgKeepAlive = "NB_DISABLE_WG_KEEP_ALIVE" ) func IsForceRelayed() bool { @@ -18,6 +19,10 @@ func IsForceRelayed() bool { return strings.EqualFold(os.Getenv(EnvKeyNBForceRelay), "true") } +func isWgKeepAliveDisabled() bool { + return strings.EqualFold(os.Getenv(EnvKeyNBDisableWgKeepAlive), "true") +} + // OverrideRelayURLs returns the relay server URL list set in // NB_HOME_RELAY_SERVERS (comma-separated) and a boolean indicating whether // the override is active. When the env var is unset, the boolean is false diff --git a/client/internal/peer/wg_watcher.go b/client/internal/peer/wg_watcher.go index 39e3d3264..0a07d65a6 100644 --- a/client/internal/peer/wg_watcher.go +++ b/client/internal/peer/wg_watcher.go @@ -109,10 +109,15 @@ func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn } lastHandshake = *handshake + w.stateDump.WGcheckSuccess() + + if isWgKeepAliveDisabled() { + w.log.Debugf("WireGuard watcher waiting for peer reset") + continue + } resetTime := time.Until(handshake.Add(checkPeriod)) timer.Reset(resetTime) - w.stateDump.WGcheckSuccess() w.log.Debugf("WireGuard watcher reset timer: %v", resetTime) case <-w.resetCh: