mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 23:19:11 +02:00
[client] Hand off dialed connections to the sweeper atomically
WrapDialContext and WrapConn registered the dial and the connection independently, so a sweep landing between the dial finishing and WrapConn cancelled only the dial registration: the connection dialed on the old network entered the fresh registry and survived the network change. Replace the pair with a Dial handle. Sweep marks pending dials under the sweeper mutex, and WrapConn decides under the same mutex: a swept dial's connection is closed and ErrSwept returned, so the caller redials on the new network; otherwise the connection transfers to the registry with no window in between.
This commit is contained in:
@@ -28,14 +28,14 @@ func WithCustomDialer(_ bool, _ string) grpc.DialOption {
|
|||||||
// dial options in order, so the later context dialer wins.
|
// dial options in order, so the later context dialer wins.
|
||||||
func WithSweeper(sweeper *netsweep.Sweeper) grpc.DialOption {
|
func WithSweeper(sweeper *netsweep.Sweeper) grpc.DialOption {
|
||||||
return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) {
|
return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) {
|
||||||
ctx, releaseDial := sweeper.WrapDialContext(ctx)
|
dial := sweeper.StartDial(ctx)
|
||||||
defer releaseDial()
|
defer dial.Release()
|
||||||
|
|
||||||
conn, err := dialContext(ctx, addr)
|
conn, err := dialContext(dial.Ctx(), addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return sweeper.WrapConn(conn), nil
|
return dial.WrapConn(conn)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+78
-26
@@ -8,12 +8,18 @@ package netsweep
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// 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.
|
||||||
|
var ErrSwept = errors.New("netsweep: connection swept by network change")
|
||||||
|
|
||||||
// sweptConn deregisters itself from the sweeper when closed.
|
// sweptConn deregisters itself from the sweeper when closed.
|
||||||
type sweptConn struct {
|
type sweptConn struct {
|
||||||
net.Conn
|
net.Conn
|
||||||
@@ -31,7 +37,7 @@ func (c *sweptConn) Close() error {
|
|||||||
type Sweeper struct {
|
type Sweeper struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
conns map[uint64]net.Conn
|
conns map[uint64]net.Conn
|
||||||
dials map[uint64]context.CancelFunc
|
dials map[uint64]*Dial
|
||||||
nextID uint64
|
nextID uint64
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,51 +45,94 @@ type Sweeper struct {
|
|||||||
func New() *Sweeper {
|
func New() *Sweeper {
|
||||||
return &Sweeper{
|
return &Sweeper{
|
||||||
conns: make(map[uint64]net.Conn),
|
conns: make(map[uint64]net.Conn),
|
||||||
dials: make(map[uint64]context.CancelFunc),
|
dials: make(map[uint64]*Dial),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// WrapConn registers conn and returns a wrapper that deregisters it on Close.
|
// Dial tracks one dial from start to connection registration. It hands the
|
||||||
func (s *Sweeper) WrapConn(conn net.Conn) net.Conn {
|
// dialed connection to the sweeper atomically, so a sweep can never fall
|
||||||
|
// between the dial finishing and the connection being registered.
|
||||||
|
type Dial struct {
|
||||||
|
sweeper *Sweeper
|
||||||
|
ctx context.Context
|
||||||
|
cancel context.CancelFunc
|
||||||
|
id uint64
|
||||||
|
done bool // set by Sweep, WrapConn or Release; guarded by sweeper.mu
|
||||||
|
}
|
||||||
|
|
||||||
|
// StartDial registers an in-flight dial. Dial with Ctx, hand the result to
|
||||||
|
// WrapConn, and Release the dial when the attempt is over, typically deferred.
|
||||||
|
func (s *Sweeper) StartDial(ctx context.Context) *Dial {
|
||||||
if s == nil {
|
if s == nil {
|
||||||
return conn
|
return &Dial{ctx: ctx}
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
d := &Dial{sweeper: s, ctx: ctx, cancel: cancel}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
d.id = s.nextID
|
||||||
|
s.nextID++
|
||||||
|
s.dials[d.id] = d
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ctx returns the dial's context. 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// WrapConn hands conn over to the sweeper. If a sweep ran since StartDial,
|
||||||
|
// the connection belongs to the old network: it is closed and ErrSwept is
|
||||||
|
// returned. Otherwise conn is registered against the next sweep and returned
|
||||||
|
// wrapped, deregistering itself on Close. Call it once, before Release.
|
||||||
|
func (d *Dial) WrapConn(conn net.Conn) (net.Conn, error) {
|
||||||
|
s := d.sweeper
|
||||||
|
if s == nil {
|
||||||
|
return conn, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
|
if d.done {
|
||||||
|
s.mu.Unlock()
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
log.Debugf("swept dial close error: %v", err)
|
||||||
|
}
|
||||||
|
return nil, ErrSwept
|
||||||
|
}
|
||||||
|
d.done = true
|
||||||
|
delete(s.dials, d.id)
|
||||||
id := s.nextID
|
id := s.nextID
|
||||||
s.nextID++
|
s.nextID++
|
||||||
s.conns[id] = conn
|
s.conns[id] = conn
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
|
|
||||||
return &sweptConn{Conn: conn, sweeper: s, id: id}
|
return &sweptConn{Conn: conn, sweeper: s, id: id}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// WrapDialContext derives a context that Sweep cancels. The returned release
|
// Release ends the dial's registration and cancels its context. It is
|
||||||
// must be called when the dial finishes, typically deferred.
|
// idempotent and safe after WrapConn, so callers can defer it.
|
||||||
func (s *Sweeper) WrapDialContext(ctx context.Context) (context.Context, context.CancelFunc) {
|
func (d *Dial) Release() {
|
||||||
|
s := d.sweeper
|
||||||
if s == nil {
|
if s == nil {
|
||||||
return ctx, func() {}
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
|
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
id := s.nextID
|
d.done = true
|
||||||
s.nextID++
|
delete(s.dials, d.id)
|
||||||
s.dials[id] = cancel
|
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
|
|
||||||
release := func() {
|
d.cancel()
|
||||||
s.mu.Lock()
|
|
||||||
delete(s.dials, id)
|
|
||||||
s.mu.Unlock()
|
|
||||||
cancel()
|
|
||||||
}
|
|
||||||
return ctx, release
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sweep closes every registered connection, aborts every in-flight dial, and
|
// Sweep closes every registered connection, aborts every in-flight dial, and
|
||||||
// returns how many connections it closed.
|
// 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 {
|
func (s *Sweeper) Sweep() int {
|
||||||
if s == nil {
|
if s == nil {
|
||||||
return 0
|
return 0
|
||||||
@@ -93,13 +142,16 @@ func (s *Sweeper) Sweep() int {
|
|||||||
conns := s.conns
|
conns := s.conns
|
||||||
dials := s.dials
|
dials := s.dials
|
||||||
s.conns = make(map[uint64]net.Conn)
|
s.conns = make(map[uint64]net.Conn)
|
||||||
s.dials = make(map[uint64]context.CancelFunc)
|
s.dials = make(map[uint64]*Dial)
|
||||||
|
for _, d := range dials {
|
||||||
|
d.done = true
|
||||||
|
}
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
|
|
||||||
if len(dials) > 0 {
|
if len(dials) > 0 {
|
||||||
log.Debugf("aborting %d in-flight dials", len(dials))
|
log.Debugf("aborting %d in-flight dials", len(dials))
|
||||||
for _, cancel := range dials {
|
for _, d := range dials {
|
||||||
cancel()
|
d.cancel()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -12,8 +12,8 @@ import (
|
|||||||
func TestSweepClosesRegisteredConns(t *testing.T) {
|
func TestSweepClosesRegisteredConns(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
c1 := sweeper.WrapConn(connPair(t))
|
c1 := wrap(t, sweeper, connPair(t))
|
||||||
c2 := sweeper.WrapConn(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.Sweep(), "both live connections should be closed")
|
||||||
|
|
||||||
@@ -30,7 +30,7 @@ func TestSweepClosesRegisteredConns(t *testing.T) {
|
|||||||
func TestCloseDeregisters(t *testing.T) {
|
func TestCloseDeregisters(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
conn := sweeper.WrapConn(connPair(t))
|
conn := wrap(t, sweeper, connPair(t))
|
||||||
require.NoError(t, conn.Close())
|
require.NoError(t, conn.Close())
|
||||||
|
|
||||||
assert.Equal(t, 0, sweeper.Sweep(), "closed connection must leave the registry")
|
assert.Equal(t, 0, sweeper.Sweep(), "closed connection must leave the registry")
|
||||||
@@ -39,7 +39,7 @@ func TestCloseDeregisters(t *testing.T) {
|
|||||||
func TestCloseIsIdempotent(t *testing.T) {
|
func TestCloseIsIdempotent(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
conn := sweeper.WrapConn(connPair(t))
|
conn := wrap(t, sweeper, connPair(t))
|
||||||
require.NoError(t, conn.Close())
|
require.NoError(t, conn.Close())
|
||||||
assert.Error(t, conn.Close(), "double close surfaces the underlying error but must not panic")
|
assert.Error(t, conn.Close(), "double close surfaces the underlying error but must not panic")
|
||||||
}
|
}
|
||||||
@@ -47,50 +47,86 @@ func TestCloseIsIdempotent(t *testing.T) {
|
|||||||
func TestSweepOnlyAffectsOlderConns(t *testing.T) {
|
func TestSweepOnlyAffectsOlderConns(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
_ = sweeper.WrapConn(connPair(t))
|
_ = wrap(t, sweeper, connPair(t))
|
||||||
assert.Equal(t, 1, sweeper.Sweep())
|
assert.Equal(t, 1, sweeper.Sweep())
|
||||||
|
|
||||||
// A connection dialed after the sweep must survive until the next one.
|
// A connection dialed after the sweep must survive until the next one.
|
||||||
_ = sweeper.WrapConn(connPair(t))
|
_ = wrap(t, sweeper, connPair(t))
|
||||||
assert.Equal(t, 1, sweeper.Sweep(), "post-sweep connection belongs to the next sweep")
|
assert.Equal(t, 1, sweeper.Sweep(), "post-sweep connection belongs to the next sweep")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSweepAbortsInFlightDials(t *testing.T) {
|
func TestSweepAbortsInFlightDials(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
dialCtx, release := sweeper.WrapDialContext(context.Background())
|
dial := sweeper.StartDial(context.Background())
|
||||||
defer release()
|
defer dial.Release()
|
||||||
|
|
||||||
sweeper.Sweep()
|
sweeper.Sweep()
|
||||||
|
|
||||||
assert.ErrorIs(t, dialCtx.Err(), context.Canceled, "sweep must cancel the in-flight dial context")
|
assert.ErrorIs(t, dial.Ctx().Err(), context.Canceled, "sweep must cancel the in-flight dial context")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestReleasedDialIsNotAborted(t *testing.T) {
|
func TestReleasedDialIsNotAborted(t *testing.T) {
|
||||||
sweeper := New()
|
sweeper := New()
|
||||||
|
|
||||||
// Simulate a dial that finished before the sweep.
|
// Simulate a dial that finished before the sweep.
|
||||||
_, release := sweeper.WrapDialContext(context.Background())
|
released := sweeper.StartDial(context.Background())
|
||||||
release()
|
released.Release()
|
||||||
|
|
||||||
// A dial still in flight during the sweep.
|
// A dial still in flight during the sweep.
|
||||||
pendingCtx, pendingRelease := sweeper.WrapDialContext(context.Background())
|
pending := sweeper.StartDial(context.Background())
|
||||||
defer pendingRelease()
|
defer pending.Release()
|
||||||
|
|
||||||
sweeper.Sweep()
|
sweeper.Sweep()
|
||||||
assert.ErrorIs(t, pendingCtx.Err(), context.Canceled, "pending dial must be aborted")
|
assert.ErrorIs(t, pending.Ctx().Err(), context.Canceled, "pending dial must be aborted")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSweepBetweenDialAndHandoffClosesConn(t *testing.T) {
|
||||||
|
sweeper := New()
|
||||||
|
|
||||||
|
dial := sweeper.StartDial(context.Background())
|
||||||
|
defer dial.Release()
|
||||||
|
|
||||||
|
// 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")
|
||||||
|
|
||||||
|
wrapped, err := dial.WrapConn(conn)
|
||||||
|
require.ErrorIs(t, err, ErrSwept)
|
||||||
|
require.Nil(t, wrapped)
|
||||||
|
|
||||||
|
buf := make([]byte, 1)
|
||||||
|
_, 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")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNilSweeperIsNoop(t *testing.T) {
|
func TestNilSweeperIsNoop(t *testing.T) {
|
||||||
var sweeper *Sweeper
|
var sweeper *Sweeper
|
||||||
|
|
||||||
conn := connPair(t)
|
conn := connPair(t)
|
||||||
assert.Equal(t, conn, sweeper.WrapConn(conn), "nil sweeper must return the conn unchanged")
|
dial := sweeper.StartDial(context.Background())
|
||||||
assert.Equal(t, 0, sweeper.Sweep(), "nil sweeper closes nothing")
|
defer dial.Release()
|
||||||
|
|
||||||
ctx, release := sweeper.WrapDialContext(context.Background())
|
wrapped, err := dial.WrapConn(conn)
|
||||||
release()
|
require.NoError(t, err)
|
||||||
assert.NoError(t, ctx.Err(), "nil sweeper must not cancel the dial context")
|
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")
|
||||||
|
}
|
||||||
|
|
||||||
|
// wrap registers conn with the sweeper through a completed dial.
|
||||||
|
func wrap(t *testing.T, sweeper *Sweeper, conn net.Conn) net.Conn {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dial := sweeper.StartDial(context.Background())
|
||||||
|
defer dial.Release()
|
||||||
|
|
||||||
|
wrapped, err := dial.WrapConn(conn)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return wrapped
|
||||||
}
|
}
|
||||||
|
|
||||||
// connPair dials a loopback TCP connection and keeps the accepted peer open
|
// connPair dials a loopback TCP connection and keeps the accepted peer open
|
||||||
|
|||||||
@@ -400,8 +400,9 @@ func (c *Client) Close() error {
|
|||||||
func (c *Client) connect(ctx context.Context) (*RelayAddr, error) {
|
func (c *Client) connect(ctx context.Context) (*RelayAddr, error) {
|
||||||
// A sweep cancels this context, so a dial started on the old network
|
// A sweep cancels this context, so a dial started on the old network
|
||||||
// aborts instead of waiting out its handshake timeout.
|
// aborts instead of waiting out its handshake timeout.
|
||||||
ctx, releaseDial := c.sweeper.WrapDialContext(ctx)
|
dial := c.sweeper.StartDial(ctx)
|
||||||
defer releaseDial()
|
defer dial.Release()
|
||||||
|
ctx = dial.Ctx()
|
||||||
|
|
||||||
mode := transportModeFromEnv()
|
mode := transportModeFromEnv()
|
||||||
dialers := c.getDialers(mode)
|
dialers := c.getDialers(mode)
|
||||||
@@ -433,7 +434,10 @@ func (c *Client) connect(ctx context.Context) (*RelayAddr, error) {
|
|||||||
c.transport = tc.Protocol()
|
c.transport = tc.Protocol()
|
||||||
}
|
}
|
||||||
|
|
||||||
conn = c.sweeper.WrapConn(conn)
|
conn, err := dial.WrapConn(conn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("register connection: %w", err)
|
||||||
|
}
|
||||||
c.relayConn = conn
|
c.relayConn = conn
|
||||||
c.datagramFallbackTriggered.Store(false)
|
c.datagramFallbackTriggered.Store(false)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user