diff --git a/client/android/client.go b/client/android/client.go index bd34cc92a..71bbe4380 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -302,12 +302,13 @@ 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. +// NotifyNetworkChange marks the management, signal and relay connections +// stale after the OS switched networks and schedules a sweep that cuts +// whatever has not redialed on the new network by then. The engine and the +// TUN device stay untouched. func (c *Client) NotifyNetworkChange() { - n := c.sweeper.Sweep() - log.Infof("network change: swept %d connections", n) + c.sweeper.MarkNetworkChange() + log.Infof("network change: connections marked stale") } // DebugBundle generates a debug bundle, uploads it, and returns the upload key. diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 1c3ad7705..f92f085ab 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -209,12 +209,13 @@ 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. +// NotifyNetworkChange marks the management, signal and relay connections +// stale after the OS switched networks and schedules a sweep that cuts +// whatever has not redialed on the new network by then. The engine and the +// TUN device stay untouched. func (c *Client) NotifyNetworkChange() { - n := c.sweeper.Sweep() - log.Infof("network change: swept %d connections", n) + c.sweeper.MarkNetworkChange() + log.Infof("network change: connections marked stale") } // Stop the internal client and free the resources diff --git a/client/netsweep/netsweep.go b/client/netsweep/netsweep.go index 41fe7fac8..ac8cd2d05 100644 --- a/client/netsweep/netsweep.go +++ b/client/netsweep/netsweep.go @@ -11,10 +11,21 @@ import ( "errors" "net" "sync" + "time" log "github.com/sirupsen/logrus" ) +// DefaultSweepDelay absorbs network flapping while the OS settles on a +// default network before the stale registrations are cut. +const DefaultSweepDelay = 500 * time.Millisecond + +// Config customizes a Sweeper. The zero value applies the defaults. +type Config struct { + // SweepDelay overrides DefaultSweepDelay when positive. + SweepDelay time.Duration +} + // ErrSwept reports that a dial finished after a network change swept its // registration. The connection is already closed; the caller must treat it // as a failed dial and redial on the new network. @@ -24,6 +35,11 @@ var ErrSwept = errors.New("netsweep: connection swept by network change") // draw from the same counter, so an id is unique across both registries. type sweepID uint64 +type connEntry struct { + conn net.Conn + gen uint64 +} + // Dial tracks one dial from start to connection registration. It hands the // dialed connection to the sweeper atomically, so a sweep can never fall // between the dial finishing and the connection being registered. @@ -32,10 +48,11 @@ type Dial struct { ctx context.Context cancel context.CancelFunc id sweepID - done bool // set by Sweep, WrapConn or Release; guarded by sweeper.mu + done bool // set by a sweep, WrapConn or Release; guarded by sweeper.mu + gen uint64 } -// Ctx returns the dial's context. Sweep cancels it, so a dial started on the +// Ctx returns the dial's context. A sweep cancels it, so a dial started on the // old network aborts instead of waiting out its handshake timeout. func (d *Dial) Ctx() context.Context { return d.ctx @@ -69,20 +86,33 @@ func (c *sweptConn) Close() error { return c.Conn.Close() } -// Sweeper registers live connections and in-flight dials so Sweep can cut -// everything that started before the network changed. +// Sweeper registers live connections and in-flight dials so the +// network-change sweep can cut everything registered before the change. type Sweeper struct { - mu sync.Mutex - conns map[sweepID]net.Conn - dials map[sweepID]*Dial - nextID sweepID + mu sync.Mutex + conns map[sweepID]connEntry + dials map[sweepID]*Dial + nextID sweepID + gen uint64 + timer *time.Timer + sweepDelay time.Duration } -// New creates an empty sweeper. +// New creates an empty sweeper with the default configuration. func New() *Sweeper { + return NewWithConfig(Config{}) +} + +// NewWithConfig creates an empty sweeper customized by cfg. +func NewWithConfig(cfg Config) *Sweeper { + delay := cfg.SweepDelay + if delay <= 0 { + delay = DefaultSweepDelay + } return &Sweeper{ - conns: make(map[sweepID]net.Conn), - dials: make(map[sweepID]*Dial), + conns: make(map[sweepID]connEntry), + dials: make(map[sweepID]*Dial), + sweepDelay: delay, } } @@ -99,6 +129,7 @@ func (s *Sweeper) StartDial(ctx context.Context) *Dial { s.mu.Lock() d.id = s.nextID s.nextID++ + d.gen = s.gen s.dials[d.id] = d s.mu.Unlock() @@ -127,28 +158,61 @@ func (d *Dial) WrapConn(conn net.Conn) (net.Conn, error) { delete(s.dials, d.id) id := s.nextID s.nextID++ - s.conns[id] = conn + // The conn inherits the dial's generation: the socket was bound to the + // network that was default when the dial started, not when it finished. + s.conns[id] = connEntry{conn: conn, gen: d.gen} s.mu.Unlock() return &sweptConn{Conn: conn, sweeper: s, id: id}, nil } -// Sweep closes every registered connection, aborts every in-flight dial, and -// returns how many connections it closed. A dial whose connection was not -// yet handed to WrapConn is marked, so the late WrapConn closes it instead -// of registering it. -func (s *Sweeper) Sweep() int { +// MarkNetworkChange records that the OS switched networks: everything +// registered so far becomes stale, and a sweep is (re)scheduled after the +// configured delay to cut whatever is still stale by then. Owners that +// redialed in the meantime hold fresh-generation registrations and survive, +// so no cancellation is needed around the sweep. +func (s *Sweeper) MarkNetworkChange() { + if s == nil { + return + } + + s.mu.Lock() + s.gen++ + cutoff := s.gen + if s.timer != nil { + s.timer.Stop() + } + s.timer = time.AfterFunc(s.sweepDelay, func() { + n := s.sweep(cutoff) + log.Infof("network change sweep: closed %d stale connections", n) + }) + s.mu.Unlock() +} + +// sweep closes the registered connections and aborts the in-flight dials +// older than cutoff, and returns how many connections it closed. A dial +// whose connection was not yet handed to WrapConn is marked, so the late +// WrapConn closes it instead of registering it. +func (s *Sweeper) sweep(cutoff uint64) int { if s == nil { return 0 } s.mu.Lock() - conns := s.conns - dials := s.dials - s.conns = make(map[sweepID]net.Conn) - s.dials = make(map[sweepID]*Dial) - for _, d := range dials { - d.done = true + var conns []net.Conn + for id, e := range s.conns { + if e.gen < cutoff { + delete(s.conns, id) + conns = append(conns, e.conn) + } + } + var dials []*Dial + for id, d := range s.dials { + if d.gen < cutoff { + d.done = true + delete(s.dials, id) + dials = append(dials, d) + } } s.mu.Unlock() diff --git a/client/netsweep/netsweep_test.go b/client/netsweep/netsweep_test.go index a7ee0ad6a..88d660c2d 100644 --- a/client/netsweep/netsweep_test.go +++ b/client/netsweep/netsweep_test.go @@ -2,8 +2,10 @@ package netsweep import ( "context" + "math" "net" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -15,7 +17,7 @@ func TestSweepClosesRegisteredConns(t *testing.T) { c1 := wrap(t, sweeper, connPair(t)) c2 := wrap(t, sweeper, connPair(t)) - assert.Equal(t, 2, sweeper.Sweep(), "both live connections should be closed") + assert.Equal(t, 2, sweeper.sweepAll(), "both live connections should be closed") // The wrappers must report closed now. buf := make([]byte, 1) @@ -24,7 +26,7 @@ func TestSweepClosesRegisteredConns(t *testing.T) { _, err = c2.Read(buf) assert.Error(t, err, "second connection should be unusable after the sweep") - assert.Equal(t, 0, sweeper.Sweep(), "second sweep should find nothing") + assert.Equal(t, 0, sweeper.sweepAll(), "second sweep should find nothing") } func TestCloseDeregisters(t *testing.T) { @@ -33,7 +35,7 @@ func TestCloseDeregisters(t *testing.T) { conn := wrap(t, sweeper, connPair(t)) require.NoError(t, conn.Close()) - assert.Equal(t, 0, sweeper.Sweep(), "closed connection must leave the registry") + assert.Equal(t, 0, sweeper.sweepAll(), "closed connection must leave the registry") } func TestCloseIsIdempotent(t *testing.T) { @@ -48,11 +50,11 @@ func TestSweepOnlyAffectsOlderConns(t *testing.T) { sweeper := New() _ = wrap(t, sweeper, connPair(t)) - assert.Equal(t, 1, sweeper.Sweep()) + assert.Equal(t, 1, sweeper.sweepAll()) // A connection dialed after the sweep must survive until the next one. _ = wrap(t, sweeper, connPair(t)) - assert.Equal(t, 1, sweeper.Sweep(), "post-sweep connection belongs to the next sweep") + assert.Equal(t, 1, sweeper.sweepAll(), "post-sweep connection belongs to the next sweep") } func TestSweepAbortsInFlightDials(t *testing.T) { @@ -61,7 +63,7 @@ func TestSweepAbortsInFlightDials(t *testing.T) { dial := sweeper.StartDial(context.Background()) defer dial.Release() - sweeper.Sweep() + sweeper.sweepAll() assert.ErrorIs(t, dial.Ctx().Err(), context.Canceled, "sweep must cancel the in-flight dial context") } @@ -77,7 +79,7 @@ func TestReleasedDialIsNotAborted(t *testing.T) { pending := sweeper.StartDial(context.Background()) defer pending.Release() - sweeper.Sweep() + sweeper.sweepAll() assert.ErrorIs(t, pending.Ctx().Err(), context.Canceled, "pending dial must be aborted") } @@ -90,7 +92,7 @@ func TestSweepBetweenDialAndHandoffClosesConn(t *testing.T) { // The dial succeeds on the old network, then the sweep lands before the // connection is handed over. conn := connPair(t) - assert.Equal(t, 0, sweeper.Sweep(), "the connection is not registered yet") + assert.Equal(t, 0, sweeper.sweepAll(), "the connection is not registered yet") wrapped, err := dial.WrapConn(conn) require.ErrorIs(t, err, ErrSwept) @@ -100,7 +102,73 @@ func TestSweepBetweenDialAndHandoffClosesConn(t *testing.T) { _, err = conn.Read(buf) assert.Error(t, err, "the old-network connection must be closed, not leaked") - assert.Equal(t, 0, sweeper.Sweep(), "nothing may leak into the next sweep") + assert.Equal(t, 0, sweeper.sweepAll(), "nothing may leak into the next sweep") +} + +func TestMarkNetworkChangeSparesFreshConns(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond}) + + stale := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + _ = wrap(t, sweeper, connPair(t)) + + _ = stale.SetReadDeadline(time.Now().Add(time.Second)) + buf := make([]byte, 1) + _, err := stale.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "stale connection must be closed by the delayed sweep") + + assert.Equal(t, 1, sweeper.sweepAll(), "the fresh connection must survive the stale sweep") +} + +func TestMarkNetworkChangeAbortsStaleDials(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond}) + + stale := sweeper.StartDial(context.Background()) + defer stale.Release() + sweeper.MarkNetworkChange() + fresh := sweeper.StartDial(context.Background()) + defer fresh.Release() + + assert.Eventually(t, func() bool { + return stale.Ctx().Err() != nil + }, time.Second, 5*time.Millisecond, "stale dial must be aborted by the delayed sweep") + assert.NoError(t, fresh.Ctx().Err(), "post-mark dial must not be aborted") +} + +func TestConnInheritsDialGeneration(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond}) + + // The dial starts before the network change but completes after it: the + // socket is bound to the old network, so the sweep must still cut it. + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + sweeper.MarkNetworkChange() + + wrapped, err := dial.WrapConn(connPair(t)) + require.NoError(t, err) + + _ = wrapped.SetReadDeadline(time.Now().Add(time.Second)) + buf := make([]byte, 1) + _, err = wrapped.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "old-generation connection must be swept") +} + +func TestRepeatedMarksCoalesce(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond}) + + first := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + second := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + _ = wrap(t, sweeper, connPair(t)) + + buf := make([]byte, 1) + for _, conn := range []net.Conn{first, second} { + _ = conn.SetReadDeadline(time.Now().Add(time.Second)) + _, err := conn.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "every pre-mark connection must be swept by the rescheduled sweep") + } + assert.Equal(t, 1, sweeper.sweepAll(), "only the newest-generation connection may remain") } func TestNilSweeperIsNoop(t *testing.T) { @@ -114,7 +182,7 @@ func TestNilSweeperIsNoop(t *testing.T) { require.NoError(t, err) assert.Equal(t, conn, wrapped, "nil sweeper must return the conn unchanged") assert.NoError(t, dial.Ctx().Err(), "nil sweeper must not cancel the dial context") - assert.Equal(t, 0, sweeper.Sweep(), "nil sweeper closes nothing") + assert.Equal(t, 0, sweeper.sweepAll(), "nil sweeper closes nothing") } // wrap registers conn with the sweeper through a completed dial. @@ -166,3 +234,8 @@ func connPair(t *testing.T) net.Conn { return conn } + +// sweepAll cuts every registration regardless of generation. +func (s *Sweeper) sweepAll() int { + return s.sweep(math.MaxUint64) +}