Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 28 additions & 17 deletions drpcmanager/active_streams.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Expand All @@ -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.
Expand All @@ -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) })
}
}

Expand Down
38 changes: 30 additions & 8 deletions drpcmanager/active_streams_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -83,29 +83,51 @@ 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"))

// must not panic
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)
Expand Down
10 changes: 3 additions & 7 deletions drpcmanager/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@ import (
"context"
"fmt"
"io"
"sync"
"sync/atomic"

"github.com/zeebo/errs"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
}

Expand All @@ -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():

Expand Down Expand Up @@ -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()

Expand Down
6 changes: 3 additions & 3 deletions drpcmanager/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 streamadd 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()
Expand Down
Loading