From aace5c268b369633367a6e6c14483cf9d1d34f2c Mon Sep 17 00:00:00 2001 From: Shubham Dhama Date: Wed, 22 Jul 2026 22:09:48 +0530 Subject: [PATCH] drpcmanager: derive shutdown waiting from active streams Each manageStream goroutine already corresponds to one activeStreams entry, so keeping a separate WaitGroup duplicated lifecycle state and forced Add to reach into Manager synchronization. Keep canceled streams registered until their manage goroutines remove them, then close a completion channel once shutdown has begun and the registry is empty. This preserves the atomic admission-versus-close gate while letting Manager.Close wait on the registry itself. Co-authored-by: Cursor --- drpcmanager/active_streams.go | 45 +++++++++++++++++++----------- drpcmanager/active_streams_test.go | 38 +++++++++++++++++++------ drpcmanager/manager.go | 10 ++----- drpcmanager/manager_test.go | 6 ++-- 4 files changed, 64 insertions(+), 35 deletions(-) diff --git a/drpcmanager/active_streams.go b/drpcmanager/active_streams.go index 62a5c63..f985457 100644 --- a/drpcmanager/active_streams.go +++ b/drpcmanager/active_streams.go @@ -16,20 +16,20 @@ type activeStreams struct { streams map[uint64]*drpcstream.Stream closed bool closeErr error + done chan struct{} + doneOnce sync.Once } func newActiveStreams() *activeStreams { return &activeStreams{ streams: make(map[uint64]*drpcstream.Stream), + done: make(chan struct{}), } } // Add adds a stream. It returns an error if the collection is closed or if a -// stream with the same ID already exists. If wg is non-nil, it is incremented -// under the same lock that checks closed, so wg.Add and the closed check are -// atomic with respect to Close (which sets closed before Manager.Close calls -// wg.Wait). -func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream, wg *sync.WaitGroup) error { +// stream with the same ID already exists. +func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream) error { if stream == nil { return managerClosed.New("stream can't be nil") } @@ -44,21 +44,17 @@ func (r *activeStreams) Add(id uint64, stream *drpcstream.Stream, wg *sync.WaitG return managerClosed.New("duplicate stream id") } r.streams[id] = stream - if wg != nil { - wg.Add(1) - } return nil } -// Remove removes a stream. It is a no-op if the stream is not present or if -// the collection has been closed. +// Remove removes a stream. If the collection is closed, removing the last +// stream unblocks Wait. func (r *activeStreams) Remove(id uint64) { r.mu.Lock() defer r.mu.Unlock() - if r.streams != nil { - delete(r.streams, id) - } + delete(r.streams, id) + r.signalDone() } // Get returns the stream for the given ID and whether it was found. @@ -73,17 +69,32 @@ func (r *activeStreams) Get(id uint64) (*drpcstream.Stream, bool) { return s, ok } -// Close cancels all active streams with the given error, clears the -// collection, and marks it as closed to prevent future Add calls. +// Close cancels all active streams with the given error and marks the +// collection as closed to prevent future Add calls. Streams remain registered +// until their manage goroutines exit and remove them. func (r *activeStreams) Close(err error) { r.mu.Lock() defer r.mu.Unlock() r.closed = true r.closeErr = err - for id, s := range r.streams { + for _, s := range r.streams { s.Cancel(err) - delete(r.streams, id) + } + r.signalDone() +} + +// Wait blocks until Close has been called and every registered stream has been +// removed. +func (r *activeStreams) Wait() { + <-r.done +} + +// signalDone closes done once shutdown has begun and no streams remain. The +// caller must hold r.mu. +func (r *activeStreams) signalDone() { + if r.closed && len(r.streams) == 0 { + r.doneOnce.Do(func() { close(r.done) }) } } diff --git a/drpcmanager/active_streams_test.go b/drpcmanager/active_streams_test.go index fba38d5..5b61bb2 100644 --- a/drpcmanager/active_streams_test.go +++ b/drpcmanager/active_streams_test.go @@ -29,7 +29,7 @@ func TestActiveStreams_AddAndGet(t *testing.T) { streams := newActiveStreams() s := testStream(t, 1) - assert.NoError(t, streams.Add(1, s, nil)) + assert.NoError(t, streams.Add(1, s)) got, ok := streams.Get(1) assert.That(t, ok) @@ -48,7 +48,7 @@ func TestActiveStreams_Remove(t *testing.T) { streams := newActiveStreams() s := testStream(t, 1) - assert.NoError(t, streams.Add(1, s, nil)) + assert.NoError(t, streams.Add(1, s)) assert.Equal(t, streams.Len(), 1) streams.Remove(1) @@ -70,8 +70,8 @@ func TestActiveStreams_DuplicateAdd(t *testing.T) { s1 := testStream(t, 1) s2 := testStream(t, 1) - assert.NoError(t, streams.Add(1, s1, nil)) - assert.Error(t, streams.Add(1, s2, nil)) + assert.NoError(t, streams.Add(1, s1)) + assert.Error(t, streams.Add(1, s2)) // original stream is still present got, ok := streams.Get(1) @@ -83,14 +83,14 @@ func TestActiveStreams_AddAfterClose(t *testing.T) { streams := newActiveStreams() streams.Close(errors.New("closed")) - err := streams.Add(1, testStream(t, 1), nil) + err := streams.Add(1, testStream(t, 1)) assert.Error(t, err) } func TestActiveStreams_RemoveAfterClose(t *testing.T) { streams := newActiveStreams() s := testStream(t, 1) - assert.NoError(t, streams.Add(1, s, nil)) + assert.NoError(t, streams.Add(1, s)) streams.Close(errors.New("closed")) @@ -98,14 +98,36 @@ func TestActiveStreams_RemoveAfterClose(t *testing.T) { streams.Remove(1) } +func TestActiveStreams_Wait(t *testing.T) { + streams := newActiveStreams() + assert.NoError(t, streams.Add(1, testStream(t, 1))) + + streams.Close(errors.New("closed")) + select { + case <-streams.done: + t.Fatal("Wait completed before the active stream was removed") + default: + } + + streams.Remove(1) + streams.Wait() +} + +func TestActiveStreams_WaitWithoutStreams(t *testing.T) { + streams := newActiveStreams() + + streams.Close(errors.New("closed")) + streams.Wait() +} + func TestActiveStreams_Len(t *testing.T) { streams := newActiveStreams() assert.Equal(t, streams.Len(), 0) - assert.NoError(t, streams.Add(1, testStream(t, 1), nil)) + assert.NoError(t, streams.Add(1, testStream(t, 1))) assert.Equal(t, streams.Len(), 1) - assert.NoError(t, streams.Add(2, testStream(t, 2), nil)) + assert.NoError(t, streams.Add(2, testStream(t, 2))) assert.Equal(t, streams.Len(), 2) streams.Remove(1) diff --git a/drpcmanager/manager.go b/drpcmanager/manager.go index c602e0e..724ad3a 100644 --- a/drpcmanager/manager.go +++ b/drpcmanager/manager.go @@ -7,7 +7,6 @@ import ( "context" "fmt" "io" - "sync" "sync/atomic" "github.com/zeebo/errs" @@ -71,8 +70,6 @@ type Manager struct { // next client stream ID, incremented atomically lastStreamID atomic.Uint64 - wg sync.WaitGroup // tracks active manageStream goroutines - // streams tracks active streams. streams *activeStreams recvPool *drpcstream.BufferPool @@ -324,7 +321,7 @@ func (m *Manager) newStream(ctx context.Context, sid uint64, kind drpc.StreamKin stream := drpcstream.NewWithOptions(ctx, sid, m.wr, m.recvPool, opts) - if err := m.streams.Add(sid, stream, &m.wg); err != nil { + if err := m.streams.Add(sid, stream); err != nil { return nil, err } @@ -342,13 +339,12 @@ func (m *Manager) newStream(ctx context.Context, sid uint64, kind drpc.StreamKin // manageStream watches the context and the stream and returns when the stream // is finished, canceling the stream if the context is canceled. func (m *Manager) manageStream(ctx context.Context, stream *drpcstream.Stream) { - defer m.wg.Done() + defer m.streams.Remove(stream.ID()) defer func() { if m.metrics.ShouldRecord() { m.metrics.StreamsTerminated.Inc(1) } }() - defer m.streams.Remove(stream.ID()) select { case <-stream.Finished(): @@ -392,7 +388,7 @@ func (m *Manager) Close() error { m.terminate(drpc.ClosedError.Wrap(managerClosed.New("Close called"))) <-m.wr.Done() // wait for writer goroutine to exit - m.wg.Wait() // wait for all stream goroutines + m.streams.Wait() // wait for all stream goroutines m.sigs.read.Wait() // wait for reader goroutine to exit m.sigs.tport.Wait() diff --git a/drpcmanager/manager_test.go b/drpcmanager/manager_test.go index 6c6df48..b49cc92 100644 --- a/drpcmanager/manager_test.go +++ b/drpcmanager/manager_test.go @@ -591,9 +591,9 @@ func TestManager_ServerClientHangupCancels(t *testing.T) { // TestManager_ConcurrentCloseAndNewClientStream exercises the race between // Manager.Close (terminate → stop writer, close transport, close streams) and -// Manager.NewClientStream (newStream → create stream, add to streams, wg.Add). -// Under -race this would fail without wg.Add being atomic with the closed -// check in activeStreams.Add. +// Manager.NewClientStream (newStream → create stream → add to streams). The +// activeStreams lock ensures a stream is either registered before Close starts +// waiting for the collection to drain or rejected after shutdown begins. func TestManager_ConcurrentCloseAndNewClientStream(t *testing.T) { for i := 0; i < 100; i++ { cconn, sconn := net.Pipe()