From 2aab1b5ac290b83261a5b0781288a64898f77a91 Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 11:42:13 +0200 Subject: [PATCH 1/7] GO-7556: recover from stale connections after sleep/wake After sleep/wake, clients kept reusing dead pre-sleep connections: new streams hung in the handshake or RPCs waited for replies for 10-30s+, and pool.Flush could leave pre-flush peers usable. net/pool: - Flush swaps in a fresh incoming/outgoing cache pair and closes the old pair in the background: peers in parallel, outgoing first, so pre-flush dials are cancelled even if an incoming teardown hangs. - Lookups merge the pair ctx and retry until their result comes from the pair that is still current, so a Get after Flush never returns or waits out a pre-flush dial; ErrClosed surfaces only when the pool is closed. A non-blocking fast path (ocache.Peeker) keeps cached lookups allocation-free and faster than before. - AddPeer retries on a flushed pair or a transient entry. Incompatible-version verdicts survive Flush. Flush/Close are serialized; Close waits, bounded, for pending teardowns. Prometheus collectors are registered once and shared across pairs. net/peer: - Lazy per-peer cleanup owner moves synchronous sub-conn closes (release, failed handshake, gc) off the caller path. In-flight closes still count toward the open limiter; the MultiConn is closed only on stall evidence. Conns doomed by gc are never reused. net/transport: - yamux Open honours ctx. - QUIC and yamux stream Read/Write report a dead connection or session as ErrConnClosed (wrapped). - Optional WriteTimeouter. net/secureservice/handshake: - OutgoingProtoHandshakeWithCloser: a cancelled handshake hands the conn to a closer exactly once; the legacy function keeps its contract. - HandshakeError unwraps to the underlying error. --- app/ocache/metrics.go | 108 +- app/ocache/metrics_test.go | 41 +- app/ocache/ocache.go | 41 + app/ocache/ocache_test.go | 76 + net/peer/cleanup.go | 174 +++ net/peer/cleanup_test.go | 581 ++++++++ net/peer/peer.go | 163 ++- net/pool/pool.go | 402 +++++- net/pool/pool_bench_test.go | 116 ++ net/pool/pool_flush_test.go | 1567 +++++++++++++++++++++ net/pool/pool_test.go | 89 +- net/pool/poolservice.go | 128 +- net/secureservice/handshake/handshake.go | 33 +- net/secureservice/handshake/proto.go | 65 +- net/secureservice/handshake/proto_test.go | 129 ++ net/transport/iroh/conn.go | 5 + net/transport/quic/conn.go | 55 +- net/transport/quic/conn_errors_test.go | 164 +++ net/transport/transport.go | 7 + net/transport/webtransport/conn.go | 5 + net/transport/yamux/conn.go | 201 ++- net/transport/yamux/conn_test.go | 308 ++++ net/transport/yamux/yamux.go | 4 +- 23 files changed, 4215 insertions(+), 247 deletions(-) create mode 100644 net/peer/cleanup.go create mode 100644 net/peer/cleanup_test.go create mode 100644 net/pool/pool_bench_test.go create mode 100644 net/pool/pool_flush_test.go create mode 100644 net/transport/quic/conn_errors_test.go create mode 100644 net/transport/yamux/conn_test.go diff --git a/app/ocache/metrics.go b/app/ocache/metrics.go index 74697b343..b009ca4d6 100644 --- a/app/ocache/metrics.go +++ b/app/ocache/metrics.go @@ -10,49 +10,15 @@ func WithPrometheus(reg *prometheus.Registry, namespace, subsystem string) Optio if reg == nil { return nil } - if subsystem == "" { - subsystem = "cache" - } - nameSplit := strings.Split(namespace, ".") - subSplit := strings.Split(subsystem, ".") - namespace = strings.Join(nameSplit, "_") - subsystem = strings.Join(subSplit, "_") - return func(cache *oCache) { + c := NewPrometheusCollectors(namespace, subsystem, cache.Len) + c.MustRegister(reg) cache.metrics = &metrics{ - hit: prometheus.NewCounter(prometheus.CounterOpts{ - Namespace: namespace, - Subsystem: subsystem, - Name: "hit", - Help: "cache hit count", - }), - miss: prometheus.NewCounter(prometheus.CounterOpts{ - Namespace: namespace, - Subsystem: subsystem, - Name: "miss", - Help: "cache miss count", - }), - gc: prometheus.NewCounter(prometheus.CounterOpts{ - Namespace: namespace, - Subsystem: subsystem, - Name: "gc", - Help: "garbage collected count", - }), - size: prometheus.NewGaugeFunc(prometheus.GaugeOpts{ - Namespace: namespace, - Subsystem: subsystem, - Name: "size", - Help: "cache size", - }, func() float64 { - return float64(cache.Len()) - }), + hit: c.Hit, + miss: c.Miss, + gc: c.GC, + size: c.Size, } - reg.MustRegister( - cache.metrics.hit, - cache.metrics.miss, - cache.metrics.gc, - cache.metrics.size, - ) } } @@ -67,6 +33,68 @@ func WithPrometheusMetrics(hit, miss, gc prometheus.Counter, size prometheus.Gau } } +// PrometheusCollectors are the collectors a cache reports through: the ones +// WithPrometheus builds and registers, exposed for a caller that recreates +// its cache and so must register them once and hand them to every instance. +type PrometheusCollectors struct { + Hit, Miss, GC prometheus.Counter + Size prometheus.GaugeFunc +} + +// NewPrometheusCollectors builds unregistered collectors with the names +// WithPrometheus would register (__{hit,miss,gc,size}, +// dots turned into underscores, subsystem defaulting to "cache"). size is +// read through sizeFn, so it can resolve whichever cache is current. +func NewPrometheusCollectors(namespace, subsystem string, sizeFn func() int) PrometheusCollectors { + if subsystem == "" { + subsystem = "cache" + } + nameSplit := strings.Split(namespace, ".") + subSplit := strings.Split(subsystem, ".") + namespace = strings.Join(nameSplit, "_") + subsystem = strings.Join(subSplit, "_") + return PrometheusCollectors{ + Hit: prometheus.NewCounter(prometheus.CounterOpts{ + Namespace: namespace, + Subsystem: subsystem, + Name: "hit", + Help: "cache hit count", + }), + Miss: prometheus.NewCounter(prometheus.CounterOpts{ + Namespace: namespace, + Subsystem: subsystem, + Name: "miss", + Help: "cache miss count", + }), + GC: prometheus.NewCounter(prometheus.CounterOpts{ + Namespace: namespace, + Subsystem: subsystem, + Name: "gc", + Help: "garbage collected count", + }), + Size: prometheus.NewGaugeFunc(prometheus.GaugeOpts{ + Namespace: namespace, + Subsystem: subsystem, + Name: "size", + Help: "cache size", + }, func() float64 { + return float64(sizeFn()) + }), + } +} + +// MustRegister registers the collectors with reg; like prometheus it panics +// on a second registration of the same names. +func (c PrometheusCollectors) MustRegister(reg prometheus.Registerer) { + reg.MustRegister(c.Hit, c.Miss, c.GC, c.Size) +} + +// Option makes a cache report through these collectors without registering +// anything. +func (c PrometheusCollectors) Option() Option { + return WithPrometheusMetrics(c.Hit, c.Miss, c.GC, c.Size) +} + type metrics struct { hit prometheus.Counter miss prometheus.Counter diff --git a/app/ocache/metrics_test.go b/app/ocache/metrics_test.go index e14d670d4..dfb6fc9e6 100644 --- a/app/ocache/metrics_test.go +++ b/app/ocache/metrics_test.go @@ -2,10 +2,11 @@ package ocache import ( "context" - "github.com/prometheus/client_golang/prometheus" - "github.com/stretchr/testify/require" "strings" "testing" + + "github.com/prometheus/client_golang/prometheus" + "github.com/stretchr/testify/require" ) func TestWithPrometheus_MetricsConvertsDots(t *testing.T) { @@ -17,3 +18,39 @@ func TestWithPrometheus_MetricsConvertsDots(t *testing.T) { require.NoError(t, err) require.True(t, strings.Contains(cache.metrics.hit.Desc().String(), "some_name_some_system_hit")) } + +func TestWithPrometheus_Registers(t *testing.T) { + reg := prometheus.NewRegistry() + cache := New(func(ctx context.Context, id string) (value Object, err error) { + return &testObject{}, nil + }, WithPrometheus(reg, "some.name", "some.system")) + _, err := cache.Get(context.Background(), "id") + require.NoError(t, err) + families, err := reg.Gather() + require.NoError(t, err) + values := map[string]float64{} + for _, mf := range families { + m := mf.GetMetric()[0] + if m.GetGauge() != nil { + values[mf.GetName()] = m.GetGauge().GetValue() + } else { + values[mf.GetName()] = m.GetCounter().GetValue() + } + } + require.Equal(t, map[string]float64{ + "some_name_some_system_hit": 0, + "some_name_some_system_miss": 1, + "some_name_some_system_gc": 0, + "some_name_some_system_size": 1, + }, values) + // the same names cannot be registered twice: the reason a caller that + // recreates its cache goes through NewPrometheusCollectors instead + require.Panics(t, func() { + New(nil, WithPrometheus(reg, "some.name", "some.system")) + }) + require.NotPanics(t, func() { + c := NewPrometheusCollectors("some.name", "some.system", func() int { return 0 }) + New(nil, c.Option()) + New(nil, c.Option()) + }) +} diff --git a/app/ocache/ocache.go b/app/ocache/ocache.go index 6b6bde665..3e8edb3a1 100644 --- a/app/ocache/ocache.go +++ b/app/ocache/ocache.go @@ -251,6 +251,47 @@ func (c *oCache) Pick(ctx context.Context, id string) (value Object, err error) return val.waitLoad(ctx, id) } +// Peeker is the non-blocking read the cache returned by New offers on top of +// OCache; kept off that interface so other implementations stay valid. +type Peeker interface { + // Peek returns the value for id only if it is loaded and not being + // closed, without loading, waiting or allocating: the hot path for + // callers that handle a miss themselves. A hit counts as a cache hit and, + // with touch, refreshes the GC deadline like Get; a miss counts nothing + // (ok=false also for a loading entry, a closing one or a closed cache). + Peek(id string, touch bool) (value Object, ok bool) +} + +func (c *oCache) Peek(id string, touch bool) (value Object, ok bool) { + c.mu.Lock() + e, exists := c.data[id] + if c.closed || !exists || e.isClosing() { + c.mu.Unlock() + return nil, false + } + select { + case <-e.load: + default: + // still loading + c.mu.Unlock() + return nil, false + } + // value and loadErr are written before load closes; a failed load deletes + // its entry under c.mu, so a non-nil loadErr here means the entry is on + // its way out + if e.loadErr != nil || e.value == nil { + c.mu.Unlock() + return nil, false + } + if touch { + e.lastUsage = time.Now() + } + value = e.value + c.mu.Unlock() + c.metricsGet(true) + return value, true +} + // ctx is the cancellable load context Get created together with the entry. func (c *oCache) load(ctx context.Context, id string, e *entry) { defer func() { diff --git a/app/ocache/ocache_test.go b/app/ocache/ocache_test.go index 577bdb954..bfa7e2f0a 100644 --- a/app/ocache/ocache_test.go +++ b/app/ocache/ocache_test.go @@ -10,6 +10,7 @@ import ( "testing" "time" + "github.com/prometheus/client_golang/prometheus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -1401,3 +1402,78 @@ func TestOCache_ForEachAfterClose(t *testing.T) { _, err := c.Pick(ctx, "id") require.ErrorIs(t, err, ErrClosed) } + +func TestOCache_Peek(t *testing.T) { + t.Run("hit touches and counts, miss counts nothing", func(t *testing.T) { + reg := prometheus.NewRegistry() + obj := NewTestObject("a", true, nil) + c := New(func(ctx context.Context, id string) (Object, error) { + return obj, nil + }, WithTTL(time.Hour), WithGCPeriod(0), WithPrometheus(reg, "peek", "test")).(*oCache) + _, ok := c.Peek("a", true) + require.False(t, ok, "nothing loaded yet") + _, err := c.Get(ctx, "a") + require.NoError(t, err) + c.mu.Lock() + c.data["a"].lastUsage = time.Now().Add(-time.Minute) + c.mu.Unlock() + v, ok := c.Peek("a", false) + require.True(t, ok) + require.Same(t, obj, v) + c.mu.Lock() + require.Less(t, c.data["a"].lastUsage, time.Now().Add(-30*time.Second), "Pick-like peek must not refresh the deadline") + c.mu.Unlock() + v, ok = c.Peek("a", true) + require.True(t, ok) + require.Same(t, obj, v) + c.mu.Lock() + require.Greater(t, c.data["a"].lastUsage, time.Now().Add(-time.Second), "Get-like peek refreshes the deadline") + c.mu.Unlock() + families, err := reg.Gather() + require.NoError(t, err) + values := map[string]float64{} + for _, mf := range families { + if m := mf.GetMetric()[0]; m.GetCounter() != nil { + values[mf.GetName()] = m.GetCounter().GetValue() + } + } + // one miss from Get's load, two hits from the two peeks + require.Equal(t, float64(1), values["peek_test_miss"]) + require.Equal(t, float64(2), values["peek_test_hit"]) + }) + t.Run("loading, closing and closed are misses", func(t *testing.T) { + loading := make(chan struct{}) + release := make(chan struct{}) + closeCh := make(chan struct{}) + obj := NewTestObject("a", false, closeCh) + c := New(func(ctx context.Context, id string) (Object, error) { + close(loading) + <-release + return obj, nil + }, WithTTL(time.Hour), WithGCPeriod(0)).(*oCache) + go func() { _, _ = c.Get(ctx, "a") }() + <-loading + _, ok := c.Peek("a", true) + require.False(t, ok, "loading entry") + close(release) + require.Eventually(t, func() bool { _, ok := c.Peek("a", false); return ok }, time.Second, time.Millisecond) + + // a closer holds the entry: Remove blocks in obj.Close until closeCh + removed := make(chan struct{}) + go func() { + _, _ = c.Remove(ctx, "a") + close(removed) + }() + require.Eventually(t, func() bool { _, ok := c.Peek("a", false); return !ok }, time.Second, time.Millisecond) + close(closeCh) + <-removed + + require.NoError(t, c.Add("b", NewTestObject("b", true, nil))) + v, ok := c.Peek("b", false) + require.True(t, ok) + require.NotNil(t, v) + require.NoError(t, c.Close()) + _, ok = c.Peek("b", false) + require.False(t, ok, "closed cache") + }) +} diff --git a/net/peer/cleanup.go b/net/peer/cleanup.go new file mode 100644 index 000000000..575baf0a6 --- /dev/null +++ b/net/peer/cleanup.go @@ -0,0 +1,174 @@ +package peer + +import ( + "io" + "sync" + "sync/atomic" + "time" + + "go.uber.org/zap" + + "github.com/anyproto/any-sync/net/transport" +) + +const ( + // cleanupMaxWorkers bounds the sub-connection closes a peer runs at once + cleanupMaxWorkers = 64 + // cleanupDefaultStallTimeout is the stall threshold for a transport that + // does not report its write timeout: twice the default yamux one + cleanupDefaultStallTimeout = 20 * time.Second +) + +// closeStallTimeout is how long a single close may run before the transport +// counts as stalled. A healthy close takes about one round trip (drpc waits +// for the remote FIN), but on a congested link a FIN may wait up to the +// transport's write timeout, so the threshold is twice that. For yamux this +// is WriteTimeoutSec, which also sets ConnectionWriteTimeout and +// StreamCloseTimeout: a single stream close cannot legitimately outlast it. +func closeStallTimeout(mc transport.MultiConn) time.Duration { + if wt, ok := mc.(transport.WriteTimeouter); ok { + if d := wt.WriteTimeout(); d > 0 { + return 2 * d + } + } + return cleanupDefaultStallTimeout +} + +// cleanupOwner closes a peer's sub connections off the callers' path. Closing +// a drpc conn waits for its reader, stream manager and transport, and a yamux +// stream close sends a FIN under a write timeout, so on a stalled connection a +// synchronous close turns a caller's expired deadline into a long hang. +// +// Workers are started on demand, up to cleanupMaxWorkers, and exit once there +// is nothing left to close, so an idle peer costs no goroutines. close never +// blocks and never drops a close; closes beyond the workers wait in a pending +// list. Its size is bounded by the peer's sub conns, and the peer's open +// limiter counts every close in flight (inFlight), so a peer that closes +// faster than the transport can keep up is throttled rather than piling up. +// +// Only on evidence of a stall, every worker busy and one of them on a close +// older than stallTimeout, is the whole MultiConn closed, which makes every +// pending and further close quick. A burst or a sustained rate of closes on a +// healthy connection just queues. The check runs when a close is handed over, +// so a stall is detected on the first close queued after stallTimeout; +// meanwhile the stuck closes are still bounded by the transport's own +// timeouts, so the cost of the delay is latency only. +type cleanupOwner struct { + mc transport.MultiConn + stallTimeout time.Duration + + mu sync.Mutex + pending []io.Closer + // running is the set of live workers + running map[*cleanupWorker]struct{} + + // inflight counts closes handed over and not yet finished + inflight atomic.Int32 + escalated atomic.Bool + escalations atomic.Int64 +} + +type cleanupWorker struct { + // started is when the current close began; guarded by cleanupOwner.mu + started time.Time +} + +func newCleanupOwner(mc transport.MultiConn) *cleanupOwner { + return &cleanupOwner{ + mc: mc, + stallTimeout: closeStallTimeout(mc), + running: map[*cleanupWorker]struct{}{}, + } +} + +// close hands cl over to be closed in the background. It never blocks. +func (c *cleanupOwner) close(cl io.Closer) { + if c == nil { + // a peer built without NewPeer + go func() { _ = cl.Close() }() + return + } + c.inflight.Add(1) + c.mu.Lock() + if len(c.running) < cleanupMaxWorkers { + w := &cleanupWorker{started: time.Now()} + c.running[w] = struct{}{} + c.mu.Unlock() + go c.work(w, cl) + return + } + if !c.stalledLocked() { + c.pending = append(c.pending, cl) + c.mu.Unlock() + return + } + c.mu.Unlock() + c.escalate(cl) +} + +// stalledLocked reports whether a running close has outlived stallTimeout +func (c *cleanupOwner) stalledLocked() bool { + now := time.Now() + for w := range c.running { + if now.Sub(w.started) > c.stallTimeout { + return true + } + } + return false +} + +func (c *cleanupOwner) work(w *cleanupWorker, cl io.Closer) { + for { + _ = cl.Close() + c.inflight.Add(-1) + c.mu.Lock() + if len(c.pending) == 0 { + delete(c.running, w) + c.pending = nil + c.mu.Unlock() + return + } + cl = c.pending[0] + c.pending[0] = nil + c.pending = c.pending[1:] + w.started = time.Now() + c.mu.Unlock() + } +} + +// escalate closes the whole connection: the transport is stalled, so the +// closes queued behind it would otherwise wait out its timeouts. A dead +// transport makes cl's close quick; the goroutines spawned here are bounded by +// the sub conns alive when the MultiConn closed, as no new ones can be opened +// afterwards. +func (c *cleanupOwner) escalate(cl io.Closer) { + c.escalations.Add(1) + first := c.escalated.CompareAndSwap(false, true) + if first { + log.Warn("sub connection cleanup is stalled: closing the connection") + } + go func() { + if first { + if err := c.mc.Close(); err != nil { + log.Debug("close connection on stalled cleanup", zap.Error(err)) + } + } + _ = cl.Close() + c.inflight.Add(-1) + }() +} + +// inFlight returns the number of closes handed over and not yet finished +func (c *cleanupOwner) inFlight() int { + if c == nil { + return 0 + } + return int(c.inflight.Load()) +} + +// stats returns the number of running and pending closes +func (c *cleanupOwner) stats() (running, pending int) { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.running), len(c.pending) +} diff --git a/net/peer/cleanup_test.go b/net/peer/cleanup_test.go new file mode 100644 index 000000000..889813d0e --- /dev/null +++ b/net/peer/cleanup_test.go @@ -0,0 +1,581 @@ +package peer + +import ( + "context" + "io" + "net" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + "storj.io/drpc" + "storj.io/drpc/drpcwire" + + "github.com/anyproto/any-sync/net/connutil" + "github.com/anyproto/any-sync/net/secureservice/handshake" + "github.com/anyproto/any-sync/net/secureservice/handshake/handshakeproto" + "github.com/anyproto/any-sync/net/transport/mock_transport" +) + +type rawMsg []byte + +type rawEncoding struct{} + +func (rawEncoding) Marshal(msg drpc.Message) ([]byte, error) { return *msg.(*rawMsg), nil } + +func (rawEncoding) Unmarshal(buf []byte, msg drpc.Message) error { + *msg.(*rawMsg) = append([]byte(nil), buf...) + return nil +} + +// blockingCloser is a closer whose Close blocks until released +type blockingCloser struct { + release chan struct{} + closed chan struct{} + once sync.Once +} + +func newBlockingCloser(release chan struct{}) *blockingCloser { + return &blockingCloser{release: release, closed: make(chan struct{})} +} + +func (b *blockingCloser) Close() error { + <-b.release + b.once.Do(func() { close(b.closed) }) + return nil +} + +// blockingCloseConn models a stalled sub connection whose Close blocks on +// the transport until released +type blockingCloseConn struct { + net.Conn + release chan struct{} + closed chan struct{} + once sync.Once +} + +func newBlockingCloseConn(conn net.Conn, release chan struct{}) *blockingCloseConn { + return &blockingCloseConn{Conn: conn, release: release, closed: make(chan struct{})} +} + +func (c *blockingCloseConn) Close() error { + <-c.release + c.once.Do(func() { close(c.closed) }) + return c.Conn.Close() +} + +// quickCloser models a healthy sub conn close: about one round trip +type quickCloser struct { + d time.Duration + running *atomic.Int32 + maxSeen *atomic.Int32 + closed chan struct{} +} + +func (q *quickCloser) Close() error { + n := q.running.Add(1) + for { + m := q.maxSeen.Load() + if n <= m || q.maxSeen.CompareAndSwap(m, n) { + break + } + } + time.Sleep(q.d) + q.running.Add(-1) + close(q.closed) + return nil +} + +type writeTimeoutMC struct { + *mock_transport.MockMultiConn + wt time.Duration +} + +func (w writeTimeoutMC) WriteTimeout() time.Duration { return w.wt } + +func TestCleanupOwner(t *testing.T) { + newMC := func(t *testing.T) (*mock_transport.MockMultiConn, *atomic.Int32) { + ctrl := gomock.NewController(t) + mc := mock_transport.NewMockMultiConn(ctrl) + var mcCloses atomic.Int32 + mc.EXPECT().Close().DoAndReturn(func() error { + mcCloses.Add(1) + return nil + }).AnyTimes() + return mc, &mcCloses + } + t.Run("burst on a healthy connection never closes it", func(t *testing.T) { + mc, mcCloses := newMC(t) + c := newCleanupOwner(mc) + + const burst = 300 + var running, maxSeen atomic.Int32 + closers := make([]*quickCloser, burst) + var wg sync.WaitGroup + for i := range closers { + closers[i] = &quickCloser{d: 2 * time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} + } + start := time.Now() + for _, cl := range closers { + wg.Add(1) + go func() { + defer wg.Done() + c.close(cl) + }() + } + wg.Wait() + assert.Less(t, time.Since(start), 500*time.Millisecond, "close must never block") + for _, cl := range closers { + select { + case <-cl.closed: + case <-time.After(5 * time.Second): + t.Fatal("a close was dropped") + } + } + assert.Zero(t, c.escalations.Load()) + assert.Zero(t, mcCloses.Load(), "a healthy connection must not be closed") + assert.LessOrEqual(t, int(maxSeen.Load()), cleanupMaxWorkers) + // workers exit once idle: an idle peer costs no goroutines + require.Eventually(t, func() bool { + r, p := c.stats() + return r == 0 && p == 0 + }, time.Second, time.Millisecond) + }) + t.Run("stalled transport closes the multiconn", func(t *testing.T) { + mc, mcCloses := newMC(t) + c := newCleanupOwner(mc) + c.stallTimeout = 50 * time.Millisecond + + release := make(chan struct{}) + var closers []*blockingCloser + enqueue := func() *blockingCloser { + cl := newBlockingCloser(release) + closers = append(closers, cl) + c.close(cl) + return cl + } + for i := 0; i < cleanupMaxWorkers; i++ { + enqueue() + } + // saturated but not stalled yet: the close just queues + enqueue() + r, p := c.stats() + assert.Equal(t, cleanupMaxWorkers, r) + assert.Equal(t, 1, p) + assert.Zero(t, c.escalations.Load()) + + time.Sleep(2 * c.stallTimeout) + start := time.Now() + enqueue() + enqueue() + assert.Less(t, time.Since(start), 100*time.Millisecond, "close must never block") + assert.Equal(t, int64(2), c.escalations.Load()) + require.Eventually(t, func() bool { return mcCloses.Load() == 1 }, time.Second, time.Millisecond) + + // nothing is dropped, and the multiconn is closed only once + close(release) + for _, cl := range closers { + select { + case <-cl.closed: + case <-time.After(time.Second): + t.Fatal("a queued close was dropped") + } + } + assert.Equal(t, int32(1), mcCloses.Load()) + }) + t.Run("sustained close rate on a healthy connection never closes it", func(t *testing.T) { + mc, mcCloses := newMC(t) + c := newCleanupOwner(mc) + + // far more closes than workers, arriving faster than they finish + const total = 3000 + var running, maxSeen atomic.Int32 + closers := make([]*quickCloser, total) + for i := range closers { + closers[i] = &quickCloser{d: 5 * time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} + } + var sawInFlight int + for i, cl := range closers { + c.close(cl) + if i%100 == 0 { + time.Sleep(time.Millisecond) + sawInFlight = max(sawInFlight, c.inFlight()) + } + } + for _, cl := range closers { + select { + case <-cl.closed: + case <-time.After(10 * time.Second): + t.Fatal("a close was dropped") + } + } + assert.Zero(t, c.escalations.Load()) + assert.Zero(t, mcCloses.Load(), "a healthy connection must not be closed") + // closes in flight are visible to the peer's open limiter + assert.Greater(t, sawInFlight, cleanupMaxWorkers) + require.Eventually(t, func() bool { return c.inFlight() == 0 }, time.Second, time.Millisecond) + }) + t.Run("stall threshold follows the transport write timeout", func(t *testing.T) { + mc, _ := newMC(t) + assert.Equal(t, cleanupDefaultStallTimeout, newCleanupOwner(mc).stallTimeout) + assert.Equal(t, 30*time.Second, newCleanupOwner(writeTimeoutMC{mc, 15 * time.Second}).stallTimeout) + assert.Equal(t, cleanupDefaultStallTimeout, newCleanupOwner(writeTimeoutMC{mc, 0}).stallTimeout) + }) + t.Run("nil owner", func(t *testing.T) { + var c *cleanupOwner + release := make(chan struct{}) + close(release) + cl := newBlockingCloser(release) + c.close(cl) + select { + case <-cl.closed: + case <-time.After(time.Second): + t.Fatal("close was dropped") + } + }) +} + +func TestPeer_HandshakeFailureCloseDoesNotBlock(t *testing.T) { + t.Run("protocol error", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + defer close(release) + + in, out := net.Pipe() + defer out.Close() + conn := newBlockingCloseConn(in, release) + // the remote declines the protocol, which leaves the close to openDrpcConn + go func() { + _, _ = handshake.IncomingProtoHandshake(ctx, out, handshake.ProtoChecker{ + AllowedProtoTypes: []handshakeproto.ProtoType{handshakeproto.ProtoType(100)}, + }) + }() + fx.mc.EXPECT().Open(gomock.Any()).Return(conn, nil) + + start := time.Now() + _, err := fx.AcquireDrpcConn(ctx) + require.ErrorIs(t, err, handshake.ErrRemoteIncompatibleProto) + assert.Less(t, time.Since(start), time.Second) + }) + t.Run("deadline", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + + in, out := net.Pipe() + defer out.Close() + conn := newBlockingCloseConn(in, release) + // the remote never answers + go func() { _, _ = io.Copy(io.Discard, out) }() + fx.mc.EXPECT().Open(gomock.Any()).Return(conn, nil) + + actx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + start := time.Now() + _, err := fx.AcquireDrpcConn(actx) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.Less(t, time.Since(start), time.Second) + + // the stream is still closed once the transport lets it + close(release) + select { + case <-conn.closed: + case <-time.After(time.Second): + t.Fatal("stream was not closed") + } + }) +} + +// TestPeer_RPCDeadlineWithBlockedClose is the payment-call shape: an +// established RPC hits its deadline on a stalled connection whose stream +// close blocks. The whole acquire/use/release path must return within the +// caller's budget, every time. +func TestPeer_RPCDeadlineWithBlockedClose(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + var releaseOnce sync.Once + releaseAll := func() { releaseOnce.Do(func() { close(release) }) } + defer releaseAll() + + var ( + connsMu sync.Mutex + conns []*blockingCloseConn + remotes []net.Conn + ) + defer func() { + connsMu.Lock() + defer connsMu.Unlock() + for _, r := range remotes { + _ = r.Close() + } + }() + fx.mc.EXPECT().Open(gomock.Any()).DoAndReturn(func(context.Context) (net.Conn, error) { + in, out := net.Pipe() + conn := newBlockingCloseConn(in, release) + connsMu.Lock() + conns = append(conns, conn) + remotes = append(remotes, out) + connsMu.Unlock() + // the remote handshakes, then swallows the request and never replies + go func() { + if _, err := handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker); err != nil { + return + } + _, _ = io.Copy(io.Discard, out) + }() + return conn, nil + }).AnyTimes() + + const ( + budget = 100 * time.Millisecond + repetitions = 10 + ) + for i := 0; i < repetitions; i++ { + cctx, cancel := context.WithTimeout(ctx, budget) + start := time.Now() + err := fx.DoDrpc(cctx, func(c drpc.Conn) error { + stream, err := c.NewStream(cctx, "/test.Test/Call", nil) + if err != nil { + return err + } + if err = stream.(streamRawWrite).RawWrite(drpcwire.KindMessage, []byte("req")); err != nil { + return err + } + var reply rawMsg + return stream.MsgRecv(&reply, rawEncoding{}) + }) + elapsed := time.Since(start) + cancel() + require.Error(t, err) + require.Less(t, elapsed, budget+500*time.Millisecond, "repetition %d: the call must return within its budget", i) + // the cancelled sub conn is never handed out again + fx.mu.Lock() + assert.Empty(t, fx.inactive, "repetition %d", i) + assert.Empty(t, fx.active, "repetition %d", i) + fx.mu.Unlock() + // at most one blocked close per repetition, nothing escalated + running, _ := fx.cleanup.stats() + require.LessOrEqual(t, running, i+1) + } + connsMu.Lock() + require.Len(t, conns, repetitions, "each repetition opens a fresh sub conn") + connsMu.Unlock() + require.Zero(t, fx.cleanup.escalations.Load()) + + // once the transport lets go, every stream is closed and nothing leaks + releaseAll() + connsMu.Lock() + for _, conn := range conns { + select { + case <-conn.closed: + case <-time.After(2 * time.Second): + t.Fatal("stream was not closed") + } + } + connsMu.Unlock() + require.Eventually(t, func() bool { + running, pending := fx.cleanup.stats() + return running == 0 && pending == 0 + }, 5*time.Second, 10*time.Millisecond, "cleanup workers must exit once idle") +} + +// closedConn is a released sub conn that is already closed +type closedConn struct { + closedCh chan struct{} + closes atomic.Int32 +} + +func newClosedConn() *closedConn { + ch := make(chan struct{}) + close(ch) + return &closedConn{closedCh: ch} +} + +func (c *closedConn) Close() error { c.closes.Add(1); return nil } +func (c *closedConn) Closed() <-chan struct{} { return c.closedCh } +func (c *closedConn) Unblocked() <-chan struct{} { return c.closedCh } +func (c *closedConn) NewStream(context.Context, string, drpc.Encoding) (drpc.Stream, error) { + return nil, io.EOF +} +func (c *closedConn) Invoke(context.Context, string, drpc.Encoding, drpc.Message, drpc.Message) error { + return io.EOF +} + +func TestPeer_ReleaseClosedAndCancelled(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + cctx, cancel := context.WithCancel(ctx) + cancel() + // select picks randomly between ready cases: repeat to hit both + for i := 0; i < 50; i++ { + conn := newClosedConn() + sc := &subConn{ConnUnblocked: conn} + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() + + fx.ReleaseDrpcConn(cctx, sc) + + fx.mu.Lock() + assert.Empty(t, fx.inactive, "a closed or cancelled conn is never reused") + assert.Empty(t, fx.active) + fx.mu.Unlock() + running, pending := fx.cleanup.stats() + assert.Zero(t, running+pending, "an already closed conn is not queued for cleanup") + assert.Zero(t, conn.closes.Load()) + } +} + +// TestPeer_ReleaseAfterGCDoesNotReuse: a conn gc took out of active is never +// handed out again, even while its close is still pending in the owner +func TestPeer_ReleaseAfterGCDoesNotReuse(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + defer close(release) + + in, out := net.Pipe() + defer out.Close() + go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() + conn := newBlockingCloseConn(in, release) + fx.mc.EXPECT().Open(gomock.Any()).Return(conn, nil) + fx.mc.EXPECT().Addr().Return("").AnyTimes() + + dc, err := fx.AcquireDrpcConn(ctx) + require.NoError(t, err) + // keep every cleanup worker busy, so the doomed conn's close stays + // pending and never starts (drpc signals Closed as soon as it does) + for i := 0; i < cleanupMaxWorkers; i++ { + fx.cleanup.close(newBlockingCloser(release)) + } + time.Sleep(20 * time.Millisecond) + // the idle active conn is doomed; its close waits in the owner + fx.gc(time.Millisecond) + select { + case <-dc.Closed(): + t.Fatal("close was expected to be still pending") + default: + } + + fx.ReleaseDrpcConn(ctx, dc) + fx.mu.Lock() + assert.Empty(t, fx.inactive, "a doomed conn must not be reused") + assert.Empty(t, fx.active) + fx.mu.Unlock() + select { + case got := <-fx.subConnRelease: + require.Nil(t, got, "a doomed conn must not be handed to a waiter") + default: + } +} + +// racyConn runs gc from inside Unblocked: gc dooms the conn after +// ReleaseDrpcConn has checked the flag, before it takes the peer lock +type racyConn struct { + p *peer + never chan struct{} + ready chan struct{} +} + +func (c *racyConn) Close() error { return nil } +func (c *racyConn) Closed() <-chan struct{} { return c.never } +func (c *racyConn) Unblocked() <-chan struct{} { + c.p.gc(time.Millisecond) + return c.ready +} +func (c *racyConn) NewStream(context.Context, string, drpc.Encoding) (drpc.Stream, error) { + return nil, io.EOF +} +func (c *racyConn) Invoke(context.Context, string, drpc.Encoding, drpc.Message, drpc.Message) error { + return io.EOF +} + +func TestPeer_ReleaseDoomedDuringCheck(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + fx.mc.EXPECT().Addr().Return("").AnyTimes() + a, b := net.Pipe() + defer a.Close() + defer b.Close() + rc := &racyConn{p: fx.peer, never: make(chan struct{}), ready: make(chan struct{})} + close(rc.ready) + sc := &subConn{ConnUnblocked: rc, LastUsageConn: connutil.NewLastUsageConn(a)} + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() + time.Sleep(10 * time.Millisecond) + + fx.ReleaseDrpcConn(ctx, sc) + require.True(t, sc.doomed.Load(), "gc doomed it during the release") + fx.mu.Lock() + assert.Empty(t, fx.inactive, "a doomed conn must not be re-pooled") + fx.mu.Unlock() +} + +func TestPeer_AcquireSkipsDoomedConns(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + a, b := net.Pipe() + defer a.Close() + defer b.Close() + rc := &racyConn{never: make(chan struct{}), ready: make(chan struct{})} + doomed := &subConn{ConnUnblocked: rc, LastUsageConn: connutil.NewLastUsageConn(a)} + doomed.doomed.Store(true) + fx.mu.Lock() + fx.inactive = append(fx.inactive, doomed) + fx.mu.Unlock() + + in, out := net.Pipe() + defer out.Close() + go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() + fx.mc.EXPECT().Open(gomock.Any()).Return(in, nil) + dc, err := fx.AcquireDrpcConn(ctx) + require.NoError(t, err) + assert.NotEqual(t, drpc.Conn(doomed), dc) +} + +// TestPeer_NilWakeKeepsThrottling: a waiter woken because a released conn +// was closed must not open at once while closes are still in flight +func TestPeer_NilWakeKeepsThrottling(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + defer close(release) + // closes stuck in flight on a slow transport, enough for the limiter + // to hold the next open for seconds + for i := 0; i < fx.limiter.startThreshold+20; i++ { + fx.cleanup.close(newBlockingCloser(release)) + } + + var opens atomic.Int32 + fx.mc.EXPECT().Open(gomock.Any()).DoAndReturn(func(context.Context) (net.Conn, error) { + opens.Add(1) + return nil, io.EOF + }).AnyTimes() + + actx, cancel := context.WithCancel(ctx) + defer cancel() + done := make(chan error, 1) + go func() { + _, err := fx.AcquireDrpcConn(actx) + done <- err + }() + // the waiter is throttled; a release wakes it with a nil conn + time.Sleep(50 * time.Millisecond) + fx.subConnRelease <- nil + require.Never(t, func() bool { return opens.Load() > 0 }, 300*time.Millisecond, 10*time.Millisecond, + "a nil wake must not bypass the limiter") + cancel() + select { + case err := <-done: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("acquire did not return on ctx") + } +} diff --git a/net/peer/peer.go b/net/peer/peer.go index d95713299..3ccc98919 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -3,6 +3,7 @@ package peer import ( "context" + "errors" "io" "net" "slices" @@ -56,6 +57,7 @@ func NewPeer(mc transport.MultiConn, ctrl connCtrl) (p Peer, err error) { if pr.id, err = CtxPeerId(ctx); err != nil { return } + pr.cleanup = newCleanupOwner(mc) go pr.acceptLoop() return pr, nil } @@ -96,6 +98,10 @@ type Peer interface { type subConn struct { encoding.ConnUnblocked *connutil.LastUsageConn + // doomed is set by gc when it takes an active conn away and hands its + // close to the cleanup owner: the holder must not return it for reuse, + // even though the close may not have landed yet + doomed atomic.Bool } func (s *subConn) Unblocked() <-chan struct{} { @@ -123,6 +129,9 @@ type peer struct { limiter limiter + // cleanup closes sub connections off the callers' path + cleanup *cleanupOwner + mu sync.Mutex created time.Time useSnappy bool @@ -135,12 +144,26 @@ func (p *peer) Id() string { } func (p *peer) AcquireDrpcConn(ctx context.Context) (drpc.Conn, error) { + for { + conn, retry, err := p.acquireDrpcConn(ctx) + if !retry { + return conn, err + } + } +} + +// acquireDrpcConn makes one acquisition attempt; retry means start over +// with a fresh look at the pool and a fresh limiter wait +func (p *peer) acquireDrpcConn(ctx context.Context) (conn drpc.Conn, retry bool, err error) { if p.IsClosed() { - return nil, transport.ErrConnClosed + return nil, false, transport.ErrConnClosed } p.mu.Lock() if len(p.inactive) == 0 { - wait := p.limiter.wait(len(p.active) + int(p.openingWaitCount.Load())) + // closes still in progress count too: they used to run on the + // releasing callers' goroutines while the conn was still active, + // and the owner must not turn into a way around the throttling + wait := p.limiter.wait(len(p.active) + int(p.openingWaitCount.Load()) + p.cleanup.inFlight()) p.openingWaitCount.Add(1) defer p.openingWaitCount.Add(-1) p.mu.Unlock() @@ -148,18 +171,24 @@ func (p *peer) AcquireDrpcConn(ctx context.Context) (drpc.Conn, error) { // throttle new connection opening select { case <-ctx.Done(): - return nil, ctx.Err() + return nil, false, ctx.Err() case dconn := <-p.subConnRelease: // nil conn means connection was closed, used to wake up AcquireDrpcConn - if dconn != nil { - return dconn, nil + if dconn != nil && !isDoomed(dconn) { + return dconn, false, nil } + // The released conn was closed, or gc doomed it on the way. + // Its close may still be in flight in the cleanup owner, so + // opening right away would bypass the throttling: start + // over, which picks up an inactive conn or recomputes the + // wait with the closes still in flight. + return nil, true, nil case <-wait: } } dconn, err := p.openDrpcConn(ctx) if err != nil { - return nil, err + return nil, false, err } p.mu.Lock() p.inactive = append(p.inactive, dconn) @@ -170,51 +199,33 @@ func (p *peer) AcquireDrpcConn(ctx context.Context) (drpc.Conn, error) { select { case <-res.Closed(): p.mu.Unlock() - return p.AcquireDrpcConn(ctx) + return nil, true, nil default: } + if res.doomed.Load() { + // gc took it while it was being re-pooled; its close is pending + p.mu.Unlock() + return nil, true, nil + } p.active[res] = struct{}{} p.mu.Unlock() - return res, nil + return res, false, nil +} + +func isDoomed(conn drpc.Conn) bool { + sc, ok := conn.(*subConn) + return ok && sc.doomed.Load() } // ReleaseDrpcConn releases the connection back to the pool. // you should pass the same ctx you passed to AcquireDrpcConn func (p *peer) ReleaseDrpcConn(ctx context.Context, conn drpc.Conn) { var closed bool - select { - case <-conn.Closed(): - closed = true - case <-ctx.Done(): - // in case ctx is closed the connection may be not yet closed because of the signal logic in the drpc manager - // but, we want to shortcut to avoid race conditions - _ = conn.Close() + if isDoomed(conn) { + // gc has taken it out of active and owns its close closed = true - default: - if connCasted, ok := conn.(encoding.ConnUnblocked); ok { - select { - case <-conn.Closed(): - closed = true - case <-connCasted.Unblocked(): - // semi-safe to reuse this connection - // it may be still a chance that connection will be closed in next milliseconds - // but this is a trade-off for performance - case <-time.After(time.Second / 5): - // means the connection has some unfinished work, - // e.g. not fully read stream - // we cannot reuse this connection so let's close it - _ = conn.Close() - closed = true - } - } else { - // By construction, conns returned from AcquireDrpcConn are *subConn - // which embeds encoding.ConnUnblocked. Reaching this branch means - // the caller passed a foreign conn; close it defensively instead - // of crashing the process. - log.Warn("released conn does not implement encoding.ConnUnblocked, closing", zap.String("peerId", p.id)) - _ = conn.Close() - closed = true - } + } else { + closed = p.checkReleased(ctx, conn) } if !closed { @@ -237,6 +248,11 @@ func (p *peer) ReleaseDrpcConn(ctx context.Context, conn drpc.Conn) { delete(p.active, sc) } + if !closed && sc.doomed.Load() { + // gc doomed it after the check above; doomed is set under p.mu, so + // this re-check is final + closed = true + } if !closed { // put it back into the pool p.inactive = append(p.inactive, sc) @@ -254,6 +270,52 @@ func (p *peer) ReleaseDrpcConn(ctx context.Context, conn drpc.Conn) { } } +// checkReleased reports whether a released conn is closed or must not be +// reused; a conn that must not is closed in the background +func (p *peer) checkReleased(ctx context.Context, conn drpc.Conn) (closed bool) { + select { + case <-conn.Closed(): + closed = true + case <-ctx.Done(): + // in case ctx is closed the connection may be not yet closed because of the signal logic in the drpc manager + // but, we want to shortcut to avoid race conditions: the conn is never reused, and it is closed in the + // background, since a drpc close waits for its reader and the transport and the caller is past its deadline + select { + case <-conn.Closed(): + // both were ready: nothing left to close + default: + p.cleanup.close(conn) + } + closed = true + default: + if connCasted, ok := conn.(encoding.ConnUnblocked); ok { + select { + case <-conn.Closed(): + closed = true + case <-connCasted.Unblocked(): + // semi-safe to reuse this connection + // it may be still a chance that connection will be closed in next milliseconds + // but this is a trade-off for performance + case <-time.After(time.Second / 5): + // means the connection has some unfinished work, + // e.g. not fully read stream + // we cannot reuse this connection so let's close it + p.cleanup.close(conn) + closed = true + } + } else { + // By construction, conns returned from AcquireDrpcConn are *subConn + // which embeds encoding.ConnUnblocked. Reaching this branch means + // the caller passed a foreign conn; close it defensively instead + // of crashing the process. + log.Warn("released conn does not implement encoding.ConnUnblocked, closing", zap.String("peerId", p.id)) + p.cleanup.close(conn) + closed = true + } + } + return closed +} + func (p *peer) DoDrpc(ctx context.Context, do func(conn drpc.Conn) error) error { conn, err := p.AcquireDrpcConn(ctx) if err != nil { @@ -276,13 +338,10 @@ func (p *peer) openDrpcConn(ctx context.Context) (*subConn, error) { return nil, err } lastUsageConn := connutil.NewLastUsageConn(conn) - proto, err := handshake.OutgoingProtoHandshake(ctx, lastUsageConn, defaultHandshakeProto) + // on any error the handshake hands the stream to the cleanup owner, once: + // on a stalled transport the close blocks, so it never runs here + proto, err := handshake.OutgoingProtoHandshakeWithCloser(ctx, lastUsageConn, defaultHandshakeProto, p.closeSubConn) if err != nil { - // OutgoingProtoHandshake closes the conn on I/O errors and ctx - // cancellation, but returns without closing on some protocol-level - // errors (incompatible/declined/unexpected proto). Close here so the - // sub-stream we opened above never leaks. Double close is harmless. - _ = lastUsageConn.Close() return nil, err } bufSize := p.ctrl.DrpcConfig().Stream.MaxMsgSizeMb * (1 << 20) @@ -299,6 +358,10 @@ func (p *peer) openDrpcConn(ctx context.Context) (*subConn, error) { }, nil } +func (p *peer) closeSubConn(conn net.Conn) { + p.cleanup.close(conn) +} + func (p *peer) acceptLoop() { var exitErr error defer func() { @@ -324,7 +387,7 @@ func (p *peer) acceptLoop() { p.incomingCount.Add(1) defer p.incomingCount.Add(-1) serveErr := p.serve(conn) - if serveErr != io.EOF && serveErr != transport.ErrConnClosed { + if serveErr != io.EOF && !errors.Is(serveErr, transport.ErrConnClosed) { log.InfoCtx(p.Context(), "serve connection error", zap.Error(serveErr)) } }() @@ -388,11 +451,12 @@ func (p *peer) TryClose(objectTTL time.Duration) (res bool, err error) { func (p *peer) gc(ttl time.Duration) (aliveCount int) { // drpc conn Close blocks until its reader unwinds, which on a stalled stream // takes until the yamux stream close timeout: collect the doomed conns and - // close them after releasing the lock + // hand them to the cleanup owner after releasing the lock, so a stalled + // peer does not hold up the GC pass of every other peer var toClose []*subConn defer func() { for _, conn := range toClose { - _ = conn.Close() + p.cleanup.close(conn) } }() p.mu.Lock() @@ -430,6 +494,7 @@ func (p *peer) gc(ttl time.Duration) (aliveCount int) { } if act.LastUsage().Before(minLastUsage) { log.Warn("close active connection because no activity", zap.String("peerId", p.id), zap.String("addr", p.Addr())) + act.doomed.Store(true) toClose = append(toClose, act) delete(p.active, act) continue diff --git a/net/pool/pool.go b/net/pool/pool.go index 6b4c46c9d..339f5d87a 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -5,7 +5,11 @@ import ( "context" "fmt" "math/rand" + "sync" + "sync/atomic" + "time" + "github.com/prometheus/client_golang/prometheus" "go.uber.org/zap" "github.com/anyproto/any-sync/app/debugstat" @@ -22,7 +26,9 @@ type Pool interface { Get(ctx context.Context, id string) (peer.Peer, error) // GetOneOf searches at least one existing connection in outgoing or creates a new one from a randomly selected id from given list GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error) - // AddPeer adds incoming peer to the pool + // AddPeer adds incoming peer to the pool. The pool evicts peers by + // instance (ocache.RemoveSame), so peer.Peer implementations must be + // comparable (pointers) AddPeer(ctx context.Context, p peer.Peer) (err error) // Pick checks if a connection with the peer exists, without dialing. // For a peer whose last dial failed it returns the cached dial error. @@ -35,9 +41,46 @@ type poolStats struct { PeerStats []*peer.Stat `json:"peerStats"` } +// caches is one immutable pair of peer caches. Flush replaces the whole pair, +// so a lookup that snapshots it never mixes a new incoming cache with an old +// outgoing one. +type caches struct { + incoming ocache.OCache + outgoing ocache.OCache + // the same two caches as ocache.Peeker, asserted once at build (ocache.New + // always returns one) so the hit path does no type assertion + peekIncoming ocache.Peeker + peekOutgoing ocache.Peeker + // ctx is cancelled the moment the pair stops being current (swap or pool + // Close), before its caches close; lookup merges it into the caller's ctx + // so a lookup blocked on this pair (a dial, a wait behind a GC TryClose) + // is cut short and retried on the current one + ctx context.Context + cancel context.CancelFunc +} + type pool struct { - outgoing ocache.OCache - incoming ocache.OCache + // current is the published cache pair. Lookups load it once per call and + // re-check it afterwards (see lookup); Flush swaps it under swapMu and + // closes the replaced pair in the background; Close leaves it in place, + // closed, so lookups on a closed pool fail with ocache.ErrClosed. + current atomic.Pointer[caches] + // swapMu orders the swap against AddPeer (read side, see addIncoming) and + // against Close; closed is set under it, after which Flush is a no-op + swapMu sync.RWMutex + closed bool + // closing counts the replaced pairs whose teardown (caches and peers) is + // still running in the background; Close waits for them, bounded + closing sync.WaitGroup + // newCaches builds a fresh pair with its own ctx; set by the service at Init + newCaches func() *caches + // closeTimeout bounds the ocache close passes and Close as a whole + closeTimeout time.Duration + // incomingMiss is the incoming cache's miss counter (nil without a + // registry): Get counts a miss there when the fast path finds the peer + // in outgoing, as a Get through the caches would + incomingMiss prometheus.Counter + statService debugstat.StatService closingCtx context.Context closingCancel context.CancelFunc @@ -50,10 +93,96 @@ func (p *pool) Name() (name string) { return CName } +// lookup runs f against the current pair and repeats it while the pair is +// swapped underneath: once Flush has published a new pair, a result from the +// replaced one is a pre-flush peer, or an error from a dial the swap cut +// short, and neither may reach the caller. f gets a ctx that is also +// cancelled by the swap, so a lookup blocked on the old pair does not wait +// out a dead dial. One retry is not enough under back-to-back flushes, so +// this loops until a result comes from the pair that is still current +// afterwards. Bounded by ctx: if flushes keep coming faster than a dial +// completes no lookup can succeed, and the caller gets its ctx error (the +// recovery worker does not flush in a tight loop). ErrClosed reaches the +// caller only when the pool itself is closed. +func (p *pool) lookup(ctx context.Context, f func(ctx context.Context, c *caches) (peer.Peer, error)) (peer.Peer, error) { + for { + c := p.current.Load() + pr, err := func() (peer.Peer, error) { + // deferred so a loader panic re-raised by ocache does not leave + // the registration on the pair ctx behind + lctx, cancel := context.WithCancel(ctx) + defer cancel() + defer context.AfterFunc(c.ctx, cancel)() + return f(lctx, c) + }() + // read before re-checking current: Flush publishes the new pair before + // it cancels the old one, so a pair found cancelled and then still + // current can only have been cancelled by Close + cancelled := c.ctx.Err() != nil + if p.current.Load() == c { + if cancelled && ctx.Err() == nil { + return nil, ocache.ErrClosed + } + return pr, err + } + if err = ctx.Err(); err != nil { + return nil, err + } + } +} + +// fast is the hit path: a live peer already loaded in the current pair is +// returned without blocking, allocating or touching the pair ctx (servers, +// which never flush, pay for this on every call). Anything else — a miss, an +// entry still loading or closing, a cached dial error, a closed peer, a pair +// swapped while reading — returns nil and the caller takes lookup, which +// handles those cases. Rechecking current after the read gives the same +// guarantee as lookup's post-check: the peer comes from a pair that was +// current after it was read. touch refreshes the GC deadline (Get) or not +// (Pick); hits are counted either way, as the caches themselves do. Get's +// incoming miss is counted too, so the series are the same as before. +func (p *pool) fast(id string, touch bool) peer.Peer { + c := p.current.Load() + v, ok := c.peekIncoming.Peek(id, touch) + if !ok { + if v, ok = c.peekOutgoing.Peek(id, touch); !ok { + return nil + } + if touch && p.incomingMiss != nil { + p.incomingMiss.Inc() + } + } + if pr, isPeer := v.(peer.Peer); isPeer && !pr.IsClosed() && p.current.Load() == c { + return pr + } + return nil +} + +// discard closes pr (if not closed yet) and evicts it from source, in the +// background and outside any pool or cache lock, so the caller never waits +// past its ctx on an entry a GC TryClose holds or on a transport teardown +// (RemoveSame closes the value itself; a second Close is idempotent). +// RemoveSame never touches a replacement installed under the same id. The +// returned channel closes once the attempt is done. +func (p *pool) discard(source ocache.OCache, pr peer.Peer) <-chan struct{} { + done := make(chan struct{}) + go func() { + defer close(done) + if p.closingCtx.Err() != nil { + // pool shutdown: cache.Close evicts whatever is left + return + } + _, _ = source.RemoveSame(p.closingCtx, pr.Id(), pr) + }() + return done +} + // evictOnClose removes the peer from the cache as soon as its underlying // connection dies, instead of waiting for the next Get or the GC to notice. // When the whole pool is shutting down, cache.Close already evicts every peer, -// so per-peer removal is skipped. It never outlives the peer. +// so per-peer removal is skipped. It never outlives the peer. cache is the +// instance the peer was published into: after a Flush that is no longer the +// current one, and RemoveSame on it fails fast with ErrClosed. func (p *pool) evictOnClose(pr peer.Peer, cache ocache.OCache, inbound bool) { select { case <-pr.CloseChan(): @@ -89,73 +218,140 @@ func (p *pool) evictOnClose(pr peer.Peer, cache ocache.OCache, inbound bool) { }) } -func (p *pool) Get(ctx context.Context, id string) (pr peer.Peer, err error) { - // if we have incoming connection - try to reuse it - if pr, err = p.get(ctx, p.incoming, id); err != nil { - // or try to get or create outgoing - return p.get(ctx, p.outgoing, id) - } - return -} - -func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.Peer, error) { - v, err := source.Get(ctx, id) - if err != nil { - return nil, err - } - pr, err := getPeer(v) - if err != nil { - return nil, err - } - if !pr.IsClosed() { +func (p *pool) Get(ctx context.Context, id string) (peer.Peer, error) { + if pr := p.fast(id, true); pr != nil { return pr, nil } - _, _ = source.Remove(ctx, id) - return p.Get(ctx, id) + return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { + // if we have incoming connection - try to reuse it + if pr, err = p.get(ctx, c.incoming, id); err != nil { + // or try to get or create outgoing + return p.get(ctx, c.outgoing, id) + } + return + }) } -func (p *pool) Flush(ctx context.Context) error { - p.incoming.ForEach(func(v ocache.Object) (isContinue bool) { - pr, err := getPeer(v) +func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.Peer, error) { + for { + v, err := source.Get(ctx, id) if err != nil { - return true + return nil, err } - _, _ = p.incoming.Remove(ctx, pr.Id()) - return true - }) - p.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { pr, err := getPeer(v) if err != nil { - return true + return nil, err + } + if !pr.IsClosed() { + return pr, nil + } + // The entry must be gone before redialing or source.Get would return + // the same instance again: wait (bounded by ctx) for the background + // discard, so a teardown that blocks never runs on this path. + select { + case <-p.discard(source, pr): + case <-ctx.Done(): + return nil, ctx.Err() + } + // with a done ctx source.Get can return the closed value again + if err = ctx.Err(); err != nil { + return nil, err + } + } +} + +// Flush invalidates every pooled connection: it builds a fresh cache pair, +// publishes it and closes the replaced pair in the background, so it never +// waits on a GC TryClose or on transport teardown. From the moment it returns +// no lookup hands out a pre-flush peer (see lookup), and a dial that was in +// flight is cancelled instead of being waited out. Cached dial errors go with +// the old pair, except incompatible-version verdicts (see errObject), which +// are carried over so their backoff survives (one still loading at the swap +// is not: its verdict lands in the old pair and the fresh one redials once). +// A no-op once the pool is closed; concurrent flushes are serialized, so each +// replaced pair is closed exactly once. +func (p *pool) Flush(ctx context.Context) error { + p.swapMu.Lock() + if p.closed { + p.swapMu.Unlock() + return nil + } + old := p.current.Load() + fresh := p.newCaches() + old.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { + if eo, ok := v.(*errObject); ok && eo.keepOnFlush() { + // cheap and non-blocking: a fresh cache has no closers + _ = fresh.outgoing.Add(eo.id, eo) } - _, _ = p.outgoing.Remove(ctx, pr.Id()) return true }) + p.current.Store(fresh) + old.cancel() + // under swapMu like closed: Close sets closed and then waits, so no Add + // can follow its Wait + p.closing.Add(1) + p.swapMu.Unlock() + go func() { + defer p.closing.Done() + peers, _ := closeCaches(old) + peers.Wait() + }() return nil } +// closeCaches tears down a pair that is no longer current. Each loaded peer is +// closed on its own goroutine first, because ocache.Close closes entries one +// at a time with no ctx and one hung teardown would hold the rest back; the +// Close passes then cancel the in-flight dials (outgoing first, so a hung +// incoming peer cannot delay that) and close the peers a second time, which +// peer.Close tolerates (the pool relies on that already, see discard). A +// RemoveSame per peer would not do: once a cache is marked closed every +// RemoveSame is refused. Known gap: a peer a GC TryClose holds past +// closeTimeout closes only when TryClose returns (ocache escalates the +// decline), never if it never returns. Returns once the caches are closed; +// the WaitGroup tracks the per-peer closes still running, and err is the +// outgoing cache's close error. +func closeCaches(c *caches) (peers *sync.WaitGroup, err error) { + peers = &sync.WaitGroup{} + for _, cache := range []ocache.OCache{c.outgoing, c.incoming} { + cache.ForEach(func(v ocache.Object) (isContinue bool) { + if pr, ok := v.(peer.Peer); ok { + peers.Add(1) + go func() { + defer peers.Done() + _ = pr.Close() + }() + } + return true + }) + } + err = c.outgoing.Close() + if e := c.incoming.Close(); e != nil { + log.Warn("close incoming cache error", zap.Error(e)) + } + return peers, err +} + func (p *pool) getIfActive(ctx context.Context, peerIds []string) peer.Peer { for _, peerId := range peerIds { - // a cached errObject (failed dial) only disqualifies this peerId, - // not the rest of the scan - if v, err := p.incoming.Pick(ctx, peerId); err == nil { - if pr, err := getPeer(v); err == nil { - if !pr.IsClosed() { - return pr - } - _, _ = p.incoming.Remove(ctx, peerId) - } + if pr := p.fast(peerId, false); pr != nil { + return pr } - if v, err := p.outgoing.Pick(ctx, peerId); err == nil { - if pr, err := getPeer(v); err == nil { - if !pr.IsClosed() { - return pr - } - _, _ = p.outgoing.Remove(ctx, peerId) + } + pr, _ := p.lookup(ctx, func(ctx context.Context, c *caches) (peer.Peer, error) { + for _, peerId := range peerIds { + // a cached errObject (failed dial) only disqualifies this peerId, + // not the rest of the scan + if pr, err := p.pick(ctx, c.incoming, peerId); err == nil { + return pr, nil + } + if pr, err := p.pick(ctx, c.outgoing, peerId); err == nil { + return pr, nil } } - } - return nil + return nil, errPeerNotFound + }) + return pr } func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error) { @@ -188,28 +384,90 @@ func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error return nil, lastErr } -func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) (err error) { - err = p.incoming.Add(pr.Id(), pr) - if err == ocache.ErrExists { - // in case when an incoming connection with a peer already exists, we close and remove an existing connection - if v, e := p.incoming.Pick(ctx, pr.Id()); e == nil { - _ = v.Close() - _, _ = p.incoming.Remove(ctx, pr.Id()) - err = p.incoming.Add(pr.Id(), pr) +// AddPeer adds an incoming peer. pr must be of a comparable type (a pointer): +// the pool evicts it by instance. +func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { + // bounds the retries on an entry that is neither pickable nor gone: one + // mid-close, whose teardown may be slow (then ErrExists, as before) + const retries = 3 + attempts := 0 + for { + c, err := p.addIncoming(pr) + if err != ocache.ErrExists { + return err + } + // An incoming connection with this peer already exists: close and + // remove it, then add again. This runs outside swapMu because the + // removal can wait on a GC TryClose that holds the entry, which must + // not stall a Flush and every AddPeer queued behind it. + v, e := c.incoming.Pick(ctx, pr.Id()) + if e != nil { + if err = ctx.Err(); err != nil { + return err + } + if p.current.Load() != c { + // the pair was flushed meanwhile: add to the current one + continue + } + if e == ocache.ErrClosed { + return e + } + // The entry was a transient one: a concurrent Get(id) creates a + // loading entry that the incoming loader fails with ErrNotExists + // (Pick waited that out), or the previous connection is mid-close. + if attempts++; attempts <= retries { + continue + } + return ocache.ErrExists + } + // The old connection's teardown (Close, then the instance-safe + // removal that never touches a replacement) runs in the background + // and is waited for only as long as ctx allows: a hung transport must + // not stall the accept path. + old, isPeer := v.(peer.Peer) + if !isPeer { + _, _ = c.incoming.RemoveSame(ctx, pr.Id(), v) + } else { + select { + case <-p.discard(c.incoming, old): + case <-ctx.Done(): + return ctx.Err() + } + } + if err = ctx.Err(); err != nil { + return err } } - if err == nil { - go p.evictOnClose(pr, p.incoming, true) +} + +// addIncoming adds pr to the current incoming cache and returns the pair it +// used. The read lock spans the Add, which never blocks, so a peer accepted +// after a Flush published its pair can only land in that pair, never in one +// that is about to be closed; Add therefore fails with ErrClosed only when the +// pool is closed (Close leaves the closed pair current). Returns ErrExists +// without starting a watcher. +func (p *pool) addIncoming(pr peer.Peer) (*caches, error) { + p.swapMu.RLock() + defer p.swapMu.RUnlock() + c := p.current.Load() + if err := c.incoming.Add(pr.Id(), pr); err != nil { + return c, err } - return err + go p.evictOnClose(pr, c.incoming, true) + return c, nil } func (p *pool) Pick(ctx context.Context, id string) (pr peer.Peer, err error) { - // check if connection with peer exist without dial - if pr, err = p.pick(ctx, p.incoming, id); err != nil { - return p.pick(ctx, p.outgoing, id) + if pr = p.fast(id, false); pr != nil { + return pr, nil } - return + return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { + // check if connection with peer exist without dial + if pr, err = p.pick(ctx, c.incoming, id); err != nil { + return p.pick(ctx, c.outgoing, id) + } + return + }) } func (p *pool) pick(ctx context.Context, source ocache.OCache, id string) (peer.Peer, error) { @@ -226,18 +484,20 @@ func (p *pool) pick(ctx context.Context, source ocache.OCache, id string) (peer. if !pr.IsClosed() { return pr, nil } - return nil, fmt.Errorf("failed to pick connection with peer: peer not found") + p.discard(source, pr) + return nil, errPeerNotFound } func (p *pool) ProvideStat() any { peerStats := make([]*peer.Stat, 0) - p.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { + c := p.current.Load() + c.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { if p, ok := v.(peer.StatProvider); ok { peerStats = append(peerStats, p.ProvideStat()) } return true }) - p.incoming.ForEach(func(v ocache.Object) (isContinue bool) { + c.incoming.ForEach(func(v ocache.Object) (isContinue bool) { if p, ok := v.(peer.StatProvider); ok { peerStats = append(peerStats, p.ProvideStat()) } @@ -254,6 +514,8 @@ func (p *pool) StatType() string { return CName } +var errPeerNotFound = fmt.Errorf("failed to pick connection with peer: peer not found") + func getPeer(val ocache.Object) (pr peer.Peer, err error) { switch v := val.(type) { case peer.Peer: diff --git a/net/pool/pool_bench_test.go b/net/pool/pool_bench_test.go new file mode 100644 index 000000000..f9ce64e6d --- /dev/null +++ b/net/pool/pool_bench_test.go @@ -0,0 +1,116 @@ +package pool + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/prometheus/client_golang/prometheus" + + "github.com/anyproto/any-sync/app" + "github.com/anyproto/any-sync/metric" + "github.com/anyproto/any-sync/net/peer" +) + +// The hit path: the peer is already pooled, no Flush runs, a prometheus +// registry is attached as on the servers. Comparable across branches, so the +// helpers below depend only on the fixtures that exist on main (testPeer, +// dialerMock). + +type benchMetric struct { + metric.Metric + reg *prometheus.Registry +} + +func (m *benchMetric) Init(a *app.App) error { m.reg = prometheus.NewRegistry(); return nil } +func (m *benchMetric) Name() string { return metric.CName } +func (m *benchMetric) Run(ctx context.Context) error { return nil } +func (m *benchMetric) Close(ctx context.Context) error { return nil } +func (m *benchMetric) Registry() *prometheus.Registry { return m.reg } + +// benchPeer never closes itself through TryClose, so the GC cannot evict it +// mid-benchmark, and its Close is idempotent like the real peer's +type benchPeer struct { + *testPeer + once sync.Once +} + +func newBenchPeer(id string) *benchPeer { + return &benchPeer{testPeer: newTestPeer(id)} +} + +func (p *benchPeer) Close() error { + p.once.Do(func() { _ = p.testPeer.Close() }) + return nil +} + +func (p *benchPeer) TryClose(time.Duration) (bool, error) { return false, nil } + +func benchPool(b *testing.B, op func(ctx context.Context, s Service) error) { + s := New() + a := new(app.App) + a.Register(s) + a.Register(&dialerMock{dial: func(ctx context.Context, id string) (peer.Peer, error) { + return newBenchPeer(id), nil + }}) + a.Register(&benchMetric{}) + if err := a.Start(context.Background()); err != nil { + b.Fatal(err) + } + defer func() { _ = a.Close(context.Background()) }() + ctx := context.Background() + if _, err := s.Get(ctx, "out"); err != nil { + b.Fatal(err) + } + if err := s.AddPeer(ctx, newBenchPeer("in")); err != nil { + b.Fatal(err) + } + b.Run("serial", func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if err := op(ctx, s); err != nil { + b.Fatal(err) + } + } + }) + b.Run("parallel", func(b *testing.B) { + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + if err := op(ctx, s); err != nil { + b.Error(err) + return + } + } + }) + }) +} + +func BenchmarkPool_GetIncoming(b *testing.B) { + benchPool(b, func(ctx context.Context, s Service) error { _, err := s.Get(ctx, "in"); return err }) +} + +func BenchmarkPool_GetOutgoing(b *testing.B) { + benchPool(b, func(ctx context.Context, s Service) error { _, err := s.Get(ctx, "out"); return err }) +} + +func BenchmarkPool_PickIncoming(b *testing.B) { + benchPool(b, func(ctx context.Context, s Service) error { _, err := s.Pick(ctx, "in"); return err }) +} + +func BenchmarkPool_PickOutgoing(b *testing.B) { + benchPool(b, func(ctx context.Context, s Service) error { _, err := s.Pick(ctx, "out"); return err }) +} + +func BenchmarkPool_GetOneOf(b *testing.B) { + ids := []string{"x1", "x2", "out"} + benchPool(b, func(ctx context.Context, s Service) error { _, err := s.GetOneOf(ctx, ids); return err }) +} + +// servers pass request contexts, which are cancellable +func BenchmarkPool_GetIncomingCancellableCtx(b *testing.B) { + rctx, cancel := context.WithTimeout(context.Background(), time.Hour) + defer cancel() + benchPool(b, func(_ context.Context, s Service) error { _, err := s.Get(rctx, "in"); return err }) +} diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go new file mode 100644 index 000000000..61bbd4588 --- /dev/null +++ b/net/pool/pool_flush_test.go @@ -0,0 +1,1567 @@ +package pool + +import ( + "context" + "fmt" + "runtime" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + atomic2 "go.uber.org/atomic" + + "github.com/anyproto/any-sync/app" + "github.com/anyproto/any-sync/app/ocache" + "github.com/anyproto/any-sync/metric" + "github.com/anyproto/any-sync/net/peer" + "github.com/anyproto/any-sync/net/peerobserver" + "github.com/anyproto/any-sync/net/secureservice/handshake" +) + +// ctlPeer is a test peer whose TryClose, Close and IsClosed can be paused or +// made to decline, to drive the races between Flush, the ocache GC and the +// lookups +type ctlPeer struct { + *testPeer + tryClose func() (bool, error) + closeHook func(call int32) + isClosedHook func(call int32) + // idHook pauses inside addIncoming's Add, between the snapshot and the + // insert (Id is the first thing the pool asks a new incoming peer) + idHook func(call int32) + closeCalls atomic.Int32 + isClosedCall atomic.Int32 + idCall atomic.Int32 +} + +func newCtlPeer(id string) *ctlPeer { + return &ctlPeer{testPeer: newTestPeer(id)} +} + +func (c *ctlPeer) TryClose(objectTTL time.Duration) (bool, error) { + if c.tryClose != nil { + return c.tryClose() + } + return c.testPeer.TryClose(objectTTL) +} + +func (c *ctlPeer) Close() error { + call := c.closeCalls.Add(1) + if c.closeHook != nil { + c.closeHook(call) + } + return c.testPeer.Close() +} + +func (c *ctlPeer) IsClosed() bool { + call := c.isClosedCall.Add(1) + if c.isClosedHook != nil { + c.isClosedHook(call) + } + return c.testPeer.IsClosed() +} + +func (c *ctlPeer) Id() string { + call := c.idCall.Add(1) + if c.idHook != nil { + c.idHook(call) + } + return c.testPeer.Id() +} + +var _ peer.Peer = (*ctlPeer)(nil) + +// newRelease returns a gate channel and an idempotent func that opens it +func newRelease() (chan struct{}, func()) { + ch := make(chan struct{}) + var once sync.Once + return ch, func() { once.Do(func() { close(ch) }) } +} + +// pairTracker wraps the pool's cache factory to number every pair it builds +type pairTracker struct { + mu sync.Mutex + seq map[*caches]int64 + pairs []*caches +} + +func trackPairs(p *pool) *pairTracker { + tr := &pairTracker{seq: map[*caches]int64{}} + cur := p.current.Load() + tr.seq[cur] = 0 + tr.pairs = append(tr.pairs, cur) + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + tr.mu.Lock() + defer tr.mu.Unlock() + tr.seq[c] = int64(len(tr.pairs)) + tr.pairs = append(tr.pairs, c) + return c + } + return tr +} + +func (tr *pairTracker) seqOf(c *caches) int64 { + tr.mu.Lock() + defer tr.mu.Unlock() + return tr.seq[c] +} + +func (tr *pairTracker) created() int { + tr.mu.Lock() + defer tr.mu.Unlock() + return len(tr.pairs) +} + +// all returns a copy of the pairs built so far, oldest first +func (tr *pairTracker) all() []*caches { + tr.mu.Lock() + defer tr.mu.Unlock() + return append([]*caches(nil), tr.pairs...) +} + +// isClosed reports whether both caches of the pair are closed, without +// touching them: Pick fails with ErrClosed on a closed cache and with +// ErrNotExists on an open one +func (c *caches) isClosed() bool { + _, inErr := c.incoming.Pick(ctx, "\x00probe") + _, outErr := c.outgoing.Pick(ctx, "\x00probe") + return inErr == ocache.ErrClosed && outErr == ocache.ErrClosed +} + +func inCurrent(p *pool, pr peer.Peer) bool { + v, err := p.current.Load().incoming.Pick(ctx, pr.Id()) + return err == nil && v == ocache.Object(pr) +} + +// startFlusher flushes the pool every period on a background goroutine until +// stop is called; stop also waits for the goroutine. flushes counts the +// completed flushes. +func startFlusher(t *testing.T, fx *fixture, period time.Duration) (flushes *atomic.Int64, stop func()) { + flushes = &atomic.Int64{} + quit := make(chan struct{}) + done := make(chan struct{}) + go func() { + defer close(done) + for { + select { + case <-quit: + return + default: + } + assert.NoError(t, fx.Flush(ctx)) + flushes.Add(1) + time.Sleep(period) + } + }() + var once sync.Once + return flushes, func() { + once.Do(func() { + close(quit) + <-done + }) + } +} + +// fromClosePrepass reports whether the caller runs on one of closeCaches's +// per-peer goroutines (the parallel teardown) rather than on the cache's own +// close pass, which closes the same peer a second time on the closing +// goroutine; test seam for "Close waits for the parallel closes" +func fromClosePrepass() bool { + buf := make([]byte, 1<<14) + n := runtime.Stack(buf, false) + return strings.Contains(string(buf[:n]), "pool.closeCaches.func") +} + +// hookCtx is a pair ctx whose Err can be paused: a seam between lookup's f +// returning and its read of the pair state +type hookCtx struct { + context.Context + onErr func() +} + +func (c *hookCtx) Err() error { + if c.onErr != nil { + c.onErr() + } + return c.Context.Err() +} + +func TestPool_FlushSwap(t *testing.T) { + t.Run("peer restored by a declined TryClose across flush is never returned", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + old := newCtlPeer("p1") + inTryClose := make(chan struct{}) + releaseTryClose, doReleaseTryClose := newRelease() + // released on every exit, so a failing assertion fails instead of + // hanging fx.Finish behind the paused hook + defer doReleaseTryClose() + old.tryClose = func() (bool, error) { + close(inTryClose) + <-releaseTryClose + return false, nil // a recently used sub conn keeps the peer alive + } + fresh := newTestPeer("p1") + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + return old, nil + } + return fresh, nil + } + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.Equal(t, peer.Peer(old), pr) + oldPair := p.current.Load() + + // the GC path: the entry is held in closing while TryClose runs + gcDone := make(chan struct{}) + go func() { + defer close(gcDone) + _, _ = oldPair.outgoing.TryRemove("p1") + }() + <-inTryClose + + // the swap does not wait on the closer holding the entry + flushed := make(chan error, 1) + go func() { flushed <- fx.Flush(ctx) }() + select { + case err = <-flushed: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("Flush waited on TryClose") + } + + doReleaseTryClose() + <-gcDone + + // the decline lands in a closed cache: escalated to a Close there, + // and the current pair never saw the peer + _, err = fx.Pick(ctx, "p1") + require.Error(t, err) + require.Nil(t, p.getIfActive(ctx, []string{"p1"})) + pr, err = fx.Get(ctx, "p1") + require.NoError(t, err) + assert.Equal(t, peer.Peer(fresh), pr) + pr, err = fx.GetOneOf(ctx, []string{"p1"}) + require.NoError(t, err) + assert.Equal(t, peer.Peer(fresh), pr) + require.Eventually(t, old.IsClosed, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return oldPair.outgoing.Len() == 0 }, time.Second, 10*time.Millisecond) + }) + t.Run("dial published after flush is closed without any lookup", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + late := newTestPeer("p1") + dialStarted := make(chan struct{}) + releaseDial, doReleaseDial := newRelease() + // released on every exit, so a failing assertion fails instead of + // hanging fx.Finish behind the paused hook + defer doReleaseDial() + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + close(dialStarted) + <-releaseDial // ignores the cancel old.Close sends + return late, nil + } + oldPair := p.current.Load() + loaded := make(chan error, 1) + go func() { + // a background loader whose result nobody looks at + _, err := oldPair.outgoing.Get(ctx, "p1") + loaded <- err + }() + <-dialStarted + require.NoError(t, fx.Flush(ctx)) + doReleaseDial() + // either the dial lands in the already closed pair (the load closes + // it and reports ErrClosed) or it is published a moment before the + // background Close, whose pass closes it + <-loaded + require.Eventually(t, late.IsClosed, time.Second, 10*time.Millisecond) + _, err := fx.Pick(ctx, "p1") + require.Error(t, err) + require.Eventually(t, func() bool { return oldPair.outgoing.Len() == 0 }, time.Second, 10*time.Millisecond) + require.Equal(t, 0, p.current.Load().outgoing.Len()) + // the rejected peer still gets exactly one Closed event + require.Eventually(t, func() bool { return len(obs.getClosed()) == 1 }, time.Second, 10*time.Millisecond) + assert.False(t, obs.getClosed()[0].Inbound) + require.Never(t, func() bool { return len(obs.getClosed()) > 1 }, 100*time.Millisecond, 10*time.Millisecond) + }) + t.Run("peer held by a declined TryClose is closed once TryClose returns", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + old := newCtlPeer("p1") + inTryClose := make(chan struct{}) + releaseTryClose, doReleaseTryClose := newRelease() + // released on every exit, so a failing assertion fails instead of + // hanging fx.Finish behind the paused hook + defer doReleaseTryClose() + old.tryClose = func() (bool, error) { + close(inTryClose) + <-releaseTryClose + return false, nil + } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return old, nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + oldPair := p.current.Load() + + gcDone := make(chan struct{}) + go func() { + defer close(gcDone) + _, _ = oldPair.outgoing.TryRemove("p1") + }() + <-inTryClose + require.NoError(t, fx.Flush(ctx)) + // accepted trade-off: the closer holding the entry owns the close, + // so the peer stays open until TryClose returns; the GC then closes + // it into the closed cache instead of restoring it + require.Never(t, old.IsClosed, 50*time.Millisecond, 10*time.Millisecond) + doReleaseTryClose() + <-gcDone + require.Eventually(t, old.IsClosed, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return oldPair.outgoing.Len() == 0 }, time.Second, 10*time.Millisecond) + }) + t.Run("concurrent waiters on a load spanning flush redial once", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + + late := newTestPeer("p1") + dialStarted := make(chan struct{}) + releaseDial, doReleaseDial := newRelease() + // released on every exit, so a failing assertion fails instead of + // hanging fx.Finish behind the paused hook + defer doReleaseDial() + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + close(dialStarted) + <-releaseDial + return late, nil + } + // a new peer per dial, like a real dialer: a dial the replaced + // pair cancels has its result closed, and handing that same + // object out again would never produce an open peer + return newTestPeer(peerId), nil + } + const waiters = 10 + results := make(chan peer.Peer, waiters) + get := func() { + pr, err := fx.Get(ctx, "p1") + assert.NoError(t, err) + results <- pr + } + go get() + <-dialStarted + for i := 1; i < waiters; i++ { + go get() + } + require.NoError(t, fx.Flush(ctx)) + doReleaseDial() + var fresh peer.Peer + for i := 0; i < waiters; i++ { + select { + case pr := <-results: + require.NotNil(t, pr) + assert.NotSame(t, late, pr, "a stale peer was returned") + if fresh == nil { + fresh = pr + } + assert.Same(t, fresh, pr, "all waiters share the redial on the fresh pair") + case <-time.After(5 * time.Second): + t.Fatal("waiter did not return") + } + } + // one redial on the fresh pair; in the window between the swap and + // the background Close a redial can also start on the replaced pair, + // which that Close then cancels + assert.GreaterOrEqual(t, dials.Load(), int32(2)) + assert.LessOrEqual(t, dials.Load(), int32(3)) + require.Eventually(t, late.IsClosed, time.Second, 10*time.Millisecond) + assert.False(t, fresh.IsClosed()) + }) + t.Run("get issued after flush does not wait on the pre-flush dial", func(t *testing.T) { + // the gap the generation design had: a Get after Flush joined the + // in-flight pre-flush dial and waited it out + fx := newFixture(t) + defer fx.Finish() + + var dials atomic.Int32 + releaseDial, doReleaseDial := newRelease() + defer doReleaseDial() + var cancelled atomic.Bool + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + // a dial on a dead path: returns only when cancelled + select { + case <-releaseDial: + case <-ctx.Done(): + cancelled.Store(true) + return nil, ctx.Err() + } + } + return newTestPeer(peerId), nil + } + first := make(chan error, 1) + go func() { + _, err := fx.Get(ctx, "p1") + first <- err + }() + require.Eventually(t, func() bool { return dials.Load() == 1 }, time.Second, time.Millisecond) + require.NoError(t, fx.Flush(ctx)) + + gctx, cancel := context.WithTimeout(ctx, 2*time.Second) + defer cancel() + start := time.Now() + pr, err := fx.Get(gctx, "p1") + require.NoError(t, err) + require.NotNil(t, pr) + require.Less(t, time.Since(start), 500*time.Millisecond, "waited on the stale dial") + // the old pair cancelled the stale dial and the first caller + // redialed too, sharing the fresh peer + require.NoError(t, <-first) + require.True(t, cancelled.Load()) + require.Equal(t, int32(2), dials.Load()) + }) + t.Run("live peer returned by the replaced pair is retried on the current one", func(t *testing.T) { + // the post-check: between the swap and old.Close the old pair still + // hands out live pre-flush peers; a lookup that took its result from + // there must not return it + fx := newFixture(t) + defer fx.Finish() + + old := newCtlPeer("p1") + inIsClosed := make(chan struct{}) + releaseIsClosed, doReleaseIsClosed := newRelease() + defer doReleaseIsClosed() + old.isClosedHook = func(call int32) { + if call == 1 { + close(inIsClosed) + <-releaseIsClosed + } + } + // the teardown of old is held too, so it stays a live stale peer + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + old.closeHook = func(int32) { <-releaseClose } + fresh := newTestPeer("p1") + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + return old, nil + } + return fresh, nil + } + res := make(chan peer.Peer, 1) + go func() { + pr, err := fx.Get(ctx, "p1") + assert.NoError(t, err) + res <- pr + }() + // the lookup has old in hand and is about to return it + <-inIsClosed + require.NoError(t, fx.Flush(ctx)) + doReleaseIsClosed() + select { + case pr := <-res: + require.Equal(t, peer.Peer(fresh), pr) + case <-time.After(2 * time.Second): + t.Fatal("Get did not return") + } + require.Equal(t, int32(2), dials.Load()) + require.False(t, old.testPeer.IsClosed()) + + // the same for Pick: nothing from the old pair + old2 := newCtlPeer("p2") + inIsClosed2 := make(chan struct{}) + releaseIsClosed2, doReleaseIsClosed2 := newRelease() + defer doReleaseIsClosed2() + old2.isClosedHook = func(call int32) { + if call == 1 { + close(inIsClosed2) + <-releaseIsClosed2 + } + } + old2.closeHook = func(int32) { <-releaseClose } + require.NoError(t, fx.AddPeer(ctx, old2)) + pick := make(chan error, 1) + go func() { + _, err := fx.Pick(ctx, "p2") + pick <- err + }() + <-inIsClosed2 + require.NoError(t, fx.Flush(ctx)) + doReleaseIsClosed2() + require.ErrorIs(t, <-pick, ocache.ErrNotExists) + }) + t.Run("get on a closed peer whose teardown blocks returns within ctx", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + // a dead peer still in the cache (its watcher has not run yet) whose + // Close never returns + dead := newCtlPeer("p1") + close(dead.closed) + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + dead.closeHook = func(int32) { <-releaseClose } + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + + // the teardown hangs: Get must not run or wait on it past its own + // deadline + done := make(chan error, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + _, gErr := fx.Get(gctx, "p1") + done <- gErr + }() + select { + case err := <-done: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(2 * time.Second): + t.Fatal("Get waited on the blocked close") + } + }) + t.Run("flush racing AddPeer never serves a stale peer", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + for i := 0; i < 100; i++ { + id := fmt.Sprintf("p%d", i) + tp := newTestPeer(id) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + assert.NoError(t, fx.AddPeer(ctx, tp)) + }() + go func() { + defer wg.Done() + _ = fx.Flush(ctx) + }() + wg.Wait() + // with no lookup at all, the peer is either current or closed + require.Eventually(t, func() bool { return tp.IsClosed() || inCurrent(p, tp) }, time.Second, time.Millisecond) + pr, err := fx.Pick(ctx, id) + if err == nil { + // served only from the current pair, which the flush left alone + require.Equal(t, peer.Peer(tp), pr) + require.True(t, inCurrent(p, tp)) + require.False(t, tp.IsClosed()) + } else { + // added to the replaced pair: closed by the flush + require.Eventually(t, tp.IsClosed, time.Second, time.Millisecond) + } + } + }) + t.Run("flush keeps a peer added after it", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + old := newTestPeer("old") + require.NoError(t, fx.AddPeer(ctx, old)) + require.NoError(t, fx.Flush(ctx)) + cur := newTestPeer("cur") + require.NoError(t, fx.AddPeer(ctx, cur)) + require.NoError(t, fx.Flush(ctx)) + require.Eventually(t, old.IsClosed, time.Second, 10*time.Millisecond) + // added after the first flush, it is stale for the second: closed; + // a peer added after the last flush is untouched + require.Eventually(t, cur.IsClosed, time.Second, 10*time.Millisecond) + last := newTestPeer("last") + require.NoError(t, fx.AddPeer(ctx, last)) + require.Never(t, last.IsClosed, 100*time.Millisecond, 10*time.Millisecond) + pr, err := fx.Pick(ctx, "last") + require.NoError(t, err) + assert.Equal(t, peer.Peer(last), pr) + }) + t.Run("incompatible version verdict survives flush", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + eo := &errObject{id: "p1", err: handshake.ErrIncompatibleVersion, createdTime: atomic2.NewTime(time.Now())} + require.NoError(t, p.current.Load().outgoing.Add("p1", eo)) + require.NoError(t, fx.Flush(ctx)) + require.NoError(t, fx.Flush(ctx)) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + t.Error("must not redial an incompatible peer") + return nil, nil + } + _, err := fx.Get(ctx, "p1") + require.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + _, err = fx.Pick(ctx, "p1") + require.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + // the same verdict object, so its backoff clock is not reset + v, err := p.current.Load().outgoing.Pick(ctx, "p1") + require.NoError(t, err) + require.Same(t, eo, v) + }) + t.Run("cached dial error is dropped by flush", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + require.NoError(t, p.current.Load().outgoing.Add("p1", &errObject{id: "p1", err: assert.AnError, createdTime: atomic2.NewTime(time.Now())})) + _, err := fx.Pick(ctx, "p1") + require.ErrorIs(t, err, assert.AnError) + + require.NoError(t, fx.Flush(ctx)) + require.Equal(t, 0, p.current.Load().outgoing.Len()) + fresh := newTestPeer("p1") + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return fresh, nil + } + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + assert.Equal(t, peer.Peer(fresh), pr) + }) + t.Run("incoming replacement via AddPeer after flush", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + + old := newTestPeer("p1") + require.NoError(t, fx.AddPeer(ctx, old)) + require.NoError(t, fx.Flush(ctx)) + _, err := fx.Pick(ctx, "p1") + require.Error(t, err) + require.Eventually(t, old.IsClosed, time.Second, 10*time.Millisecond) + + // the peer reconnects: the new incoming peer lands in the fresh pair + // at once, whatever the old pair's teardown is still doing + repl := newTestPeer("p1") + require.NoError(t, fx.AddPeer(ctx, repl)) + pr, err := fx.Pick(ctx, "p1") + require.NoError(t, err) + assert.Equal(t, peer.Peer(repl), pr) + pr, err = fx.Get(ctx, "p1") + require.NoError(t, err) + assert.Equal(t, peer.Peer(repl), pr) + require.Eventually(t, func() bool { return len(obs.getClosed()) == 1 }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 1 }, 100*time.Millisecond, 10*time.Millisecond) + }) + t.Run("replacement installed before old cleanup finishes survives", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + old := newCtlPeer("p1") + inFirstClose := make(chan struct{}) + releaseFirstClose, doReleaseFirstClose := newRelease() + // released on every exit, so a failing assertion fails instead of + // hanging fx.Finish behind the paused hook + defer doReleaseFirstClose() + old.closeHook = func(call int32) { + // the flush's background close is the first one: hold it, so the + // old pair's teardown overlaps the replacement + if call == 1 { + close(inFirstClose) + <-releaseFirstClose + } + } + require.NoError(t, fx.AddPeer(ctx, old)) + require.NoError(t, fx.Flush(ctx)) + <-inFirstClose + + repl := newTestPeer("p1") + require.NoError(t, fx.AddPeer(ctx, repl)) + doReleaseFirstClose() + + require.Eventually(t, old.IsClosed, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { + pk, err := fx.Pick(ctx, "p1") + return err != nil || pk != peer.Peer(repl) + }, 200*time.Millisecond, 10*time.Millisecond) + require.False(t, repl.IsClosed()) + require.Equal(t, 1, p.current.Load().incoming.Len()) + }) + t.Run("repeated flush closes every replaced pair once", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + tr := trackPairs(p) + + var peers []*testPeer + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + tp := newTestPeer(peerId) + peers = append(peers, tp) + return tp, nil + } + for i := 0; i < 3; i++ { + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.Equal(t, peer.Peer(peers[len(peers)-1]), pr) + require.NoError(t, fx.Flush(ctx)) + require.NoError(t, fx.Flush(ctx)) + } + require.Len(t, peers, 3) + for _, tp := range peers { + require.Eventually(t, tp.IsClosed, time.Second, 10*time.Millisecond) + } + pairs := tr.all() + require.Len(t, pairs, 7) + cur := p.current.Load() + require.Same(t, pairs[6], cur) + for _, c := range pairs[:6] { + require.Eventually(t, c.isClosed, time.Second, 10*time.Millisecond) + } + // the current pair is open and empty + require.Equal(t, 0, cur.outgoing.Len()) + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.Equal(t, peer.Peer(peers[3]), pr) + }) + t.Run("concurrent flushes never leak a pair", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + tr := trackPairs(p) + + var wg sync.WaitGroup + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + assert.NoError(t, fx.Flush(ctx)) + }() + } + wg.Wait() + pairs := tr.all() + require.Len(t, pairs, 51) + cur := p.current.Load() + require.Same(t, pairs[50], cur) + for _, c := range pairs[:50] { + require.Eventually(t, c.isClosed, time.Second, 10*time.Millisecond) + } + require.False(t, cur.isClosed()) + }) + t.Run("flush then shutdown", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + p := fx.Service.(*poolService).pool + tr := trackPairs(p) + in := newTestPeer("in") + out := newCtlPeer("out") + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return out, nil + } + require.NoError(t, fx.AddPeer(ctx, in)) + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + require.NoError(t, fx.Flush(ctx)) + fx.Finish() + require.Eventually(t, in.IsClosed, time.Second, 10*time.Millisecond) + require.Eventually(t, out.IsClosed, time.Second, 10*time.Millisecond) + // a flush on a closed pool is a no-op: no pair is built, the closed + // pair stays current and lookups fail with ErrClosed + cur := p.current.Load() + require.NoError(t, fx.Flush(ctx)) + require.Equal(t, 2, tr.created()) + require.Same(t, cur, p.current.Load()) + for _, c := range tr.all() { + require.True(t, c.isClosed()) + } + _, err = fx.Pick(ctx, "in") + require.ErrorIs(t, err, ocache.ErrClosed) + _, err = fx.Get(ctx, "out") + require.ErrorIs(t, err, ocache.ErrClosed) + require.ErrorIs(t, fx.AddPeer(ctx, newTestPeer("late")), ocache.ErrClosed) + // idempotent + require.NoError(t, fx.Service.Close(ctx)) + }) + t.Run("close waits for the pairs a flush is still closing", func(t *testing.T) { + fx := newFixture(t) + slow := newCtlPeer("p1") + slow.closeHook = func(int32) { time.Sleep(100 * time.Millisecond) } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return slow, nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.NoError(t, fx.Flush(ctx)) + fx.Finish() + // no teardown outlives the pool + require.True(t, slow.IsClosed()) + }) + t.Run("close is bounded when a flushed pair never finishes closing", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 100 * time.Millisecond + }) + hung := newCtlPeer("p1") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + hung.closeHook = func(int32) { <-releaseClose } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return hung, nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.NoError(t, fx.Flush(ctx)) + start := time.Now() + fx.Finish() + require.Less(t, time.Since(start), 2*time.Second) + require.False(t, hung.IsClosed()) + }) + t.Run("hung peer close does not hold back the others", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + peers := map[string]*ctlPeer{} + for _, id := range []string{"p1", "p2", "p3", "p4"} { + peers[id] = newCtlPeer(id) + } + peers["p1"].closeHook = func(int32) { <-releaseClose } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return peers[peerId], nil + } + for id := range peers { + _, err := fx.Get(ctx, id) + require.NoError(t, err) + } + require.NoError(t, fx.Flush(ctx)) + // well within closeTimeout, so not thanks to the pass giving up on p1 + require.Eventually(t, func() bool { + return peers["p2"].IsClosed() && peers["p3"].IsClosed() && peers["p4"].IsClosed() + }, 100*time.Millisecond, time.Millisecond) + require.False(t, peers["p1"].IsClosed()) + doReleaseClose() + require.Eventually(t, peers["p1"].IsClosed, time.Second, 10*time.Millisecond) + // the parallel close and the cache's own pass both reach each peer: + // the second call is the idempotent no-op the pool relies on + for _, pr := range peers { + require.LessOrEqual(t, pr.closeCalls.Load(), int32(2), pr.Id()) + } + }) + t.Run("hung incoming close does not delay cancelling a pre-flush dial", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + in := newCtlPeer("in") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + in.closeHook = func(int32) { <-releaseClose } + require.NoError(t, fx.AddPeer(ctx, in)) + + var dials atomic.Int32 + var cancelled atomic.Bool + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + // a dead path: returns only when cancelled + <-ctx.Done() + cancelled.Store(true) + return nil, ctx.Err() + } + return newTestPeer(peerId), nil + } + type result struct { + pr peer.Peer + err error + } + first := make(chan result, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 3*time.Second) + defer cancel() + pr, err := fx.Get(gctx, "dead") + first <- result{pr, err} + }() + require.Eventually(t, func() bool { return dials.Load() == 1 }, time.Second, time.Millisecond) + start := time.Now() + require.NoError(t, fx.Flush(ctx)) + res := <-first + require.NoError(t, res.err) + require.False(t, res.pr.IsClosed()) + require.Less(t, time.Since(start), 500*time.Millisecond, "pre-flush Get waited out the dead dial") + require.True(t, cancelled.Load()) + require.Equal(t, int32(2), dials.Load()) + }) + t.Run("hung incoming close does not keep a late dial alive", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + in := newCtlPeer("in") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + in.closeHook = func(int32) { <-releaseClose } + require.NoError(t, fx.AddPeer(ctx, in)) + + late := newTestPeer("out") + var dials atomic.Int32 + var cancelled atomic.Bool + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + // a slow handshake that completes despite the cancellation + select { + case <-ctx.Done(): + cancelled.Store(true) + case <-time.After(400 * time.Millisecond): + } + return late, nil + } + return newTestPeer(peerId), nil + } + first := make(chan peer.Peer, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 3*time.Second) + defer cancel() + pr, err := fx.Get(gctx, "out") + assert.NoError(t, err) + first <- pr + }() + require.Eventually(t, func() bool { return dials.Load() == 1 }, time.Second, time.Millisecond) + start := time.Now() + require.NoError(t, fx.Flush(ctx)) + pr := <-first + require.NotSame(t, late, pr) + require.Less(t, time.Since(start), 300*time.Millisecond, "pre-flush Get waited out the stale dial") + require.True(t, cancelled.Load()) + // published into the old outgoing cache, which closed at once + require.Eventually(t, late.IsClosed, time.Second, 10*time.Millisecond) + require.Equal(t, int32(2), dials.Load()) + }) + t.Run("lookup blocked behind a GC TryClose on the replaced pair retries at once", func(t *testing.T) { + // nothing in the old pair's Close can cut this wait short (the GC + // owns the entry): only the pair ctx that the swap cancels does + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + old := newCtlPeer("p1") + inTryClose := make(chan struct{}) + releaseTryClose, doReleaseTryClose := newRelease() + defer doReleaseTryClose() + old.tryClose = func() (bool, error) { + close(inTryClose) + <-releaseTryClose + return false, nil + } + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + return old, nil + } + return newTestPeer(peerId), nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + oldPair := p.current.Load() + go func() { _, _ = oldPair.outgoing.TryRemove("p1") }() + <-inTryClose + + got := make(chan peer.Peer, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 3*time.Second) + defer cancel() + pr, err := fx.Get(gctx, "p1") + assert.NoError(t, err) + got <- pr + }() + // the Get is parked on the closing entry + require.Never(t, func() bool { return len(got) > 0 }, 50*time.Millisecond, 5*time.Millisecond) + start := time.Now() + require.NoError(t, fx.Flush(ctx)) + select { + case pr := <-got: + require.NotSame(t, old, pr) + require.Less(t, time.Since(start), 500*time.Millisecond) + case <-time.After(2 * time.Second): + t.Fatal("Get stayed parked on the replaced pair") + } + }) + t.Run("GetOneOf never returns a peer from the replaced pair", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + old := newCtlPeer("g") + fresh := newTestPeer("g") + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if dials.Add(1) == 1 { + return old, nil + } + return fresh, nil + } + _, err := fx.Get(ctx, "g") + require.NoError(t, err) + // keep old live so the retry, not its close, decides the result + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + old.closeHook = func(int32) { <-releaseClose } + paused := make(chan struct{}) + releaseIsClosed, doReleaseIsClosed := newRelease() + defer doReleaseIsClosed() + var armed atomic.Bool + armed.Store(true) + old.isClosedHook = func(int32) { + if armed.CompareAndSwap(true, false) { + close(paused) + <-releaseIsClosed + } + } + got := make(chan peer.Peer, 1) + go func() { + pr, err := fx.GetOneOf(ctx, []string{"g"}) + assert.NoError(t, err) + got <- pr + }() + // getIfActive has old in hand + <-paused + require.NoError(t, fx.Flush(ctx)) + doReleaseIsClosed() + require.Same(t, fresh, <-got) + }) + t.Run("add in flight across flush is never rejected", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + tp := newCtlPeer("p1") + inAdd := make(chan struct{}) + releaseAdd, doReleaseAdd := newRelease() + defer doReleaseAdd() + tp.idHook = func(call int32) { + if call == 1 { + close(inAdd) + <-releaseAdd + } + } + added := make(chan error, 1) + go func() { added <- fx.AddPeer(ctx, tp) }() + <-inAdd + flushed := make(chan struct{}) + go func() { + _ = fx.Flush(ctx) + close(flushed) + }() + // the swap waits for the add that already picked its pair, so the + // peer can neither land in a closed cache nor be dropped + select { + case <-flushed: + t.Fatal("Flush did not wait for the add in flight") + case <-time.After(50 * time.Millisecond): + } + doReleaseAdd() + require.NoError(t, <-added) + <-flushed + // accepted before the flush completed: closed by it + require.Eventually(t, tp.IsClosed, time.Second, 10*time.Millisecond) + }) + t.Run("get parked in a dial gets ErrClosed when the pool closes", func(t *testing.T) { + fx := newFixture(t) + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + <-ctx.Done() + return nil, ctx.Err() + } + got := make(chan error, 1) + go func() { + _, err := fx.Get(ctx, "x") + got <- err + }() + require.Eventually(t, func() bool { return dials.Load() == 1 }, time.Second, time.Millisecond) + fx.Finish() + // the pair was cancelled but not replaced: the pool is closing, there + // is nothing to retry on + require.ErrorIs(t, <-got, ocache.ErrClosed) + }) + t.Run("a flush landing between the lookup and its post-check is a retry, not a closed pool", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // pairs built from here on carry the seam + var armed atomic.Bool + inErr := make(chan struct{}) + releaseErr, doReleaseErr := newRelease() + defer doReleaseErr() + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + c.ctx = &hookCtx{Context: c.ctx, onErr: func() { + if armed.CompareAndSwap(true, false) { + close(inErr) + <-releaseErr + } + }} + return c + } + require.NoError(t, fx.Flush(ctx)) + // a cached verdict takes the slow path (the fast path only serves peers) + eo := &errObject{id: "inc", err: handshake.ErrIncompatibleVersion, createdTime: atomic2.NewTime(time.Now())} + require.NoError(t, p.current.Load().outgoing.Add("inc", eo)) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + t.Error("must not redial an incompatible peer") + return nil, nil + } + armed.Store(true) + got := make(chan error, 1) + go func() { + _, err := fx.Get(ctx, "inc") + got <- err + }() + // f has returned; lookup is reading the pair state + <-inErr + require.NoError(t, fx.Flush(ctx)) + doReleaseErr() + // the pair is cancelled and replaced: a retry on the fresh pair, which + // carries the verdict, never ErrClosed + require.ErrorIs(t, <-got, handshake.ErrIncompatibleVersion) + }) + t.Run("add whose duplicate is flushed between Add and Pick retries on the current pair", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + require.NoError(t, fx.AddPeer(ctx, newTestPeer("d"))) + repl := newCtlPeer("d") + old := p.current.Load() + repl.idHook = func(call int32) { + if call == 2 { + // the Pick after ErrExists: the pair holding the duplicate + // is replaced and closed under AddPeer's feet + assert.NoError(t, fx.Flush(ctx)) + assert.Eventually(t, old.isClosed, time.Second, time.Millisecond) + } + } + require.NoError(t, fx.AddPeer(ctx, repl)) + require.True(t, inCurrent(p, repl)) + require.False(t, repl.IsClosed()) + }) + t.Run("add racing a Get on the same id succeeds", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // a pair whose incoming loader can be held open, standing in for the + // instant ErrNotExists one: the Get leaves a loading entry that Add + // trips over and Pick waits out + gate := make(chan struct{}) + entered := make(chan struct{}) + var once sync.Once + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + _ = c.incoming.Close() + c.incoming = ocache.New(func(ctx context.Context, id string) (ocache.Object, error) { + once.Do(func() { close(entered) }) + <-gate + return nil, ocache.ErrNotExists + }, ocache.WithGCPeriod(0)) + c.peekIncoming = mustPeeker(c.incoming) + return c + } + require.NoError(t, fx.Flush(ctx)) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + go func() { _, _ = fx.Get(ctx, "d") }() + <-entered + tp := newTestPeer("d") + added := make(chan error, 1) + go func() { added <- fx.AddPeer(ctx, tp) }() + require.Never(t, func() bool { return len(added) > 0 }, 50*time.Millisecond, 5*time.Millisecond) + close(gate) + require.NoError(t, <-added) + require.True(t, inCurrent(p, tp)) + }) + t.Run("close waits for the parallel peer closes of every pair", func(t *testing.T) { + for _, flushed := range []bool{true, false} { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 5 * time.Second + }) + pr := newCtlPeer("p1") + var finished atomic.Bool + pr.closeHook = func(int32) { + // only the parallel close is slow; the cache's own pass + // closes the peer a second time and returns at once + if fromClosePrepass() { + time.Sleep(200 * time.Millisecond) + finished.Store(true) + } + } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return pr, nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + if flushed { + require.NoError(t, fx.Flush(ctx)) + } + fx.Finish() + require.True(t, finished.Load(), "flushed=%v: Close returned before the parallel close finished", flushed) + } + }) + t.Run("close is bounded by its ctx with a hung peer", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 5 * time.Second + }) + hung := newCtlPeer("p1") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + hung.closeHook = func(int32) { <-releaseClose } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return hung, nil + } + _, err := fx.Get(ctx, "p1") + require.NoError(t, err) + cctx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + start := time.Now() + require.NoError(t, fx.Service.Close(cctx)) + require.Less(t, time.Since(start), time.Second) + require.False(t, hung.IsClosed()) + }) + t.Run("hung incoming close does not hold back the other incoming peers", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + hung := newCtlPeer("hung") + hung.closeHook = func(int32) { <-releaseClose } + require.NoError(t, fx.AddPeer(ctx, hung)) + var others []*testPeer + for i := 0; i < 10; i++ { + tp := newTestPeer(fmt.Sprintf("in%d", i)) + others = append(others, tp) + require.NoError(t, fx.AddPeer(ctx, tp)) + } + require.NoError(t, fx.Flush(ctx)) + // well within closeTimeout, so not thanks to the pass giving up + require.Eventually(t, func() bool { + for _, tp := range others { + if !tp.IsClosed() { + return false + } + } + return true + }, 100*time.Millisecond, time.Millisecond) + require.False(t, hung.IsClosed()) + }) + t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + in := newTestPeer("in") + out := newTestPeer("out") + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return out, nil + } + require.NoError(t, fx.AddPeer(ctx, in)) + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + // both a flush and a later lookup try to discard them: each peer + // still reports a single Closed event + require.NoError(t, fx.Flush(ctx)) + _, _ = fx.Pick(ctx, "in") + _, _ = fx.Pick(ctx, "out") + require.NoError(t, fx.Flush(ctx)) + + require.Eventually(t, func() bool { return len(obs.getClosed()) == 2 }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 2 }, 100*time.Millisecond, 10*time.Millisecond) + byInbound := map[bool]peerobserver.Event{} + for _, ev := range obs.getClosed() { + byInbound[ev.Inbound] = ev + } + assert.Equal(t, "in", byInbound[true].PeerId) + assert.Equal(t, "out", byInbound[false].PeerId) + }) +} + +// taggedPeer carries the number of the pair that was current when it was dialed +type taggedPeer struct { + *testPeer + tag int64 +} + +func TestPool_FlushStorm(t *testing.T) { + t.Run("back-to-back flushes never fail a lookup or serve a stale peer", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + tr := trackPairs(p) + + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + // shorter than the flush period, or no dial would ever complete + // before the next swap (a livelock both designs share) + time.Sleep(200 * time.Microsecond) + return &taggedPeer{testPeer: newTestPeer(peerId), tag: tr.seqOf(p.current.Load())}, nil + } + // the workers run until the flusher has done wantFlushes (about + // 300ms on an idle machine; with GOMAXPROCS=1 the busy workers starve + // it, so a fixed duration would not do) + const wantFlushes = 60 + flushes, stopFlusher := startFlusher(t, fx, time.Millisecond) + defer stopFlusher() + cap := time.Now().Add(20 * time.Second) + var stale, failed, ok atomic.Int32 + var wg sync.WaitGroup + for w := 0; w < 8; w++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; flushes.Load() < wantFlushes && time.Now().Before(cap); i++ { + before := tr.seqOf(p.current.Load()) + gctx, cancel := context.WithTimeout(ctx, 5*time.Second) + pr, err := fx.Get(gctx, "a") + switch { + case err != nil: + failed.Add(1) + t.Errorf("Get failed while the pool is open: %v", err) + case pr.(*taggedPeer).tag < before: + // a peer from a pair replaced before the Get started + stale.Add(1) + default: + ok.Add(1) + } + if i%50 == 0 { + // Pick and GetOneOf take the same path + if _, err = fx.Pick(gctx, "a"); err != nil { + assert.NotErrorIs(t, err, ocache.ErrClosed) + } + if _, err = fx.GetOneOf(gctx, []string{"a"}); err != nil { + t.Errorf("GetOneOf failed while the pool is open: %v", err) + } + } + cancel() + } + }() + } + wg.Wait() + stopFlusher() + t.Logf("flushes=%d ok=%d failed=%d stale=%d", flushes.Load(), ok.Load(), failed.Load(), stale.Load()) + require.GreaterOrEqual(t, flushes.Load(), int64(wantFlushes)) + require.Greater(t, ok.Load(), int32(100)) + require.Zero(t, failed.Load()) + require.Zero(t, stale.Load()) + }) + t.Run("peers added during swaps land in the current pair and stay open", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + + _, stopFlusher := startFlusher(t, fx, time.Millisecond) + defer stopFlusher() + var peers []*testPeer + for i := 0; i < 500; i++ { + tp := newTestPeer(fmt.Sprintf("p%d", i)) + peers = append(peers, tp) + // never rejected with ErrClosed while the pool is open + require.NoError(t, fx.AddPeer(ctx, tp)) + if i%100 == 0 { + time.Sleep(time.Millisecond) + } + } + stopFlusher() + // every peer is now either in the current pair, open, or was + // flushed and closed; none is both current and killed + for _, tp := range peers { + require.Eventually(t, func() bool { return tp.IsClosed() || inCurrent(p, tp) }, time.Second, time.Millisecond, tp.Id()) + if inCurrent(p, tp) { + require.False(t, tp.IsClosed(), tp.Id()) + } + } + // with the swaps over, a new peer is current and stays open + last := newTestPeer("last") + require.NoError(t, fx.AddPeer(ctx, last)) + require.True(t, inCurrent(p, last)) + require.Never(t, last.IsClosed, 50*time.Millisecond, 10*time.Millisecond) + }) + t.Run("a dial that never completes under a flush storm ends with the ctx error", func(t *testing.T) { + // every swap cancels the lookup's dial and the retry starts another: + // the livelock is bounded by the caller's ctx and is not reported + // as a closed pool + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = time.Second + }) + defer fx.Finish() + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + <-ctx.Done() + return nil, ctx.Err() + } + _, stopFlusher := startFlusher(t, fx, 2*time.Millisecond) + defer stopFlusher() + got := make(chan error, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 200*time.Millisecond) + defer cancel() + _, err := fx.Get(gctx, "s") + got <- err + }() + select { + case err := <-got: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(5 * time.Second): + t.Fatal("Get did not return with its ctx") + } + stopFlusher() + require.Greater(t, dials.Load(), int32(1), "the dial was never retried") + }) + t.Run("a kept verdict is never redialed while flushes run", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + eo := &errObject{id: "inc", err: handshake.ErrIncompatibleVersion, createdTime: atomic2.NewTime(time.Now())} + require.NoError(t, p.current.Load().outgoing.Add("inc", eo)) + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + // the verdict must be in the fresh pair before that pair is published + stop := make(chan struct{}) + var wg sync.WaitGroup + for w := 0; w < 8; w++ { + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-stop: + return + default: + } + _, err := fx.Get(ctx, "inc") + assert.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + } + }() + } + for i := 0; i < 300; i++ { + require.NoError(t, fx.Flush(ctx)) + } + close(stop) + wg.Wait() + require.Zero(t, dials.Load()) + }) +} + +// testMetric exposes a registry the way the metric component does; the other +// methods are never called by the pool +type testMetric struct { + metric.Metric + reg *prometheus.Registry +} + +func (m *testMetric) Init(a *app.App) error { return nil } +func (m *testMetric) Name() string { return metric.CName } +func (m *testMetric) Run(ctx context.Context) error { return nil } +func (m *testMetric) Close(ctx context.Context) error { return nil } +func (m *testMetric) Registry() *prometheus.Registry { return m.reg } + +func TestPool_FlushMetrics(t *testing.T) { + reg := prometheus.NewRegistry() + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + a.Register(&testMetric{reg: reg}) + }) + defer fx.Finish() + + gather := func() map[string]float64 { + families, err := reg.Gather() + require.NoError(t, err) + out := map[string]float64{} + for _, mf := range families { + require.Len(t, mf.GetMetric(), 1) + m := mf.GetMetric()[0] + if m.GetGauge() != nil { + out[mf.GetName()] = m.GetGauge().GetValue() + } else { + out[mf.GetName()] = m.GetCounter().GetValue() + } + } + return out + } + expectedNames := []string{ + "netpool_outgoing_hit", "netpool_outgoing_miss", "netpool_outgoing_gc", "netpool_outgoing_size", + "netpool_incoming_hit", "netpool_incoming_miss", "netpool_incoming_gc", "netpool_incoming_size", + } + names := func(values map[string]float64) (out []string) { + for name := range values { + out = append(out, name) + } + return + } + require.ElementsMatch(t, expectedNames, names(gather())) + + // recreating the caches re-registers nothing + for i := 0; i < 100; i++ { + require.NoError(t, fx.Flush(ctx)) + } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + _, err = fx.Get(ctx, "out") + require.NoError(t, err) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("in1"))) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("in2"))) + values := gather() + require.ElementsMatch(t, expectedNames, names(values)) + // size reads the current pair; the counters are continuous across pairs + assert.Equal(t, float64(1), values["netpool_outgoing_size"]) + assert.Equal(t, float64(2), values["netpool_incoming_size"]) + assert.Equal(t, float64(1), values["netpool_outgoing_miss"]) + // Get tries incoming first (a miss), then hits outgoing on the second call + assert.Equal(t, float64(1), values["netpool_outgoing_hit"]) + assert.Equal(t, float64(2), values["netpool_incoming_miss"]) + + require.NoError(t, fx.Flush(ctx)) + values = gather() + assert.Equal(t, float64(0), values["netpool_outgoing_size"]) + assert.Equal(t, float64(0), values["netpool_incoming_size"]) + assert.Equal(t, float64(1), values["netpool_outgoing_miss"]) +} + +func TestPool_FlushGoroutines(t *testing.T) { + // every pair owns two GC tickers and every flush a closer goroutine: all + // of them must be gone once the pool is closed + runtime.GC() + before := runtime.NumGoroutine() + fx := newFixture(t) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + for i := 0; i < 30; i++ { + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("in"))) + require.NoError(t, fx.Flush(ctx)) + } + fx.Finish() + // polled from this goroutine so the baseline is comparable (Eventually + // would add its own) + deadline := time.Now().Add(5 * time.Second) + for runtime.NumGoroutine() > before { + if time.Now().After(deadline) { + t.Fatalf("goroutines before=%d after=%d", before, runtime.NumGoroutine()) + } + time.Sleep(10 * time.Millisecond) + } +} diff --git a/net/pool/pool_test.go b/net/pool/pool_test.go index 526ca1a00..e56b175a9 100644 --- a/net/pool/pool_test.go +++ b/net/pool/pool_test.go @@ -132,29 +132,6 @@ func TestPool_Flush(t *testing.T) { _, err = fx.Pick(ctx, "peer2") assert.Error(t, err) }) - t.Run("concurrent flush operations", func(t *testing.T) { - fx := newFixture(t) - defer fx.Finish() - p1 := newTestPeer("peer1") - p2 := newTestPeer("peer2") - require.NoError(t, fx.AddPeer(ctx, p1)) - require.NoError(t, fx.AddPeer(ctx, p2)) - done := make(chan error, 2) - go func() { - done <- fx.Flush(ctx) - }() - go func() { - done <- fx.Flush(ctx) - }() - err1 := <-done - err2 := <-done - assert.NoError(t, err1) - assert.NoError(t, err2) - _, err := fx.Pick(ctx, "peer1") - assert.Error(t, err) - _, err = fx.Pick(ctx, "peer2") - assert.Error(t, err) - }) t.Run("flush then get should work correctly", func(t *testing.T) { fx := newFixture(t) defer fx.Finish() @@ -202,37 +179,44 @@ func TestPool_Flush(t *testing.T) { assert.Len(t, poolStat.PeerStats, 2) err = fx.Flush(ctx) require.NoError(t, err) - stat = statProvider.ProvideStat() - poolStat, ok = stat.(*poolStats) - require.True(t, ok) - assert.Len(t, poolStat.PeerStats, 0) + // stats read the fresh pair; the old peers close in the background + assert.Len(t, statProvider.ProvideStat().(*poolStats).PeerStats, 0) }) - t.Run("flush does not remove peers loading during flush", func(t *testing.T) { + t.Run("peer dialed across flush is rejected and redialed", func(t *testing.T) { fx := newFixture(t) defer fx.Finish() dialStarted := make(chan struct{}) blockDial := make(chan struct{}) - loadingPeer := newTestPeer("loading-peer") + latePeer := newTestPeer("loading-peer") + var dials atomic2.Int32 fx.Dialer.dial = func(ctx context.Context, peerId string) (peer peer.Peer, err error) { - close(dialStarted) - <-blockDial - return loadingPeer, nil + if dials.Add(1) == 1 { + close(dialStarted) + <-blockDial + return latePeer, nil + } + // a new peer per dial, like a real dialer (see the flush tests) + return newTestPeer(peerId), nil } resultChan := make(chan peer.Peer, 1) go func() { p, err := fx.Get(ctx, "loading-peer") - require.NoError(t, err) + assert.NoError(t, err) resultChan <- p }() <-dialStarted - err := fx.Flush(ctx) - require.NoError(t, err) + require.NoError(t, fx.Flush(ctx)) close(blockDial) + // the dial started before the flush: its peer is published late but + // rejected by the very lookup waiting on it, which dials again p := <-resultChan - require.Equal(t, loadingPeer, p) + require.NotNil(t, p) + require.NotSame(t, latePeer, p) + require.False(t, p.IsClosed()) + require.Eventually(t, latePeer.IsClosed, time.Second, 10*time.Millisecond) pickedPeer, err := fx.Pick(ctx, "loading-peer") require.NoError(t, err) - assert.Equal(t, loadingPeer, pickedPeer) + assert.Equal(t, p, pickedPeer) }) } @@ -371,7 +355,7 @@ func TestPool_GetOneOf(t *testing.T) { assert.Equal(t, 1, calls) assert.Nil(t, p) - val, err := fx.Service.(*poolService).outgoing.Pick(ctx, "1") + val, err := fx.Service.(*poolService).current.Load().outgoing.Pick(ctx, "1") require.NoError(t, err) errObj := val.(*errObject) closed, err := errObj.TryClose(time.Minute) @@ -565,6 +549,12 @@ func newFixture(t *testing.T) *fixture { } func newFixtureWithObserver(t *testing.T, obs peerobserver.Observer) *fixture { + return newFixtureCfg(t, obs, nil) +} + +// newFixtureCfg lets a test tune the service (closeTimeout) before Init and +// register extra components (a metric registry) +func newFixtureCfg(t *testing.T, obs peerobserver.Observer, cfg func(ps *poolService, a *app.App)) *fixture { fx := &fixture{ Service: New(), Dialer: &dialerMock{}, @@ -575,6 +565,9 @@ func newFixtureWithObserver(t *testing.T, obs peerobserver.Observer) *fixture { if obs != nil { a.Register(peerobserver.New(obs)) } + if cfg != nil { + cfg(fx.Service.(*poolService), a) + } require.NoError(t, a.Start(context.Background())) fx.a = a fx.t = t @@ -639,6 +632,7 @@ var _ peer.Peer = (*testPeer)(nil) type testPeer struct { id string + closeMu sync.Mutex closed chan struct{} created time.Time subConnections int @@ -697,9 +691,12 @@ func (t *testPeer) TryClose(objectTTL time.Duration) (res bool, err error) { } func (t *testPeer) Close() error { + // the pool may close a rejected peer from several paths at once; + // idempotent and silent like the real peer (its MultiConn.Close is) + t.closeMu.Lock() + defer t.closeMu.Unlock() select { case <-t.closed: - return fmt.Errorf("already closed") default: close(t.closed) } @@ -736,7 +733,7 @@ func TestPool_EvictsOutgoingOnClose(t *testing.T) { require.NoError(t, tp.Close()) require.Eventually(t, func() bool { - return fx.Service.(*poolService).outgoing.Len() == 0 + return fx.Service.(*poolService).current.Load().outgoing.Len() == 0 }, time.Second, 10*time.Millisecond) } @@ -750,7 +747,7 @@ func TestPool_EvictsIncomingOnClose(t *testing.T) { require.NoError(t, tp.Close()) require.Eventually(t, func() bool { - return fx.Service.(*poolService).incoming.Len() == 0 + return fx.Service.(*poolService).current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) } @@ -761,11 +758,11 @@ func TestPool_EvictOnClose_ExitsOnShutdownWithoutEviction(t *testing.T) { // the peer is actually in the cache, so a wrong eviction would be observable tp := newTestPeer("inc1") // peer stays alive (CloseChan never fires) - require.NoError(t, p.incoming.Add(tp.Id(), tp)) + require.NoError(t, p.current.Load().incoming.Add(tp.Id(), tp)) done := make(chan struct{}) go func() { - p.evictOnClose(tp, p.incoming, true) + p.evictOnClose(tp, p.current.Load().incoming, true) close(done) }() @@ -784,7 +781,7 @@ func TestPool_EvictOnClose_ExitsOnShutdownWithoutEviction(t *testing.T) { t.Fatal("watcher did not exit on shutdown") } require.False(t, tp.IsClosed(), "shutdown path must not close the peer") - pk, err := p.incoming.Pick(ctx, tp.Id()) + pk, err := p.current.Load().incoming.Pick(ctx, tp.Id()) require.NoError(t, err, "shutdown path must leave the peer for cache.Close to evict") require.Equal(t, tp, pk) } @@ -884,12 +881,12 @@ func TestPool_PeerObserver(t *testing.T) { p := fx.Service.(*poolService).pool tp := newTestPeer("inc1") - require.NoError(t, p.incoming.Add(tp.Id(), tp)) + require.NoError(t, p.current.Load().incoming.Add(tp.Id(), tp)) require.NoError(t, tp.Close()) // a RemoveSame that begins pool shutdown while the watcher is inside // it: the watcher must re-check and swallow the event - cache := &shutdownOnRemoveSame{OCache: p.incoming, cancel: p.closingCancel} + cache := &shutdownOnRemoveSame{OCache: p.current.Load().incoming, cancel: p.closingCancel} done := make(chan struct{}) go func() { p.evictOnClose(tp, cache, true) diff --git a/net/pool/poolservice.go b/net/pool/poolservice.go index ea08660a6..6951318b9 100644 --- a/net/pool/poolservice.go +++ b/net/pool/poolservice.go @@ -3,11 +3,11 @@ package pool import ( "context" "errors" + "fmt" "time" "github.com/prometheus/client_golang/prometheus" "go.uber.org/atomic" - "go.uber.org/zap" "github.com/anyproto/any-sync/app" "github.com/anyproto/any-sync/app/debugstat" @@ -23,10 +23,14 @@ const ( CName = "common.net.pool" ) +// closeTimeout bounds a cache close pass and Close's wait on the pairs an +// earlier Flush is still closing (see pool.closeTimeout) +const closeTimeout = 10 * time.Second + var log = logger.NewNamed(CName) func New() Service { - return &poolService{} + return &poolService{pool: &pool{closeTimeout: closeTimeout}} } type Service interface { @@ -47,65 +51,144 @@ type poolService struct { func (p *poolService) Init(a *app.App) (err error) { p.dialer = a.MustComponent("net.peerservice").(dialer) - p.pool = &pool{} + if p.pool.closeTimeout <= 0 { + p.pool.closeTimeout = closeTimeout + } p.pool.closingCtx, p.pool.closingCancel = context.WithCancel(context.Background()) if m := a.Component(metric.CName); m != nil { p.metricReg = m.(metric.Metric).Registry() } - p.pool.outgoing = ocache.New( + // Flush recreates the caches, so the collectors are built and registered + // once here and shared by every instance (WithPrometheus would register + // the same names again and panic). The names stay + // netpool_{outgoing,incoming}_{hit,miss,gc,size}; size reads the current + // cache, so it is registered only once a pair is published. ocache skips a + // nil option. + var outgoing, incoming ocache.PrometheusCollectors + var outgoingMetrics, incomingMetrics ocache.Option + if p.metricReg != nil { + outgoing = ocache.NewPrometheusCollectors("netpool", "outgoing", func() int { + return p.pool.current.Load().outgoing.Len() + }) + incoming = ocache.NewPrometheusCollectors("netpool", "incoming", func() int { + return p.pool.current.Load().incoming.Len() + }) + outgoingMetrics, incomingMetrics = outgoing.Option(), incoming.Option() + p.pool.incomingMiss = incoming.Miss + } + p.pool.newCaches = func() *caches { + return p.newCaches(outgoingMetrics, incomingMetrics) + } + p.pool.current.Store(p.pool.newCaches()) + if p.metricReg != nil { + outgoing.MustRegister(p.metricReg) + incoming.MustRegister(p.metricReg) + } + comp, ok := a.Component(debugstat.CName).(debugstat.StatService) + if !ok { + comp = debugstat.NewNoOp() + } + p.statService = comp + p.statService.AddProvider(p) + p.pool.observer = peerobserver.FromApp(a) + return nil +} + +// newCaches builds one cache pair. The outgoing loader binds its watcher to +// the cache it loads into, not to whichever pair is current when the dial +// finishes: after a Flush that is a different one. +func (p *poolService) newCaches(outgoingMetrics, incomingMetrics ocache.Option) *caches { + c := &caches{} + c.ctx, c.cancel = context.WithCancel(context.Background()) + c.outgoing = ocache.New( func(ctx context.Context, id string) (value ocache.Object, err error) { value, err = p.dialer.Dial(ctx, id) if err != nil { if errors.Is(err, handshake.ErrIncompatibleVersion) { - return &errObject{err: err, createdTime: atomic.NewTime(time.Now())}, nil + return &errObject{id: id, err: err, createdTime: atomic.NewTime(time.Now())}, nil } return value, err } if pr, ok := value.(peer.Peer); ok { - go p.pool.evictOnClose(pr, p.pool.outgoing, false) + go p.pool.evictOnClose(pr, c.outgoing, false) } return value, nil }, ocache.WithLogger(log.Sugar()), ocache.WithGCPeriod(time.Minute/2), ocache.WithTTL(time.Minute), - ocache.WithPrometheus(p.metricReg, "netpool", "outgoing"), + ocache.WithCloseTimeout(p.pool.closeTimeout), + outgoingMetrics, ) - p.pool.incoming = ocache.New( + c.incoming = ocache.New( func(ctx context.Context, id string) (value ocache.Object, err error) { return nil, ocache.ErrNotExists }, ocache.WithLogger(log.Sugar()), ocache.WithGCPeriod(time.Minute/2), ocache.WithTTL(time.Minute), - ocache.WithPrometheus(p.metricReg, "netpool", "incoming"), + ocache.WithCloseTimeout(p.pool.closeTimeout), + incomingMetrics, ) - comp, ok := a.Component(debugstat.CName).(debugstat.StatService) + c.peekIncoming, c.peekOutgoing = mustPeeker(c.incoming), mustPeeker(c.outgoing) + return c +} + +// mustPeeker asserts the hit-path read on a cache; every cache the pool builds +// comes from ocache.New, so a failure is a programming error +func mustPeeker(c ocache.OCache) ocache.Peeker { + pk, ok := c.(ocache.Peeker) if !ok { - comp = debugstat.NewNoOp() + panic(fmt.Sprintf("pool: cache %T does not implement ocache.Peeker", c)) } - p.statService = comp - p.statService.AddProvider(p) - p.pool.observer = peerobserver.FromApp(a) - return nil + return pk } func (p *pool) Run(ctx context.Context) (err error) { return nil } +// Close closes the current pair, waits for its peer teardowns and for the +// pairs earlier flushes are still closing, all bounded by ctx and closeTimeout; +// on a timeout the teardown goroutine keeps running in the background (it +// ends when the hung peer close does, which may be never) and Close returns +// nil. Flush is a no-op from here on. Idempotent. func (p *pool) Close(ctx context.Context) (err error) { + p.swapMu.Lock() + if p.closed { + p.swapMu.Unlock() + return nil + } + p.closed = true + p.swapMu.Unlock() if p.closingCancel != nil { p.closingCancel() } p.statService.RemoveProvider(p) - if e := p.incoming.Close(); e != nil { - log.Warn("close incoming cache error", zap.Error(e)) + cur := p.current.Load() + // lookups blocked on the current pair fail now with ErrClosed (see lookup) + cur.cancel() + done := make(chan error, 1) + go func() { + peers, err := closeCaches(cur) + peers.Wait() + p.closing.Wait() + done <- err + }() + timer := time.NewTimer(p.closeTimeout) + defer timer.Stop() + select { + case err = <-done: + case <-ctx.Done(): + log.Warn("pool close: ctx done before every peer closed") + case <-timer.C: + log.Warn("pool close: timed out waiting for peers to close") } - return p.outgoing.Close() + return err } type errObject struct { + id string err error createdTime *atomic.Time } @@ -114,6 +197,15 @@ func (e *errObject) Error() error { return e.err } +// keepOnFlush reports whether Flush carries this cached error over into the +// fresh cache. An incompatible-version verdict survives: it says nothing +// about the network, and dropping it would defeat its 20-minute backoff on +// every recovery. Only published verdicts are carried; one whose dial is +// still in flight at the swap stays with the old pair (one extra dial). +func (e *errObject) keepOnFlush() bool { + return errors.Is(e.err, handshake.ErrIncompatibleVersion) +} + func (e *errObject) Close() (err error) { return } diff --git a/net/secureservice/handshake/handshake.go b/net/secureservice/handshake/handshake.go index 845c8de76..11f92a89d 100644 --- a/net/secureservice/handshake/handshake.go +++ b/net/secureservice/handshake/handshake.go @@ -4,7 +4,9 @@ import ( "encoding/binary" "errors" "io" + "net" "sync" + "sync/atomic" "golang.org/x/exp/slices" @@ -39,6 +41,13 @@ func (he HandshakeError) Error() string { return he.e.String() } +// Unwrap exposes the underlying error (e.g. the TLS transport's EOF, reset or +// timeout), so callers can match it with errors.Is/As. Protocol-level +// handshake errors carry none and unwrap to nil. +func (he HandshakeError) Unwrap() error { + return he.Err +} + var ( ErrUnexpectedPayload = HandshakeError{e: handshakeproto.Error_UnexpectedPayload} ErrDeadlineExceeded = HandshakeError{e: handshakeproto.Error_DeadlineExceeded} @@ -80,7 +89,12 @@ func newHandshake() *handshake { } type handshake struct { - conn io.ReadWriteCloser + conn io.ReadWriteCloser + // closeConn, when set, replaces conn.Close for closing a net.Conn + closeConn func(net.Conn) + // abandoned, when set and true, means the caller gave up and closes the + // conn itself: an error is then neither acknowledged nor closed here + abandoned *atomic.Bool remoteCred *handshakeproto.Credentials remoteProto *handshakeproto.Proto remoteAck *handshakeproto.Ack @@ -107,9 +121,14 @@ func (h *handshake) writeProto(proto *handshakeproto.Proto) (err error) { } func (h *handshake) tryWriteErrAndClose(err error) { + if h.abandoned != nil && h.abandoned.Load() { + // the caller has given up and closed the conn: writing an ack to a + // silent peer could only block + return + } if err == ErrUnexpectedPayload { // if we got unexpected message - just close the connection - _ = h.conn.Close() + h.close() return } var ackErr handshakeproto.Error @@ -119,6 +138,14 @@ func (h *handshake) tryWriteErrAndClose(err error) { ackErr = handshakeproto.Error_Unexpected } _ = h.writeAck(ackErr) + h.close() +} + +func (h *handshake) close() { + if nc, ok := h.conn.(net.Conn); ok && h.closeConn != nil { + h.closeConn(nc) + return + } _ = h.conn.Close() } @@ -187,6 +214,8 @@ func (h *handshake) readMsg(allowedTypes ...byte) (msg message, err error) { func (h *handshake) release() { h.buf = h.buf[:0] h.conn = nil + h.closeConn = nil + h.abandoned = nil h.localAck.Error = 0 h.remoteAck.Error = 0 h.remoteCred.Type = 0 diff --git a/net/secureservice/handshake/proto.go b/net/secureservice/handshake/proto.go index caedafd92..17d63530e 100644 --- a/net/secureservice/handshake/proto.go +++ b/net/secureservice/handshake/proto.go @@ -3,6 +3,8 @@ package handshake import ( "context" "net" + "sync/atomic" + "time" "golang.org/x/exp/slices" @@ -14,34 +16,91 @@ type ProtoChecker struct { SupportedEncodings []handshakeproto.Encoding } +// OutgoingProtoHandshake negotiates the sub-connection protocol. +// +// Contract: on an I/O error or ctx cancellation the conn is closed; on a +// protocol-level error (incompatible, declined or unexpected proto) it is +// left to the caller. On cancellation the function returns at once and the +// close happens in the background: a stream close can block on the transport +// (a yamux FIN waits up to the connection write timeout), and the caller is +// typically racing a deadline. I/O errors close it asynchronously as well, so +// the conn may still be open when the function returns; it is unusable +// afterwards either way. func OutgoingProtoHandshake(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto) (*handshakeproto.Proto, error) { + return OutgoingProtoHandshakeWithCloser(ctx, conn, proto, nil) +} + +// OutgoingProtoHandshakeWithCloser is OutgoingProtoHandshake with every close +// of conn going through closeConn, which must not block (e.g. it hands the +// conn to a bounded cleanup worker). With a closer the conn is handed to it +// exactly once on any error, protocol-level ones included, so the caller +// never closes it itself. A nil closeConn keeps OutgoingProtoHandshake's +// contract and closes in a new goroutine. +func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto, closeConn func(net.Conn)) (*handshakeproto.Proto, error) { if ctx == nil { ctx = context.Background() } + closeOnAnyErr := closeConn != nil + if closeConn == nil { + closeConn = closeAsync + } else { + var handedOff atomic.Bool + hook := closeConn + closeConn = func(c net.Conn) { + if handedOff.CompareAndSwap(false, true) { + hook(c) + } + } + } h := newHandshake() done := make(chan struct{}) var ( err error remoteProto *handshakeproto.Proto + // claimed is taken by whichever side finishes first: the handshake + // goroutine (its result is returned) or the cancelled caller (which + // then closes the conn itself) + claimed atomic.Bool ) go func() { defer close(done) - remoteProto, err = outgoingProtoHandshake(h, conn, proto) + remoteProto, err = outgoingProtoHandshake(h, conn, proto, closeConn, &claimed) + if err != nil && closeOnAnyErr { + // a no-op if the handshake or an abandoning caller closed it + closeConn(conn) + } + claimed.CompareAndSwap(false, true) }() select { case <-done: return remoteProto, err case <-ctx.Done(): - _ = conn.Close() + if !claimed.CompareAndSwap(false, true) { + // the handshake finished first and is returning right now + <-done + return remoteProto, err + } + // The deadline unblocks a pending read at once where supported; the + // close is what reliably ends the handshake everywhere (a yamux write + // waiting for the send loop ignores deadlines, and some conns have + // no deadlines at all). Neither blocks the caller. + _ = conn.SetDeadline(time.Now()) + closeConn(conn) return nil, ctx.Err() } } +func closeAsync(conn net.Conn) { + go func() { _ = conn.Close() }() +} + var noEncodings = []handshakeproto.Encoding{handshakeproto.Encoding_None} -func outgoingProtoHandshake(h *handshake, conn net.Conn, proto *handshakeproto.Proto) (remoteProto *handshakeproto.Proto, err error) { +func outgoingProtoHandshake(h *handshake, conn net.Conn, proto *handshakeproto.Proto, closeConn func(net.Conn), abandoned *atomic.Bool) (remoteProto *handshakeproto.Proto, err error) { defer h.release() h.conn = conn + h.closeConn = closeConn + h.abandoned = abandoned localProto := proto if err = h.writeProto(localProto); err != nil { h.tryWriteErrAndClose(err) diff --git a/net/secureservice/handshake/proto_test.go b/net/secureservice/handshake/proto_test.go index fd50fd7aa..8599d703b 100644 --- a/net/secureservice/handshake/proto_test.go +++ b/net/secureservice/handshake/proto_test.go @@ -1,6 +1,12 @@ package handshake import ( + "context" + "errors" + "io" + "net" + "os" + "sync" "testing" "time" @@ -201,3 +207,126 @@ func TestEndToEndProto(t *testing.T) { t.Log("dur", time.Since(st)) }) } + +// blockingCloseConn models a stream whose Close blocks on the transport +type blockingCloseConn struct { + net.Conn + closeCalled chan struct{} + release chan struct{} + once sync.Once +} + +func (c *blockingCloseConn) Close() error { + c.once.Do(func() { close(c.closeCalled) }) + <-c.release + return c.Conn.Close() +} + +func TestOutgoingProtoHandshake_CancelDoesNotWaitForClose(t *testing.T) { + c1, c2 := net.Pipe() + defer c2.Close() + conn := &blockingCloseConn{Conn: c1, closeCalled: make(chan struct{}), release: make(chan struct{})} + defer close(conn.release) + // the remote reads the proto but never answers + go func() { _, _ = io.Copy(io.Discard, c2) }() + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + start := time.Now() + _, err := OutgoingProtoHandshake(ctx, conn, &handshakeproto.Proto{Proto: 1}) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.Less(t, time.Since(start), time.Second, "the caller must not wait on the blocked close") + + // the conn is still closed, off the caller's path + select { + case <-conn.closeCalled: + case <-time.After(time.Second): + t.Fatal("abandoned handshake conn was not closed") + } +} + +// noDeadlineConn models a conn whose SetDeadline is a no-op (as on wasm) +type noDeadlineConn struct { + net.Conn +} + +func (noDeadlineConn) SetDeadline(time.Time) error { return nil } + +func TestOutgoingProtoHandshakeWithCloser_Cancel(t *testing.T) { + c1, c2 := net.Pipe() + defer c2.Close() + conn := noDeadlineConn{Conn: c1} + + // the remote reads the proto, then never answers + remoteRead := make(chan error, 1) + go func() { + h := newHandshake() + h.conn = c2 + if _, err := h.readMsg(msgTypeProto); err != nil { + remoteRead <- err + return + } + // whatever comes next must be the close, never an ack + _, err := c2.Read(make([]byte, 1)) + remoteRead <- err + }() + + closed := make(chan net.Conn, 1) + closer := func(c net.Conn) { + closed <- c + _ = c.Close() + } + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + _, err := OutgoingProtoHandshakeWithCloser(ctx, conn, &handshakeproto.Proto{Proto: 1}, closer) + require.ErrorIs(t, err, context.DeadlineExceeded) + + // a conn without deadline support is still closed, through the closer + select { + case c := <-closed: + assert.Equal(t, net.Conn(conn), c) + case <-time.After(time.Second): + t.Fatal("conn was not handed to the closer") + } + select { + case err = <-remoteRead: + assert.ErrorIs(t, err, io.EOF, "an abandoned handshake must not write an ack") + case <-time.After(time.Second): + t.Fatal("remote did not observe the close") + } +} + +func TestOutgoingProtoHandshakeWithCloser_IOErrorUsesCloser(t *testing.T) { + c1, c2 := net.Pipe() + // the remote goes away mid-handshake + go func() { + h := newHandshake() + h.conn = c2 + _, _ = h.readMsg(msgTypeProto) + _ = c2.Close() + }() + closed := make(chan struct{}, 1) + closer := func(c net.Conn) { + closed <- struct{}{} + _ = c.Close() + } + _, err := OutgoingProtoHandshakeWithCloser(context.Background(), c1, &handshakeproto.Proto{Proto: 1}, closer) + require.Error(t, err) + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("I/O error close did not go through the closer") + } +} + +func TestHandshakeError_Unwrap(t *testing.T) { + err := error(HandshakeError{Err: io.EOF}) + assert.ErrorIs(t, err, io.EOF) + assert.ErrorIs(t, HandshakeError{Err: os.ErrDeadlineExceeded}, os.ErrDeadlineExceeded) + // protocol-level sentinels keep matching by value and wrap nothing + assert.ErrorIs(t, ErrIncompatibleVersion, ErrIncompatibleVersion) + assert.NotErrorIs(t, ErrIncompatibleProto, ErrIncompatibleVersion) + assert.Nil(t, errors.Unwrap(ErrIncompatibleVersion)) + // a wrapped transport error is not mistaken for a protocol sentinel + assert.NotErrorIs(t, HandshakeError{Err: io.EOF}, ErrIncompatibleVersion) +} diff --git a/net/transport/iroh/conn.go b/net/transport/iroh/conn.go index efe7f6d69..a75404fc7 100644 --- a/net/transport/iroh/conn.go +++ b/net/transport/iroh/conn.go @@ -43,6 +43,11 @@ type irohMultiConn struct { bytesWritten atomic.Int64 } +// WriteTimeout implements transport.WriteTimeouter +func (c *irohMultiConn) WriteTimeout() time.Duration { + return c.writeTimeout +} + func (c *irohMultiConn) Context() context.Context { return c.cctx } diff --git a/net/transport/quic/conn.go b/net/transport/quic/conn.go index 4d163eeac..174b82f42 100644 --- a/net/transport/quic/conn.go +++ b/net/transport/quic/conn.go @@ -59,13 +59,25 @@ func (q *quicMultiConn) BytesWritten() int64 { return q.bytesWritten.Load() } +// WriteTimeout implements transport.WriteTimeouter +func (q *quicMultiConn) WriteTimeout() time.Duration { + return q.writeTimeout +} + func (q *quicMultiConn) Context() context.Context { return q.cctx } -// isConnDead reports whether err means the underlying QUIC connection is gone -// (idle timeout, peer-initiated close, or an already-closed connection), as -// opposed to a transient or stream-level error. +// isConnDead reports whether err means the underlying QUIC connection is gone, +// as opposed to a transient or stream-level error. +// +// Matching relies on net.ErrClosed: in quic-go (v0.63) every connection-level +// error unwraps to it, so this covers idle timeout, stateless reset, +// handshake timeout, transport errors, version negotiation failure and any +// application close code, besides an already-closed connection. A +// *quic.StreamError (a reset or cancelled stream) does not unwrap to it and +// is deliberately not matched: the connection itself is fine. The explicit +// checks below only document the cases that matter most. func isConnDead(err error) bool { if err == nil { return false @@ -74,13 +86,42 @@ func isConnDead(err error) bool { if errors.As(err, &idle) { return true } - var appErr *quic.ApplicationError - if errors.As(err, &appErr) && appErr.ErrorCode == 2 { + var reset *quic.StatelessResetError + if errors.As(err, &reset) { return true } return errors.Is(err, quic.ErrServerClosed) || errors.Is(err, net.ErrClosed) } +// connDeadError is a stream error caused by the whole connection going away. +// It matches transport.ErrConnClosed, so callers classify it like a failed +// Open or Accept, and still unwraps to the original quic error for telemetry. +type connDeadError struct { + cause error +} + +func (e connDeadError) Error() string { + return transport.ErrConnClosed.Error() + ": " + e.cause.Error() +} + +func (e connDeadError) Unwrap() []error { + return []error{transport.ErrConnClosed, e.cause} +} + +// wrapConnDead normalizes a stream Read/Write error: one meaning the +// connection is dead is wrapped into connDeadError, anything else (io.EOF, +// stream resets, deadlines) is returned as is. +func wrapConnDead(err error) error { + if err == nil || !isConnDead(err) { + return err + } + var already connDeadError + if errors.As(err, &already) { + return err + } + return connDeadError{cause: err} +} + func (q *quicMultiConn) Accept() (conn net.Conn, err error) { stream, err := q.connection.AcceptStream(context.Background()) if err != nil { @@ -201,7 +242,7 @@ func (q quicNetConn) Write(b []byte) (n int, err error) { if n > 0 && q.bytesWritten != nil { q.bytesWritten.Add(int64(n)) } - return + return n, wrapConnDead(err) } func (q quicNetConn) Read(b []byte) (n int, err error) { @@ -209,7 +250,7 @@ func (q quicNetConn) Read(b []byte) (n int, err error) { if n > 0 && q.bytesRead != nil { q.bytesRead.Add(int64(n)) } - return + return n, wrapConnDead(err) } func (q quicNetConn) LocalAddr() net.Addr { diff --git a/net/transport/quic/conn_errors_test.go b/net/transport/quic/conn_errors_test.go new file mode 100644 index 000000000..5943b1bc8 --- /dev/null +++ b/net/transport/quic/conn_errors_test.go @@ -0,0 +1,164 @@ +package quic + +import ( + "errors" + "fmt" + "io" + "net" + "os" + "testing" + "time" + + "github.com/quic-go/quic-go" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/anyproto/any-sync/net/transport" +) + +func TestWrapConnDead(t *testing.T) { + dead := []struct { + name string + err error + }{ + {"idle timeout", &quic.IdleTimeoutError{}}, + {"stateless reset", &quic.StatelessResetError{}}, + {"handshake timeout", &quic.HandshakeTimeoutError{}}, + {"transport error", &quic.TransportError{ErrorCode: quic.InternalError, Remote: true}}, + {"application close code 2", &quic.ApplicationError{ErrorCode: 2, Remote: true}}, + {"wrapped idle timeout", fmt.Errorf("read: %w", &quic.IdleTimeoutError{})}, + {"net.ErrClosed", net.ErrClosed}, + } + for _, tc := range dead { + t.Run(tc.name, func(t *testing.T) { + err := wrapConnDead(tc.err) + assert.True(t, errors.Is(err, transport.ErrConnClosed), "must match ErrConnClosed") + assert.True(t, errors.Is(err, net.ErrClosed), "must keep matching net.ErrClosed") + // the original error stays reachable + assert.True(t, errors.Is(err, tc.err) || errors.Is(err, errors.Unwrap(tc.err))) + assert.Contains(t, err.Error(), tc.err.Error()) + // wrapping is idempotent + assert.Equal(t, err, wrapConnDead(err)) + }) + } + t.Run("errors.As finds the quic error", func(t *testing.T) { + err := wrapConnDead(&quic.StatelessResetError{}) + var reset *quic.StatelessResetError + assert.True(t, errors.As(err, &reset)) + + err = wrapConnDead(&quic.ApplicationError{ErrorCode: 2, Remote: true}) + var appErr *quic.ApplicationError + require.True(t, errors.As(err, &appErr)) + assert.Equal(t, quic.ApplicationErrorCode(2), appErr.ErrorCode) + + err = wrapConnDead(&quic.TransportError{ErrorCode: quic.ProtocolViolation}) + var trErr *quic.TransportError + require.True(t, errors.As(err, &trErr)) + assert.Equal(t, quic.ProtocolViolation, trErr.ErrorCode) + }) + + notDead := []struct { + name string + err error + }{ + {"nil", nil}, + {"eof", io.EOF}, + {"stream reset", &quic.StreamError{StreamID: 4, ErrorCode: 0, Remote: true}}, + {"local stream cancel", &quic.StreamError{StreamID: 4, ErrorCode: 0}}, + {"write deadline", os.ErrDeadlineExceeded}, + {"other", errors.New("other")}, + } + for _, tc := range notDead { + t.Run("not dead: "+tc.name, func(t *testing.T) { + err := wrapConnDead(tc.err) + assert.Equal(t, tc.err, err, "must be returned unchanged") + if err != nil { + assert.False(t, errors.Is(err, transport.ErrConnClosed)) + } + }) + } +} + +func TestQuicNetConn_ConnCloseNormalized(t *testing.T) { + fxS := newFixture(t) + defer fxS.finish(t) + fxC := newFixture(t) + defer fxC.finish(t) + + mcC, err := fxC.Dial(ctx, fxS.addr) + require.NoError(t, err) + var mcS transport.MultiConn + select { + case mcS = <-fxS.accepter.mcs: + case <-time.After(time.Second * 5): + t.Fatal("timeout") + } + + conn, err := mcC.Open(ctx) + require.NoError(t, err) + _, err = conn.Write([]byte("hello")) + require.NoError(t, err) + sConn, err := mcS.Accept() + require.NoError(t, err) + buf := make([]byte, 5) + _, err = io.ReadFull(sConn, buf) + require.NoError(t, err) + + // the server closes the whole connection (application code 2) + require.NoError(t, mcS.Close()) + select { + case <-mcC.CloseChan(): + case <-time.After(5 * time.Second): + t.Fatal("client did not observe the close") + } + + _, err = conn.Read(buf) + require.Error(t, err) + assert.ErrorIs(t, err, transport.ErrConnClosed) + var appErr *quic.ApplicationError + require.ErrorAs(t, err, &appErr) + assert.Equal(t, quic.ApplicationErrorCode(2), appErr.ErrorCode) + + _, err = conn.Write([]byte("again")) + require.Error(t, err) + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.ErrorAs(t, err, &appErr) +} + +func TestQuicNetConn_StreamResetNotNormalized(t *testing.T) { + fxS := newFixture(t) + defer fxS.finish(t) + fxC := newFixture(t) + defer fxC.finish(t) + + mcC, err := fxC.Dial(ctx, fxS.addr) + require.NoError(t, err) + var mcS transport.MultiConn + select { + case mcS = <-fxS.accepter.mcs: + case <-time.After(time.Second * 5): + t.Fatal("timeout") + } + + conn, err := mcC.Open(ctx) + require.NoError(t, err) + _, err = conn.Write([]byte("hello")) + require.NoError(t, err) + sConn, err := mcS.Accept() + require.NoError(t, err) + buf := make([]byte, 5) + _, err = io.ReadFull(sConn, buf) + require.NoError(t, err) + + // only the stream goes away: the server cancels its read side, which + // makes the client's writes fail with a stream error + sConn.(quicNetConn).CancelRead(42) + require.Eventually(t, func() bool { + _, err = conn.Write([]byte("more")) + return err != nil + }, 5*time.Second, 10*time.Millisecond) + var streamErr *quic.StreamError + require.ErrorAs(t, err, &streamErr) + assert.False(t, errors.Is(err, transport.ErrConnClosed)) + assert.False(t, mcC.IsClosed()) +} diff --git a/net/transport/transport.go b/net/transport/transport.go index cb0797b9b..e8df49328 100644 --- a/net/transport/transport.go +++ b/net/transport/transport.go @@ -72,6 +72,13 @@ type MultiConn interface { BytesWritten() int64 } +// WriteTimeouter is optionally implemented by a MultiConn whose writes are +// bounded by a timeout; consumers use it to tell a slow close from a stalled +// transport +type WriteTimeouter interface { + WriteTimeout() time.Duration +} + type Accepter interface { Accept(mc MultiConn) (err error) } diff --git a/net/transport/webtransport/conn.go b/net/transport/webtransport/conn.go index 59a598f22..5745355b2 100644 --- a/net/transport/webtransport/conn.go +++ b/net/transport/webtransport/conn.go @@ -97,6 +97,11 @@ func (m *wtMultiConn) BytesWritten() int64 { return m.bytesWritten.Load() } +// WriteTimeout implements transport.WriteTimeouter +func (m *wtMultiConn) WriteTimeout() time.Duration { + return m.writeTimeout +} + func (m *wtMultiConn) Context() context.Context { return m.cctx } diff --git a/net/transport/yamux/conn.go b/net/transport/yamux/conn.go index 82e2d5fc6..8475044ba 100644 --- a/net/transport/yamux/conn.go +++ b/net/transport/yamux/conn.go @@ -2,8 +2,11 @@ package yamux import ( "context" + "errors" "io" "net" + "sync" + "sync/atomic" "time" "github.com/hashicorp/yamux" @@ -14,27 +17,147 @@ import ( ) func NewMultiConn(cctx context.Context, luConn *connutil.LastUsageConn, addr string, sess *yamux.Session) transport.MultiConn { + return newMultiConn(cctx, luConn, addr, sess, 0) +} + +func newMultiConn(cctx context.Context, luConn *connutil.LastUsageConn, addr string, sess *yamux.Session, writeTimeout time.Duration) *yamuxConn { cctx = peer.CtxWithPeerAddr(cctx, transport.Yamux+"://"+sess.RemoteAddr().String()) return &yamuxConn{ - ctx: cctx, - luConn: luConn, - addr: addr, - Session: sess, + ctx: cctx, + luConn: luConn, + addr: addr, + Session: sess, + writeTimeout: writeTimeout, + backlogFreed: make(chan struct{}), } } +// maxAbandonedOpens bounds the Session.Open helpers a single connection +// keeps running for callers that have already given up (see Open) +const maxAbandonedOpens = 16 + type yamuxConn struct { ctx context.Context luConn *connutil.LastUsageConn addr string *yamux.Session + // writeTimeout is the configured WriteTimeoutSec, which yamux uses as + // both ConnectionWriteTimeout and StreamCloseTimeout: a stream close is + // bounded by it, and the peer's cleanup owner derives its stall + // threshold from it (see WriteTimeout) + writeTimeout time.Duration + + backlogMu sync.Mutex + // abandonedOpens counts Open helpers still running after their caller + // gave up + abandonedOpens int + // backlogFreed is closed and replaced whenever abandonedOpens drops + backlogFreed chan struct{} +} + +type openResult struct { + conn net.Conn + err error } +// Open opens a new stream, bounded by ctx. yamux's Session.Open takes no +// context and blocks while too many SYNs are unacknowledged, which on a +// silent connection lasts until StreamOpenTimeout closes the session. It runs +// in a helper goroutine instead; a caller whose ctx ends leaves the helper +// behind, and the helper closes the stream if it arrives late. Opens in +// progress are not limited (a congested but healthy link must not be +// throttled further), only the abandoned helpers are: while +// maxAbandonedOpens of them are running, Open waits, bounded by ctx, for one +// to finish. The cap is checked before the helper starts, so callers racing +// past it together can overshoot it by their number. func (y *yamuxConn) Open(ctx context.Context) (conn net.Conn, err error) { - if conn, err = y.Session.Open(); err != nil { - return + if conn, err = y.open(ctx); err != nil { + return nil, err + } + return y.wrapStream(conn), nil +} + +func (y *yamuxConn) open(ctx context.Context) (conn net.Conn, err error) { + if ctx.Done() == nil { + // a context that can never end needs no helper + return y.Session.Open() + } + if err = ctx.Err(); err != nil { + return nil, err + } + if err = y.waitBacklog(ctx); err != nil { + return nil, err } - return + var ( + // claimed is taken by whichever side finishes first: the helper + // (it delivers the result) or the caller giving up (the helper then + // owns the stream and closes it) + claimed atomic.Bool + res = make(chan openResult, 1) + ) + go func() { + stream, sErr := y.Session.Open() + if claimed.CompareAndSwap(false, true) { + res <- openResult{conn: stream, err: sErr} + return + } + // the caller is gone: close the late stream. The helper counts as + // abandoned until the close returns. + if stream != nil { + _ = stream.Close() + } + y.backlogMu.Lock() + y.abandonedOpens-- + close(y.backlogFreed) + y.backlogFreed = make(chan struct{}) + y.backlogMu.Unlock() + }() + select { + case r := <-res: + return r.conn, r.err + case <-ctx.Done(): + if claimed.CompareAndSwap(false, true) { + y.backlogMu.Lock() + y.abandonedOpens++ + y.backlogMu.Unlock() + return nil, ctx.Err() + } + // the helper won the race and is delivering right now + r := <-res + return r.conn, r.err + } +} + +// waitBacklog waits until fewer than maxAbandonedOpens helpers are running +func (y *yamuxConn) waitBacklog(ctx context.Context) error { + for { + y.backlogMu.Lock() + if y.abandonedOpens < maxAbandonedOpens { + y.backlogMu.Unlock() + return nil + } + freed := y.backlogFreed + y.backlogMu.Unlock() + select { + case <-freed: + case <-ctx.Done(): + return ctx.Err() + case <-y.Session.CloseChan(): + return yamux.ErrSessionShutdown + } + } +} + +// abandoned returns the number of Open helpers left behind by their callers +func (y *yamuxConn) abandoned() int { + y.backlogMu.Lock() + defer y.backlogMu.Unlock() + return y.abandonedOpens +} + +// WriteTimeout implements transport.WriteTimeouter +func (y *yamuxConn) WriteTimeout() time.Duration { + return y.writeTimeout } func (y *yamuxConn) LastUsage() time.Time { @@ -64,5 +187,67 @@ func (y *yamuxConn) Accept() (conn net.Conn, err error) { } return } - return + return y.wrapStream(conn), nil +} + +func (y *yamuxConn) wrapStream(conn net.Conn) net.Conn { + return yamuxStream{Conn: conn, sess: y.Session} +} + +// yamuxStream normalizes stream errors caused by the whole session going +// away, like the QUIC transport does: on shutdown yamux force-closes every +// stream, so a pending Read returns a bare io.EOF and a Write +// ErrSessionShutdown or ErrStreamClosed, indistinguishable from a normal +// remote close. +type yamuxStream struct { + net.Conn + sess *yamux.Session +} + +func (s yamuxStream) Read(b []byte) (n int, err error) { + n, err = s.Conn.Read(b) + return n, s.wrapSessionDead(err) +} + +func (s yamuxStream) Write(b []byte) (n int, err error) { + n, err = s.Conn.Write(b) + return n, s.wrapSessionDead(err) +} + +// wrapSessionDead wraps err into sessionDeadError when it was caused by the +// session shutting down. io.EOF, a stream reset and a closed stream count +// only while the session is closed: on a live session they are stream-level +// outcomes (a remote close or reset) and are returned unchanged. +func (s yamuxStream) wrapSessionDead(err error) error { + if err == nil { + return nil + } + var already sessionDeadError + if errors.As(err, &already) { + return err + } + switch { + case errors.Is(err, yamux.ErrSessionShutdown): + case errors.Is(err, io.EOF), errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed): + if !s.sess.IsClosed() { + return err + } + default: + return err + } + return sessionDeadError{cause: err} +} + +// sessionDeadError matches transport.ErrConnClosed and still unwraps to the +// original yamux error +type sessionDeadError struct { + cause error +} + +func (e sessionDeadError) Error() string { + return transport.ErrConnClosed.Error() + ": " + e.cause.Error() +} + +func (e sessionDeadError) Unwrap() []error { + return []error{transport.ErrConnClosed, e.cause} } diff --git a/net/transport/yamux/conn_test.go b/net/transport/yamux/conn_test.go new file mode 100644 index 000000000..b9578dbde --- /dev/null +++ b/net/transport/yamux/conn_test.go @@ -0,0 +1,308 @@ +package yamux + +import ( + "context" + "io" + "net" + "testing" + "time" + + "github.com/hashicorp/yamux" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/anyproto/any-sync/net/connutil" + "github.com/anyproto/any-sync/net/transport" +) + +// newSessionPair returns a client MultiConn whose SYN semaphore holds a single +// slot, and the server session, which acknowledges a SYN only on Accept +func newSessionPair(t *testing.T) (*yamuxConn, *yamux.Session) { + cc, sc := net.Pipe() + conf := yamux.DefaultConfig() + conf.AcceptBacklog = 1 + conf.LogOutput = io.Discard + client, err := yamux.Client(cc, conf) + require.NoError(t, err) + server, err := yamux.Server(sc, conf) + require.NoError(t, err) + t.Cleanup(func() { + _ = client.Close() + _ = server.Close() + }) + mc := newMultiConn(context.Background(), connutil.NewLastUsageConn(cc), "pipe", client, time.Second) + return mc, server +} + +func TestYamuxConn_OpenHonoursContext(t *testing.T) { + mc, _ := newSessionPair(t) + + // the only SYN slot is taken by a stream the server never accepts + first, err := mc.Open(ctx) + require.NoError(t, err) + defer first.Close() + + const attempts = 50 + for i := 0; i < attempts; i++ { + octx, cancel := context.WithTimeout(ctx, 10*time.Millisecond) + start := time.Now() + conn, oErr := mc.Open(octx) + cancel() + require.Nil(t, conn) + require.ErrorIs(t, oErr, context.DeadlineExceeded) + require.Less(t, time.Since(start), time.Second, "Open must return on ctx") + } + // the helpers left behind are capped per connection; past the cap Open + // waits for the backlog within its ctx and leaves no helper + assert.Equal(t, maxAbandonedOpens, mc.abandoned()) + + // a caller waiting on the backlog returns as soon as the session goes + // away, and the session going away releases every helper + waitErr := make(chan error, 1) + go func() { + lctx, cancel := context.WithTimeout(ctx, time.Minute) + defer cancel() + _, oErr := mc.Open(lctx) + waitErr <- oErr + }() + time.Sleep(20 * time.Millisecond) + require.NoError(t, mc.Session.Close()) + select { + case oErr := <-waitErr: + require.Error(t, oErr) + require.NotErrorIs(t, oErr, context.DeadlineExceeded) + case <-time.After(5 * time.Second): + t.Fatal("waiting Open did not return on session close") + } + require.Eventually(t, func() bool { return mc.abandoned() == 0 }, 5*time.Second, 10*time.Millisecond) +} + +func TestYamuxConn_OpenWaitsForBacklog(t *testing.T) { + mc, _ := newSessionPair(t) + // simulate a full backlog of abandoned helpers + mc.backlogMu.Lock() + mc.abandonedOpens = maxAbandonedOpens + mc.backlogMu.Unlock() + + opened := make(chan error, 1) + go func() { + lctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + conn, oErr := mc.Open(lctx) + if oErr == nil { + _ = conn.Close() + } + opened <- oErr + }() + select { + case <-opened: + t.Fatal("Open must wait while the backlog is full") + case <-time.After(50 * time.Millisecond): + } + // one helper finishes: the waiting caller proceeds and opens + mc.backlogMu.Lock() + mc.abandonedOpens-- + close(mc.backlogFreed) + mc.backlogFreed = make(chan struct{}) + mc.backlogMu.Unlock() + select { + case oErr := <-opened: + require.NoError(t, oErr) + case <-time.After(5 * time.Second): + t.Fatal("Open did not proceed once the backlog dropped") + } +} + +func TestYamuxConn_OpensInProgressAreNotThrottled(t *testing.T) { + mc, _ := newSessionPair(t) + + first, err := mc.Open(ctx) + require.NoError(t, err) + defer first.Close() + + // many callers wait on a congested link, none of them gave up + const waiting = 3 * maxAbandonedOpens + results := make(chan error, waiting) + octx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + for i := 0; i < waiting; i++ { + go func() { + conn, oErr := mc.Open(octx) + if oErr == nil { + _ = conn.Close() + } + results <- oErr + }() + } + // a caller with a short deadline is not refused because of them + time.Sleep(50 * time.Millisecond) + sctx, scancel := context.WithTimeout(ctx, 10*time.Millisecond) + _, err = mc.Open(sctx) + scancel() + require.ErrorIs(t, err, context.DeadlineExceeded) + + // every waiting open returns once the session goes away, and no helper + // is left behind. (Letting the server catch up instead would hit a yamux + // quirk: the abandoned helper's stream, closed before the remote accepts + // it, is never acknowledged and holds the only SYN slot.) + require.NoError(t, mc.Session.Close()) + for i := 0; i < waiting; i++ { + select { + case <-results: + case <-time.After(5 * time.Second): + t.Fatal("open did not return") + } + } + require.Eventually(t, func() bool { return mc.abandoned() == 0 }, 5*time.Second, 10*time.Millisecond) +} + +func TestYamuxConn_OpenClosesLateStream(t *testing.T) { + mc, server := newSessionPair(t) + + first, err := mc.Open(ctx) + require.NoError(t, err) + defer first.Close() + + octx, cancel := context.WithTimeout(ctx, 10*time.Millisecond) + _, err = mc.Open(octx) + cancel() + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Equal(t, 1, mc.abandoned()) + + // accepting the first stream frees the SYN slot: the helper's stream + // arrives after its caller left, so the helper closes it + _, err = server.Accept() + require.NoError(t, err) + late, err := server.Accept() + require.NoError(t, err) + require.NoError(t, late.SetReadDeadline(time.Now().Add(5*time.Second))) + _, err = io.ReadAll(late) + require.NoError(t, err, "the late stream must be closed by the helper") + require.Eventually(t, func() bool { return mc.abandoned() == 0 }, 5*time.Second, 10*time.Millisecond) + // Not asserted: further opens on this session. When the helper's FIN + // reaches the remote before its Accept, yamux never acknowledges the + // stream and its SYN slot stays taken until StreamOpenTimeout (a known + // upstream quirk). +} + +func TestYamuxConn_OpenCancelledContext(t *testing.T) { + mc, _ := newSessionPair(t) + cctx, cancel := context.WithCancel(ctx) + cancel() + _, err := mc.Open(cctx) + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, mc.abandoned()) +} + +func TestYamuxConn_OpenClosedSession(t *testing.T) { + mc, _ := newSessionPair(t) + require.NoError(t, mc.Session.Close()) + _, err := mc.Open(ctx) + require.ErrorIs(t, err, yamux.ErrSessionShutdown) +} + +// streamPair opens a stream from mc and accepts it on server +func streamPair(t *testing.T, mc *yamuxConn, server *yamux.Session) (client, remote net.Conn) { + accepted := make(chan net.Conn, 1) + go func() { + s, err := server.Accept() + if err == nil { + accepted <- s + } + }() + client, err := mc.Open(ctx) + require.NoError(t, err) + // the server sees the stream only once something is written + _, err = client.Write([]byte("x")) + require.NoError(t, err) + select { + case remote = <-accepted: + case <-time.After(5 * time.Second): + t.Fatal("stream not accepted") + } + buf := make([]byte, 1) + _, err = io.ReadFull(remote, buf) + require.NoError(t, err) + return client, remote +} + +func readErr(t *testing.T, conn net.Conn) chan error { + res := make(chan error, 1) + go func() { + _, err := conn.Read(make([]byte, 16)) + res <- err + }() + return res +} + +func waitErr(t *testing.T, ch chan error) error { + select { + case err := <-ch: + return err + case <-time.After(5 * time.Second): + t.Fatal("read did not return") + return nil + } +} + +func TestYamuxStream_SessionDeathNormalized(t *testing.T) { + t.Run("local session close mid-read", func(t *testing.T) { + mc, server := newSessionPair(t) + client, _ := streamPair(t, mc, server) + res := readErr(t, client) + time.Sleep(20 * time.Millisecond) + require.NoError(t, mc.Session.Close()) + err := waitErr(t, res) + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.ErrorIs(t, err, io.EOF, "the original error stays reachable") + + _, err = client.Write([]byte("more")) + assert.ErrorIs(t, err, transport.ErrConnClosed) + // yamux force-closed the stream on shutdown + assert.ErrorIs(t, err, yamux.ErrStreamClosed) + }) + t.Run("remote session close mid-read", func(t *testing.T) { + mc, server := newSessionPair(t) + client, _ := streamPair(t, mc, server) + res := readErr(t, client) + time.Sleep(20 * time.Millisecond) + require.NoError(t, server.Close()) + err := waitErr(t, res) + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.True(t, mc.Session.IsClosed()) + }) + t.Run("accepted stream, remote session close", func(t *testing.T) { + mc, server := newSessionPair(t) + accepted := make(chan net.Conn, 1) + go func() { + if s, err := mc.Accept(); err == nil { + accepted <- s + } + }() + remote, err := server.Open() + require.NoError(t, err) + _, err = remote.Write([]byte("x")) + require.NoError(t, err) + var local net.Conn + select { + case local = <-accepted: + case <-time.After(5 * time.Second): + t.Fatal("stream not accepted") + } + _, err = io.ReadFull(local, make([]byte, 1)) + require.NoError(t, err) + res := readErr(t, local) + time.Sleep(20 * time.Millisecond) + require.NoError(t, server.Close()) + assert.ErrorIs(t, waitErr(t, res), transport.ErrConnClosed) + }) + t.Run("remote stream close is a plain EOF", func(t *testing.T) { + mc, server := newSessionPair(t) + client, remote := streamPair(t, mc, server) + res := readErr(t, client) + require.NoError(t, remote.Close()) + err := waitErr(t, res) + assert.Equal(t, io.EOF, err, "a normal stream close must stay io.EOF") + assert.False(t, mc.Session.IsClosed()) + }) +} diff --git a/net/transport/yamux/yamux.go b/net/transport/yamux/yamux.go index ef921ad7a..28be3adb4 100644 --- a/net/transport/yamux/yamux.go +++ b/net/transport/yamux/yamux.go @@ -132,7 +132,7 @@ func (y *yamuxTransport) Dial(ctx context.Context, addr string) (mc transport.Mu if err != nil { return } - mc = NewMultiConn(cctx, luc, addr, sess) + mc = newMultiConn(cctx, luc, addr, sess, time.Duration(y.conf.WriteTimeoutSec)*time.Second) return } @@ -179,7 +179,7 @@ func (y *yamuxTransport) accept(conn net.Conn) { log.Info("incoming connection yamux session error", zap.Error(err), zap.String("remoteAddr", conn.RemoteAddr().String())) return } - mc := NewMultiConn(cctx, luc, conn.RemoteAddr().String(), sess) + mc := newMultiConn(cctx, luc, conn.RemoteAddr().String(), sess, time.Duration(y.conf.WriteTimeoutSec)*time.Second) if err = y.accepter.Accept(mc); err != nil { log.Info("connection accept error", zap.Error(err), zap.String("remoteAddr", conn.RemoteAddr().String())) } From 950bebaf864ee8cf5cf9c280c628a55d58addfa9 Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 12:59:45 +0200 Subject: [PATCH 2/7] net/peer: drop cleanup owner stall escalation and WriteTimeouter Each sub-conn close is bounded by the transport (yamux write/close timeouts, non-blocking QUIC/iroh/webtransport close), and in-flight closes count toward the peer's open limiter, so a backlog throttles new opens instead of needing a stall detector that kills the whole multiconn. Remove the escalation, transport.WriteTimeouter and the yamux writeTimeout plumbing; keep lazy capped workers, the pending list and inFlight. Log once when a single close runs past a minute. --- net/peer/cleanup.go | 133 ++++-------- net/peer/cleanup_test.go | 324 +++++++++++++++++++++++------ net/peer/peer.go | 2 +- net/transport/iroh/conn.go | 5 - net/transport/quic/conn.go | 5 - net/transport/transport.go | 7 - net/transport/webtransport/conn.go | 5 - net/transport/yamux/conn.go | 15 -- net/transport/yamux/conn_test.go | 4 +- net/transport/yamux/yamux.go | 4 +- 10 files changed, 308 insertions(+), 196 deletions(-) diff --git a/net/peer/cleanup.go b/net/peer/cleanup.go index 575baf0a6..769646f6e 100644 --- a/net/peer/cleanup.go +++ b/net/peer/cleanup.go @@ -7,33 +7,16 @@ import ( "time" "go.uber.org/zap" - - "github.com/anyproto/any-sync/net/transport" ) const ( // cleanupMaxWorkers bounds the sub-connection closes a peer runs at once cleanupMaxWorkers = 64 - // cleanupDefaultStallTimeout is the stall threshold for a transport that - // does not report its write timeout: twice the default yamux one - cleanupDefaultStallTimeout = 20 * time.Second + // cleanupSlowClose is how long a single close may run before it is + // logged as hung; nothing else happens to it + cleanupSlowClose = time.Minute ) -// closeStallTimeout is how long a single close may run before the transport -// counts as stalled. A healthy close takes about one round trip (drpc waits -// for the remote FIN), but on a congested link a FIN may wait up to the -// transport's write timeout, so the threshold is twice that. For yamux this -// is WriteTimeoutSec, which also sets ConnectionWriteTimeout and -// StreamCloseTimeout: a single stream close cannot legitimately outlast it. -func closeStallTimeout(mc transport.MultiConn) time.Duration { - if wt, ok := mc.(transport.WriteTimeouter); ok { - if d := wt.WriteTimeout(); d > 0 { - return 2 * d - } - } - return cleanupDefaultStallTimeout -} - // cleanupOwner closes a peer's sub connections off the callers' path. Closing // a drpc conn waits for its reader, stream manager and transport, and a yamux // stream close sends a FIN under a write timeout, so on a stalled connection a @@ -42,43 +25,42 @@ func closeStallTimeout(mc transport.MultiConn) time.Duration { // Workers are started on demand, up to cleanupMaxWorkers, and exit once there // is nothing left to close, so an idle peer costs no goroutines. close never // blocks and never drops a close; closes beyond the workers wait in a pending -// list. Its size is bounded by the peer's sub conns, and the peer's open -// limiter counts every close in flight (inFlight), so a peer that closes -// faster than the transport can keep up is throttled rather than piling up. +// list. // -// Only on evidence of a stall, every worker busy and one of them on a close -// older than stallTimeout, is the whole MultiConn closed, which makes every -// pending and further close quick. A burst or a sustained rate of closes on a -// healthy connection just queues. The check runs when a close is handed over, -// so a stall is detected on the first close queued after stallTimeout; -// meanwhile the stuck closes are still bounded by the transport's own -// timeouts, so the cost of the delay is latency only. +// The pending list is not capped, and nothing escalates on a slow close (one +// running longer than cleanupSlowClose is only logged). The backlog is +// bounded in practice, not by construction: +// - every close comes from a sub conn this peer opened, and the peer's open +// limiter counts closes in flight (inFlight): past its threshold each +// new open waits 100ms per extra conn, so opens settle at about 10/s per +// peer while closes fall behind, and the list grows ever more slowly; +// - each close is bounded by the transport: yamux stream closes by +// StreamCloseTimeout and ConnectionWriteTimeout (both WriteTimeoutSec, +// never 0), while QUIC, iroh and webtransport closes do not block; +// - callers give up on their own deadlines rather than opening forever. +// +// Keepalive is not a bound: it can be disabled. type cleanupOwner struct { - mc transport.MultiConn - stallTimeout time.Duration + peerId string mu sync.Mutex pending []io.Closer - // running is the set of live workers - running map[*cleanupWorker]struct{} + workers int // inflight counts closes handed over and not yet finished - inflight atomic.Int32 - escalated atomic.Bool - escalations atomic.Int64 -} + inflight atomic.Int32 -type cleanupWorker struct { - // started is when the current close began; guarded by cleanupOwner.mu - started time.Time + // slowClose and onSlowClose are fields for tests + slowClose time.Duration + onSlowClose func(cl io.Closer) } -func newCleanupOwner(mc transport.MultiConn) *cleanupOwner { - return &cleanupOwner{ - mc: mc, - stallTimeout: closeStallTimeout(mc), - running: map[*cleanupWorker]struct{}{}, +func newCleanupOwner(peerId string) *cleanupOwner { + c := &cleanupOwner{peerId: peerId, slowClose: cleanupSlowClose} + c.onSlowClose = func(io.Closer) { + log.Warn("sub connection close is taking too long", zap.String("peerId", c.peerId), zap.Duration("after", c.slowClose)) } + return c } // close hands cl over to be closed in the background. It never blocks. @@ -90,40 +72,23 @@ func (c *cleanupOwner) close(cl io.Closer) { } c.inflight.Add(1) c.mu.Lock() - if len(c.running) < cleanupMaxWorkers { - w := &cleanupWorker{started: time.Now()} - c.running[w] = struct{}{} - c.mu.Unlock() - go c.work(w, cl) - return - } - if !c.stalledLocked() { - c.pending = append(c.pending, cl) + if c.workers < cleanupMaxWorkers { + c.workers++ c.mu.Unlock() + go c.work(cl) return } + c.pending = append(c.pending, cl) c.mu.Unlock() - c.escalate(cl) -} - -// stalledLocked reports whether a running close has outlived stallTimeout -func (c *cleanupOwner) stalledLocked() bool { - now := time.Now() - for w := range c.running { - if now.Sub(w.started) > c.stallTimeout { - return true - } - } - return false } -func (c *cleanupOwner) work(w *cleanupWorker, cl io.Closer) { +func (c *cleanupOwner) work(cl io.Closer) { for { - _ = cl.Close() + c.closeOne(cl) c.inflight.Add(-1) c.mu.Lock() if len(c.pending) == 0 { - delete(c.running, w) + c.workers-- c.pending = nil c.mu.Unlock() return @@ -131,31 +96,15 @@ func (c *cleanupOwner) work(w *cleanupWorker, cl io.Closer) { cl = c.pending[0] c.pending[0] = nil c.pending = c.pending[1:] - w.started = time.Now() c.mu.Unlock() } } -// escalate closes the whole connection: the transport is stalled, so the -// closes queued behind it would otherwise wait out its timeouts. A dead -// transport makes cl's close quick; the goroutines spawned here are bounded by -// the sub conns alive when the MultiConn closed, as no new ones can be opened -// afterwards. -func (c *cleanupOwner) escalate(cl io.Closer) { - c.escalations.Add(1) - first := c.escalated.CompareAndSwap(false, true) - if first { - log.Warn("sub connection cleanup is stalled: closing the connection") - } - go func() { - if first { - if err := c.mc.Close(); err != nil { - log.Debug("close connection on stalled cleanup", zap.Error(err)) - } - } - _ = cl.Close() - c.inflight.Add(-1) - }() +// closeOne closes cl, logging once if the close outlives slowClose +func (c *cleanupOwner) closeOne(cl io.Closer) { + timer := time.AfterFunc(c.slowClose, func() { c.onSlowClose(cl) }) + _ = cl.Close() + timer.Stop() } // inFlight returns the number of closes handed over and not yet finished @@ -170,5 +119,5 @@ func (c *cleanupOwner) inFlight() int { func (c *cleanupOwner) stats() (running, pending int) { c.mu.Lock() defer c.mu.Unlock() - return len(c.running), len(c.pending) + return c.workers, len(c.pending) } diff --git a/net/peer/cleanup_test.go b/net/peer/cleanup_test.go index 889813d0e..8a26e7907 100644 --- a/net/peer/cleanup_test.go +++ b/net/peer/cleanup_test.go @@ -18,7 +18,6 @@ import ( "github.com/anyproto/any-sync/net/connutil" "github.com/anyproto/any-sync/net/secureservice/handshake" "github.com/anyproto/any-sync/net/secureservice/handshake/handshakeproto" - "github.com/anyproto/any-sync/net/transport/mock_transport" ) type rawMsg []byte @@ -90,27 +89,9 @@ func (q *quickCloser) Close() error { return nil } -type writeTimeoutMC struct { - *mock_transport.MockMultiConn - wt time.Duration -} - -func (w writeTimeoutMC) WriteTimeout() time.Duration { return w.wt } - func TestCleanupOwner(t *testing.T) { - newMC := func(t *testing.T) (*mock_transport.MockMultiConn, *atomic.Int32) { - ctrl := gomock.NewController(t) - mc := mock_transport.NewMockMultiConn(ctrl) - var mcCloses atomic.Int32 - mc.EXPECT().Close().DoAndReturn(func() error { - mcCloses.Add(1) - return nil - }).AnyTimes() - return mc, &mcCloses - } - t.Run("burst on a healthy connection never closes it", func(t *testing.T) { - mc, mcCloses := newMC(t) - c := newCleanupOwner(mc) + t.Run("burst never blocks the caller and never drops a close", func(t *testing.T) { + c := newCleanupOwner("p1") const burst = 300 var running, maxSeen atomic.Int32 @@ -136,8 +117,6 @@ func TestCleanupOwner(t *testing.T) { t.Fatal("a close was dropped") } } - assert.Zero(t, c.escalations.Load()) - assert.Zero(t, mcCloses.Load(), "a healthy connection must not be closed") assert.LessOrEqual(t, int(maxSeen.Load()), cleanupMaxWorkers) // workers exit once idle: an idle peer costs no goroutines require.Eventually(t, func() bool { @@ -145,39 +124,27 @@ func TestCleanupOwner(t *testing.T) { return r == 0 && p == 0 }, time.Second, time.Millisecond) }) - t.Run("stalled transport closes the multiconn", func(t *testing.T) { - mc, mcCloses := newMC(t) - c := newCleanupOwner(mc) - c.stallTimeout = 50 * time.Millisecond - + t.Run("saturated workers queue and drain", func(t *testing.T) { + c := newCleanupOwner("p1") release := make(chan struct{}) + var releaseOnce sync.Once + doRelease := func() { releaseOnce.Do(func() { close(release) }) } + defer doRelease() + var closers []*blockingCloser - enqueue := func() *blockingCloser { + for i := 0; i < cleanupMaxWorkers+5; i++ { cl := newBlockingCloser(release) closers = append(closers, cl) + start := time.Now() c.close(cl) - return cl - } - for i := 0; i < cleanupMaxWorkers; i++ { - enqueue() + require.Less(t, time.Since(start), 100*time.Millisecond, "close must never block") } - // saturated but not stalled yet: the close just queues - enqueue() r, p := c.stats() assert.Equal(t, cleanupMaxWorkers, r) - assert.Equal(t, 1, p) - assert.Zero(t, c.escalations.Load()) - - time.Sleep(2 * c.stallTimeout) - start := time.Now() - enqueue() - enqueue() - assert.Less(t, time.Since(start), 100*time.Millisecond, "close must never block") - assert.Equal(t, int64(2), c.escalations.Load()) - require.Eventually(t, func() bool { return mcCloses.Load() == 1 }, time.Second, time.Millisecond) + assert.Equal(t, 5, p) + assert.Equal(t, cleanupMaxWorkers+5, c.inFlight()) - // nothing is dropped, and the multiconn is closed only once - close(release) + doRelease() for _, cl := range closers { select { case <-cl.closed: @@ -185,11 +152,42 @@ func TestCleanupOwner(t *testing.T) { t.Fatal("a queued close was dropped") } } - assert.Equal(t, int32(1), mcCloses.Load()) + require.Eventually(t, func() bool { + r, p := c.stats() + return r == 0 && p == 0 && c.inFlight() == 0 + }, time.Second, time.Millisecond) }) - t.Run("sustained close rate on a healthy connection never closes it", func(t *testing.T) { - mc, mcCloses := newMC(t) - c := newCleanupOwner(mc) + t.Run("a hung close blocks neither the caller nor other closes", func(t *testing.T) { + c := newCleanupOwner("p1") + hang := make(chan struct{}) + defer close(hang) + hung := newBlockingCloser(hang) + start := time.Now() + c.close(hung) + require.Less(t, time.Since(start), 100*time.Millisecond) + + var running, maxSeen atomic.Int32 + var closers []*quickCloser + for i := 0; i < 200; i++ { + cl := &quickCloser{d: time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} + closers = append(closers, cl) + c.close(cl) + } + for _, cl := range closers { + select { + case <-cl.closed: + case <-time.After(5 * time.Second): + t.Fatal("a close was held up by the hung one") + } + } + // only the hung close is left, on its own worker + require.Eventually(t, func() bool { + r, p := c.stats() + return r == 1 && p == 0 && c.inFlight() == 1 + }, time.Second, time.Millisecond) + }) + t.Run("sustained close rate", func(t *testing.T) { + c := newCleanupOwner("p1") // far more closes than workers, arriving faster than they finish const total = 3000 @@ -213,24 +211,35 @@ func TestCleanupOwner(t *testing.T) { t.Fatal("a close was dropped") } } - assert.Zero(t, c.escalations.Load()) - assert.Zero(t, mcCloses.Load(), "a healthy connection must not be closed") + assert.LessOrEqual(t, int(maxSeen.Load()), cleanupMaxWorkers) // closes in flight are visible to the peer's open limiter assert.Greater(t, sawInFlight, cleanupMaxWorkers) require.Eventually(t, func() bool { return c.inFlight() == 0 }, time.Second, time.Millisecond) }) - t.Run("stall threshold follows the transport write timeout", func(t *testing.T) { - mc, _ := newMC(t) - assert.Equal(t, cleanupDefaultStallTimeout, newCleanupOwner(mc).stallTimeout) - assert.Equal(t, 30*time.Second, newCleanupOwner(writeTimeoutMC{mc, 15 * time.Second}).stallTimeout) - assert.Equal(t, cleanupDefaultStallTimeout, newCleanupOwner(writeTimeoutMC{mc, 0}).stallTimeout) + t.Run("slow close is logged once", func(t *testing.T) { + c := newCleanupOwner("p1") + c.slowClose = 20 * time.Millisecond + var logged atomic.Int32 + c.onSlowClose = func(io.Closer) { logged.Add(1) } + release := make(chan struct{}) + cl := newBlockingCloser(release) + c.close(cl) + require.Eventually(t, func() bool { return logged.Load() == 1 }, time.Second, time.Millisecond) + require.Never(t, func() bool { return logged.Load() > 1 }, 100*time.Millisecond, 10*time.Millisecond) + close(release) + <-cl.closed + // a quick close is not logged + quick := newBlockingCloser(release) + c.close(quick) + <-quick.closed + require.Never(t, func() bool { return logged.Load() > 1 }, 50*time.Millisecond, 10*time.Millisecond) }) t.Run("nil owner", func(t *testing.T) { var c *cleanupOwner release := make(chan struct{}) - close(release) cl := newBlockingCloser(release) - c.close(cl) + returnsWithin(t, 100*time.Millisecond, "close must never block", func() { c.close(cl) }) + close(release) select { case <-cl.closed: case <-time.After(time.Second): @@ -359,14 +368,13 @@ func TestPeer_RPCDeadlineWithBlockedClose(t *testing.T) { assert.Empty(t, fx.inactive, "repetition %d", i) assert.Empty(t, fx.active, "repetition %d", i) fx.mu.Unlock() - // at most one blocked close per repetition, nothing escalated + // at most one blocked close per repetition running, _ := fx.cleanup.stats() require.LessOrEqual(t, running, i+1) } connsMu.Lock() require.Len(t, conns, repetitions, "each repetition opens a fresh sub conn") connsMu.Unlock() - require.Zero(t, fx.cleanup.escalations.Load()) // once the transport lets go, every stream is closed and nothing leaks releaseAll() @@ -579,3 +587,195 @@ func TestPeer_NilWakeKeepsThrottling(t *testing.T) { t.Fatal("acquire did not return on ctx") } } + +// TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight: a close that never +// returns counts towards the open limiter while it is in flight, and nothing +// more +func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + hang := make(chan struct{}) + var hangOnce sync.Once + unhang := func() { hangOnce.Do(func() { close(hang) }) } + defer unhang() + // one past the limiter threshold: the next open waits slowDownStep + var hung []*blockingCloser + for i := 0; i <= fx.limiter.startThreshold; i++ { + cl := newBlockingCloser(hang) + hung = append(hung, cl) + fx.cleanup.close(cl) + } + + var opens atomic.Int32 + in, out := net.Pipe() + defer out.Close() + go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() + fx.mc.EXPECT().Open(gomock.Any()).DoAndReturn(func(context.Context) (net.Conn, error) { + opens.Add(1) + return in, nil + }).Times(1) + + actx, cancel := context.WithTimeout(ctx, fx.limiter.slowDownStep/2) + _, err := fx.AcquireDrpcConn(actx) + cancel() + require.ErrorIs(t, err, context.DeadlineExceeded, "throttled while the closes are in flight") + require.Zero(t, opens.Load()) + + unhang() + for _, cl := range hung { + <-cl.closed + } + require.Eventually(t, func() bool { return fx.cleanup.inFlight() == 0 }, time.Second, time.Millisecond) + actx, cancel = context.WithTimeout(ctx, fx.limiter.slowDownStep/2) + defer cancel() + _, err = fx.AcquireDrpcConn(actx) + require.NoError(t, err, "no throttling once the closes are done") + require.Equal(t, int32(1), opens.Load()) +} + +// pendingConn is a released sub conn that is not closed yet, never unblocks, +// and whose Close blocks until released +type pendingConn struct { + closedCh chan struct{} + release chan struct{} + closeCalls atomic.Int32 + once sync.Once +} + +func newPendingConn(release chan struct{}) *pendingConn { + return &pendingConn{closedCh: make(chan struct{}), release: release} +} + +func (c *pendingConn) Close() error { + c.closeCalls.Add(1) + <-c.release + c.once.Do(func() { close(c.closedCh) }) + return nil +} +func (c *pendingConn) Closed() <-chan struct{} { return c.closedCh } +func (c *pendingConn) Unblocked() <-chan struct{} { return nil } +func (c *pendingConn) NewStream(context.Context, string, drpc.Encoding) (drpc.Stream, error) { + return nil, io.EOF +} +func (c *pendingConn) Invoke(context.Context, string, drpc.Encoding, drpc.Message, drpc.Message) error { + return io.EOF +} + +// returnsWithin fails the test, instead of hanging it, when fn blocks +func returnsWithin(t *testing.T, d time.Duration, msg string, fn func()) { + done := make(chan struct{}) + go func() { + defer close(done) + fn() + }() + select { + case <-done: + case <-time.After(d): + t.Fatal(msg) + } +} + +func waitClosed(t *testing.T, c *pendingConn) { + select { + case <-c.closedCh: + case <-time.After(time.Second): + t.Fatal("the conn was never closed") + } +} + +func TestPeer_ReleaseClosesInBackground(t *testing.T) { + newActive := func(fx *fixture, release chan struct{}) (*subConn, *pendingConn) { + pc := newPendingConn(release) + sc := &subConn{ConnUnblocked: pc} + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() + return sc, pc + } + t.Run("cancelled ctx", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + sc, pc := newActive(fx, release) + cctx, cancel := context.WithCancel(ctx) + cancel() + + returnsWithin(t, 100*time.Millisecond, "the blocked close must not run on the caller", func() { + fx.ReleaseDrpcConn(cctx, sc) + }) + require.Equal(t, 1, fx.cleanup.inFlight()) + fx.mu.Lock() + assert.Empty(t, fx.inactive) + assert.Empty(t, fx.active) + fx.mu.Unlock() + + close(release) + waitClosed(t, pc) + require.Eventually(t, func() bool { return fx.cleanup.inFlight() == 0 }, time.Second, time.Millisecond) + }) + t.Run("never unblocked", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + sc, pc := newActive(fx, release) + + // it waits the 200ms reuse window, then hands the close over + returnsWithin(t, 300*time.Millisecond, "the blocked close must not run on the caller", func() { + fx.ReleaseDrpcConn(ctx, sc) + }) + require.Equal(t, 1, fx.cleanup.inFlight()) + fx.mu.Lock() + assert.Empty(t, fx.inactive, "an unfinished conn is not reused") + fx.mu.Unlock() + + close(release) + waitClosed(t, pc) + }) +} + +func TestPeer_GCClosesInBackground(t *testing.T) { + newSub := func(t *testing.T, release chan struct{}) (*subConn, *pendingConn) { + a, b := net.Pipe() + t.Cleanup(func() { _ = a.Close(); _ = b.Close() }) + pc := newPendingConn(release) + // a LastUsageConn never used reports a zero last usage: expired + return &subConn{ConnUnblocked: pc, LastUsageConn: connutil.NewLastUsageConn(a)}, pc + } + t.Run("expired inactive conn", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + sc, pc := newSub(t, release) + fx.mu.Lock() + fx.inactive = append(fx.inactive, sc) + fx.mu.Unlock() + + returnsWithin(t, 100*time.Millisecond, "gc must not wait on the close", func() { + fx.gc(time.Millisecond) + }) + fx.mu.Lock() + assert.Empty(t, fx.inactive) + fx.mu.Unlock() + require.Equal(t, 1, fx.cleanup.inFlight()) + close(release) + waitClosed(t, pc) + }) + t.Run("doomed active conn", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + fx.mc.EXPECT().Addr().Return("").AnyTimes() + release := make(chan struct{}) + sc, pc := newSub(t, release) + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() + + returnsWithin(t, 100*time.Millisecond, "gc must not wait on the close", func() { + fx.gc(time.Millisecond) + }) + require.True(t, sc.doomed.Load()) + require.Equal(t, 1, fx.cleanup.inFlight()) + close(release) + waitClosed(t, pc) + }) +} diff --git a/net/peer/peer.go b/net/peer/peer.go index 3ccc98919..6dc4aa505 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -57,7 +57,7 @@ func NewPeer(mc transport.MultiConn, ctrl connCtrl) (p Peer, err error) { if pr.id, err = CtxPeerId(ctx); err != nil { return } - pr.cleanup = newCleanupOwner(mc) + pr.cleanup = newCleanupOwner(pr.id) go pr.acceptLoop() return pr, nil } diff --git a/net/transport/iroh/conn.go b/net/transport/iroh/conn.go index a75404fc7..efe7f6d69 100644 --- a/net/transport/iroh/conn.go +++ b/net/transport/iroh/conn.go @@ -43,11 +43,6 @@ type irohMultiConn struct { bytesWritten atomic.Int64 } -// WriteTimeout implements transport.WriteTimeouter -func (c *irohMultiConn) WriteTimeout() time.Duration { - return c.writeTimeout -} - func (c *irohMultiConn) Context() context.Context { return c.cctx } diff --git a/net/transport/quic/conn.go b/net/transport/quic/conn.go index 174b82f42..5482f22e8 100644 --- a/net/transport/quic/conn.go +++ b/net/transport/quic/conn.go @@ -59,11 +59,6 @@ func (q *quicMultiConn) BytesWritten() int64 { return q.bytesWritten.Load() } -// WriteTimeout implements transport.WriteTimeouter -func (q *quicMultiConn) WriteTimeout() time.Duration { - return q.writeTimeout -} - func (q *quicMultiConn) Context() context.Context { return q.cctx } diff --git a/net/transport/transport.go b/net/transport/transport.go index e8df49328..cb0797b9b 100644 --- a/net/transport/transport.go +++ b/net/transport/transport.go @@ -72,13 +72,6 @@ type MultiConn interface { BytesWritten() int64 } -// WriteTimeouter is optionally implemented by a MultiConn whose writes are -// bounded by a timeout; consumers use it to tell a slow close from a stalled -// transport -type WriteTimeouter interface { - WriteTimeout() time.Duration -} - type Accepter interface { Accept(mc MultiConn) (err error) } diff --git a/net/transport/webtransport/conn.go b/net/transport/webtransport/conn.go index 5745355b2..59a598f22 100644 --- a/net/transport/webtransport/conn.go +++ b/net/transport/webtransport/conn.go @@ -97,11 +97,6 @@ func (m *wtMultiConn) BytesWritten() int64 { return m.bytesWritten.Load() } -// WriteTimeout implements transport.WriteTimeouter -func (m *wtMultiConn) WriteTimeout() time.Duration { - return m.writeTimeout -} - func (m *wtMultiConn) Context() context.Context { return m.cctx } diff --git a/net/transport/yamux/conn.go b/net/transport/yamux/conn.go index 8475044ba..fa03d249a 100644 --- a/net/transport/yamux/conn.go +++ b/net/transport/yamux/conn.go @@ -17,17 +17,12 @@ import ( ) func NewMultiConn(cctx context.Context, luConn *connutil.LastUsageConn, addr string, sess *yamux.Session) transport.MultiConn { - return newMultiConn(cctx, luConn, addr, sess, 0) -} - -func newMultiConn(cctx context.Context, luConn *connutil.LastUsageConn, addr string, sess *yamux.Session, writeTimeout time.Duration) *yamuxConn { cctx = peer.CtxWithPeerAddr(cctx, transport.Yamux+"://"+sess.RemoteAddr().String()) return &yamuxConn{ ctx: cctx, luConn: luConn, addr: addr, Session: sess, - writeTimeout: writeTimeout, backlogFreed: make(chan struct{}), } } @@ -41,11 +36,6 @@ type yamuxConn struct { luConn *connutil.LastUsageConn addr string *yamux.Session - // writeTimeout is the configured WriteTimeoutSec, which yamux uses as - // both ConnectionWriteTimeout and StreamCloseTimeout: a stream close is - // bounded by it, and the peer's cleanup owner derives its stall - // threshold from it (see WriteTimeout) - writeTimeout time.Duration backlogMu sync.Mutex // abandonedOpens counts Open helpers still running after their caller @@ -155,11 +145,6 @@ func (y *yamuxConn) abandoned() int { return y.abandonedOpens } -// WriteTimeout implements transport.WriteTimeouter -func (y *yamuxConn) WriteTimeout() time.Duration { - return y.writeTimeout -} - func (y *yamuxConn) LastUsage() time.Time { return y.luConn.LastUsage() } diff --git a/net/transport/yamux/conn_test.go b/net/transport/yamux/conn_test.go index b9578dbde..55cd71c60 100644 --- a/net/transport/yamux/conn_test.go +++ b/net/transport/yamux/conn_test.go @@ -30,8 +30,8 @@ func newSessionPair(t *testing.T) (*yamuxConn, *yamux.Session) { _ = client.Close() _ = server.Close() }) - mc := newMultiConn(context.Background(), connutil.NewLastUsageConn(cc), "pipe", client, time.Second) - return mc, server + mc := NewMultiConn(context.Background(), connutil.NewLastUsageConn(cc), "pipe", client) + return mc.(*yamuxConn), server } func TestYamuxConn_OpenHonoursContext(t *testing.T) { diff --git a/net/transport/yamux/yamux.go b/net/transport/yamux/yamux.go index 28be3adb4..ef921ad7a 100644 --- a/net/transport/yamux/yamux.go +++ b/net/transport/yamux/yamux.go @@ -132,7 +132,7 @@ func (y *yamuxTransport) Dial(ctx context.Context, addr string) (mc transport.Mu if err != nil { return } - mc = newMultiConn(cctx, luc, addr, sess, time.Duration(y.conf.WriteTimeoutSec)*time.Second) + mc = NewMultiConn(cctx, luc, addr, sess) return } @@ -179,7 +179,7 @@ func (y *yamuxTransport) accept(conn net.Conn) { log.Info("incoming connection yamux session error", zap.Error(err), zap.String("remoteAddr", conn.RemoteAddr().String())) return } - mc := newMultiConn(cctx, luc, conn.RemoteAddr().String(), sess, time.Duration(y.conf.WriteTimeoutSec)*time.Second) + mc := NewMultiConn(cctx, luc, conn.RemoteAddr().String(), sess) if err = y.accepter.Accept(mc); err != nil { log.Info("connection accept error", zap.Error(err), zap.String("remoteAddr", conn.RemoteAddr().String())) } From e8a47e906e9570620af6bb4d99cc7d7ebd42b330 Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 15:28:19 +0200 Subject: [PATCH 3/7] net/pool, net/peer: address #802 review pool: - AddPeer returns ErrClosed during shutdown, bounds retries, waits for a mid-close entry (ocache Peeker.WaitClosing) and never waits on a flushed pair. - Get restarts from incoming after evicting a dead outgoing peer, so a live incoming connection is used instead of a second dial. - Flush reports Closed for the old pair's peers before it returns; watchers skip peers Flush already reported (once per instance). - Old caches close concurrently; Peek counts no metrics, hits counted once. peer: - Replace the cleanup owner with closeAsync + an in-flight counter. - One throttling deadline per AcquireDrpcConn call (no waiter starvation). - Release and failed-handshake closes count toward the open limiter; gc closes don't. - OutgoingProtoHandshake closes synchronously again; only the closer variant hands the close off. transport: quic and yamux share transport.NewConnClosedError. --- app/ocache/ocache.go | 24 +- app/ocache/ocache_test.go | 6 +- net/peer/cleanup.go | 123 ----- .../{cleanup_test.go => closeasync_test.go} | 501 ++++++++++-------- net/peer/limiter.go | 13 +- net/peer/peer.go | 124 +++-- net/peerobserver/peerobserver.go | 9 +- net/pool/pool.go | 319 +++++++---- net/pool/pool_flush_test.go | 478 +++++++++++++++++ net/pool/pool_test.go | 6 +- net/pool/poolservice.go | 7 +- net/secureservice/handshake/proto.go | 70 +-- net/secureservice/handshake/proto_test.go | 114 +++- net/transport/quic/conn.go | 26 +- net/transport/transport.go | 25 + net/transport/transport_test.go | 16 + net/transport/yamux/conn.go | 24 +- 17 files changed, 1314 insertions(+), 571 deletions(-) delete mode 100644 net/peer/cleanup.go rename net/peer/{cleanup_test.go => closeasync_test.go} (61%) diff --git a/app/ocache/ocache.go b/app/ocache/ocache.go index 3e8edb3a1..f29bb91d6 100644 --- a/app/ocache/ocache.go +++ b/app/ocache/ocache.go @@ -256,10 +256,16 @@ func (c *oCache) Pick(ctx context.Context, id string) (value Object, err error) type Peeker interface { // Peek returns the value for id only if it is loaded and not being // closed, without loading, waiting or allocating: the hot path for - // callers that handle a miss themselves. A hit counts as a cache hit and, - // with touch, refreshes the GC deadline like Get; a miss counts nothing - // (ok=false also for a loading entry, a closing one or a closed cache). + // callers that handle a miss themselves. With touch a hit refreshes the + // GC deadline like Get. ok=false also for a loading entry, a closing one + // or a closed cache. Peek counts no metrics: the caller, which decides + // whether the result is used, accounts for it. Peek(id string, touch bool) (value Object, ok bool) + // WaitClosing blocks while the entry for id is being closed, bounded by + // ctx, and returns at once when there is no such entry or it is not + // closing. It is the wait a caller needs before it can add a replacement + // for a value whose removal is still running. + WaitClosing(ctx context.Context, id string) error } func (c *oCache) Peek(id string, touch bool) (value Object, ok bool) { @@ -288,10 +294,20 @@ func (c *oCache) Peek(id string, touch bool) (value Object, ok bool) { } value = e.value c.mu.Unlock() - c.metricsGet(true) return value, true } +func (c *oCache) WaitClosing(ctx context.Context, id string) error { + c.mu.Lock() + e, ok := c.data[id] + c.mu.Unlock() + if !ok { + return nil + } + _, err := e.waitClose(ctx, id) + return err +} + // ctx is the cancellable load context Get created together with the entry. func (c *oCache) load(ctx context.Context, id string, e *entry) { defer func() { diff --git a/app/ocache/ocache_test.go b/app/ocache/ocache_test.go index bfa7e2f0a..725d770c4 100644 --- a/app/ocache/ocache_test.go +++ b/app/ocache/ocache_test.go @@ -1404,7 +1404,7 @@ func TestOCache_ForEachAfterClose(t *testing.T) { } func TestOCache_Peek(t *testing.T) { - t.Run("hit touches and counts, miss counts nothing", func(t *testing.T) { + t.Run("hit touches, nothing is counted", func(t *testing.T) { reg := prometheus.NewRegistry() obj := NewTestObject("a", true, nil) c := New(func(ctx context.Context, id string) (Object, error) { @@ -1437,9 +1437,9 @@ func TestOCache_Peek(t *testing.T) { values[mf.GetName()] = m.GetCounter().GetValue() } } - // one miss from Get's load, two hits from the two peeks + // one miss from Get's load; the peeks count nothing require.Equal(t, float64(1), values["peek_test_miss"]) - require.Equal(t, float64(2), values["peek_test_hit"]) + require.Equal(t, float64(0), values["peek_test_hit"]) }) t.Run("loading, closing and closed are misses", func(t *testing.T) { loading := make(chan struct{}) diff --git a/net/peer/cleanup.go b/net/peer/cleanup.go deleted file mode 100644 index 769646f6e..000000000 --- a/net/peer/cleanup.go +++ /dev/null @@ -1,123 +0,0 @@ -package peer - -import ( - "io" - "sync" - "sync/atomic" - "time" - - "go.uber.org/zap" -) - -const ( - // cleanupMaxWorkers bounds the sub-connection closes a peer runs at once - cleanupMaxWorkers = 64 - // cleanupSlowClose is how long a single close may run before it is - // logged as hung; nothing else happens to it - cleanupSlowClose = time.Minute -) - -// cleanupOwner closes a peer's sub connections off the callers' path. Closing -// a drpc conn waits for its reader, stream manager and transport, and a yamux -// stream close sends a FIN under a write timeout, so on a stalled connection a -// synchronous close turns a caller's expired deadline into a long hang. -// -// Workers are started on demand, up to cleanupMaxWorkers, and exit once there -// is nothing left to close, so an idle peer costs no goroutines. close never -// blocks and never drops a close; closes beyond the workers wait in a pending -// list. -// -// The pending list is not capped, and nothing escalates on a slow close (one -// running longer than cleanupSlowClose is only logged). The backlog is -// bounded in practice, not by construction: -// - every close comes from a sub conn this peer opened, and the peer's open -// limiter counts closes in flight (inFlight): past its threshold each -// new open waits 100ms per extra conn, so opens settle at about 10/s per -// peer while closes fall behind, and the list grows ever more slowly; -// - each close is bounded by the transport: yamux stream closes by -// StreamCloseTimeout and ConnectionWriteTimeout (both WriteTimeoutSec, -// never 0), while QUIC, iroh and webtransport closes do not block; -// - callers give up on their own deadlines rather than opening forever. -// -// Keepalive is not a bound: it can be disabled. -type cleanupOwner struct { - peerId string - - mu sync.Mutex - pending []io.Closer - workers int - - // inflight counts closes handed over and not yet finished - inflight atomic.Int32 - - // slowClose and onSlowClose are fields for tests - slowClose time.Duration - onSlowClose func(cl io.Closer) -} - -func newCleanupOwner(peerId string) *cleanupOwner { - c := &cleanupOwner{peerId: peerId, slowClose: cleanupSlowClose} - c.onSlowClose = func(io.Closer) { - log.Warn("sub connection close is taking too long", zap.String("peerId", c.peerId), zap.Duration("after", c.slowClose)) - } - return c -} - -// close hands cl over to be closed in the background. It never blocks. -func (c *cleanupOwner) close(cl io.Closer) { - if c == nil { - // a peer built without NewPeer - go func() { _ = cl.Close() }() - return - } - c.inflight.Add(1) - c.mu.Lock() - if c.workers < cleanupMaxWorkers { - c.workers++ - c.mu.Unlock() - go c.work(cl) - return - } - c.pending = append(c.pending, cl) - c.mu.Unlock() -} - -func (c *cleanupOwner) work(cl io.Closer) { - for { - c.closeOne(cl) - c.inflight.Add(-1) - c.mu.Lock() - if len(c.pending) == 0 { - c.workers-- - c.pending = nil - c.mu.Unlock() - return - } - cl = c.pending[0] - c.pending[0] = nil - c.pending = c.pending[1:] - c.mu.Unlock() - } -} - -// closeOne closes cl, logging once if the close outlives slowClose -func (c *cleanupOwner) closeOne(cl io.Closer) { - timer := time.AfterFunc(c.slowClose, func() { c.onSlowClose(cl) }) - _ = cl.Close() - timer.Stop() -} - -// inFlight returns the number of closes handed over and not yet finished -func (c *cleanupOwner) inFlight() int { - if c == nil { - return 0 - } - return int(c.inflight.Load()) -} - -// stats returns the number of running and pending closes -func (c *cleanupOwner) stats() (running, pending int) { - c.mu.Lock() - defer c.mu.Unlock() - return c.workers, len(c.pending) -} diff --git a/net/peer/cleanup_test.go b/net/peer/closeasync_test.go similarity index 61% rename from net/peer/cleanup_test.go rename to net/peer/closeasync_test.go index 8a26e7907..bc4dfcb95 100644 --- a/net/peer/cleanup_test.go +++ b/net/peer/closeasync_test.go @@ -69,183 +69,92 @@ func (c *blockingCloseConn) Close() error { // quickCloser models a healthy sub conn close: about one round trip type quickCloser struct { - d time.Duration - running *atomic.Int32 - maxSeen *atomic.Int32 - closed chan struct{} + d time.Duration + closed chan struct{} } func (q *quickCloser) Close() error { - n := q.running.Add(1) - for { - m := q.maxSeen.Load() - if n <= m || q.maxSeen.CompareAndSwap(m, n) { - break - } - } time.Sleep(q.d) - q.running.Add(-1) close(q.closed) return nil } -func TestCleanupOwner(t *testing.T) { - t.Run("burst never blocks the caller and never drops a close", func(t *testing.T) { - c := newCleanupOwner("p1") +// sendWake hands v to a waiter blocked in AcquireDrpcConn, failing the test +// instead of hanging it when nobody is waiting +func sendWake(t *testing.T, fx *fixture, v drpc.Conn) { + select { + case fx.subConnRelease <- v: + case <-time.After(5 * time.Second): + t.Fatal("no waiter took the wake-up") + } +} - const burst = 300 - var running, maxSeen atomic.Int32 - closers := make([]*quickCloser, burst) - var wg sync.WaitGroup - for i := range closers { - closers[i] = &quickCloser{d: 2 * time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} - } - start := time.Now() - for _, cl := range closers { - wg.Add(1) - go func() { - defer wg.Done() - c.close(cl) - }() - } - wg.Wait() - assert.Less(t, time.Since(start), 500*time.Millisecond, "close must never block") - for _, cl := range closers { - select { - case <-cl.closed: - case <-time.After(5 * time.Second): - t.Fatal("a close was dropped") - } - } - assert.LessOrEqual(t, int(maxSeen.Load()), cleanupMaxWorkers) - // workers exit once idle: an idle peer costs no goroutines - require.Eventually(t, func() bool { - r, p := c.stats() - return r == 0 && p == 0 - }, time.Second, time.Millisecond) - }) - t.Run("saturated workers queue and drain", func(t *testing.T) { - c := newCleanupOwner("p1") - release := make(chan struct{}) - var releaseOnce sync.Once - doRelease := func() { releaseOnce.Do(func() { close(release) }) } - defer doRelease() - - var closers []*blockingCloser - for i := 0; i < cleanupMaxWorkers+5; i++ { - cl := newBlockingCloser(release) - closers = append(closers, cl) - start := time.Now() - c.close(cl) - require.Less(t, time.Since(start), 100*time.Millisecond, "close must never block") - } - r, p := c.stats() - assert.Equal(t, cleanupMaxWorkers, r) - assert.Equal(t, 5, p) - assert.Equal(t, cleanupMaxWorkers+5, c.inFlight()) +// waitWaiter waits until an AcquireDrpcConn call sits in the throttle select +func waitWaiter(t *testing.T, fx *fixture) { + require.Eventually(t, func() bool { return fx.openingWaitCount.Load() == 1 }, 5*time.Second, time.Millisecond, + "the caller is not throttled") +} - doRelease() - for _, cl := range closers { - select { - case <-cl.closed: - case <-time.After(time.Second): - t.Fatal("a queued close was dropped") - } - } - require.Eventually(t, func() bool { - r, p := c.stats() - return r == 0 && p == 0 && c.inFlight() == 0 - }, time.Second, time.Millisecond) - }) - t.Run("a hung close blocks neither the caller nor other closes", func(t *testing.T) { - c := newCleanupOwner("p1") +func TestPeer_CloseAsync(t *testing.T) { + t.Run("never blocks the caller, never drops a close", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() hang := make(chan struct{}) defer close(hang) hung := newBlockingCloser(hang) - start := time.Now() - c.close(hung) - require.Less(t, time.Since(start), 100*time.Millisecond) + returnsWithin(t, 100*time.Millisecond, "close must never block", func() { fx.closeAsync(hung, true) }) - var running, maxSeen atomic.Int32 + // a hung close holds up no other close var closers []*quickCloser - for i := 0; i < 200; i++ { - cl := &quickCloser{d: time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} - closers = append(closers, cl) - c.close(cl) - } + returnsWithin(t, time.Second, "close must never block", func() { + for i := 0; i < 300; i++ { + cl := &quickCloser{d: time.Millisecond, closed: make(chan struct{})} + closers = append(closers, cl) + fx.closeAsync(cl, true) + } + }) for _, cl := range closers { select { case <-cl.closed: case <-time.After(5 * time.Second): - t.Fatal("a close was held up by the hung one") + t.Fatal("a close was dropped or held up") } } - // only the hung close is left, on its own worker - require.Eventually(t, func() bool { - r, p := c.stats() - return r == 1 && p == 0 && c.inFlight() == 1 - }, time.Second, time.Millisecond) + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 1 }, time.Second, time.Millisecond, + "only the hung close is still counted") }) - t.Run("sustained close rate", func(t *testing.T) { - c := newCleanupOwner("p1") - - // far more closes than workers, arriving faster than they finish - const total = 3000 - var running, maxSeen atomic.Int32 - closers := make([]*quickCloser, total) - for i := range closers { - closers[i] = &quickCloser{d: 5 * time.Millisecond, running: &running, maxSeen: &maxSeen, closed: make(chan struct{})} - } - var sawInFlight int - for i, cl := range closers { - c.close(cl) - if i%100 == 0 { - time.Sleep(time.Millisecond) - sawInFlight = max(sawInFlight, c.inFlight()) - } - } - for _, cl := range closers { - select { - case <-cl.closed: - case <-time.After(10 * time.Second): - t.Fatal("a close was dropped") - } - } - assert.LessOrEqual(t, int(maxSeen.Load()), cleanupMaxWorkers) - // closes in flight are visible to the peer's open limiter - assert.Greater(t, sawInFlight, cleanupMaxWorkers) - require.Eventually(t, func() bool { return c.inFlight() == 0 }, time.Second, time.Millisecond) + t.Run("only release closes are counted", func(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + counted, uncounted := newBlockingCloser(release), newBlockingCloser(release) + fx.closeAsync(counted, true) + fx.closeAsync(uncounted, false) + assert.Equal(t, int32(1), fx.churnClosing.Load()) + close(release) + <-counted.closed + <-uncounted.closed + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, time.Second, time.Millisecond) }) t.Run("slow close is logged once", func(t *testing.T) { - c := newCleanupOwner("p1") - c.slowClose = 20 * time.Millisecond + fx := newFixture(t, "p1") + defer fx.finish() var logged atomic.Int32 - c.onSlowClose = func(io.Closer) { logged.Add(1) } + fx.slowClose = 20 * time.Millisecond + fx.onSlowClose = func() { logged.Add(1) } release := make(chan struct{}) cl := newBlockingCloser(release) - c.close(cl) + fx.closeAsync(cl, false) require.Eventually(t, func() bool { return logged.Load() == 1 }, time.Second, time.Millisecond) require.Never(t, func() bool { return logged.Load() > 1 }, 100*time.Millisecond, 10*time.Millisecond) close(release) <-cl.closed // a quick close is not logged quick := newBlockingCloser(release) - c.close(quick) + fx.closeAsync(quick, false) <-quick.closed require.Never(t, func() bool { return logged.Load() > 1 }, 50*time.Millisecond, 10*time.Millisecond) }) - t.Run("nil owner", func(t *testing.T) { - var c *cleanupOwner - release := make(chan struct{}) - cl := newBlockingCloser(release) - returnsWithin(t, 100*time.Millisecond, "close must never block", func() { c.close(cl) }) - close(release) - select { - case <-cl.closed: - case <-time.After(time.Second): - t.Fatal("close was dropped") - } - }) } func TestPeer_HandshakeFailureCloseDoesNotBlock(t *testing.T) { @@ -369,8 +278,7 @@ func TestPeer_RPCDeadlineWithBlockedClose(t *testing.T) { assert.Empty(t, fx.active, "repetition %d", i) fx.mu.Unlock() // at most one blocked close per repetition - running, _ := fx.cleanup.stats() - require.LessOrEqual(t, running, i+1) + require.LessOrEqual(t, int(fx.churnClosing.Load()), i+1) } connsMu.Lock() require.Len(t, conns, repetitions, "each repetition opens a fresh sub conn") @@ -387,10 +295,8 @@ func TestPeer_RPCDeadlineWithBlockedClose(t *testing.T) { } } connsMu.Unlock() - require.Eventually(t, func() bool { - running, pending := fx.cleanup.stats() - return running == 0 && pending == 0 - }, 5*time.Second, 10*time.Millisecond, "cleanup workers must exit once idle") + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, + 5*time.Second, 10*time.Millisecond, "background closes must finish") } // closedConn is a released sub conn that is already closed @@ -434,44 +340,37 @@ func TestPeer_ReleaseClosedAndCancelled(t *testing.T) { assert.Empty(t, fx.inactive, "a closed or cancelled conn is never reused") assert.Empty(t, fx.active) fx.mu.Unlock() - running, pending := fx.cleanup.stats() - assert.Zero(t, running+pending, "an already closed conn is not queued for cleanup") + assert.Zero(t, fx.churnClosing.Load(), "an already closed conn is not closed again") assert.Zero(t, conn.closes.Load()) } } // TestPeer_ReleaseAfterGCDoesNotReuse: a conn gc took out of active is never -// handed out again, even while its close is still pending in the owner +// handed out again, even while its close is still pending func TestPeer_ReleaseAfterGCDoesNotReuse(t *testing.T) { fx := newFixture(t, "p1") defer fx.finish() + fx.mc.EXPECT().Addr().Return("").AnyTimes() release := make(chan struct{}) defer close(release) + a, b := net.Pipe() + defer a.Close() + defer b.Close() + // a conn that would be reusable (unblocked), whose close blocks and does + // not report Closed until it returns + pc := newPendingConn(release) + pc.unblocked = make(chan struct{}) + close(pc.unblocked) + sc := &subConn{ConnUnblocked: pc, LastUsageConn: connutil.NewLastUsageConn(a)} + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() - in, out := net.Pipe() - defer out.Close() - go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() - conn := newBlockingCloseConn(in, release) - fx.mc.EXPECT().Open(gomock.Any()).Return(conn, nil) - fx.mc.EXPECT().Addr().Return("").AnyTimes() - - dc, err := fx.AcquireDrpcConn(ctx) - require.NoError(t, err) - // keep every cleanup worker busy, so the doomed conn's close stays - // pending and never starts (drpc signals Closed as soon as it does) - for i := 0; i < cleanupMaxWorkers; i++ { - fx.cleanup.close(newBlockingCloser(release)) - } - time.Sleep(20 * time.Millisecond) - // the idle active conn is doomed; its close waits in the owner + // the idle active conn is doomed; its close is pending fx.gc(time.Millisecond) - select { - case <-dc.Closed(): - t.Fatal("close was expected to be still pending") - default: - } + require.True(t, sc.doomed.Load()) - fx.ReleaseDrpcConn(ctx, dc) + fx.ReleaseDrpcConn(ctx, sc) fx.mu.Lock() assert.Empty(t, fx.inactive, "a doomed conn must not be reused") assert.Empty(t, fx.active) @@ -517,7 +416,6 @@ func TestPeer_ReleaseDoomedDuringCheck(t *testing.T) { fx.mu.Lock() fx.active[sc] = struct{}{} fx.mu.Unlock() - time.Sleep(10 * time.Millisecond) fx.ReleaseDrpcConn(ctx, sc) require.True(t, sc.doomed.Load(), "gc doomed it during the release") @@ -526,28 +424,6 @@ func TestPeer_ReleaseDoomedDuringCheck(t *testing.T) { fx.mu.Unlock() } -func TestPeer_AcquireSkipsDoomedConns(t *testing.T) { - fx := newFixture(t, "p1") - defer fx.finish() - a, b := net.Pipe() - defer a.Close() - defer b.Close() - rc := &racyConn{never: make(chan struct{}), ready: make(chan struct{})} - doomed := &subConn{ConnUnblocked: rc, LastUsageConn: connutil.NewLastUsageConn(a)} - doomed.doomed.Store(true) - fx.mu.Lock() - fx.inactive = append(fx.inactive, doomed) - fx.mu.Unlock() - - in, out := net.Pipe() - defer out.Close() - go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() - fx.mc.EXPECT().Open(gomock.Any()).Return(in, nil) - dc, err := fx.AcquireDrpcConn(ctx) - require.NoError(t, err) - assert.NotEqual(t, drpc.Conn(doomed), dc) -} - // TestPeer_NilWakeKeepsThrottling: a waiter woken because a released conn // was closed must not open at once while closes are still in flight func TestPeer_NilWakeKeepsThrottling(t *testing.T) { @@ -558,7 +434,7 @@ func TestPeer_NilWakeKeepsThrottling(t *testing.T) { // closes stuck in flight on a slow transport, enough for the limiter // to hold the next open for seconds for i := 0; i < fx.limiter.startThreshold+20; i++ { - fx.cleanup.close(newBlockingCloser(release)) + fx.closeAsync(newBlockingCloser(release), true) } var opens atomic.Int32 @@ -575,8 +451,8 @@ func TestPeer_NilWakeKeepsThrottling(t *testing.T) { done <- err }() // the waiter is throttled; a release wakes it with a nil conn - time.Sleep(50 * time.Millisecond) - fx.subConnRelease <- nil + waitWaiter(t, fx) + sendWake(t, fx, nil) require.Never(t, func() bool { return opens.Load() > 0 }, 300*time.Millisecond, 10*time.Millisecond, "a nil wake must not bypass the limiter") cancel() @@ -598,12 +474,12 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { var hangOnce sync.Once unhang := func() { hangOnce.Do(func() { close(hang) }) } defer unhang() - // one past the limiter threshold: the next open waits slowDownStep + // ten past the limiter threshold: the next open waits 10 slowDownSteps var hung []*blockingCloser - for i := 0; i <= fx.limiter.startThreshold; i++ { + for i := 0; i < fx.limiter.startThreshold+10; i++ { cl := newBlockingCloser(hang) hung = append(hung, cl) - fx.cleanup.close(cl) + fx.closeAsync(cl, true) } var opens atomic.Int32 @@ -615,7 +491,7 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { return in, nil }).Times(1) - actx, cancel := context.WithTimeout(ctx, fx.limiter.slowDownStep/2) + actx, cancel := context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) _, err := fx.AcquireDrpcConn(actx) cancel() require.ErrorIs(t, err, context.DeadlineExceeded, "throttled while the closes are in flight") @@ -625,8 +501,8 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { for _, cl := range hung { <-cl.closed } - require.Eventually(t, func() bool { return fx.cleanup.inFlight() == 0 }, time.Second, time.Millisecond) - actx, cancel = context.WithTimeout(ctx, fx.limiter.slowDownStep/2) + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, time.Second, time.Millisecond) + actx, cancel = context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) defer cancel() _, err = fx.AcquireDrpcConn(actx) require.NoError(t, err, "no throttling once the closes are done") @@ -637,6 +513,7 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { // and whose Close blocks until released type pendingConn struct { closedCh chan struct{} + unblocked chan struct{} release chan struct{} closeCalls atomic.Int32 once sync.Once @@ -653,7 +530,7 @@ func (c *pendingConn) Close() error { return nil } func (c *pendingConn) Closed() <-chan struct{} { return c.closedCh } -func (c *pendingConn) Unblocked() <-chan struct{} { return nil } +func (c *pendingConn) Unblocked() <-chan struct{} { return c.unblocked } func (c *pendingConn) NewStream(context.Context, string, drpc.Encoding) (drpc.Stream, error) { return nil, io.EOF } @@ -703,7 +580,7 @@ func TestPeer_ReleaseClosesInBackground(t *testing.T) { returnsWithin(t, 100*time.Millisecond, "the blocked close must not run on the caller", func() { fx.ReleaseDrpcConn(cctx, sc) }) - require.Equal(t, 1, fx.cleanup.inFlight()) + require.Equal(t, int32(1), fx.churnClosing.Load()) fx.mu.Lock() assert.Empty(t, fx.inactive) assert.Empty(t, fx.active) @@ -711,7 +588,7 @@ func TestPeer_ReleaseClosesInBackground(t *testing.T) { close(release) waitClosed(t, pc) - require.Eventually(t, func() bool { return fx.cleanup.inFlight() == 0 }, time.Second, time.Millisecond) + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, time.Second, time.Millisecond) }) t.Run("never unblocked", func(t *testing.T) { fx := newFixture(t, "p1") @@ -723,7 +600,7 @@ func TestPeer_ReleaseClosesInBackground(t *testing.T) { returnsWithin(t, 300*time.Millisecond, "the blocked close must not run on the caller", func() { fx.ReleaseDrpcConn(ctx, sc) }) - require.Equal(t, 1, fx.cleanup.inFlight()) + require.Equal(t, int32(1), fx.churnClosing.Load()) fx.mu.Lock() assert.Empty(t, fx.inactive, "an unfinished conn is not reused") fx.mu.Unlock() @@ -756,7 +633,7 @@ func TestPeer_GCClosesInBackground(t *testing.T) { fx.mu.Lock() assert.Empty(t, fx.inactive) fx.mu.Unlock() - require.Equal(t, 1, fx.cleanup.inFlight()) + require.Zero(t, fx.churnClosing.Load(), "gc closes are not counted by the limiter") close(release) waitClosed(t, pc) }) @@ -774,8 +651,212 @@ func TestPeer_GCClosesInBackground(t *testing.T) { fx.gc(time.Millisecond) }) require.True(t, sc.doomed.Load()) - require.Equal(t, 1, fx.cleanup.inFlight()) + require.Zero(t, fx.churnClosing.Load(), "gc closes are not counted by the limiter") close(release) waitClosed(t, pc) }) } + +// openPipe makes Open return a fresh handshaking conn each time +func openPipe(t *testing.T, fx *fixture) { + fx.mc.EXPECT().Open(gomock.Any()).DoAndReturn(func(context.Context) (net.Conn, error) { + in, out := net.Pipe() + t.Cleanup(func() { _ = out.Close() }) + go func() { _, _ = handshake.IncomingProtoHandshake(ctx, out, defaultProtoChecker) }() + return in, nil + }).AnyTimes() +} + +// TestPeer_WakeUpsDoNotStarveWaiters: waiters woken again and again by +// releases keep their first throttling deadline instead of restarting it +func TestPeer_WakeUpsDoNotStarveWaiters(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + openPipe(t, fx) + + // released conns closing slowly, about 10 in flight: a 1s throttle + stop := make(chan struct{}) + var churn sync.WaitGroup + churn.Add(1) + go func() { + defer churn.Done() + ticker := time.NewTicker(100 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-stop: + return + case <-ticker.C: + } + release := make(chan struct{}) + fx.closeAsync(newBlockingCloser(release), true) + time.AfterFunc(2*time.Second, func() { close(release) }) + // the release wakes a waiter with a nil conn + select { + case fx.subConnRelease <- nil: + default: + } + } + }() + defer func() { + close(stop) + churn.Wait() + }() + for i := 0; i < fx.limiter.startThreshold+10; i++ { + release := make(chan struct{}) + fx.closeAsync(newBlockingCloser(release), true) + time.AfterFunc(2*time.Second, func() { close(release) }) + } + + const waiters = 5 + results := make(chan error, waiters) + for i := 0; i < waiters; i++ { + go func() { + actx, cancel := context.WithTimeout(ctx, 4*time.Second) + defer cancel() + _, err := fx.AcquireDrpcConn(actx) + results <- err + }() + } + for i := 0; i < waiters; i++ { + select { + case err := <-results: + require.NoError(t, err, "a waiter starved") + case <-time.After(10 * time.Second): + t.Fatal("waiter did not return") + } + } +} + +// TestPeer_GCDoesNotThrottleOpens: closes started by gc are not counted by +// the open limiter, so a gc pass does not delay opens on a healthy peer +func TestPeer_GCDoesNotThrottleOpens(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + openPipe(t, fx) + release := make(chan struct{}) + defer close(release) + // expired inactive conns whose closes hang + for i := 0; i < fx.limiter.startThreshold+10; i++ { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + fx.mu.Lock() + fx.inactive = append(fx.inactive, &subConn{ConnUnblocked: newPendingConn(release), LastUsageConn: connutil.NewLastUsageConn(a)}) + fx.mu.Unlock() + } + fx.gc(time.Millisecond) + fx.mu.Lock() + require.Empty(t, fx.inactive) + fx.mu.Unlock() + + // counted, these closes would hold the open for 10 slowDownSteps + actx, cancel := context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) + defer cancel() + _, err := fx.AcquireDrpcConn(actx) + require.NoError(t, err, "gc closes must not throttle opens") +} + +// TestPeer_FailedHandshakeChurnIsThrottled: closes of sub conns that failed +// the handshake count towards the open limiter, so a peer whose handshakes +// keep failing on a stalled transport cannot spin up opens (and hung close +// goroutines) at the callers' rate +func TestPeer_FailedHandshakeChurnIsThrottled(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + defer close(release) + var opens atomic.Int32 + fx.mc.EXPECT().Open(gomock.Any()).DoAndReturn(func(context.Context) (net.Conn, error) { + opens.Add(1) + in, out := net.Pipe() + go func() { + // the remote declines the protocol, then goes away + _, _ = handshake.IncomingProtoHandshake(ctx, out, handshake.ProtoChecker{AllowedProtoTypes: []handshakeproto.ProtoType{100}}) + _ = out.Close() + }() + return newBlockingCloseConn(in, release), nil + }).AnyTimes() + + // callers retry as fast as they can for a second + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + actx, cancel := context.WithDeadline(ctx, deadline) + _, _ = fx.AcquireDrpcConn(actx) + cancel() + } + // unthrottled this is hundreds; throttled, past the 10-conn threshold + // each open waits 100ms more than the one before + assert.Less(t, int(opens.Load()), 30) + assert.Equal(t, opens.Load(), fx.churnClosing.Load(), "every failed open's close is counted while it hangs") +} + +// TestPeer_WakeWithDoomedConnIsNotHandedOut: gc can doom a released conn +// after ReleaseDrpcConn checked it and before the hand-off to a waiter; the +// waiter must not take it +func TestPeer_WakeWithDoomedConnIsNotHandedOut(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + defer close(release) + for i := 0; i < fx.limiter.startThreshold+20; i++ { + fx.closeAsync(newBlockingCloser(release), true) + } + fx.mc.EXPECT().Open(gomock.Any()).Return(nil, io.EOF).AnyTimes() + + actx, cancel := context.WithCancel(ctx) + defer cancel() + type result struct { + conn drpc.Conn + err error + } + done := make(chan result, 1) + go func() { + conn, err := fx.AcquireDrpcConn(actx) + done <- result{conn, err} + }() + waitWaiter(t, fx) + doomed := &subConn{ConnUnblocked: newPendingConn(release)} + doomed.doomed.Store(true) + sendWake(t, fx, doomed) + // the waiter starts over and keeps waiting rather than using it + waitWaiter(t, fx) + cancel() + select { + case res := <-done: + require.ErrorIs(t, res.err, context.Canceled) + require.Nil(t, res.conn) + case <-time.After(5 * time.Second): + t.Fatal("acquire did not return on ctx") + } +} + +// TestPeer_ForeignConnReleaseCloseIsCounted: a released conn that is not a +// sub conn of this peer is closed in the background and counted +func TestPeer_ForeignConnReleaseCloseIsCounted(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + release := make(chan struct{}) + foreign := &foreignConn{pendingConn: newPendingConn(release)} + returnsWithin(t, 100*time.Millisecond, "the blocked close must not run on the caller", func() { + fx.ReleaseDrpcConn(ctx, foreign) + }) + require.Equal(t, int32(1), fx.churnClosing.Load()) + close(release) + waitClosed(t, foreign.pendingConn) + require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, time.Second, time.Millisecond) +} + +// foreignConn is a drpc.Conn without Unblocked +type foreignConn struct { + pendingConn *pendingConn +} + +func (c *foreignConn) Close() error { return c.pendingConn.Close() } +func (c *foreignConn) Closed() <-chan struct{} { return c.pendingConn.Closed() } +func (c *foreignConn) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) (drpc.Stream, error) { + return c.pendingConn.NewStream(ctx, rpc, enc) +} +func (c *foreignConn) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, in, out drpc.Message) error { + return c.pendingConn.Invoke(ctx, rpc, enc, in, out) +} diff --git a/net/peer/limiter.go b/net/peer/limiter.go index 4c8e15894..58ab94d3b 100644 --- a/net/peer/limiter.go +++ b/net/peer/limiter.go @@ -10,9 +10,16 @@ type limiter struct { } func (l limiter) wait(count int) <-chan time.Time { - if count > l.startThreshold { - wait := l.slowDownStep * time.Duration(count-l.startThreshold) - return time.After(wait) + if d := l.delay(count); d > 0 { + return time.After(d) } return nil } + +// delay is how long to hold off a new conn when count are already open +func (l limiter) delay(count int) time.Duration { + if count > l.startThreshold { + return l.slowDownStep * time.Duration(count-l.startThreshold) + } + return 0 +} diff --git a/net/peer/peer.go b/net/peer/peer.go index 6dc4aa505..969fd04a3 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -57,7 +57,10 @@ func NewPeer(mc transport.MultiConn, ctrl connCtrl) (p Peer, err error) { if pr.id, err = CtxPeerId(ctx); err != nil { return } - pr.cleanup = newCleanupOwner(pr.id) + pr.slowClose = slowCloseWarn + pr.onSlowClose = func() { + log.Warn("sub connection close is taking too long", zap.String("peerId", pr.id), zap.Duration("after", pr.slowClose)) + } go pr.acceptLoop() return pr, nil } @@ -98,9 +101,9 @@ type Peer interface { type subConn struct { encoding.ConnUnblocked *connutil.LastUsageConn - // doomed is set by gc when it takes an active conn away and hands its - // close to the cleanup owner: the holder must not return it for reuse, - // even though the close may not have landed yet + // doomed is set by gc when it takes an active conn away and closes it in + // the background: the holder must not return it for reuse, even though + // the close may not have landed yet doomed atomic.Bool } @@ -129,8 +132,14 @@ type peer struct { limiter limiter - // cleanup closes sub connections off the callers' path - cleanup *cleanupOwner + // churnClosing counts background closes of sub conns callers churned + // through (released unusable, or failed in the handshake) that have not + // finished yet; the open limiter counts them as sub conns + churnClosing atomic.Int32 + // slowClose and onSlowClose log a background close that hangs; fields + // so tests can shorten them + slowClose time.Duration + onSlowClose func() mu sync.Mutex created time.Time @@ -144,26 +153,37 @@ func (p *peer) Id() string { } func (p *peer) AcquireDrpcConn(ctx context.Context) (drpc.Conn, error) { + // one throttling deadline for the whole call: a retry may shorten the + // wait, never push it back, or steady wake-ups would starve the caller + var deadline time.Time for { - conn, retry, err := p.acquireDrpcConn(ctx) + conn, retry, err := p.acquireDrpcConn(ctx, &deadline) if !retry { return conn, err } } } -// acquireDrpcConn makes one acquisition attempt; retry means start over -// with a fresh look at the pool and a fresh limiter wait -func (p *peer) acquireDrpcConn(ctx context.Context) (conn drpc.Conn, retry bool, err error) { +// acquireDrpcConn makes one acquisition attempt; retry means start over with +// a fresh look at the pool and the limiter wait recomputed (see deadline) +func (p *peer) acquireDrpcConn(ctx context.Context, deadline *time.Time) (conn drpc.Conn, retry bool, err error) { if p.IsClosed() { return nil, false, transport.ErrConnClosed } p.mu.Lock() if len(p.inactive) == 0 { - // closes still in progress count too: they used to run on the - // releasing callers' goroutines while the conn was still active, - // and the owner must not turn into a way around the throttling - wait := p.limiter.wait(len(p.active) + int(p.openingWaitCount.Load()) + p.cleanup.inFlight()) + // released conns still closing in the background count too, so + // that closing them off the releasing callers' path does not bypass + // the throttling + var wait <-chan time.Time + if delay := p.limiter.delay(len(p.active) + int(p.openingWaitCount.Load()) + int(p.churnClosing.Load())); delay > 0 { + if until := time.Now().Add(delay); deadline.IsZero() || until.Before(*deadline) { + *deadline = until + } + timer := time.NewTimer(time.Until(*deadline)) + defer timer.Stop() + wait = timer.C + } p.openingWaitCount.Add(1) defer p.openingWaitCount.Add(-1) p.mu.Unlock() @@ -178,10 +198,9 @@ func (p *peer) acquireDrpcConn(ctx context.Context) (conn drpc.Conn, retry bool, return dconn, false, nil } // The released conn was closed, or gc doomed it on the way. - // Its close may still be in flight in the cleanup owner, so - // opening right away would bypass the throttling: start - // over, which picks up an inactive conn or recomputes the - // wait with the closes still in flight. + // Its close may still be running, so opening right away + // would bypass the throttling: start over, which picks up an + // inactive conn or waits out the (never later) deadline. return nil, true, nil case <-wait: } @@ -202,11 +221,8 @@ func (p *peer) acquireDrpcConn(ctx context.Context) (conn drpc.Conn, retry bool, return nil, true, nil default: } - if res.doomed.Load() { - // gc took it while it was being re-pooled; its close is pending - p.mu.Unlock() - return nil, true, nil - } + // never doomed: gc dooms active conns only, and ReleaseDrpcConn + // re-checks the flag under p.mu before re-pooling p.active[res] = struct{}{} p.mu.Unlock() return res, false, nil @@ -284,7 +300,7 @@ func (p *peer) checkReleased(ctx context.Context, conn drpc.Conn) (closed bool) case <-conn.Closed(): // both were ready: nothing left to close default: - p.cleanup.close(conn) + p.closeAsync(conn, true) } closed = true default: @@ -300,7 +316,7 @@ func (p *peer) checkReleased(ctx context.Context, conn drpc.Conn) (closed bool) // means the connection has some unfinished work, // e.g. not fully read stream // we cannot reuse this connection so let's close it - p.cleanup.close(conn) + p.closeAsync(conn, true) closed = true } } else { @@ -309,7 +325,7 @@ func (p *peer) checkReleased(ctx context.Context, conn drpc.Conn) (closed bool) // the caller passed a foreign conn; close it defensively instead // of crashing the process. log.Warn("released conn does not implement encoding.ConnUnblocked, closing", zap.String("peerId", p.id)) - p.cleanup.close(conn) + p.closeAsync(conn, true) closed = true } } @@ -338,8 +354,8 @@ func (p *peer) openDrpcConn(ctx context.Context) (*subConn, error) { return nil, err } lastUsageConn := connutil.NewLastUsageConn(conn) - // on any error the handshake hands the stream to the cleanup owner, once: - // on a stalled transport the close blocks, so it never runs here + // on any error the handshake hands the stream over once, to be closed in + // the background: on a stalled transport the close blocks proto, err := handshake.OutgoingProtoHandshakeWithCloser(ctx, lastUsageConn, defaultHandshakeProto, p.closeSubConn) if err != nil { return nil, err @@ -359,7 +375,53 @@ func (p *peer) openDrpcConn(ctx context.Context) (*subConn, error) { } func (p *peer) closeSubConn(conn net.Conn) { - p.cleanup.close(conn) + // counted: one per failed open, so a peer whose handshakes keep failing + // on a stalled transport opens ever more slowly + p.closeAsync(conn, true) +} + +// slowCloseWarn is how long a background close may run before it is logged. +// It is above the default transport bounds; with a yamux WriteTimeoutSec of +// a minute or more a legitimate close can reach it, which only logs. +const slowCloseWarn = time.Minute + +// closeAsync closes c off the caller's path, one goroutine per close. +// Closing a drpc conn waits for its reader, stream manager and transport, +// and a yamux stream close sends a FIN under a write timeout, so on a +// stalled connection a synchronous close turns a caller's expired deadline +// into a long hang. +// +// The goroutines are not capped but self-limited by the open limiter. Each +// close is bounded by the transport (yamux StreamCloseTimeout and +// ConnectionWriteTimeout, both WriteTimeoutSec and never 0; QUIC, iroh and +// webtransport closes do not block), and every close comes from a sub conn +// this peer opened. counted marks a close of a conn a caller churned through +// (ReleaseDrpcConn of an unusable conn, or a failed handshake): the open +// limiter counts those like sub conns, so with closes of duration T in flight +// the number settles around 10+sqrt(10*T), T in seconds, instead of growing +// with the callers' rate. Closes from gc are not counted: they say nothing about how +// fast callers churn. A close running longer than slowClose is logged once; +// nothing else happens to it. +// +// Nothing waits for these goroutines: peer.Close and pool.Close return while +// they run, and they finish promptly once the connection is closed. +func (p *peer) closeAsync(c io.Closer, counted bool) { + if counted { + p.churnClosing.Add(1) + } + go func() { + if counted { + defer p.churnClosing.Add(-1) + } + var timer *time.Timer + if p.onSlowClose != nil { + timer = time.AfterFunc(p.slowClose, p.onSlowClose) + } + _ = c.Close() + if timer != nil { + timer.Stop() + } + }() } func (p *peer) acceptLoop() { @@ -451,12 +513,12 @@ func (p *peer) TryClose(objectTTL time.Duration) (res bool, err error) { func (p *peer) gc(ttl time.Duration) (aliveCount int) { // drpc conn Close blocks until its reader unwinds, which on a stalled stream // takes until the yamux stream close timeout: collect the doomed conns and - // hand them to the cleanup owner after releasing the lock, so a stalled + // close them in the background after releasing the lock, so a stalled // peer does not hold up the GC pass of every other peer var toClose []*subConn defer func() { for _, conn := range toClose { - p.cleanup.close(conn) + p.closeAsync(conn, false) } }() p.mu.Lock() diff --git a/net/peerobserver/peerobserver.go b/net/peerobserver/peerobserver.go index 717345807..886104209 100644 --- a/net/peerobserver/peerobserver.go +++ b/net/peerobserver/peerobserver.go @@ -129,9 +129,12 @@ type Event struct { // observed, subsequent suppressed attempts are not. An inbound connection // that fails before it becomes a peer produces no event. // -// Calls arrive concurrently from dial callers, transport accept loops and -// per-connection watcher goroutines; implementations must be safe for -// concurrent use. Dial-path events (KindDialStarted, outbound KindConnected, +// Calls arrive concurrently from dial callers, transport accept loops, +// per-connection watcher goroutines and, for the KindClosed of every pooled +// connection a pool.Flush invalidates, synchronously on the goroutine that +// called Flush before it returns (so for that caller those KindClosed precede +// the KindConnected of its redials; the caller must not hold a lock the +// observer takes); implementations must be safe for concurrent use. Dial-path events (KindDialStarted, outbound KindConnected, // KindDialFailed) run inside the pool's single-flight load for that peer: an // implementation must never call the pool for the peer such an event names — // the load is still open and the call blocks until the caller's context dies. diff --git a/net/pool/pool.go b/net/pool/pool.go index 339f5d87a..ff5e028bf 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -3,6 +3,7 @@ package pool import ( "context" + "errors" "fmt" "math/rand" "sync" @@ -33,7 +34,12 @@ type Pool interface { // Pick checks if a connection with the peer exists, without dialing. // For a peer whose last dial failed it returns the cached dial error. Pick(ctx context.Context, id string) (pr peer.Peer, err error) - // Flush removes all connections from the pool + // Flush invalidates every pooled connection and every dial in flight: it + // swaps in an empty pool, reports Closed for each pooled peer before it + // returns and tears the old connections down in the background. Later + // lookups redial. Callers should coalesce the triggers (heart's recovery + // worker does, with a 6s window), since each Flush cancels the dials of + // the one before. Flush(ctx context.Context) error } @@ -57,6 +63,35 @@ type caches struct { // is cut short and retried on the current one ctx context.Context cancel context.CancelFunc + // reported holds the peers whose Closed event Flush delivered itself, so + // their watchers do not report them a second time (see Flush); written + // once, by Flush, under reportedMu + reportedMu sync.Mutex + reported map[peer.Peer]struct{} +} + +// cache returns the incoming or the outgoing cache of the pair +func (c *caches) cache(inbound bool) ocache.OCache { + if inbound { + return c.incoming + } + return c.outgoing +} + +// reportedByFlush reports whether Flush delivered this peer's Closed event +func (c *caches) reportedByFlush(pr peer.Peer) bool { + c.reportedMu.Lock() + defer c.reportedMu.Unlock() + _, ok := c.reported[pr] + return ok +} + +// fastMetrics are the counters the hit path bumps itself (nil without a +// registry): ocache.Peek counts nothing, so a Get or Pick served by fast or +// by the caches counts exactly once per cache, as a Get through the caches +// did before the fast path existed +type fastMetrics struct { + incomingHit, incomingMiss, outgoingHit prometheus.Counter } type pool struct { @@ -76,10 +111,7 @@ type pool struct { newCaches func() *caches // closeTimeout bounds the ocache close passes and Close as a whole closeTimeout time.Duration - // incomingMiss is the incoming cache's miss counter (nil without a - // registry): Get counts a miss there when the fast path finds the peer - // in outgoing, as a Get through the caches would - incomingMiss prometheus.Counter + metrics *fastMetrics statService debugstat.StatService closingCtx context.Context @@ -139,23 +171,34 @@ func (p *pool) lookup(ctx context.Context, f func(ctx context.Context, c *caches // handles those cases. Rechecking current after the read gives the same // guarantee as lookup's post-check: the peer comes from a pair that was // current after it was read. touch refreshes the GC deadline (Get) or not -// (Pick); hits are counted either way, as the caches themselves do. Get's -// incoming miss is counted too, so the series are the same as before. +// (Pick). Metrics are counted only for a peer actually returned, and as the +// caches would have counted them (a hit; for Get also the incoming miss +// before an outgoing hit), so a fall-through to lookup counts nothing twice. func (p *pool) fast(id string, touch bool) peer.Peer { c := p.current.Load() - v, ok := c.peekIncoming.Peek(id, touch) - if !ok { + v, inbound := c.peekIncoming.Peek(id, touch) + if !inbound { + var ok bool if v, ok = c.peekOutgoing.Peek(id, touch); !ok { return nil } - if touch && p.incomingMiss != nil { - p.incomingMiss.Inc() - } } - if pr, isPeer := v.(peer.Peer); isPeer && !pr.IsClosed() && p.current.Load() == c { - return pr + pr, isPeer := v.(peer.Peer) + if !isPeer || pr.IsClosed() || p.current.Load() != c { + return nil } - return nil + if m := p.metrics; m != nil { + switch { + case inbound: + m.incomingHit.Inc() + default: + if touch { + m.incomingMiss.Inc() + } + m.outgoingHit.Inc() + } + } + return pr } // discard closes pr (if not closed yet) and evicts it from source, in the @@ -178,12 +221,14 @@ func (p *pool) discard(source ocache.OCache, pr peer.Peer) <-chan struct{} { } // evictOnClose removes the peer from the cache as soon as its underlying -// connection dies, instead of waiting for the next Get or the GC to notice. -// When the whole pool is shutting down, cache.Close already evicts every peer, -// so per-peer removal is skipped. It never outlives the peer. cache is the -// instance the peer was published into: after a Flush that is no longer the -// current one, and RemoveSame on it fails fast with ErrClosed. -func (p *pool) evictOnClose(pr peer.Peer, cache ocache.OCache, inbound bool) { +// connection dies, instead of waiting for the next Get or the GC to notice, +// and reports the Closed event unless Flush already did. When the whole pool +// is shutting down, cache.Close already evicts every peer, so per-peer removal +// is skipped. It never outlives the peer. c is the pair the peer was +// published into: after a Flush that is no longer the current one, and +// RemoveSame on it fails fast with ErrClosed. +func (p *pool) evictOnClose(pr peer.Peer, c *caches, inbound bool) { + cache := c.cache(inbound) select { case <-pr.CloseChan(): case <-p.closingCtx.Done(): @@ -207,8 +252,11 @@ func (p *pool) evictOnClose(pr peer.Peer, cache ocache.OCache, inbound bool) { // redial); removing by id alone would close that live replacement. _, _ = cache.RemoveSame(p.closingCtx, pr.Id(), pr) // RemoveSame can park behind another closer; re-check so no Closed is - // delivered once pool shutdown has begun - if p.closingCtx.Err() != nil { + // delivered once pool shutdown has begun. Checked after the removal: Flush + // marks the peers it saw and reports under reportedMu, and the removal + // and Flush's snapshot are ordered by the cache lock, so a peer Flush saw + // is marked by the time this runs and one it did not see is reported here + if p.closingCtx.Err() != nil || c.reportedByFlush(pr) { return } p.observer.Notify(peerobserver.Event{ @@ -223,53 +271,67 @@ func (p *pool) Get(ctx context.Context, id string) (peer.Peer, error) { return pr, nil } return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { - // if we have incoming connection - try to reuse it - if pr, err = p.get(ctx, c.incoming, id); err != nil { + for { + // if we have incoming connection - try to reuse it + if pr, err = p.get(ctx, c.incoming, id); err == nil { + return pr, nil + } // or try to get or create outgoing - return p.get(ctx, c.outgoing, id) + if pr, err = p.get(ctx, c.outgoing, id); err != errRedial { + return pr, err + } + // a closed peer was evicted: start over from incoming, where a + // live connection may have arrived meanwhile, before dialing (a + // ctx that ended meanwhile fails the next Get at once) } - return }) } +// errRedial is get's verdict after it evicted a closed peer: look again, +// starting from the incoming cache +var errRedial = errors.New("closed peer evicted") + func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.Peer, error) { - for { - v, err := source.Get(ctx, id) - if err != nil { - return nil, err - } - pr, err := getPeer(v) - if err != nil { - return nil, err - } - if !pr.IsClosed() { - return pr, nil - } - // The entry must be gone before redialing or source.Get would return - // the same instance again: wait (bounded by ctx) for the background - // discard, so a teardown that blocks never runs on this path. - select { - case <-p.discard(source, pr): - case <-ctx.Done(): - return nil, ctx.Err() - } - // with a done ctx source.Get can return the closed value again - if err = ctx.Err(); err != nil { - return nil, err - } + v, err := source.Get(ctx, id) + if err != nil { + return nil, err + } + pr, err := getPeer(v) + if err != nil { + return nil, err + } + if !pr.IsClosed() { + return pr, nil } + // The entry must be gone before looking again or source.Get would return + // the same instance: wait (bounded by ctx) for the background discard, so + // a teardown that blocks never runs on this path. + select { + case <-p.discard(source, pr): + case <-ctx.Done(): + return nil, ctx.Err() + } + return nil, errRedial } // Flush invalidates every pooled connection: it builds a fresh cache pair, -// publishes it and closes the replaced pair in the background, so it never -// waits on a GC TryClose or on transport teardown. From the moment it returns -// no lookup hands out a pre-flush peer (see lookup), and a dial that was in -// flight is cancelled instead of being waited out. Cached dial errors go with -// the old pair, except incompatible-version verdicts (see errObject), which -// are carried over so their backoff survives (one still loading at the swap -// is not: its verdict lands in the old pair and the fresh one redials once). -// A no-op once the pool is closed; concurrent flushes are serialized, so each -// replaced pair is closed exactly once. +// publishes it, cancels the old pair (which cuts every dial in flight short, +// see lookup) and closes it in the background, so it never waits on a GC +// TryClose or on transport teardown. From the moment it returns no lookup +// hands out a pre-flush peer. The Closed event of every peer it found pooled +// is delivered on the calling goroutine before it returns (observers must not +// be blocked by a lock the caller holds), so for the caller's own later calls +// a Closed(X) precedes the Connected(X) of the redial; a concurrent Get or +// Accept can still produce its Connected(X) first. The peers' own watchers +// skip those events (see evictOnClose); a peer a GC TryClose holds at that +// moment is reported by its watcher once it closes, and once pool shutdown +// has begun the events are suppressed like the watchers' (the peers close +// anyway). Cached dial errors go with the old pair, except +// incompatible-version verdicts (see errObject), which are carried over so +// their backoff survives (one still loading at the swap is not: its verdict +// lands in the old pair and the fresh one redials once). A no-op once the +// pool is closed; concurrent flushes are serialized, so each replaced pair is +// closed exactly once. func (p *pool) Flush(ctx context.Context) error { p.swapMu.Lock() if p.closed { @@ -291,26 +353,58 @@ func (p *pool) Flush(ctx context.Context) error { // can follow its Wait p.closing.Add(1) p.swapMu.Unlock() + // Snapshot and mark under reportedMu: a watcher that evicts one of these + // peers concurrently checks the mark after its removal, which the cache + // lock orders against this snapshot, so each instance is reported once. + // The events themselves go out without any lock held. + type flushedPeer struct { + pr peer.Peer + inbound bool + } + var flushed []flushedPeer + old.reportedMu.Lock() + old.reported = map[peer.Peer]struct{}{} + for _, inbound := range []bool{true, false} { + old.cache(inbound).ForEach(func(v ocache.Object) (isContinue bool) { + if pr, ok := v.(peer.Peer); ok { + old.reported[pr] = struct{}{} + flushed = append(flushed, flushedPeer{pr: pr, inbound: inbound}) + } + return true + }) + } + old.reportedMu.Unlock() go func() { defer p.closing.Done() peers, _ := closeCaches(old) peers.Wait() }() + for _, f := range flushed { + if p.closingCtx.Err() != nil { + // a Close that began meanwhile: Closed is suppressed from here on + break + } + p.observer.Notify(peerobserver.Event{ + Kind: peerobserver.KindClosed, + PeerId: f.pr.Id(), + Inbound: f.inbound, + }) + } return nil } // closeCaches tears down a pair that is no longer current. Each loaded peer is // closed on its own goroutine first, because ocache.Close closes entries one // at a time with no ctx and one hung teardown would hold the rest back; the -// Close passes then cancel the in-flight dials (outgoing first, so a hung -// incoming peer cannot delay that) and close the peers a second time, which -// peer.Close tolerates (the pool relies on that already, see discard). A -// RemoveSame per peer would not do: once a cache is marked closed every -// RemoveSame is refused. Known gap: a peer a GC TryClose holds past -// closeTimeout closes only when TryClose returns (ocache escalates the -// decline), never if it never returns. Returns once the caches are closed; -// the WaitGroup tracks the per-peer closes still running, and err is the -// outgoing cache's close error. +// two caches are then closed concurrently (each Close cancels its in-flight +// loads and closes its peers a second time, which peer.Close tolerates; the +// pool relies on that already, see discard), so a hung peer in one cache +// never delays the other. A RemoveSame per peer would not do: once a cache +// is marked closed every RemoveSame is refused. Known gap: a peer a GC +// TryClose holds past closeTimeout closes only when TryClose returns (ocache +// escalates the decline), never if it never returns. Returns once both caches +// are closed; the WaitGroup tracks the per-peer closes still running, and err +// is the outgoing cache's close error. func closeCaches(c *caches) (peers *sync.WaitGroup, err error) { peers = &sync.WaitGroup{} for _, cache := range []ocache.OCache{c.outgoing, c.incoming} { @@ -325,8 +419,10 @@ func closeCaches(c *caches) (peers *sync.WaitGroup, err error) { return true }) } + incomingClosed := make(chan error, 1) + go func() { incomingClosed <- c.incoming.Close() }() err = c.outgoing.Close() - if e := c.incoming.Close(); e != nil { + if e := <-incomingClosed; e != nil { log.Warn("close incoming cache error", zap.Error(e)) } return peers, err @@ -387,11 +483,19 @@ func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error // AddPeer adds an incoming peer. pr must be of a comparable type (a pointer): // the pool evicts it by instance. func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { - // bounds the retries on an entry that is neither pickable nor gone: one - // mid-close, whose teardown may be slow (then ErrExists, as before) + // Bounds the passes over an entry for the same id that is still there + // after this call dealt with it: one whose close another closer holds + // (waited for below), or one another AddPeer keeps replacing. Then + // ErrExists, as before. A swap does not count: the add simply moves to + // the current pair, as many times as flushes come (in practice the caller + // coalesces its flushes). const retries = 3 - attempts := 0 - for { + for attempt := 0; ; attempt++ { + if p.closingCtx.Err() != nil { + // shutting down: nothing is evicted any more (discard is a no-op), + // so there is nothing to retry towards + return ocache.ErrClosed + } c, err := p.addIncoming(pr) if err != ocache.ErrExists { return err @@ -401,45 +505,56 @@ func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // removal can wait on a GC TryClose that holds the entry, which must // not stall a Flush and every AddPeer queued behind it. v, e := c.incoming.Pick(ctx, pr.Id()) + if p.current.Load() != c { + // the pair was flushed meanwhile: whatever it holds is the + // flush's to tear down; add to the current pair instead + continue + } + if e == ocache.ErrClosed { + return e + } + if attempt == retries { + return ocache.ErrExists + } if e != nil { if err = ctx.Err(); err != nil { return err } - if p.current.Load() != c { - // the pair was flushed meanwhile: add to the current one - continue - } - if e == ocache.ErrClosed { - return e - } - // The entry was a transient one: a concurrent Get(id) creates a - // loading entry that the incoming loader fails with ErrNotExists - // (Pick waited that out), or the previous connection is mid-close. - if attempts++; attempts <= retries { - continue + // The entry is transient: a loading one a concurrent Get(id) + // created (Pick waited it out; the incoming loader fails it) or + // the previous connection mid-close. Pick does not wait for a + // close, so wait here: a remote that reconnects while its old + // connection is still being torn down must get in once that is + // done, not be refused within microseconds. Bounded by ctx and by + // the pair staying current. + if err = p.waitClosing(ctx, c, pr.Id()); err != nil && c.ctx.Err() == nil { + return err } - return ocache.ErrExists + continue } // The old connection's teardown (Close, then the instance-safe // removal that never touches a replacement) runs in the background - // and is waited for only as long as ctx allows: a hung transport must - // not stall the accept path. - old, isPeer := v.(peer.Peer) - if !isPeer { - _, _ = c.incoming.RemoveSame(ctx, pr.Id(), v) - } else { - select { - case <-p.discard(c.incoming, old): - case <-ctx.Done(): - return ctx.Err() - } - } - if err = ctx.Err(); err != nil { - return err + // and is waited for only as long as ctx allows and the pair stays + // current: a hung transport must not stall the accept path. The + // incoming cache holds peers only. + select { + case <-p.discard(c.incoming, v.(peer.Peer)): + case <-c.ctx.Done(): + case <-ctx.Done(): + return ctx.Err() } } } +// waitClosing waits for the incoming entry for id of pair c to finish +// closing, bounded by ctx and by c staying current +func (p *pool) waitClosing(ctx context.Context, c *caches, id string) error { + lctx, cancel := context.WithCancel(ctx) + defer cancel() + defer context.AfterFunc(c.ctx, cancel)() + return c.peekIncoming.WaitClosing(lctx, id) +} + // addIncoming adds pr to the current incoming cache and returns the pair it // used. The read lock spans the Add, which never blocks, so a peer accepted // after a Flush published its pair can only land in that pair, never in one @@ -453,7 +568,7 @@ func (p *pool) addIncoming(pr peer.Peer) (*caches, error) { if err := c.incoming.Add(pr.Id(), pr); err != nil { return c, err } - go p.evictOnClose(pr, c.incoming, true) + go p.evictOnClose(pr, c, true) return c, nil } diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go index 61bbd4588..4cd2387b4 100644 --- a/net/pool/pool_flush_test.go +++ b/net/pool/pool_flush_test.go @@ -1230,6 +1230,10 @@ func TestPool_FlushSwap(t *testing.T) { require.NoError(t, fx.Service.Close(cctx)) require.Less(t, time.Since(start), time.Second) require.False(t, hung.IsClosed()) + // the teardown finishes once the peer lets it; a second Close is a no-op + doReleaseClose() + fx.Finish() + require.Eventually(t, hung.IsClosed, time.Second, 10*time.Millisecond) }) t.Run("hung incoming close does not hold back the other incoming peers", func(t *testing.T) { fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { @@ -1259,6 +1263,343 @@ func TestPool_FlushSwap(t *testing.T) { }, 100*time.Millisecond, time.Millisecond) require.False(t, hung.IsClosed()) }) + t.Run("add during shutdown returns ErrClosed without spinning", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + require.NoError(t, fx.AddPeer(ctx, newTestPeer("d"))) + // shutdown has begun: nothing is evicted any more, so the duplicate + // can never be added; it must not be retried either (this used to + // loop ~340k times in 200ms) + p.closingCancel() + dup := newCtlPeer("d") + got := make(chan error, 1) + go func() { got <- fx.AddPeer(context.Background(), dup) }() + select { + case err := <-got: + require.ErrorIs(t, err, ocache.ErrClosed) + case <-time.After(2 * time.Second): + t.Fatal("AddPeer did not return") + } + require.LessOrEqual(t, dup.idCall.Load(), int32(2), "AddPeer kept retrying") + }) + t.Run("add never waits on a flushed pair's teardown", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + old := newCtlPeer("d") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + old.closeHook = func(int32) { <-releaseClose } + require.NoError(t, fx.AddPeer(ctx, old)) + repl := newCtlPeer("d") + repl.idHook = func(call int32) { + if call == 2 { + // between ErrExists and the Pick: the pair holding old is + // replaced; old's hung teardown is now the flush's problem + assert.NoError(t, fx.Flush(ctx)) + } + } + got := make(chan error, 1) + go func() { got <- fx.AddPeer(context.Background(), repl) }() + select { + case err := <-got: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("AddPeer waited on the old pair's teardown") + } + require.True(t, inCurrent(p, repl)) + require.False(t, repl.IsClosed()) + }) + t.Run("a live incoming arriving while a dead outgoing is evicted is used instead of a dial", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + dead := newCtlPeer("p1") + close(dead.closed) + inClose := make(chan struct{}) + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + dead.closeHook = func(call int32) { + if call == 1 { + close(inClose) + } + <-releaseClose + } + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + got := make(chan peer.Peer, 1) + go func() { + gctx, cancel := context.WithTimeout(ctx, 3*time.Second) + defer cancel() + pr, err := fx.Get(gctx, "p1") + assert.NoError(t, err) + got <- pr + }() + // the Get is evicting dead; the peer connects to us meanwhile + <-inClose + live := newTestPeer("p1") + require.NoError(t, fx.AddPeer(ctx, live)) + doReleaseClose() + require.Same(t, live, <-got) + require.Zero(t, dials.Load(), "dialed although an incoming connection was available") + }) + t.Run("flush reports Closed before it returns, once per instance, before the redial", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // Connected comes from peerservice, before the peer reaches the pool + connected := func(id string, inbound bool) { + p.observer.Notify(peerobserver.Event{Kind: peerobserver.KindConnected, PeerId: id, Inbound: inbound}) + } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + connected(peerId, false) + return newTestPeer(peerId), nil + } + for round := 0; round < 3; round++ { + _, err := fx.Get(ctx, "x") + require.NoError(t, err) + connected("y", true) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("y"))) + require.NoError(t, fx.Flush(ctx)) + // delivered synchronously, in both directions + require.Len(t, obs.getClosed(), 2*(round+1)) + } + // the watchers of the flushed peers must not report them again + require.Never(t, func() bool { return len(obs.getClosed()) > 6 }, 200*time.Millisecond, 10*time.Millisecond) + for _, id := range []string{"x", "y"} { + kinds := obs.kindsFor(id) + require.Len(t, kinds, 6, id) + for i, k := range kinds { + want := peerobserver.KindConnected + if i%2 == 1 { + want = peerobserver.KindClosed + } + require.Equal(t, want, k, "%s: event %d", id, i) + } + } + }) + t.Run("hung outgoing close does not delay closing the incoming cache", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 2 * time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + hung := newCtlPeer("out") + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + hung.closeHook = func(int32) { <-releaseClose } + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return hung, nil + } + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + in := newTestPeer("in") + require.NoError(t, fx.AddPeer(ctx, in)) + oldPair := p.current.Load() + require.NoError(t, fx.Flush(ctx)) + // well within closeTimeout: the incoming cache does not queue behind + // the outgoing pass that is stuck on hung + require.Eventually(t, func() bool { + _, err := oldPair.incoming.Pick(ctx, "in") + return err == ocache.ErrClosed && in.IsClosed() + }, 500*time.Millisecond, time.Millisecond) + require.False(t, hung.IsClosed()) + }) + t.Run("add gives up on an entry that cannot be evicted", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // an incoming cache that never removes anything stands in for an + // entry some other closer keeps hold of + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + c.incoming = &stickyCache{OCache: c.incoming, peek: c.peekIncoming} + c.peekIncoming = mustPeeker(c.incoming) + return c + } + require.NoError(t, fx.Flush(ctx)) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("d"))) + dup := newCtlPeer("d") + got := make(chan error, 1) + go func() { got <- fx.AddPeer(context.Background(), dup) }() + select { + case err := <-got: + require.ErrorIs(t, err, ocache.ErrExists) + case <-time.After(2 * time.Second): + t.Fatal("AddPeer kept retrying an entry it cannot evict") + } + // one pass plus the three retries, two Id calls each (Add, Pick) + require.Equal(t, int32(8), dup.idCall.Load()) + }) + t.Run("add waits for the old connection's close before giving up", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + old := newCtlPeer("d") + inClose := make(chan struct{}) + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + old.closeHook = func(call int32) { + if call == 1 { + close(inClose) + } + <-releaseClose + } + require.NoError(t, fx.AddPeer(ctx, old)) + // the connection dies: its watcher evicts it, and the close hangs + close(old.closed) + <-inClose + // the remote reconnects while the old entry is still mid-close + repl := newTestPeer("d") + got := make(chan error, 1) + go func() { + actx, cancel := context.WithTimeout(ctx, 3*time.Second) + defer cancel() + got <- fx.AddPeer(actx, repl) + }() + require.Never(t, func() bool { return len(got) > 0 }, 100*time.Millisecond, 10*time.Millisecond) + doReleaseClose() + require.NoError(t, <-got) + require.True(t, inCurrent(p, repl)) + require.False(t, repl.IsClosed()) + }) + t.Run("add parked on an old connection's close moves on when the pair is flushed", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + old := newCtlPeer("d") + inClose := make(chan struct{}) + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + old.closeHook = func(call int32) { + if call == 1 { + close(inClose) + } + <-releaseClose + } + require.NoError(t, fx.AddPeer(ctx, old)) + repl := newTestPeer("d") + got := make(chan error, 1) + // Accept passes a Background ctx: only the pair can end the wait + go func() { got <- fx.AddPeer(context.Background(), repl) }() + <-inClose + require.NoError(t, fx.Flush(ctx)) + select { + case err := <-got: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("AddPeer stayed parked on the flushed pair's teardown") + } + require.True(t, inCurrent(p, repl)) + }) + t.Run("a watcher evicting a peer while flush snapshots it reports it once", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + // the watcher's RemoveSame starts a Flush and lets it snapshot the + // still-present peer, then removes; Flush may only mark after that + snapshotDone := make(chan struct{}) + proceed := make(chan struct{}) + var flushed sync.WaitGroup + var once sync.Once + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + hc := &hookedCache{OCache: inner, peek: peek} + hc.onForEach = func() { + once.Do(func() { + close(snapshotDone) + <-proceed + }) + } + hc.onRemoveSame = func() { + flushed.Add(1) + go func() { + defer flushed.Done() + assert.NoError(t, fx.Flush(ctx)) + }() + <-snapshotDone + } + return hc + }) + tp := newTestPeer("x") + require.NoError(t, fx.AddPeer(ctx, tp)) + require.NoError(t, tp.Close()) + <-snapshotDone + // the watcher's removal runs now, with Flush holding the snapshot + time.Sleep(20 * time.Millisecond) + close(proceed) + flushed.Wait() + require.Eventually(t, func() bool { return len(obs.getClosed()) == 1 }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 1 }, 200*time.Millisecond, 10*time.Millisecond) + }) + t.Run("flush racing close reports nothing after shutdown began", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureCfg(t, obs, func(ps *poolService, a *app.App) { + ps.closeTimeout = 2 * time.Second + }) + p := fx.Service.(*poolService).pool + inSnapshot := make(chan struct{}) + proceed := make(chan struct{}) + var once sync.Once + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onForEach: func() { + once.Do(func() { + close(inSnapshot) + <-proceed + }) + }} + }) + require.NoError(t, fx.AddPeer(ctx, newTestPeer("x"))) + flushed := make(chan struct{}) + go func() { + assert.NoError(t, fx.Flush(ctx)) + close(flushed) + }() + <-inSnapshot + // Close begins while Flush holds its snapshot and has not reported yet + closed := make(chan struct{}) + go func() { + _ = fx.Service.Close(ctx) + close(closed) + }() + require.Eventually(t, func() bool { return p.closingCtx.Err() != nil }, time.Second, time.Millisecond) + close(proceed) + <-flushed + <-closed + require.Empty(t, obs.getClosed()) + }) + t.Run("get whose ctx ends as the eviction completes does not dial", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + for i := 0; i < 20; i++ { + gctx, cancel := context.WithCancel(ctx) + dead := newCtlPeer("p1") + close(dead.closed) + // the eviction's Close cancels the Get's ctx as it completes + dead.closeHook = func(int32) { cancel() } + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + _, err := fx.Get(gctx, "p1") + require.ErrorIs(t, err, context.Canceled) + cancel() + } + require.Zero(t, dials.Load(), "dialed with a done ctx after the eviction") + }) t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { obs := &poolEventRecorder{} fx := newFixtureWithObserver(t, obs) @@ -1565,3 +1906,140 @@ func TestPool_FlushGoroutines(t *testing.T) { time.Sleep(10 * time.Millisecond) } } + +// kindsFor returns the kinds of every event recorded for peerId, in order +func (r *poolEventRecorder) kindsFor(peerId string) (kinds []peerobserver.Kind) { + r.mu.Lock() + defer r.mu.Unlock() + for _, ev := range r.events { + if ev.PeerId == peerId { + kinds = append(kinds, ev.Kind) + } + } + return +} + +// TestPool_FastPathMetrics pins the series of the hit path against what a Get +// or Pick through the caches counts: exactly once per cache per call, also +// when the fast path finds a closed peer and falls through to lookup +func TestPool_FastPathMetrics(t *testing.T) { + reg := prometheus.NewRegistry() + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + a.Register(&testMetric{reg: reg}) + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + counters := func() map[string]float64 { + families, err := reg.Gather() + require.NoError(t, err) + out := map[string]float64{} + for _, mf := range families { + if m := mf.GetMetric()[0]; m.GetCounter() != nil { + out[mf.GetName()] = m.GetCounter().GetValue() + } + } + return out + } + expect := func(step string, incomingHit, incomingMiss, outgoingHit, outgoingMiss float64) { + got := counters() + assert.Equal(t, map[string]float64{ + "netpool_incoming_hit": incomingHit, "netpool_incoming_miss": incomingMiss, "netpool_incoming_gc": 0, + "netpool_outgoing_hit": outgoingHit, "netpool_outgoing_miss": outgoingMiss, "netpool_outgoing_gc": 0, + }, got, step) + } + _, err := fx.Get(ctx, "out") // incoming miss, outgoing miss (dial) + require.NoError(t, err) + expect("dial", 0, 1, 0, 1) + _, err = fx.Get(ctx, "out") // fast: incoming miss, outgoing hit + require.NoError(t, err) + expect("fast get", 0, 2, 1, 1) + _, err = fx.Pick(ctx, "out") // fast: Pick counts the hit only + require.NoError(t, err) + expect("fast pick", 0, 2, 2, 1) + _, err = fx.GetOneOf(ctx, []string{"zz", "out"}) // a Pick miss counts nothing, then the hit + require.NoError(t, err) + expect("fast getoneof", 0, 2, 3, 1) + + // a closed peer in the cache: fast counts nothing and lookup takes over + // (incoming miss, outgoing hit on the closed entry, then after the + // eviction incoming miss and outgoing miss for the dial) + dead := newTestPeer("d") + require.NoError(t, dead.Close()) + require.NoError(t, p.current.Load().outgoing.Add("d", dead)) + _, err = fx.Get(ctx, "d") + require.NoError(t, err) + expect("fallback", 0, 4, 4, 2) + + require.NoError(t, fx.AddPeer(ctx, newTestPeer("in"))) + _, err = fx.Get(ctx, "in") // fast: incoming hit + require.NoError(t, err) + _, err = fx.Pick(ctx, "in") + require.NoError(t, err) + expect("incoming", 2, 4, 4, 2) +} + +// stickyCache is an OCache whose RemoveSame never removes anything +type stickyCache struct { + ocache.OCache + peek ocache.Peeker +} + +func (c *stickyCache) RemoveSame(ctx context.Context, id string, value ocache.Object) (bool, error) { + return false, nil +} + +func (c *stickyCache) Peek(id string, touch bool) (ocache.Object, bool) { + return c.peek.Peek(id, touch) +} + +func (c *stickyCache) WaitClosing(ctx context.Context, id string) error { + return c.peek.WaitClosing(ctx, id) +} + +// hookedCache is an incoming cache whose RemoveSame and ForEach can be +// intercepted: the seams for racing a Flush against a watcher's eviction +type hookedCache struct { + ocache.OCache + peek ocache.Peeker + onRemoveSame func() + onForEach func() +} + +func (c *hookedCache) RemoveSame(ctx context.Context, id string, value ocache.Object) (bool, error) { + if c.onRemoveSame != nil { + c.onRemoveSame() + } + return c.OCache.RemoveSame(ctx, id, value) +} + +func (c *hookedCache) ForEach(f func(v ocache.Object) bool) { + c.OCache.ForEach(f) + if c.onForEach != nil { + c.onForEach() + } +} + +func (c *hookedCache) Peek(id string, touch bool) (ocache.Object, bool) { + return c.peek.Peek(id, touch) +} + +func (c *hookedCache) WaitClosing(ctx context.Context, id string) error { + return c.peek.WaitClosing(ctx, id) +} + +// installIncoming makes every pair the pool builds from now on use wrap(inner) +// as its incoming cache, and flushes once so the current pair has it +func installIncoming(t *testing.T, fx *fixture, wrap func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache) { + p := fx.Service.(*poolService).pool + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + c.incoming = wrap(c.incoming, c.peekIncoming) + c.peekIncoming = mustPeeker(c.incoming) + return c + } + require.NoError(t, fx.Flush(ctx)) +} diff --git a/net/pool/pool_test.go b/net/pool/pool_test.go index e56b175a9..41d73a05e 100644 --- a/net/pool/pool_test.go +++ b/net/pool/pool_test.go @@ -762,7 +762,7 @@ func TestPool_EvictOnClose_ExitsOnShutdownWithoutEviction(t *testing.T) { done := make(chan struct{}) go func() { - p.evictOnClose(tp, p.current.Load().incoming, true) + p.evictOnClose(tp, p.current.Load(), true) close(done) }() @@ -889,7 +889,7 @@ func TestPool_PeerObserver(t *testing.T) { cache := &shutdownOnRemoveSame{OCache: p.current.Load().incoming, cancel: p.closingCancel} done := make(chan struct{}) go func() { - p.evictOnClose(tp, cache, true) + p.evictOnClose(tp, &caches{incoming: cache}, true) close(done) }() select { @@ -1035,7 +1035,7 @@ func TestPool_EvictsOutgoingClosedDuringLoad(t *testing.T) { require.NoError(t, tp.Close()) done := make(chan struct{}) go func() { - p.evictOnClose(tp, cache, false) + p.evictOnClose(tp, &caches{outgoing: cache}, false) close(done) }() diff --git a/net/pool/poolservice.go b/net/pool/poolservice.go index 6951318b9..62a0936ab 100644 --- a/net/pool/poolservice.go +++ b/net/pool/poolservice.go @@ -51,9 +51,6 @@ type poolService struct { func (p *poolService) Init(a *app.App) (err error) { p.dialer = a.MustComponent("net.peerservice").(dialer) - if p.pool.closeTimeout <= 0 { - p.pool.closeTimeout = closeTimeout - } p.pool.closingCtx, p.pool.closingCancel = context.WithCancel(context.Background()) if m := a.Component(metric.CName); m != nil { p.metricReg = m.(metric.Metric).Registry() @@ -74,7 +71,7 @@ func (p *poolService) Init(a *app.App) (err error) { return p.pool.current.Load().incoming.Len() }) outgoingMetrics, incomingMetrics = outgoing.Option(), incoming.Option() - p.pool.incomingMiss = incoming.Miss + p.pool.metrics = &fastMetrics{incomingHit: incoming.Hit, incomingMiss: incoming.Miss, outgoingHit: outgoing.Hit} } p.pool.newCaches = func() *caches { return p.newCaches(outgoingMetrics, incomingMetrics) @@ -110,7 +107,7 @@ func (p *poolService) newCaches(outgoingMetrics, incomingMetrics ocache.Option) return value, err } if pr, ok := value.(peer.Peer); ok { - go p.pool.evictOnClose(pr, c.outgoing, false) + go p.pool.evictOnClose(pr, c, false) } return value, nil }, diff --git a/net/secureservice/handshake/proto.go b/net/secureservice/handshake/proto.go index 17d63530e..821cdb3b8 100644 --- a/net/secureservice/handshake/proto.go +++ b/net/secureservice/handshake/proto.go @@ -16,40 +16,52 @@ type ProtoChecker struct { SupportedEncodings []handshakeproto.Encoding } -// OutgoingProtoHandshake negotiates the sub-connection protocol. -// -// Contract: on an I/O error or ctx cancellation the conn is closed; on a +// OutgoingProtoHandshake negotiates the sub-connection protocol. On an I/O +// error or ctx cancellation the conn is closed before it returns; on a // protocol-level error (incompatible, declined or unexpected proto) it is -// left to the caller. On cancellation the function returns at once and the -// close happens in the background: a stream close can block on the transport -// (a yamux FIN waits up to the connection write timeout), and the caller is -// typically racing a deadline. I/O errors close it asynchronously as well, so -// the conn may still be open when the function returns; it is unusable -// afterwards either way. +// left to the caller. The close is synchronous and can block on the +// transport; OutgoingProtoHandshakeWithCloser moves it off the caller's path. func OutgoingProtoHandshake(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto) (*handshakeproto.Proto, error) { - return OutgoingProtoHandshakeWithCloser(ctx, conn, proto, nil) + if ctx == nil { + ctx = context.Background() + } + h := newHandshake() + done := make(chan struct{}) + var ( + err error + remoteProto *handshakeproto.Proto + ) + go func() { + defer close(done) + remoteProto, err = outgoingProtoHandshake(h, conn, proto, nil, nil) + }() + select { + case <-done: + return remoteProto, err + case <-ctx.Done(): + _ = conn.Close() + return nil, ctx.Err() + } } -// OutgoingProtoHandshakeWithCloser is OutgoingProtoHandshake with every close -// of conn going through closeConn, which must not block (e.g. it hands the -// conn to a bounded cleanup worker). With a closer the conn is handed to it -// exactly once on any error, protocol-level ones included, so the caller -// never closes it itself. A nil closeConn keeps OutgoingProtoHandshake's -// contract and closes in a new goroutine. +// OutgoingProtoHandshakeWithCloser is OutgoingProtoHandshake for a caller +// racing a deadline: every close of conn goes through closeConn, which must +// not block (e.g. it closes the conn in the background), and on any error, +// protocol-level ones included, the conn is handed to it exactly once, so the +// caller never closes it itself. On cancellation it returns at once. A nil +// closeConn falls back to OutgoingProtoHandshake. func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto, closeConn func(net.Conn)) (*handshakeproto.Proto, error) { + if closeConn == nil { + return OutgoingProtoHandshake(ctx, conn, proto) + } if ctx == nil { ctx = context.Background() } - closeOnAnyErr := closeConn != nil - if closeConn == nil { - closeConn = closeAsync - } else { - var handedOff atomic.Bool - hook := closeConn - closeConn = func(c net.Conn) { - if handedOff.CompareAndSwap(false, true) { - hook(c) - } + var handedOff atomic.Bool + hook := closeConn + closeConn = func(c net.Conn) { + if handedOff.CompareAndSwap(false, true) { + hook(c) } } h := newHandshake() @@ -65,7 +77,7 @@ func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto go func() { defer close(done) remoteProto, err = outgoingProtoHandshake(h, conn, proto, closeConn, &claimed) - if err != nil && closeOnAnyErr { + if err != nil { // a no-op if the handshake or an abandoning caller closed it closeConn(conn) } @@ -90,10 +102,6 @@ func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto } } -func closeAsync(conn net.Conn) { - go func() { _ = conn.Close() }() -} - var noEncodings = []handshakeproto.Encoding{handshakeproto.Encoding_None} func outgoingProtoHandshake(h *handshake, conn net.Conn, proto *handshakeproto.Proto, closeConn func(net.Conn), abandoned *atomic.Bool) (remoteProto *handshakeproto.Proto, err error) { diff --git a/net/secureservice/handshake/proto_test.go b/net/secureservice/handshake/proto_test.go index 8599d703b..8e35d32d2 100644 --- a/net/secureservice/handshake/proto_test.go +++ b/net/secureservice/handshake/proto_test.go @@ -7,6 +7,7 @@ import ( "net" "os" "sync" + "sync/atomic" "testing" "time" @@ -222,7 +223,7 @@ func (c *blockingCloseConn) Close() error { return c.Conn.Close() } -func TestOutgoingProtoHandshake_CancelDoesNotWaitForClose(t *testing.T) { +func TestOutgoingProtoHandshakeWithCloser_CancelDoesNotWaitForClose(t *testing.T) { c1, c2 := net.Pipe() defer c2.Close() conn := &blockingCloseConn{Conn: c1, closeCalled: make(chan struct{}), release: make(chan struct{})} @@ -233,7 +234,8 @@ func TestOutgoingProtoHandshake_CancelDoesNotWaitForClose(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) defer cancel() start := time.Now() - _, err := OutgoingProtoHandshake(ctx, conn, &handshakeproto.Proto{Proto: 1}) + closer := func(c net.Conn) { go func() { _ = c.Close() }() } + _, err := OutgoingProtoHandshakeWithCloser(ctx, conn, &handshakeproto.Proto{Proto: 1}, closer) require.ErrorIs(t, err, context.DeadlineExceeded) assert.Less(t, time.Since(start), time.Second, "the caller must not wait on the blocked close") @@ -245,6 +247,50 @@ func TestOutgoingProtoHandshake_CancelDoesNotWaitForClose(t *testing.T) { } } +// recordCloseConn records whether Close has returned +type recordCloseConn struct { + net.Conn + closed atomic.Bool + // delay makes the close slow, so an asynchronous one is visibly late + delay time.Duration +} + +func (c *recordCloseConn) Close() error { + time.Sleep(c.delay) + err := c.Conn.Close() + c.closed.Store(true) + return err +} + +// TestOutgoingProtoHandshake_ClosesSynchronously pins the exported contract: +// the conn is closed by the time OutgoingProtoHandshake returns +func TestOutgoingProtoHandshake_ClosesSynchronously(t *testing.T) { + t.Run("cancel", func(t *testing.T) { + c1, c2 := net.Pipe() + defer c2.Close() + conn := &recordCloseConn{Conn: c1} + go func() { _, _ = io.Copy(io.Discard, c2) }() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + _, err := OutgoingProtoHandshake(ctx, conn, &handshakeproto.Proto{Proto: 1}) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.True(t, conn.closed.Load(), "closed before returning") + }) + t.Run("I/O error", func(t *testing.T) { + c1, c2 := net.Pipe() + conn := &recordCloseConn{Conn: c1} + go func() { + h := newHandshake() + h.conn = c2 + _, _ = h.readMsg(msgTypeProto) + _ = c2.Close() + }() + _, err := OutgoingProtoHandshake(context.Background(), conn, &handshakeproto.Proto{Proto: 1}) + require.Error(t, err) + assert.True(t, conn.closed.Load(), "closed before returning") + }) +} + // noDeadlineConn models a conn whose SetDeadline is a no-op (as on wasm) type noDeadlineConn struct { net.Conn @@ -298,6 +344,9 @@ func TestOutgoingProtoHandshakeWithCloser_Cancel(t *testing.T) { func TestOutgoingProtoHandshakeWithCloser_IOErrorUsesCloser(t *testing.T) { c1, c2 := net.Pipe() + // a conn whose own Close blocks: only the closer may close it + conn := &blockingCloseConn{Conn: c1, closeCalled: make(chan struct{}), release: make(chan struct{})} + defer close(conn.release) // the remote goes away mid-handshake go func() { h := newHandshake() @@ -305,18 +354,63 @@ func TestOutgoingProtoHandshakeWithCloser_IOErrorUsesCloser(t *testing.T) { _, _ = h.readMsg(msgTypeProto) _ = c2.Close() }() - closed := make(chan struct{}, 1) - closer := func(c net.Conn) { - closed <- struct{}{} - _ = c.Close() - } - _, err := OutgoingProtoHandshakeWithCloser(context.Background(), c1, &handshakeproto.Proto{Proto: 1}, closer) - require.Error(t, err) + handed := make(chan net.Conn, 2) + closer := func(c net.Conn) { handed <- c } + res := make(chan error, 1) + go func() { + _, err := OutgoingProtoHandshakeWithCloser(context.Background(), conn, &handshakeproto.Proto{Proto: 1}, closer) + res <- err + }() select { - case <-closed: + case err := <-res: + require.Error(t, err) case <-time.After(time.Second): + t.Fatal("the handshake ran the blocking Close instead of the closer") + } + select { + case c := <-handed: + assert.Equal(t, net.Conn(conn), c) + default: t.Fatal("I/O error close did not go through the closer") } + assert.Empty(t, handed, "handed over exactly once") + select { + case <-conn.closeCalled: + t.Fatal("Close must be left to the closer") + default: + } +} + +func TestOutgoingProtoHandshakeWithCloser_ProtocolErrorHandsOver(t *testing.T) { + c1, c2 := net.Pipe() + defer c2.Close() + // the remote declines the protocol: a protocol-level error, which the + // legacy function leaves to the caller + go func() { + _, _ = IncomingProtoHandshake(context.Background(), c2, newProtoChecker(100)) + }() + handed := make(chan net.Conn, 2) + _, err := OutgoingProtoHandshakeWithCloser(context.Background(), c1, &handshakeproto.Proto{Proto: 1}, func(c net.Conn) { handed <- c }) + require.ErrorIs(t, err, ErrRemoteIncompatibleProto) + select { + case c := <-handed: + assert.Equal(t, c1, c) + case <-time.After(time.Second): + t.Fatal("the conn was not handed to the closer") + } + assert.Empty(t, handed, "handed over exactly once") +} + +func TestOutgoingProtoHandshakeWithCloser_NilCloserIsSynchronous(t *testing.T) { + c1, c2 := net.Pipe() + defer c2.Close() + conn := &recordCloseConn{Conn: c1, delay: 50 * time.Millisecond} + go func() { _, _ = io.Copy(io.Discard, c2) }() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + _, err := OutgoingProtoHandshakeWithCloser(ctx, conn, &handshakeproto.Proto{Proto: 1}, nil) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.True(t, conn.closed.Load(), "a nil closer keeps the synchronous close") } func TestHandshakeError_Unwrap(t *testing.T) { diff --git a/net/transport/quic/conn.go b/net/transport/quic/conn.go index 5482f22e8..f773688af 100644 --- a/net/transport/quic/conn.go +++ b/net/transport/quic/conn.go @@ -88,33 +88,15 @@ func isConnDead(err error) bool { return errors.Is(err, quic.ErrServerClosed) || errors.Is(err, net.ErrClosed) } -// connDeadError is a stream error caused by the whole connection going away. -// It matches transport.ErrConnClosed, so callers classify it like a failed -// Open or Accept, and still unwraps to the original quic error for telemetry. -type connDeadError struct { - cause error -} - -func (e connDeadError) Error() string { - return transport.ErrConnClosed.Error() + ": " + e.cause.Error() -} - -func (e connDeadError) Unwrap() []error { - return []error{transport.ErrConnClosed, e.cause} -} - // wrapConnDead normalizes a stream Read/Write error: one meaning the -// connection is dead is wrapped into connDeadError, anything else (io.EOF, -// stream resets, deadlines) is returned as is. +// connection is dead is wrapped with transport.NewConnClosedError, so callers +// classify it like a failed Open or Accept while the original quic error stays +// reachable; anything else (io.EOF, stream resets, deadlines) is returned as is. func wrapConnDead(err error) error { if err == nil || !isConnDead(err) { return err } - var already connDeadError - if errors.As(err, &already) { - return err - } - return connDeadError{cause: err} + return transport.NewConnClosedError(err) } func (q *quicMultiConn) Accept() (conn net.Conn, err error) { diff --git a/net/transport/transport.go b/net/transport/transport.go index cb0797b9b..927c41b6d 100644 --- a/net/transport/transport.go +++ b/net/transport/transport.go @@ -3,6 +3,7 @@ package transport import ( "context" + "errors" "net" "time" ) @@ -22,6 +23,30 @@ func (connClosedError) Error() string { return "transport connection closed" } func (connClosedError) Unwrap() error { return net.ErrClosed } +// NewConnClosedError wraps cause, an error a transport got because the whole +// connection went away, so that it matches ErrConnClosed while errors.As +// still finds the original error. A cause that is already wrapped is +// returned as is. +func NewConnClosedError(cause error) error { + var already connClosedCauseError + if errors.As(cause, &already) { + return cause + } + return connClosedCauseError{cause: cause} +} + +type connClosedCauseError struct { + cause error +} + +func (e connClosedCauseError) Error() string { + return ErrConnClosed.Error() + ": " + e.cause.Error() +} + +func (e connClosedCauseError) Unwrap() []error { + return []error{ErrConnClosed, e.cause} +} + const ( Yamux = "yamux" Quic = "quic" diff --git a/net/transport/transport_test.go b/net/transport/transport_test.go index 8f75c0932..96ac071ca 100644 --- a/net/transport/transport_test.go +++ b/net/transport/transport_test.go @@ -15,3 +15,19 @@ func TestErrConnClosed(t *testing.T) { // distinctive text so it can be told apart from other connection errors in logs assert.Equal(t, "transport connection closed", ErrConnClosed.Error()) } + +type causeErr struct{} + +func (causeErr) Error() string { return "cause" } + +func TestNewConnClosedError(t *testing.T) { + cause := &causeErr{} + err := NewConnClosedError(cause) + assert.ErrorIs(t, err, ErrConnClosed) + assert.ErrorIs(t, err, net.ErrClosed) + var got *causeErr + assert.True(t, errors.As(err, &got), "the original error stays reachable") + assert.Equal(t, "transport connection closed: cause", err.Error()) + // idempotent + assert.Equal(t, err, NewConnClosedError(err)) +} diff --git a/net/transport/yamux/conn.go b/net/transport/yamux/conn.go index fa03d249a..cb4b8a36e 100644 --- a/net/transport/yamux/conn.go +++ b/net/transport/yamux/conn.go @@ -199,18 +199,14 @@ func (s yamuxStream) Write(b []byte) (n int, err error) { return n, s.wrapSessionDead(err) } -// wrapSessionDead wraps err into sessionDeadError when it was caused by the -// session shutting down. io.EOF, a stream reset and a closed stream count +// wrapSessionDead wraps err with transport.NewConnClosedError when it was +// caused by the session shutting down. io.EOF, a stream reset and a closed stream count // only while the session is closed: on a live session they are stream-level // outcomes (a remote close or reset) and are returned unchanged. func (s yamuxStream) wrapSessionDead(err error) error { if err == nil { return nil } - var already sessionDeadError - if errors.As(err, &already) { - return err - } switch { case errors.Is(err, yamux.ErrSessionShutdown): case errors.Is(err, io.EOF), errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed): @@ -220,19 +216,5 @@ func (s yamuxStream) wrapSessionDead(err error) error { default: return err } - return sessionDeadError{cause: err} -} - -// sessionDeadError matches transport.ErrConnClosed and still unwraps to the -// original yamux error -type sessionDeadError struct { - cause error -} - -func (e sessionDeadError) Error() string { - return transport.ErrConnClosed.Error() + ": " + e.cause.Error() -} - -func (e sessionDeadError) Unwrap() []error { - return []error{transport.ErrConnClosed, e.cause} + return transport.NewConnClosedError(err) } From b672339796f06c1cfe1b2aca9a15e0449f2b4c1e Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 16:43:18 +0200 Subject: [PATCH 4/7] net/pool, net/peerservice: address #802 recheck - peerservice: a dial cancelled by its ctx (e.g. Flush) tries no more addresses and records no QUIC demotion outcome - pool: Get probes incoming via Peek (never loads into the incoming cache); the redial loop stops when the pair is replaced or the pool shuts down - pool: Flush swaps don't consume AddPeer retries; pick leaves eviction to the watcher; drop keepOnFlush - pool: only the Flush that replaced a pair walks it, once, after the swap, so concurrent Flushes report each peer exactly once - yamux: a write timeout is ErrConnClosed only once the session closed --- net/peerservice/dialoutcome_test.go | 89 ++++++ net/peerservice/peerservice.go | 18 +- net/pool/pool.go | 211 ++++++++----- net/pool/pool_flush_test.go | 447 +++++++++++++++++++++++++--- net/pool/poolservice.go | 20 +- net/transport/yamux/conn.go | 20 +- net/transport/yamux/conn_test.go | 65 ++++ net/transport/yamux/export_test.go | 8 + 8 files changed, 729 insertions(+), 149 deletions(-) create mode 100644 net/transport/yamux/export_test.go diff --git a/net/peerservice/dialoutcome_test.go b/net/peerservice/dialoutcome_test.go index 14c81330b..a2fe92f9f 100644 --- a/net/peerservice/dialoutcome_test.go +++ b/net/peerservice/dialoutcome_test.go @@ -3,6 +3,7 @@ package peerservice import ( "context" "fmt" + "sync/atomic" "testing" quicgo "github.com/quic-go/quic-go" @@ -11,6 +12,7 @@ import ( "go.uber.org/mock/gomock" "github.com/anyproto/any-sync/app" + "github.com/anyproto/any-sync/net/pool" "github.com/anyproto/any-sync/net/quicdemotion" "github.com/anyproto/any-sync/net/transport" ) @@ -226,3 +228,90 @@ func TestPeerService_DialOutcome(t *testing.T) { "reported as-is; that webtransport proves udp works is the component's business") }) } + +func TestPeerService_CancelledDialReportsNoOutcome(t *testing.T) { + const peerId = "p1" + t.Run("a dial the pool flush cancels is not reported and tries no further address", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true).AnyTimes() + // quic first; the first attempt is the dial a Flush catches in flight + // and only ends with its ctx; the yamux fallback must not be tried + // with the dead ctx (no expectation: gomock fails on a call) + dialStarted := make(chan struct{}) + var dials atomic.Int32 + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + if dials.Add(1) == 1 { + close(dialStarted) + <-ctx.Done() + return nil, ctx.Err() + } + return fx.mockMC(peerId), nil + }).Times(2) + pl := fx.a.MustComponent(pool.CName).(pool.Service) + got := make(chan error, 1) + go func() { + _, err := pl.Get(ctx, peerId) + got <- err + }() + <-dialStarted + require.NoError(t, pl.Flush(ctx)) + require.NoError(t, <-got) + // the cancelled attempt left no trace; the redial reported normally + o := stub.only(t) + assert.Equal(t, transport.Quic, o.SucceededScheme) + assert.False(t, o.FallbackFailed) + assert.False(t, o.QuicTimedOut) + }) + t.Run("a dial the caller cancels is not reported", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true) + dctx, cancel := context.WithCancel(ctx) + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + cancel() + return nil, ctx.Err() + }) + _, err := fx.Dial(dctx, peerId) + require.ErrorIs(t, err, context.Canceled) + require.Empty(t, stub.outcomes) + }) + t.Run("an accepted dial is reported even if its ctx ends right after", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true) + dctx, cancel := context.WithCancel(ctx) + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + // the connection is made; the caller gives up as it returns + cancel() + return fx.mockMC(peerId), nil + }) + pr, err := fx.Dial(dctx, peerId) + require.NoError(t, err) + require.NotNil(t, pr) + o := stub.only(t) + assert.Equal(t, transport.Quic, o.SucceededScheme) + }) + t.Run("a cancelled dial does not stop later dials from being reported", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true) + fx.nodeConf.EXPECT().PeerAddresses("p2").Return([]string{"yamux://203.0.113.2:1111", "quic://203.0.113.2:1112"}, true) + dctx, cancel := context.WithCancel(ctx) + cancel() + _, err := fx.Dial(dctx, peerId) + require.ErrorIs(t, err, context.Canceled) + require.Empty(t, stub.outcomes) + // a real fallback failure for another peer is still recorded + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.2:1112").Return(nil, fmt.Errorf("refused")) + fx.yamux.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.2:1111").Return(nil, fmt.Errorf("refused")) + _, err = fx.Dial(ctx, "p2") + require.Error(t, err) + o := stub.only(t) + assert.Equal(t, "p2", o.PeerId) + assert.True(t, o.FallbackFailed) + }) +} diff --git a/net/peerservice/peerservice.go b/net/peerservice/peerservice.go index 6df209faa..b65deada2 100644 --- a/net/peerservice/peerservice.go +++ b/net/peerservice/peerservice.go @@ -162,10 +162,16 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er // Reported once the dial is fully resolved: a connection that is opened // and then rejected (a stale address pointing at another peer) reached // nobody, so it is neither a working fallback nor evidence about quic - // toward the peer we asked for. + // toward the peer we asked for. A dial the caller cancelled (a pool Flush + // on wake, a request that gave up) says nothing about any transport + // either: it is not reported at all, or every peer's quic would be held + // back for the fallback window on the strength of a dead context. That + // drops a QuicTimedOut recorded before the cancel on purpose: without the + // yamux try that the cancel cut short there is no proof the path works + // without udp, which is what a strike requires. dialAccepted := false defer func() { - if p.demotion == nil { + if p.demotion == nil || (!dialAccepted && ctx.Err() != nil) { return } if !dialAccepted { @@ -175,6 +181,12 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er }() err = ErrAddrsNotFound for _, addr := range ordered { + if ctx.Err() != nil { + // cancelled: the remaining addresses would only fail at once + // with the same dead context + err = ctx.Err() + break + } sch := scheme(addr) if mc, err = p.dialAddr(ctx, addr); err == nil { connAddr = addr @@ -182,6 +194,8 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er break } addrErrs = append(addrErrs, err) + // a failure under a done ctx is classified like any other: the + // outcome of a cancelled dial is never reported, see above switch { case sch == transport.Quic && quic.IsDialDegraded(err): outcome.QuicTimedOut = true diff --git a/net/pool/pool.go b/net/pool/pool.go index ff5e028bf..d33189198 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -86,10 +86,10 @@ func (c *caches) reportedByFlush(pr peer.Peer) bool { return ok } -// fastMetrics are the counters the hit path bumps itself (nil without a -// registry): ocache.Peek counts nothing, so a Get or Pick served by fast or -// by the caches counts exactly once per cache, as a Get through the caches -// did before the fast path existed +// fastMetrics are the counters the pool bumps itself where it reads a cache +// through Peek, which counts nothing (nil without a registry): a Get or Pick +// counts exactly once per cache, whichever path serves it, the series being +// those a Get through both caches produced before type fastMetrics struct { incomingHit, incomingMiss, outgoingHit prometheus.Counter } @@ -273,29 +273,72 @@ func (p *pool) Get(ctx context.Context, id string) (peer.Peer, error) { return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { for { // if we have incoming connection - try to reuse it - if pr, err = p.get(ctx, c.incoming, id); err == nil { + if pr, err = p.getIncoming(ctx, c, id); err == nil { return pr, nil } - // or try to get or create outgoing - if pr, err = p.get(ctx, c.outgoing, id); err != errRedial { - return pr, err + if err != errRedial { + // or try to get or create outgoing + if pr, err = p.get(ctx, c.outgoing, id); err != errRedial { + return pr, err + } + } + // A closed peer was evicted: start over from incoming, where a + // live connection may have arrived meanwhile, before dialing. Not + // on a pair that was replaced meanwhile: the swap cancels this + // ctx through an AfterFunc, which runs asynchronously, so the pair + // is checked directly or the old cache could be dialed into with + // a ctx about to die (lookup retries on the current pair). And + // not during shutdown, when nothing is evicted any more and the + // closed peer would be found again and again. A caller ctx that + // ended was seen by the eviction's select already. + if c.ctx.Err() != nil || p.closingCtx.Err() != nil { + return nil, ocache.ErrClosed } - // a closed peer was evicted: start over from incoming, where a - // live connection may have arrived meanwhile, before dialing (a - // ctx that ended meanwhile fails the next Get at once) } }) } -// errRedial is get's verdict after it evicted a closed peer: look again, +// errRedial is the verdict after a closed peer was evicted: look again, // starting from the incoming cache var errRedial = errors.New("closed peer evicted") +// getIncoming reads the incoming cache without loading: its entries come from +// AddPeer only, so a load would just insert a failing entry for the time of +// the probe (and make every AddPeer for that id trip over it). Counted like a +// Get on the cache would be: a hit, or a miss before the outgoing lookup. +func (p *pool) getIncoming(ctx context.Context, c *caches, id string) (peer.Peer, error) { + v, ok := c.peekIncoming.Peek(id, true) + if !ok && c.ctx.Err() != nil { + // a replaced or closing pair: no verdict about the peer, and not a + // miss (a Get on the cache would have failed with ErrClosed uncounted) + return nil, ocache.ErrClosed + } + if m := p.metrics; m != nil { + if ok { + m.incomingHit.Inc() + } else { + m.incomingMiss.Inc() + } + } + if !ok { + return nil, ocache.ErrNotExists + } + return p.live(ctx, c.incoming, v) +} + func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.Peer, error) { v, err := source.Get(ctx, id) if err != nil { return nil, err } + return p.live(ctx, source, v) +} + +// live resolves a cached value to an open peer. A closed one is evicted first: +// the entry must be gone before looking again or the cache would return the +// same instance, so this waits (bounded by ctx) for the background discard, +// and a teardown that blocks never runs on this path. +func (p *pool) live(ctx context.Context, source ocache.OCache, v ocache.Object) (peer.Peer, error) { pr, err := getPeer(v) if err != nil { return nil, err @@ -303,9 +346,6 @@ func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.P if !pr.IsClosed() { return pr, nil } - // The entry must be gone before looking again or source.Get would return - // the same instance: wait (bounded by ctx) for the background discard, so - // a teardown that blocks never runs on this path. select { case <-p.discard(source, pr): case <-ctx.Done(): @@ -326,12 +366,12 @@ func (p *pool) get(ctx context.Context, source ocache.OCache, id string) (peer.P // skip those events (see evictOnClose); a peer a GC TryClose holds at that // moment is reported by its watcher once it closes, and once pool shutdown // has begun the events are suppressed like the watchers' (the peers close -// anyway). Cached dial errors go with the old pair, except -// incompatible-version verdicts (see errObject), which are carried over so -// their backoff survives (one still loading at the swap is not: its verdict -// lands in the old pair and the fresh one redials once). A no-op once the -// pool is closed; concurrent flushes are serialized, so each replaced pair is -// closed exactly once. +// anyway). Cached dial verdicts (today only incompatible-version ones, see +// the loader) are carried over so their backoff survives; one still loading +// at the swap is not, and the fresh pair redials once. A no-op once the pool +// is closed. Concurrent flushes are safe: swapMu serializes the swaps, and +// each replaced pair is walked, reported and closed exactly once, by the +// Flush that replaced it. func (p *pool) Flush(ctx context.Context) error { p.swapMu.Lock() if p.closed { @@ -340,8 +380,10 @@ func (p *pool) Flush(ctx context.Context) error { } old := p.current.Load() fresh := p.newCaches() + // the verdicts must be in the fresh pair before it is published, so this + // one read of the old outgoing cache happens before the swap old.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { - if eo, ok := v.(*errObject); ok && eo.keepOnFlush() { + if eo, ok := v.(*errObject); ok { // cheap and non-blocking: a fresh cache has no closers _ = fresh.outgoing.Add(eo.id, eo) } @@ -353,33 +395,17 @@ func (p *pool) Flush(ctx context.Context) error { // can follow its Wait p.closing.Add(1) p.swapMu.Unlock() - // Snapshot and mark under reportedMu: a watcher that evicts one of these - // peers concurrently checks the mark after its removal, which the cache - // lock orders against this snapshot, so each instance is reported once. - // The events themselves go out without any lock held. - type flushedPeer struct { - pr peer.Peer - inbound bool - } - var flushed []flushedPeer - old.reportedMu.Lock() - old.reported = map[peer.Peer]struct{}{} - for _, inbound := range []bool{true, false} { - old.cache(inbound).ForEach(func(v ocache.Object) (isContinue bool) { - if pr, ok := v.(peer.Peer); ok { - old.reported[pr] = struct{}{} - flushed = append(flushed, flushedPeer{pr: pr, inbound: inbound}) - } - return true - }) - } - old.reportedMu.Unlock() + // The walk over the old pair comes after the swap: addIncoming adds under + // the read lock, so every incoming peer of the old pair is visible now, + // and only this Flush (the one that replaced the pair) walks it, so the + // marks are set exactly once. It serves the marking, the Closed events + // and the parallel pre-close alike. + peers := old.snapshot(true) go func() { defer p.closing.Done() - peers, _ := closeCaches(old) - peers.Wait() + _ = closeCaches(old, peers) }() - for _, f := range flushed { + for _, f := range peers { if p.closingCtx.Err() != nil { // a Close that began meanwhile: Closed is suppressed from here on break @@ -393,39 +419,64 @@ func (p *pool) Flush(ctx context.Context) error { return nil } -// closeCaches tears down a pair that is no longer current. Each loaded peer is -// closed on its own goroutine first, because ocache.Close closes entries one -// at a time with no ctx and one hung teardown would hold the rest back; the -// two caches are then closed concurrently (each Close cancels its in-flight -// loads and closes its peers a second time, which peer.Close tolerates; the -// pool relies on that already, see discard), so a hung peer in one cache -// never delays the other. A RemoveSame per peer would not do: once a cache -// is marked closed every RemoveSame is refused. Known gap: a peer a GC -// TryClose holds past closeTimeout closes only when TryClose returns (ocache -// escalates the decline), never if it never returns. Returns once both caches -// are closed; the WaitGroup tracks the per-peer closes still running, and err -// is the outgoing cache's close error. -func closeCaches(c *caches) (peers *sync.WaitGroup, err error) { - peers = &sync.WaitGroup{} - for _, cache := range []ocache.OCache{c.outgoing, c.incoming} { - cache.ForEach(func(v ocache.Object) (isContinue bool) { +type snapshotPeer struct { + pr peer.Peer + inbound bool +} + +// snapshot lists the loaded peers of both caches (not the ones another closer +// holds). With mark, every peer found is recorded as reported by Flush (see +// evictOnClose) under reportedMu, which stays held across the walk so the +// marks are complete by the time a watcher reads them. +func (c *caches) snapshot(mark bool) (peers []snapshotPeer) { + if mark { + c.reportedMu.Lock() + defer c.reportedMu.Unlock() + c.reported = map[peer.Peer]struct{}{} + } + for _, inbound := range []bool{true, false} { + c.cache(inbound).ForEach(func(v ocache.Object) (isContinue bool) { if pr, ok := v.(peer.Peer); ok { - peers.Add(1) - go func() { - defer peers.Done() - _ = pr.Close() - }() + if mark { + c.reported[pr] = struct{}{} + } + peers = append(peers, snapshotPeer{pr: pr, inbound: inbound}) } return true }) } + return peers +} + +// closeCaches tears down a pair that is no longer current and returns once +// its caches and the given peers are closed. Each peer is closed on its own +// goroutine first, because ocache.Close closes entries one at a time with no +// ctx and one hung teardown would hold the rest back; the two caches are then +// closed concurrently (each Close cancels its in-flight loads and closes its +// peers a second time, which peer.Close tolerates; the pool relies on that +// already, see discard), so a hung peer in one cache never delays the other. +// A RemoveSame per peer would not do: once a cache is marked closed every +// RemoveSame is refused. Known gap: a peer a GC TryClose holds past +// closeTimeout closes only when TryClose returns (ocache escalates the +// decline), never if it never returns. err is the outgoing cache's close +// error. +func closeCaches(c *caches, peers []snapshotPeer) (err error) { + var wg sync.WaitGroup + for _, sp := range peers { + wg.Add(1) + go func() { + defer wg.Done() + _ = sp.pr.Close() + }() + } incomingClosed := make(chan error, 1) go func() { incomingClosed <- c.incoming.Close() }() err = c.outgoing.Close() if e := <-incomingClosed; e != nil { log.Warn("close incoming cache error", zap.Error(e)) } - return peers, err + wg.Wait() + return err } func (p *pool) getIfActive(ctx context.Context, peerIds []string) peer.Peer { @@ -484,13 +535,14 @@ func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error // the pool evicts it by instance. func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // Bounds the passes over an entry for the same id that is still there - // after this call dealt with it: one whose close another closer holds - // (waited for below), or one another AddPeer keeps replacing. Then + // after this call dealt with it: one another closer holds (its close is + // waited for below), or one another AddPeer keeps replacing. Then // ErrExists, as before. A swap does not count: the add simply moves to // the current pair, as many times as flushes come (in practice the caller // coalesces its flushes). const retries = 3 - for attempt := 0; ; attempt++ { + attempts := 0 + for { if p.closingCtx.Err() != nil { // shutting down: nothing is evicted any more (discard is a no-op), // so there is nothing to retry towards @@ -513,20 +565,19 @@ func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { if e == ocache.ErrClosed { return e } - if attempt == retries { + if attempts++; attempts > retries { return ocache.ErrExists } if e != nil { if err = ctx.Err(); err != nil { return err } - // The entry is transient: a loading one a concurrent Get(id) - // created (Pick waited it out; the incoming loader fails it) or - // the previous connection mid-close. Pick does not wait for a - // close, so wait here: a remote that reconnects while its old - // connection is still being torn down must get in once that is - // done, not be refused within microseconds. Bounded by ctx and by - // the pair staying current. + // The entry is mid-close (nothing else is ever in this cache + // unloaded: it is filled by AddPeer alone, never by a load). Pick + // does not wait for a close, so wait here: a remote that + // reconnects while its old connection is still being torn down + // must get in once that is done, not be refused within + // microseconds. Bounded by ctx and by the pair staying current. if err = p.waitClosing(ctx, c, pr.Id()); err != nil && c.ctx.Err() == nil { return err } @@ -599,7 +650,7 @@ func (p *pool) pick(ctx context.Context, source ocache.OCache, id string) (peer. if !pr.IsClosed() { return pr, nil } - p.discard(source, pr) + // a closed peer is on its way out: its watcher evicts it return nil, errPeerNotFound } diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go index 4cd2387b4..639f6d529 100644 --- a/net/pool/pool_flush_test.go +++ b/net/pool/pool_flush_test.go @@ -615,24 +615,25 @@ func TestPool_FlushSwap(t *testing.T) { require.NoError(t, err) require.Same(t, eo, v) }) - t.Run("cached dial error is dropped by flush", func(t *testing.T) { + t.Run("every cached verdict is carried over by flush", func(t *testing.T) { + // the loader caches incompatible-version verdicts only, so Flush + // carries whatever errObject it finds, without inspecting it fx := newFixture(t) defer fx.Finish() p := fx.Service.(*poolService).pool - require.NoError(t, p.current.Load().outgoing.Add("p1", &errObject{id: "p1", err: assert.AnError, createdTime: atomic2.NewTime(time.Now())})) - _, err := fx.Pick(ctx, "p1") - require.ErrorIs(t, err, assert.AnError) - + eo := &errObject{id: "p1", err: assert.AnError, createdTime: atomic2.NewTime(time.Now())} + require.NoError(t, p.current.Load().outgoing.Add("p1", eo)) require.NoError(t, fx.Flush(ctx)) - require.Equal(t, 0, p.current.Load().outgoing.Len()) - fresh := newTestPeer("p1") fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { - return fresh, nil + t.Error("must not redial a peer with a cached verdict") + return nil, nil } - pr, err := fx.Get(ctx, "p1") + _, err := fx.Get(ctx, "p1") + require.ErrorIs(t, err, assert.AnError) + v, err := p.current.Load().outgoing.Pick(ctx, "p1") require.NoError(t, err) - assert.Equal(t, peer.Peer(fresh), pr) + require.Same(t, eo, v) }) t.Run("incoming replacement via AddPeer after flush", func(t *testing.T) { obs := &poolEventRecorder{} @@ -1148,41 +1149,32 @@ func TestPool_FlushSwap(t *testing.T) { require.True(t, inCurrent(p, repl)) require.False(t, repl.IsClosed()) }) - t.Run("add racing a Get on the same id succeeds", func(t *testing.T) { + t.Run("get never loads into the incoming cache", func(t *testing.T) { + // incoming entries come from AddPeer alone: a Get that finds none + // must not leave a loading entry behind for an AddPeer to trip over fx := newFixture(t) defer fx.Finish() p := fx.Service.(*poolService).pool - // a pair whose incoming loader can be held open, standing in for the - // instant ErrNotExists one: the Get leaves a loading entry that Add - // trips over and Pick waits out - gate := make(chan struct{}) - entered := make(chan struct{}) - var once sync.Once - orig := p.newCaches - p.newCaches = func() *caches { - c := orig() - _ = c.incoming.Close() - c.incoming = ocache.New(func(ctx context.Context, id string) (ocache.Object, error) { - once.Do(func() { close(entered) }) - <-gate - return nil, ocache.ErrNotExists - }, ocache.WithGCPeriod(0)) - c.peekIncoming = mustPeeker(c.incoming) - return c - } - require.NoError(t, fx.Flush(ctx)) + var loads atomic.Int32 + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onGet: func() { loads.Add(1) }} + }) fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { return newTestPeer(peerId), nil } - go func() { _, _ = fx.Get(ctx, "d") }() - <-entered - tp := newTestPeer("d") - added := make(chan error, 1) - go func() { added <- fx.AddPeer(ctx, tp) }() - require.Never(t, func() bool { return len(added) > 0 }, 50*time.Millisecond, 5*time.Millisecond) - close(gate) - require.NoError(t, <-added) - require.True(t, inCurrent(p, tp)) + for i := 0; i < 50; i++ { + id := fmt.Sprintf("p%d", i) + tp := newTestPeer(id) + done := make(chan struct{}) + go func() { + _, _ = fx.Get(ctx, id) + close(done) + }() + require.NoError(t, fx.AddPeer(ctx, tp)) + <-done + require.True(t, inCurrent(p, tp)) + } + require.Zero(t, loads.Load(), "Get loaded into the incoming cache") }) t.Run("close waits for the parallel peer closes of every pair", func(t *testing.T) { for _, flushed := range []bool{true, false} { @@ -1600,6 +1592,342 @@ func TestPool_FlushSwap(t *testing.T) { } require.Zero(t, dials.Load(), "dialed with a done ctx after the eviction") }) + t.Run("swaps do not use up the add retries", func(t *testing.T) { + for _, swaps := range []int{3, 5} { + fx := newFixture(t) + p := fx.Service.(*poolService).pool + // every pair the pool builds already holds a connection for "d", + // and each Pick of the duplicate triggers a flush, swaps times + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + require.NoError(t, c.incoming.Add("d", newTestPeer("d"))) + return c + } + require.NoError(t, fx.Flush(ctx)) + repl := newCtlPeer("d") + var flushes atomic.Int32 + repl.idHook = func(call int32) { + // calls 2, 4, 6...: the Pick after each ErrExists + if call%2 == 0 && int(flushes.Load()) < swaps { + flushes.Add(1) + assert.NoError(t, fx.Flush(ctx)) + } + } + got := make(chan error, 1) + go func() { got <- fx.AddPeer(context.Background(), repl) }() + select { + case err := <-got: + require.NoError(t, err, "swaps=%d", swaps) + case <-time.After(3 * time.Second): + t.Fatalf("swaps=%d: AddPeer did not return", swaps) + } + require.Equal(t, int32(swaps), flushes.Load()) + require.True(t, inCurrent(p, repl)) + require.False(t, repl.IsClosed()) + fx.Finish() + } + }) + t.Run("a replacement storm for one id ends with ErrExists", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + // every eviction of the duplicate is followed by another connection + // taking its place, as two remotes reconnecting in lockstep would + var replaced atomic.Int32 + var inner ocache.OCache + installIncoming(t, fx, func(in ocache.OCache, peek ocache.Peeker) ocache.OCache { + inner = in + return &hookedCache{OCache: in, peek: peek} + }) + p := fx.Service.(*poolService).pool + hc := p.current.Load().incoming.(*hookedCache) + hc.afterRemoveSame = func() { + // the removal has landed: a new duplicate takes the id at once + if inner.Add("d", newTestPeer("d")) == nil { + replaced.Add(1) + } + } + require.NoError(t, inner.Add("d", newTestPeer("d"))) + repl := newCtlPeer("d") + got := make(chan error, 1) + go func() { got <- fx.AddPeer(context.Background(), repl) }() + select { + case err := <-got: + require.ErrorIs(t, err, ocache.ErrExists) + case <-time.After(5 * time.Second): + t.Fatal("AddPeer did not terminate") + } + // the first pass and two retries each evicted one duplicate; the + // fourth pass hit the bound before evicting another + require.Equal(t, int32(3), replaced.Load()) + }) + t.Run("get finding a closed peer during shutdown returns ErrClosed at once", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + dead := newTestPeer("p1") + require.NoError(t, dead.Close()) + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + // shutdown has begun but the pair is not cancelled yet: nothing is + // evicted any more, so the closed peer stays where it is + p.closingCancel() + got := make(chan error, 1) + go func() { + _, err := fx.Get(ctx, "p1") + got <- err + }() + select { + case err := <-got: + require.ErrorIs(t, err, ocache.ErrClosed) + case <-time.After(2 * time.Second): + t.Fatal("Get spun on the closed peer") + } + require.Zero(t, dials.Load()) + }) + t.Run("pick leaves the eviction of a closed peer to its watcher", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // a closed peer without a watcher: nothing but a discard could close + // or remove it + dead := newCtlPeer("p1") + close(dead.closed) + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + for i := 0; i < 10; i++ { + _, err := fx.Pick(ctx, "p1") + require.Error(t, err) + require.Nil(t, p.getIfActive(ctx, []string{"p1"})) + } + require.Never(t, func() bool { return dead.closeCalls.Load() > 0 }, 100*time.Millisecond, 10*time.Millisecond) + require.Equal(t, 1, p.current.Load().outgoing.Len()) + }) + t.Run("concurrent flushes report every peer exactly once before they return", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + // flush A is held inside its walk of the pair it replaced while flush + // B replaces the pair A published and a peer is added in between + inWalk := make(chan struct{}) + releaseWalk, doReleaseWalk := newRelease() + defer doReleaseWalk() + // a CAS, not sync.Once: Once would park B's walk behind A's paused one + var armed atomic.Bool + armed.Store(true) + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onForEach: func() { + if armed.CompareAndSwap(true, false) { + close(inWalk) + <-releaseWalk + } + }} + }) + x := newTestPeer("x") + require.NoError(t, fx.AddPeer(ctx, x)) + aDone := make(chan struct{}) + go func() { + assert.NoError(t, fx.Flush(ctx)) + close(aDone) + }() + <-inWalk + // A has swapped and is walking the old pair; y lands in A's fresh pair + y := newTestPeer("y") + require.NoError(t, fx.AddPeer(ctx, y)) + require.NoError(t, fx.Flush(ctx)) + // B returned: y was reported by B, before A even finished its walk + require.Equal(t, 1, len(obs.kindsFor("y"))) + require.Empty(t, obs.kindsFor("x")) + doReleaseWalk() + <-aDone + require.Equal(t, 1, len(obs.kindsFor("x"))) + require.Eventually(t, func() bool { return x.IsClosed() && y.IsClosed() }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 2 }, 200*time.Millisecond, 10*time.Millisecond) + }) + t.Run("a peer added while flush waits for the swap is reported before flush returns", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + x := newCtlPeer("x") + inAdd := make(chan struct{}) + releaseAdd, doReleaseAdd := newRelease() + defer doReleaseAdd() + x.idHook = func(call int32) { + if call == 1 { + close(inAdd) + <-releaseAdd + } + } + added := make(chan error, 1) + go func() { added <- fx.AddPeer(ctx, x) }() + <-inAdd + // the add holds the read lock: Flush queues behind it, so x lands in + // the pair Flush replaces and must be in its walk + flushed := make(chan struct{}) + go func() { + assert.NoError(t, fx.Flush(ctx)) + close(flushed) + }() + time.Sleep(20 * time.Millisecond) + doReleaseAdd() + require.NoError(t, <-added) + <-flushed + require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("x")) + require.Eventually(t, x.IsClosed, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 1 }, 200*time.Millisecond, 10*time.Millisecond) + }) + t.Run("a probe of a pair flushed under it counts no incoming miss", func(t *testing.T) { + reg := prometheus.NewRegistry() + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + a.Register(&testMetric{reg: reg}) + }) + defer fx.Finish() + // the pair is replaced while the Get probes its incoming cache: the + // probe finds nothing, but that says nothing about the peer and a Get + // on the cache would have failed uncounted + // the first peek is the fast path's (it counts nothing itself); the + // second is the lookup's probe, which the flush lands under + var peeks atomic.Int32 + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onPeek: func() { + if peeks.Add(1) == 2 { + assert.NoError(t, fx.Flush(ctx)) + } + }} + }) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + families, err := reg.Gather() + require.NoError(t, err) + for _, mf := range families { + if mf.GetName() == "netpool_incoming_miss" { + // the fast path's peek and the probe of the replaced pair count + // nothing; the retry on the fresh pair counts the one miss + require.Equal(t, float64(1), mf.GetMetric()[0].GetCounter().GetValue()) + } + } + }) + t.Run("a flush landing during an eviction never dials with the dead ctx", func(t *testing.T) { + // the swap cancels the lookup's ctx while its eviction completes; the + // redial loop must notice before it looks again, whichever of the two + // the select picked, or the old pair would be dialed with a dead ctx + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + var deadDials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if ctx.Err() != nil { + deadDials.Add(1) + return nil, ctx.Err() + } + return newTestPeer(peerId), nil + } + for i := 0; i < 40; i++ { + id := fmt.Sprintf("p%d", i) + dead := newCtlPeer(id) + close(dead.closed) + dead.closeHook = func(int32) { assert.NoError(t, fx.Flush(ctx)) } + require.NoError(t, p.current.Load().outgoing.Add(id, dead)) + pr, err := fx.Get(ctx, id) + require.NoError(t, err) + require.False(t, pr.IsClosed()) + } + require.Zero(t, deadDials.Load(), "dialed with a cancelled ctx") + }) + t.Run("a live incoming arriving while a dead incoming is evicted is used instead of a dial", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // the replacement lands the moment the dead peer's removal completes, + // before the Get looks again: it must find it instead of dialing + live := newTestPeer("p1") + var inner ocache.OCache + installIncoming(t, fx, func(in ocache.OCache, peek ocache.Peeker) ocache.OCache { + inner = in + return &hookedCache{OCache: in, peek: peek, afterRemoveSame: func() { + _ = in.Add("p1", live) + }} + }) + dead := newTestPeer("p1") + require.NoError(t, dead.Close()) + // a dead incoming peer without a watcher + require.NoError(t, inner.Add("p1", dead)) + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.Same(t, live, pr) + require.Zero(t, dials.Load(), "dialed although an incoming connection arrived") + require.True(t, inCurrent(p, live)) + }) + t.Run("the incoming probe refreshes the deadline for Get and not for Pick", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + var mu sync.Mutex + var touches []bool + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &peekRecorder{OCache: inner, peek: peek, record: func(touch bool) { + mu.Lock() + touches = append(touches, touch) + mu.Unlock() + }} + }) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newTestPeer(peerId), nil + } + _, err := fx.Get(ctx, "out") + require.NoError(t, err) + // the fast path's peek and the lookup's probe, both Get-like + require.Equal(t, []bool{true, true}, touches) + touches = nil + _, err = fx.Pick(ctx, "out") + require.NoError(t, err) + require.Equal(t, []bool{false}, touches) + }) + t.Run("an incoming peer found by the probe counts a hit", func(t *testing.T) { + reg := prometheus.NewRegistry() + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + a.Register(&testMetric{reg: reg}) + }) + defer fx.Finish() + // the peer connects between the fast path's miss and the probe + tp := newTestPeer("in") + var peeks atomic.Int32 + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onPeek: func() { + if peeks.Add(1) == 2 { + assert.NoError(t, fx.AddPeer(ctx, tp)) + } + }} + }) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + t.Error("must not dial: an incoming connection exists") + return nil, nil + } + pr, err := fx.Get(ctx, "in") + require.NoError(t, err) + require.Same(t, tp, pr) + families, err := reg.Gather() + require.NoError(t, err) + values := map[string]float64{} + for _, mf := range families { + if m := mf.GetMetric()[0]; m.GetCounter() != nil { + values[mf.GetName()] = m.GetCounter().GetValue() + } + } + assert.Equal(t, float64(1), values["netpool_incoming_hit"]) + assert.Equal(t, float64(0), values["netpool_incoming_miss"]) + assert.Equal(t, float64(0), values["netpool_outgoing_hit"]) + }) t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { obs := &poolEventRecorder{} fx := newFixtureWithObserver(t, obs) @@ -2003,16 +2331,30 @@ func (c *stickyCache) WaitClosing(ctx context.Context, id string) error { // intercepted: the seams for racing a Flush against a watcher's eviction type hookedCache struct { ocache.OCache - peek ocache.Peeker - onRemoveSame func() - onForEach func() + peek ocache.Peeker + onRemoveSame func() + afterRemoveSame func() + onForEach func() + onGet func() + onPeek func() } func (c *hookedCache) RemoveSame(ctx context.Context, id string, value ocache.Object) (bool, error) { if c.onRemoveSame != nil { c.onRemoveSame() } - return c.OCache.RemoveSame(ctx, id, value) + ok, err := c.OCache.RemoveSame(ctx, id, value) + if c.afterRemoveSame != nil { + c.afterRemoveSame() + } + return ok, err +} + +func (c *hookedCache) Get(ctx context.Context, id string) (ocache.Object, error) { + if c.onGet != nil { + c.onGet() + } + return c.OCache.Get(ctx, id) } func (c *hookedCache) ForEach(f func(v ocache.Object) bool) { @@ -2023,6 +2365,9 @@ func (c *hookedCache) ForEach(f func(v ocache.Object) bool) { } func (c *hookedCache) Peek(id string, touch bool) (ocache.Object, bool) { + if c.onPeek != nil { + c.onPeek() + } return c.peek.Peek(id, touch) } @@ -2043,3 +2388,19 @@ func installIncoming(t *testing.T, fx *fixture, wrap func(inner ocache.OCache, p } require.NoError(t, fx.Flush(ctx)) } + +// peekRecorder is an incoming cache that records the touch flag of every Peek +type peekRecorder struct { + ocache.OCache + peek ocache.Peeker + record func(touch bool) +} + +func (c *peekRecorder) Peek(id string, touch bool) (ocache.Object, bool) { + c.record(touch) + return c.peek.Peek(id, touch) +} + +func (c *peekRecorder) WaitClosing(ctx context.Context, id string) error { + return c.peek.WaitClosing(ctx, id) +} diff --git a/net/pool/poolservice.go b/net/pool/poolservice.go index 62a0936ab..fce209510 100644 --- a/net/pool/poolservice.go +++ b/net/pool/poolservice.go @@ -59,8 +59,9 @@ func (p *poolService) Init(a *app.App) (err error) { // once here and shared by every instance (WithPrometheus would register // the same names again and panic). The names stay // netpool_{outgoing,incoming}_{hit,miss,gc,size}; size reads the current - // cache, so it is registered only once a pair is published. ocache skips a - // nil option. + // pair, so it is registered only once one is published. The hit path + // reads the caches through Peek, which counts nothing, and counts for + // itself through fastMetrics. ocache skips a nil option. var outgoing, incoming ocache.PrometheusCollectors var outgoingMetrics, incomingMetrics ocache.Option if p.metricReg != nil { @@ -167,8 +168,7 @@ func (p *pool) Close(ctx context.Context) (err error) { cur.cancel() done := make(chan error, 1) go func() { - peers, err := closeCaches(cur) - peers.Wait() + err := closeCaches(cur, cur.snapshot(false)) p.closing.Wait() done <- err }() @@ -184,6 +184,9 @@ func (p *pool) Close(ctx context.Context) (err error) { return err } +// errObject is a cached dial verdict. The loader stores only +// incompatible-version ones: they say nothing about the network, so Flush +// carries them over and their 20-minute backoff survives a recovery. type errObject struct { id string err error @@ -194,15 +197,6 @@ func (e *errObject) Error() error { return e.err } -// keepOnFlush reports whether Flush carries this cached error over into the -// fresh cache. An incompatible-version verdict survives: it says nothing -// about the network, and dropping it would defeat its 20-minute backoff on -// every recovery. Only published verdicts are carried; one whose dial is -// still in flight at the swap stays with the old pair (one extra dial). -func (e *errObject) keepOnFlush() bool { - return errors.Is(e.err, handshake.ErrIncompatibleVersion) -} - func (e *errObject) Close() (err error) { return } diff --git a/net/transport/yamux/conn.go b/net/transport/yamux/conn.go index cb4b8a36e..e698c1aa4 100644 --- a/net/transport/yamux/conn.go +++ b/net/transport/yamux/conn.go @@ -138,13 +138,6 @@ func (y *yamuxConn) waitBacklog(ctx context.Context) error { } } -// abandoned returns the number of Open helpers left behind by their callers -func (y *yamuxConn) abandoned() int { - y.backlogMu.Lock() - defer y.backlogMu.Unlock() - return y.abandonedOpens -} - func (y *yamuxConn) LastUsage() time.Time { return y.luConn.LastUsage() } @@ -200,16 +193,21 @@ func (s yamuxStream) Write(b []byte) (n int, err error) { } // wrapSessionDead wraps err with transport.NewConnClosedError when it was -// caused by the session shutting down. io.EOF, a stream reset and a closed stream count -// only while the session is closed: on a live session they are stream-level -// outcomes (a remote close or reset) and are returned unchanged. +// caused by the session dying. ErrSessionShutdown always is. io.EOF, a stream +// reset, a closed stream and a connection write timeout count only while the +// session is closed: on a live session they are stream-level outcomes (a +// remote close or reset, or a send that waited out ConnectionWriteTimeout on a +// slow but live peer, which yamux does not treat as fatal) and are returned +// unchanged. A truly stalled session ends through missed keepalives, after +// which its streams fail with errors covered here. func (s yamuxStream) wrapSessionDead(err error) error { if err == nil { return nil } switch { case errors.Is(err, yamux.ErrSessionShutdown): - case errors.Is(err, io.EOF), errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed): + case errors.Is(err, io.EOF), errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed), + errors.Is(err, yamux.ErrConnectionWriteTimeout): if !s.sess.IsClosed() { return err } diff --git a/net/transport/yamux/conn_test.go b/net/transport/yamux/conn_test.go index 55cd71c60..b8b524cf5 100644 --- a/net/transport/yamux/conn_test.go +++ b/net/transport/yamux/conn_test.go @@ -2,6 +2,7 @@ package yamux import ( "context" + "errors" "io" "net" "testing" @@ -306,3 +307,67 @@ func TestYamuxStream_SessionDeathNormalized(t *testing.T) { assert.False(t, mc.Session.IsClosed()) }) } + +func TestYamuxStream_SessionFatalErrorsNormalized(t *testing.T) { + mc, _ := newSessionPair(t) + s := yamuxStream{sess: mc.Session} + err := s.wrapSessionDead(yamux.ErrSessionShutdown) + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.ErrorIs(t, err, yamux.ErrSessionShutdown, "the original error stays reachable") + // stream-level outcomes on a live session pass through unchanged + for _, cause := range []error{io.EOF, yamux.ErrConnectionReset, yamux.ErrStreamClosed, yamux.ErrConnectionWriteTimeout, yamux.ErrTimeout} { + assert.Equal(t, cause, s.wrapSessionDead(cause)) + } + // once the session is closed they mean it died + require.NoError(t, mc.Session.Close()) + for _, cause := range []error{io.EOF, yamux.ErrConnectionReset, yamux.ErrStreamClosed, yamux.ErrConnectionWriteTimeout} { + err := s.wrapSessionDead(cause) + assert.ErrorIs(t, err, transport.ErrConnClosed, cause.Error()) + assert.ErrorIs(t, err, cause) + } +} + +// slowReader reads slowly, so the writer's sends queue up +type slowReader struct{ net.Conn } + +func (s *slowReader) Read(b []byte) (int, error) { + time.Sleep(10 * time.Millisecond) + if len(b) > 4096 { + b = b[:4096] + } + return s.Conn.Read(b) +} + +// TestYamuxStream_WriteTimeoutOnLiveSession: a write that waits out +// ConnectionWriteTimeout on a slow but live peer is not a dead connection +func TestYamuxStream_WriteTimeoutOnLiveSession(t *testing.T) { + c1, c2 := net.Pipe() + cfg := yamux.DefaultConfig() + cfg.ConnectionWriteTimeout = 200 * time.Millisecond + cfg.EnableKeepAlive = false + cfg.LogOutput = io.Discard + client, err := yamux.Client(c1, cfg) + require.NoError(t, err) + server, err := yamux.Server(&slowReader{c2}, cfg) + require.NoError(t, err) + t.Cleanup(func() { + _ = client.Close() + _ = server.Close() + }) + go func() { + for { + st, aErr := server.Accept() + if aErr != nil { + return + } + go func() { _, _ = io.Copy(io.Discard, st) }() + } + }() + mc := NewMultiConn(context.Background(), connutil.NewLastUsageConn(c1), "pipe", client) + st, err := mc.Open(ctx) + require.NoError(t, err) + _, err = st.Write(make([]byte, 256*1024)) + require.ErrorIs(t, err, yamux.ErrConnectionWriteTimeout) + assert.False(t, errors.Is(err, transport.ErrConnClosed), "a slow live peer is not a dead connection") + assert.False(t, client.IsClosed()) +} diff --git a/net/transport/yamux/export_test.go b/net/transport/yamux/export_test.go new file mode 100644 index 000000000..31cc28c73 --- /dev/null +++ b/net/transport/yamux/export_test.go @@ -0,0 +1,8 @@ +package yamux + +// abandoned returns the number of Open helpers left behind by their callers +func (y *yamuxConn) abandoned() int { + y.backlogMu.Lock() + defer y.backlogMu.Unlock() + return y.abandonedOpens +} From d7e424bbb56110a4d26fc26846a50ac267edbf69 Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 18:28:46 +0200 Subject: [PATCH 5/7] net/pool, net/peer, transport: address #802 review round 3 pool: - Get dials outgoing only on a real incoming miss; the outgoing loader refuses to dial for a replaced pair and dials under the pair's ctx - Peek reports miss/busy/hit; Get waits out a busy incoming entry (e.g. a GC TryClose that declines) instead of dialing a duplicate - no comparability requirement on peer.Peer (value-level check, by-id fallback) - per-peer close concurrency capped at 256 - invariants documented at the top of pool.go peer: - an RPC error the caller didn't cause becomes ErrConnClosed when the sub conn closed, was doomed, or the session died (drpc otherwise reports a dead yamux session as context.Canceled) - a conn handed to a waiter is claimed under the peer lock - OutgoingProtoHandshake shares the closer variant's body (sync close) transport/peerservice: - yamux reads keep plain io.EOF; Open and closed-session errors are ErrConnClosed with the original cause kept - cancellation is recorded per dial attempt; Accept bounds AddPeer at 10s --- app/ocache/ocache.go | 68 ++- app/ocache/ocache_test.go | 65 ++- net/peer/closeasync_test.go | 75 ++- net/peer/connlost_ext_test.go | 167 +++++++ net/peer/peer.go | 79 ++- net/peerobserver/peerobserver.go | 5 +- net/peerservice/dialoutcome_test.go | 19 + net/peerservice/peerservice.go | 29 +- net/peerservice/peerservice_test.go | 35 ++ net/pool/pool.go | 258 +++++++--- net/pool/pool_flush_test.go | 570 +++++++++++++++++++++- net/pool/pool_test.go | 5 +- net/pool/poolservice.go | 15 +- net/secureservice/handshake/proto.go | 36 +- net/secureservice/handshake/proto_test.go | 56 ++- net/transport/quic/conn_errors_test.go | 42 +- net/transport/yamux/conn.go | 32 +- net/transport/yamux/conn_test.go | 49 +- 18 files changed, 1407 insertions(+), 198 deletions(-) create mode 100644 net/peer/connlost_ext_test.go diff --git a/app/ocache/ocache.go b/app/ocache/ocache.go index f29bb91d6..bab7ef35a 100644 --- a/app/ocache/ocache.go +++ b/app/ocache/ocache.go @@ -133,8 +133,10 @@ type OCache interface { // RemoveSame closes and removes the object only if the value currently // stored under id is exactly the given one (pointer identity). It lets a // caller evict a specific instance it owns without racing a newer value - // that has replaced it under the same id. Returns ok=true only when this - // call performed the removal. + // that has replaced it under the same id. A value of a non-comparable + // type has no identity: for such values RemoveSame removes whatever is + // stored under id. Returns ok=true only when this call performed the + // removal. RemoveSame(ctx context.Context, id string, value Object) (ok bool, err error) // TryRemove tries to close and to remove the object. ok reports whether // this call removed it; (false, nil) means the object declined to close, @@ -251,16 +253,30 @@ func (c *oCache) Pick(ctx context.Context, id string) (value Object, err error) return val.waitLoad(ctx, id) } +// PeekState is Peek's verdict about an id. +type PeekState int + +const ( + // PeekMiss: no entry for id (or the cache is closed) + PeekMiss PeekState = iota + // PeekBusy: an entry exists but is still loading or is being closed; a + // Get or Pick would wait for it + PeekBusy + // PeekHit: a loaded value, returned + PeekHit +) + // Peeker is the non-blocking read the cache returned by New offers on top of // OCache; kept off that interface so other implementations stay valid. type Peeker interface { // Peek returns the value for id only if it is loaded and not being // closed, without loading, waiting or allocating: the hot path for // callers that handle a miss themselves. With touch a hit refreshes the - // GC deadline like Get. ok=false also for a loading entry, a closing one - // or a closed cache. Peek counts no metrics: the caller, which decides - // whether the result is used, accounts for it. - Peek(id string, touch bool) (value Object, ok bool) + // GC deadline like Get. The state tells a miss from an entry that is + // loading or closing (a caller that must not act on a false miss waits + // for the latter, see WaitClosing). Peek counts no metrics: the caller, + // which decides whether the result is used, accounts for it. + Peek(id string, touch bool) (value Object, state PeekState) // WaitClosing blocks while the entry for id is being closed, bounded by // ctx, and returns at once when there is no such entry or it is not // closing. It is the wait a caller needs before it can add a replacement @@ -268,33 +284,37 @@ type Peeker interface { WaitClosing(ctx context.Context, id string) error } -func (c *oCache) Peek(id string, touch bool) (value Object, ok bool) { +func (c *oCache) Peek(id string, touch bool) (value Object, state PeekState) { c.mu.Lock() e, exists := c.data[id] - if c.closed || !exists || e.isClosing() { + if c.closed || !exists { + c.mu.Unlock() + return nil, PeekMiss + } + if e.isClosing() { c.mu.Unlock() - return nil, false + return nil, PeekBusy } select { case <-e.load: default: // still loading c.mu.Unlock() - return nil, false + return nil, PeekBusy } // value and loadErr are written before load closes; a failed load deletes // its entry under c.mu, so a non-nil loadErr here means the entry is on // its way out if e.loadErr != nil || e.value == nil { c.mu.Unlock() - return nil, false + return nil, PeekBusy } if touch { e.lastUsage = time.Now() } value = e.value c.mu.Unlock() - return value, true + return value, PeekHit } func (c *oCache) WaitClosing(ctx context.Context, id string) error { @@ -421,7 +441,7 @@ func (c *oCache) RemoveSame(ctx context.Context, id string, value Object) (ok bo // if this call is the one that transitions it to closing. If e was already // replaced under the same id it is in a closed state and remove() is a // no-op, so a stale caller can never close the newer value that took the id. - same := exists && value != nil && e.value == value + same := exists && value != nil && sameObject(e.value, value) c.mu.Unlock() if !same { return false, ErrNotExists @@ -429,6 +449,28 @@ func (c *oCache) RemoveSame(ctx context.Context, id string, value Object) (ok bo return c.removeCtx(ctx, e) } +// sameObject reports whether stored is the very instance given. Pointer +// implementations (the usual kind) compare by identity. A value that is not +// comparable has no identity to check and must not panic the comparison: it +// is treated as the stored one, so RemoveSame degrades to Remove by id for +// such values. Checked on the values, not the types: a struct with an +// interface field is comparable as a type and still panics when that field +// holds a slice. +func sameObject(stored, given Object) bool { + if stored == nil || given == nil { + // a still-loading entry has no value yet; nothing matches it + return false + } + sv, gv := reflect.ValueOf(stored), reflect.ValueOf(given) + if sv.Type() != gv.Type() { + return false + } + if !sv.Comparable() || !gv.Comparable() { + return true + } + return stored == given +} + func (c *oCache) TryRemove(id string) (ok bool, err error) { c.mu.Lock() diff --git a/app/ocache/ocache_test.go b/app/ocache/ocache_test.go index 725d770c4..3b2adfeba 100644 --- a/app/ocache/ocache_test.go +++ b/app/ocache/ocache_test.go @@ -1410,21 +1410,21 @@ func TestOCache_Peek(t *testing.T) { c := New(func(ctx context.Context, id string) (Object, error) { return obj, nil }, WithTTL(time.Hour), WithGCPeriod(0), WithPrometheus(reg, "peek", "test")).(*oCache) - _, ok := c.Peek("a", true) - require.False(t, ok, "nothing loaded yet") + _, st := c.Peek("a", true) + require.Equal(t, PeekMiss, st, "nothing loaded yet") _, err := c.Get(ctx, "a") require.NoError(t, err) c.mu.Lock() c.data["a"].lastUsage = time.Now().Add(-time.Minute) c.mu.Unlock() - v, ok := c.Peek("a", false) - require.True(t, ok) + v, st := c.Peek("a", false) + require.Equal(t, PeekHit, st) require.Same(t, obj, v) c.mu.Lock() require.Less(t, c.data["a"].lastUsage, time.Now().Add(-30*time.Second), "Pick-like peek must not refresh the deadline") c.mu.Unlock() - v, ok = c.Peek("a", true) - require.True(t, ok) + v, st = c.Peek("a", true) + require.Equal(t, PeekHit, st) require.Same(t, obj, v) c.mu.Lock() require.Greater(t, c.data["a"].lastUsage, time.Now().Add(-time.Second), "Get-like peek refreshes the deadline") @@ -1453,10 +1453,10 @@ func TestOCache_Peek(t *testing.T) { }, WithTTL(time.Hour), WithGCPeriod(0)).(*oCache) go func() { _, _ = c.Get(ctx, "a") }() <-loading - _, ok := c.Peek("a", true) - require.False(t, ok, "loading entry") + _, st := c.Peek("a", true) + require.Equal(t, PeekBusy, st, "loading entry") close(release) - require.Eventually(t, func() bool { _, ok := c.Peek("a", false); return ok }, time.Second, time.Millisecond) + require.Eventually(t, func() bool { _, st := c.Peek("a", false); return st == PeekHit }, time.Second, time.Millisecond) // a closer holds the entry: Remove blocks in obj.Close until closeCh removed := make(chan struct{}) @@ -1464,16 +1464,53 @@ func TestOCache_Peek(t *testing.T) { _, _ = c.Remove(ctx, "a") close(removed) }() - require.Eventually(t, func() bool { _, ok := c.Peek("a", false); return !ok }, time.Second, time.Millisecond) + require.Eventually(t, func() bool { _, st := c.Peek("a", false); return st == PeekBusy }, time.Second, time.Millisecond) close(closeCh) <-removed + _, st = c.Peek("a", false) + require.Equal(t, PeekMiss, st, "removed entry") require.NoError(t, c.Add("b", NewTestObject("b", true, nil))) - v, ok := c.Peek("b", false) - require.True(t, ok) + v, st := c.Peek("b", false) + require.Equal(t, PeekHit, st) require.NotNil(t, v) require.NoError(t, c.Close()) - _, ok = c.Peek("b", false) - require.False(t, ok, "closed cache") + _, st = c.Peek("b", false) + require.Equal(t, PeekMiss, st, "closed cache") }) } + +// value types with value receivers: sliceObject is not comparable as a +// type, ifaceObject only as a value when payload holds a slice +type sliceObject struct { + tags []string +} + +func (sliceObject) Close() error { return nil } +func (sliceObject) TryClose(time.Duration) (bool, error) { return true, nil } + +type ifaceObject struct { + payload any +} + +func (ifaceObject) Close() error { return nil } +func (ifaceObject) TryClose(time.Duration) (bool, error) { return true, nil } + +func TestOCache_SameObject(t *testing.T) { + a, b := &testObject{name: "a"}, &testObject{name: "b"} + // a still-loading entry has no value: nothing matches it, and it never + // panics + require.False(t, sameObject(nil, a)) + require.False(t, sameObject(a, nil)) + // pointers compare by identity + require.True(t, sameObject(a, a)) + require.False(t, sameObject(a, b)) + // different types never match + require.False(t, sameObject(a, sliceObject{})) + // values that cannot be compared are taken as the stored one + require.True(t, sameObject(sliceObject{tags: []string{"x"}}, sliceObject{tags: []string{"y"}})) + // comparable as a type, not as a value: still no panic + require.True(t, sameObject(ifaceObject{payload: []string{"x"}}, ifaceObject{payload: []string{"y"}})) + require.False(t, sameObject(ifaceObject{payload: "x"}, ifaceObject{payload: "y"})) + require.True(t, sameObject(ifaceObject{payload: "x"}, ifaceObject{payload: "x"})) +} diff --git a/net/peer/closeasync_test.go b/net/peer/closeasync_test.go index bc4dfcb95..98c296b5d 100644 --- a/net/peer/closeasync_test.go +++ b/net/peer/closeasync_test.go @@ -102,7 +102,7 @@ func TestPeer_CloseAsync(t *testing.T) { hang := make(chan struct{}) defer close(hang) hung := newBlockingCloser(hang) - returnsWithin(t, 100*time.Millisecond, "close must never block", func() { fx.closeAsync(hung, true) }) + returnsWithin(t, time.Second, "close must never block", func() { fx.closeAsync(hung, true) }) // a hung close holds up no other close var closers []*quickCloser @@ -271,7 +271,7 @@ func TestPeer_RPCDeadlineWithBlockedClose(t *testing.T) { elapsed := time.Since(start) cancel() require.Error(t, err) - require.Less(t, elapsed, budget+500*time.Millisecond, "repetition %d: the call must return within its budget", i) + require.Less(t, elapsed, budget+2*time.Second, "repetition %d: the call must return within its budget", i) // the cancelled sub conn is never handed out again fx.mu.Lock() assert.Empty(t, fx.inactive, "repetition %d", i) @@ -474,9 +474,9 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { var hangOnce sync.Once unhang := func() { hangOnce.Do(func() { close(hang) }) } defer unhang() - // ten past the limiter threshold: the next open waits 10 slowDownSteps + // fifty past the limiter threshold: the next open waits 50 slowDownSteps var hung []*blockingCloser - for i := 0; i < fx.limiter.startThreshold+10; i++ { + for i := 0; i < fx.limiter.startThreshold+50; i++ { cl := newBlockingCloser(hang) hung = append(hung, cl) fx.closeAsync(cl, true) @@ -491,7 +491,7 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { return in, nil }).Times(1) - actx, cancel := context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) + actx, cancel := context.WithTimeout(ctx, 10*fx.limiter.slowDownStep) _, err := fx.AcquireDrpcConn(actx) cancel() require.ErrorIs(t, err, context.DeadlineExceeded, "throttled while the closes are in flight") @@ -502,7 +502,7 @@ func TestPeer_HungCloseThrottlesAcquireOnlyWhileInFlight(t *testing.T) { <-cl.closed } require.Eventually(t, func() bool { return fx.churnClosing.Load() == 0 }, time.Second, time.Millisecond) - actx, cancel = context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) + actx, cancel = context.WithTimeout(ctx, 20*fx.limiter.slowDownStep) defer cancel() _, err = fx.AcquireDrpcConn(actx) require.NoError(t, err, "no throttling once the closes are done") @@ -577,7 +577,7 @@ func TestPeer_ReleaseClosesInBackground(t *testing.T) { cctx, cancel := context.WithCancel(ctx) cancel() - returnsWithin(t, 100*time.Millisecond, "the blocked close must not run on the caller", func() { + returnsWithin(t, time.Second, "the blocked close must not run on the caller", func() { fx.ReleaseDrpcConn(cctx, sc) }) require.Equal(t, int32(1), fx.churnClosing.Load()) @@ -597,7 +597,7 @@ func TestPeer_ReleaseClosesInBackground(t *testing.T) { sc, pc := newActive(fx, release) // it waits the 200ms reuse window, then hands the close over - returnsWithin(t, 300*time.Millisecond, "the blocked close must not run on the caller", func() { + returnsWithin(t, time.Second, "the blocked close must not run on the caller", func() { fx.ReleaseDrpcConn(ctx, sc) }) require.Equal(t, int32(1), fx.churnClosing.Load()) @@ -627,7 +627,7 @@ func TestPeer_GCClosesInBackground(t *testing.T) { fx.inactive = append(fx.inactive, sc) fx.mu.Unlock() - returnsWithin(t, 100*time.Millisecond, "gc must not wait on the close", func() { + returnsWithin(t, time.Second, "gc must not wait on the close", func() { fx.gc(time.Millisecond) }) fx.mu.Lock() @@ -647,7 +647,7 @@ func TestPeer_GCClosesInBackground(t *testing.T) { fx.active[sc] = struct{}{} fx.mu.Unlock() - returnsWithin(t, 100*time.Millisecond, "gc must not wait on the close", func() { + returnsWithin(t, time.Second, "gc must not wait on the close", func() { fx.gc(time.Millisecond) }) require.True(t, sc.doomed.Load()) @@ -712,7 +712,7 @@ func TestPeer_WakeUpsDoNotStarveWaiters(t *testing.T) { results := make(chan error, waiters) for i := 0; i < waiters; i++ { go func() { - actx, cancel := context.WithTimeout(ctx, 4*time.Second) + actx, cancel := context.WithTimeout(ctx, 15*time.Second) defer cancel() _, err := fx.AcquireDrpcConn(actx) results <- err @@ -722,7 +722,7 @@ func TestPeer_WakeUpsDoNotStarveWaiters(t *testing.T) { select { case err := <-results: require.NoError(t, err, "a waiter starved") - case <-time.After(10 * time.Second): + case <-time.After(30 * time.Second): t.Fatal("waiter did not return") } } @@ -737,7 +737,7 @@ func TestPeer_GCDoesNotThrottleOpens(t *testing.T) { release := make(chan struct{}) defer close(release) // expired inactive conns whose closes hang - for i := 0; i < fx.limiter.startThreshold+10; i++ { + for i := 0; i < fx.limiter.startThreshold+50; i++ { a, b := net.Pipe() defer a.Close() defer b.Close() @@ -750,8 +750,8 @@ func TestPeer_GCDoesNotThrottleOpens(t *testing.T) { require.Empty(t, fx.inactive) fx.mu.Unlock() - // counted, these closes would hold the open for 10 slowDownSteps - actx, cancel := context.WithTimeout(ctx, 5*fx.limiter.slowDownStep) + // counted, these closes would hold the open for 50 slowDownSteps + actx, cancel := context.WithTimeout(ctx, 20*fx.limiter.slowDownStep) defer cancel() _, err := fx.AcquireDrpcConn(actx) require.NoError(t, err, "gc closes must not throttle opens") @@ -838,7 +838,7 @@ func TestPeer_ForeignConnReleaseCloseIsCounted(t *testing.T) { defer fx.finish() release := make(chan struct{}) foreign := &foreignConn{pendingConn: newPendingConn(release)} - returnsWithin(t, 100*time.Millisecond, "the blocked close must not run on the caller", func() { + returnsWithin(t, time.Second, "the blocked close must not run on the caller", func() { fx.ReleaseDrpcConn(ctx, foreign) }) require.Equal(t, int32(1), fx.churnClosing.Load()) @@ -860,3 +860,46 @@ func (c *foreignConn) NewStream(ctx context.Context, rpc string, enc drpc.Encodi func (c *foreignConn) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, in, out drpc.Message) error { return c.pendingConn.Invoke(ctx, rpc, enc, in, out) } + +// TestPeer_HandedOffConnIsNotTakenByGC: a conn handed straight to a waiter +// counts as just used, so a gc pass right after does not doom it under the +// new holder, however long it sat idle before +func TestPeer_HandedOffConnIsNotTakenByGC(t *testing.T) { + fx := newFixture(t, "p1") + defer fx.finish() + fx.mc.EXPECT().Addr().Return("").AnyTimes() + release := make(chan struct{}) + defer close(release) + for i := 0; i < fx.limiter.startThreshold+20; i++ { + fx.closeAsync(newBlockingCloser(release), true) + } + a, b := net.Pipe() + defer a.Close() + defer b.Close() + // never read or written: idle since forever + sc := &subConn{ConnUnblocked: newPendingConn(release), LastUsageConn: connutil.NewLastUsageConn(a)} + fx.mu.Lock() + fx.active[sc] = struct{}{} + fx.mu.Unlock() + + got := make(chan drpc.Conn, 1) + go func() { + dc, err := fx.AcquireDrpcConn(ctx) + assert.NoError(t, err) + got <- dc + }() + waitWaiter(t, fx) + sendWake(t, fx, sc) + select { + case dc := <-got: + require.Equal(t, drpc.Conn(sc), dc) + case <-time.After(5 * time.Second): + t.Fatal("the waiter did not take the conn") + } + fx.gc(time.Minute) + assert.False(t, sc.doomed.Load(), "gc must not take a conn its holder has just acquired") + fx.mu.Lock() + _, active := fx.active[sc] + fx.mu.Unlock() + assert.True(t, active) +} diff --git a/net/peer/connlost_ext_test.go b/net/peer/connlost_ext_test.go new file mode 100644 index 000000000..e993f5481 --- /dev/null +++ b/net/peer/connlost_ext_test.go @@ -0,0 +1,167 @@ +package peer_test + +import ( + "context" + "errors" + "io" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/anyproto/any-sync/net/peer" + "github.com/anyproto/any-sync/net/rpc" + "github.com/anyproto/any-sync/net/rpc/rpctest/multiconntest" + "github.com/anyproto/any-sync/net/secureservice/handshake/handshakeproto" + "github.com/anyproto/any-sync/net/transport" +) + +// silentCtrl serves sub conns by reading requests and never answering +type silentCtrl struct{} + +func (silentCtrl) DrpcConfig() rpc.Config { + return rpc.Config{Stream: rpc.StreamConfig{MaxMsgSizeMb: 1}} +} + +func (silentCtrl) ServeConn(ctx context.Context, conn net.Conn) error { + _, err := io.Copy(io.Discard, conn) + return err +} + +// closingCtrl ends each sub stream after its first request: a remote that +// closes the sub stream on a live session +type closingCtrl struct{} + +func (closingCtrl) DrpcConfig() rpc.Config { + return rpc.Config{Stream: rpc.StreamConfig{MaxMsgSizeMb: 1}} +} + +func (closingCtrl) ServeConn(ctx context.Context, conn net.Conn) error { + _, _ = conn.Read(make([]byte, 1)) + time.Sleep(50 * time.Millisecond) + return conn.Close() +} + +// TestPeer_RPCOnDeadYamuxSessionIsConnClosed: over real yamux, an RPC in +// flight when the session dies fails with transport.ErrConnClosed, not the +// context.Canceled drpc makes of the stream's io.EOF; a caller's own +// cancellation stays context.Canceled +func TestPeer_RPCOnDeadYamuxSessionIsConnClosed(t *testing.T) { + for _, tc := range []struct { + name string + close func(serv, client transport.MultiConn) + }{ + {"remote close", func(serv, _ transport.MultiConn) { _ = serv.Close() }}, + {"local close", func(_, client transport.MultiConn) { _ = client.Close() }}, + } { + t.Run(tc.name, func(t *testing.T) { + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, silentCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + res := make(chan error, 1) + go func() { + res <- dc.Invoke(context.Background(), "/x/y", nil, &handshakeproto.Proto{Proto: 1}, &handshakeproto.Proto{}) + }() + time.Sleep(100 * time.Millisecond) + tc.close(mcS, mcC) + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the RPC did not return") + } + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.False(t, errors.Is(err, context.Canceled), "must not look like the caller's cancellation: %v", err) + }) + } + t.Run("sub conn closed locally mid-RPC", func(t *testing.T) { + // what gc or a release does to a sub conn while an RPC runs on it + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, silentCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + res := make(chan error, 1) + go func() { + res <- dc.Invoke(context.Background(), "/x/y", nil, &handshakeproto.Proto{Proto: 1}, &handshakeproto.Proto{}) + }() + time.Sleep(100 * time.Millisecond) + go func() { _ = dc.Close() }() + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the RPC did not return") + } + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.False(t, mcC.IsClosed(), "only the sub conn closed") + }) + t.Run("remote ends only the sub stream", func(t *testing.T) { + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, closingCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + res := make(chan error, 1) + go func() { + res <- dc.Invoke(context.Background(), "/x/y", nil, &handshakeproto.Proto{Proto: 1}, &handshakeproto.Proto{}) + }() + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the RPC did not return") + } + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.False(t, errors.Is(err, context.Canceled), "must not look like the caller's cancellation: %v", err) + assert.False(t, mcC.IsClosed(), "the session is alive") + }) + t.Run("caller cancel stays canceled", func(t *testing.T) { + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, silentCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + cctx, cancel := context.WithCancel(context.Background()) + res := make(chan error, 1) + go func() { + res <- dc.Invoke(cctx, "/x/y", nil, &handshakeproto.Proto{Proto: 1}, &handshakeproto.Proto{}) + }() + time.Sleep(100 * time.Millisecond) + cancel() + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the RPC did not return") + } + assert.ErrorIs(t, err, context.Canceled) + assert.NotErrorIs(t, err, transport.ErrConnClosed) + }) +} diff --git a/net/peer/peer.go b/net/peer/peer.go index 969fd04a3..40feaa185 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -4,6 +4,7 @@ package peer import ( "context" "errors" + "fmt" "io" "net" "slices" @@ -101,16 +102,69 @@ type Peer interface { type subConn struct { encoding.ConnUnblocked *connutil.LastUsageConn + // mc is the connection the sub conn runs over; nil in some tests + mc transport.MultiConn // doomed is set by gc when it takes an active conn away and closes it in // the background: the holder must not return it for reuse, even though // the close may not have landed yet doomed atomic.Bool + // acquiredAt (unix nanos) is when a caller last took the conn; gc + // counts it as usage, so a conn that sat idle in the pool is not taken + // away from the caller that has just acquired it + acquiredAt atomic.Int64 } func (s *subConn) Unblocked() <-chan struct{} { return s.ConnUnblocked.Unblocked() } +// idleSince is the later of the last read or write and the last acquisition +func (s *subConn) idleSince() time.Time { + last := s.LastUsage() + if at := s.acquiredAt.Load(); at > last.UnixNano() { + return time.Unix(0, at) + } + return last +} + +// Invoke reports an RPC cut short by the sub conn closing as +// transport.ErrConnClosed (see connLost) +func (s *subConn) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, in, out drpc.Message) error { + return s.connLost(ctx, s.ConnUnblocked.Invoke(ctx, rpc, enc, in, out)) +} + +// NewStream reports a stream refused because the sub conn closed as +// transport.ErrConnClosed (see connLost) +func (s *subConn) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) (drpc.Stream, error) { + stream, err := s.ConnUnblocked.NewStream(ctx, rpc, enc) + return stream, s.connLost(ctx, err) +} + +// connLost reports an error the caller did not cause, returned while this sub +// conn is closed or doomed, as transport.ErrConnClosed: the RPC ended because +// the sub conn did, whether the whole connection died, the remote ended just +// this sub stream on a live session, or gc or a release closed it here. drpc +// shows such an end as context.Canceled (its stand-in for a transport +// io.EOF) or as "manager closed"; callers must be able to tell either from +// their own cancellation. context.Canceled is kept out of the error chain on +// purpose; other causes stay reachable. +func (s *subConn) connLost(ctx context.Context, err error) error { + if err == nil || ctx.Err() != nil { + return err + } + select { + case <-s.Closed(): + default: + if !s.doomed.Load() && (s.mc == nil || !s.mc.IsClosed()) { + return err + } + } + if errors.Is(err, context.Canceled) { + return transport.NewConnClosedError(fmt.Errorf("rpc ended by the sub conn closing: %v", err)) + } + return transport.NewConnClosedError(err) +} + type peer struct { id string @@ -194,7 +248,7 @@ func (p *peer) acquireDrpcConn(ctx context.Context, deadline *time.Time) (conn d return nil, false, ctx.Err() case dconn := <-p.subConnRelease: // nil conn means connection was closed, used to wake up AcquireDrpcConn - if dconn != nil && !isDoomed(dconn) { + if dconn != nil && p.claimHandoff(dconn) { return dconn, false, nil } // The released conn was closed, or gc doomed it on the way. @@ -223,11 +277,29 @@ func (p *peer) acquireDrpcConn(ctx context.Context, deadline *time.Time) (conn d } // never doomed: gc dooms active conns only, and ReleaseDrpcConn // re-checks the flag under p.mu before re-pooling + res.acquiredAt.Store(time.Now().UnixNano()) p.active[res] = struct{}{} p.mu.Unlock() return res, false, nil } +// claimHandoff takes a conn a releaser handed over directly. The check runs +// under p.mu, where gc dooms conns: one doomed on the way is refused, and the +// claimed one counts as just used, so a gc right after does not take it. +func (p *peer) claimHandoff(conn drpc.Conn) bool { + sc, ok := conn.(*subConn) + if !ok { + return true + } + p.mu.Lock() + defer p.mu.Unlock() + if sc.doomed.Load() { + return false + } + sc.acquiredAt.Store(time.Now().UnixNano()) + return true +} + func isDoomed(conn drpc.Conn) bool { sc, ok := conn.(*subConn) return ok && sc.doomed.Load() @@ -371,6 +443,7 @@ func (p *peer) openDrpcConn(ctx context.Context) (*subConn, error) { return &subConn{ ConnUnblocked: encoding.WrapConnEncoding(drpcConn, isSnappy), LastUsageConn: lastUsageConn, + mc: p.MultiConn, }, nil } @@ -532,7 +605,7 @@ func (p *peer) gc(ttl time.Duration) (aliveCount int) { hasClosed = true default: } - if in.LastUsage().Before(minLastUsage) { + if in.idleSince().Before(minLastUsage) { toClose = append(toClose, in) p.inactive[i] = nil hasClosed = true @@ -554,7 +627,7 @@ func (p *peer) gc(ttl time.Duration) (aliveCount int) { continue default: } - if act.LastUsage().Before(minLastUsage) { + if act.idleSince().Before(minLastUsage) { log.Warn("close active connection because no activity", zap.String("peerId", p.id), zap.String("addr", p.Addr())) act.doomed.Store(true) toClose = append(toClose, act) diff --git a/net/peerobserver/peerobserver.go b/net/peerobserver/peerobserver.go index 886104209..21c36615e 100644 --- a/net/peerobserver/peerobserver.go +++ b/net/peerobserver/peerobserver.go @@ -134,8 +134,9 @@ type Event struct { // connection a pool.Flush invalidates, synchronously on the goroutine that // called Flush before it returns (so for that caller those KindClosed precede // the KindConnected of its redials; the caller must not hold a lock the -// observer takes); implementations must be safe for concurrent use. Dial-path events (KindDialStarted, outbound KindConnected, -// KindDialFailed) run inside the pool's single-flight load for that peer: an +// observer takes); implementations must be safe for concurrent use. Dial-path +// events (KindDialStarted, outbound KindConnected, KindDialFailed) run inside +// the pool's single-flight load for that peer: an // implementation must never call the pool for the peer such an event names — // the load is still open and the call blocks until the caller's context dies. // A pool call for a DIFFERENT peer is allowed, but it can synchronously diff --git a/net/peerservice/dialoutcome_test.go b/net/peerservice/dialoutcome_test.go index a2fe92f9f..ad7742a3b 100644 --- a/net/peerservice/dialoutcome_test.go +++ b/net/peerservice/dialoutcome_test.go @@ -295,6 +295,25 @@ func TestPeerService_CancelledDialReportsNoOutcome(t *testing.T) { o := stub.only(t) assert.Equal(t, transport.Quic, o.SucceededScheme) }) + t.Run("a dial whose addresses all failed on their own is reported even if ctx ends after", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true) + dctx, cancel := context.WithCancel(ctx) + defer cancel() + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").Return(nil, &quicgo.HandshakeTimeoutError{}) + fx.yamux.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1111").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + // refused for real; the caller gives up as it returns + cancel() + return nil, fmt.Errorf("refused") + }) + _, err := fx.Dial(dctx, peerId) + require.Error(t, err) + o := stub.only(t) + assert.True(t, o.QuicTimedOut) + assert.True(t, o.FallbackFailed) + }) t.Run("a cancelled dial does not stop later dials from being reported", func(t *testing.T) { fx, stub := newFixtureWithStubDemotion(t) defer fx.finish(t) diff --git a/net/peerservice/peerservice.go b/net/peerservice/peerservice.go index b65deada2..e3348481d 100644 --- a/net/peerservice/peerservice.go +++ b/net/peerservice/peerservice.go @@ -61,8 +61,17 @@ type peerService struct { // afterwards, so it is read without locking; its zero value is a no-op observer peerobserver.Notifier mu sync.RWMutex + // acceptAddTimeout bounds how long Accept waits for the pool to make + // room for an incoming peer; zero means acceptAddTimeoutDefault + acceptAddTimeout time.Duration } +// acceptAddTimeoutDefault bounds Accept's pool.AddPeer: replacing an old +// incoming connection with the same peer waits for the old one's teardown, +// which a hung transport could otherwise stretch indefinitely while holding +// up the transport's accept path +const acceptAddTimeoutDefault = 10 * time.Second + func (p *peerService) Init(a *app.App) (err error) { if comp := a.Component(yamux.CName); comp != nil { p.yamux = comp.(transport.Transport) @@ -168,10 +177,13 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er // back for the fallback window on the strength of a dead context. That // drops a QuicTimedOut recorded before the cancel on purpose: without the // yamux try that the cancel cut short there is no proof the path works - // without udp, which is what a strike requires. + // without udp, which is what a strike requires. Cancellation is recorded + // where it cut an attempt short, not when the dial returns: a dial whose + // addresses all failed on their own is reported even if ctx ends after. dialAccepted := false + cancelled := false defer func() { - if p.demotion == nil || (!dialAccepted && ctx.Err() != nil) { + if p.demotion == nil || (!dialAccepted && cancelled) { return } if !dialAccepted { @@ -185,6 +197,7 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er // cancelled: the remaining addresses would only fail at once // with the same dead context err = ctx.Err() + cancelled = true break } sch := scheme(addr) @@ -194,6 +207,10 @@ func (p *peerService) Dial(ctx context.Context, peerId string) (pr peer.Peer, er break } addrErrs = append(addrErrs, err) + if cerr := ctx.Err(); cerr != nil && errors.Is(err, cerr) { + // the attempt was cut short by the caller, not by the path + cancelled = true + } // a failure under a done ctx is classified like any other: the // outcome of a cancelled dial is never reported, see above switch { @@ -354,7 +371,13 @@ func (p *peerService) Accept(mc transport.MultiConn) (err error) { Inbound: true, ProtoVersion: protoVersion, }) - if err = p.pool.AddPeer(context.Background(), pr); err != nil { + addTimeout := p.acceptAddTimeout + if addTimeout <= 0 { + addTimeout = acceptAddTimeoutDefault + } + addCtx, cancel := context.WithTimeout(context.Background(), addTimeout) + defer cancel() + if err = p.pool.AddPeer(addCtx, pr); err != nil { _ = pr.Close() p.observer.Notify(peerobserver.Event{ Kind: peerobserver.KindClosed, diff --git a/net/peerservice/peerservice_test.go b/net/peerservice/peerservice_test.go index 3ef02d297..9a27c1de9 100644 --- a/net/peerservice/peerservice_test.go +++ b/net/peerservice/peerservice_test.go @@ -295,6 +295,41 @@ func TestPeerService_Accept(t *testing.T) { require.NoError(t, fx.Accept(mc)) } +func TestPeerService_AcceptBoundedByHungOldPeer(t *testing.T) { + fx := newFixture(t) + defer fx.finish(t) + release := make(chan struct{}) + var releaseOnce sync.Once + doRelease := func() { releaseOnce.Do(func() { close(release) }) } + // released before the fixture closes the pool, which closes the old peer too + defer doRelease() + fx.PeerService.(*peerService).acceptAddTimeout = 200 * time.Millisecond + + // the old incoming connection, whose transport hangs on close + cctx := peer.CtxWithProtoVersion(peer.CtxWithPeerId(ctx, "p1"), 13) + old := mock_transport.NewMockMultiConn(fx.ctrl) + old.EXPECT().Context().Return(cctx).AnyTimes() + old.EXPECT().Addr().Return("yamux://192.0.2.7:3333").AnyTimes() + old.EXPECT().IsClosed().Return(false).AnyTimes() + old.EXPECT().Accept().Return(nil, fmt.Errorf("test")).AnyTimes() + old.EXPECT().CloseChan().Return((<-chan struct{})(nil)).AnyTimes() + old.EXPECT().Close().DoAndReturn(func() error { + <-release + return nil + }).AnyTimes() + require.NoError(t, fx.Accept(old)) + + // the same peer reconnects: Accept must not wait on the hung teardown + done := make(chan error, 1) + go func() { done <- fx.Accept(fx.mockMC("p1")) }() + select { + case err := <-done: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(5 * time.Second): + t.Fatal("Accept is held up by the old connection's teardown") + } +} + func TestPeerService_PeerObserver(t *testing.T) { // public (non-local) addrs: the global preferQuic order applies var addrs = []string{ diff --git a/net/pool/pool.go b/net/pool/pool.go index d33189198..1c3dfefc9 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "math/rand" + "reflect" "sync" "sync/atomic" "time" @@ -21,15 +22,29 @@ import ( "github.com/anyproto/any-sync/net/secureservice/handshake" ) +// The pool keeps one immutable pair of caches (incoming, outgoing) behind an +// atomic pointer. Invariants: +// - a lookup reads one pair and re-checks after the read that it is still +// current, so no result comes from a replaced pair (lookup, fast); +// - Flush swaps the pair under swapMu, cancels the old one (cutting every +// lookup on it short) and closes it in the background; AddPeer adds under +// the read side of swapMu, so a peer never lands in a replaced pair; +// - the Flush that replaced a pair walks it once afterwards: it reports the +// Closed events and marks the peers, and a peer's watcher reports only a +// peer that is not marked, so every instance is reported exactly once; +// - Close marks the pool closed under swapMu, after which Flush is a no-op +// and lookups fail with ocache.ErrClosed. + // Pool creates and caches outgoing connection type Pool interface { // Get lookups to peer in existing connections or creates and outgoing new one Get(ctx context.Context, id string) (peer.Peer, error) // GetOneOf searches at least one existing connection in outgoing or creates a new one from a randomly selected id from given list GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error) - // AddPeer adds incoming peer to the pool. The pool evicts peers by - // instance (ocache.RemoveSame), so peer.Peer implementations must be - // comparable (pointers) + // AddPeer adds incoming peer to the pool. The pool tracks peers by + // instance; an implementation whose type is not comparable (not a + // pointer) is tracked by id instead, which is only weaker when two + // connections for one id overlap AddPeer(ctx context.Context, p peer.Peer) (err error) // Pick checks if a connection with the peer exists, without dialing. // For a peer whose last dial failed it returns the cached dial error. @@ -37,9 +52,8 @@ type Pool interface { // Flush invalidates every pooled connection and every dial in flight: it // swaps in an empty pool, reports Closed for each pooled peer before it // returns and tears the old connections down in the background. Later - // lookups redial. Callers should coalesce the triggers (heart's recovery - // worker does, with a 6s window), since each Flush cancels the dials of - // the one before. + // lookups redial. Callers should coalesce their triggers, since each + // Flush cancels the dials of the one before. Flush(ctx context.Context) error } @@ -63,11 +77,37 @@ type caches struct { // is cut short and retried on the current one ctx context.Context cancel context.CancelFunc - // reported holds the peers whose Closed event Flush delivered itself, so - // their watchers do not report them a second time (see Flush); written - // once, by Flush, under reportedMu + // reported holds the peers (see peerKey) whose Closed event Flush + // delivered itself, so their watchers do not report them a second time + // (see Flush); written once, by Flush, under reportedMu reportedMu sync.Mutex - reported map[peer.Peer]struct{} + reported map[any]struct{} +} + +// peerKey identifies a pooled peer instance in a map: the peer itself when it +// is comparable (a pointer, as every implementation in this module is), else +// its id and direction, which within one pair name a single entry at a time. +// Never panics on a non-comparable implementation: checked on the value, since +// a type with an interface field is comparable while the value may not be. +func peerKey(pr peer.Peer, inbound bool) any { + if reflect.ValueOf(pr).Comparable() { + return pr + } + return struct { + id string + inbound bool + }{pr.Id(), inbound} +} + +// bind derives from ctx a context that also ends when the pair stops being +// current; release frees it and must be called once the work is done +func (c *caches) bind(ctx context.Context) (bound context.Context, release func()) { + bound, cancel := context.WithCancel(ctx) + stop := context.AfterFunc(c.ctx, cancel) + return bound, func() { + stop() + cancel() + } } // cache returns the incoming or the outgoing cache of the pair @@ -79,10 +119,10 @@ func (c *caches) cache(inbound bool) ocache.OCache { } // reportedByFlush reports whether Flush delivered this peer's Closed event -func (c *caches) reportedByFlush(pr peer.Peer) bool { +func (c *caches) reportedByFlush(pr peer.Peer, inbound bool) bool { c.reportedMu.Lock() defer c.reportedMu.Unlock() - _, ok := c.reported[pr] + _, ok := c.reported[peerKey(pr, inbound)] return ok } @@ -142,9 +182,8 @@ func (p *pool) lookup(ctx context.Context, f func(ctx context.Context, c *caches pr, err := func() (peer.Peer, error) { // deferred so a loader panic re-raised by ocache does not leave // the registration on the pair ctx behind - lctx, cancel := context.WithCancel(ctx) - defer cancel() - defer context.AfterFunc(c.ctx, cancel)() + lctx, release := c.bind(ctx) + defer release() return f(lctx, c) }() // read before re-checking current: Flush publishes the new pair before @@ -174,18 +213,30 @@ func (p *pool) lookup(ctx context.Context, f func(ctx context.Context, c *caches // (Pick). Metrics are counted only for a peer actually returned, and as the // caches would have counted them (a hit; for Get also the incoming miss // before an outgoing hit), so a fall-through to lookup counts nothing twice. -func (p *pool) fast(id string, touch bool) peer.Peer { +// missed is the pair whose caches both held nothing for id while it stayed +// current (nil otherwise): Get skips its own incoming probe once on that very +// pair (the miss is counted here, as a Get would have). +func (p *pool) fast(id string, touch bool) (pr peer.Peer, missed *caches) { c := p.current.Load() - v, inbound := c.peekIncoming.Peek(id, touch) + v, in := c.peekIncoming.Peek(id, touch) + inbound := in == ocache.PeekHit if !inbound { - var ok bool - if v, ok = c.peekOutgoing.Peek(id, touch); !ok { - return nil + var out ocache.PeekState + if v, out = c.peekOutgoing.Peek(id, touch); out != ocache.PeekHit { + // a busy entry (loading, or a close that may yet be declined) + // is not a miss: the lookup waits for it + if in != ocache.PeekMiss || out != ocache.PeekMiss || p.current.Load() != c { + return nil, nil + } + if touch && p.metrics != nil { + p.metrics.incomingMiss.Inc() + } + return nil, c } } pr, isPeer := v.(peer.Peer) if !isPeer || pr.IsClosed() || p.current.Load() != c { - return nil + return nil, nil } if m := p.metrics; m != nil { switch { @@ -198,7 +249,7 @@ func (p *pool) fast(id string, touch bool) peer.Peer { m.outgoingHit.Inc() } } - return pr + return pr, nil } // discard closes pr (if not closed yet) and evicts it from source, in the @@ -256,7 +307,7 @@ func (p *pool) evictOnClose(pr peer.Peer, c *caches, inbound bool) { // marks the peers it saw and reports under reportedMu, and the removal // and Flush's snapshot are ordered by the cache lock, so a peer Flush saw // is marked by the time this runs and one it did not see is reported here - if p.closingCtx.Err() != nil || c.reportedByFlush(pr) { + if p.closingCtx.Err() != nil || c.reportedByFlush(pr, inbound) { return } p.observer.Notify(peerobserver.Event{ @@ -267,31 +318,47 @@ func (p *pool) evictOnClose(pr peer.Peer, c *caches, inbound bool) { } func (p *pool) Get(ctx context.Context, id string) (peer.Peer, error) { - if pr := p.fast(id, true); pr != nil { + pr, missed := p.fast(id, true) + if pr != nil { return pr, nil } return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { for { - // if we have incoming connection - try to reuse it - if pr, err = p.getIncoming(ctx, c, id); err == nil { - return pr, nil + // if we have incoming connection - try to reuse it (the fast path + // just did when it found this very pair empty) + if c == missed { + missed = nil + err = ocache.ErrNotExists + } else { + pr, err = p.getIncoming(ctx, c, id) } - if err != errRedial { - // or try to get or create outgoing + switch { + case err == nil: + return pr, nil + case err == ocache.ErrNotExists: + // or try to get or create outgoing. Only on a real miss: an + // ErrClosed (the pair was replaced; lookup retries on the + // current one) or a ctx error must not start a dial, the + // former into the old pair with a ctx the swap has not + // cancelled yet (that runs through an AfterFunc, + // asynchronously), the latter with a dead one if pr, err = p.get(ctx, c.outgoing, id); err != errRedial { return pr, err } + case err != errRedial: + return nil, err } // A closed peer was evicted: start over from incoming, where a // live connection may have arrived meanwhile, before dialing. Not - // on a pair that was replaced meanwhile: the swap cancels this - // ctx through an AfterFunc, which runs asynchronously, so the pair - // is checked directly or the old cache could be dialed into with - // a ctx about to die (lookup retries on the current pair). And - // not during shutdown, when nothing is evicted any more and the - // closed peer would be found again and again. A caller ctx that - // ended was seen by the eviction's select already. - if c.ctx.Err() != nil || p.closingCtx.Err() != nil { + // with a ctx that ended meanwhile (the eviction's select picks at + // random when both are ready) and not during shutdown, when + // nothing is evicted any more and the closed peer would be found + // again and again. A replaced pair is caught by the incoming probe + // and by the loader, which never dials for one. + if err = ctx.Err(); err != nil { + return nil, err + } + if p.closingCtx.Err() != nil { return nil, ocache.ErrClosed } } @@ -307,20 +374,36 @@ var errRedial = errors.New("closed peer evicted") // the probe (and make every AddPeer for that id trip over it). Counted like a // Get on the cache would be: a hit, or a miss before the outgoing lookup. func (p *pool) getIncoming(ctx context.Context, c *caches, id string) (peer.Peer, error) { - v, ok := c.peekIncoming.Peek(id, true) - if !ok && c.ctx.Err() != nil { + v, state := c.peekIncoming.Peek(id, true) + // An entry another closer holds may be a GC TryClose the live peer + // declines: wait it out (bounded by ctx and the pair, like a Get on the + // cache would) and look again before dialing a second connection to the + // peer. Bounded: a closer that keeps coming back is treated as a miss. + for attempt := 0; state == ocache.PeekBusy && attempt < 3; attempt++ { + bctx, release := c.bind(ctx) + err := c.peekIncoming.WaitClosing(bctx, id) + release() + if err != nil { + if c.ctx.Err() != nil { + return nil, ocache.ErrClosed + } + return nil, err + } + v, state = c.peekIncoming.Peek(id, true) + } + if state != ocache.PeekHit && c.ctx.Err() != nil { // a replaced or closing pair: no verdict about the peer, and not a // miss (a Get on the cache would have failed with ErrClosed uncounted) return nil, ocache.ErrClosed } if m := p.metrics; m != nil { - if ok { + if state == ocache.PeekHit { m.incomingHit.Inc() } else { m.incomingMiss.Inc() } } - if !ok { + if state != ocache.PeekHit { return nil, ocache.ErrNotExists } return p.live(ctx, c.incoming, v) @@ -354,24 +437,19 @@ func (p *pool) live(ctx context.Context, source ocache.OCache, v ocache.Object) return nil, errRedial } -// Flush invalidates every pooled connection: it builds a fresh cache pair, -// publishes it, cancels the old pair (which cuts every dial in flight short, -// see lookup) and closes it in the background, so it never waits on a GC -// TryClose or on transport teardown. From the moment it returns no lookup -// hands out a pre-flush peer. The Closed event of every peer it found pooled -// is delivered on the calling goroutine before it returns (observers must not -// be blocked by a lock the caller holds), so for the caller's own later calls -// a Closed(X) precedes the Connected(X) of the redial; a concurrent Get or -// Accept can still produce its Connected(X) first. The peers' own watchers -// skip those events (see evictOnClose); a peer a GC TryClose holds at that -// moment is reported by its watcher once it closes, and once pool shutdown -// has begun the events are suppressed like the watchers' (the peers close -// anyway). Cached dial verdicts (today only incompatible-version ones, see -// the loader) are carried over so their backoff survives; one still loading -// at the swap is not, and the fresh pair redials once. A no-op once the pool -// is closed. Concurrent flushes are safe: swapMu serializes the swaps, and -// each replaced pair is walked, reported and closed exactly once, by the -// Flush that replaced it. +// Flush invalidates every pooled connection: it builds a fresh pair, carries +// the cached dial verdicts over (incompatible-version ones, whose backoff must +// survive; one still loading at the swap is lost and redialed once), publishes +// it, cancels the old pair, which cuts every dial in flight short, and closes +// it in the background, so it never waits on a GC TryClose or on transport +// teardown. The Closed event of every peer found pooled is delivered on the +// calling goroutine before it returns (observers must not be blocked by a lock +// the caller holds): for the caller's own later calls a Closed(X) precedes the +// Connected(X) of the redial; a concurrent Get or Accept can still produce its +// Connected first. A peer a GC TryClose holds at that moment is reported by +// its watcher once it closes; once pool shutdown has begun the events are +// suppressed like the watchers'. A no-op once the pool is closed; concurrent +// flushes are safe (see the invariants at the top of the file). func (p *pool) Flush(ctx context.Context) error { p.swapMu.Lock() if p.closed { @@ -424,6 +502,15 @@ type snapshotPeer struct { inbound bool } +// closePeer is the pre-close of one peer (named: the tests tell this close +// from the cache's own pass by it) +func closePeer(pr peer.Peer) { + _ = pr.Close() +} + +// maxPeerClosers caps how many peers closeCaches closes at once +const maxPeerClosers = 256 + // snapshot lists the loaded peers of both caches (not the ones another closer // holds). With mark, every peer found is recorded as reported by Flush (see // evictOnClose) under reportedMu, which stays held across the walk so the @@ -432,13 +519,13 @@ func (c *caches) snapshot(mark bool) (peers []snapshotPeer) { if mark { c.reportedMu.Lock() defer c.reportedMu.Unlock() - c.reported = map[peer.Peer]struct{}{} + c.reported = map[any]struct{}{} } for _, inbound := range []bool{true, false} { c.cache(inbound).ForEach(func(v ocache.Object) (isContinue bool) { if pr, ok := v.(peer.Peer); ok { if mark { - c.reported[pr] = struct{}{} + c.reported[peerKey(pr, inbound)] = struct{}{} } peers = append(peers, snapshotPeer{pr: pr, inbound: inbound}) } @@ -449,9 +536,10 @@ func (c *caches) snapshot(mark bool) (peers []snapshotPeer) { } // closeCaches tears down a pair that is no longer current and returns once -// its caches and the given peers are closed. Each peer is closed on its own -// goroutine first, because ocache.Close closes entries one at a time with no -// ctx and one hung teardown would hold the rest back; the two caches are then +// its caches and the given peers are closed. The peers are closed in +// parallel first (up to maxPeerClosers at a time), because ocache.Close +// closes entries one at a time with no ctx and one hung teardown would hold +// the rest back; the two caches are then // closed concurrently (each Close cancels its in-flight loads and closes its // peers a second time, which peer.Close tolerates; the pool relies on that // already, see discard), so a hung peer in one cache never delays the other. @@ -462,13 +550,28 @@ func (c *caches) snapshot(mark bool) (peers []snapshotPeer) { // error. func closeCaches(c *caches, peers []snapshotPeer) (err error) { var wg sync.WaitGroup - for _, sp := range peers { - wg.Add(1) - go func() { - defer wg.Done() - _ = sp.pr.Close() - }() - } + // The pre-close runs at most maxPeerClosers peers at a time: a server + // pool holds tens of thousands, and one goroutine each would be a burst + // that outlives a Close which gave up. The dispatch runs on its own + // goroutine so the cache closes below (which cancel the in-flight dials) + // are not held back behind hung peers waiting for a slot; its own wg + // count keeps Wait from returning before every peer has been dispatched. + wg.Add(1) + go func() { + defer wg.Done() + slots := make(chan struct{}, maxPeerClosers) + for _, sp := range peers { + slots <- struct{}{} + wg.Add(1) + go func() { + defer func() { + <-slots + wg.Done() + }() + closePeer(sp.pr) + }() + } + }() incomingClosed := make(chan error, 1) go func() { incomingClosed <- c.incoming.Close() }() err = c.outgoing.Close() @@ -481,7 +584,7 @@ func closeCaches(c *caches, peers []snapshotPeer) (err error) { func (p *pool) getIfActive(ctx context.Context, peerIds []string) peer.Peer { for _, peerId := range peerIds { - if pr := p.fast(peerId, false); pr != nil { + if pr, _ := p.fast(peerId, false); pr != nil { return pr } } @@ -531,8 +634,8 @@ func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error return nil, lastErr } -// AddPeer adds an incoming peer. pr must be of a comparable type (a pointer): -// the pool evicts it by instance. +// AddPeer adds an incoming peer. The pool evicts it by instance (see peerKey +// and ocache.RemoveSame for non-comparable implementations). func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // Bounds the passes over an entry for the same id that is still there // after this call dealt with it: one another closer holds (its close is @@ -600,9 +703,8 @@ func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // waitClosing waits for the incoming entry for id of pair c to finish // closing, bounded by ctx and by c staying current func (p *pool) waitClosing(ctx context.Context, c *caches, id string) error { - lctx, cancel := context.WithCancel(ctx) - defer cancel() - defer context.AfterFunc(c.ctx, cancel)() + lctx, release := c.bind(ctx) + defer release() return c.peekIncoming.WaitClosing(lctx, id) } @@ -624,7 +726,7 @@ func (p *pool) addIncoming(pr peer.Peer) (*caches, error) { } func (p *pool) Pick(ctx context.Context, id string) (pr peer.Peer, err error) { - if pr = p.fast(id, false); pr != nil { + if pr, _ = p.fast(id, false); pr != nil { return pr, nil } return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go index 639f6d529..7f08bfcd7 100644 --- a/net/pool/pool_flush_test.go +++ b/net/pool/pool_flush_test.go @@ -3,6 +3,7 @@ package pool import ( "context" "fmt" + net2 "net" "runtime" "strings" "sync" @@ -14,6 +15,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" atomic2 "go.uber.org/atomic" + "storj.io/drpc" "github.com/anyproto/any-sync/app" "github.com/anyproto/any-sync/app/ocache" @@ -176,7 +178,9 @@ func startFlusher(t *testing.T, fx *fixture, period time.Duration) (flushes *ato func fromClosePrepass() bool { buf := make([]byte, 1<<14) n := runtime.Stack(buf, false) - return strings.Contains(string(buf[:n]), "pool.closeCaches.func") + // closePeer exists for this: it is the one frame the pre-close has and + // the cache's own pass has not + return strings.Contains(string(buf[:n]), "pool.closePeer") } // hookCtx is a pair ctx whose Err can be paused: a seam between lookup's f @@ -1261,8 +1265,7 @@ func TestPool_FlushSwap(t *testing.T) { p := fx.Service.(*poolService).pool require.NoError(t, fx.AddPeer(ctx, newTestPeer("d"))) // shutdown has begun: nothing is evicted any more, so the duplicate - // can never be added; it must not be retried either (this used to - // loop ~340k times in 200ms) + // can never be added; it must not be retried either p.closingCancel() dup := newCtlPeer("d") got := make(chan error, 1) @@ -1580,13 +1583,16 @@ func TestPool_FlushSwap(t *testing.T) { return newTestPeer(peerId), nil } for i := 0; i < 20; i++ { + // a distinct id per round: the Get returns on its ctx while the + // entry may still be closing + id := fmt.Sprintf("p%d", i) gctx, cancel := context.WithCancel(ctx) - dead := newCtlPeer("p1") + dead := newCtlPeer(id) close(dead.closed) // the eviction's Close cancels the Get's ctx as it completes dead.closeHook = func(int32) { cancel() } - require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) - _, err := fx.Get(gctx, "p1") + require.NoError(t, p.current.Load().outgoing.Add(id, dead)) + _, err := fx.Get(gctx, id) require.ErrorIs(t, err, context.Canceled) cancel() } @@ -1884,10 +1890,19 @@ func TestPool_FlushSwap(t *testing.T) { fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { return newTestPeer(peerId), nil } + // a dead outgoing peer keeps the fast path from settling the miss by + // itself, so the lookup's own probe runs too + dead := newTestPeer("out") + require.NoError(t, dead.Close()) + require.NoError(t, fx.Service.(*poolService).pool.current.Load().outgoing.Add("out", dead)) _, err := fx.Get(ctx, "out") require.NoError(t, err) - // the fast path's peek and the lookup's probe, both Get-like - require.Equal(t, []bool{true, true}, touches) + // the fast path's peek and the lookup's probes (before and after the + // eviction), all Get-like + require.GreaterOrEqual(t, len(touches), 2) + for _, touch := range touches { + require.True(t, touch) + } touches = nil _, err = fx.Pick(ctx, "out") require.NoError(t, err) @@ -1899,7 +1914,9 @@ func TestPool_FlushSwap(t *testing.T) { a.Register(&testMetric{reg: reg}) }) defer fx.Finish() - // the peer connects between the fast path's miss and the probe + // the peer connects between the fast path's miss and the probe (a + // dead outgoing peer under the id keeps the fast path from settling + // the miss by itself) tp := newTestPeer("in") var peeks atomic.Int32 installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { @@ -1909,6 +1926,9 @@ func TestPool_FlushSwap(t *testing.T) { } }} }) + dead := newTestPeer("in") + require.NoError(t, dead.Close()) + require.NoError(t, fx.Service.(*poolService).pool.current.Load().outgoing.Add("in", dead)) fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { t.Error("must not dial: an incoming connection exists") return nil, nil @@ -1928,6 +1948,392 @@ func TestPool_FlushSwap(t *testing.T) { assert.Equal(t, float64(0), values["netpool_incoming_miss"]) assert.Equal(t, float64(0), values["netpool_outgoing_hit"]) }) + t.Run("a get whose ctx ends while a dead incoming is evicted does not dial", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + var deadDials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + if ctx.Err() != nil { + deadDials.Add(1) + return nil, ctx.Err() + } + return newTestPeer(peerId), nil + } + // the caller's ctx ends the moment the eviction's removal lands, so + // both the discard and the ctx are ready for the select + var cancelFn atomic.Value + var inner ocache.OCache + installIncoming(t, fx, func(in ocache.OCache, peek ocache.Peeker) ocache.OCache { + inner = in + return &hookedCache{OCache: in, peek: peek, afterRemoveSame: func() { + cancelFn.Load().(context.CancelFunc)() + }} + }) + for i := 0; i < 50; i++ { + id := fmt.Sprintf("p%d", i) + gctx, cancel := context.WithCancel(ctx) + cancelFn.Store(cancel) + dead := newTestPeer(id) + require.NoError(t, dead.Close()) + require.NoError(t, inner.Add(id, dead)) + _, err := fx.Get(gctx, id) + require.ErrorIs(t, err, context.Canceled) + cancel() + } + require.Zero(t, deadDials.Load(), "dialed with a dead ctx") + require.Equal(t, 0, p.current.Load().outgoing.Len()) + }) + t.Run("a flush landing at the incoming probe never dials into the replaced pair", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + var dials atomic.Int32 + dialedInto := map[*caches]int{} + var mu sync.Mutex + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + mu.Lock() + dialedInto[p.current.Load()]++ + mu.Unlock() + return newTestPeer(peerId), nil + } + // the second peek of each round is the lookup's probe: the pair is + // replaced right there, so the probe answers ErrClosed, not a miss + var armed atomic.Int32 + installIncoming(t, fx, func(in ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: in, peek: peek, onPeek: func() { + if armed.Add(-1) == 0 { + assert.NoError(t, fx.Flush(ctx)) + } + }} + }) + for i := 0; i < 100; i++ { + id := fmt.Sprintf("p%d", i) + before := dials.Load() + armed.Store(2) + pr, err := fx.Get(ctx, id) + require.NoError(t, err) + require.False(t, pr.IsClosed()) + // exactly one dial, into the pair that is current afterwards + require.Equal(t, int32(1), dials.Load()-before) + require.True(t, inCurrentOutgoing(p, pr)) + } + }) + t.Run("the parallel pre-close never exceeds its cap", func(t *testing.T) { + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 300 * time.Millisecond + }) + defer fx.Finish() + // every peer's first close hangs; the cache's own serial pass adds at + // most one blocked call on top of the capped workers + var inflight, peak atomic.Int32 + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + const n = 3 * maxPeerClosers + peers := make([]*ctlPeer, 0, n) + for i := 0; i < n; i++ { + pr := newCtlPeer(fmt.Sprintf("in%d", i)) + pr.closeHook = func(call int32) { + if call != 1 { + return + } + cur := inflight.Add(1) + for { + old := peak.Load() + if cur <= old || peak.CompareAndSwap(old, cur) { + break + } + } + <-releaseClose + inflight.Add(-1) + } + peers = append(peers, pr) + require.NoError(t, fx.AddPeer(ctx, pr)) + } + require.NoError(t, fx.Flush(ctx)) + require.Eventually(t, func() bool { return inflight.Load() >= maxPeerClosers }, time.Second, time.Millisecond) + require.Never(t, func() bool { return peak.Load() > maxPeerClosers+1 }, 100*time.Millisecond, time.Millisecond) + doReleaseClose() + require.Eventually(t, func() bool { + for _, pr := range peers { + if !pr.IsClosed() { + return false + } + } + return true + }, 5*time.Second, 10*time.Millisecond) + require.LessOrEqual(t, peak.Load(), int32(maxPeerClosers+1)) + }) + t.Run("a non-comparable peer implementation is pooled, flushed, evicted and reported", func(t *testing.T) { + // such a peer is tracked by id instead of by instance: everything + // works and nothing panics, as long as two connections for one id do + // not overlap (then the id-based eviction and marking are weaker, + // see peerKey and ocache.RemoveSame) + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // the watcher path: a dying incoming peer is evicted and reported + a := newValuePeer("a") + require.NoError(t, fx.AddPeer(ctx, a)) + a.close() + // the watcher removes the entry first and reports after + require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return len(obs.kindsFor("a")) == 1 }, time.Second, 10*time.Millisecond) + require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("a")) + // the replacement path through RemoveSame: no panic, the old one is + // closed and reported; by id the old one's watcher may take the + // replacement down with it, so its fate is not asserted + b1, b2 := newValuePeer("b"), newValuePeer("b") + require.NoError(t, fx.AddPeer(ctx, b1)) + require.NoError(t, fx.AddPeer(ctx, b2)) + require.True(t, b1.IsClosed()) + require.Eventually(t, func() bool { return len(obs.kindsFor("b")) >= 1 }, time.Second, 10*time.Millisecond) + b2.close() + require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) + // the discard path: a dead outgoing peer found by a lookup (added + // without a watcher, so no second instance of the id is live while + // the lookup evicts it and redials) + dead := newValuePeer("c") + dead.close() + require.NoError(t, p.current.Load().outgoing.Add("c", dead)) + fresh := newValuePeer("c") + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return fresh, nil + } + pr, err := fx.Get(ctx, "c") + require.NoError(t, err) + require.False(t, pr.IsClosed()) + require.Equal(t, 1, p.current.Load().outgoing.Len()) + // the flush path: the remaining peer is reported before Flush returns + require.NoError(t, fx.Flush(ctx)) + require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("c")) + require.Eventually(t, fresh.IsClosed, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 4 }, 200*time.Millisecond, 10*time.Millisecond) + }) + t.Run("a lookup parked in an eviction when the pool closes gets ErrClosed", func(t *testing.T) { + // the pair stays current but is cancelled by Close: the ctx error the + // eviction's select returns must surface as ErrClosed + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 200 * time.Millisecond + }) + p := fx.Service.(*poolService).pool + dead := newCtlPeer("p1") + close(dead.closed) + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + dead.closeHook = func(int32) { <-releaseClose } + require.NoError(t, p.current.Load().outgoing.Add("p1", dead)) + got := make(chan error, 1) + go func() { + _, err := fx.Get(ctx, "p1") + got <- err + }() + require.Eventually(t, func() bool { return dead.closeCalls.Load() == 1 }, time.Second, time.Millisecond) + fx.Finish() + require.ErrorIs(t, <-got, ocache.ErrClosed) + }) + t.Run("an incompatible-version verdict from the dialer survives flush", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return nil, handshake.ErrIncompatibleVersion + } + _, err := fx.Get(ctx, "p1") + require.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + require.NoError(t, fx.Flush(ctx)) + _, err = fx.Get(ctx, "p1") + require.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + _, err = fx.Pick(ctx, "p1") + require.ErrorIs(t, err, handshake.ErrIncompatibleVersion) + require.Equal(t, int32(1), dials.Load(), "the verdict was redialed after the flush") + }) + t.Run("a dial completing after flush is evicted from the pair it loaded into", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // each pair's outgoing cache records the ids RemoveSame is called for + var mu sync.Mutex + removed := map[*hookedCache][]string{} + var hooks []*hookedCache + installOutgoing(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + hc := &hookedCache{OCache: inner, peek: peek} + hc.onRemoveSameID = func(id string) { + mu.Lock() + removed[hc] = append(removed[hc], id) + mu.Unlock() + } + mu.Lock() + hooks = append(hooks, hc) + mu.Unlock() + return hc + }) + oldPair := p.current.Load() + oldHook := oldPair.outgoing.(*hookedCache) + late := newTestPeer("p1") + dialStarted := make(chan struct{}) + releaseDial, doReleaseDial := newRelease() + defer doReleaseDial() + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + close(dialStarted) + <-releaseDial // completes despite the cancellation + return late, nil + } + loaded := make(chan struct{}) + go func() { + _, _ = oldPair.outgoing.Get(ctx, "p1") + close(loaded) + }() + <-dialStarted + require.NoError(t, fx.Flush(ctx)) + newHook := p.current.Load().outgoing.(*hookedCache) + doReleaseDial() + <-loaded + // the late peer is closed by the old pair; its watcher evicts it from + // the old pair, not from the one that is current now + require.Eventually(t, late.IsClosed, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(removed[oldHook]) == 1 + }, time.Second, 10*time.Millisecond) + mu.Lock() + defer mu.Unlock() + require.Equal(t, []string{"p1"}, removed[oldHook]) + require.Empty(t, removed[newHook]) + }) + t.Run("a GC TryClose the incoming peer declines does not cause a dial", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + in := newCtlPeer("p1") + inTryClose := make(chan struct{}) + releaseTryClose, doReleaseTryClose := newRelease() + defer doReleaseTryClose() + in.tryClose = func() (bool, error) { + close(inTryClose) + <-releaseTryClose + return false, nil // in use: stays + } + require.NoError(t, fx.AddPeer(ctx, in)) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + t.Error("dialed although an incoming connection exists") + return nil, nil + } + // the GC holds the entry in closing while TryClose runs + gcDone := make(chan struct{}) + go func() { + defer close(gcDone) + _, _ = p.current.Load().incoming.TryRemove("p1") + }() + <-inTryClose + got := make(chan peer.Peer, 1) + go func() { + pr, err := fx.Get(ctx, "p1") + assert.NoError(t, err) + got <- pr + }() + // the Get waits for the close to resolve instead of dialing + require.Never(t, func() bool { return len(got) > 0 }, 50*time.Millisecond, 5*time.Millisecond) + doReleaseTryClose() + <-gcDone + require.Same(t, in, <-got) + }) + t.Run("the loader never dials for a replaced pair", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + var dials atomic.Int32 + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + dials.Add(1) + return newTestPeer(peerId), nil + } + // the pair is replaced as the outgoing load starts, before the loader + // runs: no dial for the old pair, one for the fresh one + var armed atomic.Bool + installOutgoing(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onGet: func() { + if armed.CompareAndSwap(true, false) { + assert.NoError(t, fx.Flush(ctx)) + } + }} + }) + armed.Store(true) + pr, err := fx.Get(ctx, "p1") + require.NoError(t, err) + require.False(t, pr.IsClosed()) + require.Equal(t, int32(1), dials.Load()) + }) + t.Run("a dial in flight ends with its pair even for a direct load", func(t *testing.T) { + // the dial ctx is bound to the pair by the loader itself, not only + // through the lookup's ctx or through the cache's Close (which + // cancels loads too, but is held back here) + fx := newFixtureCfg(t, nil, func(ps *poolService, a *app.App) { + ps.closeTimeout = 2 * time.Second + }) + defer fx.Finish() + p := fx.Service.(*poolService).pool + releaseClose, doReleaseClose := newRelease() + defer doReleaseClose() + var armed atomic.Bool + installOutgoing(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onClose: func() { + if armed.CompareAndSwap(true, false) { + <-releaseClose + } + }} + }) + dialStarted := make(chan struct{}) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + close(dialStarted) + <-ctx.Done() + return nil, ctx.Err() + } + oldPair := p.current.Load() + loaded := make(chan error, 1) + go func() { + _, err := oldPair.outgoing.Get(ctx, "p1") + loaded <- err + }() + <-dialStarted + // the old pair's outgoing Close is held: only the pair ctx can end + // the dial now + armed.Store(true) + require.NoError(t, fx.Flush(ctx)) + select { + case err := <-loaded: + require.Error(t, err) + case <-time.After(time.Second): + t.Fatal("the dial outlived its pair") + } + }) + t.Run("a peer comparable by type but not by value is tracked without panicking", func(t *testing.T) { + // a struct with an interface field: the type is comparable, the value + // is not once the field holds a slice; == and a map insert panic + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + p := fx.Service.(*poolService).pool + a := newIfacePeer("a") + require.NoError(t, fx.AddPeer(ctx, a)) + a.close() + require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return len(obs.kindsFor("a")) == 1 }, time.Second, 10*time.Millisecond) + b := newIfacePeer("b") + require.NoError(t, fx.AddPeer(ctx, b)) + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return newIfacePeer(peerId), nil + } + _, err := fx.Get(ctx, "c") + require.NoError(t, err) + require.NoError(t, fx.Flush(ctx)) + require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("b")) + require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("c")) + require.Eventually(t, b.IsClosed, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.getClosed()) > 3 }, 200*time.Millisecond, 10*time.Millisecond) + }) t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { obs := &poolEventRecorder{} fx := newFixtureWithObserver(t, obs) @@ -2226,7 +2632,7 @@ func TestPool_FlushGoroutines(t *testing.T) { fx.Finish() // polled from this goroutine so the baseline is comparable (Eventually // would add its own) - deadline := time.Now().Add(5 * time.Second) + deadline := time.Now().Add(15 * time.Second) for runtime.NumGoroutine() > before { if time.Now().After(deadline) { t.Fatalf("goroutines before=%d after=%d", before, runtime.NumGoroutine()) @@ -2319,7 +2725,7 @@ func (c *stickyCache) RemoveSame(ctx context.Context, id string, value ocache.Ob return false, nil } -func (c *stickyCache) Peek(id string, touch bool) (ocache.Object, bool) { +func (c *stickyCache) Peek(id string, touch bool) (ocache.Object, ocache.PeekState) { return c.peek.Peek(id, touch) } @@ -2333,16 +2739,21 @@ type hookedCache struct { ocache.OCache peek ocache.Peeker onRemoveSame func() + onRemoveSameID func(id string) afterRemoveSame func() onForEach func() onGet func() onPeek func() + onClose func() } func (c *hookedCache) RemoveSame(ctx context.Context, id string, value ocache.Object) (bool, error) { if c.onRemoveSame != nil { c.onRemoveSame() } + if c.onRemoveSameID != nil { + c.onRemoveSameID(id) + } ok, err := c.OCache.RemoveSame(ctx, id, value) if c.afterRemoveSame != nil { c.afterRemoveSame() @@ -2357,6 +2768,13 @@ func (c *hookedCache) Get(ctx context.Context, id string) (ocache.Object, error) return c.OCache.Get(ctx, id) } +func (c *hookedCache) Close() error { + if c.onClose != nil { + c.onClose() + } + return c.OCache.Close() +} + func (c *hookedCache) ForEach(f func(v ocache.Object) bool) { c.OCache.ForEach(f) if c.onForEach != nil { @@ -2364,7 +2782,7 @@ func (c *hookedCache) ForEach(f func(v ocache.Object) bool) { } } -func (c *hookedCache) Peek(id string, touch bool) (ocache.Object, bool) { +func (c *hookedCache) Peek(id string, touch bool) (ocache.Object, ocache.PeekState) { if c.onPeek != nil { c.onPeek() } @@ -2375,6 +2793,19 @@ func (c *hookedCache) WaitClosing(ctx context.Context, id string) error { return c.peek.WaitClosing(ctx, id) } +// installOutgoing is installIncoming for the outgoing cache +func installOutgoing(t *testing.T, fx *fixture, wrap func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache) { + p := fx.Service.(*poolService).pool + orig := p.newCaches + p.newCaches = func() *caches { + c := orig() + c.outgoing = wrap(c.outgoing, c.peekOutgoing) + c.peekOutgoing = mustPeeker(c.outgoing) + return c + } + require.NoError(t, fx.Flush(ctx)) +} + // installIncoming makes every pair the pool builds from now on use wrap(inner) // as its incoming cache, and flushes once so the current pair has it func installIncoming(t *testing.T, fx *fixture, wrap func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache) { @@ -2396,7 +2827,7 @@ type peekRecorder struct { record func(touch bool) } -func (c *peekRecorder) Peek(id string, touch bool) (ocache.Object, bool) { +func (c *peekRecorder) Peek(id string, touch bool) (ocache.Object, ocache.PeekState) { c.record(touch) return c.peek.Peek(id, touch) } @@ -2404,3 +2835,116 @@ func (c *peekRecorder) Peek(id string, touch bool) (ocache.Object, bool) { func (c *peekRecorder) WaitClosing(ctx context.Context, id string) error { return c.peek.WaitClosing(ctx, id) } + +func inCurrentOutgoing(p *pool, pr peer.Peer) bool { + v, err := p.current.Load().outgoing.Pick(ctx, pr.Id()) + return err == nil && v == ocache.Object(pr) +} + +// valuePeer is a peer.Peer implementation that is not comparable (a slice +// field, value receivers): == on two of them panics, so the pool must never +// compare peers directly +type valuePeer struct { + id string + mu *sync.Mutex + closed chan struct{} + tags []string +} + +func newValuePeer(id string) valuePeer { + return valuePeer{id: id, mu: &sync.Mutex{}, closed: make(chan struct{}), tags: []string{id}} +} + +// close is idempotent under concurrent callers (the pool closes a flushed +// peer from the pre-close and from the cache's pass) +func (v valuePeer) close() { + v.mu.Lock() + defer v.mu.Unlock() + select { + case <-v.closed: + default: + close(v.closed) + } +} + +func (v valuePeer) Id() string { return v.id } +func (v valuePeer) Addr() string { return "" } +func (v valuePeer) Close() error { v.close(); return nil } +func (v valuePeer) TryClose(time.Duration) (bool, error) { v.close(); return true, nil } +func (v valuePeer) IsClosed() bool { + select { + case <-v.closed: + return true + default: + return false + } +} +func (v valuePeer) CloseChan() <-chan struct{} { return v.closed } +func (v valuePeer) SetTTL(time.Duration) {} +func (v valuePeer) DoDrpc(context.Context, func(conn drpc.Conn) error) error { + return fmt.Errorf("not implemented") +} +func (v valuePeer) AcquireDrpcConn(context.Context) (drpc.Conn, error) { + return nil, fmt.Errorf("not implemented") +} +func (v valuePeer) ReleaseDrpcConn(context.Context, drpc.Conn) {} +func (v valuePeer) Context() context.Context { return ctx } +func (v valuePeer) Accept() (net2.Conn, error) { return nil, fmt.Errorf("not implemented") } +func (v valuePeer) Open(context.Context) (net2.Conn, error) { + return nil, fmt.Errorf("not implemented") +} + +var _ peer.Peer = valuePeer{} + +// ifacePeer is comparable as a type (no slice or map fields of its own) but +// not as a value: its payload field holds a slice, so == on two of them, or a +// map insert, panics. Value receivers, like valuePeer. +type ifacePeer struct { + id string + mu *sync.Mutex + closed chan struct{} + payload any +} + +func newIfacePeer(id string) ifacePeer { + return ifacePeer{id: id, mu: &sync.Mutex{}, closed: make(chan struct{}), payload: []string{id}} +} + +func (v ifacePeer) close() { + v.mu.Lock() + defer v.mu.Unlock() + select { + case <-v.closed: + default: + close(v.closed) + } +} + +func (v ifacePeer) Id() string { return v.id } +func (v ifacePeer) Addr() string { return "" } +func (v ifacePeer) Close() error { v.close(); return nil } +func (v ifacePeer) TryClose(time.Duration) (bool, error) { v.close(); return true, nil } +func (v ifacePeer) IsClosed() bool { + select { + case <-v.closed: + return true + default: + return false + } +} +func (v ifacePeer) CloseChan() <-chan struct{} { return v.closed } +func (v ifacePeer) SetTTL(time.Duration) {} +func (v ifacePeer) DoDrpc(context.Context, func(conn drpc.Conn) error) error { + return fmt.Errorf("not implemented") +} +func (v ifacePeer) AcquireDrpcConn(context.Context) (drpc.Conn, error) { + return nil, fmt.Errorf("not implemented") +} +func (v ifacePeer) ReleaseDrpcConn(context.Context, drpc.Conn) {} +func (v ifacePeer) Context() context.Context { return ctx } +func (v ifacePeer) Accept() (net2.Conn, error) { return nil, fmt.Errorf("not implemented") } +func (v ifacePeer) Open(context.Context) (net2.Conn, error) { + return nil, fmt.Errorf("not implemented") +} + +var _ peer.Peer = ifacePeer{} diff --git a/net/pool/pool_test.go b/net/pool/pool_test.go index 41d73a05e..1c9f6d93a 100644 --- a/net/pool/pool_test.go +++ b/net/pool/pool_test.go @@ -633,6 +633,7 @@ var _ peer.Peer = (*testPeer)(nil) type testPeer struct { id string closeMu sync.Mutex + closes int closed chan struct{} created time.Time subConnections int @@ -692,9 +693,11 @@ func (t *testPeer) TryClose(objectTTL time.Duration) (res bool, err error) { func (t *testPeer) Close() error { // the pool may close a rejected peer from several paths at once; - // idempotent and silent like the real peer (its MultiConn.Close is) + // idempotent and silent like the real peer (its MultiConn.Close is). + // closes counts every call, for tests that pin how often that happens t.closeMu.Lock() defer t.closeMu.Unlock() + t.closes++ select { case <-t.closed: default: diff --git a/net/pool/poolservice.go b/net/pool/poolservice.go index fce209510..821e8f047 100644 --- a/net/pool/poolservice.go +++ b/net/pool/poolservice.go @@ -92,15 +92,22 @@ func (p *poolService) Init(a *app.App) (err error) { return nil } -// newCaches builds one cache pair. The outgoing loader binds its watcher to -// the cache it loads into, not to whichever pair is current when the dial -// finishes: after a Flush that is a different one. +// newCaches builds one cache pair. The outgoing loader belongs to its pair: it +// never dials for a pair that has been replaced or is closing (the load fails +// with ErrClosed, and a lookup retries on the current pair), the dial itself +// ends with the pair, and the watcher it starts is bound to the cache it +// loads into, not to whichever pair is current when the dial finishes. func (p *poolService) newCaches(outgoingMetrics, incomingMetrics ocache.Option) *caches { c := &caches{} c.ctx, c.cancel = context.WithCancel(context.Background()) c.outgoing = ocache.New( func(ctx context.Context, id string) (value ocache.Object, err error) { - value, err = p.dialer.Dial(ctx, id) + if c.ctx.Err() != nil { + return nil, ocache.ErrClosed + } + dctx, release := c.bind(ctx) + defer release() + value, err = p.dialer.Dial(dctx, id) if err != nil { if errors.Is(err, handshake.ErrIncompatibleVersion) { return &errObject{id: id, err: err, createdTime: atomic.NewTime(time.Now())}, nil diff --git a/net/secureservice/handshake/proto.go b/net/secureservice/handshake/proto.go index 821cdb3b8..1b6052c78 100644 --- a/net/secureservice/handshake/proto.go +++ b/net/secureservice/handshake/proto.go @@ -22,26 +22,7 @@ type ProtoChecker struct { // left to the caller. The close is synchronous and can block on the // transport; OutgoingProtoHandshakeWithCloser moves it off the caller's path. func OutgoingProtoHandshake(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto) (*handshakeproto.Proto, error) { - if ctx == nil { - ctx = context.Background() - } - h := newHandshake() - done := make(chan struct{}) - var ( - err error - remoteProto *handshakeproto.Proto - ) - go func() { - defer close(done) - remoteProto, err = outgoingProtoHandshake(h, conn, proto, nil, nil) - }() - select { - case <-done: - return remoteProto, err - case <-ctx.Done(): - _ = conn.Close() - return nil, ctx.Err() - } + return outgoingProtoHandshakeCloser(ctx, conn, proto, closeSync, false) } // OutgoingProtoHandshakeWithCloser is OutgoingProtoHandshake for a caller @@ -54,6 +35,17 @@ func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto if closeConn == nil { return OutgoingProtoHandshake(ctx, conn, proto) } + return outgoingProtoHandshakeCloser(ctx, conn, proto, closeConn, true) +} + +func closeSync(conn net.Conn) { + _ = conn.Close() +} + +// outgoingProtoHandshakeCloser runs the handshake, closing conn through +// closeConn at most once: on an I/O error, on cancellation, and with +// closeOnAnyErr on a protocol-level error as well +func outgoingProtoHandshakeCloser(ctx context.Context, conn net.Conn, proto *handshakeproto.Proto, closeConn func(net.Conn), closeOnAnyErr bool) (*handshakeproto.Proto, error) { if ctx == nil { ctx = context.Background() } @@ -77,7 +69,7 @@ func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto go func() { defer close(done) remoteProto, err = outgoingProtoHandshake(h, conn, proto, closeConn, &claimed) - if err != nil { + if err != nil && closeOnAnyErr { // a no-op if the handshake or an abandoning caller closed it closeConn(conn) } @@ -95,7 +87,7 @@ func OutgoingProtoHandshakeWithCloser(ctx context.Context, conn net.Conn, proto // The deadline unblocks a pending read at once where supported; the // close is what reliably ends the handshake everywhere (a yamux write // waiting for the send loop ignores deadlines, and some conns have - // no deadlines at all). Neither blocks the caller. + // no deadlines at all). _ = conn.SetDeadline(time.Now()) closeConn(conn) return nil, ctx.Err() diff --git a/net/secureservice/handshake/proto_test.go b/net/secureservice/handshake/proto_test.go index 8e35d32d2..150d90b7f 100644 --- a/net/secureservice/handshake/proto_test.go +++ b/net/secureservice/handshake/proto_test.go @@ -233,11 +233,18 @@ func TestOutgoingProtoHandshakeWithCloser_CancelDoesNotWaitForClose(t *testing.T ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) defer cancel() - start := time.Now() closer := func(c net.Conn) { go func() { _ = c.Close() }() } - _, err := OutgoingProtoHandshakeWithCloser(ctx, conn, &handshakeproto.Proto{Proto: 1}, closer) - require.ErrorIs(t, err, context.DeadlineExceeded) - assert.Less(t, time.Since(start), time.Second, "the caller must not wait on the blocked close") + res := make(chan error, 1) + go func() { + _, err := OutgoingProtoHandshakeWithCloser(ctx, conn, &handshakeproto.Proto{Proto: 1}, closer) + res <- err + }() + select { + case err := <-res: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(3 * time.Second): + t.Fatal("the caller must not wait on the blocked close") + } // the conn is still closed, off the caller's path select { @@ -424,3 +431,44 @@ func TestHandshakeError_Unwrap(t *testing.T) { // a wrapped transport error is not mistaken for a protocol sentinel assert.NotErrorIs(t, HandshakeError{Err: io.EOF}, ErrIncompatibleVersion) } + +// countingConn counts writes and closes +type countingConn struct { + writes, closes atomic.Int32 +} + +func (c *countingConn) Read([]byte) (int, error) { return 0, io.EOF } +func (c *countingConn) Write(b []byte) (int, error) { c.writes.Add(1); return len(b), nil } +func (c *countingConn) Close() error { c.closes.Add(1); return nil } +func (c *countingConn) LocalAddr() net.Addr { return nil } +func (c *countingConn) RemoteAddr() net.Addr { return nil } +func (c *countingConn) SetDeadline(time.Time) error { return nil } +func (c *countingConn) SetReadDeadline(time.Time) error { return nil } +func (c *countingConn) SetWriteDeadline(time.Time) error { return nil } + +func TestHandshake_AbandonedSkipsReplyAndClose(t *testing.T) { + conn := &countingConn{} + var abandoned atomic.Bool + abandoned.Store(true) + h := &handshake{conn: conn, abandoned: &abandoned, localAck: &handshakeproto.Ack{}} + h.tryWriteErrAndClose(ErrUnexpected) + assert.Zero(t, conn.writes.Load(), "no ack to a caller that gave up") + assert.Zero(t, conn.closes.Load(), "the caller closes it") + + abandoned.Store(false) + h.tryWriteErrAndClose(ErrUnexpected) + assert.Equal(t, int32(1), conn.writes.Load()) + assert.Equal(t, int32(1), conn.closes.Load()) +} + +func TestHandshake_ReleaseClearsCloser(t *testing.T) { + h := newHandshake() + h.conn = &countingConn{} + h.closeConn = func(net.Conn) {} + h.abandoned = &atomic.Bool{} + h.release() + // a recycled handshake must not carry a previous caller's closer + assert.Nil(t, h.closeConn) + assert.Nil(t, h.abandoned) + assert.Nil(t, h.conn) +} diff --git a/net/transport/quic/conn_errors_test.go b/net/transport/quic/conn_errors_test.go index 5943b1bc8..7a167caec 100644 --- a/net/transport/quic/conn_errors_test.go +++ b/net/transport/quic/conn_errors_test.go @@ -85,8 +85,7 @@ func TestQuicNetConn_ConnCloseNormalized(t *testing.T) { fxC := newFixture(t) defer fxC.finish(t) - mcC, err := fxC.Dial(ctx, fxS.addr) - require.NoError(t, err) + mcC := dialRetry(t, fxC, fxS.addr) var mcS transport.MultiConn select { case mcS = <-fxS.accepter.mcs: @@ -131,12 +130,11 @@ func TestQuicNetConn_StreamResetNotNormalized(t *testing.T) { fxC := newFixture(t) defer fxC.finish(t) - mcC, err := fxC.Dial(ctx, fxS.addr) - require.NoError(t, err) + mcC := dialRetry(t, fxC, fxS.addr) var mcS transport.MultiConn select { case mcS = <-fxS.accepter.mcs: - case <-time.After(time.Second * 5): + case <-time.After(30 * time.Second): t.Fatal("timeout") } @@ -153,12 +151,38 @@ func TestQuicNetConn_StreamResetNotNormalized(t *testing.T) { // only the stream goes away: the server cancels its read side, which // makes the client's writes fail with a stream error sConn.(quicNetConn).CancelRead(42) - require.Eventually(t, func() bool { - _, err = conn.Write([]byte("more")) - return err != nil - }, 5*time.Second, 10*time.Millisecond) + // the STOP_SENDING takes a round trip: keep writing until it lands + writeErr := make(chan error, 1) + go func() { + for { + if _, wErr := conn.Write([]byte("more")); wErr != nil { + writeErr <- wErr + return + } + time.Sleep(5 * time.Millisecond) + } + }() + select { + case err = <-writeErr: + case <-time.After(30 * time.Second): + t.Fatal("the stream reset never reached the writer") + } var streamErr *quic.StreamError require.ErrorAs(t, err, &streamErr) assert.False(t, errors.Is(err, transport.ErrConnClosed)) assert.False(t, mcC.IsClosed()) } + +// dialRetry dials, retrying a dial that timed out: under a loaded -race run a +// loopback QUIC handshake can idle out, which is not what these tests are about +func dialRetry(t *testing.T, fx *fixture, addr string) transport.MultiConn { + var err error + for i := 0; i < 3; i++ { + var mc transport.MultiConn + if mc, err = fx.Dial(ctx, addr); err == nil { + return mc + } + } + require.NoError(t, err) + return nil +} diff --git a/net/transport/yamux/conn.go b/net/transport/yamux/conn.go index e698c1aa4..92190dd43 100644 --- a/net/transport/yamux/conn.go +++ b/net/transport/yamux/conn.go @@ -62,6 +62,10 @@ type openResult struct { // past it together can overshoot it by their number. func (y *yamuxConn) Open(ctx context.Context) (conn net.Conn, err error) { if conn, err = y.open(ctx); err != nil { + if errors.Is(err, yamux.ErrSessionShutdown) { + // like Accept and the QUIC transport + err = transport.NewConnClosedError(err) + } return nil, err } return y.wrapStream(conn), nil @@ -75,6 +79,9 @@ func (y *yamuxConn) open(ctx context.Context) (conn net.Conn, err error) { if err = ctx.Err(); err != nil { return nil, err } + // The helper is needed even with free SYN slots: Session.Open also sends + // the SYN through the session's send queue, which on a congested link + // waits up to ConnectionWriteTimeout regardless of ctx. if err = y.waitBacklog(ctx); err != nil { return nil, err } @@ -193,24 +200,33 @@ func (s yamuxStream) Write(b []byte) (n int, err error) { } // wrapSessionDead wraps err with transport.NewConnClosedError when it was -// caused by the session dying. ErrSessionShutdown always is. io.EOF, a stream -// reset, a closed stream and a connection write timeout count only while the -// session is closed: on a live session they are stream-level outcomes (a -// remote close or reset, or a send that waited out ConnectionWriteTimeout on a -// slow but live peer, which yamux does not treat as fatal) and are returned -// unchanged. A truly stalled session ends through missed keepalives, after -// which its streams fail with errors covered here. +// caused by the session dying. ErrSessionShutdown always is. A stream reset, a +// closed stream and a connection write timeout count only while the session is +// closed: on a live session they are stream-level outcomes (a remote reset, or +// a send that waited out ConnectionWriteTimeout on a slow but live peer, which +// yamux does not treat as fatal) and are returned unchanged; on a closed one +// they are reported as caused by ErrSessionShutdown, with the original error +// still in the chain. A truly stalled session +// ends through missed keepalives, after which its streams fail with errors +// covered here. +// +// io.EOF is never wrapped: a stream that got its data and FIN before the +// session died must still end in a plain io.EOF (io.ReadAll, io.Copy and +// err == io.EOF checks rely on it). An RPC cut short by the session dying is +// classified above the transport, in the peer's sub conn. func (s yamuxStream) wrapSessionDead(err error) error { if err == nil { return nil } switch { case errors.Is(err, yamux.ErrSessionShutdown): - case errors.Is(err, io.EOF), errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed), + case errors.Is(err, yamux.ErrConnectionReset), errors.Is(err, yamux.ErrStreamClosed), errors.Is(err, yamux.ErrConnectionWriteTimeout): if !s.sess.IsClosed() { return err } + // keeps the cause (and a write timeout's Timeout()) reachable + err = errors.Join(yamux.ErrSessionShutdown, err) default: return err } diff --git a/net/transport/yamux/conn_test.go b/net/transport/yamux/conn_test.go index b8b524cf5..a8a5f79aa 100644 --- a/net/transport/yamux/conn_test.go +++ b/net/transport/yamux/conn_test.go @@ -200,6 +200,12 @@ func TestYamuxConn_OpenClosedSession(t *testing.T) { require.NoError(t, mc.Session.Close()) _, err := mc.Open(ctx) require.ErrorIs(t, err, yamux.ErrSessionShutdown) + require.ErrorIs(t, err, transport.ErrConnClosed) + // the helper path as well + cctx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + _, err = mc.Open(cctx) + require.ErrorIs(t, err, transport.ErrConnClosed) } // streamPair opens a stream from mc and accepts it on server @@ -254,13 +260,13 @@ func TestYamuxStream_SessionDeathNormalized(t *testing.T) { time.Sleep(20 * time.Millisecond) require.NoError(t, mc.Session.Close()) err := waitErr(t, res) - assert.ErrorIs(t, err, transport.ErrConnClosed) - assert.ErrorIs(t, err, io.EOF, "the original error stays reachable") + // a read ends in a plain io.EOF: the peer's sub conn, not the + // transport, classifies an RPC cut short this way + assert.Equal(t, io.EOF, err) _, err = client.Write([]byte("more")) assert.ErrorIs(t, err, transport.ErrConnClosed) - // yamux force-closed the stream on shutdown - assert.ErrorIs(t, err, yamux.ErrStreamClosed) + assert.ErrorIs(t, err, yamux.ErrSessionShutdown) }) t.Run("remote session close mid-read", func(t *testing.T) { mc, server := newSessionPair(t) @@ -269,7 +275,7 @@ func TestYamuxStream_SessionDeathNormalized(t *testing.T) { time.Sleep(20 * time.Millisecond) require.NoError(t, server.Close()) err := waitErr(t, res) - assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.Equal(t, io.EOF, err) assert.True(t, mc.Session.IsClosed()) }) t.Run("accepted stream, remote session close", func(t *testing.T) { @@ -295,7 +301,9 @@ func TestYamuxStream_SessionDeathNormalized(t *testing.T) { res := readErr(t, local) time.Sleep(20 * time.Millisecond) require.NoError(t, server.Close()) - assert.ErrorIs(t, waitErr(t, res), transport.ErrConnClosed) + assert.Equal(t, io.EOF, waitErr(t, res)) + _, err = local.Write([]byte("more")) + assert.ErrorIs(t, err, transport.ErrConnClosed) }) t.Run("remote stream close is a plain EOF", func(t *testing.T) { mc, server := newSessionPair(t) @@ -320,11 +328,36 @@ func TestYamuxStream_SessionFatalErrorsNormalized(t *testing.T) { } // once the session is closed they mean it died require.NoError(t, mc.Session.Close()) - for _, cause := range []error{io.EOF, yamux.ErrConnectionReset, yamux.ErrStreamClosed, yamux.ErrConnectionWriteTimeout} { + for _, cause := range []error{yamux.ErrConnectionReset, yamux.ErrStreamClosed, yamux.ErrConnectionWriteTimeout} { err := s.wrapSessionDead(cause) assert.ErrorIs(t, err, transport.ErrConnClosed, cause.Error()) - assert.ErrorIs(t, err, cause) + assert.ErrorIs(t, err, yamux.ErrSessionShutdown) + assert.ErrorIs(t, err, cause, "the original error stays reachable") + assert.NotErrorIs(t, err, io.EOF) } + // io.EOF never is: a stream that got its FIN must end in a plain EOF + assert.Equal(t, io.EOF, s.wrapSessionDead(io.EOF)) +} + +// TestYamuxStream_ReadAllAfterFINThenSessionDeath: data and FIN received +// before the session dies still read as a complete stream ending in io.EOF +func TestYamuxStream_ReadAllAfterFINThenSessionDeath(t *testing.T) { + mc, server := newSessionPair(t) + client, remote := streamPair(t, mc, server) + payload := []byte("the whole message") + _, err := remote.Write(payload) + require.NoError(t, err) + require.NoError(t, remote.Close()) + // let the data and FIN arrive, then the session dies + time.Sleep(50 * time.Millisecond) + require.NoError(t, server.Close()) + require.Eventually(t, mc.Session.IsClosed, 5*time.Second, time.Millisecond) + + got, err := io.ReadAll(client) + require.NoError(t, err, "io.ReadAll sees a plain io.EOF") + assert.Equal(t, payload, got) + _, err = client.Read(make([]byte, 1)) + assert.True(t, err == io.EOF, "err == io.EOF must hold, got %v", err) } // slowReader reads slowly, so the writer's sends queue up From d397dd7f433123bfa9965d8a3c9f99269bb18fbe Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 19:10:21 +0200 Subject: [PATCH 6/7] net/peer, net/pool: keep server error codes; wrap open streams peer: - never treat an RPC error carrying a drpc code as a lost connection; the ErrConnClosed wrapper also exposes Code() - streams from NewStream report a dead sub conn on MsgSend/MsgRecv/ CloseSend as ErrConnClosed (caller cancellation stays Canceled) - NewConnClosedError doesn't double-wrap; acceptLoop uses errors.Is pool: - Pick and getIfActive return not-found at once when both caches miss (0 allocs, as on main) - non-comparable peers are held behind a pool-owned pointer, so each pooled instance has exact identity (no by-id fallback) - Flush/Close before Init are no-ops --- app/ocache/ocache.go | 20 ++-- app/ocache/ocache_test.go | 8 +- net/peer/closeasync_test.go | 68 +++++++++++-- net/peer/connlost_ext_test.go | 56 +++++++++++ net/peer/peer.go | 47 ++++++++- net/peerservice/dialoutcome_test.go | 34 ++++++- net/pool/pool.go | 144 +++++++++++++++++----------- net/pool/pool_bench_test.go | 11 +++ net/pool/pool_flush_test.go | 100 ++++++++++++++++--- net/pool/poolservice.go | 9 +- net/transport/transport.go | 15 ++- net/transport/transport_test.go | 11 +++ net/transport/yamux/conn_test.go | 42 ++++++++ 13 files changed, 457 insertions(+), 108 deletions(-) diff --git a/app/ocache/ocache.go b/app/ocache/ocache.go index bab7ef35a..83b3e7e93 100644 --- a/app/ocache/ocache.go +++ b/app/ocache/ocache.go @@ -133,10 +133,9 @@ type OCache interface { // RemoveSame closes and removes the object only if the value currently // stored under id is exactly the given one (pointer identity). It lets a // caller evict a specific instance it owns without racing a newer value - // that has replaced it under the same id. A value of a non-comparable - // type has no identity: for such values RemoveSame removes whatever is - // stored under id. Returns ok=true only when this call performed the - // removal. + // that has replaced it under the same id. A value that is not comparable + // has no identity and is never matched (store such values behind a + // pointer). Returns ok=true only when this call performed the removal. RemoveSame(ctx context.Context, id string, value Object) (ok bool, err error) // TryRemove tries to close and to remove the object. ok reports whether // this call removed it; (false, nil) means the object declined to close, @@ -452,22 +451,19 @@ func (c *oCache) RemoveSame(ctx context.Context, id string, value Object) (ok bo // sameObject reports whether stored is the very instance given. Pointer // implementations (the usual kind) compare by identity. A value that is not // comparable has no identity to check and must not panic the comparison: it -// is treated as the stored one, so RemoveSame degrades to Remove by id for -// such values. Checked on the values, not the types: a struct with an -// interface field is comparable as a type and still panics when that field -// holds a slice. +// never matches, so nothing is removed by mistake (a caller that needs +// instance-safe removal of such values stores them behind a pointer). Checked +// on the values, not the types: a struct with an interface field is +// comparable as a type and still panics when that field holds a slice. func sameObject(stored, given Object) bool { if stored == nil || given == nil { // a still-loading entry has no value yet; nothing matches it return false } sv, gv := reflect.ValueOf(stored), reflect.ValueOf(given) - if sv.Type() != gv.Type() { + if sv.Type() != gv.Type() || !sv.Comparable() || !gv.Comparable() { return false } - if !sv.Comparable() || !gv.Comparable() { - return true - } return stored == given } diff --git a/app/ocache/ocache_test.go b/app/ocache/ocache_test.go index 3b2adfeba..d8a699209 100644 --- a/app/ocache/ocache_test.go +++ b/app/ocache/ocache_test.go @@ -1507,10 +1507,10 @@ func TestOCache_SameObject(t *testing.T) { require.False(t, sameObject(a, b)) // different types never match require.False(t, sameObject(a, sliceObject{})) - // values that cannot be compared are taken as the stored one - require.True(t, sameObject(sliceObject{tags: []string{"x"}}, sliceObject{tags: []string{"y"}})) - // comparable as a type, not as a value: still no panic - require.True(t, sameObject(ifaceObject{payload: []string{"x"}}, ifaceObject{payload: []string{"y"}})) + // values that cannot be compared have no identity: never a match + require.False(t, sameObject(sliceObject{tags: []string{"x"}}, sliceObject{tags: []string{"y"}})) + // comparable as a type, not as a value: still no panic, no match + require.False(t, sameObject(ifaceObject{payload: []string{"x"}}, ifaceObject{payload: []string{"y"}})) require.False(t, sameObject(ifaceObject{payload: "x"}, ifaceObject{payload: "y"})) require.True(t, sameObject(ifaceObject{payload: "x"}, ifaceObject{payload: "x"})) } diff --git a/net/peer/closeasync_test.go b/net/peer/closeasync_test.go index 98c296b5d..bf1b8a496 100644 --- a/net/peer/closeasync_test.go +++ b/net/peer/closeasync_test.go @@ -2,6 +2,7 @@ package peer import ( "context" + "errors" "io" "net" "sync" @@ -13,11 +14,13 @@ import ( "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" "storj.io/drpc" + "storj.io/drpc/drpcerr" "storj.io/drpc/drpcwire" "github.com/anyproto/any-sync/net/connutil" "github.com/anyproto/any-sync/net/secureservice/handshake" "github.com/anyproto/any-sync/net/secureservice/handshake/handshakeproto" + "github.com/anyproto/any-sync/net/transport" ) type rawMsg []byte @@ -175,10 +178,10 @@ func TestPeer_HandshakeFailureCloseDoesNotBlock(t *testing.T) { }() fx.mc.EXPECT().Open(gomock.Any()).Return(conn, nil) - start := time.Now() - _, err := fx.AcquireDrpcConn(ctx) - require.ErrorIs(t, err, handshake.ErrRemoteIncompatibleProto) - assert.Less(t, time.Since(start), time.Second) + returnsWithin(t, 3*time.Second, "the blocked close must not run on the caller", func() { + _, err := fx.AcquireDrpcConn(ctx) + assert.ErrorIs(t, err, handshake.ErrRemoteIncompatibleProto) + }) }) t.Run("deadline", func(t *testing.T) { fx := newFixture(t, "p1") @@ -194,10 +197,10 @@ func TestPeer_HandshakeFailureCloseDoesNotBlock(t *testing.T) { actx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) defer cancel() - start := time.Now() - _, err := fx.AcquireDrpcConn(actx) - require.ErrorIs(t, err, context.DeadlineExceeded) - assert.Less(t, time.Since(start), time.Second) + returnsWithin(t, 3*time.Second, "the blocked close must not run on the caller", func() { + _, err := fx.AcquireDrpcConn(actx) + assert.ErrorIs(t, err, context.DeadlineExceeded) + }) // the stream is still closed once the transport lets it close(release) @@ -903,3 +906,52 @@ func TestPeer_HandedOffConnIsNotTakenByGC(t *testing.T) { fx.mu.Unlock() assert.True(t, active) } + +// scriptedConn is a sub conn whose Invoke returns a set error and whose +// Closed fires only when told +type scriptedConn struct { + closedCh chan struct{} + invokeErr error +} + +func (c *scriptedConn) Close() error { return nil } +func (c *scriptedConn) Closed() <-chan struct{} { return c.closedCh } +func (c *scriptedConn) Unblocked() <-chan struct{} { return nil } +func (c *scriptedConn) NewStream(context.Context, string, drpc.Encoding) (drpc.Stream, error) { + return nil, c.invokeErr +} +func (c *scriptedConn) Invoke(context.Context, string, drpc.Encoding, drpc.Message, drpc.Message) error { + return c.invokeErr +} + +func TestSubConn_ConnLost(t *testing.T) { + errManagerClosed := errors.New("manager closed: Close called") + t.Run("doomed before its close lands", func(t *testing.T) { + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errManagerClosed}} + sc.doomed.Store(true) + err := sc.Invoke(ctx, "/x", nil, nil, nil) + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.ErrorIs(t, err, errManagerClosed) + }) + t.Run("live sub conn passes errors through", func(t *testing.T) { + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errManagerClosed}} + assert.Equal(t, errManagerClosed, sc.Invoke(ctx, "/x", nil, nil, nil)) + }) + t.Run("a coded server reply is never a connection loss", func(t *testing.T) { + closed := make(chan struct{}) + close(closed) + coded := drpcerr.WithCode(errors.New("space is deleted"), 1003) + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: closed, invokeErr: coded}} + err := sc.Invoke(ctx, "/x", nil, nil, nil) + assert.Equal(t, coded, err) + assert.Equal(t, uint64(1003), drpcerr.Code(err)) + }) + t.Run("the caller's own cancellation stays", func(t *testing.T) { + closed := make(chan struct{}) + close(closed) + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: closed, invokeErr: context.Canceled}} + cctx, cancel := context.WithCancel(ctx) + cancel() + assert.Equal(t, context.Canceled, sc.Invoke(cctx, "/x", nil, nil, nil)) + }) +} diff --git a/net/peer/connlost_ext_test.go b/net/peer/connlost_ext_test.go index e993f5481..f32363f31 100644 --- a/net/peer/connlost_ext_test.go +++ b/net/peer/connlost_ext_test.go @@ -136,6 +136,62 @@ func TestPeer_RPCOnDeadYamuxSessionIsConnClosed(t *testing.T) { assert.False(t, errors.Is(err, context.Canceled), "must not look like the caller's cancellation: %v", err) assert.False(t, mcC.IsClosed(), "the session is alive") }) + t.Run("stream receive on a dead session", func(t *testing.T) { + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, silentCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + st, err := dc.NewStream(context.Background(), "/x/y", nil) + require.NoError(t, err) + require.NoError(t, st.MsgSend(&handshakeproto.Proto{Proto: 1}, nil)) + res := make(chan error, 1) + go func() { res <- st.MsgRecv(&handshakeproto.Proto{}, nil) }() + time.Sleep(100 * time.Millisecond) + _ = mcS.Close() + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the receive did not return") + } + assert.ErrorIs(t, err, transport.ErrConnClosed) + assert.False(t, errors.Is(err, context.Canceled), "must not look like the caller's cancellation: %v", err) + // a send after it as well + assert.ErrorIs(t, st.MsgSend(&handshakeproto.Proto{Proto: 1}, nil), transport.ErrConnClosed) + }) + t.Run("stream cancelled by its caller stays canceled", func(t *testing.T) { + mcS, mcC := multiconntest.MultiConnPair( + peer.CtxWithPeerId(context.Background(), "client"), + peer.CtxWithPeerId(context.Background(), "server"), + ) + _, err := peer.NewPeer(mcS, silentCtrl{}) + require.NoError(t, err) + pr, err := peer.NewPeer(mcC, silentCtrl{}) + require.NoError(t, err) + defer pr.Close() + dc, err := pr.AcquireDrpcConn(context.Background()) + require.NoError(t, err) + sctx, cancel := context.WithCancel(context.Background()) + st, err := dc.NewStream(sctx, "/x/y", nil) + require.NoError(t, err) + res := make(chan error, 1) + go func() { res <- st.MsgRecv(&handshakeproto.Proto{}, nil) }() + time.Sleep(100 * time.Millisecond) + cancel() + select { + case err = <-res: + case <-time.After(10 * time.Second): + t.Fatal("the receive did not return") + } + assert.ErrorIs(t, err, context.Canceled) + assert.NotErrorIs(t, err, transport.ErrConnClosed) + }) t.Run("caller cancel stays canceled", func(t *testing.T) { mcS, mcC := multiconntest.MultiConnPair( peer.CtxWithPeerId(context.Background(), "client"), diff --git a/net/peer/peer.go b/net/peer/peer.go index 40feaa185..6b298e964 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -15,6 +15,7 @@ import ( "go.uber.org/zap" "storj.io/drpc" "storj.io/drpc/drpcconn" + "storj.io/drpc/drpcerr" "storj.io/drpc/drpcmanager" "storj.io/drpc/drpcstream" "storj.io/drpc/drpcwire" @@ -134,10 +135,45 @@ func (s *subConn) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, in, } // NewStream reports a stream refused because the sub conn closed as -// transport.ErrConnClosed (see connLost) +// transport.ErrConnClosed (see connLost), and so does the returned stream +// for a send or receive cut short the same way func (s *subConn) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) (drpc.Stream, error) { stream, err := s.ConnUnblocked.NewStream(ctx, rpc, enc) - return stream, s.connLost(ctx, err) + if err != nil { + return nil, s.connLost(ctx, err) + } + return connLostStream{Stream: stream, sc: s, ctx: ctx}, nil +} + +// connLostStream passes a stream's errors through its sub conn's connLost, +// judged against the ctx the caller opened the stream with +type connLostStream struct { + drpc.Stream + sc *subConn + ctx context.Context +} + +func (s connLostStream) MsgSend(msg drpc.Message, enc drpc.Encoding) error { + return s.sc.connLost(s.ctx, s.Stream.MsgSend(msg, enc)) +} + +func (s connLostStream) MsgRecv(msg drpc.Message, enc drpc.Encoding) error { + return s.sc.connLost(s.ctx, s.Stream.MsgRecv(msg, enc)) +} + +func (s connLostStream) CloseSend() error { + return s.sc.connLost(s.ctx, s.Stream.CloseSend()) +} + +// RawWrite forwards the encoding layer's raw write +func (s connLostStream) RawWrite(kind drpcwire.Kind, data []byte) error { + rw, ok := s.Stream.(interface { + RawWrite(kind drpcwire.Kind, data []byte) error + }) + if !ok { + return fmt.Errorf("stream does not support raw writes") + } + return s.sc.connLost(s.ctx, rw.RawWrite(kind, data)) } // connLost reports an error the caller did not cause, returned while this sub @@ -147,9 +183,10 @@ func (s *subConn) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) // shows such an end as context.Canceled (its stand-in for a transport // io.EOF) or as "manager closed"; callers must be able to tell either from // their own cancellation. context.Canceled is kept out of the error chain on -// purpose; other causes stay reachable. +// purpose; other causes stay reachable. A reply carrying a drpc error code is +// the server's answer, never a connection loss, and is returned unchanged. func (s *subConn) connLost(ctx context.Context, err error) error { - if err == nil || ctx.Err() != nil { + if err == nil || ctx.Err() != nil || drpcerr.Code(err) != 0 { return err } select { @@ -500,7 +537,7 @@ func (p *peer) closeAsync(c io.Closer, counted bool) { func (p *peer) acceptLoop() { var exitErr error defer func() { - if exitErr != transport.ErrConnClosed { + if !errors.Is(exitErr, transport.ErrConnClosed) { log.Warn("accept error: close connection", zap.Error(exitErr)) _ = p.MultiConn.Close() } diff --git a/net/peerservice/dialoutcome_test.go b/net/peerservice/dialoutcome_test.go index ad7742a3b..448ed5e77 100644 --- a/net/peerservice/dialoutcome_test.go +++ b/net/peerservice/dialoutcome_test.go @@ -5,6 +5,7 @@ import ( "fmt" "sync/atomic" "testing" + "time" quicgo "github.com/quic-go/quic-go" "github.com/stretchr/testify/assert" @@ -237,7 +238,13 @@ func TestPeerService_CancelledDialReportsNoOutcome(t *testing.T) { fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true).AnyTimes() // quic first; the first attempt is the dial a Flush catches in flight // and only ends with its ctx; the yamux fallback must not be tried - // with the dead ctx (no expectation: gomock fails on a call) + // with the dead ctx + var yamuxDials atomic.Int32 + fx.yamux.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1111").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + yamuxDials.Add(1) + return nil, fmt.Errorf("must not be dialed") + }).AnyTimes() dialStarted := make(chan struct{}) var dials atomic.Int32 fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").DoAndReturn( @@ -257,7 +264,13 @@ func TestPeerService_CancelledDialReportsNoOutcome(t *testing.T) { }() <-dialStarted require.NoError(t, pl.Flush(ctx)) - require.NoError(t, <-got) + select { + case err := <-got: + require.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("the get did not return") + } + assert.Zero(t, yamuxDials.Load(), "no address is tried with a dead ctx") // the cancelled attempt left no trace; the redial reported normally o := stub.only(t) assert.Equal(t, transport.Quic, o.SucceededScheme) @@ -295,6 +308,23 @@ func TestPeerService_CancelledDialReportsNoOutcome(t *testing.T) { o := stub.only(t) assert.Equal(t, transport.Quic, o.SucceededScheme) }) + t.Run("a fallback attempt cut short by the caller is not reported", func(t *testing.T) { + fx, stub := newFixtureWithStubDemotion(t) + defer fx.finish(t) + fx.nodeConf.EXPECT().PeerAddresses(peerId).Return(demotionAddrs, true) + dctx, cancel := context.WithCancel(ctx) + defer cancel() + fx.quic.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1112").Return(nil, &quicgo.HandshakeTimeoutError{}) + fx.yamux.MockTransport.EXPECT().Dial(gomock.Any(), "203.0.113.1:1111").DoAndReturn( + func(ctx context.Context, addr string) (transport.MultiConn, error) { + // the caller gives up while the fallback is being dialed + cancel() + return nil, fmt.Errorf("dial: %w", ctx.Err()) + }) + _, err := fx.Dial(dctx, peerId) + require.ErrorIs(t, err, context.Canceled) + require.Empty(t, stub.outcomes, "a fallback the caller cut short proves nothing about it") + }) t.Run("a dial whose addresses all failed on their own is reported even if ctx ends after", func(t *testing.T) { fx, stub := newFixtureWithStubDemotion(t) defer fx.finish(t) diff --git a/net/pool/pool.go b/net/pool/pool.go index 1c3dfefc9..7b3af9fb5 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -42,9 +42,8 @@ type Pool interface { // GetOneOf searches at least one existing connection in outgoing or creates a new one from a randomly selected id from given list GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error) // AddPeer adds incoming peer to the pool. The pool tracks peers by - // instance; an implementation whose type is not comparable (not a - // pointer) is tracked by id instead, which is only weaker when two - // connections for one id overlap + // instance; an implementation that is not comparable (not a pointer) is + // kept behind a pointer of the pool's own, so it behaves the same AddPeer(ctx context.Context, p peer.Peer) (err error) // Pick checks if a connection with the peer exists, without dialing. // For a peer whose last dial failed it returns the cached dial error. @@ -77,26 +76,29 @@ type caches struct { // is cut short and retried on the current one ctx context.Context cancel context.CancelFunc - // reported holds the peers (see peerKey) whose Closed event Flush - // delivered itself, so their watchers do not report them a second time - // (see Flush); written once, by Flush, under reportedMu + // reported holds the stored objects (see wrap) of the peers whose Closed + // event Flush delivered itself, so their watchers do not report them a + // second time (see Flush); written once, by Flush, under reportedMu reportedMu sync.Mutex - reported map[any]struct{} + reported map[ocache.Object]struct{} } -// peerKey identifies a pooled peer instance in a map: the peer itself when it -// is comparable (a pointer, as every implementation in this module is), else -// its id and direction, which within one pair name a single entry at a time. -// Never panics on a non-comparable implementation: checked on the value, since -// a type with an interface field is comparable while the value may not be. -func peerKey(pr peer.Peer, inbound bool) any { +// pooledPeer is the stored form of a peer whose own value is not comparable: +// the pool tracks peers by instance (map keys, ocache.RemoveSame), so such a +// peer is kept behind this pointer, which is. getPeer unwraps it; the methods +// the pool needs (Id, Close, CloseChan...) are promoted. +type pooledPeer struct { + peer.Peer +} + +// wrap returns the object the pool stores for pr: pr itself when it is +// comparable (a pointer, as every implementation in this module is), else a +// pooledPeer. Only ever called where a peer enters a cache, not on hit paths. +func wrap(pr peer.Peer) ocache.Object { if reflect.ValueOf(pr).Comparable() { return pr } - return struct { - id string - inbound bool - }{pr.Id(), inbound} + return &pooledPeer{Peer: pr} } // bind derives from ctx a context that also ends when the pair stops being @@ -118,11 +120,12 @@ func (c *caches) cache(inbound bool) ocache.OCache { return c.outgoing } -// reportedByFlush reports whether Flush delivered this peer's Closed event -func (c *caches) reportedByFlush(pr peer.Peer, inbound bool) bool { +// reportedByFlush reports whether Flush delivered the Closed event of the +// peer stored as v +func (c *caches) reportedByFlush(v ocache.Object) bool { c.reportedMu.Lock() defer c.reportedMu.Unlock() - _, ok := c.reported[peerKey(pr, inbound)] + _, ok := c.reported[v] return ok } @@ -224,8 +227,9 @@ func (p *pool) fast(id string, touch bool) (pr peer.Peer, missed *caches) { var out ocache.PeekState if v, out = c.peekOutgoing.Peek(id, touch); out != ocache.PeekHit { // a busy entry (loading, or a close that may yet be declined) - // is not a miss: the lookup waits for it - if in != ocache.PeekMiss || out != ocache.PeekMiss || p.current.Load() != c { + // is not a miss: the lookup waits for it. Nor is a cancelled pair + // (replaced, or the pool closed): the lookup tells which. + if in != ocache.PeekMiss || out != ocache.PeekMiss || c.ctx.Err() != nil || p.current.Load() != c { return nil, nil } if touch && p.metrics != nil { @@ -258,7 +262,7 @@ func (p *pool) fast(id string, touch bool) (pr peer.Peer, missed *caches) { // (RemoveSame closes the value itself; a second Close is idempotent). // RemoveSame never touches a replacement installed under the same id. The // returned channel closes once the attempt is done. -func (p *pool) discard(source ocache.OCache, pr peer.Peer) <-chan struct{} { +func (p *pool) discard(source ocache.OCache, stored ocache.Object) <-chan struct{} { done := make(chan struct{}) go func() { defer close(done) @@ -266,7 +270,9 @@ func (p *pool) discard(source ocache.OCache, pr peer.Peer) <-chan struct{} { // pool shutdown: cache.Close evicts whatever is left return } - _, _ = source.RemoveSame(p.closingCtx, pr.Id(), pr) + if pr, err := getPeer(stored); err == nil { + _, _ = source.RemoveSame(p.closingCtx, pr.Id(), stored) + } }() return done } @@ -278,7 +284,11 @@ func (p *pool) discard(source ocache.OCache, pr peer.Peer) <-chan struct{} { // is skipped. It never outlives the peer. c is the pair the peer was // published into: after a Flush that is no longer the current one, and // RemoveSame on it fails fast with ErrClosed. -func (p *pool) evictOnClose(pr peer.Peer, c *caches, inbound bool) { +func (p *pool) evictOnClose(stored ocache.Object, c *caches, inbound bool) { + pr, err := getPeer(stored) + if err != nil { + return + } cache := c.cache(inbound) select { case <-pr.CloseChan(): @@ -301,13 +311,13 @@ func (p *pool) evictOnClose(pr peer.Peer, c *caches, inbound bool) { // Remove only if the cache still holds THIS peer. A newer connection for // the same id may have replaced pr (incoming AddPeer re-add, or outgoing // redial); removing by id alone would close that live replacement. - _, _ = cache.RemoveSame(p.closingCtx, pr.Id(), pr) + _, _ = cache.RemoveSame(p.closingCtx, pr.Id(), stored) // RemoveSame can park behind another closer; re-check so no Closed is // delivered once pool shutdown has begun. Checked after the removal: Flush // marks the peers it saw and reports under reportedMu, and the removal // and Flush's snapshot are ordered by the cache lock, so a peer Flush saw // is marked by the time this runs and one it did not see is reported here - if p.closingCtx.Err() != nil || c.reportedByFlush(pr, inbound) { + if p.closingCtx.Err() != nil || c.reportedByFlush(stored) { return } p.observer.Notify(peerobserver.Event{ @@ -430,7 +440,7 @@ func (p *pool) live(ctx context.Context, source ocache.OCache, v ocache.Object) return pr, nil } select { - case <-p.discard(source, pr): + case <-p.discard(source, v): case <-ctx.Done(): return nil, ctx.Err() } @@ -452,11 +462,12 @@ func (p *pool) live(ctx context.Context, source ocache.OCache, v ocache.Object) // flushes are safe (see the invariants at the top of the file). func (p *pool) Flush(ctx context.Context) error { p.swapMu.Lock() - if p.closed { + old := p.current.Load() + if p.closed || old == nil { + // closed, or never initialised: nothing to replace p.swapMu.Unlock() return nil } - old := p.current.Load() fresh := p.newCaches() // the verdicts must be in the fresh pair before it is published, so this // one read of the old outgoing cache happens before the swap @@ -499,13 +510,14 @@ func (p *pool) Flush(ctx context.Context) error { type snapshotPeer struct { pr peer.Peer + stored ocache.Object inbound bool } // closePeer is the pre-close of one peer (named: the tests tell this close // from the cache's own pass by it) -func closePeer(pr peer.Peer) { - _ = pr.Close() +func closePeer(v ocache.Object) { + _ = v.Close() } // maxPeerClosers caps how many peers closeCaches closes at once @@ -519,15 +531,15 @@ func (c *caches) snapshot(mark bool) (peers []snapshotPeer) { if mark { c.reportedMu.Lock() defer c.reportedMu.Unlock() - c.reported = map[any]struct{}{} + c.reported = map[ocache.Object]struct{}{} } for _, inbound := range []bool{true, false} { c.cache(inbound).ForEach(func(v ocache.Object) (isContinue bool) { - if pr, ok := v.(peer.Peer); ok { + if pr, err := getPeer(v); err == nil { if mark { - c.reported[peerKey(pr, inbound)] = struct{}{} + c.reported[v] = struct{}{} } - peers = append(peers, snapshotPeer{pr: pr, inbound: inbound}) + peers = append(peers, snapshotPeer{pr: pr, stored: v, inbound: inbound}) } return true }) @@ -568,7 +580,7 @@ func closeCaches(c *caches, peers []snapshotPeer) (err error) { <-slots wg.Done() }() - closePeer(sp.pr) + closePeer(sp.stored) }() } }() @@ -583,10 +595,22 @@ func closeCaches(c *caches, peers []snapshotPeer) (err error) { } func (p *pool) getIfActive(ctx context.Context, peerIds []string) peer.Peer { - for _, peerId := range peerIds { - if pr, _ := p.fast(peerId, false); pr != nil { + // when every id missed on one and the same pair there is nothing a + // lookup could find either + var missedAll *caches + for i, peerId := range peerIds { + pr, missed := p.fast(peerId, false) + if pr != nil { return pr } + if i == 0 { + missedAll = missed + } else if missed != missedAll { + missedAll = nil + } + } + if missedAll != nil || len(peerIds) == 0 { + return nil } pr, _ := p.lookup(ctx, func(ctx context.Context, c *caches) (peer.Peer, error) { for _, peerId := range peerIds { @@ -634,8 +658,7 @@ func (p *pool) GetOneOf(ctx context.Context, peerIds []string) (peer.Peer, error return nil, lastErr } -// AddPeer adds an incoming peer. The pool evicts it by instance (see peerKey -// and ocache.RemoveSame for non-comparable implementations). +// AddPeer adds an incoming peer. The pool evicts it by instance (see wrap). func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // Bounds the passes over an entry for the same id that is still there // after this call dealt with it: one another closer holds (its close is @@ -645,13 +668,14 @@ func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // coalesces its flushes). const retries = 3 attempts := 0 + stored := wrap(pr) for { if p.closingCtx.Err() != nil { // shutting down: nothing is evicted any more (discard is a no-op), // so there is nothing to retry towards return ocache.ErrClosed } - c, err := p.addIncoming(pr) + c, err := p.addIncoming(pr, stored) if err != ocache.ErrExists { return err } @@ -692,7 +716,7 @@ func (p *pool) AddPeer(ctx context.Context, pr peer.Peer) error { // current: a hung transport must not stall the accept path. The // incoming cache holds peers only. select { - case <-p.discard(c.incoming, v.(peer.Peer)): + case <-p.discard(c.incoming, v): case <-c.ctx.Done(): case <-ctx.Done(): return ctx.Err() @@ -714,21 +738,27 @@ func (p *pool) waitClosing(ctx context.Context, c *caches, id string) error { // that is about to be closed; Add therefore fails with ErrClosed only when the // pool is closed (Close leaves the closed pair current). Returns ErrExists // without starting a watcher. -func (p *pool) addIncoming(pr peer.Peer) (*caches, error) { +func (p *pool) addIncoming(pr peer.Peer, stored ocache.Object) (*caches, error) { p.swapMu.RLock() defer p.swapMu.RUnlock() c := p.current.Load() - if err := c.incoming.Add(pr.Id(), pr); err != nil { + if err := c.incoming.Add(pr.Id(), stored); err != nil { return c, err } - go p.evictOnClose(pr, c, true) + go p.evictOnClose(stored, c, true) return c, nil } func (p *pool) Pick(ctx context.Context, id string) (pr peer.Peer, err error) { - if pr, _ = p.fast(id, false); pr != nil { + pr, missed := p.fast(id, false) + if pr != nil { return pr, nil } + if missed != nil { + // both caches empty on a pair that stayed current: a lookup would + // find nothing either + return nil, ocache.ErrNotExists + } return p.lookup(ctx, func(ctx context.Context, c *caches) (pr peer.Peer, err error) { // check if connection with peer exist without dial if pr, err = p.pick(ctx, c.incoming, id); err != nil { @@ -759,18 +789,16 @@ func (p *pool) pick(ctx context.Context, source ocache.OCache, id string) (peer. func (p *pool) ProvideStat() any { peerStats := make([]*peer.Stat, 0) c := p.current.Load() - c.outgoing.ForEach(func(v ocache.Object) (isContinue bool) { - if p, ok := v.(peer.StatProvider); ok { - peerStats = append(peerStats, p.ProvideStat()) - } - return true - }) - c.incoming.ForEach(func(v ocache.Object) (isContinue bool) { - if p, ok := v.(peer.StatProvider); ok { - peerStats = append(peerStats, p.ProvideStat()) + collect := func(v ocache.Object) (isContinue bool) { + if pr, err := getPeer(v); err == nil { + if sp, ok := pr.(peer.StatProvider); ok { + peerStats = append(peerStats, sp.ProvideStat()) + } } return true - }) + } + c.outgoing.ForEach(collect) + c.incoming.ForEach(collect) return &poolStats{PeerStats: peerStats} } @@ -786,6 +814,8 @@ var errPeerNotFound = fmt.Errorf("failed to pick connection with peer: peer not func getPeer(val ocache.Object) (pr peer.Peer, err error) { switch v := val.(type) { + case *pooledPeer: + pr = v.Peer case peer.Peer: pr = v case *errObject: diff --git a/net/pool/pool_bench_test.go b/net/pool/pool_bench_test.go index f9ce64e6d..07481a473 100644 --- a/net/pool/pool_bench_test.go +++ b/net/pool/pool_bench_test.go @@ -2,6 +2,7 @@ package pool import ( "context" + "fmt" "sync" "testing" "time" @@ -108,6 +109,16 @@ func BenchmarkPool_GetOneOf(b *testing.B) { benchPool(b, func(ctx context.Context, s Service) error { _, err := s.GetOneOf(ctx, ids); return err }) } +// a Pick for a peer that is not pooled: servers probe connectivity this way +func BenchmarkPool_PickMiss(b *testing.B) { + benchPool(b, func(ctx context.Context, s Service) error { + if _, err := s.Pick(ctx, "absent"); err == nil { + return fmt.Errorf("unexpected hit") + } + return nil + }) +} + // servers pass request contexts, which are cancellable func BenchmarkPool_GetIncomingCancellableCtx(b *testing.B) { rctx, cancel := context.WithTimeout(context.Background(), time.Hour) diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go index 7f08bfcd7..9883bfe00 100644 --- a/net/pool/pool_flush_test.go +++ b/net/pool/pool_flush_test.go @@ -139,7 +139,11 @@ func (c *caches) isClosed() bool { func inCurrent(p *pool, pr peer.Peer) bool { v, err := p.current.Load().incoming.Pick(ctx, pr.Id()) - return err == nil && v == ocache.Object(pr) + if err != nil { + return false + } + got, err := getPeer(v) + return err == nil && got == pr } // startFlusher flushes the pool every period on a background goroutine until @@ -436,7 +440,7 @@ func TestPool_FlushSwap(t *testing.T) { pr, err := fx.Get(gctx, "p1") require.NoError(t, err) require.NotNil(t, pr) - require.Less(t, time.Since(start), 500*time.Millisecond, "waited on the stale dial") + require.Less(t, time.Since(start), 2*time.Second, "waited on the stale dial") // the old pair cancelled the stale dial and the first caller // redialed too, sharing the fresh peer require.NoError(t, <-first) @@ -899,7 +903,7 @@ func TestPool_FlushSwap(t *testing.T) { res := <-first require.NoError(t, res.err) require.False(t, res.pr.IsClosed()) - require.Less(t, time.Since(start), 500*time.Millisecond, "pre-flush Get waited out the dead dial") + require.Less(t, time.Since(start), 2*time.Second, "pre-flush Get waited out the dead dial") require.True(t, cancelled.Load()) require.Equal(t, int32(2), dials.Load()) }) @@ -923,7 +927,7 @@ func TestPool_FlushSwap(t *testing.T) { select { case <-ctx.Done(): cancelled.Store(true) - case <-time.After(400 * time.Millisecond): + case <-time.After(10 * time.Second): } return late, nil } @@ -942,7 +946,7 @@ func TestPool_FlushSwap(t *testing.T) { require.NoError(t, fx.Flush(ctx)) pr := <-first require.NotSame(t, late, pr) - require.Less(t, time.Since(start), 300*time.Millisecond, "pre-flush Get waited out the stale dial") + require.Less(t, time.Since(start), 2*time.Second, "pre-flush Get waited out the stale dial") require.True(t, cancelled.Load()) // published into the old outgoing cache, which closed at once require.Eventually(t, late.IsClosed, time.Second, 10*time.Millisecond) @@ -993,7 +997,7 @@ func TestPool_FlushSwap(t *testing.T) { select { case pr := <-got: require.NotSame(t, old, pr) - require.Less(t, time.Since(start), 500*time.Millisecond) + require.Less(t, time.Since(start), 2*time.Second) case <-time.After(2 * time.Second): t.Fatal("Get stayed parked on the replaced pair") } @@ -2082,22 +2086,21 @@ func TestPool_FlushSwap(t *testing.T) { require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) require.Eventually(t, func() bool { return len(obs.kindsFor("a")) == 1 }, time.Second, 10*time.Millisecond) require.Equal(t, []peerobserver.Kind{peerobserver.KindClosed}, obs.kindsFor("a")) - // the replacement path through RemoveSame: no panic, the old one is - // closed and reported; by id the old one's watcher may take the - // replacement down with it, so its fate is not asserted + // the replacement path through RemoveSame: the old one is closed and + // reported, the replacement survives its stale watcher b1, b2 := newValuePeer("b"), newValuePeer("b") require.NoError(t, fx.AddPeer(ctx, b1)) require.NoError(t, fx.AddPeer(ctx, b2)) require.True(t, b1.IsClosed()) - require.Eventually(t, func() bool { return len(obs.kindsFor("b")) >= 1 }, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return len(obs.kindsFor("b")) == 1 }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return b2.IsClosed() || p.current.Load().incoming.Len() != 1 }, 100*time.Millisecond, 10*time.Millisecond) b2.close() - require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 }, time.Second, 10*time.Millisecond) + require.Eventually(t, func() bool { return p.current.Load().incoming.Len() == 0 && len(obs.kindsFor("b")) == 2 }, time.Second, 10*time.Millisecond) // the discard path: a dead outgoing peer found by a lookup (added - // without a watcher, so no second instance of the id is live while - // the lookup evicts it and redials) + // without a watcher, stored as the pool would store it) dead := newValuePeer("c") dead.close() - require.NoError(t, p.current.Load().outgoing.Add("c", dead)) + require.NoError(t, p.current.Load().outgoing.Add("c", wrap(dead))) fresh := newValuePeer("c") fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { return fresh, nil @@ -2112,6 +2115,59 @@ func TestPool_FlushSwap(t *testing.T) { require.Eventually(t, fresh.IsClosed, time.Second, 10*time.Millisecond) require.Never(t, func() bool { return len(obs.getClosed()) > 4 }, 200*time.Millisecond, 10*time.Millisecond) }) + t.Run("a non-comparable peer's stale watcher never touches its replacement and flush marks never hide it", func(t *testing.T) { + obs := &poolEventRecorder{} + fx := newFixtureWithObserver(t, obs) + defer fx.Finish() + p := fx.Service.(*poolService).pool + // the old instance's removal is held while the replacement lands and + // a flush marks the replacement: the stale watcher must neither close + // the replacement nor be silenced by the mark on it + // a CAS gate, not sync.Once: the replacement's own removal goes + // through the same seam and must not queue behind the held one + var armed atomic.Bool + armed.Store(true) + inRemove := make(chan struct{}) + releaseRemove, doReleaseRemove := newRelease() + defer doReleaseRemove() + installIncoming(t, fx, func(inner ocache.OCache, peek ocache.Peeker) ocache.OCache { + return &hookedCache{OCache: inner, peek: peek, onRemoveSame: func() { + if armed.CompareAndSwap(true, false) { + close(inRemove) + <-releaseRemove + } + }} + }) + a := newValuePeer("x") + require.NoError(t, fx.AddPeer(ctx, a)) + a.close() + // a's watcher is inside its RemoveSame now + <-inRemove + b := newValuePeer("x") + require.NoError(t, fx.AddPeer(ctx, b)) + require.NoError(t, fx.Flush(ctx)) + // Flush reported b (and closes it); a's watcher still owes a's event + require.Len(t, obs.kindsFor("x"), 1) + doReleaseRemove() + require.Eventually(t, func() bool { return len(obs.kindsFor("x")) == 2 }, time.Second, 10*time.Millisecond) + require.Never(t, func() bool { return len(obs.kindsFor("x")) > 2 }, 200*time.Millisecond, 10*time.Millisecond) + + // and without a flush: the stale watcher's removal leaves the live + // replacement in place + c1 := newValuePeer("y") + require.NoError(t, fx.AddPeer(ctx, c1)) + c2 := newValuePeer("y") + require.NoError(t, fx.AddPeer(ctx, c2)) + require.True(t, c1.IsClosed()) + require.Eventually(t, func() bool { return len(obs.kindsFor("y")) == 1 }, time.Second, 10*time.Millisecond) + require.False(t, c2.IsClosed()) + v, err := p.current.Load().incoming.Pick(ctx, "y") + require.NoError(t, err) + got, err := getPeer(v) + require.NoError(t, err) + require.Equal(t, "y", got.Id()) + require.False(t, got.IsClosed()) + }) t.Run("a lookup parked in an eviction when the pool closes gets ErrClosed", func(t *testing.T) { // the pair stays current but is cancelled by Close: the ctx error the // eviction's select returns must surface as ErrClosed @@ -2334,6 +2390,22 @@ func TestPool_FlushSwap(t *testing.T) { require.Eventually(t, b.IsClosed, time.Second, 10*time.Millisecond) require.Never(t, func() bool { return len(obs.getClosed()) > 3 }, 200*time.Millisecond, 10*time.Millisecond) }) + t.Run("a miss on both caches allocates nothing for Pick and GetOneOf's scan", func(t *testing.T) { + fx := newFixture(t) + defer fx.Finish() + p := fx.Service.(*poolService).pool + require.Zero(t, testing.AllocsPerRun(100, func() { + if _, err := fx.Pick(ctx, "absent"); err == nil { + t.Error("unexpected hit") + } + })) + ids := []string{"absent1", "absent2"} + require.Zero(t, testing.AllocsPerRun(100, func() { + if p.getIfActive(ctx, ids) != nil { + t.Error("unexpected hit") + } + })) + }) t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { obs := &poolEventRecorder{} fx := newFixtureWithObserver(t, obs) diff --git a/net/pool/poolservice.go b/net/pool/poolservice.go index 821e8f047..ee964deb9 100644 --- a/net/pool/poolservice.go +++ b/net/pool/poolservice.go @@ -115,7 +115,8 @@ func (p *poolService) newCaches(outgoingMetrics, incomingMetrics ocache.Option) return value, err } if pr, ok := value.(peer.Peer); ok { - go p.pool.evictOnClose(pr, c, false) + value = wrap(pr) + go p.pool.evictOnClose(value, c, false) } return value, nil }, @@ -165,12 +166,16 @@ func (p *pool) Close(ctx context.Context) (err error) { return nil } p.closed = true + cur := p.current.Load() p.swapMu.Unlock() + if cur == nil { + // never initialised: nothing to close + return nil + } if p.closingCancel != nil { p.closingCancel() } p.statService.RemoveProvider(p) - cur := p.current.Load() // lookups blocked on the current pair fail now with ErrClosed (see lookup) cur.cancel() done := make(chan error, 1) diff --git a/net/transport/transport.go b/net/transport/transport.go index 927c41b6d..f03713c18 100644 --- a/net/transport/transport.go +++ b/net/transport/transport.go @@ -6,6 +6,8 @@ import ( "errors" "net" "time" + + "storj.io/drpc/drpcerr" ) var ( @@ -25,11 +27,10 @@ func (connClosedError) Unwrap() error { return net.ErrClosed } // NewConnClosedError wraps cause, an error a transport got because the whole // connection went away, so that it matches ErrConnClosed while errors.As -// still finds the original error. A cause that is already wrapped is -// returned as is. +// still finds the original error and drpcerr.Code still finds its code. A +// cause that already matches ErrConnClosed is returned as is. func NewConnClosedError(cause error) error { - var already connClosedCauseError - if errors.As(cause, &already) { + if errors.Is(cause, ErrConnClosed) { return cause } return connClosedCauseError{cause: cause} @@ -47,6 +48,12 @@ func (e connClosedCauseError) Unwrap() []error { return []error{ErrConnClosed, e.cause} } +// Code exposes the cause's drpc error code: drpcerr.Code does not follow a +// multi-error Unwrap +func (e connClosedCauseError) Code() uint64 { + return drpcerr.Code(e.cause) +} + const ( Yamux = "yamux" Quic = "quic" diff --git a/net/transport/transport_test.go b/net/transport/transport_test.go index 96ac071ca..7e65337fa 100644 --- a/net/transport/transport_test.go +++ b/net/transport/transport_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/stretchr/testify/assert" + "storj.io/drpc/drpcerr" ) func TestErrConnClosed(t *testing.T) { @@ -31,3 +32,13 @@ func TestNewConnClosedError(t *testing.T) { // idempotent assert.Equal(t, err, NewConnClosedError(err)) } + +func TestNewConnClosedError_CodeAndBare(t *testing.T) { + // a drpc error code survives the wrapping + err := NewConnClosedError(drpcerr.WithCode(errors.New("coded"), 7)) + assert.ErrorIs(t, err, ErrConnClosed) + assert.Equal(t, uint64(7), drpcerr.Code(err)) + assert.Zero(t, drpcerr.Code(NewConnClosedError(errors.New("plain")))) + // a bare ErrConnClosed is not wrapped again + assert.Equal(t, ErrConnClosed, NewConnClosedError(ErrConnClosed)) +} diff --git a/net/transport/yamux/conn_test.go b/net/transport/yamux/conn_test.go index a8a5f79aa..07f0aab67 100644 --- a/net/transport/yamux/conn_test.go +++ b/net/transport/yamux/conn_test.go @@ -404,3 +404,45 @@ func TestYamuxStream_WriteTimeoutOnLiveSession(t *testing.T) { assert.False(t, errors.Is(err, transport.ErrConnClosed), "a slow live peer is not a dead connection") assert.False(t, client.IsClosed()) } + +// BenchmarkYamuxConn_Open compares an open with a context that can never end +// (Session.Open called directly, as before the ctx-bound Open) with one that +// can be cancelled (helper goroutine and result channel) +func BenchmarkYamuxConn_Open(b *testing.B) { + run := func(b *testing.B, octx context.Context) { + cc, sc := net.Pipe() + conf := yamux.DefaultConfig() + conf.LogOutput = io.Discard + client, err := yamux.Client(cc, conf) + require.NoError(b, err) + server, err := yamux.Server(sc, conf) + require.NoError(b, err) + defer client.Close() + defer server.Close() + go func() { + for { + s, aErr := server.Accept() + if aErr != nil { + return + } + _ = s.Close() + } + }() + mc := NewMultiConn(context.Background(), connutil.NewLastUsageConn(cc), "pipe", client) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + conn, oErr := mc.Open(octx) + if oErr != nil { + b.Fatal(oErr) + } + _ = conn.Close() + } + } + b.Run("background ctx", func(b *testing.B) { run(b, context.Background()) }) + b.Run("cancellable ctx", func(b *testing.B) { + cctx, cancel := context.WithCancel(context.Background()) + defer cancel() + run(b, cctx) + }) +} From 4392a71d58bc3bc617fff7aef7e45f615ad18cf2 Mon Sep 17 00:00:00 2001 From: Roman Khafizianov Date: Fri, 2 Oct 2026 22:33:02 +0200 Subject: [PATCH 7/7] net/peer, net/pool: pass stream io.EOF through; unwrap fast-path peers - connLost rewrites only cancellation, drpc closed errors and closed-network errors on a closed/doomed sub conn; io.EOF (exact), coded and application errors pass through, so `err == io.EOF` works - fast() returns the peer, never the pool's pooledPeer holder - TestPool_FlushStorm: flush count follows the work done, not a background ticker that -race could starve --- net/peer/closeasync_test.go | 35 ++++++++++++++++++--- net/peer/connlost_ext_test.go | 11 +++++-- net/peer/peer.go | 43 ++++++++++++++++++++------ net/pool/pool.go | 6 ++-- net/pool/pool_flush_test.go | 58 ++++++++++++++++++++++++++++++----- 5 files changed, 125 insertions(+), 28 deletions(-) diff --git a/net/peer/closeasync_test.go b/net/peer/closeasync_test.go index bf1b8a496..d40bcdfe1 100644 --- a/net/peer/closeasync_test.go +++ b/net/peer/closeasync_test.go @@ -925,17 +925,35 @@ func (c *scriptedConn) Invoke(context.Context, string, drpc.Encoding, drpc.Messa } func TestSubConn_ConnLost(t *testing.T) { - errManagerClosed := errors.New("manager closed: Close called") + errClosed := drpc.ClosedError.New("closed") t.Run("doomed before its close lands", func(t *testing.T) { - sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errManagerClosed}} + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errClosed}} sc.doomed.Store(true) err := sc.Invoke(ctx, "/x", nil, nil, nil) assert.ErrorIs(t, err, transport.ErrConnClosed) - assert.ErrorIs(t, err, errManagerClosed) + assert.ErrorIs(t, err, errClosed) }) t.Run("live sub conn passes errors through", func(t *testing.T) { - sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errManagerClosed}} - assert.Equal(t, errManagerClosed, sc.Invoke(ctx, "/x", nil, nil, nil)) + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: make(chan struct{}), invokeErr: errClosed}} + assert.Equal(t, errClosed, sc.Invoke(ctx, "/x", nil, nil, nil)) + }) + t.Run("other errors are untouched on a closed sub conn", func(t *testing.T) { + closed := make(chan struct{}) + close(closed) + appErr := errors.New("application error") + for _, cause := range []error{io.EOF, appErr} { + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: closed, invokeErr: cause}} + assert.True(t, sc.Invoke(ctx, "/x", nil, nil, nil) == cause, "%v must be returned as is", cause) + } + }) + t.Run("a stream's normal end stays exactly io.EOF", func(t *testing.T) { + closed := make(chan struct{}) + close(closed) + sc := &subConn{ConnUnblocked: &scriptedConn{closedCh: closed}} + st := connLostStream{Stream: eofStream{}, sc: sc, ctx: ctx} + assert.True(t, st.MsgRecv(nil, nil) == io.EOF, "callers compare with ==") + assert.True(t, st.MsgSend(nil, nil) == io.EOF) + assert.True(t, st.CloseSend() == io.EOF) }) t.Run("a coded server reply is never a connection loss", func(t *testing.T) { closed := make(chan struct{}) @@ -955,3 +973,10 @@ func TestSubConn_ConnLost(t *testing.T) { assert.Equal(t, context.Canceled, sc.Invoke(cctx, "/x", nil, nil, nil)) }) } + +// eofStream is a drpc stream the server has ended normally +type eofStream struct{ drpc.Stream } + +func (eofStream) MsgRecv(drpc.Message, drpc.Encoding) error { return io.EOF } +func (eofStream) MsgSend(drpc.Message, drpc.Encoding) error { return io.EOF } +func (eofStream) CloseSend() error { return io.EOF } diff --git a/net/peer/connlost_ext_test.go b/net/peer/connlost_ext_test.go index f32363f31..66fcae209 100644 --- a/net/peer/connlost_ext_test.go +++ b/net/peer/connlost_ext_test.go @@ -160,10 +160,15 @@ func TestPeer_RPCOnDeadYamuxSessionIsConnClosed(t *testing.T) { case <-time.After(10 * time.Second): t.Fatal("the receive did not return") } - assert.ErrorIs(t, err, transport.ErrConnClosed) + // drpc ends the receive either as a cancellation, reported as + // ErrConnClosed, or with a plain io.EOF, the same value a normal end of + // stream gives, which is left as is since callers compare it with ==. + // Never a bare context.Canceled. + assert.True(t, err == io.EOF || errors.Is(err, transport.ErrConnClosed), "got %v", err) assert.False(t, errors.Is(err, context.Canceled), "must not look like the caller's cancellation: %v", err) - // a send after it as well - assert.ErrorIs(t, st.MsgSend(&handshakeproto.Proto{Proto: 1}, nil), transport.ErrConnClosed) + err = st.MsgSend(&handshakeproto.Proto{Proto: 1}, nil) + assert.True(t, err == io.EOF || errors.Is(err, transport.ErrConnClosed), "got %v", err) + assert.False(t, errors.Is(err, context.Canceled), "got %v", err) }) t.Run("stream cancelled by its caller stays canceled", func(t *testing.T) { mcS, mcC := multiconntest.MultiConnPair( diff --git a/net/peer/peer.go b/net/peer/peer.go index 6b298e964..3ea9f69e3 100644 --- a/net/peer/peer.go +++ b/net/peer/peer.go @@ -176,17 +176,21 @@ func (s connLostStream) RawWrite(kind drpcwire.Kind, data []byte) error { return s.sc.connLost(s.ctx, rw.RawWrite(kind, data)) } -// connLost reports an error the caller did not cause, returned while this sub -// conn is closed or doomed, as transport.ErrConnClosed: the RPC ended because -// the sub conn did, whether the whole connection died, the remote ended just -// this sub stream on a live session, or gc or a release closed it here. drpc -// shows such an end as context.Canceled (its stand-in for a transport -// io.EOF) or as "manager closed"; callers must be able to tell either from -// their own cancellation. context.Canceled is kept out of the error chain on -// purpose; other causes stay reachable. A reply carrying a drpc error code is -// the server's answer, never a connection loss, and is returned unchanged. +// connLost reports an error that says the sub conn ended under the caller, +// returned while this sub conn is closed or doomed, as +// transport.ErrConnClosed: the RPC ended because the sub conn did, whether +// the whole connection died, the remote ended just this sub stream on a live +// session, or gc or a release closed it here. drpc shows such an end as +// context.Canceled (its stand-in for a transport io.EOF), as a "manager +// closed" or drpc.ClosedError, or as a transport error matching +// net.ErrClosed; callers must be able to tell those from their own +// cancellation. Every other error is returned untouched, even on a closed sub +// conn: a plain io.EOF (the server ended the stream normally; callers compare +// it with ==), a reply carrying a drpc error code, any application error. +// context.Canceled is kept out of the error chain on purpose; other causes +// stay reachable. func (s *subConn) connLost(ctx context.Context, err error) error { - if err == nil || ctx.Err() != nil || drpcerr.Code(err) != 0 { + if err == nil || err == io.EOF || ctx.Err() != nil || drpcerr.Code(err) != 0 || !subConnEnded(err) { return err } select { @@ -202,6 +206,25 @@ func (s *subConn) connLost(ctx context.Context, err error) error { return transport.NewConnClosedError(err) } +// drpcManagerClosed is the name of drpc's (unexported) error class for a +// terminated manager +const drpcManagerClosed = "manager closed" + +// subConnEnded reports whether err is how drpc or the transport report the +// sub conn ending +func subConnEnded(err error) bool { + if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) || drpc.ClosedError.Has(err) { + return true + } + var named interface{ Name() (string, bool) } + if errors.As(err, &named) { + if name, ok := named.Name(); ok && name == drpcManagerClosed { + return true + } + } + return false +} + type peer struct { id string diff --git a/net/pool/pool.go b/net/pool/pool.go index 7b3af9fb5..078839c0d 100644 --- a/net/pool/pool.go +++ b/net/pool/pool.go @@ -238,8 +238,10 @@ func (p *pool) fast(id string, touch bool) (pr peer.Peer, missed *caches) { return nil, c } } - pr, isPeer := v.(peer.Peer) - if !isPeer || pr.IsClosed() || p.current.Load() != c { + // getPeer, not a type assertion: a stored pooledPeer satisfies peer.Peer + // too and must not reach the caller + pr, err := getPeer(v) + if err != nil || pr.IsClosed() || p.current.Load() != c { return nil, nil } if m := p.metrics; m != nil { diff --git a/net/pool/pool_flush_test.go b/net/pool/pool_flush_test.go index 9883bfe00..0e7cc2f5e 100644 --- a/net/pool/pool_flush_test.go +++ b/net/pool/pool_flush_test.go @@ -2406,6 +2406,43 @@ func TestPool_FlushSwap(t *testing.T) { } })) }) + t.Run("a non-comparable peer comes back as itself from every lookup", func(t *testing.T) { + // the pool stores it behind its own pointer; no path may return that + fx := newFixture(t) + defer fx.Finish() + in := newValuePeer("in") + require.NoError(t, fx.AddPeer(ctx, in)) + out := newValuePeer("out") + fx.Dialer.dial = func(ctx context.Context, peerId string) (peer.Peer, error) { + return out, nil + } + isValue := func(pr peer.Peer) { + _, ok := pr.(valuePeer) + require.True(t, ok, "got %T", pr) + } + for i := 0; i < 2; i++ { // the second round takes the fast path + pr, err := fx.Get(ctx, "in") + require.NoError(t, err) + isValue(pr) + pr, err = fx.Get(ctx, "out") + require.NoError(t, err) + isValue(pr) + pr, err = fx.Pick(ctx, "in") + require.NoError(t, err) + isValue(pr) + pr, err = fx.Pick(ctx, "out") + require.NoError(t, err) + isValue(pr) + pr, err = fx.GetOneOf(ctx, []string{"absent", "out"}) + require.NoError(t, err) + isValue(pr) + pr, err = fx.GetOneOf(ctx, []string{"absent", "in"}) + require.NoError(t, err) + isValue(pr) + } + stats := fx.Service.(*poolService).ProvideStat().(*poolStats) + require.Empty(t, stats.PeerStats, "valuePeer is no StatProvider") + }) t.Run("connected and closed pairing for flushed peers", func(t *testing.T) { obs := &poolEventRecorder{} fx := newFixtureWithObserver(t, obs) @@ -2457,20 +2494,25 @@ func TestPool_FlushStorm(t *testing.T) { time.Sleep(200 * time.Microsecond) return &taggedPeer{testPeer: newTestPeer(peerId), tag: tr.seqOf(p.current.Load())}, nil } - // the workers run until the flusher has done wantFlushes (about - // 300ms on an idle machine; with GOMAXPROCS=1 the busy workers starve - // it, so a fixed duration would not do) - const wantFlushes = 60 + // The workers drive the flushes themselves (one every 20 Gets each) + // on top of a 1ms background flusher, so the number of swaps the + // Gets race is proportional to the work and cannot be starved by the + // scheduler, whatever the machine load (a busy -race run at + // GOMAXPROCS=1 used to starve a background flusher). + const perWorker, flushEvery = 200, 20 flushes, stopFlusher := startFlusher(t, fx, time.Millisecond) defer stopFlusher() - cap := time.Now().Add(20 * time.Second) var stale, failed, ok atomic.Int32 var wg sync.WaitGroup for w := 0; w < 8; w++ { wg.Add(1) go func() { defer wg.Done() - for i := 0; flushes.Load() < wantFlushes && time.Now().Before(cap); i++ { + for i := 0; i < perWorker; i++ { + if i%flushEvery == flushEvery-1 { + assert.NoError(t, fx.Flush(ctx)) + flushes.Add(1) + } before := tr.seqOf(p.current.Load()) gctx, cancel := context.WithTimeout(ctx, 5*time.Second) pr, err := fx.Get(gctx, "a") @@ -2500,8 +2542,8 @@ func TestPool_FlushStorm(t *testing.T) { wg.Wait() stopFlusher() t.Logf("flushes=%d ok=%d failed=%d stale=%d", flushes.Load(), ok.Load(), failed.Load(), stale.Load()) - require.GreaterOrEqual(t, flushes.Load(), int64(wantFlushes)) - require.Greater(t, ok.Load(), int32(100)) + require.GreaterOrEqual(t, flushes.Load(), int64(8*perWorker/flushEvery)) + require.Equal(t, int32(8*perWorker), ok.Load()) require.Zero(t, failed.Load()) require.Zero(t, stale.Load()) })