diff --git a/client/android/client.go b/client/android/client.go index 05dde7126..eac8fa246 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -26,6 +26,7 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/net" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/formatter" "github.com/netbirdio/netbird/route" @@ -78,6 +79,9 @@ type Client struct { // ConnectClient, which distributes it to every reconnection loop. netState *netstate.State + // sweeper also outlives engine restarts; NotifyNetworkChange sweeps it. + sweeper *netsweep.Sweeper + stateMu sync.RWMutex connectClient *internal.ConnectClient config *profilemanager.Config @@ -149,6 +153,7 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd ctxCancelLock: &sync.Mutex{}, networkChangeListener: networkChangeListener, netState: netstate.New(), + sweeper: netsweep.New(), } } @@ -161,6 +166,14 @@ func (c *Client) SetNetworkAvailable(available bool) { c.recorder.SetNetworkAvailable(available) } +// NotifyNetworkChange cuts the management, signal and relay connections +// after the OS switched networks, so the reconnect loops redial immediately +// on the new one. The engine and the TUN device stay untouched. +func (c *Client) NotifyNetworkChange() { + n := c.sweeper.Sweep() + log.Infof("network change: swept %d connections", n) +} + // Run start the internal client. It is a blocker function func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroidTV bool, dns *DNSList, dnsReadyListener DnsReadyListener, envList *EnvList) error { exportEnvList(envList) @@ -198,7 +211,8 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid } // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetworkState(c.netState)) + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper)) c.setState(cfg, cacheDir, cfgFile, connectClient) // This path runs the interactive SSO flow, so reaching here means the peer // is authenticated again — release the latch Status() reports from. Clear @@ -239,7 +253,8 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetworkState(c.netState)) + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper)) c.setState(cfg, cacheDir, cfgFile, connectClient) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } diff --git a/client/grpc/dialer_generic.go b/client/grpc/dialer_generic.go index 479575996..41aca1808 100644 --- a/client/grpc/dialer_generic.go +++ b/client/grpc/dialer_generic.go @@ -16,28 +16,47 @@ import ( "google.golang.org/grpc" nbnet "github.com/netbirdio/netbird/client/net" + "github.com/netbirdio/netbird/client/netsweep" ) func WithCustomDialer(_ bool, _ string) grpc.DialOption { + return grpc.WithContextDialer(dialContext) +} + +// WithSweeper dials like WithCustomDialer but registers connections and +// dials with the sweeper. Append it after WithCustomDialer: gRPC applies +// dial options in order, so the later context dialer wins. +func WithSweeper(sweeper *netsweep.Sweeper) grpc.DialOption { return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) { - if runtime.GOOS == "linux" { - currentUser, err := user.Current() - if err != nil { - return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err) - } + ctx, releaseDial := sweeper.WrapDialContext(ctx) + defer releaseDial() - // the custom dialer requires root permissions which are not required for use cases run as non-root - if currentUser.Uid != "0" { - log.Debug("Not running as root, using standard dialer") - dialer := &net.Dialer{} - return dialer.DialContext(ctx, "tcp", addr) - } - } - - conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr) + conn, err := dialContext(ctx, addr) if err != nil { - return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err) + return nil, err } - return conn, nil + return sweeper.WrapConn(conn), nil }) } + +func dialContext(ctx context.Context, addr string) (net.Conn, error) { + if runtime.GOOS == "linux" { + currentUser, err := user.Current() + if err != nil { + return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err) + } + + // the custom dialer requires root permissions which are not required for use cases run as non-root + if currentUser.Uid != "0" { + log.Debug("Not running as root, using standard dialer") + dialer := &net.Dialer{} + return dialer.DialContext(ctx, "tcp", addr) + } + } + + conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr) + if err != nil { + return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err) + } + return conn, nil +} diff --git a/client/grpc/dialer_js.go b/client/grpc/dialer_js.go index b89ec3c21..8863756d7 100644 --- a/client/grpc/dialer_js.go +++ b/client/grpc/dialer_js.go @@ -3,6 +3,7 @@ package grpc import ( "google.golang.org/grpc" + "github.com/netbirdio/netbird/client/netsweep" "github.com/netbirdio/netbird/util/wsproxy/client" ) @@ -11,3 +12,8 @@ import ( func WithCustomDialer(tlsEnabled bool, component string) grpc.DialOption { return client.WithWebSocketDialer(tlsEnabled, component) } + +// WithSweeper is a no-op on WASM/JS: there is no network change signal. +func WithSweeper(_ *netsweep.Sweeper) grpc.DialOption { + return grpc.EmptyDialOption{} +} diff --git a/client/internal/connect.go b/client/internal/connect.go index 01f3ed6e5..e45ecca44 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -39,6 +39,7 @@ import ( "github.com/netbirdio/netbird/client/internal/updater/installer" nbnet "github.com/netbirdio/netbird/client/net" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" cProto "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ssh" sshconfig "github.com/netbirdio/netbird/client/ssh/config" @@ -76,6 +77,10 @@ type ConnectClient struct { // availability. Nil (the default) disables gating; mobile platforms // inject it via WithNetworkState. netState *netstate.State + + // sweeper cuts the management, signal and relay connections on network + // change; nil disables it. + sweeper *netsweep.Sweeper } // ConnectClientOption configures optional ConnectClient behavior. @@ -87,6 +92,11 @@ func WithNetworkState(netState *netstate.State) ConnectClientOption { return func(c *ConnectClient) { c.netState = netState } } +// WithSweeper injects the network change sweeper. +func WithSweeper(sweeper *netsweep.Sweeper) ConnectClientOption { + return func(c *ConnectClient) { c.sweeper = sweeper } +} + func NewConnectClient( ctx context.Context, config *profilemanager.Config, @@ -312,7 +322,8 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan }() log.Debugf("connecting to the Management service %s", c.config.ManagementURL.Host) - mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled, mgm.WithNetworkState(c.netState)) + mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled, + mgm.WithNetworkState(c.netState), mgm.WithSweeper(c.sweeper)) if err != nil { // On daemon shutdown / Down() the parent context is cancelled // and the dial fails with "context canceled". Wrapping that @@ -387,7 +398,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan }() // with the global Netbird config in hand connect (just a connection, no stream yet) Signal - signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey, c.netState) + signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey, c.netState, c.sweeper) if err != nil { log.Error(err) return wrapErr(err) @@ -424,7 +435,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan } relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU, - relayClient.WithNetworkState(c.netState)) + relayClient.WithNetworkState(c.netState), relayClient.WithSweeper(c.sweeper)) c.statusRecorder.SetRelayMgr(relayManager) if len(relayURLs) > 0 { if token != nil { @@ -712,7 +723,7 @@ func selectMTU(localMTU uint16, peerMTU int32) uint16 { } // connectToSignal creates Signal Service client and established a connection -func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key, netState *netstate.State) (*signal.GrpcClient, error) { +func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key, netState *netstate.State, sweeper *netsweep.Sweeper) (*signal.GrpcClient, error) { var sigTLSEnabled bool if wtConfig.Signal.Protocol == mgmProto.HostConfig_HTTPS { sigTLSEnabled = true @@ -720,7 +731,8 @@ func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourP sigTLSEnabled = false } - signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled, signal.WithNetworkState(netState)) + signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled, + signal.WithNetworkState(netState), signal.WithSweeper(sweeper)) if err != nil { log.Errorf("error while connecting to the Signal Exchange Service %s: %s", wtConfig.Signal.Uri, err) return nil, gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Signal Service : %s", err) diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 546a0454d..347b71ea5 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -22,6 +22,7 @@ import ( "github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/formatter" "github.com/netbirdio/netbird/route" @@ -79,6 +80,8 @@ type Client struct { // the engine lifecycle. Run injects it into each new ConnectClient, which // distributes it to every reconnection loop. netState *netstate.State + // sweeper also outlives engine restarts; NotifyNetworkChange sweeps it. + sweeper *netsweep.Sweeper // preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked) preloadedConfig *profilemanager.Config @@ -102,6 +105,7 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV networkChangeListener: networkChangeListener, dnsManager: dnsManager, netState: netstate.New(), + sweeper: netsweep.New(), } } @@ -177,7 +181,8 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { c.onHostDnsFn = func([]string) {} cfg.WgIface = interfaceName - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetworkState(c.netState)) + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper)) c.setState(cfg, connectClient) // Persist the latest sync response so DebugBundle can include the network // map. On iOS this is backed by disk to keep it out of the constrained @@ -196,6 +201,14 @@ func (c *Client) SetNetworkAvailable(available bool) { c.recorder.SetNetworkAvailable(available) } +// NotifyNetworkChange cuts the management, signal and relay connections +// after the OS switched networks, so the reconnect loops redial immediately +// on the new one. The engine and the TUN device stay untouched. +func (c *Client) NotifyNetworkChange() { + n := c.sweeper.Sweep() + log.Infof("network change: swept %d connections", n) +} + // Stop the internal client and free the resources func (c *Client) Stop() { c.ctxCancelLock.Lock() diff --git a/shared/management/client/grpc.go b/shared/management/client/grpc.go index 498044560..be81c097e 100644 --- a/shared/management/client/grpc.go +++ b/shared/management/client/grpc.go @@ -22,6 +22,7 @@ import ( nbgrpc "github.com/netbirdio/netbird/client/grpc" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/encryption" "github.com/netbirdio/netbird/shared/management/domain" @@ -67,6 +68,9 @@ type GrpcClient struct { // availability; nil (the default) disables gating. netState *netstate.State + // sweeper cuts the transport connections on network change; nil disables it. + sweeper *netsweep.Sweeper + // syncStreamErr holds the last Sync stream error, or nil while the stream // is established and healthy. GetServerKey succeeds even when the peer // cannot sync (e.g. the server returns "settings not found"), so the @@ -125,16 +129,34 @@ func WithNetworkState(netState *netstate.State) ClientOption { return func(c *GrpcClient) { c.netState = netState } } +// WithSweeper injects the network change sweeper. +func WithSweeper(sweeper *netsweep.Sweeper) ClientOption { + return func(c *GrpcClient) { c.sweeper = sweeper } +} + // NewClient creates a new client to Management service func NewClient(ctx context.Context, addr string, ourPrivateKey wgtypes.Key, tlsEnabled bool, opts ...ClientOption) (*GrpcClient, error) { - var conn *grpc.ClientConn + // Options apply before dialing: the sweeper must wrap the first connection too. + c := &GrpcClient{ + key: ourPrivateKey, + ctx: ctx, + connStateCallbackLock: sync.RWMutex{}, + serverURL: addr, + } + for _, opt := range opts { + opt(c) + } var extraOpts []grpc.DialOption if maxSize := MaxRecvMsgSize(); maxSize > 0 { extraOpts = append(extraOpts, grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(maxSize))) log.Infof("management gRPC max receive message size set to %d bytes", maxSize) } + if c.sweeper != nil { + extraOpts = append(extraOpts, nbgrpc.WithSweeper(c.sweeper)) + } + var conn *grpc.ClientConn operation := func() error { var err error conn, err = nbgrpc.CreateConnection(ctx, addr, tlsEnabled, wsproxy.ManagementComponent, extraOpts...) @@ -150,19 +172,8 @@ func NewClient(ctx context.Context, addr string, ourPrivateKey wgtypes.Key, tlsE return nil, err } - realClient := proto.NewManagementServiceClient(conn) - - c := &GrpcClient{ - key: ourPrivateKey, - realClient: realClient, - ctx: ctx, - conn: conn, - connStateCallbackLock: sync.RWMutex{}, - serverURL: addr, - } - for _, opt := range opts { - opt(c) - } + c.conn = conn + c.realClient = proto.NewManagementServiceClient(conn) return c, nil } diff --git a/shared/relay/client/client.go b/shared/relay/client/client.go index 8d4aa6020..cec8393a0 100644 --- a/shared/relay/client/client.go +++ b/shared/relay/client/client.go @@ -14,6 +14,7 @@ import ( log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/netsweep" auth "github.com/netbirdio/netbird/shared/relay/auth/hmac" "github.com/netbirdio/netbird/shared/relay/client/dialer" netErr "github.com/netbirdio/netbird/shared/relay/client/dialer/net" @@ -184,6 +185,10 @@ type Client struct { // datagram-sized transport is avoided on subsequent connects. Shared via // the manager. transportFallback *transportFallback + + // sweeper cuts the relay connection on network change; the read loop + // reports the disconnect and the guard reconnects. Shared via the manager. + sweeper *netsweep.Sweeper // datagramFallbackTriggered guards a single fallback per connection so a // burst of oversized datagrams triggers one reconnect, not many. datagramFallbackTriggered atomic.Bool @@ -393,6 +398,11 @@ func (c *Client) Close() error { } func (c *Client) connect(ctx context.Context) (*RelayAddr, error) { + // A sweep cancels this context, so a dial started on the old network + // aborts instead of waiting out its handshake timeout. + ctx, releaseDial := c.sweeper.WrapDialContext(ctx) + defer releaseDial() + mode := transportModeFromEnv() dialers := c.getDialers(mode) @@ -417,6 +427,7 @@ func (c *Client) connect(ctx context.Context) (*RelayAddr, error) { return nil, fmt.Errorf("dial via FQDN: %w", err) } } + conn = c.sweeper.WrapConn(conn) c.relayConn = conn c.datagramFallbackTriggered.Store(false) if tc, ok := conn.(transportConn); ok { diff --git a/shared/relay/client/manager.go b/shared/relay/client/manager.go index 9162b8277..80e38ae2d 100644 --- a/shared/relay/client/manager.go +++ b/shared/relay/client/manager.go @@ -13,6 +13,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" relayAuth "github.com/netbirdio/netbird/shared/relay/auth/hmac" ) @@ -72,6 +73,11 @@ func WithNetworkState(netState *netstate.State) ManagerOption { return func(m *Manager) { m.netState = netState } } +// WithSweeper injects the network change sweeper. +func WithSweeper(sweeper *netsweep.Sweeper) ManagerOption { + return func(m *Manager) { m.sweeper = sweeper } +} + // Manager is a manager for the relay client instances. It establishes one persistent connection to the given relay URL // and automatically reconnect to them in case disconnection. // The manager also manage temporary relay connection. If a client wants to communicate with a client on a @@ -100,6 +106,7 @@ type Manager struct { mtu uint16 maxBackoffInterval time.Duration netState *netstate.State + sweeper *netsweep.Sweeper cleanupInterval time.Duration keepUnusedServerTime time.Duration @@ -136,6 +143,7 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin for _, opt := range opts { opt(m) } + m.serverPicker.Sweeper = m.sweeper m.serverPicker.ServerURLs.Store(serverURLs) m.reconnectGuard = NewGuard(m.serverPicker, m.maxBackoffInterval, m.netState) return m @@ -362,6 +370,7 @@ func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string relayClient := NewClientWithServerIP(serverAddress, serverIP, m.tokenStore, m.peerID, m.mtu) relayClient.SetTransportFallback(m.transportFallback) + relayClient.sweeper = m.sweeper err := relayClient.Connect(m.ctx) if err != nil { rt.Lock() diff --git a/shared/relay/client/picker.go b/shared/relay/client/picker.go index bb721e4ad..72789fadc 100644 --- a/shared/relay/client/picker.go +++ b/shared/relay/client/picker.go @@ -9,6 +9,7 @@ import ( log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/netsweep" auth "github.com/netbirdio/netbird/shared/relay/auth/hmac" ) @@ -30,6 +31,7 @@ type ServerPicker struct { MTU uint16 ConnectionTimeout time.Duration TransportFallback *transportFallback + Sweeper *netsweep.Sweeper } func (sp *ServerPicker) PickServer(parentCtx context.Context) (*Client, error) { @@ -73,6 +75,7 @@ func (sp *ServerPicker) startConnection(ctx context.Context, resultChan chan con log.Infof("try to connecting to relay server: %s", url) relayClient := NewClient(url, sp.TokenStore, sp.PeerID, sp.MTU) relayClient.SetTransportFallback(sp.TransportFallback) + relayClient.sweeper = sp.Sweeper err := relayClient.Connect(ctx) resultChan <- connResult{ RelayClient: relayClient, diff --git a/shared/signal/client/grpc.go b/shared/signal/client/grpc.go index 7eddf6045..d17c7c958 100644 --- a/shared/signal/client/grpc.go +++ b/shared/signal/client/grpc.go @@ -20,6 +20,7 @@ import ( nbgrpc "github.com/netbirdio/netbird/client/grpc" "github.com/netbirdio/netbird/client/netstate" + "github.com/netbirdio/netbird/client/netsweep" "github.com/netbirdio/netbird/encryption" "github.com/netbirdio/netbird/shared/management/client" "github.com/netbirdio/netbird/shared/signal/proto" @@ -70,6 +71,9 @@ type GrpcClient struct { // availability; nil (the default) disables gating. netState *netstate.State + // sweeper cuts the transport connections on network change; nil disables it. + sweeper *netsweep.Sweeper + onReconnectedListenerFn func() decryptionWorker *Worker @@ -93,7 +97,6 @@ type GrpcClient struct { watchdogWg sync.WaitGroup } -// NewClient creates a new Signal client // ClientOption configures optional GrpcClient behavior. type ClientOption func(*GrpcClient) @@ -103,12 +106,34 @@ func WithNetworkState(netState *netstate.State) ClientOption { return func(c *GrpcClient) { c.netState = netState } } -func NewClient(ctx context.Context, addr string, key wgtypes.Key, tlsEnabled bool, opts ...ClientOption) (*GrpcClient, error) { - var conn *grpc.ClientConn +// WithSweeper injects the network change sweeper. +func WithSweeper(sweeper *netsweep.Sweeper) ClientOption { + return func(c *GrpcClient) { c.sweeper = sweeper } +} +// NewClient creates a new Signal client +func NewClient(ctx context.Context, addr string, key wgtypes.Key, tlsEnabled bool, opts ...ClientOption) (*GrpcClient, error) { + // Options apply before dialing: the sweeper must wrap the first connection too. + c := &GrpcClient{ + ctx: ctx, + key: key, + mux: sync.Mutex{}, + status: StreamDisconnected, + connStateCallbackLock: sync.RWMutex{}, + } + for _, opt := range opts { + opt(c) + } + + var extraOpts []grpc.DialOption + if c.sweeper != nil { + extraOpts = append(extraOpts, nbgrpc.WithSweeper(c.sweeper)) + } + + var conn *grpc.ClientConn operation := func() error { var err error - conn, err = nbgrpc.CreateConnection(ctx, addr, tlsEnabled, wsproxy.SignalComponent) + conn, err = nbgrpc.CreateConnection(ctx, addr, tlsEnabled, wsproxy.SignalComponent, extraOpts...) if err != nil { return fmt.Errorf("create connection: %w", err) } @@ -123,18 +148,8 @@ func NewClient(ctx context.Context, addr string, key wgtypes.Key, tlsEnabled boo log.Debugf("connected to Signal Service: %v", conn.Target()) - c := &GrpcClient{ - realClient: proto.NewSignalExchangeClient(conn), - ctx: ctx, - signalConn: conn, - key: key, - mux: sync.Mutex{}, - status: StreamDisconnected, - connStateCallbackLock: sync.RWMutex{}, - } - for _, opt := range opts { - opt(c) - } + c.signalConn = conn + c.realClient = proto.NewSignalExchangeClient(conn) return c, nil }